From dc0480c2959ca13e9478baec566250b2f7c2ed3a Mon Sep 17 00:00:00 2001 From: pythongosssss <125205205+pythongosssss@users.noreply.github.com> Date: Sun, 14 Jul 2024 10:56:22 +0100 Subject: [PATCH] Better lora examples --- py/better_combos.py | 45 ++++++++++++++++++++++++-- web/js/betterCombos.js | 20 ++++++++++-- web/js/modelInfo.js | 73 +++++++++++++++++++++++++++++++++++------- 3 files changed, 122 insertions(+), 16 deletions(-) diff --git a/py/better_combos.py b/py/better_combos.py index 46f3efa..99e3e01 100644 --- a/py/better_combos.py +++ b/py/better_combos.py @@ -66,12 +66,45 @@ async def get_examples(request): file_path_no_ext = os.path.splitext(file_path)[0] examples = [] + + if os.path.isfile(file_path_no_ext + ".txt"): + examples += ["notes"] + if os.path.isdir(file_path_no_ext): - examples += map(lambda t: os.path.relpath(t, file_path_no_ext), - glob.glob(file_path_no_ext + "/*.txt")) + examples += sorted(map(lambda t: os.path.relpath(t, file_path_no_ext), + glob.glob(file_path_no_ext + "/*.txt"))) return web.json_response(examples) +@PromptServer.instance.routes.post("/pysssss/examples/{name}") +async def save_example(request): + name = request.match_info["name"] + pos = name.index("/") + type = name[0:pos] + name = name[pos+1:] + body = await request.json() + example_name = body["name"] + example = body["example"] + + file_path = folder_paths.get_full_path( + type, name) + if not file_path: + return web.Response(status=404) + + if not example_name.endswith(".txt"): + example_name += ".txt" + + file_path_no_ext = os.path.splitext(file_path)[0] + file_name = os.path.split(file_path_no_ext)[1] + example_path = os.path.join(file_path_no_ext, file_name) + example_file = os.path.join(example_path, example_name) + if not os.path.exists(example_path): + os.mkdir(example_path) + with open(example_file, 'w', encoding='utf8') as f: + f.write(example) + + return web.Response(status=201) + def populate_items(names, type): for idx, item_name in enumerate(names): @@ -99,11 +132,16 @@ def populate_items(names, type): class LoraLoaderWithImages(LoraLoader): + RETURN_TYPES = ("MODEL", "CLIP", "STRING") + @classmethod def INPUT_TYPES(s): types = super().INPUT_TYPES() names = types["required"]["lora_name"][0] populate_items(names, "loras") + + types["optional"] = { "prompt": ("HIDDEN",) } + return types @classmethod @@ -119,7 +157,8 @@ class LoraLoaderWithImages(LoraLoader): def load_lora(self, **kwargs): kwargs["lora_name"] = kwargs["lora_name"]["content"] - return super().load_lora(**kwargs) + prompt = kwargs.pop("prompt", "") + return (*super().load_lora(**kwargs), prompt) class CheckpointLoaderSimpleWithImages(CheckpointLoaderSimple): diff --git a/web/js/betterCombos.js b/web/js/betterCombos.js index a3e080b..dd0689f 100644 --- a/web/js/betterCombos.js +++ b/web/js/betterCombos.js @@ -246,8 +246,14 @@ app.registerExtension({ const v = this.widgets[0].value.content; const pos = v.lastIndexOf("."); const name = v.substr(0, pos); - - const example = await (await get("view", `/${name}/${exampleList.value}`)).text(); + let exampleName = exampleList.value; + let viewPath = `/${name}`; + if (exampleName === "notes") { + viewPath += ".txt"; + } else { + viewPath += `/${exampleName}`; + } + const example = await (await get("view", viewPath)).text(); if (!exampleWidget) { exampleWidget = ComfyWidgets["STRING"](this, "prompt", ["STRING", { multiline: true }], app).widget; exampleWidget.inputEl.readOnly = true; @@ -273,6 +279,7 @@ app.registerExtension({ } catch (error) {} } exampleList.options.values = ["[none]", ...examples]; + exampleList.value = exampleList.options.values[+!!examples.length]; exampleList.callback(); exampleList.disabled = !examples.length; app.graph.setDirtyCanvas(true, true); @@ -297,6 +304,15 @@ app.registerExtension({ modelWidget.callback(); }, 30); }; + + if (isLora) { + // Prevent adding HIDDEN inputs + const addInput = nodeType.prototype.addInput ?? LGraphNode.prototype.addInput; + nodeType.prototype.addInput = function (_, type) { + if (type === "HIDDEN") return; + return addInput.apply(this, arguments); + }; + } } const getExtraMenuOptions = nodeType.prototype.getExtraMenuOptions; diff --git a/web/js/modelInfo.js b/web/js/modelInfo.js index 838dcbd..6e7b8a4 100644 --- a/web/js/modelInfo.js +++ b/web/js/modelInfo.js @@ -1,4 +1,5 @@ import { app } from "../../../scripts/app.js"; +import { api } from "../../../scripts/api.js"; import { $el } from "../../../scripts/ui.js"; import { ModelInfoDialog } from "./common/modelInfoDialog.js"; @@ -128,6 +129,22 @@ export class LoraInfoDialog extends ModelInfoDialog { const info = await p; if (info) { + const textArea = $el("textarea", { + textContent: info.trainedWords.join(", "), + style: { + whiteSpace: "pre-wrap", + margin: "10px 0", + color: "#fff", + background: "#222", + padding: "5px", + borderRadius: "5px", + maxHeight: "250px", + overflow: "auto", + display: "block", + border: "none", + width: "calc(100% - 10px)", + }, + }); $el( "p", { @@ -135,18 +152,17 @@ export class LoraInfoDialog extends ModelInfoDialog { textContent: "Trained Words: ", }, [ - $el("pre", { - textContent: info.trainedWords.join(", "), + textArea, + $el("button", { + onclick: async () => { + await this.saveAsExample(textArea.value, "trainedwords.txt"); + }, + textContent: "Save as Example", style: { - whiteSpace: "pre-wrap", - margin: "10px 0", - background: "#222", - padding: "5px", - borderRadius: "5px", - maxHeight: "250px", - overflow: "auto", + fontSize: "14px", }, }), + $el("hr"), ] ); $el("div", { @@ -160,16 +176,43 @@ export class LoraInfoDialog extends ModelInfoDialog { } } + async saveAsExample(example, name = "example.txt") { + if (!example.length) { + return; + } + try { + name = prompt("Enter example name", name); + if (!name) return; + + await api.fetchApi("/pysssss/examples/" + encodeURIComponent(`${this.type}/${this.name}`), { + method: "POST", + body: JSON.stringify({ + name, + example, + }), + headers: { + "content-type": "application/json", + }, + }); + alert("Saved!"); + } catch (error) { + console.error(error); + alert("Error saving: " + error); + } + } + createButtons() { const btns = super.createButtons(); - + function tagsToCsv(tags) { + return tags.map((el) => el.dataset.tag).join(", "); + } function copyTags(e, tags) { const textarea = $el("textarea", { parent: document.body, style: { position: "fixed", }, - textContent: tags.map((el) => el.dataset.tag).join(", "), + textContent: tagsToCsv(tags), }); textarea.select(); try { @@ -189,6 +232,14 @@ export class LoraInfoDialog extends ModelInfoDialog { } btns.unshift( + $el("button", { + type: "button", + textContent: "Save Selected as Example", + onclick: async (e) => { + const tags = tagsToCsv([...this.tags.querySelectorAll(".pysssss-model-tag--selected")]); + await this.saveAsExample(tags); + }, + }), $el("button", { type: "button", textContent: "Copy Selected",