dreamzero / tests /trex_track_force /test_visualization.py
zhicao's picture
Upload dreamzero source code (no model checkpoints)
fbd9366 verified
Raw
History Blame Contribute Delete
8.35 kB
from pathlib import Path
from types import ModuleType, SimpleNamespace
import sys
import numpy as np
import torch
from groot.vla.experiment import trex_wandb_video_callback as callback_module
from groot.vla.experiment import trex_eval_utils
from groot.vla.experiment.trex_track_force_eval_utils import (
TrexTrackForceVisualization,
_decode_future_video,
overlay_track_motion,
)
from groot.vla.experiment.trex_wandb_video_callback import TrexWandbVideoCallback
def test_track_motion_overlay_outputs_only_three_view_panels():
frames = np.zeros((8, 100, 200, 3), dtype=np.uint8)
tracks = np.full((1, 16, 250, 2), 0.5, dtype=np.float32)
tracks[0, :, :, 0] += np.linspace(-0.2, 0.2, 16)[:, None]
visualization = TrexTrackForceVisualization(
video_frames=frames,
predicted_track_xy=tracks,
target_track_xy=tracks.copy(),
target_track_visibility=np.ones((1, 16, 250), dtype=np.float32),
anchor_timestamps=np.array([0.0], dtype=np.float64),
frame_track_indices=np.array(
[2, 4, 6, 8, 10, 12, 14, 15], dtype=np.int64
),
)
overlaid = overlay_track_motion(visualization, trail_steps=8)
assert overlaid.shape == (8, 50, 300, 3)
assert np.count_nonzero(overlaid[:, :, :100]) > 0
assert np.count_nonzero(overlaid[:, :, 100:200]) > 0
assert np.count_nonzero(overlaid[:, :, 200:]) > 0
def test_track_motion_overlay_rejects_mismatched_video_indices():
visualization = TrexTrackForceVisualization(
video_frames=np.zeros((2, 100, 200, 3), dtype=np.uint8),
predicted_track_xy=np.zeros((1, 16, 250, 2), dtype=np.float32),
target_track_xy=np.zeros((1, 16, 250, 2), dtype=np.float32),
target_track_visibility=np.ones((1, 16, 250), dtype=np.float32),
anchor_timestamps=np.array([0.0], dtype=np.float64),
frame_track_indices=np.array([1], dtype=np.int64),
)
try:
overlay_track_motion(visualization)
except ValueError as exc:
assert "video and frame-track index counts" in str(exc)
else:
raise AssertionError("mismatched frame-track indices must fail")
def test_track_video_starts_with_model_gt_anchor():
conditioning = np.full((250, 2), 0.5, dtype=np.float32)
anchored = np.repeat(conditioning[None], 16, axis=0)
anchored[:, :, 0] += np.linspace(0.0, 0.1, 16)[:, None]
np.testing.assert_allclose(anchored[0], conditioning)
visualization = TrexTrackForceVisualization(
video_frames=np.zeros((9, 100, 200, 3), dtype=np.uint8),
predicted_track_xy=anchored[None],
target_track_xy=anchored[None],
target_track_visibility=np.ones((1, 16, 250), dtype=np.float32),
anchor_timestamps=np.array([0.0], dtype=np.float64),
frame_track_indices=np.array(
[-1, 2, 4, 6, 8, 10, 12, 14, 15], dtype=np.int64
),
conditioning_frame=np.zeros((100, 200, 3), dtype=np.uint8),
conditioning_track_xy=conditioning[None],
conditioning_track_visibility=np.ones(
(1, 250), dtype=np.float32
),
)
overlaid = overlay_track_motion(visualization)
assert overlaid.shape == (9, 50, 300, 3)
assert np.count_nonzero(overlaid[0]) > 0
def test_callback_saves_plain_and_track_overlay_videos_separately(
tmp_path, monkeypatch
):
visualization = TrexTrackForceVisualization(
video_frames=np.zeros((8, 100, 200, 3), dtype=np.uint8),
predicted_track_xy=np.zeros((1, 16, 250, 2), dtype=np.float32),
target_track_xy=np.zeros((1, 16, 250, 2), dtype=np.float32),
target_track_visibility=np.ones((1, 16, 250), dtype=np.float32),
anchor_timestamps=np.array([0.0], dtype=np.float64),
frame_track_indices=np.array(
[2, 4, 6, 8, 10, 12, 14, 15], dtype=np.int64
),
)
monkeypatch.setattr(
callback_module,
"run_trex_track_force_prediction",
lambda *args, **kwargs: visualization,
)
monkeypatch.setattr(
callback_module,
"overlay_track_motion",
lambda *args, **kwargs: np.full_like(
visualization.video_frames, 255
),
)
written: dict[Path, np.ndarray] = {}
def fake_mimsave(path, frames, **kwargs):
written[Path(path)] = np.stack(frames)
monkeypatch.setattr(callback_module.imageio, "mimsave", fake_mimsave)
logged: list[dict] = []
fake_wandb = ModuleType("wandb")
fake_wandb.run = object()
fake_wandb.define_metric = lambda *args, **kwargs: None
fake_wandb.Video = lambda path, **kwargs: {"path": path, **kwargs}
fake_wandb.log = logged.append
monkeypatch.setitem(sys.modules, "wandb", fake_wandb)
callback = TrexWandbVideoCallback.__new__(TrexWandbVideoCallback)
callback.episode_index = 0
callback.num_chunks = 1
callback.every_n_steps = 500
callback.fps = 10
callback.eval_bf16 = False
callback.use_dataset_prompt = False
callback.prompt = "test task"
callback.overlay_tracks = True
callback.save_tracks = True
callback.track_trail_steps = 8
callback.reconstruction_inference_steps = 1
callback.eval_video_dir = tmp_path / "eval_videos"
callback.eval_track_video_dir = tmp_path / "eval_track_videos"
callback.eval_track_dir = tmp_path / "eval_tracks"
callback._episode = SimpleNamespace(
episode_index=0,
get_task=lambda row: "dataset task",
)
callback._eval_transform = object()
class FakeModel:
training = True
def eval(self):
self.training = False
def train(self):
self.training = True
model = FakeModel()
state = SimpleNamespace(is_world_process_zero=True, global_step=500)
callback.on_train_begin(None, state, None)
callback.on_step_end(None, state, None, model=model)
plain_path = callback.eval_video_dir / "train_step_000500.mp4"
track_path = callback.eval_track_video_dir / "train_step_000500.mp4"
assert np.count_nonzero(written[plain_path]) == 0
assert np.all(written[track_path] == 255)
assert (callback.eval_track_dir / "train_step_000500.npz").is_file()
assert TrexWandbVideoCallback.VIDEO_METRIC in logged[0]
assert TrexWandbVideoCallback.TRACK_VIDEO_METRIC in logged[0]
assert model.training
def test_episode_frame_reads_are_cached_across_visualization_chunks(monkeypatch):
opens = 0
class FakeCapture:
def __init__(self, path):
nonlocal opens
opens += 1
self.position = 0
def isOpened(self):
return True
def set(self, prop, value):
self.position = int(value)
def read(self):
value = self.position
self.position += 1
return True, np.full((2, 3, 3), value, dtype=np.uint8)
def release(self):
pass
monkeypatch.setattr(trex_eval_utils.cv2, "VideoCapture", FakeCapture)
episode = trex_eval_utils.TrexEpisode.__new__(
trex_eval_utils.TrexEpisode
)
episode.length = 10
episode.video_dirs = {"video.head_left": "fake.mp4"}
episode._frame_cache = {"video.head_left": {}}
first = episode.get_frames([0, 1, 1, 2], "video.head_left")
second = episode.get_frames([1, 2], "video.head_left")
assert opens == 1
assert first[:, 0, 0, 0].tolist() == [0, 1, 1, 2]
assert second[:, 0, 0, 0].tolist() == [1, 2]
def test_track_force_video_decode_reads_tiling_from_policy_config():
seen = {}
class FakeVAE:
def decode(self, latents, **kwargs):
seen["shape"] = tuple(latents.shape)
seen.update(kwargs)
return torch.zeros(1, 3, 9, 4, 8)
action_head = SimpleNamespace(
config=SimpleNamespace(
tiled=False,
tile_size_height=20,
tile_size_width=21,
tile_stride_height=10,
tile_stride_width=11,
),
vae=FakeVAE(),
)
model = SimpleNamespace(action_head=action_head)
frames = _decode_future_video(
model,
torch.zeros(1, 4, 3, 2, 2),
torch.zeros(1, 4, 2, 2, 2),
)
assert frames.shape == (8, 4, 8, 3)
assert seen == {
"shape": (1, 4, 3, 2, 2),
"tiled": False,
"tile_size": (20, 21),
"tile_stride": (10, 11),
}