498 lines
16 KiB
Python
Executable File
498 lines
16 KiB
Python
Executable File
#!/usr/bin/env python3
|
|
"""Generate a LoRA trigger word dictionary HTML page from .safetensors metadata.
|
|
|
|
Usage:
|
|
lora-trigger-sheet # default: ~/mnt/nimo.loc/ComfyUI/models/loras
|
|
lora-trigger-sheet /path/to/loras # specify directory
|
|
lora-trigger-sheet . # current directory
|
|
lora-trigger-sheet /path/to/loras /path/to/output # specify both input and output
|
|
"""
|
|
|
|
import json
|
|
import os
|
|
import re
|
|
import struct
|
|
import sys
|
|
from pathlib import Path
|
|
|
|
import hashlib
|
|
import time
|
|
from urllib.error import HTTPError, URLError
|
|
from urllib.request import urlopen, Request
|
|
|
|
|
|
def read_safetensors_meta(filepath):
|
|
try:
|
|
with open(filepath, "rb") as f:
|
|
header_len = struct.unpack("<Q", f.read(8))[0]
|
|
raw = f.read(header_len)
|
|
parsed = json.loads(raw)
|
|
except Exception:
|
|
return None
|
|
|
|
# Metadata is inside "__metadata__" key
|
|
if "__metadata__" in parsed:
|
|
return parsed["__metadata__"]
|
|
return parsed
|
|
|
|
|
|
def get_tool(meta):
|
|
if meta is None:
|
|
return "error"
|
|
|
|
if meta.get("ss_sd_scripts_commit_hash"):
|
|
return "sd-scripts"
|
|
|
|
net = (meta.get("ss_training_network") or meta.get("ss_network_module") or "").lower()
|
|
if "kohya" in net:
|
|
return "kohya"
|
|
if "sd-scripts" in net or "sd_scripts" in net:
|
|
return "sd-scripts"
|
|
|
|
if meta.get("ss_learning_rate"):
|
|
return "sd-scripts"
|
|
|
|
if net:
|
|
return net
|
|
|
|
return "unknown"
|
|
|
|
|
|
def get_base_model(meta):
|
|
if meta is None:
|
|
return "error"
|
|
|
|
base = (meta.get("ss_base_model_version")
|
|
or meta.get("ss_base_model")
|
|
or meta.get("modelspec.architecture", "")
|
|
.replace("/lora", "")
|
|
or "")
|
|
if base:
|
|
return base
|
|
|
|
return "unknown"
|
|
|
|
|
|
def get_tags(meta):
|
|
if meta is None:
|
|
return None, 0
|
|
|
|
raw = meta.get("ss_tag_frequency")
|
|
if not raw:
|
|
return None, 0
|
|
|
|
# ss_tag_frequency is often a JSON string; parse if needed
|
|
if isinstance(raw, str):
|
|
try:
|
|
raw = json.loads(raw)
|
|
except (json.JSONDecodeError, TypeError):
|
|
return None, 0
|
|
|
|
tags = {}
|
|
total = 0
|
|
|
|
if isinstance(raw, dict):
|
|
for category, tag_dict in raw.items():
|
|
if isinstance(tag_dict, dict):
|
|
for tag, count in tag_dict.items():
|
|
total += count
|
|
tags[tag] = tags.get(tag, 0) + count
|
|
elif isinstance(tag_dict, (int, float)):
|
|
total += tag_dict
|
|
|
|
if not tags:
|
|
if all(isinstance(v, (int, float)) for v in raw.values()):
|
|
tags = raw
|
|
total = sum(raw.values())
|
|
else:
|
|
return None, 0
|
|
|
|
return tags, total
|
|
|
|
|
|
def shorten_base_model(name):
|
|
name = re.sub(r"[-_][Vv]\d+(?:[._-]\d+)?$", "", name)
|
|
name = re.sub(r"[-_][Vv]\d+[a-z]?\d*$", "", name)
|
|
replacements = {
|
|
"stable-diffusion-xl-v1-base": "sdxl_base_v1-0",
|
|
"stable diffusion xl base v1.0": "sdxl_base_v1-0",
|
|
"sdxl base v1.0": "sdxl_base_v1-0",
|
|
"sdxl base v0.9": "sdxl_base_v0-9",
|
|
"flux.1": "flux1",
|
|
"flux": "flux1",
|
|
}
|
|
for old, new in replacements.items():
|
|
if name.lower() == old.lower() or name.lower().startswith(old.lower()):
|
|
return new
|
|
return name
|
|
|
|
|
|
def html_id():
|
|
import hashlib, time
|
|
return hashlib.md5(str(time.monotonic_ns()).encode()).hexdigest()[:8]
|
|
|
|
|
|
# --- Source URL lookup cache ---
|
|
|
|
CACHE_DIR = os.path.expanduser("~/.cache")
|
|
CACHE_FILE = os.path.join(CACHE_DIR, "lora-sources.json")
|
|
|
|
|
|
def load_source_cache():
|
|
if os.path.exists(CACHE_FILE):
|
|
try:
|
|
with open(CACHE_FILE) as f:
|
|
return json.load(f)
|
|
except (json.JSONDecodeError, IOError):
|
|
return {}
|
|
return {}
|
|
|
|
|
|
def save_source_cache(cache):
|
|
os.makedirs(CACHE_DIR, exist_ok=True)
|
|
with open(CACHE_FILE, "w") as f:
|
|
json.dump(cache, f, indent=2)
|
|
|
|
|
|
def file_blake2b_hash(filepath):
|
|
h = hashlib.blake2b(digest_size=32)
|
|
with open(filepath, "rb") as f:
|
|
while True:
|
|
chunk = f.read(65536)
|
|
if not chunk:
|
|
break
|
|
h.update(chunk)
|
|
return h.hexdigest()
|
|
|
|
|
|
def _fetch_json(url, timeout=10):
|
|
req = Request(url, headers={"User-Agent": "lora-trigger-sheet/1.0"})
|
|
with urlopen(req, timeout=timeout) as resp:
|
|
return json.loads(resp.read().decode())
|
|
|
|
|
|
def lookup_source(filepath, filename, cache):
|
|
now = int(time.time())
|
|
cached = cache.get(filename)
|
|
if cached:
|
|
url = cached.get("url")
|
|
label = cached.get("label")
|
|
if url and (now - cached.get("last_checked", 0)) < 86400 * 30:
|
|
return url, label
|
|
|
|
# CivitAI model-version lookup by BLAKE2b file hash
|
|
try:
|
|
h = file_blake2b_hash(filepath)
|
|
data = _fetch_json(
|
|
f"https://civitai.com/api/v1/model-versions/by-hash/{h}"
|
|
)
|
|
if isinstance(data, dict):
|
|
model_id = data.get("modelId")
|
|
version_id = data.get("id")
|
|
if model_id:
|
|
url = f"https://civitai.com/models/{model_id}"
|
|
if version_id:
|
|
url += f"?modelVersionId={version_id}"
|
|
cache[filename] = {
|
|
"url": url, "label": "CivitAI",
|
|
"model_name": data.get("name", ""),
|
|
"last_checked": now,
|
|
}
|
|
return url, "CivitAI"
|
|
except (HTTPError, URLError, json.JSONDecodeError, OSError):
|
|
pass
|
|
|
|
# Fallback: HuggingFace model search by filename
|
|
try:
|
|
q = (
|
|
filename.replace(".safetensors", "")
|
|
.replace("_", " ")
|
|
.replace("-", " ")
|
|
)
|
|
data = _fetch_json(
|
|
"https://huggingface.co/api/models"
|
|
f"?search={q}&sort=downloads&direction=-1&limit=3"
|
|
)
|
|
if isinstance(data, list) and len(data) > 0:
|
|
model_id = data[0].get("modelId", "")
|
|
if model_id:
|
|
url = f"https://huggingface.co/{model_id}"
|
|
cache[filename] = {
|
|
"url": url, "label": "HuggingFace",
|
|
"model_name": model_id, "last_checked": now,
|
|
}
|
|
return url, "HuggingFace"
|
|
except (HTTPError, URLError, json.JSONDecodeError, OSError):
|
|
pass
|
|
|
|
cache[filename] = {"url": None, "label": None, "last_checked": now}
|
|
return None, None
|
|
|
|
|
|
def generate_html(loras, sources=None):
|
|
if sources is None:
|
|
sources = {}
|
|
rows = []
|
|
no_trigger_files = []
|
|
|
|
for fname, meta in loras:
|
|
tool = get_tool(meta)
|
|
base = shorten_base_model(get_base_model(meta))
|
|
tags, total = get_tags(meta)
|
|
|
|
fname_no_ext = re.sub(r"\.safetensors$", "", fname)
|
|
|
|
if tags:
|
|
sorted_tags = sorted(tags.items(), key=lambda x: -x[1])
|
|
top8 = sorted_tags[:8]
|
|
rest = sorted_tags[8:]
|
|
has_rest = len(rest) > 0
|
|
|
|
top8_html = ", ".join(
|
|
f"<code>{_e(t)}</code> <span class=\"c\">({c}x)</span>"
|
|
for t, c in top8
|
|
)
|
|
rest_html = ""
|
|
more_id = ""
|
|
if has_rest:
|
|
more_id = "m_" + html_id()
|
|
more_count = sum(c for _, c in rest)
|
|
rest_html = " <span class=\"ml\" onclick=\"t('" + more_id + "',this)\" data-n=\"" + str(more_count) + "\">" + str(more_count) + " more…</span>"
|
|
rest_list = ", ".join(
|
|
f"<code>{_e(t)}</code> <span class=\"c\">({c}x)</span>"
|
|
for t, c in rest
|
|
)
|
|
rest_html += "\n <span class=\"ht\" id=\"" + more_id + "\">" + rest_list + "</span>"
|
|
|
|
triggers = top8_html + rest_html
|
|
else:
|
|
no_trigger_files.append((fname_no_ext, tool, base))
|
|
triggers = "<span class=\"n\">(none)</span>"
|
|
|
|
src_url, src_label = sources.get(fname, (None, None))
|
|
rows.append((fname_no_ext, tool, base, triggers, src_url or "", src_label or ""))
|
|
|
|
num_files = len(rows)
|
|
|
|
table_rows = ""
|
|
for fname, tool, base, triggers, src_url, src_label in rows:
|
|
if src_url:
|
|
src_cell = f'<a href="{_e(src_url)}" target="_blank" rel="noopener">{_e(src_label)}</a>'
|
|
else:
|
|
src_cell = '<span class="n">—</span>'
|
|
table_rows += (
|
|
f" <tr>\n"
|
|
f" <td class=\"fn\">{_e(fname)}</td>\n"
|
|
f" <td class=\"tl\">{_e(tool)}</td>\n"
|
|
f" <td class=\"bm\">{_e(base)}</td>\n"
|
|
f" <td class=\"tr\">{triggers}</td>\n"
|
|
f" <td class=\"sr\">{src_cell}</td>\n"
|
|
f" </tr>\n"
|
|
)
|
|
|
|
html = f"""<!DOCTYPE html>
|
|
<html lang="en">
|
|
<head><meta charset="UTF-8"><meta name="viewport" content="width=device-width,initial-scale=1.0">
|
|
<title>LoRA Info Sheet</title>
|
|
<link rel="icon" type="image/png" href="@img/nude.png">
|
|
<style>
|
|
*,*::before,*::after{{box-sizing:border-box;margin:0;padding:0}}
|
|
body{{font-family:-apple-system,BlinkMacSystemFont,'Segoe UI',Roboto,sans-serif;background:#0d1117;color:#c9d1d9;padding:24px;line-height:1.5;font-size:14px}}
|
|
h1{{font-size:1.5rem;margin-bottom:4px;color:#f0f6fc}}
|
|
.sub{{color:#8b949e;margin-bottom:16px;font-size:0.85rem}}
|
|
.bar{{display:flex;gap:12px;margin-bottom:12px;flex-wrap:wrap;align-items:center}}
|
|
#q{{flex:1;min-width:200px;padding:8px 12px;border:1px solid #30363d;border-radius:6px;background:#161b22;color:#c9d1d9;font-size:0.85rem;outline:none}}
|
|
#q:focus{{border-color:#58a6ff}}
|
|
#cnt{{color:#8b949e;font-size:0.85rem}}
|
|
table{{width:100%;border-collapse:collapse;font-size:0.82rem}}
|
|
th{{text-align:left;padding:9px 10px;border-bottom:2px solid #30363d;color:#f0f6fc;font-weight:600;cursor:pointer;user-select:none;white-space:nowrap;background:#161b22;position:sticky;top:0;z-index:1}}
|
|
th:hover{{color:#58a6ff}}
|
|
th .ar{{margin-left:3px;color:#58a6ff}}
|
|
td{{padding:7px 10px;border-bottom:1px solid #21262d;vertical-align:top;word-break:break-word}}
|
|
tr:hover td{{background:#161b22}}
|
|
.fn{{color:#58a6ff;font-weight:500;white-space:nowrap}}
|
|
.tl,.bm{{color:#8b949e}}
|
|
.tr code{{background:#1f2937;padding:1px 5px;border-radius:3px;font-size:0.78rem;color:#f0c674;white-space:nowrap}}
|
|
.c{{color:#8b949e;font-size:0.72rem;margin-right:3px}}
|
|
.n{{color:#484f58;font-style:italic}}
|
|
.ml{{color:#58a6ff;cursor:pointer;font-size:0.78rem;text-decoration:none;border-bottom:1px dotted #58a6ff}}
|
|
.ml:hover{{border-bottom:1px solid #58a6ff}}
|
|
.ht{{display:none}}
|
|
mark.hl{{background:#264f78;color:#f0f6fc;border-radius:2px;padding:0 2px}}
|
|
.sr a{{color:#58a6ff;text-decoration:none;font-size:0.78rem}}
|
|
.c0{{width:24%}}.c1{{width:8%}}.c2{{width:12%}}.c3{{width:48%}}.c4{{width:8%}}
|
|
@media(max-width:768px){{.c1,.c2,.c4,td:nth-child(2),td:nth-child(3),td:nth-child(5){{display:none}}.c0{{width:30%}}.c3{{width:70%}}}}
|
|
.foot{{margin-top:20px;color:#8b949e}}
|
|
.foot h2{{font-size:1rem;color:#f0f6fc;margin-bottom:6px}}
|
|
.foot ul{{list-style:none;columns:3;font-size:0.82rem}}
|
|
.foot li{{padding:1px 0}}
|
|
</style></head>
|
|
<body>
|
|
<h1>LoRA Info Sheet</h1>
|
|
<p class="sub">{num_files} LoRA files · <span id="vc">{num_files}</span> visible</p>
|
|
<div class="bar"><input type="text" id="q" placeholder="Search filename, trigger, base model..." autofocus> <span id="cnt">{num_files} / {num_files}</span></div>
|
|
<table><thead><tr>
|
|
<th class="c0" onclick="s(0)">File <span class="ar">▲</span></th>
|
|
<th class="c1" onclick="s(1)">Tool <span class="ar"></span></th>
|
|
<th class="c2" onclick="s(2)">Base <span class="ar"></span></th>
|
|
<th class="c3" onclick="s(3)">Trigger(s) <span class="ar"></span></th>
|
|
<th class="c4" onclick="s(4)">Source <span class="ar"></span></th>
|
|
</tr></thead><tbody id="b">
|
|
{table_rows}</tbody></table>
|
|
|
|
<script>
|
|
let col = 0, desc = true;
|
|
const b = document.getElementById('b');
|
|
const q = document.getElementById('q');
|
|
const cnt = document.getElementById('cnt');
|
|
const vc = document.getElementById('vc');
|
|
|
|
function s(i){{
|
|
if (col == i) desc = !desc; else {{ col = i; desc = i == 0; }}
|
|
let rows = [...b.children];
|
|
rows.sort((a,c)=>{{
|
|
let av = a.children[col].textContent.trim().toLowerCase();
|
|
let cv = c.children[col].textContent.trim().toLowerCase();
|
|
if (col == 0 || col == 3) {{
|
|
let am = parseInt(av.match(/\\d+/)?.[0]||'0');
|
|
let cm = parseInt(cv.match(/\\d+/)?.[0]||'0');
|
|
if (!isNaN(am) && !isNaN(cm)) return desc ? cm-am : am-cm;
|
|
}}
|
|
return desc ? av.localeCompare(cv) : cv.localeCompare(av);
|
|
}});
|
|
rows.forEach(r => b.appendChild(r));
|
|
let arrows = document.querySelectorAll('.ar');
|
|
arrows.forEach(a => a.innerHTML = '');
|
|
document.querySelectorAll('th')[col+1].querySelector('.ar').innerHTML = desc ? '▲' : '▼';
|
|
}}
|
|
|
|
function t(id,el){{
|
|
let sp = document.getElementById(id);
|
|
if (sp.classList.contains('ht')){{ sp.classList.remove('ht'); el.style.display='none'; }}
|
|
}}
|
|
|
|
function highlight(el, str) {{
|
|
if (!str) {{
|
|
el.childNodes.forEach(c => {{
|
|
if (c.nodeType === 3) return;
|
|
if (c.classList) c.classList.remove('hl');
|
|
if (c.querySelectorAll) c.querySelectorAll('.hl').forEach(h => h.classList.remove('hl'));
|
|
}});
|
|
return;
|
|
}}
|
|
el.childNodes.forEach(c => {{
|
|
if (c.nodeType === 3 && c.textContent.trim()) {{
|
|
let idx = c.textContent.toLowerCase().indexOf(str);
|
|
if (idx === -1) return;
|
|
let span = document.createElement('span');
|
|
let before = c.textContent.slice(0, idx);
|
|
let match = c.textContent.slice(idx, idx + str.length);
|
|
let after = c.textContent.slice(idx + str.length);
|
|
let mk = document.createElement('mark');
|
|
mk.className = 'hl';
|
|
mk.textContent = match;
|
|
if (before) span.appendChild(document.createTextNode(before));
|
|
span.appendChild(mk);
|
|
if (after) span.appendChild(document.createTextNode(after));
|
|
el.replaceChild(span, c);
|
|
}} else if (c.nodeType === 1) {{
|
|
highlight(c, str);
|
|
}}
|
|
}});
|
|
}}
|
|
|
|
q.addEventListener('input', ()=>{{
|
|
let v = q.value.toLowerCase();
|
|
let rows = b.children;
|
|
let vis = 0;
|
|
document.querySelectorAll('.hl').forEach(h => {{
|
|
let p = h.parentNode;
|
|
let txt = document.createTextNode(h.textContent);
|
|
p.replaceChild(txt, h);
|
|
p.normalize();
|
|
}});
|
|
for (let r of rows) {{
|
|
let ok = r.textContent.toLowerCase().includes(v);
|
|
r.style.display = ok ? '' : 'none';
|
|
if (ok && v) highlight(r, v);
|
|
if (ok) vis++;
|
|
}}
|
|
cnt.textContent = vis + ' / ' + rows.length;
|
|
vc.textContent = vis;
|
|
}});
|
|
</script>
|
|
</body></html>
|
|
"""
|
|
|
|
return html
|
|
|
|
|
|
def _e(s):
|
|
"""Escape HTML entities."""
|
|
s = str(s)
|
|
s = s.replace("&", "&")
|
|
s = s.replace("<", "<")
|
|
s = s.replace(">", ">")
|
|
s = s.replace('"', """)
|
|
return s
|
|
|
|
|
|
def main():
|
|
args = sys.argv[1:]
|
|
|
|
if len(args) >= 1:
|
|
target = args[0]
|
|
else:
|
|
target = os.path.expanduser("~/mnt/nimo.loc/ComfyUI/models/loras")
|
|
|
|
target = os.path.abspath(os.path.expanduser(target))
|
|
|
|
if not os.path.isdir(target):
|
|
print(f"Error: {target} is not a directory", file=sys.stderr)
|
|
sys.exit(1)
|
|
|
|
safetensors = sorted([
|
|
f for f in os.listdir(target)
|
|
if f.endswith(".safetensors") and not f.startswith("._")
|
|
])
|
|
|
|
if not safetensors:
|
|
print(f"No .safetensors files found in {target}", file=sys.stderr)
|
|
sys.exit(1)
|
|
|
|
loras = []
|
|
for fname in safetensors:
|
|
fpath = os.path.join(target, fname)
|
|
meta = read_safetensors_meta(fpath)
|
|
loras.append((fname, meta))
|
|
|
|
# Look up source URLs with disk-backed cache
|
|
cache = load_source_cache()
|
|
sources = {}
|
|
total = len(loras)
|
|
for i, (fname, _) in enumerate(loras, 1):
|
|
fpath = os.path.join(target, fname)
|
|
url, label = lookup_source(fpath, fname, cache)
|
|
if url:
|
|
sources[fname] = (url, label)
|
|
if i % 10 == 0 or i == total:
|
|
print(f" source lookup: {i}/{total}", file=sys.stderr)
|
|
save_source_cache(cache)
|
|
|
|
html = generate_html(loras, sources)
|
|
|
|
if len(args) >= 2:
|
|
out_path = os.path.abspath(os.path.expanduser(args[1]))
|
|
else:
|
|
out_path = os.path.join(os.path.dirname(os.path.abspath(__file__)), "lora-info-sheet.html")
|
|
|
|
with open(out_path, "w") as f:
|
|
f.write(html)
|
|
|
|
out_count = sum(1 for v in sources.values() if v[0])
|
|
print(f"Generated {out_path} with {len(loras)} LoRA entries ({out_count} with source links)")
|
|
|
|
|
|
if __name__ == "__main__":
|
|
main()
|