Files
AEmotionStudio-ComfyUI-Shad…/shaders/base.py
T
Æmotion StudioandClaude Opus 5 1a4b3a8808 fix: fill every latent channel at the source, not downstream
The generators handed back noise spanning about one channel however many the
latent had, and travel_mode walk worked around it by remixing 64 renders after
the fact. The plan was to fix shaders/base.py::expand_channels. That alone would
have changed nothing for SD 1.5 or SDXL: the collapse had three sources, and the
one that mattered at four channels was somewhere else.

- domain_warp copied one field four times with repeat(1, 4, 1, 1), so
  expand_channels returned early at four channels and never ran.
- temporal_coherent broadcast one field to every channel with .expand(): rank
  1.00 at any channel count, and it is the shader the video preset picks.
- curl_noise padded its colour path with copies of one magnitude field.

BaseNoiseGenerator.fill_channels replaces expand_channels. Every channel up to
CHANNEL_BASIS (64) is its own render at seed + 6151*c; channels past that are
QR-orthogonalised mixtures of those renders. Channel 0 stays the generator's own
draw, so a one-channel request is unchanged -- and jump, which is built from
one-channel draws, is byte-identical to before for all four generators. The
extra renders run inside a forked RNG. 6151 has no collisions among the first 64
channels under the mod-10000 seed hashing curl_noise does internally.

Effective rank for domain_warp at 4 / 24 / 128 channels goes from
1.00 / 2.13 / 2.43 to 3.91 / 22.72 / 69.56.

The travel-mode guard had to become direction-aware. Only a basis of 1 skipped
it, so once the generators were wide, drift (basis 4) would have been treated as
a remix that could not help and silently returned walk's noise. It already did
that for tensor_field and curl_noise. A basis below the widest now narrows
unconditionally. Widening is skipped once a draw exceeds 0.65 of its target,
just above the best a random remix reached (0.52-0.64, measured across all four
generators at 16, 24 and 128 channels), so walk passes the generator's noise
through instead of rendering 64 more draws and discarding them. The old 0.9
threshold would still have paid for that on 16-channel models.

Seven of the eleven goldens moved, not the one the handoff predicted; the four
that should not move held byte-identical. The suite is now pinned to the
standard pipeline, re-captured, with every standard-only input set explicitly in
NODE_DEFAULTS. The old legacy output is at the pre-collapse-fix tag.
test_legacy_golden.py becomes test_golden.py.

test_blend_calibration_is_current no longer compares difference on an absolute
tolerance: its whole curve sits below the tolerance, and after this change it
passed by 0.002. Every other mode stays within 0.049 of the table.

Co-Authored-By: Claude Opus 5 <noreply@anthropic.com>
2026-09-12 14:02:59 -07:00

276 lines
9.5 KiB
Python

