BonanDing commited on
Commit
6b25dc9
·
1 Parent(s): 072e8f7

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