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:
co-authored by
Claude Opus 5
parent
298e8bdc72
commit
33a57ae7b2
+15
-3
@@ -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
|
||||
|
||||
@@ -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
@@ -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.
Reference in New Issue
Block a user