5 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
13 changed files with 379 additions and 358 deletions
+17 -33
View File
@@ -38,37 +38,8 @@ This node can shift latent along x and y-axis.
**Usage:**
![sample](https://i.imgur.com/1Dp5dSw.png)
## TSampler with transforms (Latent Control)
This node can multiply, mirror and shift latent during generation.
**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)
## ~~TSampler with transforms (Latent Control)~~
Removed from version 2.0.0
## TSampler (Latent Control)
This node allows to combine a lot of transforms with different parameters.
@@ -125,10 +96,10 @@ Each transform node has own one-time version. They allow to make one transform a
Fixes some issues when sampling modified latent space.
**Input**
**Input:**
exactly matches the `VAE Decode` node
**Output**
**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
@@ -150,4 +121,17 @@ 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)
+5 -3
View File
@@ -7,8 +7,9 @@ NODE_CLASS_MAPPINGS = {
"LatentMirror": LatentMirror,
"LatentShift": LatentShift,
"LatentNormalize": LatentNormalize,
"TSamplerWithTransform": TSamplerWithTransform,
"TransformSampler": TransformSampler,
"TransformSampler": TSampler,
"TransformSamplerAdvanced": TSamplerAdvanced,
"TransformHijack": TransformHijack,
"MirrorTransform": MirrorTransform,
"ShiftTransform": ShiftTransform,
"MultiplyTransform": MultiplyTransform,
@@ -28,8 +29,9 @@ NODE_DISPLAY_NAME_MAPPINGS = {
"LatentMirror": "Latent mirror",
"LatentShift": "Latent shift",
"LatentNormalize": "Latent normalize",
"TSamplerWithTransform": "TSampler with transforms (Latent Control)",
"TransformSampler": "TSampler (Latent Control)",
"TransformSamplerAdvanced": "TSampler Advanced (Latent Control)",
"TransformHijack": "Transform Hijack",
"MirrorTransform": "Mirror transform",
"ShiftTransform": "Shift transform",
"MultiplyTransform": "Multiply transform",
+131 -123
View File
@@ -1,153 +1,122 @@
import { app } from "/scripts/app.js";
import {computeCanvasSize, generatePattern, recursiveLinkUpstream} from "./utils.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; // координата x для квадратика
const x = widgetX + index * squareSize;
// Устанавливаем цвет заливки и обводки
ctx.fillStyle = value === 1 ? "#222223" : "#00000000";
ctx.strokeStyle = "#222223"; // Цвет обводки для всех квадратиков
ctx.strokeStyle = "#222223";
ctx.lineWidth = Math.floor(squareSize / 16);
if (value === 1) {
// Если значение 1, закрашиваем квадрат
ctx.fillRect(x, widgetY, squareSize, squareSize);
ctx.fillRect(x, widgetY, actualSquareSize, actualSquareSize);
}
// Рисуем обводку для всех квадратиков
ctx.strokeRect(x, widgetY, squareSize, squareSize);
ctx.strokeRect(x, widgetY, actualSquareSize, actualSquareSize);
// Добавляем текст в квадратик
if (squareSize >= 24) {
ctx.font = `bold ${squareSize/3}px Arial`; // Размер шрифта адаптируем под размер квадратика
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 + squareSize/2, widgetY + squareSize/2); // Позиционируем текст по центру квадратика
ctx.fillStyle = value === 1 ? "#dbdbdc" : "#222223";
ctx.fillText(value.toString(), x + actualSquareSize/2, widgetY + actualSquareSize/2);
}
});
}
function addOffsetCanvas(node, app) {
const widget = {
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)
}
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 = []
let patterns = []
if (node.type === "OffsetCombine") {
const inputList = [...Array(node.inputs.length).keys()]
for (let i of inputList) {
const connectedNodes = recursiveLinkUpstream(node, node.inputs[i].type, 0, i)
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,
mode: connectedNode.widgets[2].value
}
patterns.push(pattern)
}
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)
},
};
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 {
const pattern = {
process_every: node.widgets[0].value,
offset: node.widgets[1].value,
mode: node.widgets[2].value}
patterns.push(pattern)
}
};
node.onResize = function (size) {
computeCanvasSize(node, size);
}
const pattern = generatePattern(patterns)
return { minWidth: 200, minHeight: 200, widget }
}
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",
@@ -157,7 +126,7 @@ app.registerExtension({
nodeType.prototype.onNodeCreated = function () {
const r = onNodeCreated ? onNodeCreated.apply(this, arguments) : undefined;
addOffsetCanvas(this, app)
addCanvas(this, app, offsetWidget)
return r;
}
@@ -173,7 +142,46 @@ app.registerExtension({
nodeType.prototype.onNodeCreated = function () {
const r = onNodeCreated ? onNodeCreated.apply(this, arguments) : undefined;
addOffsetCanvas(this, app)
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;
}
+59 -15
View File
@@ -83,26 +83,24 @@ export function generatePattern(rules) {
return pattern;
}
export function recursiveLinkUpstream(node, type, depth, index=null) {
export function recursiveLinkUpstream(node, slot_type, node_type, depth) {
depth += 1
let connections = []
if (node.type === "OffsetCombine") {
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]
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 === type) {
if (connectedNode.outputs[slotID].type === slot_type) {
connections.push([connectedNode.id, depth])
connections.push([connectedNode.id, depth])
if (connectedNode.inputs) {
const index = (connectedNode.type === "OffsetCombine") ? 0 : null
connections = connections.concat(recursiveLinkUpstream(connectedNode, type, depth, index))
}
if (connectedNode.inputs) {
const index = (connectedNode.type === node_type) ? 0 : null
connections = connections.concat(recursiveLinkUpstream(connectedNode, slot_type, node_type, depth))
}
}
}
@@ -110,3 +108,49 @@ export function recursiveLinkUpstream(node, type, depth, index=null) {
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].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)
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,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,)
+14 -14
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,26 +34,26 @@ 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
+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,)
+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,)
+3 -2
View File
@@ -1,8 +1,9 @@
from .LatentMirror import LatentMirror
from .LatentShift import LatentShift
from .LatentNormalize import LatentNormalize
from .KSamplerNodes.TSamplerWithTransform import TSamplerWithTransform
from .KSamplerNodes.TransformSampler import TransformSampler
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