Prepare DeMemWM latent action conditioning
Browse files
algorithms/dememwm/df_video.py
CHANGED
|
@@ -77,9 +77,16 @@ def _preprocess_dememwm_latent_batch(batch):
|
|
| 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 |
}
|
|
@@ -95,6 +102,7 @@ def _preprocess_dememwm_latent_batch(batch):
|
|
| 95 |
return {
|
| 96 |
"latents": latents,
|
| 97 |
"actions": actions,
|
|
|
|
| 98 |
"poses": poses,
|
| 99 |
"frame_indices": frame_indices,
|
| 100 |
"memory_segments": memory_segments,
|
|
|
|
| 77 |
f"memory_segments sum to {start} frames, but latent batch has {latents.shape[0]}"
|
| 78 |
)
|
| 79 |
|
| 80 |
+
action_conditions = actions.clone()
|
| 81 |
+
if target_length:
|
| 82 |
+
action_conditions[target_slice.start:target_slice.start + 1] = 0
|
| 83 |
+
for stream_slice in stream_slices.values():
|
| 84 |
+
action_conditions[stream_slice] = 0
|
| 85 |
+
|
| 86 |
sequence_tensors = {
|
| 87 |
"latents": latents,
|
| 88 |
"actions": actions,
|
| 89 |
+
"action_conditions": action_conditions,
|
| 90 |
"poses": poses,
|
| 91 |
"frame_indices": frame_indices,
|
| 92 |
}
|
|
|
|
| 102 |
return {
|
| 103 |
"latents": latents,
|
| 104 |
"actions": actions,
|
| 105 |
+
"action_conditions": action_conditions,
|
| 106 |
"poses": poses,
|
| 107 |
"frame_indices": frame_indices,
|
| 108 |
"memory_segments": memory_segments,
|
tests/test_dememwm_latent_dataset.py
CHANGED
|
@@ -175,6 +175,9 @@ class DeMemWMLatentDatasetTests(unittest.TestCase):
|
|
| 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"])
|
|
@@ -231,6 +234,11 @@ class DeMemWMLatentDatasetTests(unittest.TestCase):
|
|
| 231 |
preprocessed = _preprocess_dememwm_latent_batch(default_collate([sample]))
|
| 232 |
self.assertEqual(tuple(preprocessed["latents"].shape), (9, 1, 1, 1, 2))
|
| 233 |
self.assertEqual(tuple(preprocessed["actions"].shape), (9, 1, 25))
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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"])
|
|
|
|
| 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["action_conditions"][:, 0, 0].tolist(), [0.0, 101.0, 0.0, 0.0, 0.0, 0.0])
|
| 179 |
+
self.assertEqual(preprocessed["target_tensors"]["action_conditions"][:, 0, 0].tolist(), [0.0, 101.0])
|
| 180 |
+
self.assertEqual(preprocessed["stream_tensors"]["dynamic"]["action_conditions"][:, 0, 0].tolist(), [0.0, 0.0, 0.0])
|
| 181 |
self.assertEqual(preprocessed["stream_tensors"]["dynamic"]["poses"][:, 0, 0].tolist(), [203.0, 204.0, 205.0])
|
| 182 |
self.assertEqual(preprocessed["stream_tensors"]["dynamic"]["frame_indices"][:, 0].tolist(), [13, 14, 15])
|
| 183 |
self.assertIs(preprocessed["memory_masks"]["dynamic"], batch["memory_masks"]["dynamic"])
|
|
|
|
| 234 |
preprocessed = _preprocess_dememwm_latent_batch(default_collate([sample]))
|
| 235 |
self.assertEqual(tuple(preprocessed["latents"].shape), (9, 1, 1, 1, 2))
|
| 236 |
self.assertEqual(tuple(preprocessed["actions"].shape), (9, 1, 25))
|
| 237 |
+
self.assertEqual(tuple(preprocessed["action_conditions"].shape), (9, 1, 25))
|
| 238 |
+
self.assertTrue(torch.equal(preprocessed["action_conditions"][0], torch.zeros_like(preprocessed["actions"][0])))
|
| 239 |
+
self.assertTrue(torch.equal(preprocessed["action_conditions"][1:3], preprocessed["actions"][1:3]))
|
| 240 |
+
self.assertTrue(torch.equal(preprocessed["action_conditions"][3:], torch.zeros_like(preprocessed["actions"][3:])))
|
| 241 |
+
self.assertTrue(torch.equal(preprocessed["actions"][3:], torch.ones_like(preprocessed["actions"][3:])))
|
| 242 |
self.assertEqual(tuple(preprocessed["poses"].shape), (9, 1, 5))
|
| 243 |
self.assertEqual(tuple(preprocessed["frame_indices"].shape), (9, 1))
|
| 244 |
self.assertEqual(preprocessed["memory_segments"], sample["memory_segments"])
|