15 Commits
Author SHA1 Message Date
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
RomanKuschanow a89d51dd05 readme 2024-04-02 00:04:58 +03:00
RomanKuschanow 0b3641c0e6 transform hijack 2024-04-01 23:51:10 +03:00
RomanKuschanow 26c36e9c91 adding more than 2 offsets to offset combine node 2024-03-31 17:45:19 +03:00
RomanKuschanow d74687b48c refactoring 2024-03-31 15:06:15 +03:00
RomanKuschanow 5f46f8f623 remove test node 2024-03-31 15:05:28 +03:00
RomanKuschanow 872f5b371c offset widget 2024-03-31 15:04:33 +03:00
RomanKuschanow 6b3770acea test widget 2024-03-18 22:14:01 +02:00
RomanKuschanow 17c23c1203 invert factor parameter in latent interpolate node 2024-03-18 18:12:37 +02:00
RomanKuschanow 6585f9c014 Merge branch 'master' into dev 2024-03-16 14:17:40 +02:00
RomanKuschanow d3f14bb9d4 latent add and interpolate bug fix 2024-03-16 14:09:54 +02:00
RomanKuschanow 83f7a56317 offset_optional bug fix 2024-03-16 13:16:11 +02:00
RomanKuschanow 6e25170ee8 README fix 2024-03-13 21:40:09 +02:00
RomanKuschanow 1efbd7d31d latent normalize node 2024-03-11 21:14:49 +02:00
19 changed files with 628 additions and 247 deletions
+60 -43
View File
@@ -1,10 +1,12 @@
# ComfyUI-Advanced-Latent-Control
**This custom node helps to transform latent in different ways.**
**This custom nodes helps to transform latent in different ways.**
## Custom Nodes
### Latent mirror
This node can flip latent and merge original and flipped version
>You can access new features earlier by switching from the master branch to dev,
but you need to remember that there may be some issues on the dev branch and some nodes' behavior may change after release.
## Latent mirror
This node can flip latent and merge original and flipped version.
**Input:**
- `latent`
@@ -20,8 +22,8 @@ This node can flip latent and merge original and flipped version
![sample](https://i.imgur.com/YMyYorQ.png)
![sample](https://i.imgur.com/W5BasCO.png)
### Latent shift
This node can shift latent along x and y axes
## Latent shift
This node can shift latent along x and y-axis.
**Input:**
- `latent`
@@ -36,40 +38,11 @@ This node can shift latent along x and y axes
**Usage:**
![sample](https://i.imgur.com/1Dp5dSw.png)
### KSampler with transforms (Latent Control)
This node can multiply, mirror and shift latent during generation
## ~~TSampler with transforms (Latent Control)~~
Removed from version 2.0.0
**Input:**
exactly matches the base KSampler
**Fields:**
- base KSampler fields
- `start_mirror_at` – a number between 0 and 1 that indicates at what point the sampler will start mirroring
- `stop_mirror_at` – a number between 0 and 1 that indicates at what point the sampler will stop mirroring
- `mirror_mode` – can be `replace` or `combine`. `replace` will replace the latent with the transformed one, `combine` will add the original and the transformed latent and divide by 2
- `mirror_direction` – can be `none`, `vertically`, `horizontally`, `both`, `90 degree rotation` or `180 degree rotation`
- `start_shift_at` – a number between 0 and 1 that indicates at what point the sampler will start shifting
- `stop_shift_at` – a number between 0 and 1 that indicates at what point the sampler will stop shifting
- `shift_mode` – can be `replace` or `combine`. `replace` will replace the latent with the transformed one, `combine` will add the original and the transformed latent and divide by 2
- `x_shift` – a number between -1 and 1 that indicates how much the latent should be shifted
- `y_shift` – a number between -1 and 1 that indicates how much the latent should be shifted
- `start_multiplier_at` – a number between 0 and 1 that indicates at what point the sampler will start multiplying
- `stop_multiplier_at` – a number between 0 and 1 that indicates at what point the sampler will stop multiplying
- `multiplier_mode` – can be `replace` or `combine`. `replace` will replace the latent with the transformed one, `combine` will add the original and the transformed latent
- `multiplier` – multiply latent by specified number
**Output:**
exactly matches the base KSampler
**Usage:**
**You also can use those params together**
![sample](https://i.imgur.com/RMJTnWF.png)
![sample](https://i.imgur.com/fQ7UWuS.png)
![sample](https://i.imgur.com/pxWupAx.png)
![sample](https://i.imgur.com/1YkERDu.png)
### KSampler (Latent Control)
This node allows to combine a lot of transforms with different parameters
## TSampler (Latent Control)
This node allows to combine a lot of transforms with different parameters.
**Input:**
- base KSampler fields
@@ -85,7 +58,7 @@ exactly matches the base KSampler
![sample](https://i.imgur.com/PlGnAtA.png)
![sample](https://i.imgur.com/CtrBRPn.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.
There are two new transform nodes:
- Latent add
@@ -93,7 +66,7 @@ There are two new transform nodes:
They work exactly the same as LatentAdd and LatentBlend nodes from standard node pack, but also, can multiply result by specified number.
### Offset
## Offset
You can apply specific offset for transform nodes.
**Fields:**
@@ -110,11 +83,55 @@ You can apply specific offset for transform nodes.
![sample](https://i.imgur.com/MGVLfve.png)
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**.
### One time nodes
## One time nodes
Each transform node has own one-time version. They allow to make one transform action at specified step.
**Usage:**
![sample](https://i.imgur.com/Q1Vyob0.png)
## Latent normalize
Fixes some issues when sampling modified latent space.
**Input:**
exactly matches the `VAE Decode` node
**Output:**
- latent
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).
![sample](https://i.imgur.com/3FXk8n7.png)
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)
This node also change some results even if output without this node looks good.
![sample](https://i.imgur.com/kP0f6vh.png)
![sample](https://i.imgur.com/YI8ZqLd.png)
And it very slightly changes results from latent, which have not been modified.
![sample](https://i.imgur.com/xTU08xm.png)
![sample](https://i.imgur.com/yzgW7QT.png)
## Transform hijack
Allow you to use transforms with any samplers that you like.
**Inputs:**
- latent
- transforms
**Outputs:**
- latent
**Usage:**
![sample](https://i.imgur.com/YwVhHYF.png)
+10 -3
View File
@@ -1,10 +1,15 @@
from .nodes import *
WEB_DIRECTORY = "js"
NODE_CLASS_MAPPINGS = {
"LatentMirror": LatentMirror,
"LatentShift": LatentShift,
"TSamplerWithTransform": TSamplerWithTransform,
"TransformSampler": TransformSampler,
"LatentNormalize": LatentNormalize,
"TransformSampler": TSampler,
"TransformSamplerAdvanced": TSamplerAdvanced,
"TransformHijack": TransformHijack,
"MirrorTransform": MirrorTransform,
"ShiftTransform": ShiftTransform,
"MultiplyTransform": MultiplyTransform,
@@ -23,8 +28,10 @@ NODE_CLASS_MAPPINGS = {
NODE_DISPLAY_NAME_MAPPINGS = {
"LatentMirror": "Latent mirror",
"LatentShift": "Latent shift",
"TSamplerWithTransform": "TSampler with transforms (Latent Control)",
"LatentNormalize": "Latent normalize",
"TransformSampler": "TSampler (Latent Control)",
"TransformSamplerAdvanced": "TSampler Advanced (Latent Control)",
"TransformHijack": "Transform Hijack",
"MirrorTransform": "Mirror transform",
"ShiftTransform": "Shift transform",
"MultiplyTransform": "Multiply transform",
+191
View File
@@ -0,0 +1,191 @@
import { app } from "/scripts/app.js";
import {addCanvas, computeCanvasSize, generatePattern, recursiveLinkUpstream, renameNodeInputs, removeNodeInputs} from "./utils.js";
function drawSquares(ctx, widgetX, widgetY, squareSize, pattern) {
const actualSquareSize = squareSize - Math.floor(squareSize / 16);
widgetY += Math.floor(squareSize / 16) / 2;
widgetX += Math.floor(squareSize / 16) / 2;
pattern.forEach((value, index) => {
const x = widgetX + index * squareSize;
ctx.fillStyle = value === 1 ? "#222223" : "#00000000";
ctx.strokeStyle = "#222223";
ctx.lineWidth = Math.floor(squareSize / 16);
if (value === 1) {
ctx.fillRect(x, widgetY, actualSquareSize, actualSquareSize);
}
ctx.strokeRect(x, widgetY, actualSquareSize, actualSquareSize);
if (squareSize >= 24) {
ctx.font = `bold ${squareSize/3}px Arial`;
ctx.textAlign = "center";
ctx.textBaseline = "middle";
ctx.text
ctx.fillStyle = value === 1 ? "#dbdbdc" : "#222223";
ctx.fillText(value.toString(), x + actualSquareSize/2, widgetY + actualSquareSize/2);
}
});
}
const offsetWidget = {
type: "customCanvas",
name: "Offset-Canvas",
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 (connectedNodes.length !== 0) {
for (let [node_ID, depth] of connectedNodes) {
const connectedNode = node.graph._nodes_by_id[node_ID]
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)
ctx.fillStyle = "#ffffff88"
ctx.fillRect(widgetX, widgetY, backgroundWidth, backgroundHeight);
},
};
app.registerExtension({
name: "Comfy.LatentControl.TransformOffset",
async beforeRegisterNodeDef (nodeType, nodeData, app){
if (nodeData.name === "TransformOffset") {
const onNodeCreated = nodeType.prototype.onNodeCreated;
nodeType.prototype.onNodeCreated = function () {
const r = onNodeCreated ? onNodeCreated.apply(this, arguments) : undefined;
addCanvas(this, app, offsetWidget)
return r;
}
}
},
});
app.registerExtension({
name: "Comfy.LatentControl.OffsetCombine",
async beforeRegisterNodeDef (nodeType, nodeData, app){
if (nodeData.name === "OffsetCombine") {
const onNodeCreated = nodeType.prototype.onNodeCreated;
nodeType.prototype.onNodeCreated = function () {
const r = onNodeCreated ? onNodeCreated.apply(this, arguments) : undefined;
addCanvas(this, app, offsetWidget)
this.getExtraMenuOptions = function(_, options) {
options.unshift(
{
content: `add offset`,
callback: () => {
this.addInput("offset", "OFFSET")
renameNodeInputs(this, "offset")
this.setDirtyCanvas(true);
},
},
{
content: `remove offset`,
callback: () => {
removeNodeInputs(this, [this.inputs.length-1])
renameNodeInputs(this, "offset")
},
},
{
content: "remove all unconnected offsets",
callback: () => {
let indexesToRemove = []
for (let i = 0; i < this.inputs.length; i++) {
if (!this.inputs[i].link) {
indexesToRemove.push(i)
}
}
if (indexesToRemove.length) {
removeNodeInputs(this, indexesToRemove)
}
renameNodeInputs(this, "offset")
},
},
);
}
return r;
}
}
},
});
+156
View File
@@ -0,0 +1,156 @@
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) {
// Функция для вычисления наибольшего общего делителя (НОД)
while (b !== 0) {
let t = b;
b = a % b;
a = t;
}
return a;
}
function lcm(a, b) {
// Функция для вычисления наименьшего общего кратного (НОК)
return (a * b) / gcd(a, b);
}
function findPatternLength(rules) {
// Вычисление длины цикла как НОК всех process_every
return rules.map(rule => rule.process_every).reduce((acc, val) => lcm(acc, val), 1);
}
export function generatePattern(rules) {
let length = findPatternLength(rules); // Определение длины паттерна
let pattern = new Array(length).fill(0);
rules.forEach(rule => {
let offset = rule.offset % rule.process_every;
for (let i = 0; i < length; i++) {
let value = ((i + offset) % rule.process_every === 0) === (rule.mode === "process_every") ? 1 : 0;
pattern[i] = pattern[i] || value;
}
});
return pattern;
}
export function recursiveLinkUpstream(node, slot_type, node_type, depth) {
depth += 1
let connections = []
const inputList = [...Array(node.inputs.length).keys()]
for (let i of inputList) {
const link = node.inputs[i].link
if (link) {
const nodeID = node.graph.links[link].origin_id
const slotID = node.graph.links[link].origin_slot
const connectedNode = node.graph._nodes_by_id[nodeID]
if (connectedNode.outputs[slotID].type === slot_type) {
connections.push([connectedNode.id, depth])
if (connectedNode.inputs) {
const index = (connectedNode.type === node_type) ? 0 : null
connections = connections.concat(recursiveLinkUpstream(connectedNode, slot_type, node_type, depth))
}
}
}
}
return connections
}
export function renameNodeInputs(node, name) {
for (let i=0; i < node.inputs.length; i++) {
node.inputs[i].name = `${name}${i + 1}`
}
}
export function removeNodeInputs(node, indexesToRemove) {
indexesToRemove.sort((a, b) => b - a);
for (let i of indexesToRemove) {
if (node.inputs.length <= 2) { console.log("too short"); continue } // if only 2 left
node.removeInput(i)
}
node.onResize(node.size)
}
export function addCanvas(node, app, widget) {
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";
}
}
}
};
node.onResize = function (size) {
computeCanvasSize(node, size);
}
}
@@ -1,88 +0,0 @@
import comfy.samplers
from .TransformSampler import TransformSampler
from .Transforms import MirrorTransform, ShiftTransform, MultiplyTransform
MIRROR_DIRECTIONS = ["none", "vertically", "horizontally", "both", "90 degree rotation", "180 degree rotation"]
MODE = ["replace", "combine"]
class TSamplerWithTransform:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"model": ("MODEL",),
"seed": ("INT", {"default": 0, "min": 0, "max": 0xffffffffffffffff}),
"steps": ("INT", {"default": 20, "min": 1, "max": 10000}),
"cfg": ("FLOAT", {"default": 8.0, "min": 0.0, "max": 100.0, "step":0.1, "round": 0.01}),
"sampler_name": (comfy.samplers.KSampler.SAMPLERS, ),
"scheduler": (comfy.samplers.KSampler.SCHEDULERS, ),
"positive": ("CONDITIONING", ),
"negative": ("CONDITIONING", ),
"latent_image": ("LATENT", ),
"denoise": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01}),
"start_mirror_at": ("FLOAT", {"default": 0, "min": 0.0, "max": 1.0, "step": 0.01}),
"stop_mirror_at": ("FLOAT", {"default": 0, "min": 0.0, "max": 1.0, "step": 0.01}),
"mirror_mode": (MODE,),
"mirror_direction": (MIRROR_DIRECTIONS, {"default": "none"}),
"start_shift_at": ("FLOAT", {"default": 0, "min": 0.0, "max": 1.0, "step": 0.01}),
"stop_shift_at": ("FLOAT", {"default": 0, "min": 0.0, "max": 1.0, "step": 0.01}),
"shift_mode": (MODE, {"default": "replace"}),
"x_shift": ("FLOAT", {"default": 0, "min": -1, "max": 1, "step": 0.01}),
"y_shift": ("FLOAT", {"default": 0, "min": -1, "max": 1, "step": 0.01}),
"start_multiplier_at": ("FLOAT", {"default": 0, "min": 0.0, "max": 1.0, "step": 0.01}),
"stop_multiplier_at": ("FLOAT", {"default": 0, "min": 0.0, "max": 1.0, "step": 0.01}),
"multiplier_mode": (MODE, {"default": "combine"}),
"multiplier": ("FLOAT", {"default": 1, "min": -10, "max": 10, "step": 0.01}),
}
}
RETURN_TYPES = ("LATENT",)
FUNCTION = "sample"
CATEGORY = "sampling"
def sample(self,
model,
seed,
steps,
cfg,
sampler_name,
scheduler,
positive,
negative,
latent_image,
denoise=1.0,
start_mirror_at=0,
stop_mirror_at=0,
mirror_mode="replace",
mirror_direction="none",
start_shift_at=0,
stop_shift_at=0,
shift_mode="replace",
x_shift=0,
y_shift=0,
start_multiplier_at=0,
stop_multiplier_at=0,
multiplier_mode="combine",
multiplier=1):
transforms = (
MirrorTransform().process(start_mirror_at, stop_mirror_at, mirror_mode, mirror_direction) +
ShiftTransform().process(start_shift_at, stop_shift_at, shift_mode, x_shift, y_shift) +
MultiplyTransform().process(start_multiplier_at, stop_multiplier_at, multiplier_mode, multiplier))[0]
return TransformSampler().sample(
model,
seed,
steps,
cfg,
sampler_name,
scheduler,
positive,
negative,
latent_image,
transform_optional=transforms,
denoise=denoise)
+65
View File
@@ -0,0 +1,65 @@
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()
+32
View File
@@ -0,0 +1,32 @@
from .TransformContext import TransformContext
class TransformHijack:
@classmethod
def INPUT_TYPES(cls):
return {
"required" : {
"latent": ("LATENT",),
"transforms": ("TRANSFORM",)
},
}
RETURN_TYPES = ("LATENT",)
FUNCTION = "func"
CATEGORY = "sampling/transforms"
_context = None
_hijack_node_id = None
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,)
+26 -75
View File
@@ -1,86 +1,37 @@
import torch
import comfy.samplers
from latent_preview import prepare_callback as preview_callback
from .TransformContext import TransformContext
from nodes import KSampler, KSamplerAdvanced
def prepare_callback(model, steps, transforms, x0_output_dict=None):
def transform_callback(step, x0, x, total_steps):
for transform in transforms:
for i in range(x0.size()[0]):
x0[i] = transform["function"](step, x0[i], 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)
def insert_transform_input(input_types):
input_types["optional"] = {"transform_optional": ("TRANSFORM",)}
return input_types
return callback
class Transforms:
clazz = None
def sample_common(model, seed, steps, cfg, sampler_name, scheduler, positive, negative, latent, transform, 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, transform)
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
return (out,)
class TransformSampler:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"model": ("MODEL",),
"seed": ("INT", {"default": 0, "min": 0, "max": 0xffffffffffffffff}),
"steps": ("INT", {"default": 20, "min": 1, "max": 10000}),
"cfg": ("FLOAT", {"default": 8.0, "min": 0.0, "max": 100.0, "step":0.1, "round": 0.01}),
"sampler_name": (comfy.samplers.KSampler.SAMPLERS, ),
"scheduler": (comfy.samplers.KSampler.SCHEDULERS, ),
"positive": ("CONDITIONING", ),
"negative": ("CONDITIONING", ),
"latent_image": ("LATENT", ),
"denoise": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01}),
},
"optional":{
"transform_optional": ("TRANSFORM",),
}
}
def INPUT_TYPES(cls):
return insert_transform_input(cls.clazz.INPUT_TYPES())
RETURN_TYPES = ("LATENT",)
FUNCTION = "sample"
FUNCTION = "func"
CATEGORY = "sampling"
def __init__(self):
self.original_function_name = self.clazz.FUNCTION
def sample(self, model, seed, steps, cfg, sampler_name, scheduler, positive, negative, latent_image, denoise=1.0, transform_optional=None):
def func(self, **kwargs):
ctx = TransformContext()
ctx.hijack()
latent = kwargs["latent_image"]
latent["transforms"] = kwargs.pop("transform_optional")
kwargs["latent_image"] = latent
out = getattr(self, self.clazz.FUNCTION)(**kwargs)
return out
if transform_optional is None:
transform_optional = []
return sample_common(
model,
seed,
steps,
cfg,
sampler_name,
scheduler,
positive,
negative,
latent_image,
transform=transform_optional,
denoise=denoise)
def variations_factory(original_class: type, name=None) -> type:
name = name or original_class.__name__ + "Transform"
return type(name, (Transforms, original_class), {'clazz': original_class})
TSampler = variations_factory(KSampler)
TSamplerAdvanced = variations_factory(KSamplerAdvanced)
@@ -1,4 +1,6 @@
from .utils import latent_add_transform, get_offset_list
import comfy
import torch
class LatentAddTransform:
@@ -22,14 +24,14 @@ class LatentAddTransform:
CATEGORY = "sampling/transforms"
def process(self,
offset_optional,
latent,
start_at=0,
stop_at=0,
multiplier=1):
multiplier=1,
offset_optional=None):
return ([{
"params": {
"latent": latent["samples"][0],
"latent": latent["samples"][0].unsqueeze(0),
"start_at": start_at,
"stop_at": stop_at,
"multiplier": multiplier,
@@ -1,4 +1,6 @@
from .utils import latent_interpolate_transform, get_offset_list
import comfy
import torch
class LatentInterpolateTransform:
@@ -23,15 +25,15 @@ class LatentInterpolateTransform:
CATEGORY = "sampling/transforms"
def process(self,
offset_optional,
latent,
start_at=0,
stop_at=0,
factor=0.5,
multiplier=1):
multiplier=1,
offset_optional=None):
return ([{
"params": {
"latent": latent["samples"][0],
"latent": latent["samples"][0].unsqueeze(0),
"start_at": start_at,
"stop_at": stop_at,
"factor": factor,
@@ -24,11 +24,11 @@ class MirrorTransform:
CATEGORY = "sampling/transforms"
def process(self,
offset_optional,
start_at=0,
stop_at=0,
mode="replace",
direction="horizontally",):
direction="horizontally",
offset_optional=None):
return ([{
"params": {
"start_at": start_at,
@@ -22,11 +22,11 @@ class MultiplyTransform:
CATEGORY = "sampling/transforms"
def process(self,
offset_optional,
start_at=0,
stop_at=0,
mode="combine",
multiplier=1):
multiplier=1,
offset_optional=None):
return ([{
"params": {
"start_at": start_at,
@@ -1,4 +1,4 @@
from itertools import chain
class OffsetCombine:
@@ -8,6 +8,7 @@ class OffsetCombine:
"required": {
"offset1": ("OFFSET", ),
"offset2": ("OFFSET", ),
"offset": ("INT", {"default": 0, "min": -10000, "max": 10000}),
}
}
@@ -16,5 +17,10 @@ class OffsetCombine:
CATEGORY = "sampling/transforms"
def combine(self, offset1, offset2):
return (offset1 + offset2,)
def combine(self, offset, **kwargs):
offsets = sum(chain([v for k, v in kwargs.items()]), [])
for o in offsets:
o["offset"] += offset
return (offsets,)
@@ -23,12 +23,12 @@ class ShiftTransform:
CATEGORY = "sampling/transforms"
def process(self,
offset_optional,
start_at=0,
stop_at=0,
mode="replace",
x_shift=0,
y_shift=0):
y_shift=0,
offset_optional=None):
return ([{
"params": {
"start_at": start_at,
+19 -19
View File
@@ -17,14 +17,14 @@ def shift_transform(x0, params):
if params["mode"] == "replace":
if params["x_shift"] != 0:
x = torch.roll(x, shifts=int(x.size()[2] * params["x_shift"]), dims=[2])
x = torch.roll(x, shifts=int(x.size()[2] * params["x_shift"]), dims=[3])
if params["y_shift"] != 0:
x = torch.roll(x, shifts=int(x.size()[1] * params["y_shift"]), dims=[1])
x = torch.roll(x, shifts=int(x.size()[1] * params["y_shift"]), dims=[2])
elif params["mode"] == "combine":
if params["x_shift"] != 0:
x = (torch.roll(x, shifts=int(x.size()[2] * params["x_shift"]), dims=[2]) + x) / 2
x = (torch.roll(x, shifts=int(x.size()[2] * params["x_shift"]), dims=[3]) + x) / 2
if params["y_shift"] != 0:
x = (torch.roll(x, shifts=int(x.size()[1] * params["y_shift"]), dims=[1]) + x) / 2
x = (torch.roll(x, shifts=int(x.size()[1] * params["y_shift"]), dims=[2]) + x) / 2
return x
@@ -34,50 +34,50 @@ def mirror_transform(x0, params):
if params["mode"] == "replace":
if params["direction"] == "vertically":
x = torch.flip(x, [1])
elif params["direction"] == "horizontally":
x = torch.flip(x, [2])
elif params["direction"] == "horizontally":
x = torch.flip(x, [3])
elif params["direction"] == "both":
x = torch.flip(x, [1, 2])
x = torch.flip(x, [2, 3])
elif params["direction"] == "90 degree rotation":
x = torch.rot90(x, dims=[1, 2])
x = torch.rot90(x, dims=[2, 3])
elif params["direction"] == "180 degree rotation":
x = torch.rot90(torch.rot90(x, dims=[1, 2]), dims=[1, 2])
x = torch.rot90(torch.rot90(x, dims=[2, 3]), dims=[2, 3])
elif params["mode"] == "combine":
if params["direction"] == "vertically":
x = (torch.flip(x, [1]) + x) / 2
elif params["direction"] == "horizontally":
x = (torch.flip(x, [2]) + x) / 2
elif params["direction"] == "horizontally":
x = (torch.flip(x, [3]) + x) / 2
elif params["direction"] == "both":
x = (torch.flip(x, [1, 2]) + x) / 2
x = (torch.flip(x, [2, 3]) + x) / 2
elif params["direction"] == "90 degree rotation":
x = (torch.rot90(x, dims=[1, 2]) + x) / 2
x = (torch.rot90(x, dims=[2, 3]) + x) / 2
elif params["direction"] == "180 degree rotation":
x = (torch.rot90(torch.rot90(x, dims=[1, 2]), dims=[1, 2]) + x) / 2
x = (torch.rot90(torch.rot90(x, dims=[2, 3]), dims=[2, 3]) + x) / 2
return x
def latent_interpolate_transform(x0, params):
latent = params["latent"]
latent = params["latent"].to(x0.device)
if x0.shape != latent.shape:
latent.permute(0, 3, 1, 2)
latent = comfy.utils.common_upscale(latent, x0.shape[3], x0.shape[2], 'bicubic')
latent = comfy.utils.common_upscale(latent, x0.shape[3], x0.shape[2], 'bicubic', crop='center')
latent.permute(0, 2, 3, 1)
x = x0 * params["factor"] + latent * (1 - params["factor"])
x = latent * params["factor"] + x0 * (1 - params["factor"])
x *= params["multiplier"]
return x
def latent_add_transform(x0, params):
latent = params["latent"]
latent = params["latent"].to(x0.device)
if x0.shape != latent.shape:
latent.permute(0, 3, 1, 2)
latent = comfy.utils.common_upscale(latent, x0.shape[3], x0.shape[2], 'bicubic')
latent = comfy.utils.common_upscale(latent, x0.shape[3], x0.shape[2], 'bicubic', crop='center')
latent.permute(0, 2, 3, 1)
x = x0 + latent
+9 -1
View File
@@ -1,4 +1,5 @@
import torch
from nodes import PreviewImage
MIRROR_DIRECTIONS = ["vertically", "horizontally", "both"]
@@ -15,6 +16,9 @@ class LatentMirror:
"max": 10.0,
"step": 0.01
})
},
"optional": {
"vae_optional": ("VAE",)
}
}
@@ -23,7 +27,7 @@ class LatentMirror:
CATEGORY = "latent/advanced"
def mirror(self, latent, direction, multiplier):
def mirror(self, latent, direction, multiplier, vae_optional = None):
l = latent.copy()
if direction == "vertically" or direction == "both":
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"] *= multiplier
if vae_optional:
return {"result": (l,), "ui": PreviewImage().save_images(vae_optional.decode(l["samples"]))["ui"]}
return (l,)
+22
View File
@@ -0,0 +1,22 @@
import torch
from nodes import PreviewImage
class LatentNormalize:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"latent": ("LATENT",),
"vae": ("VAE",)
}
}
RETURN_TYPES = ("LATENT",)
FUNCTION = "normalize"
CATEGORY = "latent/advanced"
def normalize(self, latent, vae):
image = vae.decode(latent["samples"])
sample = vae.encode(image[:,:,:,:3])
return {"result": ({"samples": sample},), "ui": PreviewImage().save_images(image)["ui"]}
+9 -1
View File
@@ -1,4 +1,5 @@
import torch
from nodes import PreviewImage
class LatentShift:
@@ -19,6 +20,9 @@ class LatentShift:
"max": 1,
"step": 0.01
}),
},
"optional": {
"vae_optional": ("VAE",)
}
}
@@ -27,7 +31,7 @@ class LatentShift:
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()
if x_shift != 0:
@@ -35,4 +39,8 @@ class LatentShift:
if y_shift != 0:
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,)
+4 -2
View File
@@ -1,7 +1,9 @@
from .LatentMirror import LatentMirror
from .LatentShift import LatentShift
from .KSamplerNodes.TSamplerWithTransform import TSamplerWithTransform
from .KSamplerNodes.TransformSampler import TransformSampler
from .LatentNormalize import LatentNormalize
from .KSamplerNodes.TransformSampler import TSampler
from .KSamplerNodes.TransformSampler import TSamplerAdvanced
from .KSamplerNodes.TransformHijack import TransformHijack
from .KSamplerNodes.Transforms import MirrorTransform
from .KSamplerNodes.Transforms import MultiplyTransform
from .KSamplerNodes.Transforms import ShiftTransform