custom_noise: short distance noiseとtkg noiseを追加
Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
This commit is contained in:
co-authored by
Claude Fable 5
parent
5b94269770
commit
87508dfd8e
@@ -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),)
|
||||
@@ -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),)
|
||||
Reference in New Issue
Block a user