|
| 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