improve civitai integration

This commit is contained in:
Vladimir Mandic
2023-09-10 18:13:20 -04:00
parent 5a649f951a
commit 56e041c3b6
5 changed files with 57 additions and 38 deletions
+7 -4
View File
@@ -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
+1 -1
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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])