Implement DeMemWM segment split preprocessing
Browse files
algorithms/dememwm/df_video.py
CHANGED
|
@@ -26,6 +26,7 @@ import glob
|
|
| 26 |
|
| 27 |
# Utility Functions
|
| 28 |
_DEMEMWM_SEGMENT_KEYS = ("target", "anchor", "dynamic", "revisit")
|
|
|
|
| 29 |
|
| 30 |
|
| 31 |
def _segment_value_to_int(value, key):
|
|
@@ -52,6 +53,9 @@ def _preprocess_dememwm_latent_batch(batch):
|
|
| 52 |
key: _segment_value_to_int(batch["memory_segments"][key], key)
|
| 53 |
for key in _DEMEMWM_SEGMENT_KEYS
|
| 54 |
}
|
|
|
|
|
|
|
|
|
|
| 55 |
memory_masks = {key: batch["memory_masks"][key] for key in _DEMEMWM_SEGMENT_KEYS}
|
| 56 |
|
| 57 |
# Latent dataset batches are already VAE-encoded: B x T_all x C x H_lat x W_lat.
|
|
@@ -60,6 +64,33 @@ def _preprocess_dememwm_latent_batch(batch):
|
|
| 60 |
poses = rearrange(batch["poses"], "b t d -> t b d").contiguous()
|
| 61 |
frame_indices = rearrange(batch["frame_indices"], "b t -> t b").contiguous()
|
| 62 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 63 |
# Keep original image H/W for later ray geometry; latent H/W is not a substitute.
|
| 64 |
return {
|
| 65 |
"latents": latents,
|
|
@@ -67,6 +98,15 @@ def _preprocess_dememwm_latent_batch(batch):
|
|
| 67 |
"poses": poses,
|
| 68 |
"frame_indices": frame_indices,
|
| 69 |
"memory_segments": memory_segments,
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 70 |
"memory_masks": memory_masks,
|
| 71 |
"image_hw": batch["image_hw"],
|
| 72 |
}
|
|
|
|
| 26 |
|
| 27 |
# Utility Functions
|
| 28 |
_DEMEMWM_SEGMENT_KEYS = ("target", "anchor", "dynamic", "revisit")
|
| 29 |
+
_DEMEMWM_STREAM_KEYS = ("anchor", "dynamic", "revisit")
|
| 30 |
|
| 31 |
|
| 32 |
def _segment_value_to_int(value, key):
|
|
|
|
| 53 |
key: _segment_value_to_int(batch["memory_segments"][key], key)
|
| 54 |
for key in _DEMEMWM_SEGMENT_KEYS
|
| 55 |
}
|
| 56 |
+
segment_lengths = dict(memory_segments)
|
| 57 |
+
target_length = segment_lengths["target"]
|
| 58 |
+
stream_lengths = {key: segment_lengths[key] for key in _DEMEMWM_STREAM_KEYS}
|
| 59 |
memory_masks = {key: batch["memory_masks"][key] for key in _DEMEMWM_SEGMENT_KEYS}
|
| 60 |
|
| 61 |
# Latent dataset batches are already VAE-encoded: B x T_all x C x H_lat x W_lat.
|
|
|
|
| 64 |
poses = rearrange(batch["poses"], "b t d -> t b d").contiguous()
|
| 65 |
frame_indices = rearrange(batch["frame_indices"], "b t -> t b").contiguous()
|
| 66 |
|
| 67 |
+
segment_slices = {}
|
| 68 |
+
start = 0
|
| 69 |
+
for key in _DEMEMWM_SEGMENT_KEYS:
|
| 70 |
+
stop = start + segment_lengths[key]
|
| 71 |
+
segment_slices[key] = slice(start, stop)
|
| 72 |
+
start = stop
|
| 73 |
+
target_slice = segment_slices["target"]
|
| 74 |
+
stream_slices = {key: segment_slices[key] for key in _DEMEMWM_STREAM_KEYS}
|
| 75 |
+
if start != latents.shape[0]:
|
| 76 |
+
raise ValueError(
|
| 77 |
+
f"memory_segments sum to {start} frames, but latent batch has {latents.shape[0]}"
|
| 78 |
+
)
|
| 79 |
+
|
| 80 |
+
sequence_tensors = {
|
| 81 |
+
"latents": latents,
|
| 82 |
+
"actions": actions,
|
| 83 |
+
"poses": poses,
|
| 84 |
+
"frame_indices": frame_indices,
|
| 85 |
+
}
|
| 86 |
+
# Packed sequence order stays [target][anchor][dynamic][revisit] in T x B layout.
|
| 87 |
+
segments = {
|
| 88 |
+
key: {name: tensor[segment_slices[key]] for name, tensor in sequence_tensors.items()}
|
| 89 |
+
for key in _DEMEMWM_SEGMENT_KEYS
|
| 90 |
+
}
|
| 91 |
+
target_tensors = segments["target"]
|
| 92 |
+
stream_tensors = {key: segments[key] for key in _DEMEMWM_STREAM_KEYS}
|
| 93 |
+
|
| 94 |
# Keep original image H/W for later ray geometry; latent H/W is not a substitute.
|
| 95 |
return {
|
| 96 |
"latents": latents,
|
|
|
|
| 98 |
"poses": poses,
|
| 99 |
"frame_indices": frame_indices,
|
| 100 |
"memory_segments": memory_segments,
|
| 101 |
+
"segment_lengths": segment_lengths,
|
| 102 |
+
"target_length": target_length,
|
| 103 |
+
"stream_lengths": stream_lengths,
|
| 104 |
+
"segment_slices": segment_slices,
|
| 105 |
+
"target_slice": target_slice,
|
| 106 |
+
"stream_slices": stream_slices,
|
| 107 |
+
"segments": segments,
|
| 108 |
+
"target_tensors": target_tensors,
|
| 109 |
+
"stream_tensors": stream_tensors,
|
| 110 |
"memory_masks": memory_masks,
|
| 111 |
"image_hw": batch["image_hw"],
|
| 112 |
}
|
tests/test_dememwm_latent_dataset.py
CHANGED
|
@@ -3,6 +3,7 @@ import unittest
|
|
| 3 |
from pathlib import Path
|
| 4 |
|
| 5 |
import numpy as np
|
|
|
|
| 6 |
from omegaconf import OmegaConf
|
| 7 |
|
| 8 |
from datasets.video.memory_selection import select_memory_indices
|
|
@@ -138,6 +139,48 @@ class MemorySelectionTests(unittest.TestCase):
|
|
| 138 |
|
| 139 |
|
| 140 |
class DeMemWMLatentDatasetTests(unittest.TestCase):
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 141 |
def test_dataset_returns_target_anchor_dynamic_revisit_contract(self):
|
| 142 |
with tempfile.TemporaryDirectory() as tmp:
|
| 143 |
root = Path(tmp)
|
|
@@ -191,7 +234,23 @@ class DeMemWMLatentDatasetTests(unittest.TestCase):
|
|
| 191 |
self.assertEqual(tuple(preprocessed["poses"].shape), (9, 1, 5))
|
| 192 |
self.assertEqual(tuple(preprocessed["frame_indices"].shape), (9, 1))
|
| 193 |
self.assertEqual(preprocessed["memory_segments"], sample["memory_segments"])
|
|
|
|
|
|
|
|
|
|
| 194 |
self.assertTrue(all(isinstance(value, int) for value in preprocessed["memory_segments"].values()))
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 195 |
self.assertEqual(tuple(preprocessed["memory_masks"]["target"].shape), (1, 3))
|
| 196 |
self.assertEqual(tuple(preprocessed["memory_masks"]["anchor"].shape), (1, 2))
|
| 197 |
self.assertEqual(tuple(preprocessed["memory_masks"]["dynamic"].shape), (1, 2))
|
|
|
|
| 3 |
from pathlib import Path
|
| 4 |
|
| 5 |
import numpy as np
|
| 6 |
+
import torch
|
| 7 |
from omegaconf import OmegaConf
|
| 8 |
|
| 9 |
from datasets.video.memory_selection import select_memory_indices
|
|
|
|
| 139 |
|
| 140 |
|
| 141 |
class DeMemWMLatentDatasetTests(unittest.TestCase):
|
| 142 |
+
|
| 143 |
+
def test_preprocess_splits_sequence_tensors_from_memory_segments(self):
|
| 144 |
+
from algorithms.dememwm.df_video import _preprocess_dememwm_latent_batch
|
| 145 |
+
|
| 146 |
+
batch = {
|
| 147 |
+
"latents": torch.arange(6, dtype=torch.float32).view(1, 6, 1, 1, 1),
|
| 148 |
+
"actions": (100 + torch.arange(6, dtype=torch.float32)).view(1, 6, 1),
|
| 149 |
+
"poses": (200 + torch.arange(6, dtype=torch.float32)).view(1, 6, 1).repeat(1, 1, 5),
|
| 150 |
+
"frame_indices": (10 + torch.arange(6, dtype=torch.long)).view(1, 6),
|
| 151 |
+
"memory_segments": {
|
| 152 |
+
"target": torch.tensor([2]),
|
| 153 |
+
"anchor": torch.tensor([1]),
|
| 154 |
+
"dynamic": torch.tensor([3]),
|
| 155 |
+
"revisit": torch.tensor([0]),
|
| 156 |
+
},
|
| 157 |
+
"memory_masks": {
|
| 158 |
+
"target": torch.tensor([[True, False]]),
|
| 159 |
+
"anchor": torch.tensor([[True]]),
|
| 160 |
+
"dynamic": torch.tensor([[True, False, True]]),
|
| 161 |
+
"revisit": torch.zeros((1, 0), dtype=torch.bool),
|
| 162 |
+
},
|
| 163 |
+
"image_hw": torch.tensor([[360, 640]], dtype=torch.long),
|
| 164 |
+
}
|
| 165 |
+
|
| 166 |
+
preprocessed = _preprocess_dememwm_latent_batch(batch)
|
| 167 |
+
|
| 168 |
+
self.assertEqual(preprocessed["target_length"], 2)
|
| 169 |
+
self.assertEqual(preprocessed["stream_lengths"], {"anchor": 1, "dynamic": 3, "revisit": 0})
|
| 170 |
+
self.assertEqual((preprocessed["target_slice"].start, preprocessed["target_slice"].stop), (0, 2))
|
| 171 |
+
self.assertEqual(
|
| 172 |
+
{key: (slc.start, slc.stop) for key, slc in preprocessed["stream_slices"].items()},
|
| 173 |
+
{"anchor": (2, 3), "dynamic": (3, 6), "revisit": (6, 6)},
|
| 174 |
+
)
|
| 175 |
+
self.assertEqual(preprocessed["target_tensors"]["latents"][:, 0, 0, 0, 0].tolist(), [0.0, 1.0])
|
| 176 |
+
self.assertEqual(preprocessed["stream_tensors"]["dynamic"]["latents"][:, 0, 0, 0, 0].tolist(), [3.0, 4.0, 5.0])
|
| 177 |
+
self.assertEqual(preprocessed["stream_tensors"]["dynamic"]["actions"][:, 0, 0].tolist(), [103.0, 104.0, 105.0])
|
| 178 |
+
self.assertEqual(preprocessed["stream_tensors"]["dynamic"]["poses"][:, 0, 0].tolist(), [203.0, 204.0, 205.0])
|
| 179 |
+
self.assertEqual(preprocessed["stream_tensors"]["dynamic"]["frame_indices"][:, 0].tolist(), [13, 14, 15])
|
| 180 |
+
self.assertIs(preprocessed["memory_masks"]["dynamic"], batch["memory_masks"]["dynamic"])
|
| 181 |
+
self.assertEqual(tuple(preprocessed["memory_masks"]["dynamic"].shape), (1, 3))
|
| 182 |
+
self.assertEqual(preprocessed["memory_masks"]["dynamic"].device.type, "cpu")
|
| 183 |
+
|
| 184 |
def test_dataset_returns_target_anchor_dynamic_revisit_contract(self):
|
| 185 |
with tempfile.TemporaryDirectory() as tmp:
|
| 186 |
root = Path(tmp)
|
|
|
|
| 234 |
self.assertEqual(tuple(preprocessed["poses"].shape), (9, 1, 5))
|
| 235 |
self.assertEqual(tuple(preprocessed["frame_indices"].shape), (9, 1))
|
| 236 |
self.assertEqual(preprocessed["memory_segments"], sample["memory_segments"])
|
| 237 |
+
self.assertEqual(preprocessed["segment_lengths"], sample["memory_segments"])
|
| 238 |
+
self.assertEqual(preprocessed["target_length"], 3)
|
| 239 |
+
self.assertEqual(preprocessed["stream_lengths"], {"anchor": 2, "dynamic": 2, "revisit": 2})
|
| 240 |
self.assertTrue(all(isinstance(value, int) for value in preprocessed["memory_segments"].values()))
|
| 241 |
+
self.assertEqual(
|
| 242 |
+
{key: (slc.start, slc.stop) for key, slc in preprocessed["segment_slices"].items()},
|
| 243 |
+
{"target": (0, 3), "anchor": (3, 5), "dynamic": (5, 7), "revisit": (7, 9)},
|
| 244 |
+
)
|
| 245 |
+
self.assertEqual((preprocessed["target_slice"].start, preprocessed["target_slice"].stop), (0, 3))
|
| 246 |
+
self.assertEqual(
|
| 247 |
+
{key: (slc.start, slc.stop) for key, slc in preprocessed["stream_slices"].items()},
|
| 248 |
+
{"anchor": (3, 5), "dynamic": (5, 7), "revisit": (7, 9)},
|
| 249 |
+
)
|
| 250 |
+
self.assertEqual(preprocessed["target_tensors"]["frame_indices"][:, 0].tolist(), [106, 107, 108])
|
| 251 |
+
self.assertEqual(preprocessed["stream_tensors"]["anchor"]["frame_indices"][:, 0].tolist(), [100, 101])
|
| 252 |
+
self.assertEqual(preprocessed["stream_tensors"]["dynamic"]["frame_indices"][:, 0].tolist(), [104, 105])
|
| 253 |
+
self.assertEqual(preprocessed["stream_tensors"]["revisit"]["frame_indices"].shape[0], 2)
|
| 254 |
self.assertEqual(tuple(preprocessed["memory_masks"]["target"].shape), (1, 3))
|
| 255 |
self.assertEqual(tuple(preprocessed["memory_masks"]["anchor"].shape), (1, 2))
|
| 256 |
self.assertEqual(tuple(preprocessed["memory_masks"]["dynamic"].shape), (1, 2))
|