diff --git a/CHANGELOG.md b/CHANGELOG.md index d09cacf..f96fbe8 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -27,13 +27,25 @@ so the first step away from strength 0 is only as large as the shader makes it. it has instead is direct control over the large-scale structure the model actually reads -- `noise_scale` sets the band, `octaves` the roll-off, `warp_strength` the anisotropy -- and temporal coherence for free: holding the seed and advancing time - turns each mode at its own rate rather than redrawing the field, which measures as - 0.27 frame-to-frame correlation against `temporal_coherent`'s 0.26 and - `domain_warp`'s 0.00. + turns each mode at its own rate rather than redrawing the field. Adjacent frames + of a 24-channel, 8-frame draw correlate at 0.68, against `temporal_coherent`'s + 0.55 and `tensor_field`'s 0.80. The figure depends on clip length: `time` always + spans 0 to 1 across the whole clip, so a longer clip takes smaller steps and + comes out smoother frame to frame. Its strengths are **not** calibrated against real prompts the way the other four are. Treat the presets' numbers as not applying to it yet. +### Fixed +- **The first frame of a video was drawn from a different noise function than the + rest.** `domain_warp` chose between a 2D and a 3D field by testing `time != 0`, + and frame 0 is the only frame whose time is exactly 0, so it alone took the 2D + path. With temporal coherence on, adjacent frames correlate at about 0.95 -- + except frame 0 against frame 1, which measured -0.01. The choice is now made once + per draw from the latent's frame count, so a clip is one field throughout. Single + images are unaffected, and only frame 0 of a video draw changes; the + `video_nested` fixture was re-captured. + ### Performance - **The channel axis is drawn in one call instead of one per channel.** At MiniMax diff --git a/core/shader_noise.py b/core/shader_noise.py index fc06679..b78d780 100644 --- a/core/shader_noise.py +++ b/core/shader_noise.py @@ -268,6 +268,12 @@ def _generate( generator = generator or resolve_generator(shader_type) base_params = _as_dict(params) base_time = float(base_params.get("time", 0.0) or 0.0) + # Whether `time` is a real axis for this draw, rather than whether a particular + # frame's value happens to be zero. A generator that switches between a 2D and a + # 3D field must make that choice once for the whole clip: deciding it per frame + # drew frame 0 -- the only frame whose time is exactly 0.0 -- from a different + # function than every frame after it. + base_params["time_axis"] = layout["frames"] > 1 devices = [device.index if device.index is not None else torch.cuda.current_device()] \ if torch.device(device).type == "cuda" else [] diff --git a/shaders/domain_warp.py b/shaders/domain_warp.py index bfe8813..4bfba70 100644 --- a/shaders/domain_warp.py +++ b/shaders/domain_warp.py @@ -23,6 +23,17 @@ from ..core.constants import DEFAULT_CHANNELS logger = logging.getLogger(__name__) +def _time_is_an_axis(time_axis, time): + """ + Whether this draw should evaluate its field in 3D with time as the third axis. + + `time_axis` is set by core.shader_noise from the latent's frame count. When it + is absent -- a direct caller, or an older saved path -- fall back to the old + test, which is what the golden fixtures for single images pin. + """ + return (time != 0) if time_axis is None else time_axis + + @shader_generator("domain_warp", metadata={"description": "Domain warping noise for swirling patterns"}) class DomainWarpGenerator(BaseNoiseGenerator): """ @@ -88,12 +99,13 @@ class DomainWarpGenerator(BaseNoiseGenerator): coords = create_coordinate_grid(batch_size, height, width, device) # Generate domain warp noise + time_axis = params.get("time_axis", None) warp_type = int(octaves % 4) try: result = DomainWarpGenerator._domain_warp_with_phase( coords, device, octaves, current_seed, 0, warp_type, - scale, warp_strength, phase_shift, time + scale, warp_strength, phase_shift, time, time_axis ) except Exception as e: logger.warning(f"Error generating domain warp: {e}, using fallback noise") @@ -138,7 +150,7 @@ class DomainWarpGenerator(BaseNoiseGenerator): def draw(channel_seed): field = DomainWarpGenerator._domain_warp_with_phase( coords, device, octaves, channel_seed, 0, warp_type, - scale, warp_strength, phase_shift, time + scale, warp_strength, phase_shift, time, time_axis ) * contrast if applied_mask is not None: field = torch.lerp(field, field * applied_mask, shape_strength) @@ -147,7 +159,7 @@ class DomainWarpGenerator(BaseNoiseGenerator): def draw_many(channel_seeds): fields = DomainWarpGenerator._domain_warp_with_phase( coords, device, octaves, channel_seeds.reshape(-1, 1, 1, 1, 1), - 0, warp_type, scale, warp_strength, phase_shift, time + 0, warp_type, scale, warp_strength, phase_shift, time, time_axis ) * contrast if applied_mask is not None: fields = torch.lerp(fields, fields * applied_mask, shape_strength) @@ -324,7 +336,8 @@ class DomainWarpGenerator(BaseNoiseGenerator): return result, g, b @staticmethod - def _domain_warp_with_phase(p, device, octaves, seed, warp_layer, warp_type, scale, warp_strength, phase_shift, time): + def _domain_warp_with_phase(p, device, octaves, seed, warp_layer, warp_type, scale, + warp_strength, phase_shift, time, time_axis=None): """ Generate domain warp noise with phase shift parameter. @@ -360,17 +373,21 @@ class DomainWarpGenerator(BaseNoiseGenerator): if warp_type == 0: # Standard FBM - result = DomainWarpGenerator.fbm_noise(warped_p, octaves_int, time, device, seed) + result = DomainWarpGenerator.fbm_noise(warped_p, octaves_int, time, device, seed, + time_axis=time_axis) elif warp_type == 1: # Ridged FBM - result = DomainWarpGenerator.fbm_noise(warped_p, octaves_int, time, device, seed) + result = DomainWarpGenerator.fbm_noise(warped_p, octaves_int, time, device, seed, + time_axis=time_axis) result = 1.0 - torch.abs(result) elif warp_type == 2: # Turbulent FBM - result = torch.abs(DomainWarpGenerator.fbm_noise(warped_p, octaves_int, time, device, seed)) + result = torch.abs(DomainWarpGenerator.fbm_noise( + warped_p, octaves_int, time, device, seed, time_axis=time_axis)) else: # Domain warp FBM - result = DomainWarpGenerator.fbm_noise_domain_warp(warped_p, octaves_int, time, device, seed) + result = DomainWarpGenerator.fbm_noise_domain_warp( + warped_p, octaves_int, time, device, seed, time_axis=time_axis) # Normalize result. Per draw when the seed is a tensor: each channel has to # be standardised against itself, exactly as it would be on its own, or a @@ -408,7 +425,7 @@ class DomainWarpGenerator(BaseNoiseGenerator): return simplex_3d(coords, seed, corners=1) @staticmethod - def fbm_noise(p, octaves, time, device, seed, use_temporal_coherence=True): + def fbm_noise(p, octaves, time, device, seed, use_temporal_coherence=True, time_axis=None): """ Generate FBM (Fractal Brownian Motion) noise. @@ -434,7 +451,7 @@ class DomainWarpGenerator(BaseNoiseGenerator): for i in range(min(octaves, 8)): current_p = p * freq - if use_temporal_coherence and time != 0: + if use_temporal_coherence and _time_is_an_axis(time_axis, time): time_offset = time * (0.2 + i * 0.05) current_p_3d = torch.cat([ current_p, @@ -453,7 +470,7 @@ class DomainWarpGenerator(BaseNoiseGenerator): return result / max_amp @staticmethod - def fbm_noise_domain_warp(p, octaves, time, device, seed): + def fbm_noise_domain_warp(p, octaves, time, device, seed, time_axis=None): """ Generate FBM noise with domain warping applied at each octave. """ @@ -477,7 +494,7 @@ class DomainWarpGenerator(BaseNoiseGenerator): current_p = p_warped * freq - if time != 0: + if _time_is_an_axis(time_axis, time): time_offset = time * (0.2 + i * 0.05) current_p_3d = torch.cat([ current_p, diff --git a/tests/golden/video_nested.pt b/tests/golden/video_nested.pt index 3e8d1e1..4d53334 100644 Binary files a/tests/golden/video_nested.pt and b/tests/golden/video_nested.pt differ