fix: draw a clip's first frame from the same field as the rest

domain_warp chose between a 2D field and a 3D one with time as the third axis
by testing `time != 0`. core.shader_noise ramps time from 0 to 1 across a
clip, so frame 0 is the only frame whose time is exactly 0, and it alone took
the 2D path. Every video drawn with the default shader had a first frame that
came from a different function than the five, or thirty-six, after it.

With temporal coherence on, adjacent frames correlate at about 0.95. Frame 0
against frame 1 measured -0.01, which is what uncorrelated looks like.
Nudging the base time to 0.001 made it +0.95, which is what confirmed the
cause.

The choice is now made once per draw, from the latent's frame count, and
passed down as `time_axis`. A single image still takes the 2D path -- frames
== 1 -- so no image fixture moved. In a video draw only frame 0 changes;
frames 1 onward are byte-identical, checked. video_nested was re-captured.

Where `time_axis` is absent, because something called the generator directly,
the old test still applies, so nothing outside this pipeline changes
behaviour.

Also corrects the CHANGELOG's claim that domain_warp shows no frame-to-frame
correlation under temporal coherence. That number came from measuring the one
frame pair that straddled this bug. It is the most coherent of the five, not
the least: 0.96 against curl_noise's 0.96, tensor_field's 0.80, spectral's
0.68 and temporal_coherent's 0.55.

Co-Authored-By: Claude Opus 5 (1M context) <noreply@anthropic.com>
This commit is contained in:
Æmotion Studio
2026-09-17 19:34:14 -07:00
co-authored by Claude Opus 5
parent 298e8bdc72
commit 33a57ae7b2
4 changed files with 50 additions and 15 deletions
+15 -3
View File
@@ -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
+6
View File
@@ -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 []
+29 -12
View File
@@ -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,
Binary file not shown.