Added a few custom nodes
This commit is contained in:
@@ -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("&", "&").replaceAll("<", "<").replaceAll(">", ">");
|
||||
}
|
||||
function unescapeXml(safe) {
|
||||
return safe.replaceAll("&", "&").replaceAll("<", "<").replaceAll(">", ">");
|
||||
}
|
||||
|
||||
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);
|
||||
},
|
||||
});
|
||||
@@ -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,
|
||||
}
|
||||
@@ -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...
|
||||
@@ -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,
|
||||
}
|
||||
@@ -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);
|
||||
},
|
||||
}
|
||||
);
|
||||
});
|
||||
}
|
||||
},
|
||||
});
|
||||
@@ -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)
|
||||
@@ -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 };
|
||||
},
|
||||
};
|
||||
},
|
||||
});
|
||||
@@ -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,
|
||||
}
|
||||
Reference in New Issue
Block a user