Skip to content

Commit cd4873e

Browse files
committed
[bugfix] Make HYWorld memory retrieval deterministic
1 parent 3d8ac9d commit cd4873e

3 files changed

Lines changed: 163 additions & 6 deletions

File tree

fastvideo/models/dits/hyworld/retrieval_context.py

Lines changed: 39 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -22,17 +22,52 @@
2222
import math
2323

2424

25-
def generate_points_in_sphere(n_points: int, radius: float) -> torch.Tensor:
25+
def make_retrieval_generator(
26+
generator: torch.Generator | list[torch.Generator] | tuple[torch.Generator, ...] | None = None,
27+
seed: int | None = None,
28+
) -> torch.Generator:
29+
"""Create an isolated CPU generator for Monte Carlo memory retrieval.
30+
31+
The source generator is inspected only for its initial seed, so memory
32+
retrieval does not advance either the diffusion generator or the global
33+
RNG state. HYWorld currently supports one trajectory per batch, therefore
34+
a generator sequence follows its first trajectory.
35+
"""
36+
if isinstance(generator, (list, tuple)):
37+
if not generator:
38+
raise ValueError("generator list must not be empty")
39+
generator = generator[0]
40+
41+
if generator is not None:
42+
if not isinstance(generator, torch.Generator):
43+
raise TypeError("generator must be a torch.Generator or a sequence of them")
44+
retrieval_seed = generator.initial_seed()
45+
elif seed is not None:
46+
retrieval_seed = seed
47+
else:
48+
# Reading the initial seed does not consume the global RNG state.
49+
retrieval_seed = torch.initial_seed()
50+
51+
return torch.Generator(device="cpu").manual_seed(retrieval_seed)
52+
53+
54+
def generate_points_in_sphere(
55+
n_points: int,
56+
radius: float,
57+
generator: torch.Generator | None = None,
58+
) -> torch.Tensor:
2659
"""
2760
Uniformly sample points within a sphere of a specified radius.
2861
2962
:param n_points: The number of points to generate.
3063
:param radius: The radius of the sphere.
64+
:param generator: Optional generator used for all three random draws.
3165
:return: A tensor of shape (n_points, 3), representing the (x, y, z) coordinates of the points.
3266
"""
33-
samples_r = torch.rand(n_points)
34-
samples_phi = torch.rand(n_points)
35-
samples_u = torch.rand(n_points)
67+
sampling_device = generator.device if generator is not None else None
68+
samples_r = torch.rand(n_points, generator=generator, device=sampling_device)
69+
samples_phi = torch.rand(n_points, generator=generator, device=sampling_device)
70+
samples_u = torch.rand(n_points, generator=generator, device=sampling_device)
3671

3772
r = radius * torch.pow(samples_r, 1 / 3)
3873
phi = 2 * math.pi * samples_phi

fastvideo/pipelines/stages/hyworld_denoising.py

Lines changed: 4 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -18,7 +18,8 @@
1818
from fastvideo.pipelines.stages.validators import StageValidators as V
1919
from fastvideo.pipelines.stages.validators import VerificationResult
2020
from fastvideo.utils import dict_to_3d_list
21-
from fastvideo.models.dits.hyworld.retrieval_context import (generate_points_in_sphere, select_aligned_memory_frames)
21+
from fastvideo.models.dits.hyworld.retrieval_context import (generate_points_in_sphere, make_retrieval_generator,
22+
select_aligned_memory_frames)
2223
from fastvideo.models.dits.hyworld.pose import pose_to_input, compute_latent_num
2324

