Added a few custom nodes

This commit is contained in:
pythongosssss
2023-03-06 19:42:04 +00:00
committed by GitHub
parent 9ad6dbd90e
commit 0371969c69
8 changed files with 631 additions and 0 deletions
+243
View File
@@ -0,0 +1,243 @@
import { app } from "../scripts/app.js";
import { ComfyWidgets } from "../scripts/widgets.js";
// Adds support for import + export as SVG including input + output images
// Adds two context menu items to the canvas
// Supports drag + drop import
// https://codepen.io/peterhry/pen/nbMaYg
function wrapText(context, text, x, y, maxWidth, lineHeight) {
var words = text.split(" "),
line = "",
i,
test,
metrics;
for (i = 0; i < words.length; i++) {
test = words[i];
metrics = context.measureText(test);
while (metrics.width > maxWidth) {
// Determine how much of the word will fit
test = test.substring(0, test.length - 1);
metrics = context.measureText(test);
}
if (words[i] != test) {
words.splice(i + 1, 0, words[i].substr(test.length));
words[i] = test;
}
test = line + words[i] + " ";
metrics = context.measureText(test);
if (metrics.width > maxWidth && i > 0) {
context.fillText(line, x, y);
line = words[i] + " ";
y += lineHeight;
} else {
line = test;
}
}
context.fillText(line, x, y);
}
function escapeXml(unsafe) {
return unsafe.replaceAll("&", "&amp;").replaceAll("<", "&lt;").replaceAll(">", "&gt;");
}
function unescapeXml(safe) {
return safe.replaceAll("&amp;", "&").replaceAll("&lt;", "<").replaceAll("&gt;", ">");
}
let saving = false;
app.registerExtension({
name: "pysssss.ExportAsSvg",
init() {
const stringWidget = ComfyWidgets.STRING;
// Override multiline string widgets to draw text using canvas while saving as svg
ComfyWidgets.STRING = function () {
const w = stringWidget.apply(this, arguments);
if (w.widget && w.widget.type === "customtext") {
const draw = w.widget.draw;
w.widget.draw = function (ctx) {
draw.apply(this, arguments);
if (saving) {
const t = ctx.getTransform();
ctx.save();
ctx.resetTransform();
const style = document.defaultView.getComputedStyle(this.inputEl, null);
const x = parseInt(this.inputEl.style.left);
const y = parseInt(this.inputEl.style.top);
const w = parseInt(this.inputEl.style.width);
const h = parseInt(this.inputEl.style.height);
ctx.fillStyle = style.getPropertyValue("background-color");
ctx.fillRect(x, y, w, h);
ctx.fillStyle = style.getPropertyValue("color");
ctx.font = style.getPropertyValue("font");
wrapText(ctx, this.inputEl.value, x, y + t.d * 12, w, t.d * 12);
ctx.restore();
}
};
}
return w;
};
},
setup(app) {
const script = document.createElement("script");
script.onload = function () {
function exportSvg() {
// Calculate the min max bounds for the nodes on the graph
const bounds = app.graph._nodes.reduce(
(p, n) => {
if (n.pos[0] < p[0]) p[0] = n.pos[0];
if (n.pos[1] < p[1]) p[1] = n.pos[1];
const r = n.pos[0] + n.size[0];
const b = n.pos[1] + n.size[1];
if (r > p[2]) p[2] = r;
if (b > p[3]) p[3] = b;
return p;
},
[99999, 99999, -99999, -99999]
);
bounds[0] -= 100;
bounds[1] -= 100;
bounds[2] += 100;
bounds[3] += 100;
// Store current canvas values to reset after drawing
const ctx = app.canvas.ctx;
const scale = app.canvas.ds.scale;
const width = app.canvas.canvas.width;
const height = app.canvas.canvas.height;
const offset = app.canvas.ds.offset;
const svgCtx = new C2S(bounds[2] - bounds[0], bounds[3] - bounds[1]);
// Override the c2s handling of images to draw images as canvases
const drawImage = svgCtx.drawImage;
svgCtx.drawImage = function (...args) {
const image = args[0];
// If we are an image node and not a datauri then we need to replace with a canvas
// we cant convert to data uri here as it is an async process
if (image.nodeName === "IMG" && !image.src.startsWith("data:image/")) {
const canvas = document.createElement("canvas");
canvas.width = image.width;
canvas.height = image.height;
const imgCtx = canvas.getContext("2d");
imgCtx.drawImage(image, 0, 0);
args[0] = canvas;
}
return drawImage.apply(this, args);
};
// Implement missing required functions
svgCtx.getTransform = function () {
return ctx.getTransform();
};
svgCtx.resetTransform = function () {
return ctx.resetTransform();
};
svgCtx.roundRect = svgCtx.rect;
// Force the canvas to render the whole graph to the svg context
app.canvas.ds.scale = 1;
app.canvas.canvas.width = bounds[2] - bounds[0];
app.canvas.canvas.height = bounds[3] - bounds[1];
app.canvas.ds.offset = [-bounds[0], -bounds[1]];
app.canvas.ctx = svgCtx;
// Trigger saving
saving = true;
app.canvas.draw(true, true);
saving = false;
// Restore original settings
app.canvas.ds.scale = scale;
app.canvas.canvas.width = width;
app.canvas.canvas.height = height;
app.canvas.ds.offset = offset;
app.canvas.ctx = ctx;
app.canvas.draw(true, true);
// Convert to SVG, embed graph and save
const json = JSON.stringify(app.graph.serialize());
const svg = svgCtx.getSerializedSvg(true).replace("</svg>", `<desc>${escapeXml(json)}</desc></svg>`);
const blob = new Blob([svg], { type: "image/svg+xml" });
const url = URL.createObjectURL(blob);
const a = document.createElement("a");
Object.assign(a, {
href: url,
download: "workflow.svg",
style: "display: none",
});
document.body.append(a);
a.click();
setTimeout(function () {
a.remove();
window.URL.revokeObjectURL(url);
}, 0);
}
let fileInput;
function importSvg() {
if (!fileInput) {
fileInput = document.createElement("input");
Object.assign(fileInput, {
type: "file",
accept: ".svg,image/svg+xml",
style: "display: none",
onchange: () => {
app.handleFile(fileInput.files[0]);
},
});
document.body.append(fileInput);
}
fileInput.click();
}
// Override file handling to allow drag & drop of SVG
const handleFile = app.handleFile;
app.handleFile = function (file) {
if (file.type === "image/svg+xml" || file.name.endsWith(".svg")) {
const reader = new FileReader();
reader.onload = () => {
// Extract embedded workflow from desc tags
const descEnd = reader.result.lastIndexOf("</desc>");
if (descEnd !== -1) {
const descStart = reader.result.lastIndexOf("<desc>", descEnd);
if (descStart !== -1) {
const json = reader.result.substring(descStart + 6, descEnd);
this.loadGraphData(JSON.parse(unescapeXml(json)));
}
}
};
reader.readAsText(file);
} else {
return handleFile.apply(this, arguments);
}
};
// Add canvas menu options
const orig = LGraphCanvas.prototype.getCanvasMenuOptions;
LGraphCanvas.prototype.getCanvasMenuOptions = function () {
const options = orig.apply(this, arguments);
options.push(
null,
{ content: "SVG -> Import", callback: importSvg },
{ content: "SVG -> Export", callback: exportSvg }
);
return options;
};
};
script.src = "http://gliffy.github.io/canvas2svg/canvas2svg.js";
document.body.append(script);
},
});
+29
View File
@@ -0,0 +1,29 @@
import comfy.utils
class LatentUpscaleBy:
upscale_methods = ["nearest-exact", "bilinear", "area"]
crop_methods = ["disabled", "center"]
@classmethod
def INPUT_TYPES(s):
return {"required": {"samples": ("LATENT",), "upscale_method": (s.upscale_methods,),
"scale": ("FLOAT", {"default": 1.5, "min": 0.1, "max": 10, "step": 0.05}),
"crop": (s.crop_methods,)}}
RETURN_TYPES = ("LATENT",)
FUNCTION = "upscale"
CATEGORY = "latent"
def upscale(self, samples, upscale_method, scale, crop):
s = samples.copy()
w = round(samples["samples"].shape[3] * 8 * scale)
h = round(samples["samples"].shape[2] * 8 * scale)
s["samples"] = comfy.utils.common_upscale(
samples["samples"], w // 8, h // 8, upscale_method, crop)
return (s,)
NODE_CLASS_MAPPINGS = {
"LatentUpscaleBy": LatentUpscaleBy,
}
+2
View File
@@ -0,0 +1,2 @@
This will store the base64 encoded version of every input image in memory until it is removed from history
May need some kind of limit on this...
+42
View File
@@ -0,0 +1,42 @@
import numpy as np
import json
from PIL.PngImagePlugin import PngInfo
from PIL import Image
import base64
from io import BytesIO
class PreviewImage:
@classmethod
def INPUT_TYPES(s):
return {"required":
{"images": ("IMAGE", ), },
"hidden": {"prompt": "PROMPT", "extra_pnginfo": "EXTRA_PNGINFO"},
}
RETURN_TYPES = ()
FUNCTION = "save_images"
OUTPUT_NODE = True
CATEGORY = "image"
def save_images(self, images, prompt=None, extra_pnginfo=None):
paths = list()
for image in images:
i = 255. * image.cpu().numpy()
img = Image.fromarray(i.astype(np.uint8))
metadata = PngInfo()
if prompt is not None:
metadata.add_text("prompt", json.dumps(prompt))
if extra_pnginfo is not None:
for x in extra_pnginfo:
metadata.add_text(x, json.dumps(extra_pnginfo[x]))
buffered = BytesIO()
img.save(buffered, format="PNG", pnginfo=metadata, optimize=True)
paths.append("data:image/png;base64," + base64.b64encode(buffered.getvalue()).decode('ascii'))
return {"ui": {"images": paths}}
NODE_CLASS_MAPPINGS = {
"PreviewImage": PreviewImage,
}
+198
View File
@@ -0,0 +1,198 @@
import { app } from "../../scripts/app.js";
// Adds a bunch of context menu entries for quickly adding common steps
// Any node with VAE input: "Use Vae" - Will find or add a VAE Loader and attach it
// KSampler: "Add Blank Input" - Connects an EmptyLatentImage
// KSampler: "Add Hi-res Fix" - Connects a LatentUpscale + second KSampler
// KSampler: "Add 2nd Pass" - Connects a new Ckpt Loader, CLIP Texts, Upscale and Sampler
// KSampler: "Add Save Image" - Connects a VAEDecode + SaveImage
// CheckpointLoaderSimple: "Add Clip Skip" - Connects a CLIPSetLastLayer node
// CheckpointLoaderSimple, CheckpointLoader, LoraLoader - "Add LORA" - Connects a new LoraLoader
// CheckpointLoaderSimple, CheckpointLoader, LoraLoader - "Add Promps" - Connects two new CLIPTextEncodes
function addMenuHandler(nodeType, cb) {
const getOpts = nodeType.prototype.getExtraMenuOptions;
nodeType.prototype.getExtraMenuOptions = function () {
const r = getOpts.apply(this, arguments);
cb.apply(this, arguments);
return r;
};
}
function getOrAddVAELoader(node) {
let vaeNode = app.graph._nodes.find((n) => n.type === "VAELoader");
if (!vaeNode) {
vaeNode = addNode("VAELoader", node);
}
return vaeNode;
}
function addNode(name, nextTo, options) {
options = { select: true, shiftY: 0, before: false, ...(options || {}) };
const node = LiteGraph.createNode(name);
app.graph.add(node);
node.pos = [
options.before ? nextTo.pos[0] - node.size[0] - 30 : nextTo.pos[0] + nextTo.size[0] + 30,
nextTo.pos[1] + options.shiftY,
];
if (options.select) {
app.canvas.selectNode(node, false);
}
return node;
}
app.registerExtension({
name: "pysssss.QuickNodes",
async beforeRegisterNodeDef(nodeType, nodeData, app) {
if (nodeData.input && nodeData.input.required) {
const keys = Object.keys(nodeData.input.required);
for (let i = 0; i < keys.length; i++) {
if (nodeData.input.required[keys[i]][0] === "VAE") {
addMenuHandler(nodeType, function (_, options) {
options.unshift({
content: "Use VAE",
callback: () => {
getOrAddVAELoader(this).connect(0, this, i);
},
});
});
break;
}
}
}
if (nodeData.name === "KSampler") {
addMenuHandler(nodeType, function (_, options) {
options.unshift(
{
content: "Add Blank Input",
callback: () => {
const imageNode = addNode("EmptyLatentImage", this, { before: true });
imageNode.connect(0, this, 3);
},
},
{
content: "Add Hi-res Fix",
callback: () => {
const upscaleNode = addNode("LatentUpscale", this);
this.connect(0, upscaleNode, 0);
const sampleNode = addNode("KSampler", upscaleNode);
for (let i = 0; i < 3; i++) {
const l = this.getInputLink(i);
if (l) {
app.graph.getNodeById(l.origin_id).connect(l.origin_slot, sampleNode, i);
}
}
upscaleNode.connect(0, sampleNode, 3);
},
},
{
content: "Add 2nd Pass",
callback: () => {
const upscaleNode = addNode("LatentUpscale", this);
this.connect(0, upscaleNode, 0);
const ckptNode = addNode("CheckpointLoaderSimple", this);
const sampleNode = addNode("KSampler", ckptNode);
const positiveLink = this.getInputLink(1);
const negativeLink = this.getInputLink(2);
const positiveNode = positiveLink
? app.graph.add(app.graph.getNodeById(positiveLink.origin_id).clone())
: addNode("CLIPTextEncode");
const negativeNode = negativeLink
? app.graph.add(app.graph.getNodeById(negativeLink.origin_id).clone())
: addNode("CLIPTextEncode");
ckptNode.connect(0, sampleNode, 0);
ckptNode.connect(1, positiveNode, 0);
ckptNode.connect(1, negativeNode, 0);
positiveNode.connect(0, sampleNode, 1);
negativeNode.connect(0, sampleNode, 2);
upscaleNode.connect(0, sampleNode, 3);
},
},
{
content: "Add Save Image",
callback: () => {
const decodeNode = addNode("VAEDecode", this);
this.connect(0, decodeNode, 0);
getOrAddVAELoader(decodeNode).connect(0, decodeNode, 1);
const saveNode = addNode("SaveImage", decodeNode);
decodeNode.connect(0, saveNode, 0);
},
}
);
});
}
if (nodeData.name === "CheckpointLoaderSimple") {
addMenuHandler(nodeType, function (_, options) {
options.unshift({
content: "Add Clip Skip",
callback: () => {
const clipSkipNode = addNode("CLIPSetLastLayer", this);
const clipLinks = this.outputs[1].links ? this.outputs[1].links.map((l) => ({ ...graph.links[l] })) : [];
this.disconnectOutput(1);
this.connect(1, clipSkipNode, 0);
for (const clipLink of clipLinks) {
clipSkipNode.connect(0, clipLink.target_id, clipLink.target_slot);
}
}
});
});
}
if (
nodeData.name === "CheckpointLoaderSimple" ||
nodeData.name === "CheckpointLoader" ||
nodeData.name === "LoraLoader"
) {
addMenuHandler(nodeType, function (_, options) {
options.unshift(
{
content: "Add LORA",
callback: () => {
const loraNode = addNode("LoraLoader", this);
const modelLinks = this.outputs[0].links ? this.outputs[0].links.map((l) => ({ ...graph.links[l] })) : [];
const clipLinks = this.outputs[1].links ? this.outputs[1].links.map((l) => ({ ...graph.links[l] })) : [];
this.disconnectOutput(0);
this.disconnectOutput(1);
this.connect(0, loraNode, 0);
this.connect(1, loraNode, 1);
for (const modelLink of modelLinks) {
loraNode.connect(0, modelLink.target_id, modelLink.target_slot);
}
for (const clipLink of clipLinks) {
loraNode.connect(1, clipLink.target_id, clipLink.target_slot);
}
},
},
{
content: "Add Prompts",
callback: () => {
const positiveNode = addNode("CLIPTextEncode", this);
const negativeNode = addNode("CLIPTextEncode", this, { shiftY: positiveNode.size[1] + 30 });
this.connect(1, positiveNode, 0);
this.connect(1, negativeNode, 0);
},
}
);
});
}
},
});
+28
View File
@@ -0,0 +1,28 @@
This will embed the selected image into the generated images which probably isnt desired
Current fix for this is to add this to execution.py:
def prune_prompt(prompt):
pruned_prompt = {}
for n in prompt:
class_type = prompt[n]['class_type']
class_def = nodes.NODE_CLASS_MAPPINGS[class_type]
valid_inputs = class_def.INPUT_TYPES()
pruned_prompt[n] = copy.deepcopy(prompt[n])
if "inputs" in pruned_prompt[n]:
for x in pruned_prompt[n]["inputs"]:
input_type = None
if ("required" in valid_inputs and x in valid_inputs["required"]):
input_type = valid_inputs["required"][x]
elif ("optional" in valid_inputs and x in valid_inputs["optional"]):
input_type = valid_inputs["optional"][x]
if input_type is not None and input_type[0] == "B64IMAGE":
pruned_prompt[n]["inputs"][x] = None
return pruned_prompt
And call that from
if h[x] == "PROMPT":
input_data_all[x] = prune_prompt(prompt)
+54
View File
@@ -0,0 +1,54 @@
import { app } from "../scripts/app.js";
// Adds a new UploadImage node
const toBase64 = (file) =>
new Promise((resolve, reject) => {
const reader = new FileReader();
reader.readAsDataURL(file);
reader.onload = () => resolve(reader.result);
reader.onerror = (error) => reject(error);
});
app.registerExtension({
name: "Comfy.UploadImage",
async getCustomWidgets() {
return {
B64IMAGE(node) {
let uploadWidget;
const fileInput = document.createElement("input");
Object.assign(fileInput, {
type: "file",
accept: "image/jpeg,image/png",
style: "display: none",
onchange: () => {
if (fileInput.files.length) {
const img = new Image();
img.onload = () => {
node.imgs = [img];
};
toBase64(fileInput.files[0]).then((d) => {
img.src = d;
});
}
},
});
document.body.append(fileInput);
uploadWidget = node.addWidget("button", "image", "image", () => {
fileInput.click();
});
uploadWidget.serializeValue = () => {
if(node.imgs && node.imgs.length) {
return node.imgs[0].src;
}
return null;
};
return { widget: uploadWidget };
},
};
},
});
+35
View File
@@ -0,0 +1,35 @@
from PIL import Image
import numpy as np
import torch
class UploadImage:
@classmethod
def INPUT_TYPES(s):
return {"required":
{"image": ("B64IMAGE",)},
}
CATEGORY = "image"
RETURN_TYPES = ("IMAGE",)
FUNCTION = "load_image"
def load_image(self, image):
from io import BytesIO
import re
import base64
if image.startswith("data:image/"):
image_data = re.sub('^data:image/.+;base64,', '', image)
i = Image.open(BytesIO(base64.b64decode(image_data)))
else:
raise Exception("Invalid image data")
image = i.convert("RGB")
image = np.array(image).astype(np.float32) / 255.0
image = torch.from_numpy(image)[None,]
return (image,)
NODE_CLASS_MAPPINGS = {
"UploadImage": UploadImage,
}