Merge pull request #156 from qrtal/main

Add autocomplete for LoRAs
This commit is contained in:
pythongosssss
2024-01-17 19:30:37 +00:00
committed by GitHub
5 changed files with 73 additions and 4 deletions
+7
View File
@@ -1,6 +1,7 @@
from server import PromptServer
from aiohttp import web
import os
import folder_paths
dir = os.path.abspath(os.path.join(__file__, "../../user"))
if not os.path.exists(dir):
@@ -20,3 +21,9 @@ async def update_autocomplete(request):
with open(file, "w", encoding="utf-8") as f:
f.write(await request.text())
return web.Response(status=200)
@PromptServer.instance.routes.get("/pysssss/loras")
async def get_loras(request):
loras = folder_paths.get_filename_list("loras")
return web.json_response(list(map(lambda a: os.path.splitext(a)[0], loras)))
+2 -2
View File
@@ -31,7 +31,7 @@ async def save_notes(request):
name = name[pos+1:]
file_path = None
if type == "embeddings":
if type == "embeddings" or type == "loras":
name = name.lower()
files = folder_paths.get_filename_list(type)
for f in files:
@@ -67,7 +67,7 @@ async def load_metadata(request):
name = name[pos+1:]
file_path = None
if type == "embeddings":
if type == "embeddings" or type == "loras":
name = name.lower()
files = folder_paths.get_filename_list(type)
for f in files:
+60 -1
View File
@@ -4,6 +4,7 @@ import { api } from "../../../../scripts/api.js";
import { $el, ComfyDialog } from "../../../../scripts/ui.js";
import { TextAreaAutoComplete } from "./common/autocomplete.js";
import { ModelInfoDialog } from "./common/modelInfoDialog.js";
import { LoraInfoDialog } from "./modelInfo.js";
function parseCSV(csvText) {
const rows = [];
@@ -127,6 +128,13 @@ async function addCustomWords(text) {
}
}
function toggleLoras() {
[TextAreaAutoComplete.globalWords, TextAreaAutoComplete.globalWordsExclLoras] = [
TextAreaAutoComplete.globalWordsExclLoras,
TextAreaAutoComplete.globalWords,
];
}
class EmbeddingInfoDialog extends ModelInfoDialog {
async addInfo() {
super.addInfo();
@@ -269,7 +277,38 @@ app.registerExtension({
TextAreaAutoComplete.updateWords("pysssss.embeddings", words);
}
Promise.all([addEmbeddings(), addCustomWords()]);
async function addLoras() {
const loras = await api
.fetchApi("/pysssss/loras", { cache: "no-store" })
.then(res => res.json());
const words = {};
words["lora:"] = { text: "lora:" };
for (const lora of loras) {
const v = `<lora:${lora}:1.0>`;
words[v] = {
text: v,
info: () => new LoraInfoDialog(lora).show("loras", lora),
};
}
TextAreaAutoComplete.updateWords("pysssss.loras", words);
}
// store global words with/without loras
Promise.all([addEmbeddings(), addCustomWords()])
.then(() => {
TextAreaAutoComplete.globalWordsExclLoras = Object.assign(
{},
TextAreaAutoComplete.globalWords
);
})
.then(addLoras)
.then(() => {
if (!TextAreaAutoComplete.lorasEnabled) {
toggleLoras(); // off by default
}
});
const STRING = ComfyWidgets.STRING;
const SKIP_WIDGETS = new Set(["ttN xyPlot.x_values", "ttN xyPlot.y_values"]);
@@ -347,6 +386,26 @@ app.registerExtension({
}),
]
),
$el(
"label",
{
textContent: "Loras enabled ",
style: {
display: "block",
},
},
[
$el("input", {
type: "checkbox",
checked: !!TextAreaAutoComplete.lorasEnabled,
onchange: (event) => {
const checked = !!event.target.checked;
TextAreaAutoComplete.lorasEnabled = checked;
toggleLoras();
},
}),
]
),
$el(
"label",
{
+3
View File
@@ -317,6 +317,7 @@ export class TextAreaAutoComplete {
static insertOnTab = true;
static insertOnEnter = true;
static replacer = undefined;
static lorasEnabled = false;
/** @type {Record<string, Record<string, AutoCompleteEntry>>} */
static groups = {};
@@ -324,6 +325,8 @@ export class TextAreaAutoComplete {
static globalGroups = new Set();
/** @type {Record<string, AutoCompleteEntry>} */
static globalWords = {};
/** @type {Record<string, AutoCompleteEntry>} */
static globalWordsExclLoras = {};
/** @type {HTMLTextAreaElement} */
el;
+1 -1
View File
@@ -4,7 +4,7 @@ import { ModelInfoDialog } from "./common/modelInfoDialog.js";
const MAX_TAGS = 500;
class LoraInfoDialog extends ModelInfoDialog {
export class LoraInfoDialog extends ModelInfoDialog {
getTagFrequency() {
if (!this.metadata.ss_tag_frequency) return [];