From 1f6e6cf8e1e0103c6ba3cabf5478cd7e4322947a Mon Sep 17 00:00:00 2001 From: pythongosssss <125205205+pythongosssss@users.noreply.github.com> Date: Sun, 20 Aug 2023 18:34:41 +0100 Subject: [PATCH] Add lora info --- js/loraInfo.css | 81 ++++++++++++++++ js/loraInfo.js | 243 ++++++++++++++++++++++++++++++++++++++++++++++++ py/lora_info.py | 37 ++++++++ 3 files changed, 361 insertions(+) create mode 100644 js/loraInfo.css create mode 100644 js/loraInfo.js create mode 100644 py/lora_info.py diff --git a/js/loraInfo.css b/js/loraInfo.css new file mode 100644 index 0000000..dc88dba --- /dev/null +++ b/js/loraInfo.css @@ -0,0 +1,81 @@ +.pysssss-lora-info { + color: white; + font-family: sans-serif; + max-width: 90vw; +} +.pysssss-lora-content { + display: flex; + flex-direction: column; + overflow: hidden; +} +.pysssss-lora-info h2 { + text-align: center; + margin: 0 0 10px 0; +} +.pysssss-lora-info p { + margin: 5px 0; +} +.pysssss-lora-info a { + color: dodgerblue; +} +.pysssss-lora-info a:hover { + text-decoration: underline; +} +.pysssss-lora-tags-list { + display: flex; + flex-wrap: wrap; + list-style: none; + gap: 10px; + max-height: 200px; + overflow: auto; + margin: 10px 0; + padding: 0; +} +.pysssss-lora-tag { + background-color: rgb(128, 213, 247); + color: #000; + display: flex; + align-items: center; + gap: 5px; + border-radius: 5px; + padding: 2px 5px; + cursor: pointer; +} +.pysssss-lora-tag--selected span::before { + content: "✅"; + position: absolute; + background-color: dodgerblue; + left: 0; + top: 0; + right: 0; + bottom: 0; + text-align: center; +} +.pysssss-lora-tag:hover { + outline: 2px solid dodgerblue; +} +.pysssss-lora-tag p { + margin: 0; +} +.pysssss-lora-tag span { + text-align: center; + border-radius: 5px; + background-color: dodgerblue; + color: #fff; + padding: 2px; + position: relative; + min-width: 20px; + overflow: hidden; +} + +.pysssss-lora-metadata .comfy-modal-content { + max-width: 100%; +} +.pysssss-lora-metadata label { + margin-right: 1ch; + color: #ccc; +} + +.pysssss-lora-metadata span { + color: dodgerblue; +} diff --git a/js/loraInfo.js b/js/loraInfo.js new file mode 100644 index 0000000..7997bac --- /dev/null +++ b/js/loraInfo.js @@ -0,0 +1,243 @@ +import { app } from "../../../scripts/app.js"; +import { $el, ComfyDialog } from "../../../scripts/ui.js"; +import { api } from "../../../scripts/api.js"; +import { addStylesheet, getUrl } from "./common/utils.js"; + +addStylesheet(getUrl("loraInfo.css", import.meta.url)); + +const MAX_TAGS = 500; + +class LoraMetadataDialog extends ComfyDialog { + constructor(name, metadata) { + super(); + + this.element.classList.add("pysssss-lora-metadata"); + } + + show(metadata) { + super.show( + $el( + "div", + Object.keys(metadata).map((k) => + $el("div", [$el("label", { textContent: k }), $el("span", { textContent: metadata[k] })]) + ) + ) + ); + } +} + +class LoraInfoDialog extends ComfyDialog { + #metadata; + + get tagFrequency() { + if (!this.#metadata.ss_tag_frequency) return []; + + const datasets = JSON.parse(this.#metadata.ss_tag_frequency); + const tags = {}; + for (const setName in datasets) { + const set = datasets[setName]; + for (const t in set) { + if (t in tags) { + tags[t] += set[t]; + } else { + tags[t] = set[t]; + } + } + } + + return Object.entries(tags).sort((a, b) => b[1] - a[1]); + } + + get resolutions() { + let res = []; + if (this.#metadata.ss_bucket_info) { + const { buckets } = JSON.parse(this.#metadata.ss_bucket_info); + for (const { resolution, count } of Object.values(buckets)) { + res.push([count, `${resolution.join("x")} * ${count}`]); + } + } + res = res.sort((a, b) => b[0] - a[0]).map((a) => a[1]); + let r = this.#metadata.ss_resolution; + if (r) { + const s = r.split(","); + const w = s[0].replace("(", ""); + const h = s[1].replace(")", ""); + res.push(`${w.trim()}x${h.trim()} (Base res)`); + } else if ((r = this.#metadata["modelspec.resolution"])) { + res.push(r + " (Base res"); + } + if (!res.length) { + res.push("⚠️ Unknown"); + } + return res; + } + + getTagList(tags) { + return tags.map((t) => + $el( + "li.pysssss-lora-tag", + { + dataset: { + tag: t[0], + }, + $: (el) => { + el.onclick = () => { + el.classList.toggle("pysssss-lora-tag--selected"); + }; + }, + }, + [ + $el("p", { + textContent: t[0], + }), + $el("span", { + textContent: t[1], + }), + ] + ) + ); + } + + constructor(name, metadata) { + super(); + + this.element.classList.add("pysssss-lora-info"); + this.#metadata = metadata; + + let tags = this.tagFrequency; + let hasMore; + if (tags?.length) { + const c = tags.length; + let list; + if (c > MAX_TAGS) { + tags = tags.slice(0, MAX_TAGS); + hasMore = $el("p", [ + $el("span", { textContent: `⚠️ Only showing first ${MAX_TAGS} tags ` }), + $el("a", { + href: "#", + textContent: `Show all ${c}`, + onclick: () => { + list.replaceChildren(...this.getTagList(this.tagFrequency)); + hasMore.remove(); + }, + }), + ]); + } + list = $el("ol.pysssss-lora-tags-list", this.getTagList(tags)); + this.tags = $el("div", [list]); + } else { + this.tags = $el("p", { textContent: "⚠️ No tag frequency metadata found" }); + } + + const resolutions = $el("label", { textContent: "Resolution:" }, [ + $el( + "select", + this.resolutions.map((r) => $el("option", { textContent: r })) + ), + ]); + + this.content = $el( + "div.pysssss-lora-content", + [ + $el("h2", { textContent: name }), + $el("p", { + textContent: "Output Name: " + (metadata.ss_output_name || "⚠️ Unknown"), + }), + $el("p", { + textContent: "Base Model: " + (metadata.ss_sd_model_name || "⚠️ Unknown"), + }), + $el("p", { + textContent: "Clip Skip: " + (metadata.ss_clip_skip || "⚠️ Unknown"), + }), + resolutions, + this.tags, + hasMore, + ].filter(Boolean) + ); + } + + createButtons() { + const btns = super.createButtons(); + + function copyTags(e, tags) { + const textarea = $el("textarea", { + parent: document.body, + style: { + position: "fixed", + }, + textContent: tags.map((el) => el.dataset.tag).join(", "), + }); + textarea.select(); + try { + document.execCommand("copy"); + if (!e.target.dataset.text) { + e.target.dataset.text = e.target.textContent; + } + e.target.textContent = "Copied " + tags.length + " tags"; + setTimeout(() => { + e.target.textContent = e.target.dataset.text; + }, 1000); + } catch (ex) { + prompt("Copy to clipboard: Ctrl+C, Enter", text); + } finally { + document.body.removeChild(textarea); + } + } + + btns.unshift( + $el("button", { + type: "button", + textContent: "Copy Selected", + onclick: (e) => { + copyTags(e, [...this.tags.querySelectorAll(".pysssss-lora-tag--selected")]); + }, + }), + $el("button", { + type: "button", + textContent: "Copy All", + onclick: (e) => { + copyTags(e, [...this.tags.querySelectorAll(".pysssss-lora-tag")]); + }, + }), + $el("button", { + type: "button", + textContent: "View raw metadata", + onclick: (e) => { + new LoraMetadataDialog().show(this.#metadata); + }, + }) + ); + return btns; + } + + show() { + super.show(this.content); + } +} + +app.registerExtension({ + name: "pysssss.LoraInfo", + beforeRegisterNodeDef(nodeType, nodeData, app) { + if (nodeType.comfyClass === "LoraLoader" || nodeType.comfyClass === "LoraLoader|pysssss") { + const getExtraMenuOptions = nodeType.prototype.getExtraMenuOptions; + nodeType.prototype.getExtraMenuOptions = function (_, options) { + let value = this.widgets[0].value; + if (!value) { + return; + } + if (value.content) { + value = value.content; + } + options.unshift({ + content: "View info...", + callback: async () => { + const meta = await (await api.fetchApi("/pysssss/metadata/" + encodeURIComponent(`loras/${value}`))).json(); + new LoraInfoDialog(value, meta).show(); + }, + }); + + return getExtraMenuOptions?.apply(this, arguments); + }; + } + }, +}); diff --git a/py/lora_info.py b/py/lora_info.py new file mode 100644 index 0000000..b93e6fb --- /dev/null +++ b/py/lora_info.py @@ -0,0 +1,37 @@ +import json +from aiohttp import web +from server import PromptServer +import folder_paths + + +def get_metadata(filepath): + with open(filepath, "rb") as file: + # https://github.com/huggingface/safetensors#format + # 8 bytes: N, an unsigned little-endian 64-bit integer, containing the size of the header + header_size = int.from_bytes(file.read(8), "little", signed=False) + + if header_size <= 0: + raise BufferError("Invalid header size") + + header = file.read(header_size) + if header_size <= 0: + raise BufferError("Invalid header") + + header_json = json.loads(header) + return header_json["__metadata__"] if "__metadata__" in header_json else None + + +@PromptServer.instance.routes.get("/pysssss/metadata/{name}") +async def load_metadata(request): + name = request.match_info["name"] + pos = name.index("/") + type = name[0:pos] + name = name[pos+1:] + + file_path = folder_paths.get_full_path( + type, name) + if not file_path: + return web.Response(status=404) + + meta = get_metadata(file_path) + return web.json_response(meta)