13 Commits
Author SHA1 Message Date
kuschanow 57c4fb855a feat: refactor offset widget to use custom canvas; update README and version to 3.1.0 2026-08-11 14:00:57 +03:00
kuschanow 68bc65783d chore: bump version to 3.0.1 in pyproject.toml 2026-08-10 23:32:17 +03:00
kuschanow 2e7a3bf3d0 fix: update GitHub Actions condition to reflect correct repository owner 2026-08-10 23:29:53 +03:00
kuschanow 551835195e feat: update transforms to operate on model instead of latent; bug fixes; version bump to 3.0.0 2026-08-10 23:15:09 +03:00
Roman Kushanov 9e685f9f2d Merge pull request #7 from ComfyNodePRs/update-publish-yaml
Update Github Action for Publishing to Comfy Registry
2025-03-27 19:57:43 +02:00
snomiao df4d4210f1 chore(publish): update workflow for node publishing
- Added permissions to allow issue writing.
- Updated condition to run job only for specific repository owner.
- Changed action version from `main` to `v1` for stability.
2025-01-21 08:45:27 +00:00
RomanKuschanow a92091a8f2 add publisher id 2024-06-21 10:29:08 +03:00
Kuschanow Roman e8ea84b4cf Merge pull request #4 from ComfyNodePRs/pyproject
Add pyproject.toml for Custom Node Registry
2024-06-21 10:26:33 +03:00
Kuschanow Roman 589c3626e1 Merge pull request #3 from ComfyNodePRs/publish
Add Github Action for Publishing to Comfy Registry
2024-06-20 20:52:31 +03:00
snomiao 36a76e06dc chore(pyproject): Add pyproject.toml for Custom Node Registry 2024-06-14 08:10:54 +00:00
snomiao 20496551ad chore(publish): Add Github Action for Publishing to Comfy Registry 2024-06-14 08:10:54 +00:00
RomanKuschanow b88bb480fc add preview for mirror and shift 2024-05-05 12:17:04 +03:00
RomanKuschanow 7bd9377e94 fix square size 2024-04-04 18:03:25 +03:00
29 changed files with 291 additions and 318 deletions
+26
View File
@@ -0,0 +1,26 @@
name: Publish to Comfy registry
on:
workflow_dispatch:
push:
branches:
- main
- master
paths:
- "pyproject.toml"
permissions:
issues: write
jobs:
publish-node:
name: Publish Custom Node to registry
runs-on: ubuntu-latest
if: ${{ github.repository_owner == 'kuschanow' }}
steps:
- name: Check out code
uses: actions/checkout@v4
- name: Publish Custom Node
uses: Comfy-Org/publish-node-action@v1
with:
## Add your own personal access token to your Github Repository secrets and reference it here.
personal_access_token: ${{ secrets.REGISTRY_ACCESS_TOKEN }}
+27 -18
View File
@@ -19,8 +19,8 @@ This node can flip latent and merge original and flipped version.
- `latent` - `latent`
**Usage:** **Usage:**
![sample](https://i.imgur.com/YMyYorQ.png) ![sample](assets/latent-mirror-usage-1.png)
![sample](https://i.imgur.com/W5BasCO.png) ![sample](assets/latent-mirror-usage-2.png)
## Latent shift ## Latent shift
This node can shift latent along x and y-axis. This node can shift latent along x and y-axis.
@@ -36,7 +36,7 @@ This node can shift latent along x and y-axis.
- `latent` - `latent`
**Usage:** **Usage:**
![sample](https://i.imgur.com/1Dp5dSw.png) ![sample](assets/latent-shift-usage.png)
## ~~TSampler with transforms (Latent Control)~~ ## ~~TSampler with transforms (Latent Control)~~
Removed from version 2.0.0 Removed from version 2.0.0
@@ -55,8 +55,8 @@ exactly matches the base KSampler
exactly matches the base KSampler exactly matches the base KSampler
**Usage:** **Usage:**
![sample](https://i.imgur.com/PlGnAtA.png) ![sample](assets/tsampler-usage-1.png)
![sample](https://i.imgur.com/CtrBRPn.png) ![sample](assets/tsampler-usage-2.png)
Multiply, Mirror and Shift transform nodes parameters exactly match the corresponding `KSampler with transforms (Latent Control)` parameters. Multiply, Mirror and Shift transform nodes parameters exactly match the corresponding `KSampler with transforms (Latent Control)` parameters.
@@ -78,9 +78,9 @@ You can apply specific offset for transform nodes.
- `offset` - `offset`
**Usage:** **Usage:**
![sample](https://i.imgur.com/ExZacqG.png) ![sample](assets/offset-usage-1.png)
![sample](https://i.imgur.com/tR6KSmI.png) ![sample](assets/offset-usage-2.png)
![sample](https://i.imgur.com/MGVLfve.png) ![sample](assets/offset-usage-3.png)
You can combine different offsets to achieve interesting patterns. For example: You can combine different offsets to achieve interesting patterns. For example:
**0 0 0 1** and **0 0 1** give this pattern: **0 0 1 1 0 1 0 1 1 0 0 1**. **0 0 0 1** and **0 0 1** give this pattern: **0 0 1 1 0 1 0 1 1 0 0 1**.
@@ -90,7 +90,7 @@ You can combine different offsets to achieve interesting patterns. For example:
Each transform node has own one-time version. They allow to make one transform action at specified step. Each transform node has own one-time version. They allow to make one transform action at specified step.
**Usage:** **Usage:**
![sample](https://i.imgur.com/Q1Vyob0.png) ![sample](assets/one-time-nodes-usage.png)
## Latent normalize ## Latent normalize
@@ -105,33 +105,42 @@ exactly matches the `VAE Decode` node
When you multiply latent by negative or big positive (bigger than 2) number and paste this latent in sampler, you can see that the When you multiply latent by negative or big positive (bigger than 2) number and paste this latent in sampler, you can see that the
image will be generated very poorly. This is because stable diffusion cannot work with such set of numbers (meaning the numbers contained in latent). image will be generated very poorly. This is because stable diffusion cannot work with such set of numbers (meaning the numbers contained in latent).
![sample](https://i.imgur.com/3FXk8n7.png) ![sample](assets/latent-normalize-poor-result.png)
But you can prevent this behavior by sequential decode and encode latent using vae. Node `Latent normalize` make this process easier. But you can prevent this behavior by sequential decode and encode latent using vae. Node `Latent normalize` make this process easier.
![sample](https://i.imgur.com/hkFYYVh.png) ![sample](assets/latent-normalize-fix.png)
This node also change some results even if output without this node looks good. This node also change some results even if output without this node looks good.
![sample](https://i.imgur.com/kP0f6vh.png) ![sample](assets/latent-normalize-comparison-1.png)
![sample](https://i.imgur.com/YI8ZqLd.png) ![sample](assets/latent-normalize-comparison-2.png)
And it very slightly changes results from latent, which have not been modified. And it very slightly changes results from latent, which have not been modified.
![sample](https://i.imgur.com/xTU08xm.png) ![sample](assets/latent-normalize-unmodified-1.png)
![sample](https://i.imgur.com/yzgW7QT.png) ![sample](assets/latent-normalize-unmodified-2.png)
## Transform hijack ## Transform hijack
Allow you to use transforms with any samplers that you like. Allow you to use transforms with any samplers that you like.
Instead of patching the latent, this node patches the **model**: it attaches the transforms
to a cloned model via a sampler post-cfg hook. Connect the returned model to any sampler
(`KSampler`, `KSamplerAdvanced`, custom sampler nodes, etc.) and the transforms will be applied
during sampling.
**Inputs:** **Inputs:**
- latent - model
- transforms - transforms
**Outputs:** **Outputs:**
- latent - model
**Usage:** **Usage:**
![sample](https://i.imgur.com/YwVhHYF.png) ![sample](assets/transform-hijack-usage.png)
> **Breaking change in 3.0.0:** `Transform hijack` now takes and returns a `MODEL` instead of a
> `LATENT`. This replaces the old global `common_ksampler` monkey-patch, which conflicted with the
> stock `KSampler` and other custom nodes. Rewire this node to your model input/output after updating.
+2 -2
View File
@@ -38,8 +38,8 @@ NODE_DISPLAY_NAME_MAPPINGS = {
"LatentInterpolateTransform": "Latent interpolate transform", "LatentInterpolateTransform": "Latent interpolate transform",
"LatentAddTransform": "Latent add transform", "LatentAddTransform": "Latent add transform",
"OneTimeMirrorTransform": "Mirror transform (one time)", "OneTimeMirrorTransform": "Mirror transform (one time)",
"OneTimeMultiplyTransform": "Shift transform (one time)", "OneTimeMultiplyTransform": "Multiply transform (one time)",
"OneTimeShiftTransform": "Multiply transform (one time)", "OneTimeShiftTransform": "Shift transform (one time)",
"OneTimeLatentInterpolateTransform": "Latent interpolate transform (one time)", "OneTimeLatentInterpolateTransform": "Latent interpolate transform (one time)",
"OneTimeLatentAddTransform": "Latent add transform (one time)", "OneTimeLatentAddTransform": "Latent add transform (one time)",
"TransformsCombine": "Combine transforms", "TransformsCombine": "Combine transforms",
Binary file not shown.

After

Width:  |  Height:  |  Size: 502 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 264 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 755 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 777 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 732 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 597 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 787 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 789 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 352 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 314 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 315 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 300 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 565 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 355 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 204 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 290 KiB

+40 -130
View File
@@ -1,152 +1,62 @@
import { app } from "/scripts/app.js"; import { app } from "/scripts/app.js";
import {computeCanvasSize, generatePattern, recursiveLinkUpstream, renameNodeInputs, removeNodeInputs} from "./utils.js"; import {addCustomCanvasWidget, generatePattern, recursiveLinkUpstream, renameNodeInputs, removeNodeInputs} from "./utils.js";
function drawSquares(ctx, startX, startY, squareSize, pattern) {
const inset = Math.max(1, Math.floor(squareSize / 16));
const cell = squareSize - inset;
function drawSquares(ctx, widgetX, widgetY, squareSize, pattern) {
pattern.forEach((value, index) => { pattern.forEach((value, index) => {
const x = widgetX + index * squareSize; // координата x для квадратика const x = startX + index * squareSize + inset / 2;
const y = startY + inset / 2;
// Устанавливаем цвет заливки и обводки ctx.lineWidth = inset;
ctx.fillStyle = value === 1 ? "#222223" : "#00000000"; ctx.strokeStyle = "#222223";
ctx.strokeStyle = "#222223"; // Цвет обводки для всех квадратиков
ctx.lineWidth = Math.floor(squareSize / 16);
if (value === 1) { if (value === 1) {
// Если значение 1, закрашиваем квадрат ctx.fillStyle = "#222223";
ctx.fillRect(x, widgetY, squareSize, squareSize); ctx.fillRect(x, y, cell, cell);
} }
// Рисуем обводку для всех квадратиков ctx.strokeRect(x, y, cell, cell);
ctx.strokeRect(x, widgetY, squareSize, squareSize);
// Добавляем текст в квадратик
if (squareSize >= 24) { if (squareSize >= 24) {
ctx.font = `bold ${squareSize/3}px Arial`; // Размер шрифта адаптируем под размер квадратика ctx.font = `bold ${squareSize / 3}px Arial`;
ctx.textAlign = "center"; ctx.textAlign = "center";
ctx.textBaseline = "middle"; ctx.textBaseline = "middle";
ctx.text ctx.fillStyle = value === 1 ? "#dbdbdc" : "#222223";
ctx.fillStyle = value === 1 ? "#dbdbdc" : "#222223"; // Цвет текста, чтобы он контрастировал с фоном квадратика ctx.fillText(value.toString(), x + cell / 2, y + cell / 2);
ctx.fillText(value.toString(), x + squareSize/2, widgetY + squareSize/2); // Позиционируем текст по центру квадратика
} }
}); });
} }
function addOffsetCanvas(node, app) { // Builds the 0/1 preview pattern from the node's own widgets, or (for
const widget = { // OffsetCombine) from every upstream offset node feeding into it.
type: "customCanvas", function getOffsetPattern(node) {
name: "Offset-Canvas", let patterns = []
get value() {
return this.canvas.value;
},
set value(x) {
this.canvas.value = x;
},
draw: function (ctx, node, widgetWidth, widgetY) {
if (!node.canvasHeight) {
computeCanvasSize(node, node.size)
}
let patterns = [] if (node.type === "OffsetCombine") {
const connectedNodes = recursiveLinkUpstream(node, node.inputs[0].type, node.type, 0)
if (node.type === "OffsetCombine") { for (let [node_ID] of connectedNodes) {
const inputList = [...Array(node.inputs.length).keys()] const connectedNode = node.graph._nodes_by_id[node_ID]
for (let i of inputList) { if (connectedNode && connectedNode.type !== "OffsetCombine") {
const connectedNodes = recursiveLinkUpstream(node, node.inputs[i].type, 0, i) patterns.push({
if (connectedNodes.length !== 0) { process_every: connectedNode.widgets[0].value,
for (let [node_ID, depth] of connectedNodes) { offset: connectedNode.widgets[1].value + node.widgets[0].value,
const connectedNode = node.graph._nodes_by_id[node_ID] mode: connectedNode.widgets[2].value,
if (connectedNode.type !== "OffsetCombine") { })
const pattern = {
process_every: connectedNode.widgets[0].value,
offset: connectedNode.widgets[1].value + node.widgets[0].value,
mode: connectedNode.widgets[2].value
}
patterns.push(pattern)
}
}
}
}
} else {
const pattern = {
process_every: node.widgets[0].value,
offset: node.widgets[1].value,
mode: node.widgets[2].value}
patterns.push(pattern)
}
const pattern = generatePattern(patterns)
const visible = true
const t = ctx.getTransform();
const margin = 10
const widgetHeight = node.canvasHeight
const width = pattern.length * 32
const height = 32
const scale = Math.min((widgetWidth-margin*2)/width, (widgetHeight-margin*2)/height)
Object.assign(this.canvas.style, {
left: `${t.e}px`,
top: `${t.f + (widgetY*t.d)}px`,
width: `${widgetWidth * t.a}px`,
height: `${widgetHeight * t.d}px`,
position: "absolute",
zIndex: 1,
fontSize: `${t.d * 10.0}px`,
pointerEvents: "none",
});
this.canvas.hidden = !visible;
let backgroundWidth = width * scale
let backgroundHeight = height * scale
let xOffset = margin
if (backgroundWidth < widgetWidth) {
xOffset += (widgetWidth-backgroundWidth)/2 - margin
}
let yOffset = margin
if (backgroundHeight < widgetHeight) {
yOffset += (widgetHeight-backgroundHeight)/2 - margin
}
let widgetX = xOffset
widgetY = widgetY + yOffset
// Вычисляем размер квадратика, основываясь на ширине канваса и количестве элементов в списке
const squareSize = backgroundWidth / pattern.length;
// Рисуем квадратики
drawSquares(ctx, widgetX, widgetY, squareSize, pattern)
},
};
widget.canvas = document.createElement("canvas");
widget.canvas.className = "latent-control-custom-canvas";
widget.parent = node;
document.body.appendChild(widget.canvas);
node.addCustomWidget(widget);
app.canvas.onDrawBackground = function () {
for (let n in app.graph._nodes) {
n = graph._nodes[n];
for (let w in n.widgets) {
let wid = n.widgets[w];
if (Object.hasOwn(wid, "canvas")) {
wid.canvas.style.left = -8000 + "px";
wid.canvas.style.position = "absolute";
}
} }
} }
}; } else {
patterns.push({
node.onResize = function (size) { process_every: node.widgets[0].value,
computeCanvasSize(node, size); offset: node.widgets[1].value,
mode: node.widgets[2].value,
})
} }
return { minWidth: 200, minHeight: 200, widget } if (patterns.length === 0) return []
return generatePattern(patterns)
} }
app.registerExtension({ app.registerExtension({
@@ -157,7 +67,7 @@ app.registerExtension({
nodeType.prototype.onNodeCreated = function () { nodeType.prototype.onNodeCreated = function () {
const r = onNodeCreated ? onNodeCreated.apply(this, arguments) : undefined; const r = onNodeCreated ? onNodeCreated.apply(this, arguments) : undefined;
addOffsetCanvas(this, app) addCustomCanvasWidget(this, "Offset-Canvas", getOffsetPattern, drawSquares)
return r; return r;
} }
@@ -173,7 +83,7 @@ app.registerExtension({
nodeType.prototype.onNodeCreated = function () { nodeType.prototype.onNodeCreated = function () {
const r = onNodeCreated ? onNodeCreated.apply(this, arguments) : undefined; const r = onNodeCreated ? onNodeCreated.apply(this, arguments) : undefined;
addOffsetCanvas(this, app) addCustomCanvasWidget(this, "Offset-Canvas", getOffsetPattern, drawSquares)
this.getExtraMenuOptions = function(_, options) { this.getExtraMenuOptions = function(_, options) {
options.unshift( options.unshift(
+99 -65
View File
@@ -1,52 +1,3 @@
export function computeCanvasSize(node, size) {
if (node.widgets[0].last_y == null) return;
const MIN_SIZE = 64;
const inputs = node.inputs === undefined ? 0 : node.inputs.length
const outputs = node.outputs === undefined ? 0 : node.outputs.length
let y = LiteGraph.NODE_WIDGET_HEIGHT * Math.max(inputs, outputs) + 5;
let freeSpace = size[1] - y;
// Compute the height of all non customtext widgets
let widgetHeight = 0;
for (let i = 0; i < node.widgets.length; i++) {
const w = node.widgets[i];
if (w.type !== "customCanvas") {
if (w.computeSize) {
widgetHeight += w.computeSize()[1] + 4;
} else {
widgetHeight += LiteGraph.NODE_WIDGET_HEIGHT + 5;
}
}
}
// See how large the canvas can be
freeSpace -= widgetHeight;
// There isnt enough space for all the widgets, increase the size of the node
if (freeSpace < MIN_SIZE) {
freeSpace = MIN_SIZE;
node.size[1] = y + widgetHeight + freeSpace;
node.graph.setDirtyCanvas(true);
}
// Position each of the widgets
for (const w of node.widgets) {
w.y = y;
if (w.type === "customCanvas") {
y += freeSpace;
} else if (w.computeSize) {
y += w.computeSize()[1] + 4;
} else {
y += LiteGraph.NODE_WIDGET_HEIGHT + 4;
}
}
node.canvasHeight = freeSpace;
}
function gcd(a, b) { function gcd(a, b) {
// Функция для вычисления наибольшего общего делителя (НОД) // Функция для вычисления наибольшего общего делителя (НОД)
while (b !== 0) { while (b !== 0) {
@@ -83,26 +34,24 @@ export function generatePattern(rules) {
return pattern; return pattern;
} }
export function recursiveLinkUpstream(node, type, depth, index=null) { export function recursiveLinkUpstream(node, slot_type, node_type, depth) {
depth += 1 depth += 1
let connections = [] let connections = []
if (node.type === "OffsetCombine") { const inputList = [...Array(node.inputs.length).keys()]
const inputList = [...Array(node.inputs.length).keys()] for (let i of inputList) {
for (let i of inputList) { const link = node.inputs[i].link
const link = node.inputs[i].link if (link) {
if (link) { const nodeID = node.graph.links[link].origin_id
const nodeID = node.graph.links[link].origin_id const slotID = node.graph.links[link].origin_slot
const slotID = node.graph.links[link].origin_slot const connectedNode = node.graph._nodes_by_id[nodeID]
const connectedNode = node.graph._nodes_by_id[nodeID]
if (connectedNode.outputs[slotID].type === type) { if (connectedNode.outputs[slotID].type === slot_type) {
connections.push([connectedNode.id, depth]) connections.push([connectedNode.id, depth])
if (connectedNode.inputs) { if (connectedNode.inputs) {
const index = (connectedNode.type === "OffsetCombine") ? 0 : null const index = (connectedNode.type === node_type) ? 0 : null
connections = connections.concat(recursiveLinkUpstream(connectedNode, type, depth, index)) connections = connections.concat(recursiveLinkUpstream(connectedNode, slot_type, node_type, depth))
}
} }
} }
} }
@@ -125,5 +74,90 @@ export function removeNodeInputs(node, indexesToRemove) {
node.removeInput(i) node.removeInput(i)
} }
node.onResize(node.size) node.setSize(node.computeSize())
node.graph?.setDirtyCanvas(true, true)
} }
const MARGIN_X = 15; // side padding — matches the node's standard widgets
const MARGIN_Y = 10; // top/bottom padding around the grid (node-local px)
const MIN_SQUARE = 24; // smallest square that still shows the row on tiny/long patterns
const MAX_SQUARE = 72; // largest square, so a short pattern doesn't blow up
const MIN_NODE_WIDTH = 220; // enough width for the squares to read
// Square side that fits the given (full node) width and column count. Squares
// shrink to fit a long pattern (down to 1px) and are capped so a short one stays
// reasonable. Guarantees squareSize * cols <= width - 2*MARGIN_X, i.e. no overflow.
function fitSquareSize(width, cols) {
const availWidth = width - MARGIN_X * 2;
return Math.max(1, Math.min(availWidth / Math.max(cols, 1), MAX_SQUARE));
}
/**
* Adds a self-contained canvas widget that draws a pattern preview directly onto
* the LiteGraph node context. Works with the current ComfyUI frontend (Nodes v2):
* it paints in node-local coordinates, so it scales and positions correctly with
* zoom and DPI — no DOM overlay and no manual ctx.getTransform() math.
*
* The preview auto-sizes to the space available inside the node: computeSize()
* reports exactly the height the grid needs for the current width and pattern
* length, so the node reserves the right amount of room, the squares grow/shrink
* as the node gets wider/narrower or the pattern longer/shorter, and the drawing
* always stays inside the node body.
*
* @param node the LiteGraph node
* @param name widget name
* @param getPattern (node) => number[] the 0/1 pattern to render
* @param drawSquares (ctx, startX, startY, squareSize, pattern) => void
*/
export function addCustomCanvasWidget(node, name, getPattern, drawSquares) {
const widget = {
type: "customCanvas",
name,
value: undefined,
options: { serialize: false },
computeSize() {
// Use the real node width, not the passed value: LiteGraph feeds a
// widened width to selected nodes, which would inflate the reserve.
const nodeWidth = node.size[0];
const cols = getPattern(node).length || 1;
const rowHeight = Math.max(fitSquareSize(nodeWidth, cols), MIN_SQUARE);
return [nodeWidth, rowHeight + MARGIN_Y * 2];
},
draw(ctx, node, width, widgetY) {
const pattern = getPattern(node);
if (!pattern || !pattern.length) return;
// Use node.size[0] rather than the `width` argument: on a selected node
// LiteGraph passes an enlarged width, which would push the grid right and
// peg the square size at its cap. node.size[0] is stable in both states.
const nodeWidth = node.size[0];
const cols = pattern.length;
const squareSize = fitSquareSize(nodeWidth, cols);
// Same band height computeSize reserved -> the grid stays inside the node.
const bandHeight = Math.max(squareSize, MIN_SQUARE) + MARGIN_Y * 2;
const gridWidth = squareSize * cols;
const startX = (nodeWidth - gridWidth) / 2;
const startY = widgetY + (bandHeight - squareSize) / 2;
// Light backing drawn first so empty cells stay readable on dark nodes.
ctx.fillStyle = "#ffffffcc";
ctx.fillRect(startX, startY, gridWidth, squareSize);
drawSquares(ctx, startX, startY, squareSize, pattern);
},
};
node.addCustomWidget(widget);
// Widen a touch and reserve the preview height on creation. On workflow load
// LiteGraph applies the saved size after onNodeCreated, so this does not fight
// persisted sizes.
node.size[0] = Math.max(node.size[0], MIN_NODE_WIDTH);
node.setSize(node.computeSize());
return widget;
}
-65
View File
@@ -1,65 +0,0 @@
import torch
import nodes
import comfy
from latent_preview import prepare_callback as preview_callback
class TransformContext:
original_sample_function = nodes.common_ksampler
def get_transform_sample_function(self):
def prepare_callback(model, steps, x0_output_dict=None, transforms=None):
def transform_callback(step, x0, x, total_steps):
if transforms is None:
return
for transform in transforms:
for i in range(x0.size()[0]):
x0[i] = transform["function"](step, x0[i].unsqueeze(0), total_steps, transform["params"])
preview = preview_callback(model, steps, x0_output_dict)
def callback(step, x0, x, total_steps):
transform_callback(step, x0, x, total_steps)
preview(step, x0, x, total_steps)
return callback
def sample(model, seed, steps, cfg, sampler_name, scheduler, positive, negative, latent, denoise=1.0,
disable_noise=False, start_step=None, last_step=None, force_full_denoise=False):
latent_image = latent["samples"]
if disable_noise:
noise = torch.zeros(latent_image.size(), dtype=latent_image.dtype, layout=latent_image.layout, device="cpu")
else:
batch_inds = latent["batch_index"] if "batch_index" in latent else None
noise = comfy.sample.prepare_noise(latent_image, seed, batch_inds)
noise_mask = None
if "noise_mask" in latent:
noise_mask = latent["noise_mask"]
callback = prepare_callback(model, steps, transforms=latent["transforms"])
disable_pbar = not comfy.utils.PROGRESS_BAR_ENABLED
samples = comfy.sample.sample(model, noise, steps, cfg, sampler_name, scheduler, positive, negative, latent_image,
denoise=denoise, disable_noise=disable_noise, start_step=start_step,
last_step=last_step,
force_full_denoise=force_full_denoise, noise_mask=noise_mask, callback=callback,
disable_pbar=disable_pbar, seed=seed)
out = latent.copy()
out["samples"] = samples
self.unhijack()
return (out,)
return sample
def hijack(self):
nodes.common_ksampler = self.get_transform_sample_function()
def unhijack(self):
nodes.common_ksampler = TransformContext.original_sample_function
def __enter__(self):
self.hijack()
def __exit__(self, exc_type, exc_value, exc_traceback):
self.unhijack()
+7 -19
View File
@@ -1,32 +1,20 @@
from .TransformContext import TransformContext from .transform_apply import attach_transforms
class TransformHijack: class TransformHijack:
@classmethod @classmethod
def INPUT_TYPES(cls): def INPUT_TYPES(cls):
return { return {
"required" : { "required": {
"latent": ("LATENT",), "model": ("MODEL",),
"transforms": ("TRANSFORM",) "transforms": ("TRANSFORM",),
}, },
} }
RETURN_TYPES = ("LATENT",) RETURN_TYPES = ("MODEL",)
FUNCTION = "func" FUNCTION = "func"
CATEGORY = "sampling/transforms" CATEGORY = "sampling/transforms"
_context = None def func(self, model, transforms):
_hijack_node_id = None return (attach_transforms(model, transforms),)
def func(self, latent, transforms):
latent["transforms"] = transforms
if TransformHijack._context is None:
TransformHijack._hijack_node_id = id
TransformHijack._context = TransformContext()
else:
return (latent,)
TransformHijack._context.hijack()
return (latent,)
+6 -12
View File
@@ -1,4 +1,4 @@
from .TransformContext import TransformContext from .transform_apply import attach_transforms
from nodes import KSampler, KSamplerAdvanced from nodes import KSampler, KSamplerAdvanced
@@ -16,17 +16,11 @@ class Transforms:
FUNCTION = "func" FUNCTION = "func"
def __init__(self):
self.original_function_name = self.clazz.FUNCTION
def func(self, **kwargs): def func(self, **kwargs):
ctx = TransformContext() transforms = kwargs.pop("transform_optional", None)
ctx.hijack() if transforms:
latent = kwargs["latent_image"] kwargs["model"] = attach_transforms(kwargs["model"], transforms)
latent["transforms"] = kwargs.pop("transform_optional") return getattr(self, self.clazz.FUNCTION)(**kwargs)
kwargs["latent_image"] = latent
out = getattr(self, self.clazz.FUNCTION)(**kwargs)
return out
def variations_factory(original_class: type, name=None) -> type: def variations_factory(original_class: type, name=None) -> type:
@@ -34,4 +28,4 @@ def variations_factory(original_class: type, name=None) -> type:
return type(name, (Transforms, original_class), {'clazz': original_class}) return type(name, (Transforms, original_class), {'clazz': original_class})
TSampler = variations_factory(KSampler) TSampler = variations_factory(KSampler)
TSamplerAdvanced = variations_factory(KSamplerAdvanced) TSamplerAdvanced = variations_factory(KSamplerAdvanced)
+1 -5
View File
@@ -62,9 +62,7 @@ def latent_interpolate_transform(x0, params):
latent = params["latent"].to(x0.device) latent = params["latent"].to(x0.device)
if x0.shape != latent.shape: if x0.shape != latent.shape:
latent.permute(0, 3, 1, 2)
latent = comfy.utils.common_upscale(latent, x0.shape[3], x0.shape[2], 'bicubic', crop='center') latent = comfy.utils.common_upscale(latent, x0.shape[3], x0.shape[2], 'bicubic', crop='center')
latent.permute(0, 2, 3, 1)
x = latent * params["factor"] + x0 * (1 - params["factor"]) x = latent * params["factor"] + x0 * (1 - params["factor"])
x *= params["multiplier"] x *= params["multiplier"]
@@ -76,9 +74,7 @@ def latent_add_transform(x0, params):
latent = params["latent"].to(x0.device) latent = params["latent"].to(x0.device)
if x0.shape != latent.shape: if x0.shape != latent.shape:
latent.permute(0, 3, 1, 2) latent = comfy.utils.common_upscale(latent, x0.shape[3], x0.shape[2], 'bicubic', crop='center')
latent = comfy.utils.common_upscale(latent, x0.shape[3], x0.shape[2], 'bicubic', crop='center')
latent.permute(0, 2, 3, 1)
x = x0 + latent x = x0 + latent
x *= params["multiplier"] x *= params["multiplier"]
+51
View File
@@ -0,0 +1,51 @@
import torch
def apply_transforms_to_x0(x0, step, total_steps, transforms):
x = x0.clone()
for transform in transforms:
for i in range(x.size()[0]):
x[i] = transform["function"](step, x[i].unsqueeze(0), total_steps, transform["params"])
return x
def _find_step(sigma, sigmas):
# sigma is a scalar tensor for the current model evaluation, sigmas is the full
# schedule. High order samplers evaluate the model at intermediate sigmas that are
# not part of the schedule; for those we return None so the transform is applied
# exactly once per step (parity with the old per-step callback).
diff = torch.abs(sigmas - sigma.to(sigmas.device))
idx = int(torch.argmin(diff).item())
if diff[idx] <= 1e-4 * max(1.0, float(sigmas[idx].abs())):
return idx
return None
def make_post_cfg_function(transforms):
def post_cfg_function(args):
denoised = args["denoised"]
if not transforms:
return denoised
sigmas = args["model_options"].get("transformer_options", {}).get("sample_sigmas", None)
if sigmas is None:
return denoised
step = _find_step(args["sigma"], sigmas)
if step is None:
return denoised
total_steps = len(sigmas) - 1
return apply_transforms_to_x0(denoised, step, total_steps, transforms)
return post_cfg_function
def attach_transforms(model, transforms):
m = model.clone()
m.set_model_sampler_post_cfg_function(make_post_cfg_function(transforms))
return m
+9 -1
View File
@@ -1,4 +1,5 @@
import torch import torch
from nodes import PreviewImage
MIRROR_DIRECTIONS = ["vertically", "horizontally", "both"] MIRROR_DIRECTIONS = ["vertically", "horizontally", "both"]
@@ -15,6 +16,9 @@ class LatentMirror:
"max": 10.0, "max": 10.0,
"step": 0.01 "step": 0.01
}) })
},
"optional": {
"vae_optional": ("VAE",)
} }
} }
@@ -23,7 +27,7 @@ class LatentMirror:
CATEGORY = "latent/advanced" CATEGORY = "latent/advanced"
def mirror(self, latent, direction, multiplier): def mirror(self, latent, direction, multiplier, vae_optional = None):
l = latent.copy() l = latent.copy()
if direction == "vertically" or direction == "both": if direction == "vertically" or direction == "both":
l["samples"] = torch.flip(l["samples"], dims=[2]) + l["samples"] l["samples"] = torch.flip(l["samples"], dims=[2]) + l["samples"]
@@ -31,4 +35,8 @@ class LatentMirror:
l["samples"] = torch.flip(l["samples"], dims=[3]) + l["samples"] l["samples"] = torch.flip(l["samples"], dims=[3]) + l["samples"]
l["samples"] *= multiplier l["samples"] *= multiplier
if vae_optional:
return {"result": (l,), "ui": PreviewImage().save_images(vae_optional.decode(l["samples"]))["ui"]}
return (l,) return (l,)
+9 -1
View File
@@ -1,4 +1,5 @@
import torch import torch
from nodes import PreviewImage
class LatentShift: class LatentShift:
@@ -19,6 +20,9 @@ class LatentShift:
"max": 1, "max": 1,
"step": 0.01 "step": 0.01
}), }),
},
"optional": {
"vae_optional": ("VAE",)
} }
} }
@@ -27,7 +31,7 @@ class LatentShift:
CATEGORY = "latent/advanced" CATEGORY = "latent/advanced"
def shift(self, latent, x_shift, y_shift): def shift(self, latent, x_shift, y_shift, vae_optional = None):
l = latent.copy() l = latent.copy()
if x_shift != 0: if x_shift != 0:
@@ -35,4 +39,8 @@ class LatentShift:
if y_shift != 0: if y_shift != 0:
l["samples"] = torch.roll(l["samples"], shifts=int(l["samples"].size()[2] * y_shift), dims=[2]) l["samples"] = torch.roll(l["samples"], shifts=int(l["samples"].size()[2] * y_shift), dims=[2])
if vae_optional:
return {"result": (l,), "ui": PreviewImage().save_images(vae_optional.decode(l["samples"]))["ui"]}
return (l,) return (l,)
+14
View File
@@ -0,0 +1,14 @@
[project]
name = "comfyui-advanced-latent-control"
description = "This custom node helps to transform latent in different ways."
version = "3.1.0"
license = "LICENSE"
[project.urls]
Repository = "https://github.com/RomanKuschanow/ComfyUI-Advanced-Latent-Control"
# Used by Comfy Registry https://comfyregistry.org
[tool.comfy]
PublisherId = "kuschanow"
DisplayName = "ComfyUI-Advanced-Latent-Control"
Icon = ""