2425
logger = init_logger(__name__)
@@ -156,7 +157,8 @@ def forward(
156157

157158
# Generate local points if not provided
158159
if points_local is None:
159-
points_local = generate_points_in_sphere(50000, 8.0).to(device)
160+
retrieval_generator = make_retrieval_generator(batch.generator, batch.seed)
161+
points_local = generate_points_in_sphere(50000, 8.0, generator=retrieval_generator).to(device)
160162
else:
161163
points_local = points_local.to(device)
162164

Lines changed: 120 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,120 @@
1+
import numpy as np
2+
import pytest
3+
import torch
4+
5+
from fastvideo.models.dits.hyworld.retrieval_context import (
6+
generate_points_in_sphere,
7+
make_retrieval_generator,
8+
select_aligned_memory_frames,
9+
)
10+
11+
12+
def _w2c_at_x(x: float) -> np.ndarray:
13+
c2w = np.eye(4, dtype=np.float64)
14+
c2w[0, 3] = x
15+
return np.linalg.inv(c2w)
16+
17+
18+
def _symmetric_camera_path() -> np.ndarray:
19+
poses = np.stack([_w2c_at_x(0.0) for _ in range(36)])
20+
poses[4:8] = _w2c_at_x(1.0)
21+
poses[8:12] = _w2c_at_x(-1.0)
22+
poses[12:16] = _w2c_at_x(4.0)
23+
return poses
24+
25+
26+
def _select_memory(points: torch.Tensor) -> list[int]:
27+
return select_aligned_memory_frames(
28+
_symmetric_camera_path(),
29+
current_frame_idx=28,
30+
memory_frames=20,
31+
temporal_context_size=12,
32+
pred_latent_size=4,
33+
device="cpu",
34+
points_local=points,
35+
)
36+
37+
38+
def test_explicit_seed_is_independent_of_global_rng() -> None:
39+
torch.manual_seed(11)
40+
torch.rand(37)
41+
points_a = generate_points_in_sphere(4096, 8.0, generator=make_retrieval_generator(seed=1234))
42+
43+
torch.manual_seed(99)
44+
torch.rand(113)
45+
points_b = generate_points_in_sphere(4096, 8.0, generator=make_retrieval_generator(seed=1234))
46+
47+
torch.testing.assert_close(points_a, points_b, rtol=0, atol=0)
48+
assert _select_memory(points_a) == _select_memory(points_b)
49+
50+
51+
def test_retrieval_does_not_advance_source_or_global_rng() -> None:
52+
source = torch.Generator(device="cpu").manual_seed(4321)
53+
torch.rand(17, generator=source)
54+
source_state = source.get_state().clone()
55+
56+
torch.manual_seed(8765)
57+
global_state = torch.random.get_rng_state().clone()
58+
59+
points = generate_points_in_sphere(256, 2.0, generator=make_retrieval_generator(source))
60+
61+
assert points.shape == (256, 3)
62+
assert torch.equal(source.get_state(), source_state)
63+
assert torch.equal(torch.random.get_rng_state(), global_state)
64+
65+
66+
def test_same_initial_seed_ignores_source_state_progress() -> None:
67+
first = torch.Generator(device="cpu").manual_seed(31415)
68+
second = torch.Generator(device="cpu").manual_seed(31415)
69+
torch.rand(13, generator=first)
70+
torch.rand(79, generator=second)
71+
72+
points_first = generate_points_in_sphere(512, 3.0, generator=make_retrieval_generator(first))
73+
points_second = generate_points_in_sphere(512, 3.0, generator=make_retrieval_generator(second))
74+
75+
torch.testing.assert_close(points_first, points_second, rtol=0, atol=0)
76+
77+
78+
def test_source_generator_takes_precedence_over_seed() -> None:
79+
source = torch.Generator(device="cpu").manual_seed(17)
80+
81+
from_source = generate_points_in_sphere(128, 1.0, generator=make_retrieval_generator(source, seed=999))
82+
expected = generate_points_in_sphere(128, 1.0, generator=make_retrieval_generator(seed=17))
83+
84+
torch.testing.assert_close(from_source, expected, rtol=0, atol=0)
85+
86+
87+
def test_generator_list_uses_first_trajectory_seed() -> None:
88+
first = torch.Generator(device="cpu").manual_seed(7)
89+
second = torch.Generator(device="cpu").manual_seed(9)
90+
91+
from_list = generate_points_in_sphere(128, 1.0, generator=make_retrieval_generator([first, second]))
92+
from_first = generate_points_in_sphere(128, 1.0, generator=make_retrieval_generator(first))
93+
94+
torch.testing.assert_close(from_list, from_first, rtol=0, atol=0)
95+
96+
97+
def test_global_seed_fallback_is_reproducible_without_advancing_rng() -> None:
98+
torch.manual_seed(2468)
99+
global_state = torch.random.get_rng_state().clone()
100+
101+
fallback_points = generate_points_in_sphere(128, 1.0, generator=make_retrieval_generator())
102+
explicit_points = generate_points_in_sphere(128, 1.0, generator=make_retrieval_generator(seed=2468))
103+
104+
torch.testing.assert_close(fallback_points, explicit_points, rtol=0, atol=0)
105+
assert torch.equal(torch.random.get_rng_state(), global_state)
106+
107+
108+
def test_known_seeds_can_change_selected_symmetric_chunk() -> None:
109+
points_zero = generate_points_in_sphere(4096, 8.0, generator=make_retrieval_generator(seed=0))
110+
points_two = generate_points_in_sphere(4096, 8.0, generator=make_retrieval_generator(seed=2))
111+
112+
assert _select_memory(points_zero)[4:8] == [4, 5, 6, 7]
113+
assert _select_memory(points_two)[4:8] == [8, 9, 10, 11]
114+
115+
116+
def test_invalid_generator_sequences_are_rejected() -> None:
117+
with pytest.raises(ValueError, match="must not be empty"):
118+
make_retrieval_generator([])
119+
with pytest.raises(TypeError, match="torch.Generator"):
120+
make_retrieval_generator("not-a-generator")

0 commit comments

Comments
 (0)