import json import subprocess import sys import tempfile import unittest from unittest import mock from pathlib import Path import numpy as np import torch from torch import nn from omegaconf import OmegaConf from datasets.video.memory_selection import ( _dynamic_multiview_selector, _dynamic_policy, _memory_candidate_frames, select_memory_indices, ) from datasets.video.minecraft_video_dememwm_latent_dataset import MinecraftVideoDeMemWMLatentDataset def _dynamic_cfg(selection_policy="recent", **overrides): cfg = { "selection_policy": selection_policy, "multiview_selector": "fov_greedy", "scene_threshold": 2.5, "state_threshold": 2.5, "stable_threshold": 1.0, "stable_frames": 3, "min_event_gap": 8, "min_anchor_score": 2.0, "max_event_anchors": None, "b_pose": 0.3, "b_action": 0.2, } cfg.update(overrides) return cfg def _selection_cfg(**overrides): cfg = { "enabled": True, "causal": True, "max_anchor_frames": 2, "max_dynamic_frames": 2, "max_revisit_frames": 2, "pose_similarity_threshold": 0.6, "training_use_plucker": False, "training_plucker_weight": 0.1, "training_dropout": 0.0, "fov_overlap_threshold": 0.0, "min_total_selected_coverage": 0.0, "local_context_exclusion_frames": 2, "plucker_moment_radius": 30.0, "anchor_diverse_selection": False, "pose_preselect_topk": 32, "candidate_chunk_size": 0, "dynamic": _dynamic_cfg(), } cfg.update(overrides) return OmegaConf.create(cfg) def _benchmark_script_path(): return Path(__file__).resolve().parents[1] / "scripts" / "benchmark_dememwm_multiview_selection.py" def _poses(num_frames): poses = np.zeros((num_frames, 5), dtype=np.float32) poses[:, 0] = np.arange(num_frames, dtype=np.float32) poses[:, 4] = np.arange(num_frames, dtype=np.float32) * 3.0 return poses def _fov_pool(candidate_ids, inside, *, fov_values=None, plucker=None, gaps=None, poses=None): candidate_ids = np.asarray(candidate_ids, dtype=np.int64) candidates_t = torch.as_tensor(candidate_ids, dtype=torch.long) inside_t = torch.as_tensor(inside, dtype=torch.bool) if fov_values is None: fov_values_t = inside_t.float().mean(dim=1) if inside_t.shape[1] > 0 else torch.zeros((len(candidate_ids),), dtype=torch.float32) else: fov_values_t = torch.as_tensor(fov_values, dtype=torch.float32) if plucker is None: plucker_t = torch.zeros((len(candidate_ids),), dtype=torch.float32) else: plucker_t = torch.as_tensor(plucker, dtype=torch.float32) if gaps is None: gaps_t = torch.arange(len(candidate_ids), 0, -1, dtype=torch.long) else: gaps_t = torch.as_tensor(gaps, dtype=torch.long) if poses is None: candidate_poses = torch.zeros((len(candidate_ids), 5), dtype=torch.float32) candidate_poses[:, 0] = candidates_t.to(dtype=torch.float32) else: candidate_poses = torch.as_tensor(poses[candidate_ids], dtype=torch.float32) return { "candidates_t": candidates_t, "candidate_poses": candidate_poses, "inside": inside_t, "fov_values": fov_values_t, "plucker": plucker_t, "gaps": gaps_t, "positive_fov": inside_t.any(dim=1), } def _dataset_cfg(root, **overrides): cfg = { "save_dir": str(root), "precomputed_feature_dir": str(Path(root) / "vae_features"), "n_frames": 3, "n_frames_valid": 3, "frame_skip": 1, "action_cond_dim": 25, "resolution": 128, "image_hw": [360, 640], "shuffle_clips": False, "single_eval_clip": True, "memory_selection": _selection_cfg(max_anchor_frames=0, max_dynamic_frames=0, max_revisit_frames=0), } cfg.update(overrides) return OmegaConf.create(cfg) def _event_latents(num_frames, event_frame): latents = np.zeros((num_frames, 2, 1, 1), dtype=np.float32) latents[:, 0, 0, 0] = 1.0 if event_frame < num_frames: latents[event_frame:, 0, 0, 0] = 0.0 latents[event_frame:, 1, 0, 0] = 1.0 return latents def _multi_event_latents(num_frames, event_frames): latents = np.zeros((num_frames, 2, 1, 1), dtype=np.float32) event_set = {int(frame) for frame in event_frames} state = 0 for frame in range(num_frames): if frame in event_set: state = 1 - state latents[frame, state, 0, 0] = 1.0 return latents def _write_vae_feature_clip(root, feature_root, split="training", subdir="scene", stem="000001", num_frames=104, actions=None, poses=None, latents=None): video_path = Path(root) / split / subdir / f"{stem}.mp4" video_path.parent.mkdir(parents=True, exist_ok=True) video_path.write_bytes(b"") if actions is None: actions = np.ones((num_frames, 8), dtype=np.int32) if poses is None: poses = _poses(num_frames) if latents is None: latents = np.arange(num_frames * 2, dtype=np.float16).reshape(num_frames, 1, 1, 2) np.savez(video_path.with_suffix(".npz"), actions=actions, poses=poses) feature_dir = Path(feature_root) / split / subdir feature_dir.mkdir(parents=True, exist_ok=True) np.save(feature_dir / f"{stem}_vae_feature.npy", latents) with open(feature_dir / f"{stem}_vae_feature_meta.json", "w", encoding="utf-8") as handle: json.dump({"image_height": 360, "image_width": 640}, handle) class MemorySelectionTests(unittest.TestCase): def test_memory_candidate_frames_default_training_is_causal_without_local_exclusion(self): cfg = _selection_cfg(local_context_exclusion_frames=99) candidates = _memory_candidate_frames(10, np.array([5]), cfg, "training", min_candidate_frame=1) self.assertEqual(candidates.tolist(), [1, 2, 3, 4]) def test_memory_candidate_frames_training_can_be_noncausal(self): cfg = _selection_cfg(causal=False, local_context_exclusion_frames=99) candidates = _memory_candidate_frames(8, np.array([3]), cfg, "training", min_candidate_frame=1) self.assertEqual(candidates.tolist(), [1, 2, 3, 4, 5, 6, 7]) def test_memory_candidate_frames_validation_and_test_stay_causal(self): cfg = _selection_cfg(causal=False) for split in ("validation", "test"): with self.subTest(split=split): candidates = _memory_candidate_frames(8, np.array([3]), cfg, split, min_candidate_frame=1) self.assertEqual(candidates.tolist(), [1, 2]) def test_revisit_excludes_local_context_and_target_window_but_keeps_future(self): import datasets.video.memory_selection as memory_selection poses = _poses(12) cfg = _selection_cfg( causal=False, max_anchor_frames=0, max_dynamic_frames=0, max_revisit_frames=2, local_context_exclusion_frames=2, ) seen_candidates = [] def fake_pose_similarity(poses_arg, candidates, target_positions, cfg_arg, count, rng=None): seen_candidates.append(candidates.copy()) return np.asarray([1, 8], dtype=np.int64) with mock.patch.object(memory_selection, "_select_by_pose_similarity", side_effect=fake_pose_similarity): indices, masks = select_memory_indices(poses, np.array([4, 5, 6, 7]), cfg, split="training") self.assertEqual([candidates.tolist() for candidates in seen_candidates], [[0, 1, 8, 9, 10, 11]]) self.assertEqual(indices["revisit"].tolist(), [1, 8]) self.assertEqual(masks["revisit"].tolist(), [True, True]) def test_dememwm_latent_config_has_dataset_side_memory_causal_default(self): config_path = Path(__file__).resolve().parents[1] / "configurations" / "dataset" / "video_minecraft_dememwm_latent.yaml" cfg = OmegaConf.load(config_path) self.assertIs(cfg.memory_selection.causal, True) self.assertEqual(cfg.memory_selection.dynamic.selection_policy, "multiview") self.assertEqual(cfg.memory_selection.dynamic.multiview_selector, "fov_greedy") self.assertEqual(cfg.memory_selection.pose_preselect_topk, 32) self.assertEqual(cfg.memory_selection.candidate_chunk_size, 0) self.assertEqual(cfg.memory_selection.training_dropout, 0.1) def test_revisit_uses_sampled_fov_selection(self): poses = np.array( [ [0, 0, 0, 0, 0], [0, 0, 0, 0, 180], [0, 0, 0, 0, 180], [0, 0, 0, 0, 180], [0, 0, 0, 0, 180], [0, 0, 0, 0, 0], ], dtype=np.float32, ) cfg = _selection_cfg( max_anchor_frames=0, max_dynamic_frames=0, max_revisit_frames=1, fov_overlap_threshold=0.5, local_context_exclusion_frames=1, ) indices, masks = select_memory_indices(poses, np.array([5]), cfg, split="validation") self.assertEqual(indices["revisit"].tolist(), [0]) self.assertEqual(masks["revisit"].tolist(), [True]) def test_sample_points_in_sphere_is_deterministic_without_torch_rand(self): import datasets.video.memory_selection as memory_selection center = torch.tensor([1.0, 2.0, 3.0, 0.0, 0.0], dtype=torch.float32) with mock.patch.object(torch, "rand", side_effect=AssertionError("unexpected random FOV sampling")): first = memory_selection._sample_points_in_sphere(center) second = memory_selection._sample_points_in_sphere(center) self.assertEqual(tuple(first.shape), (memory_selection._FOV_NUM_POINTS, 3)) self.assertTrue(torch.equal(first, second)) def test_fov_revisit_selection_is_deterministic_across_calls(self): poses = np.array( [ [0, 0, 0, 0, 0], [0, 0, 0, 0, 180], [0, 0, 0, 0, 180], [0, 0, 0, 0, 180], [0, 0, 0, 0, 180], [0, 0, 0, 0, 0], ], dtype=np.float32, ) cfg = _selection_cfg( max_anchor_frames=0, max_dynamic_frames=0, max_revisit_frames=1, fov_overlap_threshold=0.5, local_context_exclusion_frames=1, ) first_indices, first_masks = select_memory_indices(poses, np.array([5]), cfg, split="validation") second_indices, second_masks = select_memory_indices(poses, np.array([5]), cfg, split="validation") self.assertEqual(first_indices["revisit"].tolist(), [0]) self.assertEqual(second_indices["revisit"].tolist(), first_indices["revisit"].tolist()) self.assertEqual(second_masks["revisit"].tolist(), first_masks["revisit"].tolist()) def test_empty_fov_candidate_pool_returns_empty_selection(self): import datasets.video.memory_selection as memory_selection pool = memory_selection._build_fov_candidate_pool( _poses(4), np.empty((0,), dtype=np.int64), np.array([3], dtype=np.int64), _selection_cfg(), use_plucker=True, ) selected = memory_selection._select_by_point_union_from_pool(pool, count=1) self.assertEqual(tuple(pool["candidates_t"].shape), (0,)) self.assertEqual(tuple(pool["candidate_poses"].shape), (0, 5)) self.assertEqual(tuple(pool["inside"].shape), (0, 0)) self.assertEqual(pool["positive_fov"].dtype, torch.bool) self.assertEqual(selected.tolist(), []) def test_validation_fov_revisit_selection_stays_causal(self): poses = np.zeros((8, 5), dtype=np.float32) poses[:, 0] = 100.0 poses[:, 4] = 180.0 poses[0] = np.asarray([0, 0, 0, 0, 0], dtype=np.float32) poses[5] = np.asarray([0, 0, 0, 0, 0], dtype=np.float32) poses[6] = np.asarray([0, 0, 0, 0, 0], dtype=np.float32) cfg = _selection_cfg( causal=False, max_anchor_frames=0, max_dynamic_frames=0, max_revisit_frames=2, fov_overlap_threshold=0.5, local_context_exclusion_frames=1, ) indices, masks = select_memory_indices(poses, np.array([5]), cfg, split="validation") selected = indices["revisit"][masks["revisit"]] self.assertEqual(selected.tolist(), [0]) self.assertTrue(np.all(selected < 5)) def test_training_revisit_uses_pose_similarity_threshold(self): poses = np.array( [ [0.2, 0, 0, 0, 2], [40, 0, 0, 0, 180], [45, 0, 0, 0, 180], [50, 0, 0, 0, 180], [55, 0, 0, 0, 180], [0, 0, 0, 0, 0], ], dtype=np.float32, ) cfg = _selection_cfg( max_anchor_frames=0, max_dynamic_frames=0, max_revisit_frames=1, pose_similarity_threshold=0.6, local_context_exclusion_frames=1, ) indices, masks = select_memory_indices( poses, np.array([5]), cfg, split="training", rng=np.random.default_rng(0), ) self.assertEqual(indices["revisit"].tolist(), [0]) self.assertEqual(masks["revisit"].tolist(), [True]) def test_training_revisit_pads_when_pose_matches_are_short(self): poses = np.array( [ [0.2, 0, 0, 0, 2], [40, 0, 0, 0, 180], [45, 0, 0, 0, 180], [50, 0, 0, 0, 180], [55, 0, 0, 0, 180], [0, 0, 0, 0, 0], ], dtype=np.float32, ) cfg = _selection_cfg( max_anchor_frames=0, max_dynamic_frames=0, max_revisit_frames=2, pose_similarity_threshold=0.95, local_context_exclusion_frames=1, ) indices, masks = select_memory_indices( poses, np.array([5]), cfg, split="training", rng=np.random.default_rng(0), ) self.assertEqual(indices["revisit"].tolist(), [0, -1]) self.assertEqual(masks["revisit"].tolist(), [True, False]) def test_pose_similarity_refactor_preserves_deterministic_ranked_result(self): import datasets.video.memory_selection as memory_selection poses = np.array( [ [0.1, 0, 0, 0, 1], [2.0, 0, 0, 0, 4], [40.0, 0, 0, 0, 180], [0.0, 0, 0, 0, 0], ], dtype=np.float32, ) cfg = _selection_cfg(pose_similarity_threshold=0.0, training_use_plucker=True, training_plucker_weight=0.0) selected = memory_selection._select_by_pose_similarity(poses, np.asarray([0, 1, 2]), np.asarray([3]), cfg, count=2) self.assertEqual(selected.tolist(), [0, 1]) def test_pose_fps_returns_at_most_count_frames(self): import datasets.video.memory_selection as memory_selection ids = torch.as_tensor([0, 1], dtype=torch.long) poses = torch.as_tensor([[0, 0, 0, 0, 0], [10, 0, 0, 0, 0]], dtype=torch.float32) selected = memory_selection._select_pose_fps(ids, poses, count=5) self.assertEqual(selected.tolist(), [0, 1]) def test_pose_fps_is_deterministic(self): import datasets.video.memory_selection as memory_selection ids = torch.as_tensor([0, 1, 2], dtype=torch.long) poses = torch.as_tensor([[0, 0, 0, 0, 0], [1, 0, 0, 0, 0], [60, 0, 0, 0, 0]], dtype=torch.float32) first = memory_selection._select_pose_fps(ids, poses, count=2) second = memory_selection._select_pose_fps(ids, poses, count=2) self.assertEqual(first.tolist(), second.tolist()) def test_pose_fps_caps_pose_preselect_topk_before_diversity(self): import datasets.video.memory_selection as memory_selection poses = np.zeros((6, 5), dtype=np.float32) poses[1, 0] = 1.0 poses[2, 0] = 30.0 poses[3, 0] = 60.0 cfg = _selection_cfg( pose_similarity_threshold=0.0, training_use_plucker=True, training_plucker_weight=0.0, pose_preselect_topk=2, ) selected = memory_selection._select_dynamic_multiview_pose_plucker_fps( poses, np.asarray([0, 1, 2, 3], dtype=np.int64), np.asarray([5], dtype=np.int64), cfg, count=4, ) self.assertLessEqual(len(selected), 2) self.assertTrue(set(selected.tolist()) <= {0, 1}) def test_pose_fps_pose_preselect_topk_zero_is_uncapped(self): import datasets.video.memory_selection as memory_selection poses = np.zeros((6, 5), dtype=np.float32) poses[1, 0] = 1.0 poses[2, 0] = 30.0 poses[3, 0] = 60.0 cfg = _selection_cfg( pose_similarity_threshold=0.0, training_use_plucker=True, training_plucker_weight=0.0, pose_preselect_topk=0, ) selected = memory_selection._select_dynamic_multiview_pose_plucker_fps( poses, np.asarray([0, 1, 2, 3], dtype=np.int64), np.asarray([5], dtype=np.int64), cfg, count=4, ) self.assertEqual(len(selected), 4) self.assertEqual(set(selected.tolist()), {0, 1, 2, 3}) def test_pose_fps_prefers_diverse_candidate_poses(self): import datasets.video.memory_selection as memory_selection ids = torch.as_tensor([0, 1, 2], dtype=torch.long) poses = torch.as_tensor([[0, 0, 0, 0, 0], [1, 0, 0, 0, 0], [60, 0, 0, 0, 0]], dtype=torch.float32) selected = memory_selection._select_pose_fps(ids, poses, count=2) self.assertEqual(selected.tolist(), [0, 2]) def test_pose_fps_seed_views_make_first_selection_far_when_possible(self): import datasets.video.memory_selection as memory_selection all_poses = np.zeros((3, 5), dtype=np.float32) all_poses[1, 0] = 1.0 all_poses[2, 0] = 60.0 ids = torch.as_tensor([1, 2], dtype=torch.long) candidate_poses = torch.as_tensor(all_poses[ids.numpy()], dtype=torch.float32) selected = memory_selection._select_pose_fps(ids, candidate_poses, count=1, seed_ids=np.asarray([0]), all_poses=all_poses) self.assertEqual(selected.tolist(), [2]) def test_fov_greedy_dynamic_selection_avoids_excluded_revisit_frames(self): import datasets.video.memory_selection as memory_selection pool = _fov_pool([0, 1, 2], [[True, True], [True, False], [False, True]], gaps=[3, 2, 1]) selected = memory_selection._select_dynamic_multiview_from_fov_pool(pool, count=1, excluded=np.asarray([0])) self.assertNotIn(0, selected.tolist()) self.assertEqual(len(selected), 1) def test_fov_greedy_dynamic_selection_is_deterministic(self): import datasets.video.memory_selection as memory_selection pool = _fov_pool([0, 1, 2], [[True, False, False], [False, True, False], [False, False, True]], gaps=[3, 2, 1]) first = memory_selection._select_dynamic_multiview_from_fov_pool(pool, count=2) second = memory_selection._select_dynamic_multiview_from_fov_pool(pool, count=2) self.assertEqual(first.tolist(), second.tolist()) def test_fov_greedy_dynamic_selection_uses_absolute_temporal_gap_tie_break(self): import datasets.video.memory_selection as memory_selection pool = _fov_pool( [4, 10], [[True], [True]], fov_values=[1.0, 1.0], plucker=[0.0, 0.0], gaps=[1, -5], ) selected = memory_selection._select_dynamic_multiview_from_fov_pool(pool, count=1) self.assertEqual(selected.tolist(), [4]) def test_fov_greedy_dynamic_selection_prefers_residual_coverage(self): import datasets.video.memory_selection as memory_selection pool = _fov_pool( [0, 1, 2], [[True, False], [True, False], [False, True]], plucker=[0.0, 10.0, 0.0], gaps=[3, 1, 2], ) selected = memory_selection._select_dynamic_multiview_from_fov_pool( pool, count=1, excluded=np.asarray([0]), reference_frames=np.asarray([0]), ) self.assertEqual(selected.tolist(), [2]) def test_dynamic_multiview_provided_fov_pool_keeps_excluded_reference_coverage(self): import datasets.video.memory_selection as memory_selection poses = _poses(8) cfg = _selection_cfg(dynamic=_dynamic_cfg("multiview", multiview_selector="fov_greedy")) pool = _fov_pool( [0, 2, 3], [[True, False], [True, False], [False, True]], fov_values=[0.5, 1.0, 0.5], plucker=[0.0, 0.0, 0.0], gaps=[5, 3, 2], poses=poses, ) with mock.patch.object( memory_selection, "_build_fov_candidate_pool", side_effect=AssertionError("unexpected FOV pool build"), ): selected = memory_selection._select_dynamic_multiview( poses, np.asarray([7], dtype=np.int64), cfg, count=1, excluded=np.asarray([0], dtype=np.int64), reference_frames=np.asarray([0], dtype=np.int64), split="training", fov_pool=pool, ) self.assertEqual(selected.tolist(), [3]) def test_dynamic_multiview_fov_greedy_candidates_pool_exclusion_and_determinism(self): import datasets.video.memory_selection as memory_selection poses = _poses(8) default_cfg = _selection_cfg(dynamic=_dynamic_cfg("multiview", multiview_selector="fov_greedy")) noncausal_cfg = _selection_cfg(causal=False, dynamic=_dynamic_cfg("multiview", multiview_selector="fov_greedy")) cases = [ (default_cfg, "training", np.asarray([5], dtype=np.int64), [1, 2], [1, 2]), (noncausal_cfg, "training", np.asarray([3], dtype=np.int64), [4, 5, 6, 7], [6, 7]), (noncausal_cfg, "validation", np.asarray([5], dtype=np.int64), [1, 2], [1, 2]), (noncausal_cfg, "test", np.asarray([5], dtype=np.int64), [1, 2], [1, 2]), ] def fake_select(pool, count, **kwargs): return pool["candidates_t"].cpu().numpy()[-count:] for cfg, split, targets, expected_candidates, expected_selected in cases: seen_candidates = [] def fake_build(poses_arg, candidates, target_positions, cfg_arg, *, use_plucker): seen_candidates.append((candidates.copy(), use_plucker)) return _fov_pool(candidates, np.ones((len(candidates), 1), dtype=bool), poses=poses_arg) with self.subTest(split=split), mock.patch.object(memory_selection, "_build_fov_candidate_pool", side_effect=fake_build), mock.patch.object( memory_selection, "_select_dynamic_multiview_from_fov_pool", side_effect=fake_select ): selected = memory_selection._select_dynamic_multiview(poses, targets, cfg, count=2, split=split, min_candidate_frame=1) self.assertEqual(seen_candidates[0][0].tolist(), expected_candidates) self.assertTrue(seen_candidates[0][1]) self.assertEqual(selected.tolist(), expected_selected) provided_poses = _poses(6) pool = _fov_pool([0, 1, 2, 3], [[True, False], [False, True], [True, False], [False, True]], poses=provided_poses) kwargs = dict(excluded=np.asarray([0, 1], dtype=np.int64), reference_frames=np.asarray([0, 1], dtype=np.int64), split="training", fov_pool=pool) with mock.patch.object(memory_selection, "_build_fov_candidate_pool", side_effect=AssertionError("unexpected FOV pool build")): first = memory_selection._select_dynamic_multiview(provided_poses, np.asarray([5], dtype=np.int64), default_cfg, count=2, **kwargs) second = memory_selection._select_dynamic_multiview(provided_poses, np.asarray([5], dtype=np.int64), default_cfg, count=2, **kwargs) self.assertEqual(first.tolist(), [2]) self.assertEqual(second.tolist(), first.tolist()) self.assertTrue(set(first.tolist()).isdisjoint({0, 1})) def test_dynamic_multiview_pose_plucker_fps_uses_pose_path_without_fov_masks(self): import datasets.video.memory_selection as memory_selection poses = _poses(8) cfg = _selection_cfg( pose_similarity_threshold=0.0, training_use_plucker=False, pose_preselect_topk=3, dynamic=_dynamic_cfg("multiview", multiview_selector="pose_plucker_fps"), ) reference = np.asarray([0], dtype=np.int64) with mock.patch.object(memory_selection, "_candidate_fov_masks", side_effect=AssertionError("unexpected exact FOV masks")), mock.patch.object( memory_selection, "_build_fov_candidate_pool", side_effect=AssertionError("unexpected FOV pool build") ), mock.patch.object(memory_selection, "_rank_pose_plucker_candidates", wraps=memory_selection._rank_pose_plucker_candidates) as rank_mock, mock.patch.object( memory_selection, "_select_pose_fps", wraps=memory_selection._select_pose_fps ) as fps_mock: selected = memory_selection._select_dynamic_multiview( poses, np.asarray([6], dtype=np.int64), cfg, count=2, reference_frames=reference, split="training", ) self.assertTrue(rank_mock.called) self.assertTrue(fps_mock.called) self.assertEqual(fps_mock.call_args.kwargs["seed_ids"].tolist(), [0]) self.assertEqual(selected.tolist(), sorted(selected.tolist())) self.assertLessEqual(len(selected), 2) def test_training_default_causal_revisit_does_not_select_future(self): poses = _poses(8) poses[6] = poses[4] cfg = _selection_cfg( max_anchor_frames=0, max_dynamic_frames=0, max_revisit_frames=1, pose_similarity_threshold=0.0, ) indices, masks = select_memory_indices( poses, np.array([4]), cfg, split="training", rng=np.random.default_rng(0), ) selected = indices["revisit"][masks["revisit"]] self.assertTrue(len(selected) > 0) self.assertTrue(np.all(selected < 4)) def test_training_noncausal_revisit_can_select_future_pose_match(self): poses = np.zeros((8, 5), dtype=np.float32) poses[:, 0] = 100.0 poses[:, 4] = 180.0 poses[4] = np.asarray([0, 0, 0, 0, 0], dtype=np.float32) poses[6] = np.asarray([0, 0, 0, 0, 0], dtype=np.float32) cfg = _selection_cfg( causal=False, max_anchor_frames=0, max_dynamic_frames=0, max_revisit_frames=2, pose_similarity_threshold=0.95, ) indices, masks = select_memory_indices( poses, np.array([4]), cfg, split="training", rng=np.random.default_rng(0), ) selected = indices["revisit"][masks["revisit"]] self.assertIn(6, selected.tolist()) def test_training_noncausal_revisit_excludes_target_window_frame(self): poses = np.zeros((8, 5), dtype=np.float32) poses[:, 0] = 100.0 poses[:, 4] = 180.0 poses[4] = np.asarray([0, 0, 0, 0, 0], dtype=np.float32) poses[6] = np.asarray([0, 0, 0, 0, 0], dtype=np.float32) cfg = _selection_cfg( causal=False, max_anchor_frames=0, max_dynamic_frames=0, max_revisit_frames=1, pose_similarity_threshold=0.95, ) indices, masks = select_memory_indices( poses, np.array([4]), cfg, split="training", rng=np.random.default_rng(0), ) self.assertEqual(indices["revisit"].tolist(), [6]) self.assertEqual(masks["revisit"].tolist(), [True]) def test_validation_and_test_noncausal_revisit_stays_past_only(self): import datasets.video.memory_selection as memory_selection poses = _poses(8) cfg = _selection_cfg( causal=False, max_anchor_frames=0, max_dynamic_frames=0, max_revisit_frames=1, ) for split in ("validation", "test"): seen_candidates = [] def fake_point_union(poses_arg, candidates, target_positions, cfg_arg, count, **kwargs): seen_candidates.append(candidates.copy()) return candidates[-count:].astype(np.int64, copy=False) with self.subTest(split=split), mock.patch.object(memory_selection, "_select_by_point_union", side_effect=fake_point_union): indices, masks = select_memory_indices(poses, np.array([4]), cfg, split=split) self.assertEqual([candidates.tolist() for candidates in seen_candidates], [[0, 1]]) self.assertEqual(indices["revisit"].tolist(), [1]) self.assertEqual(masks["revisit"].tolist(), [True]) def test_revisit_can_duplicate_anchor_and_dynamic_when_scored_best(self): poses = np.zeros((8, 5), dtype=np.float32) poses[:, 0] = 100.0 poses[:, 4] = 180.0 poses[0] = np.asarray([0, 0, 0, 0, 0], dtype=np.float32) poses[5] = np.asarray([0, 0, 0, 0, 0], dtype=np.float32) poses[6] = np.asarray([0, 0, 0, 0, 0], dtype=np.float32) cfg = _selection_cfg( max_anchor_frames=1, max_dynamic_frames=1, max_revisit_frames=2, pose_similarity_threshold=0.95, local_context_exclusion_frames=0, ) indices, masks = select_memory_indices( poses, np.array([6]), cfg, split="training", rng=np.random.default_rng(0), ) self.assertEqual(indices["anchor"].tolist(), [0]) self.assertEqual(indices["dynamic"].tolist(), [5]) self.assertEqual(indices["revisit"].tolist(), [0, 5]) self.assertTrue(masks["revisit"].all()) def test_dememwm_selection_is_typed_and_causal(self): poses = _poses(12) indices, masks = select_memory_indices(poses, np.array([6, 7, 8]), _selection_cfg()) self.assertEqual(indices["anchor"].tolist(), [0, 1]) self.assertEqual(indices["dynamic"].tolist(), [2, 3]) for key in ("anchor", "dynamic", "revisit"): self.assertTrue(np.all(indices[key][masks[key]] < 6)) def test_recent_dynamic_excludes_anchor_inside_local_context(self): poses = _poses(8) cfg = _selection_cfg( max_anchor_frames=1, max_dynamic_frames=1, max_revisit_frames=0, dynamic=_dynamic_cfg("recent"), ) indices, masks = select_memory_indices( poses, np.array([6]), cfg, split="training", anchor_candidate_start=5, anchor_candidate_stop=6, ) self.assertEqual(indices["anchor"].tolist(), [5]) self.assertEqual(indices["dynamic"].tolist(), [3]) self.assertEqual(masks["anchor"].tolist(), [True]) self.assertEqual(masks["dynamic"].tolist(), [True]) self.assertTrue(np.all(indices["dynamic"][masks["dynamic"]] < 6)) def test_anchor_selection_is_context_bounded_farthest_point(self): poses = _poses(10) cfg = _selection_cfg( max_anchor_frames=2, max_dynamic_frames=0, max_revisit_frames=0, anchor_diverse_selection=True, ) indices, masks = select_memory_indices( poses, np.array([8]), cfg, anchor_candidate_start=2, anchor_candidate_stop=6, ) self.assertEqual(indices["anchor"].tolist(), [2, 5]) self.assertEqual(masks["anchor"].tolist(), [True, True]) def test_training_dropout_randomizes_memory_streams_with_replacement(self): poses = _poses(12) cfg = _selection_cfg( max_anchor_frames=2, max_dynamic_frames=3, max_revisit_frames=2, training_dropout=1.0, dynamic=_dynamic_cfg("recent"), ) indices, masks = select_memory_indices( poses, np.array([8]), cfg, split="training", rng=np.random.default_rng(0), anchor_candidate_start=2, anchor_candidate_stop=6, ) self.assertEqual(indices["anchor"].tolist(), [4, 3]) self.assertEqual(indices["dynamic"].tolist(), [1, 0, 0]) self.assertEqual(indices["revisit"].tolist(), [0, 1]) for key in ("anchor", "dynamic", "revisit"): self.assertTrue(masks[key].all()) def test_event_triggered_dynamic_causal_stream_stays_before_target(self): poses = _poses(12) cfg = _selection_cfg( max_anchor_frames=0, max_dynamic_frames=3, max_revisit_frames=0, dynamic=_dynamic_cfg("event_triggered"), ) indices, masks = select_memory_indices( poses, np.array([5]), cfg, split="training", dynamic_stream=np.asarray([1, 4, 5, 7, 9], dtype=np.int64), ) selected = indices["dynamic"][masks["dynamic"]] self.assertEqual(selected.tolist(), [1]) self.assertTrue(np.all(selected < 5)) def test_event_triggered_dynamic_noncausal_stream_can_select_future(self): poses = _poses(12) cfg = _selection_cfg( causal=False, max_anchor_frames=0, max_dynamic_frames=3, max_revisit_frames=0, dynamic=_dynamic_cfg("event_triggered"), ) indices, masks = select_memory_indices( poses, np.array([5]), cfg, split="training", dynamic_stream=np.asarray([1, 4, 5, 7, 9], dtype=np.int64), ) selected = indices["dynamic"][masks["dynamic"]] self.assertEqual(selected.tolist(), [1, 7, 9]) self.assertTrue(np.any(selected > 5)) def test_event_triggered_dynamic_eval_stream_stays_past_when_noncausal(self): poses = _poses(12) cfg = _selection_cfg( causal=False, max_anchor_frames=0, max_dynamic_frames=3, max_revisit_frames=0, dynamic=_dynamic_cfg("event_triggered"), ) for split in ("validation", "test"): with self.subTest(split=split): indices, masks = select_memory_indices( poses, np.array([5]), cfg, split=split, dynamic_stream=np.asarray([1, 4, 5, 7, 9], dtype=np.int64), ) selected = indices["dynamic"][masks["dynamic"]] self.assertEqual(selected.tolist(), [1]) self.assertTrue(np.all(selected < 5)) def test_event_triggered_dynamic_can_duplicate_anchor_and_revisit(self): import datasets.video.memory_selection as memory_selection poses = _poses(8) cfg = _selection_cfg( max_anchor_frames=1, max_dynamic_frames=2, max_revisit_frames=2, anchor_diverse_selection=False, dynamic=_dynamic_cfg("event_triggered"), ) with mock.patch.object(memory_selection, "_select_revisit", return_value=np.asarray([0, 5], dtype=np.int64)): indices, masks = select_memory_indices( poses, np.array([6]), cfg, split="training", dynamic_stream=np.asarray([0, 5], dtype=np.int64), ) self.assertEqual(indices["anchor"].tolist(), [0]) self.assertEqual(indices["revisit"].tolist(), [0, 5]) self.assertEqual(indices["dynamic"].tolist(), [0, -1]) self.assertEqual(masks["dynamic"].tolist(), [True, False]) def test_recent_dynamic_remains_past_only_when_memory_selection_is_noncausal(self): poses = _poses(10) cfg = _selection_cfg( causal=False, max_anchor_frames=0, max_dynamic_frames=2, max_revisit_frames=0, dynamic=_dynamic_cfg("recent"), ) indices, masks = select_memory_indices( poses, np.array([6]), cfg, split="training", dynamic_stream=np.asarray([6, 7, 8, 9], dtype=np.int64), ) self.assertEqual(indices["dynamic"].tolist(), [2, 3]) self.assertEqual(masks["dynamic"].tolist(), [True, True]) self.assertTrue(np.all(indices["dynamic"][masks["dynamic"]] < 6)) def test_event_triggered_dynamic_selects_first_stable_post_event_frame(self): poses = np.zeros((10, 5), dtype=np.float32) actions = np.zeros((10, 25), dtype=np.float32) cfg = _selection_cfg( max_anchor_frames=0, max_dynamic_frames=2, max_revisit_frames=0, dynamic=_dynamic_cfg("event_triggered", stable_frames=2, min_event_gap=0), ) indices, masks = select_memory_indices( poses, np.array([8]), cfg, latents=_event_latents(10, event_frame=3), actions=actions, ) self.assertEqual(indices["dynamic"].tolist(), [4, -1]) self.assertEqual(masks["dynamic"].tolist(), [True, False]) def test_event_triggered_dynamic_pads_without_event_or_stable_frame(self): poses = np.zeros((10, 5), dtype=np.float32) actions = np.zeros((10, 25), dtype=np.float32) cfg = _selection_cfg( max_anchor_frames=0, max_dynamic_frames=2, max_revisit_frames=0, dynamic=_dynamic_cfg("event_triggered", stable_frames=1, min_event_gap=0), ) cases = [ _event_latents(10, event_frame=99), _event_latents(10, event_frame=7), ] for latents in cases: with self.subTest(event_frame=int(np.argmax(latents[:, 1, 0, 0] > 0.0))): indices, masks = select_memory_indices( poses, np.array([8]), cfg, latents=latents, actions=actions, ) self.assertEqual(indices["dynamic"].tolist(), [-1, -1]) self.assertEqual(masks["dynamic"].tolist(), [False, False]) def test_dynamic_policy_validation_accepts_exact_supported_values(self): for policy in ("recent", "event_triggered", "multiview"): with self.subTest(policy=policy): cfg = _selection_cfg(dynamic=_dynamic_cfg(policy)) self.assertEqual(_dynamic_policy(cfg), policy) def test_unknown_dynamic_policies_are_rejected_without_aliases(self): poses = np.zeros((8, 5), dtype=np.float32) for policy in ("fov_multiview", "fov_fps", "hybrid"): with self.subTest(policy=policy): cfg = _selection_cfg( max_anchor_frames=0, max_dynamic_frames=2, max_revisit_frames=0, dynamic=_dynamic_cfg(policy), ) with self.assertRaisesRegex(ValueError, "recent, event_triggered, multiview"): select_memory_indices(poses, np.array([6]), cfg, latents=_event_latents(8, event_frame=99)) def test_multiview_selector_validation_accepts_exact_backends(self): for selector in ("fov_greedy", "pose_plucker_fps"): with self.subTest(selector=selector): cfg = _selection_cfg(dynamic=_dynamic_cfg("multiview", multiview_selector=selector)) self.assertEqual(_dynamic_multiview_selector(cfg), selector) def test_unknown_multiview_selector_is_rejected(self): poses = np.zeros((8, 5), dtype=np.float32) cfg = _selection_cfg( max_anchor_frames=0, max_dynamic_frames=2, max_revisit_frames=0, dynamic=_dynamic_cfg("multiview", multiview_selector="fov_fps"), ) with self.assertRaisesRegex(ValueError, "fov_greedy, pose_plucker_fps"): select_memory_indices(poses, np.array([6]), cfg) def test_benchmark_script_runs_synthetic_mode_without_dataset_paths(self): with tempfile.TemporaryDirectory() as tmpdir: report_path = Path(tmpdir) / "speed_report.md" result = subprocess.run( [ sys.executable, str(_benchmark_script_path()), "--num-frames", "32", "--target-start", "16", "--target-len", "2", "--num-iters", "1", "--pose-preselect-topk", "4", "--candidate-chunk-size", "4", "--write-report", str(report_path), ], cwd=Path(__file__).resolve().parents[1], capture_output=True, text=True, ) self.assertEqual(result.returncode, 0, msg=result.stderr) self.assertIn("fov_greedy", result.stdout) self.assertIn("pose_plucker_fps", result.stdout) self.assertIn("not a substitute for a real dataset sampling benchmark", report_path.read_text(encoding="utf-8")) def test_benchmark_script_rejects_unknown_selector(self): result = subprocess.run( [ sys.executable, str(_benchmark_script_path()), "--num-frames", "32", "--target-start", "16", "--target-len", "2", "--num-iters", "1", "--pose-preselect-topk", "4", "--candidate-chunk-size", "4", "--selectors", "fov_fps", ], cwd=Path(__file__).resolve().parents[1], capture_output=True, text=True, ) self.assertNotEqual(result.returncode, 0) self.assertIn("unknown selector", result.stderr) def test_multiview_policy_selection_uses_shared_fov_pool_and_segments(self): import datasets.video.memory_selection as memory_selection poses = _poses(10) cfg = _selection_cfg( max_anchor_frames=1, max_dynamic_frames=2, max_revisit_frames=1, local_context_exclusion_frames=2, fov_overlap_threshold=0.0, min_total_selected_coverage=0.0, dynamic=_dynamic_cfg("multiview", multiview_selector="fov_greedy"), ) pool = _fov_pool( [0, 1, 2, 3, 4, 5], [ [True, False, False], [False, True, False], [False, False, True], [True, True, False], [False, False, False], [True, True, True], ], gaps=[6, 5, 4, 3, 2, 1], poses=poses, ) dynamic_calls = [] def fake_dynamic( poses_arg, target_positions, cfg_arg, count, *, excluded=None, reference_frames=None, split="training", min_candidate_frame=0, fov_pool=None, ): dynamic_calls.append( { "target_positions": target_positions.copy(), "excluded": np.asarray(excluded, dtype=np.int64).copy(), "reference_frames": np.asarray(reference_frames, dtype=np.int64).copy(), "split": split, "fov_pool": fov_pool, } ) return np.asarray([1, 2], dtype=np.int64)[:count] with mock.patch.object(memory_selection, "_build_fov_candidate_pool", return_value=pool) as build_mock, mock.patch.object( memory_selection, "_select_dynamic_multiview", side_effect=fake_dynamic ): indices, masks = select_memory_indices(poses, np.array([6, 7]), cfg, split="validation") self.assertEqual(set(indices), {"anchor", "dynamic", "revisit"}) self.assertEqual({key: len(indices[key]) for key in ("anchor", "dynamic", "revisit")}, {"anchor": 1, "dynamic": 2, "revisit": 1}) self.assertEqual({key: len(masks[key]) for key in ("anchor", "dynamic", "revisit")}, {"anchor": 1, "dynamic": 2, "revisit": 1}) self.assertEqual(build_mock.call_count, 1) self.assertEqual(build_mock.call_args.args[2].tolist(), [6, 7]) selected_revisit = indices["revisit"][masks["revisit"]] self.assertTrue(np.all(selected_revisit < 4)) self.assertEqual(len(dynamic_calls), 1) self.assertIs(dynamic_calls[0]["fov_pool"], pool) self.assertEqual(dynamic_calls[0]["target_positions"].tolist(), [6, 7]) self.assertEqual(dynamic_calls[0]["excluded"].tolist(), selected_revisit.tolist()) self.assertEqual(dynamic_calls[0]["reference_frames"].tolist(), selected_revisit.tolist()) self.assertEqual(indices["dynamic"].tolist(), [1, 2]) self.assertEqual(masks["dynamic"].tolist(), [True, True]) def test_multiview_shared_fov_pool_preselects_after_revisit_local_exclusion(self): poses = np.zeros((8, 5), dtype=np.float32) poses[:, 0] = 1000.0 poses[:, 4] = 180.0 poses[0] = np.asarray([0, 0, 0, 0, 0], dtype=np.float32) poses[5] = np.asarray([0, 0, 0, 0, 0], dtype=np.float32) poses[6] = np.asarray([0, 0, 0, 0, 0], dtype=np.float32) cfg = _selection_cfg( max_anchor_frames=0, max_dynamic_frames=1, max_revisit_frames=1, local_context_exclusion_frames=2, pose_preselect_topk=1, fov_overlap_threshold=0.0, min_total_selected_coverage=0.0, dynamic=_dynamic_cfg("multiview", multiview_selector="fov_greedy"), ) indices, masks = select_memory_indices(poses, np.asarray([6], dtype=np.int64), cfg, split="validation") self.assertEqual(indices["revisit"].tolist(), [0]) self.assertEqual(masks["revisit"].tolist(), [True]) def test_multiview_policy_dispatcher_calls_selector_with_target_positions(self): import datasets.video.memory_selection as memory_selection poses = _poses(8) cfg = _selection_cfg( max_anchor_frames=0, max_dynamic_frames=1, max_revisit_frames=0, dynamic=_dynamic_cfg("multiview", multiview_selector="fov_greedy"), ) target_positions = np.asarray([6, 7], dtype=np.int64) with mock.patch.object(memory_selection, "_select_dynamic_multiview", return_value=np.asarray([2], dtype=np.int64)) as select_mock: selected = memory_selection._select_dynamic_by_policy( 6, 1, cfg, poses=poses, target_positions=target_positions, split="validation", ) self.assertEqual(selected.tolist(), [2]) self.assertEqual(select_mock.call_args.args[1].tolist(), [6, 7]) self.assertEqual(select_mock.call_args.kwargs["split"], "validation") def test_event_dynamic_uses_nearest_anchors_per_revisit_frame(self): import datasets.video.memory_selection as memory_selection poses = np.zeros((100, 5), dtype=np.float32) actions = np.zeros((100, 25), dtype=np.float32) cfg = _selection_cfg( max_anchor_frames=0, max_dynamic_frames=4, max_revisit_frames=2, dynamic=_dynamic_cfg("event_triggered", stable_frames=1, min_event_gap=0), ) with mock.patch.object(memory_selection, "_select_revisit", return_value=np.asarray([35, 85], dtype=np.int64)): indices, masks = select_memory_indices( poses, np.array([99]), cfg, latents=_multi_event_latents(100, [10, 20, 30, 40, 70, 80, 90]), actions=actions, ) self.assertEqual(indices["revisit"].tolist(), [35, 85]) self.assertEqual(indices["dynamic"].tolist(), [31, 41, 81, 91]) self.assertEqual(masks["dynamic"].tolist(), [True, True, True, True]) def test_event_dynamic_uses_all_slots_near_single_revisit_frame(self): import datasets.video.memory_selection as memory_selection poses = np.zeros((100, 5), dtype=np.float32) actions = np.zeros((100, 25), dtype=np.float32) cfg = _selection_cfg( max_anchor_frames=0, max_dynamic_frames=4, max_revisit_frames=2, dynamic=_dynamic_cfg("event_triggered", stable_frames=1, min_event_gap=0), ) with mock.patch.object(memory_selection, "_select_revisit", return_value=np.asarray([35], dtype=np.int64)): indices, masks = select_memory_indices( poses, np.array([99]), cfg, latents=_multi_event_latents(100, [10, 20, 30, 40, 70, 80, 90]), actions=actions, ) self.assertEqual(indices["revisit"].tolist(), [35, -1]) self.assertEqual(indices["dynamic"].tolist(), [11, 21, 31, 41]) self.assertEqual(masks["dynamic"].tolist(), [True, True, True, True]) def test_event_dynamic_falls_back_to_target_nearest_without_revisit(self): import datasets.video.memory_selection as memory_selection poses = np.zeros((100, 5), dtype=np.float32) actions = np.zeros((100, 25), dtype=np.float32) cfg = _selection_cfg( max_anchor_frames=0, max_dynamic_frames=4, max_revisit_frames=2, dynamic=_dynamic_cfg("event_triggered", stable_frames=1, min_event_gap=0), ) with mock.patch.object(memory_selection, "_select_revisit", return_value=np.empty((0,), dtype=np.int64)): indices, masks = select_memory_indices( poses, np.array([99]), cfg, latents=_multi_event_latents(100, [10, 20, 30, 40, 70, 80, 90]), actions=actions, ) self.assertEqual(indices["revisit"].tolist(), [-1, -1]) self.assertEqual(indices["dynamic"].tolist(), [41, 71, 81, 91]) self.assertEqual(masks["dynamic"].tolist(), [True, True, True, True]) class DeMemWMLatentDatasetTests(unittest.TestCase): def test_dataset_requires_split_directory_without_root_fallback(self): with tempfile.TemporaryDirectory() as tmp: root = Path(tmp) (root / "sample.mp4").write_bytes(b"") with self.assertRaisesRegex(FileNotFoundError, "validation"): MinecraftVideoDeMemWMLatentDataset(_dataset_cfg(root), split="validation") def test_dataset_matches_original_w_updown_split_filtering(self): with tempfile.TemporaryDirectory() as tmp: root = Path(tmp) _write_vae_feature_clip(root, root / "vae_features", split="validation", stem="plain") _write_vae_feature_clip(root, root / "vae_features", split="validation", stem="clip_w_updown") dataset = MinecraftVideoDeMemWMLatentDataset(_dataset_cfg(root), split="validation") dataset_wo = MinecraftVideoDeMemWMLatentDataset(_dataset_cfg(root, wo_updown=True), split="validation") self.assertEqual([path.name for path in dataset.data_paths], ["clip_w_updown.mp4"]) self.assertEqual([path.name for path in dataset_wo.data_paths], ["plain.mp4"]) def test_dataset_filters_nested_validation_files_by_selected_split(self): with tempfile.TemporaryDirectory() as tmp: root = Path(tmp) _write_vae_feature_clip(root, root / "vae_features", split="validation", subdir="nested", stem="plain") _write_vae_feature_clip(root, root / "vae_features", split="validation", subdir="nested", stem="clip_w_updown") dataset = MinecraftVideoDeMemWMLatentDataset(_dataset_cfg(root), split="validation") dataset_wo = MinecraftVideoDeMemWMLatentDataset(_dataset_cfg(root, wo_updown=True), split="validation") self.assertEqual([path.name for path in dataset.data_paths], ["clip_w_updown.mp4"]) self.assertEqual([path.name for path in dataset_wo.data_paths], ["plain.mp4"]) def test_dataset_matches_original_fallback_when_split_filter_has_no_matches(self): cases = [ ("validation", False, "plain"), ("test", True, "clip_w_updown"), ] for split, wo_updown, stem in cases: with self.subTest(split=split, wo_updown=wo_updown), tempfile.TemporaryDirectory() as tmp: root = Path(tmp) _write_vae_feature_clip(root, root / "vae_features", split=split, subdir="nested", stem=stem) dataset = MinecraftVideoDeMemWMLatentDataset(_dataset_cfg(root, wo_updown=wo_updown), split=split) self.assertEqual([path.name for path in dataset.data_paths], [f"{stem}.mp4"]) def test_dataset_rejects_files_shorter_than_original_frame_skip_window(self): with tempfile.TemporaryDirectory() as tmp: root = Path(tmp) _write_vae_feature_clip(root, root / "vae_features", stem="short", num_frames=102) with self.assertRaisesRegex(ValueError, "requires at least 103"): MinecraftVideoDeMemWMLatentDataset(_dataset_cfg(root), split="training") def test_dataset_applies_original_pose_height_sanity_check(self): with tempfile.TemporaryDirectory() as tmp: root = Path(tmp) poses = _poses(104) poses[:, 1] = np.linspace(0.0, 3.0, 104, dtype=np.float32) _write_vae_feature_clip(root, root / "vae_features", stem="sample", poses=poses) dataset = MinecraftVideoDeMemWMLatentDataset(_dataset_cfg(root), split="training") with self.assertRaisesRegex(RuntimeError, "Pose height variation") as captured: dataset[0] self.assertIsInstance(captured.exception.__cause__, ValueError) def test_dataset_converts_raw_actions_to_original_25d_contract(self): with tempfile.TemporaryDirectory() as tmp: root = Path(tmp) raw_actions = np.zeros((104, 8), dtype=np.int64) raw_actions[:, 0] = 1 raw_actions[:, 3] = 11 raw_actions[:, 4] = 13 raw_actions[:, 7] = 1 _write_vae_feature_clip(root, root / "vae_features", stem="sample", actions=raw_actions) sample = MinecraftVideoDeMemWMLatentDataset(_dataset_cfg(root), split="training")[0] self.assertEqual(tuple(sample["actions"].shape), (3, 25)) self.assertTrue(torch.equal(sample["actions"][:, 11], torch.ones(3))) self.assertTrue(torch.equal(sample["actions"][:, 16], torch.ones(3))) self.assertTrue(torch.equal(sample["actions"][:, 15], -torch.ones(3))) self.assertTrue(torch.equal(sample["actions"][:, 2], torch.ones(3))) def test_dataset_reads_dememwm_vae_feature_layout(self): with tempfile.TemporaryDirectory() as tmp: root = Path(tmp) / "minecraft" feature_root = root / "vae_features" raw_actions = np.zeros((104, 8), dtype=np.int64) raw_actions[:, 0] = 1 raw_actions[:, 3] = 11 raw_actions[:, 4] = 13 raw_actions[:, 7] = 1 _write_vae_feature_clip(root, feature_root, actions=raw_actions) dataset = MinecraftVideoDeMemWMLatentDataset(_dataset_cfg(root), split="training") sample = dataset[0] self.assertEqual([path.suffix for path in dataset.data_paths], [".mp4"]) self.assertEqual(tuple(sample["latents"].shape), (3, 1, 1, 2)) self.assertEqual(tuple(sample["actions"].shape), (3, 25)) self.assertEqual(sample["image_hw"].tolist(), [360, 640]) self.assertTrue(torch.equal(sample["actions"][:, 11], torch.ones(3))) self.assertTrue(torch.equal(sample["actions"][:, 16], torch.ones(3))) self.assertTrue(torch.equal(sample["actions"][:, 15], -torch.ones(3))) self.assertTrue(torch.equal(sample["actions"][:, 2], torch.ones(3))) def test_dataset_uses_original_training_clip_window(self): with tempfile.TemporaryDirectory() as tmp: root = Path(tmp) _write_vae_feature_clip(root, root / "vae_features", stem="sample", num_frames=1500) dataset = MinecraftVideoDeMemWMLatentDataset(_dataset_cfg(root, n_frames=8), split="training") self.assertEqual(len(dataset), 1300 - 8 + 1) def test_dataset_reuses_cached_feature_arrays_for_same_file(self): with tempfile.TemporaryDirectory() as tmp: root = Path(tmp) _write_vae_feature_clip(root, root / "vae_features", stem="sample") dataset = MinecraftVideoDeMemWMLatentDataset(_dataset_cfg(root), split="training") first = dataset[0] with mock.patch("datasets.video.minecraft_video_dememwm_latent_dataset.np.load", side_effect=AssertionError("cache miss")): second = dataset[1] self.assertEqual(first["frame_indices"][:3].tolist(), [100, 101, 102]) self.assertEqual(second["frame_indices"][:3].tolist(), [101, 102, 103]) def test_preprocess_splits_sequence_tensors_from_memory_segments(self): from algorithms.dememwm.df_video import _preprocess_dememwm_latent_batch batch = { "latents": torch.arange(6, dtype=torch.float32).view(1, 6, 1, 1, 1), "actions": (100 + torch.arange(6, dtype=torch.float32)).view(1, 6, 1), "poses": (200 + torch.arange(6, dtype=torch.float32)).view(1, 6, 1).repeat(1, 1, 5), "frame_indices": (10 + torch.arange(6, dtype=torch.long)).view(1, 6), "memory_segments": { "target": torch.tensor([2]), "anchor": torch.tensor([1]), "dynamic": torch.tensor([3]), "revisit": torch.tensor([0]), }, "memory_masks": { "target": torch.tensor([[True, False]]), "anchor": torch.tensor([[True]]), "dynamic": torch.tensor([[True, False, True]]), "revisit": torch.zeros((1, 0), dtype=torch.bool), }, "image_hw": torch.tensor([[360, 640]], dtype=torch.long), } preprocessed = _preprocess_dememwm_latent_batch(batch) self.assertEqual(preprocessed["target_length"], 2) self.assertEqual(preprocessed["stream_lengths"], {"anchor": 1, "dynamic": 3, "revisit": 0}) self.assertEqual((preprocessed["target_slice"].start, preprocessed["target_slice"].stop), (0, 2)) self.assertEqual( {key: (slc.start, slc.stop) for key, slc in preprocessed["stream_slices"].items()}, {"anchor": (2, 3), "dynamic": (3, 6), "revisit": (6, 6)}, ) self.assertEqual(preprocessed["target_tensors"]["latents"][:, 0, 0, 0, 0].tolist(), [0.0, 1.0]) self.assertEqual(preprocessed["stream_tensors"]["dynamic"]["latents"][:, 0, 0, 0, 0].tolist(), [3.0, 4.0, 5.0]) self.assertEqual(preprocessed["stream_tensors"]["dynamic"]["actions"][:, 0, 0].tolist(), [103.0, 104.0, 105.0]) self.assertEqual(preprocessed["action_conditions"][:, 0, 0].tolist(), [0.0, 101.0, 0.0, 0.0, 0.0, 0.0]) self.assertEqual(preprocessed["target_tensors"]["action_conditions"][:, 0, 0].tolist(), [0.0, 101.0]) self.assertEqual(preprocessed["stream_tensors"]["dynamic"]["action_conditions"][:, 0, 0].tolist(), [0.0, 0.0, 0.0]) self.assertEqual(preprocessed["stream_tensors"]["dynamic"]["poses"][:, 0, 0].tolist(), [203.0, 204.0, 205.0]) self.assertEqual(preprocessed["frame_memory_pose"][:, 0, 0].tolist(), [200.0, 201.0, 202.0, 203.0, 204.0, 205.0]) self.assertIs(preprocessed["frame_memory_pose"], preprocessed["poses"]) self.assertEqual(preprocessed["stream_tensors"]["dynamic"]["frame_indices"][:, 0].tolist(), [13, 14, 15]) self.assertIs(preprocessed["memory_masks"]["dynamic"], batch["memory_masks"]["dynamic"]) self.assertEqual(tuple(preprocessed["memory_masks"]["dynamic"].shape), (1, 3)) self.assertEqual(preprocessed["memory_masks"]["dynamic"].device.type, "cpu") self.assertEqual(preprocessed["image_hw"].tolist(), [[360, 640]]) self.assertEqual(preprocessed["image_hw"].dtype, torch.long) def test_preprocess_expands_shared_image_hw_metadata(self): from algorithms.dememwm.df_video import _preprocess_dememwm_latent_batch batch = { "latents": torch.zeros((2, 1, 1, 1, 1), dtype=torch.float32), "actions": torch.zeros((2, 1, 1), dtype=torch.float32), "poses": torch.arange(10, dtype=torch.float32).view(2, 1, 5), "frame_indices": torch.zeros((2, 1), dtype=torch.long), "memory_segments": { "target": torch.tensor([1, 1]), "anchor": torch.tensor([0, 0]), "dynamic": torch.tensor([0, 0]), "revisit": torch.tensor([0, 0]), }, "memory_masks": { "target": torch.ones((2, 1), dtype=torch.bool), "anchor": torch.zeros((2, 0), dtype=torch.bool), "dynamic": torch.zeros((2, 0), dtype=torch.bool), "revisit": torch.zeros((2, 0), dtype=torch.bool), }, "image_hw": torch.tensor([720, 1280], dtype=torch.long), } preprocessed = _preprocess_dememwm_latent_batch(batch) self.assertEqual(tuple(preprocessed["frame_memory_pose"].shape), (1, 2, 5)) self.assertEqual(preprocessed["frame_memory_pose"][:, :, 0].tolist(), [[0.0, 5.0]]) self.assertEqual(tuple(preprocessed["image_hw"].shape), (2, 2)) self.assertEqual(preprocessed["image_hw"].tolist(), [[720, 1280], [720, 1280]]) def test_training_step_uses_latent_dict_target_loss_and_memory_metadata(self): from algorithms.dememwm.df_video import DeMemWMMinecraft, _preprocess_dememwm_latent_batch class FakeDiffusion: stabilization_level = 7 def __call__(self, x, action_cond, pose_cond, **kwargs): self.x = x self.action_cond = action_cond self.pose_cond = pose_cond self.kwargs = kwargs loss = (torch.arange(x.shape[0], device=x.device, dtype=x.dtype) + 2).view(-1, 1, 1, 1, 1) return x, loss.expand_as(x).clone() class Harness: def __init__(self): self.diffusion_model = FakeDiffusion() self.logged = [] def _preprocess_batch(self, batch): return _preprocess_dememwm_latent_batch(batch) def _generate_noise_levels(self, xs): self.noise_input_shape = tuple(xs.shape) return torch.arange(xs.shape[0], device=xs.device, dtype=torch.long).view(-1, 1) def encode(self, xs): raise AssertionError("latent dict training must not VAE-encode inputs") def reweight_loss(self, loss, weight=None): raise AssertionError("target mask path should reduce the latent dict loss directly") def log(self, name, value, **kwargs): self.logged.append((name, value, kwargs)) batch = { "latents": torch.arange(5, dtype=torch.float32).view(1, 5, 1, 1, 1), "actions": (100 + torch.arange(5, dtype=torch.float32)).view(1, 5, 1), "poses": (200 + torch.arange(5, dtype=torch.float32)).view(1, 5, 1).repeat(1, 1, 5), "frame_indices": (10 + torch.arange(5, dtype=torch.long)).view(1, 5), "memory_segments": { "target": torch.tensor([2]), "anchor": torch.tensor([1]), "dynamic": torch.tensor([1]), "revisit": torch.tensor([1]), }, "memory_masks": { "target": torch.tensor([[True, False]]), "anchor": torch.tensor([[True]]), "dynamic": torch.tensor([[False]]), "revisit": torch.tensor([[True]]), }, "image_hw": torch.tensor([[360, 640]], dtype=torch.long), } harness = Harness() output = DeMemWMMinecraft.training_step(harness, batch, 0) call = harness.diffusion_model self.assertEqual(harness.noise_input_shape, (2, 1, 1, 1, 1)) self.assertTrue(torch.equal(call.x[:, 0, 0, 0, 0], torch.arange(5, dtype=torch.float32))) self.assertEqual(call.action_cond[:, 0, 0].tolist(), [0.0, 101.0, 0.0, 0.0, 0.0]) self.assertIsNone(call.pose_cond) self.assertEqual(call.kwargs["reference_length"], 0) self.assertEqual(call.kwargs["frame_memory_segments"], {"target": 2, "anchor": 1, "dynamic": 1, "revisit": 1}) self.assertEqual(call.kwargs["noise_levels"][:, 0].tolist(), [0, 1, 7, 7, 7]) self.assertEqual(call.kwargs["frame_idx"][:, 0].tolist(), [10, 11, 12, 13, 14]) self.assertEqual(call.kwargs["frame_memory_pose"][:, 0, 0].tolist(), [200.0, 201.0, 202.0, 203.0, 204.0]) self.assertEqual(call.kwargs["image_hw"].tolist(), [[360, 640]]) self.assertEqual(call.kwargs["frame_memory_masks"]["target"].device, call.x.device) self.assertEqual(call.kwargs["frame_memory_masks"]["dynamic"].tolist(), [[False]]) self.assertEqual(float(output["loss"]), 2.0) self.assertEqual(harness.logged[0][0], "training/loss") self.assertEqual(float(harness.logged[0][1]), 2.0) def test_training_step_rejects_raw_tuple_batches(self): from algorithms.dememwm.df_video import DeMemWMMinecraft class Harness: def _preprocess_batch(self, batch): raise AssertionError("raw batches must be rejected before preprocessing") with self.assertRaisesRegex(TypeError, "latent dict batch contract"): DeMemWMMinecraft.training_step(Harness(), (None, None, None, None), 0) def test_memory_noise_helpers_route_and_keep_validation_clean(self): from algorithms.dememwm.df_video import _apply_memory_route_masks, _memory_noise_levels_for_streams class Diffusion: stabilization_level, timesteps, sampling_timesteps = 7, 10, 4 cfg = OmegaConf.create({ "memory_noise": {"enabled": True, "anchor_max_fraction": 1.0, "dynamic_max_fraction": 0.5, "revisit_max_fraction": 0.25}, "noise_route": {"anchor": "high", "dynamic": "low", "revisit": "all"}, }) query_noise = torch.tensor([[9, 1], [7, 1]], dtype=torch.long) stream_lengths = {"anchor": 1, "dynamic": 1, "revisit": 1} torch.manual_seed(0) levels = _memory_noise_levels_for_streams(cfg, Diffusion, query_noise, stream_lengths, mode="training") self.assertTrue(bool((levels["anchor"] <= torch.tensor([[7, 1]])).all())) self.assertTrue(bool((levels["dynamic"] <= torch.tensor([[3, 0]])).all())) self.assertTrue(bool((levels["revisit"] <= torch.tensor([[1, 0]])).all())) validation_levels = _memory_noise_levels_for_streams(cfg, Diffusion, query_noise, stream_lengths, mode="validation") self.assertEqual(sum(int(v.sum().item()) for v in validation_levels.values()), 0) masks = {"target": torch.ones((2, 2), dtype=torch.bool), **{key: torch.ones((2, 1), dtype=torch.bool) for key in stream_lengths}} routed = _apply_memory_route_masks(masks, query_noise, cfg, Diffusion, mode="training") self.assertEqual({key: routed[key].tolist() for key in stream_lengths}, {"anchor": [[True], [False]], "dynamic": [[False], [True]], "revisit": [[True], [True]]}) def test_memory_noise_independent_mode_ignores_query_bounds(self): from algorithms.dememwm.df_video import _memory_noise_levels_for_streams class Diffusion: stabilization_level, timesteps, sampling_timesteps = 7, 10, 4 cfg = OmegaConf.create({ "memory_noise": { "enabled": True, "mode": "independent", }, }) query_noise = torch.ones((2, 2), dtype=torch.long) stream_lengths = {"anchor": 3, "dynamic": 0, "revisit": 0} def fake_randint(low, high, size, device=None, dtype=None): self.assertEqual((low, high, size), (0, 10, (3, 2))) return torch.full(size, high - 1, device=device, dtype=dtype) with mock.patch("algorithms.dememwm.df_video.torch.randint", side_effect=fake_randint): levels = _memory_noise_levels_for_streams(cfg, Diffusion, query_noise, stream_lengths, mode="training") self.assertEqual(levels["anchor"].tolist(), [[9, 9], [9, 9], [9, 9]]) self.assertTrue(bool((levels["anchor"] > query_noise.max()).all())) validation_levels = _memory_noise_levels_for_streams(cfg, Diffusion, query_noise, stream_lengths, mode="validation") self.assertEqual(sum(int(v.sum().item()) for v in validation_levels.values()), 0) def test_trainability_controls_memory_groups_full_dit_ramp_and_vae_freeze(self): from algorithms.dememwm.df_video import ( _apply_dememwm_trainability, _apply_dememwm_optimizer_group_lrs, _dememwm_optimizer_parameter_groups, _dememwm_optimizer_parameters, _trainability_cfg, ) from algorithms.dememwm.df_base import DiffusionForcingBase class TinyReferenceAttention(nn.Module): def __init__(self): super().__init__() self.to_q = nn.Linear(2, 2, bias=False) self.query_pose_proj = nn.Linear(6, 2, bias=False) self.key_pose_proj = nn.Linear(6, 2, bias=False) class TinyBlock(nn.Module): def __init__(self): super().__init__() self.s_mlp = nn.Linear(2, 2, bias=False) self.r_adaLN_modulation = nn.Sequential(nn.SiLU(), nn.Linear(2, 2, bias=False)) self.r_mlp = nn.Linear(2, 2, bias=False) self.r_attn_anchor = TinyReferenceAttention() class TinyDiffusion(nn.Module): def __init__(self): super().__init__() self.model = nn.Module() self.model.blocks = nn.ModuleList([TinyBlock()]) self.model.final_layer = nn.Linear(2, 2, bias=False) diffusion = TinyDiffusion() vae = nn.Linear(2, 2, bias=False) cfg = OmegaConf.create({ "trainability": { "freeze_vae": True, "train_full_dit": True, "full_dit_start_step": 10, "geometry_projections": False, "lr": {"memory_modules": 8.0e-5, "base_dit": 2.0e-5}, } }) trainability = _trainability_cfg(cfg) _apply_dememwm_trainability(diffusion, vae, trainability, global_step=0) flags = {name: param.requires_grad for name, param in diffusion.named_parameters()} self.assertTrue(flags["model.blocks.0.r_attn_anchor.to_q.weight"]) self.assertFalse(flags["model.blocks.0.r_attn_anchor.query_pose_proj.weight"]) self.assertFalse(flags["model.blocks.0.r_attn_anchor.key_pose_proj.weight"]) self.assertTrue(flags["model.blocks.0.r_adaLN_modulation.1.weight"]) self.assertTrue(flags["model.blocks.0.r_mlp.weight"]) self.assertFalse(flags["model.blocks.0.s_mlp.weight"]) self.assertFalse(flags["model.final_layer.weight"]) self.assertFalse(any(param.requires_grad for param in vae.parameters())) opt_ids = {id(param) for param in _dememwm_optimizer_parameters(diffusion, vae, trainability)} self.assertEqual(opt_ids, {id(param) for param in diffusion.parameters()}) param_groups = _dememwm_optimizer_parameter_groups(diffusion, vae, trainability, default_lr=2.0e-5, global_step=0) lr_by_name = {group["name"]: group["target_lr"] for group in param_groups} warmup_start_by_name = {group["name"]: group.get("warmup_start_step", 0) for group in param_groups} self.assertEqual(set(lr_by_name), {"memory_modules", "base_dit"}) self.assertAlmostEqual(lr_by_name["memory_modules"], 8.0e-5) self.assertEqual(lr_by_name["base_dit"], 0.0) self.assertEqual(warmup_start_by_name["memory_modules"], 0) self.assertEqual(warmup_start_by_name["base_dit"], 10) optimizer = torch.optim.AdamW(param_groups) lr_owner = object.__new__(DiffusionForcingBase) lr_owner.cfg = OmegaConf.create({"lr": 2.0e-5, "warmup_steps": 10}) lr_owner._apply_optimizer_lrs(optimizer, step=0) lr_by_name = {group["name"]: group["lr"] for group in optimizer.param_groups} self.assertAlmostEqual(lr_by_name["memory_modules"], 8.0e-6) self.assertEqual(lr_by_name["base_dit"], 0.0) _apply_dememwm_optimizer_group_lrs(optimizer, trainability, default_lr=2.0e-5, global_step=10) lr_owner._apply_optimizer_lrs(optimizer, step=10) lr_by_name = {group["name"]: group["lr"] for group in optimizer.param_groups} self.assertAlmostEqual(lr_by_name["memory_modules"], 8.0e-5) self.assertAlmostEqual(lr_by_name["base_dit"], 2.0e-6) _apply_dememwm_optimizer_group_lrs(optimizer, trainability, default_lr=2.0e-5, global_step=19) lr_owner._apply_optimizer_lrs(optimizer, step=19) lr_by_name = {group["name"]: group["lr"] for group in optimizer.param_groups} self.assertAlmostEqual(lr_by_name["memory_modules"], 8.0e-5) self.assertAlmostEqual(lr_by_name["base_dit"], 2.0e-5) _apply_dememwm_trainability(diffusion, vae, trainability, global_step=10) self.assertTrue(all(param.requires_grad for param in diffusion.parameters())) self.assertFalse(any(param.requires_grad for param in vae.parameters())) memory_only = OmegaConf.create({"trainability": {"freeze_vae": True}}) _apply_dememwm_trainability(diffusion, vae, _trainability_cfg(memory_only), global_step=0) default_flags = {name: param.requires_grad for name, param in diffusion.named_parameters()} self.assertTrue(default_flags["model.blocks.0.r_attn_anchor.query_pose_proj.weight"]) self.assertFalse(default_flags["model.final_layer.weight"]) def test_validation_step_builds_online_memory_from_committed_latents(self): import algorithms.dememwm.df_video as df_video from algorithms.dememwm.df_video import DeMemWMMinecraft, _preprocess_dememwm_latent_batch revisit_calls = [] def fake_select_revisit(poses, target_positions, cfg, count, excluded, split, rng=None, min_candidate_frame=0): start = int(target_positions[0]) excluded = np.asarray(excluded, dtype=np.int64) revisit_calls.append((target_positions.tolist(), split, float(cfg.fov_overlap_threshold), excluded.tolist())) if count <= 0 or start <= 2: return np.empty((0,), dtype=np.int64) for candidate in range(max(0, start - 2), -1, -1): if candidate not in excluded: return np.asarray([candidate], dtype=np.int64) return np.empty((0,), dtype=np.int64) class FakeDiffusion: def __init__(self): self.calls = [] def sample_step(self, x, action_cond, pose_cond, curr_noise_level, next_noise_level, **kwargs): self.calls.append( { "x": x.detach().clone(), "action_cond": action_cond.detach().clone(), "pose_cond": pose_cond, "curr_noise_level": curr_noise_level.detach().clone(), "next_noise_level": next_noise_level.detach().clone(), "kwargs": kwargs, } ) return torch.full_like(x[: kwargs["frame_memory_segments"]["target"]], float(len(self.calls))) class Harness: def __init__(self): self.memory_condition_length, self.clip_noise = 3, 0.0 self.context_frames, self.frame_stack, self.chunk_size, self.n_tokens = 2, 1, 2, 2 self.cfg = OmegaConf.create( { "memory_selection": {"enabled": True, "max_anchor_frames": 1, "max_dynamic_frames": 1, "max_revisit_frames": 1, "fov_overlap_threshold": 0.75}, "memory_noise": {"enabled": True, "validation_noisy_memory": False}, } ) self.diffusion_model = FakeDiffusion() self.logged = [] self.horizons = [] self.decoded_inputs = [] self.metric_updates = [] self.logger = type("Logger", (), {"experiment": object()})() self.log_video = True self.log_memory_selection_sheet = False self.save_local = False self.local_save_dir = None def _preprocess_batch(self, batch): return _preprocess_dememwm_latent_batch(batch) def _generate_scheduling_matrix(self, horizon): self.horizons.append(horizon) return np.stack([np.full(horizon, level, dtype=np.int64) for level in (2, 1, 0)]) def encode(self, xs): raise AssertionError("latent dict validation must not VAE-encode inputs") def decode(self, xs): self.decoded_inputs.append(xs.detach().clone()) return xs + 10.0 def _update_metric_accumulators(self, xs_pred, xs_gt, valid_mask=None, eval_start=0): self.metric_updates.append((xs_pred.detach().clone(), xs_gt.detach().clone())) def _update_per_frame_metric_accumulators(self, xs_pred, xs_gt, mask, eval_start): pass def _update_revisit_metric_accumulators(self, xs_pred, xs_gt, poses, frame_indices, mask, eval_start, batch_idx, namespace): pass def log(self, name, value): self.logged.append((name, value)) batch = { "latents": torch.tensor([0, 1, 2, 3, 4, 100, 101, 102], dtype=torch.float32).view(1, 8, 1, 1, 1), "actions": (100 + torch.arange(8, dtype=torch.float32)).view(1, 8, 1), "poses": (200 + torch.arange(8, dtype=torch.float32)).view(1, 8, 1).repeat(1, 1, 5), "frame_indices": torch.tensor([[10, 11, 12, 13, 14, 900, 901, 902]], dtype=torch.long), "memory_segments": {"target": torch.tensor([5]), "anchor": torch.tensor([1]), "dynamic": torch.tensor([1]), "revisit": torch.tensor([1])}, "memory_masks": {"target": torch.ones((1, 5), dtype=torch.bool), "anchor": torch.ones((1, 1), dtype=torch.bool), "dynamic": torch.ones((1, 1), dtype=torch.bool), "revisit": torch.ones((1, 1), dtype=torch.bool)}, "image_hw": torch.tensor([[360, 640]], dtype=torch.long), } harness = Harness() video_calls = [] anchor_calls = [] original_select_anchor = df_video._select_online_anchor_indices def fake_log_video(*args, **kwargs): video_calls.append((args, kwargs)) def fake_select_anchor(n_context_frames, count, cfg, poses=None): anchor_calls.append((int(n_context_frames), int(count), tuple(poses.shape) if poses is not None else None)) return original_select_anchor(n_context_frames, count, cfg, poses=poses) with mock.patch.object(df_video, "_select_online_anchor_indices", side_effect=fake_select_anchor), mock.patch.object( df_video, "_select_revisit", side_effect=fake_select_revisit ), mock.patch.object(df_video, "log_video", side_effect=fake_log_video): loss = DeMemWMMinecraft.validation_step(harness, batch, 0, namespace="test") calls = harness.diffusion_model.calls required_kwargs = {"reference_length", "frame_memory_segments", "frame_memory_masks", "frame_memory_pose", "frame_idx", "image_hw"} self.assertEqual(harness.horizons, [2, 1]) self.assertEqual(anchor_calls, [(2, 1, (2, 5))]) self.assertEqual(revisit_calls, [([2, 3], "test", 0.75, []), ([4], "test", 0.75, [])]) self.assertEqual(len(revisit_calls), len(harness.horizons)) self.assertEqual(len(calls), 4) self.assertTrue(all(required_kwargs <= call["kwargs"].keys() for call in calls)) self.assertTrue(all(call["kwargs"]["reference_length"] == 0 for call in calls)) self.assertEqual( ([call["kwargs"]["current_frame"] for call in calls[::2]], [call["kwargs"]["frame_memory_segments"]["target"] for call in calls[::2]]), ([2, 4], [2, 2]), ) self.assertEqual([call["kwargs"]["current_frame"] for call in calls], [2, 2, 4, 4]) self.assertTrue(torch.equal(calls[0]["kwargs"]["frame_idx"], calls[1]["kwargs"]["frame_idx"])) self.assertTrue(torch.equal(calls[2]["kwargs"]["frame_idx"], calls[3]["kwargs"]["frame_idx"])) self.assertEqual(calls[0]["x"][:, 0, 0, 0, 0].tolist(), [0.0, 0.0, 0.0, 1.0]) self.assertEqual(calls[2]["x"][:, 0, 0, 0, 0].tolist(), [2.0, 0.0, 0.0, 2.0, 2.0]) self.assertEqual(calls[0]["action_cond"][:, 0, 0].tolist(), [102.0, 103.0, 0.0, 0.0]) self.assertEqual(calls[0]["kwargs"]["frame_idx"][:, 0].tolist(), [12, 13, 10, 11]) self.assertEqual(calls[2]["kwargs"]["frame_idx"][:, 0].tolist(), [13, 14, 10, 12, 12]) self.assertEqual(calls[0]["curr_noise_level"][:, 0].tolist(), [2, 2, 0, 0]) self.assertEqual(calls[2]["curr_noise_level"][:, 0].tolist(), [0, 2, 0, 0, 0]) self.assertEqual(calls[0]["kwargs"]["frame_memory_masks"]["revisit"].tolist(), [[]]) self.assertEqual(calls[2]["kwargs"]["frame_memory_masks"]["revisit"].tolist(), [[True]]) self.assertAlmostEqual(float(loss), 1.0 / 3.0) self.assertEqual([name for name, _ in harness.logged], ["test/latent_mse"]) self.assertEqual(len(harness.decoded_inputs), 2) self.assertEqual(harness.decoded_inputs[0][:, 0, 0, 0, 0].tolist(), [2.0, 4.0, 4.0]) self.assertEqual(harness.decoded_inputs[1][:, 0, 0, 0, 0].tolist(), [2.0, 3.0, 4.0]) self.assertEqual(len(harness.metric_updates), 1) metric_pred, metric_gt = harness.metric_updates[0] self.assertEqual(metric_pred[:, 0, 0, 0, 0].tolist(), [12.0, 14.0, 14.0]) self.assertEqual(metric_gt[:, 0, 0, 0, 0].tolist(), [12.0, 13.0, 14.0]) self.assertEqual(len(video_calls), 1) self.assertEqual(video_calls[0][1]["namespace"], "test_vis") self.assertEqual(video_calls[0][1]["step"], 0) def test_validation_multiview_policy_uses_current_query_pool_without_event_cache(self): import algorithms.dememwm.df_video as df_video from algorithms.dememwm.df_video import DeMemWMMinecraft, _preprocess_dememwm_latent_batch build_calls = [] revisit_calls = [] dynamic_calls = [] def fake_build_shared_fov_pool( poses, target_positions, cfg, split, *, min_candidate_frame=0, dynamic_count=0, revisit_count=0, revisit_excluded=None, ): pool = {"query": tuple(int(v) for v in target_positions)} build_calls.append( { "target_positions": target_positions.tolist(), "split": split, "min_candidate_frame": min_candidate_frame, "dynamic_count": dynamic_count, "revisit_count": revisit_count, "pool": pool, } ) return pool def fake_select_revisit(poses, target_positions, cfg, count, excluded, split, rng=None, min_candidate_frame=0, fov_pool=None): revisit_calls.append( { "target_positions": target_positions.tolist(), "split": split, "fov_pool": fov_pool, } ) return np.asarray([0], dtype=np.int64)[:count] def fake_select_dynamic_by_policy( target_start, count, cfg, poses=None, latents=None, actions=None, min_candidate_frame=0, reference_frames=None, excluded=None, target_positions=None, split="training", fov_pool=None, ): dynamic_calls.append( { "target_start": int(target_start), "target_positions": np.asarray(target_positions, dtype=np.int64).tolist(), "excluded": np.asarray(excluded, dtype=np.int64).tolist(), "reference_frames": np.asarray(reference_frames, dtype=np.int64).tolist(), "split": split, "fov_pool": fov_pool, } ) return np.asarray([1], dtype=np.int64)[:count] class FakeDiffusion: stabilization_level = 0 def __init__(self): self.calls = [] def sample_step(self, x, action_cond, pose_cond, curr_noise_level, next_noise_level, **kwargs): self.calls.append({"x": x.detach().clone(), "kwargs": kwargs}) return x[: kwargs["frame_memory_segments"]["target"]] class Harness: def __init__(self): self.memory_condition_length, self.clip_noise = 2, 0.0 self.context_frames, self.frame_stack, self.chunk_size, self.n_tokens = 2, 1, 2, 2 self.cfg = OmegaConf.create( { "memory_selection": { "enabled": True, "max_anchor_frames": 0, "max_dynamic_frames": 1, "max_revisit_frames": 1, "dynamic": _dynamic_cfg("multiview", multiview_selector="fov_greedy"), } } ) self.diffusion_model = FakeDiffusion() self.horizons = [] self.logged = [] self.logger = None self.log_video = False self.log_memory_selection_sheet = False self.save_local = False self.local_save_dir = None def _preprocess_batch(self, batch): return _preprocess_dememwm_latent_batch(batch) def _generate_scheduling_matrix(self, horizon): self.horizons.append(horizon) return np.stack([np.full(horizon, level, dtype=np.int64) for level in (1, 0)]) def decode(self, xs): return xs def _update_metric_accumulators(self, xs_pred, xs_gt, valid_mask=None, eval_start=0): pass def _update_per_frame_metric_accumulators(self, xs_pred, xs_gt, mask, eval_start): pass def _update_revisit_metric_accumulators(self, xs_pred, xs_gt, poses, frame_indices, mask, eval_start, batch_idx, namespace): pass def log(self, name, value): self.logged.append((name, value)) batch = { "latents": torch.arange(7, dtype=torch.float32).view(1, 7, 1, 1, 1), "actions": torch.zeros((1, 7, 25), dtype=torch.float32), "poses": torch.arange(7, dtype=torch.float32).view(1, 7, 1).repeat(1, 1, 5), "frame_indices": (10 + torch.arange(7, dtype=torch.long)).view(1, 7), "memory_segments": {"target": torch.tensor([5]), "anchor": torch.tensor([0]), "dynamic": torch.tensor([1]), "revisit": torch.tensor([1])}, "memory_masks": { "target": torch.ones((1, 5), dtype=torch.bool), "anchor": torch.zeros((1, 0), dtype=torch.bool), "dynamic": torch.ones((1, 1), dtype=torch.bool), "revisit": torch.ones((1, 1), dtype=torch.bool), }, "image_hw": torch.tensor([[360, 640]], dtype=torch.long), } harness = Harness() with mock.patch.object(df_video, "_new_online_event_cache", side_effect=AssertionError("unexpected event cache")), mock.patch.object( df_video, "_extend_online_event_cache", side_effect=AssertionError("unexpected event cache extension") ), mock.patch.object(df_video, "_build_shared_fov_candidate_pool", side_effect=fake_build_shared_fov_pool), mock.patch.object( df_video, "_select_revisit", side_effect=fake_select_revisit ), mock.patch.object(df_video, "_select_dynamic_by_policy", side_effect=fake_select_dynamic_by_policy): DeMemWMMinecraft.validation_step(harness, batch, 0, namespace="test") self.assertEqual(harness.horizons, [2, 1]) self.assertEqual([call["target_positions"] for call in build_calls], [[2, 3], [4]]) self.assertEqual([call["split"] for call in build_calls], ["test", "test"]) self.assertEqual([call["dynamic_count"] for call in build_calls], [1, 1]) self.assertEqual([call["revisit_count"] for call in build_calls], [1, 1]) self.assertEqual([call["target_positions"] for call in revisit_calls], [[2, 3], [4]]) self.assertEqual([call["target_positions"] for call in dynamic_calls], [[2, 3], [4]]) for built, revisit, dynamic in zip(build_calls, revisit_calls, dynamic_calls): self.assertIs(revisit["fov_pool"], built["pool"]) self.assertIs(dynamic["fov_pool"], built["pool"]) self.assertEqual([call["excluded"] for call in dynamic_calls], [[0], [0]]) self.assertEqual([call["reference_frames"] for call in dynamic_calls], [[0], [0]]) self.assertEqual([name for name, _ in harness.logged], ["test/latent_mse"]) def test_validation_event_policy_uses_committed_history_before_local_span(self): from algorithms.dememwm.df_video import DeMemWMMinecraft, _preprocess_dememwm_latent_batch class FakeDiffusion: stabilization_level = 0 def __init__(self): self.calls = [] def sample_step(self, x, action_cond, pose_cond, curr_noise_level, next_noise_level, **kwargs): self.calls.append( { "x": x.detach().clone(), "kwargs": kwargs, } ) return x[: kwargs["frame_memory_segments"]["target"]] class Harness: def __init__(self): self.memory_condition_length, self.clip_noise = 1, 0.0 self.context_frames, self.frame_stack, self.chunk_size, self.n_tokens = 4, 1, 1, 2 self.cfg = OmegaConf.create( { "memory_selection": { "enabled": True, "max_anchor_frames": 0, "max_dynamic_frames": 1, "max_revisit_frames": 0, "dynamic": _dynamic_cfg("event_triggered", stable_frames=1, min_event_gap=0), } } ) self.diffusion_model = FakeDiffusion() self.logged = [] self.logger = None self.log_video = False self.log_memory_selection_sheet = False self.save_local = False self.local_save_dir = None def _preprocess_batch(self, batch): return _preprocess_dememwm_latent_batch(batch) def _generate_scheduling_matrix(self, horizon): return np.stack([np.full(horizon, level, dtype=np.int64) for level in (1, 0)]) def decode(self, xs): return xs def _update_metric_accumulators(self, xs_pred, xs_gt, valid_mask=None, eval_start=0): pass def _update_per_frame_metric_accumulators(self, xs_pred, xs_gt, mask, eval_start): pass def _update_revisit_metric_accumulators(self, xs_pred, xs_gt, poses, frame_indices, mask, eval_start, batch_idx, namespace): pass def log(self, name, value): self.logged.append((name, value)) latents = torch.cat( [ torch.as_tensor(_event_latents(5, event_frame=2)), torch.zeros((1, 2, 1, 1), dtype=torch.float32), ], dim=0, ).unsqueeze(0) batch = { "latents": latents, "actions": torch.zeros((1, 6, 25), dtype=torch.float32), "poses": torch.zeros((1, 6, 5), dtype=torch.float32), "frame_indices": torch.tensor([[10, 11, 12, 13, 14, -1]], dtype=torch.long), "memory_segments": {"target": torch.tensor([5]), "anchor": torch.tensor([0]), "dynamic": torch.tensor([1]), "revisit": torch.tensor([0])}, "memory_masks": { "target": torch.ones((1, 5), dtype=torch.bool), "anchor": torch.zeros((1, 0), dtype=torch.bool), "dynamic": torch.ones((1, 1), dtype=torch.bool), "revisit": torch.zeros((1, 0), dtype=torch.bool), }, "image_hw": torch.tensor([[360, 640]], dtype=torch.long), } harness = Harness() DeMemWMMinecraft.validation_step(harness, batch, 0, namespace="validation") self.assertEqual(len(harness.diffusion_model.calls), 1) call = harness.diffusion_model.calls[0] self.assertEqual(call["kwargs"]["frame_idx"][:, 0].tolist(), [13, 14, 10]) self.assertEqual(call["kwargs"]["frame_memory_masks"]["dynamic"].tolist(), [[True]]) self.assertEqual(call["x"][:, 0, 0, 0, 0].tolist(), [0.0, 0.0, 1.0]) def test_interactive_generation_fails_until_packed_memory_path_exists(self): from algorithms.dememwm.df_video import DeMemWMMinecraft with self.assertRaisesRegex(NotImplementedError, "packed .* frame-memory API"): DeMemWMMinecraft.interactive(object(), None, None, None, None, None, None, None, None, None) def test_disabled_memory_selection_does_not_allocate_memory_slots(self): with tempfile.TemporaryDirectory() as tmp: root = Path(tmp) _write_vae_feature_clip(root, root / "vae_features", stem="sample") dataset = MinecraftVideoDeMemWMLatentDataset( _dataset_cfg(root, memory_selection=_selection_cfg(enabled=False)), split="training", ) sample = dataset[0] self.assertEqual(dataset.memory_condition_length, 0) self.assertEqual(sample["memory_segments"], {"target": 3, "anchor": 0, "dynamic": 0, "revisit": 0}) self.assertEqual(tuple(sample["latents"].shape), (3, 1, 1, 2)) self.assertEqual(sample["frame_indices"].tolist(), [100, 101, 102]) self.assertTrue(sample["memory_masks"]["target"].all().item()) for key in ("anchor", "dynamic", "revisit"): self.assertEqual(tuple(sample["memory_masks"][key].shape), (0,)) def test_dataset_excludes_initial_offset_from_memory_candidates(self): with tempfile.TemporaryDirectory() as tmp: root = Path(tmp) _write_vae_feature_clip( root, root / "vae_features", stem="sample", num_frames=112, actions=np.ones((112, 25), dtype=np.float32), ) dataset = MinecraftVideoDeMemWMLatentDataset( _dataset_cfg(root, memory_selection=_selection_cfg()), split="training", ) sample = dataset[0] self.assertEqual(sample["frame_indices"][:3].tolist(), [100, 101, 102]) self.assertEqual(sample["frame_indices"][3:].tolist(), [-1, -1, -1, -1, -1, -1]) for key in ("anchor", "dynamic", "revisit"): self.assertFalse(sample["memory_masks"][key].any().item()) def test_training_anchor_window_uses_initial_offset_context_length(self): with tempfile.TemporaryDirectory() as tmp: root = Path(tmp) _write_vae_feature_clip( root, root / "vae_features", stem="sample", num_frames=112, actions=np.ones((112, 25), dtype=np.float32), ) dataset = MinecraftVideoDeMemWMLatentDataset( _dataset_cfg( root, context_length=4, memory_selection=_selection_cfg(max_anchor_frames=2, max_dynamic_frames=2, max_revisit_frames=0, anchor_diverse_selection=True), ), split="training", ) sample = dataset[2] self.assertEqual(dataset.context_length, 4) self.assertEqual(len(dataset), 112 - (100 + 4 + 3) + 1) self.assertEqual(sample["frame_indices"].tolist(), [106, 107, 108, 102, 105, 102, 103]) self.assertTrue(sample["memory_masks"]["anchor"].all().item()) self.assertTrue(sample["memory_masks"]["dynamic"].all().item()) self.assertEqual(tuple(sample["memory_masks"]["revisit"].shape), (0,)) def test_validation_ignores_training_context_length_shift(self): with tempfile.TemporaryDirectory() as tmp: root = Path(tmp) _write_vae_feature_clip( root, root / "vae_features", split="validation", stem="clip_w_updown", num_frames=112, actions=np.ones((112, 25), dtype=np.float32), ) dataset = MinecraftVideoDeMemWMLatentDataset( _dataset_cfg(root, context_length=4), split="validation", ) sample = dataset[0] self.assertEqual(dataset.context_length, 0) self.assertEqual(sample["frame_indices"].tolist(), [100, 101, 102]) def test_dataset_passes_latents_actions_to_event_dynamic_selection(self): with tempfile.TemporaryDirectory() as tmp: root = Path(tmp) num_frames = 112 _write_vae_feature_clip( root, root / "vae_features", stem="sample", num_frames=num_frames, actions=np.zeros((num_frames, 25), dtype=np.float32), latents=_event_latents(num_frames, event_frame=102), ) dataset = MinecraftVideoDeMemWMLatentDataset( _dataset_cfg( root, context_length=6, memory_selection=_selection_cfg( max_anchor_frames=0, max_dynamic_frames=2, max_revisit_frames=0, dynamic=_dynamic_cfg("event_triggered", stable_frames=1, min_event_gap=0), ), ), split="training", ) sample = dataset[0] self.assertEqual(sample["frame_indices"].tolist(), [106, 107, 108, 100, 103]) self.assertEqual(sample["memory_masks"]["dynamic"].tolist(), [True, True]) def test_dataset_does_not_build_dynamic_stream_for_multiview_policy(self): import datasets.video.minecraft_video_dememwm_latent_dataset as dataset_module with tempfile.TemporaryDirectory() as tmp: root = Path(tmp) num_frames = 112 _write_vae_feature_clip( root, root / "vae_features", stem="sample", num_frames=num_frames, actions=np.zeros((num_frames, 25), dtype=np.float32), latents=_event_latents(num_frames, event_frame=102), ) dataset = MinecraftVideoDeMemWMLatentDataset( _dataset_cfg( root, context_length=6, memory_selection=_selection_cfg( max_anchor_frames=0, max_dynamic_frames=2, max_revisit_frames=0, dynamic=_dynamic_cfg("multiview"), ), ), split="training", ) with mock.patch.object(dataset_module, "_build_dynamic_stream", side_effect=AssertionError("unexpected dynamic stream build")): arrays = dataset._load_feature_arrays(dataset.data_paths[0]) self.assertEqual(arrays["dynamic_stream"].tolist(), []) def test_dataset_returns_target_anchor_dynamic_revisit_contract(self): with tempfile.TemporaryDirectory() as tmp: root = Path(tmp) num_frames = 112 _write_vae_feature_clip( root, root / "vae_features", stem="sample", num_frames=num_frames, actions=np.ones((num_frames, 25), dtype=np.float32), ) dataset = MinecraftVideoDeMemWMLatentDataset( _dataset_cfg(root, context_length=6, memory_selection=_selection_cfg(training_use_plucker=True)), split="training", ) sample = dataset[0] self.assertEqual(sample["memory_segments"], {"target": 3, "anchor": 2, "dynamic": 2, "revisit": 2}) self.assertEqual(tuple(sample["latents"].shape), (9, 1, 1, 2)) self.assertEqual(sample["frame_indices"][:3].tolist(), [106, 107, 108]) self.assertEqual(sample["frame_indices"][3:5].tolist(), [100, 101]) self.assertEqual(sample["frame_indices"][5:7].tolist(), [102, 103]) revisit_frames = sample["frame_indices"][7:9].numpy() self.assertEqual(revisit_frames.tolist(), [102, 103]) self.assertEqual(len(np.unique(revisit_frames)), 2) self.assertTrue(sample["memory_masks"]["target"].all().item()) self.assertTrue(sample["memory_masks"]["anchor"].all().item()) self.assertTrue(sample["memory_masks"]["dynamic"].all().item()) self.assertTrue(sample["memory_masks"]["revisit"].all().item()) self.assertEqual(sample["image_hw"].tolist(), [360, 640]) from torch.utils.data import default_collate from algorithms.dememwm.df_video import _preprocess_dememwm_latent_batch preprocessed = _preprocess_dememwm_latent_batch(default_collate([sample])) self.assertEqual(tuple(preprocessed["latents"].shape), (9, 1, 1, 1, 2)) self.assertEqual(tuple(preprocessed["actions"].shape), (9, 1, 25)) self.assertEqual(tuple(preprocessed["action_conditions"].shape), (9, 1, 25)) self.assertTrue(torch.equal(preprocessed["action_conditions"][0], torch.zeros_like(preprocessed["actions"][0]))) self.assertTrue(torch.equal(preprocessed["action_conditions"][1:3], preprocessed["actions"][1:3])) self.assertTrue(torch.equal(preprocessed["action_conditions"][3:], torch.zeros_like(preprocessed["actions"][3:]))) self.assertTrue(torch.equal(preprocessed["actions"][3:5], torch.ones_like(preprocessed["actions"][3:5]))) self.assertTrue(torch.equal(preprocessed["actions"][5:], torch.ones_like(preprocessed["actions"][5:]))) self.assertEqual(tuple(preprocessed["poses"].shape), (9, 1, 5)) self.assertIs(preprocessed["frame_memory_pose"], preprocessed["poses"]) self.assertEqual(preprocessed["frame_memory_pose"][:7, 0, 0].tolist(), [0.0, 1.0, 2.0, -6.0, -5.0, -4.0, -3.0]) self.assertEqual(tuple(preprocessed["frame_indices"].shape), (9, 1)) self.assertEqual(preprocessed["memory_segments"], sample["memory_segments"]) self.assertEqual(preprocessed["segment_lengths"], sample["memory_segments"]) self.assertEqual(preprocessed["target_length"], 3) self.assertEqual(preprocessed["stream_lengths"], {"anchor": 2, "dynamic": 2, "revisit": 2}) self.assertTrue(all(isinstance(value, int) for value in preprocessed["memory_segments"].values())) self.assertEqual( {key: (slc.start, slc.stop) for key, slc in preprocessed["segment_slices"].items()}, {"target": (0, 3), "anchor": (3, 5), "dynamic": (5, 7), "revisit": (7, 9)}, ) self.assertEqual((preprocessed["target_slice"].start, preprocessed["target_slice"].stop), (0, 3)) self.assertEqual( {key: (slc.start, slc.stop) for key, slc in preprocessed["stream_slices"].items()}, {"anchor": (3, 5), "dynamic": (5, 7), "revisit": (7, 9)}, ) self.assertEqual(preprocessed["target_tensors"]["frame_indices"][:, 0].tolist(), [106, 107, 108]) self.assertEqual(preprocessed["stream_tensors"]["anchor"]["frame_indices"][:, 0].tolist(), [100, 101]) self.assertEqual(preprocessed["stream_tensors"]["dynamic"]["frame_indices"][:, 0].tolist(), [102, 103]) self.assertEqual(preprocessed["stream_tensors"]["revisit"]["frame_indices"].shape[0], 2) self.assertEqual(tuple(preprocessed["memory_masks"]["target"].shape), (1, 3)) self.assertEqual(tuple(preprocessed["memory_masks"]["anchor"].shape), (1, 2)) self.assertEqual(tuple(preprocessed["memory_masks"]["dynamic"].shape), (1, 2)) self.assertEqual(tuple(preprocessed["memory_masks"]["revisit"].shape), (1, 2)) self.assertEqual(preprocessed["image_hw"].tolist(), [[360, 640]]) if __name__ == "__main__": unittest.main()