Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
8 changes: 7 additions & 1 deletion py/better_combos.py
Original file line number Diff line number Diff line change
@@ -1,6 +1,6 @@
import glob
import os
from nodes import LoraLoader, CheckpointLoaderSimple
from nodes import LoraLoader, CheckpointLoaderSimple, UNETLoader
import folder_paths
from server import PromptServer
from folder_paths import get_directory_by_type
Expand Down Expand Up @@ -161,12 +161,18 @@ def load_checkpoint(self, **kwargs):
return (*super().load_checkpoint(**kwargs), prompt)


class UNETLoaderWithImages(UNETLoader):
pass


NODE_CLASS_MAPPINGS = {
"LoraLoader|pysssss": LoraLoaderWithImages,
"CheckpointLoader|pysssss": CheckpointLoaderSimpleWithImages,
"UNETLoader|pysssss": UNETLoaderWithImages,
}

NODE_DISPLAY_NAME_MAPPINGS = {
"LoraLoader|pysssss": "Lora Loader 馃悕",
"CheckpointLoader|pysssss": "Checkpoint Loader 馃悕",
"UNETLoader|pysssss": "Load Diffusion Model 馃悕",
}
45 changes: 31 additions & 14 deletions web/js/betterCombos.js
Original file line number Diff line number Diff line change
Expand Up @@ -5,18 +5,36 @@ import { api } from "../../../scripts/api.js";

const CHECKPOINT_LOADER = "CheckpointLoader|pysssss";
const LORA_LOADER = "LoraLoader|pysssss";
const DIFFUSION_MODEL_LOADER = "UNETLoader|pysssss";
const IMAGE_WIDTH = 384;
const IMAGE_HEIGHT = 384;

function getType(node) {
if (node.comfyClass === CHECKPOINT_LOADER) {
return "checkpoints";
}
if (node.comfyClass === DIFFUSION_MODEL_LOADER) {
return "diffusion_models";
}
return "loras";
}

function getWidgetName(type) {
return type === "checkpoints" ? "ckpt_name" : "lora_name";
if (type === "checkpoints") {
return "ckpt_name";
}
if (type === "diffusion_models") {
return "unet_name";
}
return "lora_name";
}

function isImageLoader(node) {
return node?.comfyClass === LORA_LOADER || node?.comfyClass === CHECKPOINT_LOADER || node?.comfyClass === DIFFUSION_MODEL_LOADER;
}

function hasExamples(nodeData) {
return nodeData.name === LORA_LOADER || nodeData.name === CHECKPOINT_LOADER;
}

