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>
276 lines
9.5 KiB
Python
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))
|