mirror of
https://github.com/vladmandic/automatic
synced 2026-09-04 03:50:44 +02:00
fix modelloader from huggingface
This commit is contained in:
Executable
+77
@@ -0,0 +1,77 @@
|
||||
import os
|
||||
import logging
|
||||
import git
|
||||
from rich import console, progress
|
||||
|
||||
|
||||
class GitRemoteProgress(git.RemoteProgress):
|
||||
OP_CODES = ["BEGIN", "CHECKING_OUT", "COMPRESSING", "COUNTING", "END", "FINDING_SOURCES", "RECEIVING", "RESOLVING", "WRITING"]
|
||||
OP_CODE_MAP = { getattr(git.RemoteProgress, _op_code): _op_code for _op_code in OP_CODES }
|
||||
|
||||
def __init__(self, url, folder) -> None:
|
||||
super().__init__()
|
||||
self.url = url
|
||||
self.folder = folder
|
||||
self.progressbar = progress.Progress(
|
||||
progress.SpinnerColumn(),
|
||||
progress.TextColumn("[cyan][progress.description]{task.description}"),
|
||||
progress.BarColumn(),
|
||||
progress.TextColumn("[progress.percentage]{task.percentage:>3.0f}%"),
|
||||
progress.TimeRemainingColumn(),
|
||||
progress.TextColumn("[yellow]<{task.fields[url]}>"),
|
||||
progress.TextColumn("{task.fields[message]}"),
|
||||
console=console.Console(),
|
||||
transient=False,
|
||||
)
|
||||
self.progressbar.start()
|
||||
self.active_task = None
|
||||
|
||||
def __del__(self) -> None:
|
||||
self.progressbar.stop()
|
||||
|
||||
@classmethod
|
||||
def get_curr_op(cls, op_code: int) -> str:
|
||||
op_code_masked = op_code & cls.OP_MASK
|
||||
return cls.OP_CODE_MAP.get(op_code_masked, "?").title()
|
||||
|
||||
def update(self, op_code: int, cur_count: str | float, max_count: str | float | None = None, message: str | None = "") -> None:
|
||||
if op_code & self.BEGIN:
|
||||
self.curr_op = self.get_curr_op(op_code) # pylint: disable=attribute-defined-outside-init
|
||||
self.active_task = self.progressbar.add_task(description=self.curr_op, total=max_count, message=message, url=self.url)
|
||||
self.progressbar.update(task_id=self.active_task, completed=cur_count, message=message)
|
||||
if op_code & self.END:
|
||||
self.progressbar.update(task_id=self.active_task, message=f"[bright_black]{message}")
|
||||
|
||||
|
||||
def clone(url: str, folder: str):
|
||||
git.Repo.clone_from(
|
||||
url=url,
|
||||
to_path=folder,
|
||||
progress=GitRemoteProgress(url=url, folder=folder),
|
||||
multi_options=['--config core.compression=0', '--config core.loosecompression=0', '--config pack.window=0'],
|
||||
allow_unsafe_options=True,
|
||||
depth=1,
|
||||
)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
import argparse
|
||||
parser = argparse.ArgumentParser(description = 'downloader')
|
||||
parser.add_argument('--url', required=True, help="download url, required")
|
||||
parser.add_argument('--folder', required=False, help="output folder, default: autodetect")
|
||||
args = parser.parse_args()
|
||||
logging.basicConfig(level=logging.INFO, format="%(asctime)s %(levelname)s: %(message)s")
|
||||
log = logging.getLogger(__name__)
|
||||
try:
|
||||
if not args.url.startswith('http'):
|
||||
raise ValueError(f'invalid url: {args.url}')
|
||||
f = args.url.split('/')[-1].split('.')[0] if args.folder is None else args.folder
|
||||
if os.path.exists(f):
|
||||
raise FileExistsError(f'folder already exists: {f}')
|
||||
log.info(f'Clone start: url={args.url} folder={f}')
|
||||
clone(url=args.url, folder=f)
|
||||
log.info(f'Clone complete: url={args.url} folder={f}')
|
||||
except KeyboardInterrupt:
|
||||
log.warning(f'Clone cancelled: url={args.url} folder={f}')
|
||||
except Exception as e:
|
||||
log.error(f'Clone: url={args.url} {e}')
|
||||
+1
-1
@@ -3,7 +3,7 @@
|
||||
"path": "dreamshaper_8.safetensors@https://civitai.com/api/download/models/128713",
|
||||
"desc": "Showcase finetuned model based on Stable diffusion 1.5",
|
||||
"preview": "dreamshaper_8.jpg",
|
||||
"extras": "width: 512, height: 512, sampler: DEIS, steps: 20, cfg_scale: 6.0",
|
||||
"extras": "width: 768, height: 512, sampler: DEIS, steps: 20, cfg_scale: 6.0",
|
||||
"original": true
|
||||
},
|
||||
"DreamShaper SD XL Turbo": {
|
||||
|
||||
@@ -187,7 +187,7 @@ def send_image_and_dimensions(x):
|
||||
def parse_generation_parameters(infotext):
|
||||
if not isinstance(infotext, str):
|
||||
return {}
|
||||
|
||||
debug(f'Parse infotext: {infotext}')
|
||||
re_param = re.compile(r'\s*([\w ]+):\s*("(?:\\"[^,]|\\"|\\|[^\"])+"|[^,]*)(?:,|$)') # multi-word: value
|
||||
re_size = re.compile(r"^(\d+)x(\d+)$") # int x int
|
||||
sanitized = infotext.replace('prompt:', 'Prompt:').replace('negative prompt:', 'Negative prompt:').replace('Negative Prompt', 'Negative prompt') # cleanup everything in brackets so re_params can work
|
||||
@@ -196,6 +196,7 @@ def parse_generation_parameters(infotext):
|
||||
sanitized = re.sub(r'\{[^}]*\}', lambda match: ' ' * len(match.group()), sanitized)
|
||||
|
||||
params = dict(re_param.findall(sanitized))
|
||||
debug(f"Parse params: {params}")
|
||||
params = { k.strip():params[k].strip() for k in params if k.lower() not in ['hashes', 'lora', 'embeddings', 'prompt', 'negative prompt']} # remove some keys
|
||||
first_param = next(iter(params)) if params else None
|
||||
params_idx = sanitized.find(f'{first_param}:') if first_param else -1
|
||||
@@ -223,7 +224,7 @@ def parse_generation_parameters(infotext):
|
||||
params[k] = v
|
||||
params["Prompt"] = prompt.replace('Prompt:', '').strip()
|
||||
params["Negative prompt"] = negative.replace('Negative prompt:', '').strip()
|
||||
debug(f"Paste params: {params}")
|
||||
debug(f"Parse: {params}")
|
||||
return params
|
||||
|
||||
|
||||
|
||||
@@ -301,6 +301,8 @@ def load_reference(name: str):
|
||||
if len(found) > 0: # already downloaded
|
||||
model_opts = get_reference_opts(found[0]['name'])
|
||||
return True
|
||||
else:
|
||||
model_opts = get_reference_opts(name)
|
||||
if model_opts.get('skip', False):
|
||||
return True
|
||||
shared.log.debug(f'Reference: download="{name}"')
|
||||
|
||||
+4
-5
@@ -60,17 +60,15 @@ def apply_wildcards_to_prompt(prompt, all_wildcards):
|
||||
|
||||
|
||||
def apply_styles_to_extra(p, style: Style):
|
||||
global reference_style # pylint: disable=global-statement
|
||||
if style is None:
|
||||
return
|
||||
name_map = {
|
||||
'sampler': 'sampler_name',
|
||||
}
|
||||
from modules.generation_parameters_copypaste import parse_generation_parameters
|
||||
s = style.extra
|
||||
s = 'Negative prompt: ' + s if 'Negative prompt:' not in s else s
|
||||
s = 'Prompt: ' + s if 'Prompt:' not in s else s
|
||||
extra = parse_generation_parameters(reference_style) if shared.opts.extra_network_reference else {}
|
||||
extra.update(parse_generation_parameters(s))
|
||||
extra.update(parse_generation_parameters(style.extra))
|
||||
extra.pop('Prompt', None)
|
||||
extra.pop('Negative prompt', None)
|
||||
fields = []
|
||||
@@ -88,7 +86,8 @@ def apply_styles_to_extra(p, style: Style):
|
||||
fields.append(f'{k}={v}')
|
||||
else:
|
||||
skipped.append(f'{k}={v}')
|
||||
shared.log.debug(f'Applying style: name="{style.name}" extra={fields} skipped={skipped}')
|
||||
shared.log.debug(f'Applying style: name="{style.name}" extra={fields} skipped={skipped} reference={True if reference_style else False}')
|
||||
# reference_style = None
|
||||
|
||||
|
||||
class StyleDatabase:
|
||||
|
||||
Reference in New Issue
Block a user