mirror of
https://github.com/vladmandic/automatic
synced 2026-09-19 09:14:35 +02:00
improve civitai integration
This commit is contained in:
+7
-4
@@ -4,15 +4,18 @@
|
||||
|
||||
Mostly a service release
|
||||
- tons of fixes
|
||||
- update ui hints
|
||||
- update **ui hints**
|
||||
- updated **models -> civitai**
|
||||
- search and download loras
|
||||
- find previews for already downloaded models or loras
|
||||
- new option **inference mode**
|
||||
- default is standard `torch.no_grad`
|
||||
new option is `torch.inference_only` which is slightly faster and uses less vram, but only works on some gpus
|
||||
- new cmdline param `--no-metadata`
|
||||
skips reading metadata from models that are not already cached
|
||||
- updated gradio
|
||||
- styles support for subfolders
|
||||
- clean-up logging
|
||||
- updated **gradio**
|
||||
- **styles** support for subfolders
|
||||
- clean-up **logging**
|
||||
- capture system info in startup log
|
||||
- better diagnostic output
|
||||
- capture extension output
|
||||
|
||||
@@ -95,7 +95,7 @@ class LoraOnDisk:
|
||||
|
||||
def set_hash(self, v):
|
||||
self.hash = v
|
||||
self.shorthash = self.hash[0:12]
|
||||
self.shorthash = self.hash[0:10]
|
||||
if self.shorthash:
|
||||
available_lora_hash_lookup[self.shorthash] = self
|
||||
|
||||
|
||||
+9
-4
@@ -21,12 +21,17 @@ def cache(subsection):
|
||||
return s
|
||||
|
||||
|
||||
def calculate_sha256(filename):
|
||||
def calculate_sha256(filename, quiet=False):
|
||||
hash_sha256 = hashlib.sha256()
|
||||
blksize = 1024 * 1024
|
||||
with progress.open(filename, 'rb', description=f'Calculating model hash: [cyan]{filename}', auto_refresh=True) as f:
|
||||
for chunk in iter(lambda: f.read(blksize), b""):
|
||||
hash_sha256.update(chunk)
|
||||
if not quiet:
|
||||
with progress.open(filename, 'rb', description=f'Calculating model hash: [cyan]{filename}', auto_refresh=True) as f:
|
||||
for chunk in iter(lambda: f.read(blksize), b""):
|
||||
hash_sha256.update(chunk)
|
||||
else:
|
||||
with open(filename, 'rb') as f:
|
||||
for chunk in iter(lambda: f.read(blksize), b""):
|
||||
hash_sha256.update(chunk)
|
||||
return hash_sha256.hexdigest()
|
||||
|
||||
|
||||
|
||||
+15
-13
@@ -4,6 +4,7 @@ import shutil
|
||||
import importlib
|
||||
from typing import Dict
|
||||
from urllib.parse import urlparse
|
||||
import PIL.Image as Image
|
||||
from modules import shared
|
||||
from modules.upscaler import Upscaler, UpscalerLanczos, UpscalerNearest, UpscalerNone
|
||||
from modules.paths import script_path, models_path
|
||||
@@ -59,12 +60,13 @@ def download_civit_preview(model_path: str, preview_url: str):
|
||||
import rich.progress as p
|
||||
_, ext = os.path.splitext(preview_url)
|
||||
model_name, _ = os.path.splitext(os.path.basename(model_path))
|
||||
preview_file = os.path.splitext(model_path)[0] + ext
|
||||
preview_file = f'{os.path.splitext(model_path)[0]}{ext}' if '.safetensors' in model_path.lower() else f'{model_path}{ext}'
|
||||
res = f'CivitAI download: name={model_name} url={preview_url}'
|
||||
req = requests.get(preview_url, stream=True, timeout=30)
|
||||
total_size = int(req.headers.get('content-length', 0))
|
||||
block_size = 16384 # 16KB blocks
|
||||
written = 0
|
||||
img = None
|
||||
shared.state.begin('civitai-download-preview')
|
||||
try:
|
||||
with open(preview_file, 'wb') as f:
|
||||
@@ -77,13 +79,14 @@ def download_civit_preview(model_path: str, preview_url: str):
|
||||
if written < 1024: # min threshold
|
||||
os.remove(preview_file)
|
||||
raise ValueError(f'removed invalid download: bytes={written}')
|
||||
img = Image.open(preview_file)
|
||||
except Exception as e:
|
||||
shared.log.error(f'CivitAI download error: name={model_name} url={preview_url} {e}')
|
||||
if total_size == written:
|
||||
shared.log.info(f'{res} size={total_size}')
|
||||
else:
|
||||
shared.log.error(f'{res} size={total_size} written={written}')
|
||||
shared.state.end()
|
||||
if img is None:
|
||||
return res
|
||||
shared.log.info(f'{res} size={total_size} image={img.size}')
|
||||
img.close()
|
||||
return res
|
||||
|
||||
|
||||
@@ -114,7 +117,7 @@ def download_civit_model(model_url: str, model_name: str, model_path: str, model
|
||||
written = written + len(data)
|
||||
f.write(data)
|
||||
progress.update(task, advance=block_size, description="Downloading")
|
||||
if written < 1024 * 1024 * 1024: # min threshold
|
||||
if written < 1024 * 1024: # min threshold
|
||||
os.remove(model_file)
|
||||
raise ValueError(f'removed invalid download: bytes={written}')
|
||||
if preview is not None:
|
||||
@@ -160,17 +163,16 @@ def download_diffusers_model(hub_id: str, cache_dir: str = None, download_config
|
||||
except Exception as e:
|
||||
shared.log.error(f"Diffusers download error: {hub_id} {e}")
|
||||
try:
|
||||
model_info_dict = hf.model_info(hub_id).cardData # pylint: disable=no-member # TODO Diffusers is this real error?
|
||||
model_info_dict = hf.model_info(hub_id).cardData if pipeline_dir is not None else None # pylint: disable=no-member # TODO Diffusers is this real error?
|
||||
except Exception:
|
||||
model_info_dict = None
|
||||
# some checkpoints need to be downloaded as "hidden" as they just serve as pre- or post-pipelines of other pipelines
|
||||
if model_info_dict is not None and "prior" in model_info_dict:
|
||||
if model_info_dict is not None and "prior" in model_info_dict: # some checkpoints need to be downloaded as "hidden" as they just serve as pre- or post-pipelines of other pipelines
|
||||
download_dir = DiffusionPipeline.download(model_info_dict["prior"][0], **download_config)
|
||||
model_info_dict["prior"] = download_dir
|
||||
# mark prior as hidden
|
||||
with open(os.path.join(download_dir, "hidden"), "w", encoding="utf-8") as f:
|
||||
with open(os.path.join(download_dir, "hidden"), "w", encoding="utf-8") as f: # mark prior as hidden
|
||||
f.write("True")
|
||||
shared.writefile(model_info_dict, os.path.join(pipeline_dir, "model_info.json"))
|
||||
if pipeline_dir is not None:
|
||||
shared.writefile(model_info_dict, os.path.join(pipeline_dir, "model_info.json"))
|
||||
shared.state.end()
|
||||
return pipeline_dir
|
||||
|
||||
@@ -323,7 +325,7 @@ def extension_filter(ext_filter=None, ext_blacklist=None):
|
||||
return (not ext_filter or any(fp.upper().endswith(ew) for ew in ext_filter)) and (not ext_blacklist or not any(fp.upper().endswith(ew) for ew in ext_blacklist))
|
||||
return filter
|
||||
|
||||
def load_file_from_url(url: str, *, model_dir: str, progress: bool = True, file_name: str | None = None) -> str:
|
||||
def load_file_from_url(url: str, *, model_dir: str, progress: bool = True, file_name = None):
|
||||
"""Download a file from url into model_dir, using the file present if possible. Returns the path to the downloaded file."""
|
||||
os.makedirs(model_dir, exist_ok=True)
|
||||
if not file_name:
|
||||
|
||||
+25
-16
@@ -8,6 +8,7 @@ from modules.ui_common import create_refresh_button
|
||||
from modules.call_queue import wrap_gradio_gpu_call
|
||||
from modules.shared import opts, log
|
||||
import modules.errors
|
||||
import modules.hashes
|
||||
|
||||
|
||||
def create_ui():
|
||||
@@ -201,15 +202,11 @@ def create_ui():
|
||||
|
||||
def hf_download_model(hub_id: str, token, variant, revision, mirror):
|
||||
from modules.modelloader import download_diffusers_model
|
||||
try:
|
||||
download_diffusers_model(hub_id, cache_dir=opts.diffusers_dir, token=token, variant=variant, revision=revision, mirror=mirror)
|
||||
except Exception as e:
|
||||
log.error(f"Diffuser model downloaded error: model={hub_id} {e}")
|
||||
return f"Diffuser model downloaded error: model={hub_id} {e}"
|
||||
download_diffusers_model(hub_id, cache_dir=opts.diffusers_dir, token=token, variant=variant, revision=revision, mirror=mirror)
|
||||
from modules.sd_models import list_models # pylint: disable=W0621
|
||||
list_models()
|
||||
log.info(f"Diffuser model downloaded: model={hub_id}")
|
||||
return f'Diffuser model downloaded: model={hub_id}'
|
||||
log.info(f'Diffuser model downloaded: model="{hub_id}"')
|
||||
return f'Diffuser model downloaded: model="{hub_id}"'
|
||||
|
||||
with gr.Column(scale=6):
|
||||
with gr.Row():
|
||||
@@ -252,7 +249,7 @@ def create_ui():
|
||||
if tag is not None and len(tag) > 0:
|
||||
url += f'&tag={tag}'
|
||||
r = requests.get(url, timeout=60, headers=headers)
|
||||
log.debug(f'CivitAI search: name={name} tag={tag} status={r.status_code}')
|
||||
log.debug(f'CivitAI search: name="{name}" tag={tag or "none"} status={r.status_code}')
|
||||
if r.status_code != 200:
|
||||
return [], [], []
|
||||
body = r.json()
|
||||
@@ -261,6 +258,8 @@ def create_ui():
|
||||
data1 = []
|
||||
for model in data:
|
||||
found = 0
|
||||
if model_type == 'LoRA' and model['type'] == 'LORA':
|
||||
found += 1
|
||||
for variant in model['modelVersions']:
|
||||
if model_type == 'SD 1.5':
|
||||
if 'SD 1.' in variant['baseModel']:
|
||||
@@ -297,7 +296,7 @@ def create_ui():
|
||||
d['baseModel'],
|
||||
d['createdAt'],
|
||||
])
|
||||
log.debug(f'CivitAI select: model={in_data[evt.index[0]]} versions={len(data2)}')
|
||||
log.debug(f'CivitAI select: model="{in_data[evt.index[0]]}" versions={len(data2)}')
|
||||
return data2, preview_img
|
||||
|
||||
def civit_select2(evt: gr.SelectData, in_data):
|
||||
@@ -315,7 +314,7 @@ def create_ui():
|
||||
json.dumps(f['metadata']),
|
||||
f['downloadUrl'],
|
||||
])
|
||||
log.debug(f'CivitAI select: model={in_data[evt.index[0]]} files={len(data3)}')
|
||||
log.debug(f'CivitAI select: model="{in_data[evt.index[0]]}" files={len(data3)}')
|
||||
return data3
|
||||
|
||||
def civit_select3(evt: gr.SelectData, in_data):
|
||||
@@ -336,7 +335,7 @@ def create_ui():
|
||||
list_models()
|
||||
return res
|
||||
|
||||
def civit_download_previews():
|
||||
def civit_download_previews(civit_previews_rehash):
|
||||
import requests
|
||||
from modules.ui_extra_networks import extra_pages
|
||||
from modules.modelloader import download_civit_preview
|
||||
@@ -347,17 +346,26 @@ def create_ui():
|
||||
if item.get('fullname', None) is None:
|
||||
continue
|
||||
if 'card-no-preview.png' in item['preview'] and os.path.isfile(item['fullname']):
|
||||
sha = item.get('hash', None)
|
||||
if item.get('hash', None) is None:
|
||||
log.debug(f'CivitAI skipping item without hash: name={item["name"]}')
|
||||
log.debug(f'CivitAI skipping item without hash: name="{item["name"]}"')
|
||||
continue
|
||||
url = f'https://civitai.com/api/v1/model-versions/by-hash/{item["hash"]}'
|
||||
r = requests.get(url, timeout=5, headers=headers)
|
||||
log.debug(f'CivitAI search: name={item["name"]} hash={item["hash"]} status={r.status_code}')
|
||||
r = requests.get(f'https://civitai.com/api/v1/model-versions/by-hash/{sha}', timeout=5, headers=headers)
|
||||
log.debug(f'CivitAI search: name="{item["name"]}" hash={sha} status={r.status_code}')
|
||||
if r.status_code == 200:
|
||||
d = r.json()
|
||||
if d.get('images') is not None and len(d['images']) > 0 and len(d['images'][0]['url']) > 0:
|
||||
preview_url = d['images'][0]['url']
|
||||
res += download_civit_preview(item['filename'], preview_url) + '<br>'
|
||||
elif civit_previews_rehash and os.stat(item['fullname']).st_size < (1024 * 1024 * 1024):
|
||||
sha = modules.hashes.calculate_sha256(item['fullname'], quiet=True)
|
||||
r = requests.get(f'https://civitai.com/api/v1/model-versions/by-hash/{sha}', timeout=5, headers=headers)
|
||||
log.debug(f'CivitAI search: name="{item["name"]}" hash={sha} status={r.status_code}')
|
||||
if r.status_code == 200:
|
||||
d = r.json()
|
||||
if d.get('images') is not None and len(d['images']) > 0 and len(d['images'][0]['url']) > 0:
|
||||
preview_url = d['images'][0]['url']
|
||||
res += download_civit_preview(item['filename'], preview_url) + '<br>'
|
||||
return res
|
||||
|
||||
with gr.Row():
|
||||
@@ -389,6 +397,7 @@ def create_ui():
|
||||
civit_results1 = gr.DataFrame(value = None, label = 'Search results', show_label = True, interactive = False, wrap = True, overflow_row_behaviour = 'paginate', max_rows = 10, headers = civit_headers1, datatype = civit_types1, type='array')
|
||||
with gr.Row():
|
||||
civit_previews_btn = gr.Button(value="Fetch previews for existing models", variant='primary')
|
||||
civit_previews_rehash = gr.Checkbox(value=False, label="Check alternative hash")
|
||||
|
||||
civit_search_text.submit(fn=civit_search, inputs=[civit_search_text, civit_search_tag, civit_model_type], outputs=[civit_results1, civit_results2, civit_results3])
|
||||
civit_search_tag.submit(fn=civit_search, inputs=[civit_search_text, civit_search_tag, civit_model_type], outputs=[civit_results1, civit_results2, civit_results3])
|
||||
@@ -397,4 +406,4 @@ def create_ui():
|
||||
civit_results2.select(fn=civit_select2, inputs=[civit_results2], outputs=[civit_results3])
|
||||
civit_results3.select(fn=civit_select3, inputs=[civit_results3], outputs=[civit_selected, civit_name, civit_search_btn])
|
||||
civit_download_model_btn.click(fn=civit_download_model, inputs=[civit_selected, civit_name, civit_path, civit_model_type, models_image], outputs=[models_outcome])
|
||||
civit_previews_btn.click(fn=civit_download_previews, inputs=[], outputs=[models_outcome])
|
||||
civit_previews_btn.click(fn=civit_download_previews, inputs=[civit_previews_rehash], outputs=[models_outcome])
|
||||
|
||||
Reference in New Issue
Block a user