Tighten DeMemWM acceptance tests
Browse files
tests/test_dememwm_latent_dataset.py
CHANGED
|
@@ -460,17 +460,26 @@ class DeMemWMLatentDatasetTests(unittest.TestCase):
|
|
| 460 |
calls = harness.diffusion_model.calls
|
| 461 |
required_kwargs = {"reference_length", "frame_memory_segments", "frame_memory_masks", "frame_memory_pose", "frame_idx", "image_hw"}
|
| 462 |
|
| 463 |
-
self.assertEqual(
|
|
|
|
|
|
|
|
|
|
| 464 |
self.assertTrue(all(required_kwargs <= call["kwargs"].keys() for call in calls))
|
| 465 |
self.assertTrue(all(call["kwargs"]["reference_length"] == 0 for call in calls))
|
| 466 |
self.assertEqual(
|
| 467 |
([call["kwargs"]["current_frame"] for call in calls[::2]], [call["kwargs"]["frame_memory_segments"]["target"] for call in calls[::2]]),
|
| 468 |
([2, 4], [2, 1]),
|
| 469 |
)
|
|
|
|
|
|
|
|
|
|
| 470 |
self.assertEqual(calls[0]["x"][:, 0, 0, 0, 0].tolist(), [0.0, 0.0, 0.0, 1.0, 0.0])
|
|
|
|
| 471 |
self.assertEqual(calls[0]["action_cond"][:, 0, 0].tolist(), [102.0, 103.0, 0.0, 0.0, 0.0])
|
| 472 |
self.assertEqual(calls[0]["kwargs"]["frame_idx"][:, 0].tolist(), [12, 13, 10, 11, 0])
|
|
|
|
| 473 |
self.assertEqual(calls[0]["kwargs"]["frame_memory_masks"]["revisit"].tolist(), [[False]])
|
|
|
|
| 474 |
self.assertAlmostEqual(float(loss), 1.0 / 3.0)
|
| 475 |
|
| 476 |
def test_dataset_returns_target_anchor_dynamic_revisit_contract(self):
|
|
|
|
| 460 |
calls = harness.diffusion_model.calls
|
| 461 |
required_kwargs = {"reference_length", "frame_memory_segments", "frame_memory_masks", "frame_memory_pose", "frame_idx", "image_hw"}
|
| 462 |
|
| 463 |
+
self.assertEqual(harness.horizons, [2, 1])
|
| 464 |
+
self.assertEqual(selection_calls, [([2, 3], "test", 0.75), ([4], "test", 0.75)])
|
| 465 |
+
self.assertEqual(len(selection_calls), len(harness.horizons))
|
| 466 |
+
self.assertEqual(len(calls), 4)
|
| 467 |
self.assertTrue(all(required_kwargs <= call["kwargs"].keys() for call in calls))
|
| 468 |
self.assertTrue(all(call["kwargs"]["reference_length"] == 0 for call in calls))
|
| 469 |
self.assertEqual(
|
| 470 |
([call["kwargs"]["current_frame"] for call in calls[::2]], [call["kwargs"]["frame_memory_segments"]["target"] for call in calls[::2]]),
|
| 471 |
([2, 4], [2, 1]),
|
| 472 |
)
|
| 473 |
+
self.assertEqual([call["kwargs"]["current_frame"] for call in calls], [2, 2, 4, 4])
|
| 474 |
+
self.assertTrue(torch.equal(calls[0]["kwargs"]["frame_idx"], calls[1]["kwargs"]["frame_idx"]))
|
| 475 |
+
self.assertTrue(torch.equal(calls[2]["kwargs"]["frame_idx"], calls[3]["kwargs"]["frame_idx"]))
|
| 476 |
self.assertEqual(calls[0]["x"][:, 0, 0, 0, 0].tolist(), [0.0, 0.0, 0.0, 1.0, 0.0])
|
| 477 |
+
self.assertEqual(calls[2]["x"][:, 0, 0, 0, 0].tolist(), [0.0, 0.0, 2.0, 2.0])
|
| 478 |
self.assertEqual(calls[0]["action_cond"][:, 0, 0].tolist(), [102.0, 103.0, 0.0, 0.0, 0.0])
|
| 479 |
self.assertEqual(calls[0]["kwargs"]["frame_idx"][:, 0].tolist(), [12, 13, 10, 11, 0])
|
| 480 |
+
self.assertEqual(calls[2]["kwargs"]["frame_idx"][:, 0].tolist(), [14, 10, 13, 12])
|
| 481 |
self.assertEqual(calls[0]["kwargs"]["frame_memory_masks"]["revisit"].tolist(), [[False]])
|
| 482 |
+
self.assertEqual(calls[2]["kwargs"]["frame_memory_masks"]["revisit"].tolist(), [[True]])
|
| 483 |
self.assertAlmostEqual(float(loss), 1.0 / 3.0)
|
| 484 |
|
| 485 |
def test_dataset_returns_target_anchor_dynamic_revisit_contract(self):
|