function encodeRFC3986URIComponent(str) {
Expand Down Expand Up @@ -74,7 +92,7 @@ app.registerExtension({
const displayOptions = { "List (normal)": 0, "Tree (subfolders)": 1, "Thumbnails (grid)": 2 };
const displaySetting = app.ui.settings.addSetting({
id: "pysssss.Combo++.Submenu",
name: "馃悕 Lora & Checkpoint loader display mode",
name: "馃悕 Model loader display mode",
defaultValue: 1,
type: "combo",
options: (value) => {
Expand Down Expand Up @@ -157,15 +175,14 @@ app.registerExtension({
`,
parent: document.body,
});
const p1 = loadImageList("checkpoints");
const p2 = loadImageList("loras");
const modelTypes = ["checkpoints", "loras", "diffusion_models"];
const modelLists = Promise.all(modelTypes.map((type) => loadImageList(type)));

const refreshComboInNodes = app.refreshComboInNodes;
app.refreshComboInNodes = async function () {
const r = await Promise.all([
refreshComboInNodes.apply(this, arguments),
loadImageList("checkpoints").catch(() => {}),
loadImageList("loras").catch(() => {}),
...modelTypes.map((type) => loadImageList(type).catch(() => {})),
]);
return r[0];
};
Expand All @@ -192,8 +209,7 @@ app.registerExtension({

const updateMenu = async (menu, type) => {
try {
await p1;
await p2;
await modelLists;
} catch (error) {
console.error(error);
console.error("Error loading pysssss.betterCombos data");
Expand Down Expand Up @@ -365,7 +381,7 @@ app.registerExtension({
const mutationObserver = new MutationObserver((mutations) => {
const node = app.canvas.current_node;

if (!node || (node.comfyClass !== LORA_LOADER && node.comfyClass !== CHECKPOINT_LOADER)) {
if (!isImageLoader(node)) {
return;
}

Expand Down Expand Up @@ -395,14 +411,15 @@ app.registerExtension({
mutationObserver.observe(document.body, { childList: true, subtree: false });
},
async beforeRegisterNodeDef(nodeType, nodeData, app) {
const isCkpt = nodeData.name === CHECKPOINT_LOADER;
const isLora = nodeData.name === LORA_LOADER;
if (isCkpt || isLora) {
if (hasExamples(nodeData)) {
const onAdded = nodeType.prototype.onAdded;
nodeType.prototype.onAdded = function () {
onAdded?.apply(this, arguments);
const { widget: exampleList } = ComfyWidgets["COMBO"](this, "example", [[""], {}], app);
this.widgets.find((w) => w.name === "prompt").computeSize = () => [0, -4];
const promptWidget = this.widgets.find((w) => w.name === "prompt");
if (promptWidget) {
promptWidget.computeSize = () => [0, -4];
}
let exampleWidget;

const get = async (route, suffix) => {
Expand Down Expand Up @@ -500,7 +517,7 @@ app.registerExtension({
img = this.imgs[this.overIndex];
}
if (img) {
const nodes = app.graph._nodes.filter((n) => n.comfyClass === LORA_LOADER || n.comfyClass === CHECKPOINT_LOADER);
const nodes = app.graph._nodes.filter((n) => isImageLoader(n));
if (nodes.length) {
options.unshift({
content: "Save as Preview",
Expand Down
16 changes: 12 additions & 4 deletions web/js/modelInfo.js
Original file line number Diff line number Diff line change
Expand Up @@ -300,6 +300,11 @@ class CheckpointInfoDialog extends ModelInfoDialog {
}

const lookups = {};
const modelInfoTypes = {
Lora: { folderType: "loras", infoClass: LoraInfoDialog },
Checkpoint: { folderType: "checkpoints", infoClass: CheckpointInfoDialog },
"Diffusion Model": { folderType: "diffusion_models", infoClass: CheckpointInfoDialog },
};

function addInfoOption(node, type, infoClass, widgetNamePattern, opts) {
const widgets = widgetNamePattern
Expand All @@ -326,12 +331,14 @@ function addInfoOption(node, type, infoClass, widgetNamePattern, opts) {
}

function addTypeOptions(node, typeName, options) {
const type = typeName.toLowerCase() + "s";
const typeInfo = modelInfoTypes[typeName];
if (!typeInfo) return;

const { folderType: type, infoClass: cls } = typeInfo;
const values = lookups[typeName][node.type];
if (!values) return;

const widgets = Object.keys(values);
const cls = type === "loras" ? LoraInfoDialog : CheckpointInfoDialog;

const opts = [];
for (const w of widgets) {
Expand All @@ -357,9 +364,9 @@ function addTypeOptions(node, typeName, options) {
app.registerExtension({
name: "pysssss.ModelInfo",
setup() {
const addSetting = (type, defaultValue) => {
const addSetting = (type, defaultValue, idType = type) => {
app.ui.settings.addSetting({
id: `pysssss.ModelInfo.${type}Nodes`,
id: `pysssss.ModelInfo.${idType}Nodes`,
name: `馃悕 Model Info - ${type} Nodes/Widgets`,
type: "text",
defaultValue,
Expand All @@ -384,6 +391,7 @@ app.registerExtension({
"Checkpoint",
["CheckpointLoader.ckpt_name", "CheckpointLoaderSimple", "CheckpointLoader|pysssss", "Efficient Loader", "Eff. Loader SDXL"].join(",")
);
addSetting("Diffusion Model", ["UNETLoader.unet_name", "UNETLoader|pysssss"].join(","), "DiffusionModel");

app.ui.settings.addSetting({
id: `pysssss.ModelInfo.NsfwLevel`,
Expand Down