Files
automatic/modules/ui_extra_networks_checkpoints.py
T
CalamitousFelicitousness 2d5903e91b fix(networks): keep reference mtime a datetime for date sort
strptime was handed the stat datetime instead of the date string, so it
always failed and left a raw string mtime; the Date sort then crashed on
mixed str/datetime items. Parse the date string, falling back to stat mtime.
2026-07-03 21:14:50 +01:00

200 lines
7.9 KiB
Python

import os
import html
import json
import concurrent.futures
from datetime import datetime
from modules import shared, ui_extra_networks, sd_models, modelstats, paths, devices
from modules.logger import log
from modules.json_helpers import readfile
version_map = {
"QwenEdit": "Qwen",
"QwenEditPlus": "Qwen",
"Flux.1 D": "Flux",
"Flux.1 S": "Flux",
"FluxKontext": "Flux",
"SDXL 1.0": "SD XL",
"SDXL Hyper": "SD XL",
"StableDiffusion": "SD 1.5",
"StableDiffusion3": "SD 3",
"StableDiffusionXL": "SD XL",
"WanToVideo": "Wan",
"WanVACE": "Wan",
"Z": "Z-Image",
"Glm": "GLM-Image",
"Krea2": "Krea 2",
"AnimaTextTo": "Anima",
"Ideogram4": "Ideogram 4",
"Flux2": "Flux 2",
"Flux2Klein": "Flux 2 Klein",
"Flux2KleinKV": "Flux 2 Klein",
}
class ExtraNetworksPageCheckpoints(ui_extra_networks.ExtraNetworksPage):
def __init__(self):
super().__init__('Model')
def refresh(self):
shared.refresh_checkpoints()
def list_reference(self): # pylint: disable=inconsistent-return-statements
existing = [model.filename if model.type == 'safetensors' else model.name for model in sd_models.checkpoints_list.values()]
def reference_downloaded(url):
url = url.split('@')[0] if '@' in url else 'Diffusers/' + url
url = url.split('+')[0] if '+' in url else url
return any(model.endswith(url) for model in existing)
if not shared.opts.sd_checkpoint_autodownload or not shared.opts.extra_network_reference_enable:
log.debug(f'Networks: type="reference" autodownload={shared.opts.sd_checkpoint_autodownload} enable={shared.opts.extra_network_reference_enable}')
return []
count = { 'total': 0, 'ready': 0, 'hidden': 0, 'experimental': 0, 'base': 0, 'quantized': 0, 'distilled': 0, 'community': 0, 'cloud': 0, 'nunchaku': 0 }
shared.reference_models = {}
for tag in count.keys():
fn = os.path.join('data', f'reference-{tag}.json')
dct = readfile(fn, as_type="dict", silent=True)
for k, v in dct.items():
v['skip'] = 'safetensors' not in v.get('path', '')
v['tags'] = [tag.capitalize()]
size = v.get('size', 0)
if size > 0:
v['tags'].append(f'Size: {size} GB')
shared.reference_models[k] = v
for k, v in shared.reference_models.items():
count['total'] += 1
url = v['path']
if v.get('hidden', False):
count['hidden'] += 1
continue
experimental = v.get('experimental', False)
if experimental:
if shared.cmd_opts.experimental:
log.debug(f'Networks: experimental model="{k}"')
count['experimental'] += 1
else:
continue
preview = v.get('preview', v['path'])
preview_file = self.find_preview_file(os.path.join(paths.reference_path, preview))
name = os.path.normpath(os.path.join(paths.reference_path, k)).replace('\\', '/')
size = int(float(v.get('size', 0)) * 1024 * 1024 * 1024)
mtime = v.get('date', None)
_size, _mtime = modelstats.stat(preview_file)
if mtime is None:
mtime = _mtime
else:
try:
mtime = datetime.strptime(mtime, '%Y %B') # 2025 January
except Exception:
mtime = _mtime
if size == 0:
size = _size
if len(v.get("subfolder", "")) > 0:
path = f'{v.get("path", "")}+{v.get("subfolder", "")}'
else:
path = f'{v.get("path", "")}'
tag = v.get('tags', [])
if isinstance(tag, list):
tag = ', '.join(tag)
if isinstance(tag, list) and len(tag) > 0:
primary = tag[0].strip()
elif isinstance(tag, str):
primary = tag.split(',')[0].strip() if len(tag) > 0 else ''
else:
primary = ''
if ('nunchaku' in tag) and (devices.backend != 'cuda' and not shared.cmd_opts.experimental):
count['hidden'] += 1
continue
if primary in count:
count[primary] += 1
elif primary != '':
count[primary] = 1
else:
count['base'] += 1
ready = reference_downloaded(url)
version = "ready" if ready else "download"
if 'cloud' in tag :
version = 'Cloud'
if not ready and shared.opts.offline_mode:
count['hidden'] += 1
continue
if ready:
count['ready'] += 1
yield {
"type": 'Model',
"name": name,
"title": name,
"filename": url,
"preview": self.find_preview(os.path.join(paths.reference_path, preview)),
"local_preview": preview_file,
"onclick": '"' + html.escape(f"selectReference({json.dumps(path)})") + '"',
"hash": None,
"mtime": mtime,
"size": size,
"info": {},
"metadata": {},
"description": v.get('desc', ''),
"version": version,
"tags": v.get('tags', []),
}
log.debug(f'Networks: type="reference" {count}')
def create_item(self, name):
record = None
try:
checkpoint: sd_models.CheckpointInfo = sd_models.checkpoints_list.get(name)
size, mtime = modelstats.stat(checkpoint.filename)
record = {
"type": 'Model',
"name": checkpoint.name,
"title": checkpoint.title,
"filename": checkpoint.filename,
"hash": checkpoint.shorthash,
"metadata": checkpoint.metadata,
"onclick": '"' + html.escape(f"selectCheckpoint({json.dumps(name)})") + '"',
"mtime": mtime,
"size": size,
}
record['info'] = self.find_info(checkpoint.filename)
record['description'] = self.find_description(checkpoint.filename, record['info'])
version = self.find_version(checkpoint, record['info'])
if 'baseModel' in version:
record['version'] = version.get("baseModel", "")
elif '_class_name' in record['info']:
cls = record['info']['_class_name']
if isinstance(cls, list):
cls = cls[-1]
record['version'] = cls.replace('Pipeline', '').replace('Image', '')
else:
record['version'] = ''
record['version'] = version_map.get(record['version'], record['version'])
except Exception as e:
log.error(f'Networks error: type=model file="{name}" {e}')
if os.environ.get('SD_EN_DEBUG', None) is not None:
from modules import errors
errors.display(e, 'Networks')
return record
def list_items(self):
items = []
with concurrent.futures.ThreadPoolExecutor(max_workers=shared.max_workers) as executor:
future_items = {executor.submit(self.create_item, cp): cp for cp in list(sd_models.checkpoints_list.copy())}
for future in concurrent.futures.as_completed(future_items):
item = future.result()
if item is not None:
items.append(item)
for record in self.list_reference():
items.append(record)
self.update_all_previews(items)
return items
def allowed_directories_for_previews(self):
return [v for v in [shared.opts.ckpt_dir, paths.reference_path, sd_models.model_path] if v is not None]