custom_noise: short distance noiseとtkg noiseを追加

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
This commit is contained in:
laksjdjf
2026-07-04 08:06:13 +09:00
co-authored by Claude Fable 5
parent 5b94269770
commit 87508dfd8e
3 changed files with 284 additions and 2 deletions
+10 -2
View File
@@ -1,16 +1,24 @@
from .variation_noise import VariationNoise, RandomNoiseOffset, RandomNoiseVariationSimple
from .short_distance_noise import ShortDistanceNoise, SameColorNoise
from .tkg_noise import TKGRandomNoise
from ... import SYMBOL, NODE_SURFIX
NODE_CLASS_MAPPINGS = {
f"VariationNoise{NODE_SURFIX}": VariationNoise,
f"RandomNoiseOffset{NODE_SURFIX}": RandomNoiseOffset,
f"RandomNoiseVariationSimple{NODE_SURFIX}": RandomNoiseVariationSimple
f"RandomNoiseVariationSimple{NODE_SURFIX}": RandomNoiseVariationSimple,
f"TKGRandomNoise{NODE_SURFIX}": TKGRandomNoise,
f"ShortDistanceNoise{NODE_SURFIX}": ShortDistanceNoise,
f"SameColorNoise{NODE_SURFIX}": SameColorNoise,
}
NODE_DISPLAY_NAME_MAPPINGS = {
f"VariationNoise{NODE_SURFIX}": f"Variation Noise {SYMBOL}",
f"RandomNoiseOffset{NODE_SURFIX}": f"Random Noise Offset {SYMBOL}",
f"RandomNoiseVariationSimple{NODE_SURFIX}": f"Random Noise Variation Simple {SYMBOL}"
f"RandomNoiseVariationSimple{NODE_SURFIX}": f"Random Noise Variation Simple {SYMBOL}",
f"TKGRandomNoise{NODE_SURFIX}": f"TKG Random Noise {SYMBOL}",
f"ShortDistanceNoise{NODE_SURFIX}": f"Short Distance Noise {SYMBOL}",
f"SameColorNoise{NODE_SURFIX}": f"Same Color Noise {SYMBOL}",
}
__all__ = ["NODE_CLASS_MAPPINGS", "NODE_DISPLAY_NAME_MAPPINGS"]
@@ -0,0 +1,91 @@
import comfy
from ... import ROOT_NAME
import torch
CATEGORY_NAME = ROOT_NAME + "custom_noise"
class Noise_ShortDistance:
def __init__(self, seed, num_samples=32, reference_latents=None):
self.seed = seed
self.num_samples = num_samples
self.reference_latents = reference_latents
def generate_noise(self, input_latent):
assert self.reference_latents["samples"].shape == input_latent["samples"].shape, "Reference latents and input latents must have the same shape."
latent = self.reference_latents["samples"].to(input_latent["samples"].device, dtype=input_latent["samples"].dtype)
batch_inds = input_latent.get("batch_index", None)
B = latent.shape[0]
K = self.num_samples
latent_repeat = latent.unsqueeze(1).repeat(1, K, *[1 for _ in latent.shape[1:]])
noise = comfy.sample.prepare_noise(latent_repeat, self.seed, batch_inds)
diff = (latent_repeat - noise) ** 2
dist = diff.flatten(start_dim=2).sum(dim=2)
best_idx = dist.argmin(dim=1)
# gatherで最短ノイズを選択
idx_expand = best_idx.view(B, 1, *[1 for _ in latent.shape[1:]]).expand_as(latent_repeat[:, :1])
best_noise = noise.gather(1, idx_expand).squeeze(1)
return best_noise
class ShortDistanceNoise:
@classmethod
def INPUT_TYPES(s):
return {
"required":{
"seed": ("INT", {"default": 0, "min": 0, "max": 0xffffffffffffffff}),
"num_samples": ("INT", {"default": 32, "min": 1, "max": 4096}),
"reference_latents": ("LATENT", ),
}
}
RETURN_TYPES = ("NOISE",)
FUNCTION = "get_noise"
CATEGORY = CATEGORY_NAME
def get_noise(self, seed, reference_latents):
return (Noise_ShortDistance(seed, reference_latents),)
class Noise_SameColor:
def __init__(self, seed, reference_latents, strength, **kwargs):
self.seed = seed
self.reference_latents = reference_latents
self.strength = strength
self.channel_mask = torch.tensor([1.0 if kwargs.get(f"ch_{i:02d}", True) else 0.0 for i in range(16)])
def generate_noise(self, input_latent):
assert self.reference_latents["samples"].shape == input_latent["samples"].shape, "Reference latents and input latents must have the same shape."
latent = self.reference_latents["samples"].to(input_latent["samples"].device, dtype=input_latent["samples"].dtype)
batch_inds = input_latent.get("batch_index", None)
noise = comfy.sample.prepare_noise(latent, self.seed, batch_inds)
latent_mean = latent.mean(dim=1, keepdim=True)
noise_mean = noise.mean(dim=1, keepdim=True)
channel_mask = self.channel_mask.to(latent.device, dtype=latent.dtype).view(1, -1, *[1 for _ in range(len(latent.shape)-2)])
noise = noise + (latent_mean - noise_mean) * self.strength * channel_mask
return noise
class SameColorNoise:
@classmethod
def INPUT_TYPES(s):
retval = {
"required":{
"seed": ("INT", {"default": 0, "min": 0, "max": 0xffffffffffffffff}),
"reference_latents": ("LATENT", ),
"strength": ("FLOAT", {"default": 0.1, "min": -1.0, "max": 1.0, "step": 0.01}),
}
}
for i in range(16):
retval["required"][f"ch_{i:02d}"] = ("BOOLEAN", {"default": True})
return retval
RETURN_TYPES = ("NOISE",)
FUNCTION = "get_noise"
CATEGORY = CATEGORY_NAME
def get_noise(self, seed, reference_latents, strength, **kwargs):
return (Noise_SameColor(seed, reference_latents, strength, **kwargs),)
+183
View File
@@ -0,0 +1,183 @@
import comfy
from typing import NamedTuple
import torch
import torch.nn.functional as F
from ... import ROOT_NAME
CATEGORY_NAME = ROOT_NAME + "custom_noise"
def get_mean_shifted_latents(
latents: torch.Tensor,
shift: float = 0.11,
delta_shift: float = 0.1,
channels: list[float] = [0, 1, 1, 0], # list of {-1, 0, 1}
) -> torch.Tensor:
shifted_latents = latents.clone()
for idx, sign in enumerate(channels):
if sign == 0:
# skip
continue
latent_channel = shifted_latents[:, idx, :, :]
positive_ratio = (latent_channel > 0).float().mean()
target_ratio = positive_ratio + shift * sign
# gradually shift latent_channel
while True:
latent_channel += delta_shift * sign
new_positive_ratio = (latent_channel > 0).float().mean()
if new_positive_ratio >= target_ratio:
break
# replace the channel in the original latents
shifted_latents[:, idx, :, :] = latent_channel
return shifted_latents
def get_2d_gaussian(
latent_height: int,
latent_width: int,
std_dev: float,
device: torch.device,
center_x: float = 0.0,
center_y: float = 0.0,
factor: int = 8, # idk why
):
y = torch.linspace(-1, 1, steps=latent_height // factor, device=device)
x = torch.linspace(-1, 1, steps=latent_width // factor, device=device)
y_grid, x_grid = torch.meshgrid(y, x, indexing="ij")
x_grid = x_grid - center_x
y_grid = y_grid - center_y
gauss = torch.exp(-((x_grid**2 + y_grid**2) / (2 * std_dev**2)))
gauss = gauss[None, None, :, :] # add batch and channel dimensions
return gauss
def apply_tkg_noise(
latents: torch.Tensor,
shift: float = 0.11,
delta_shift: float = 0.1,
std_dev: float = 0.5,
factor: int = 8,
channels: list[float] = [0, 1, 1, 0],
):
batch_size, num_channels, latent_height, latent_width = latents.shape
shifted_latents = get_mean_shifted_latents(
latents,
shift=shift,
delta_shift=delta_shift,
channels=channels,
)
gauss_mask = get_2d_gaussian(
latent_height=latent_height,
latent_width=latent_width,
std_dev=std_dev,
center_x=0.0,
center_y=0.0,
factor=factor,
device=latents.device,
)
gauss_mask = F.interpolate(
gauss_mask,
size=(latent_height, latent_width),
mode="bilinear",
align_corners=False,
)
gauss_mask = gauss_mask.expand(batch_size, num_channels, -1, -1)
noised_latents = shifted_latents * (1 - gauss_mask) + latents * gauss_mask
return noised_latents
class ColorSet(NamedTuple):
name: str
channels: list[float]
# ref: Figure 28. Additional Result in various color Background with SD
COLOR_SETS: list[ColorSet] = [
ColorSet("green", [0, 1, 1, 0]),
ColorSet("cyan", [0, 1, 0, 0]),
ColorSet("magenta", [0, -1, -1, -1]),
ColorSet("purple", [0, 0, -1, -1]),
ColorSet("black", [-1, 0, 0, 1]),
ColorSet("orange", [-1, -1, 1, 0]),
ColorSet("white", [0, 0, 0, -1]),
ColorSet("yellow", [0, -1, 1, -1]),
]
COLOR_SET_MAP: dict[str, ColorSet] = {c.name: c for c in COLOR_SETS}
class Noise_RandomNoise:
def __init__(self, seed, color="green", shift=0.11, grid_factor=8):
self.seed = seed
self.color = color
self.shift = shift
self.grid_factor = grid_factor
def generate_noise(self, input_latent):
latent_image = input_latent["samples"]
batch_inds = input_latent["batch_index"] if "batch_index" in input_latent else None
noise = comfy.sample.prepare_noise(latent_image, self.seed, batch_inds)
color_set = COLOR_SET_MAP.get(self.color, COLOR_SET_MAP["green"])
noise = apply_tkg_noise(
noise,
shift=self.shift,
channels=color_set.channels,
factor=self.grid_factor,
)
return noise
class TKGRandomNoise:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"noise_seed": ("INT", {
"default": 0,
"min": 0,
"max": 0xffffffffffffffff,
"control_after_generate": True,
}),
"color": (
[c.name for c in COLOR_SETS],
{
"default": "green",
},
),
"shift": (
"FLOAT",
{
"default": 0.11,
"min": 0.0,
"max": 1.0,
"step": 0.01,
},
),
"grid_factor": (
"INT",
{
"default": 8,
"min": 1,
"max": 16,
"step": 1,
},
),
}
}
RETURN_TYPES = ("NOISE",)
FUNCTION = "get_noise"
CATEGORY = CATEGORY_NAME
def get_noise(self, noise_seed, color, shift, grid_factor):
return (Noise_RandomNoise(noise_seed, color, shift, grid_factor),)