"""
Base class for shader noise generators.
This module provides the abstract base class that all shader generators
must inherit from, ensuring consistent interface and shared functionality.
"""
import torch
from abc import ABC, abstractmethod
from typing import Callable, Dict, Any, Optional, Tuple
from ..utils.color_utils import apply_color_scheme
from ..utils.shape_masks import apply_shape_mask, apply_mask_to_tensor
from ..utils.noise_utils import create_coordinate_grid
from ..core.params import ShaderParams, get_param_value
from ..core.constants import CHANNEL_BASIS, DEFAULT_CHANNELS
# Seed step between a draw's channels. Prime and unrelated to
# core.shader_noise._MIX_SEED_STRIDE, and checked against the mod-10000 seed
# hashing curl_noise does internally: no two of the first 64 channels land on
# the same internal seed.
_CHANNEL_SEED_STRIDE = 6151
class BaseNoiseGenerator(ABC):
"""
Abstract base class for all shader noise generators.
Provides common functionality for coordinate grid creation,
shape mask application, and color scheme handling.
"""
@staticmethod
@abstractmethod
def generate(
batch_size: int,
height: int,
width: int,
params: ShaderParams,
device: torch.device,
seed: int = 0,
target_channels: int = DEFAULT_CHANNELS
) -> torch.Tensor:
"""
Generate noise tensor.
Args:
batch_size: Number of images in batch
height: Height of tensor
width: Width of tensor
params: Shader parameters
device: Device to create tensor on
seed: Random seed for deterministic results
target_channels: Number of output channels
Returns:
Tensor with shape [batch_size, target_channels, height, width]
"""
pass
@staticmethod
def create_coordinate_grid(
batch_size: int,
height: int,
width: int,
device: torch.device,
dtype: torch.dtype = torch.float32,
range_type: str = "unit"
) -> torch.Tensor:
"""
Create a coordinate grid for noise generation.
Args:
batch_size: Number of batches
height: Grid height
width: Grid width
device: Target device
dtype: Data type
range_type: "unit" for [0, 1], "centered" for [-0.5, 0.5], "symmetric" for [-1, 1]
Returns:
Coordinate tensor [B, H, W, 2] with (x, y) coordinates
"""
return create_coordinate_grid(batch_size, height, width, device, dtype, range_type)
@staticmethod
def apply_shape_mask(
noise: torch.Tensor,
coords: torch.Tensor,
params: ShaderParams
) -> torch.Tensor:
"""
Apply shape mask to noise tensor.
Args:
noise: Input noise tensor [B, H, W, C] or [B, C, H, W]
coords: Coordinate grid [B, H, W, 2]
params: Shader parameters containing shape_type and shape_strength
Returns:
Masked noise tensor
"""
shape_type = params.shape_type
shape_strength = params.shape_strength
time = params.time
base_seed = params.get("base_seed", 0)
if shape_type in ["none", "0"] or shape_strength <= 0:
return noise
# Generate shape mask
mask = apply_shape_mask(coords, shape_type, time, base_seed, shape_strength)
# Apply mask to noise
return apply_mask_to_tensor(noise, mask, shape_strength)
@staticmethod
def apply_color_scheme(
noise: torch.Tensor,
params: ShaderParams,
velocity_field: Optional[torch.Tensor] = None
) -> torch.Tensor:
"""
Apply color scheme to noise tensor.
Args:
noise: Input noise tensor [B, C, H, W]
params: Shader parameters containing color_scheme and color_intensity
velocity_field: Optional velocity field for direction-based coloring [B, 2, H, W]
Returns:
Color-modified noise tensor
"""
color_scheme = params.color_scheme
color_intensity = params.color_intensity
time = params.time
if color_scheme in ["none", "0"] or color_intensity <= 0:
return noise
return apply_color_scheme(noise, color_scheme, color_intensity, velocity_field, time)
@staticmethod
def apply_common_postprocessing(
noise: torch.Tensor,
params: ShaderParams,
coords: torch.Tensor,
velocity_field: Optional[torch.Tensor] = None
) -> torch.Tensor:
"""
Apply common postprocessing steps (shape mask, color scheme).
Args:
noise: Input noise tensor [B, C, H, W]
params: Shader parameters
coords: Coordinate grid [B, H, W, 2]
velocity_field: Optional velocity field for color schemes
Returns:
Postprocessed noise tensor
"""
# Apply color scheme first (operates on channels)
noise = BaseNoiseGenerator.apply_color_scheme(noise, params, velocity_field)
# Apply shape mask using the coords parameter
noise = BaseNoiseGenerator.apply_shape_mask(noise, coords, params)
return noise
@staticmethod
def normalize_to_range(
tensor: torch.Tensor,
target_min: float = -1.0,
target_max: float = 1.0
) -> torch.Tensor:
"""
Normalize tensor to target range.
Args:
tensor: Input tensor
target_min: Target minimum value
target_max: Target maximum value
Returns:
Normalized tensor
"""
t_min = tensor.min()
t_max = tensor.max()
if t_max - t_min < 1e-8:
# Avoid division by zero for constant tensors
return torch.full_like(tensor, (target_min + target_max) / 2)
normalized = (tensor - t_min) / (t_max - t_min)
return normalized * (target_max - target_min) + target_min
@staticmethod
def fill_channels(
render: Callable[[int], torch.Tensor],
base: torch.Tensor,
target_channels: int,
seed: int,
) -> torch.Tensor:
"""
Fill the channel axis with independent draws of the generator's field.
A sampler expects every latent channel to carry its own noise. Building the
extra channels out of the first one or two -- copying them, or passing them
through sin and abs -- leaves the draw spanning about one channel however
many it has: rank 1.00 for domain_warp at four. The shader then stops
steering the sample and starts overwriting it.
Args:
render: draws one field [B, 1, H, W] at the seed it is given
base: channels the generator has already drawn, kept as the first ones.
Channel 0 is the generator's own draw at `seed`, so a one-channel
request comes back exactly as it did before this existed -- which
matters, because core.shader_noise builds the travel-mode basis
from one-channel draws.
target_channels: channels to return
seed: the seed channel 0 was drawn with
Returns:
[B, target_channels, H, W]
Channels up to CHANNEL_BASIS are each rendered. Past it, the rest are
mixtures of those renders through a seeded orthogonal matrix, which keeps
the mixtures from re-correlating what the renders kept apart.
"""
if base.shape[1] >= target_channels:
return base[:, :target_channels]
seed = int(seed.item() if isinstance(seed, torch.Tensor) else seed)
rendered = min(target_channels, max(CHANNEL_BASIS, base.shape[1]))
# Generators reseed the global RNG inside each render. Forking leaves the
# caller's RNG exactly where channel 0 left it.
devices = [base.device.index if base.device.index is not None else torch.cuda.current_device()] \
if base.device.type == "cuda" else []
with torch.random.fork_rng(devices=devices):
extra = [render(seed + _CHANNEL_SEED_STRIDE * c).to(base)
for c in range(base.shape[1], rendered)]
channels = torch.cat([base, *extra], dim=1)
remaining = target_channels - rendered
if remaining <= 0:
return channels
mixer = torch.Generator(device="cpu").manual_seed(seed)
if remaining >= rendered:
weights = torch.linalg.qr(
torch.randn(remaining, rendered, generator=mixer, dtype=torch.float64)).Q
else:
weights = torch.linalg.qr(
torch.randn(rendered, remaining, generator=mixer, dtype=torch.float64)).Q.T
weights = weights / weights.norm(dim=1, keepdim=True).clamp_min(1e-12)
mixed = torch.einsum("mc,bchw->bmhw", weights.to(device=base.device, dtype=base.dtype), channels)
return torch.cat([channels, mixed], dim=1)
@staticmethod
def get_target_channels(
params: ShaderParams,
default: int = DEFAULT_CHANNELS
) -> int:
"""
Get target channel count from parameters.
Args:
params: Shader parameters
default: Default channel count
Returns:
Target number of channels
"""
return int(params.get("target_channels", default))