Validate DeMemWM DiT stream boundaries
Browse files
.exp_artifact/dememwm_remaining_fix_plan.md
CHANGED
|
@@ -232,7 +232,7 @@ Reject stale DeMemWM experiment flags
|
|
| 232 |
|
| 233 |
## Substep 4: Add Minimal DeMemWM DiT Boundary Checks
|
| 234 |
|
| 235 |
-
Status: `[
|
| 236 |
|
| 237 |
Bug:
|
| 238 |
|
|
|
|
| 232 |
|
| 233 |
## Substep 4: Add Minimal DeMemWM DiT Boundary Checks
|
| 234 |
|
| 235 |
+
Status: `[x]`
|
| 236 |
|
| 237 |
Bug:
|
| 238 |
|
algorithms/dememwm/models/dit.py
CHANGED
|
@@ -356,18 +356,42 @@ class SpatioTemporalDiTBlock(nn.Module):
|
|
| 356 |
anchor = int(frame_memory_segments.get("anchor", 0))
|
| 357 |
dynamic = int(frame_memory_segments.get("dynamic", 0))
|
| 358 |
revisit = int(frame_memory_segments.get("revisit", 0))
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 359 |
a0, a1 = target, target + anchor
|
| 360 |
d0, d1 = a1, a1 + dynamic
|
| 361 |
r0, r1 = d1, d1 + revisit
|
| 362 |
return x[:, :target], x[:, a0:a1], x[:, d0:d1], x[:, r0:r1]
|
| 363 |
|
| 364 |
def _frame_memory_stream_mask(self, frame_memory_masks, stream_name, stream_hidden):
|
|
|
|
| 365 |
if stream_hidden is None or int(stream_hidden.shape[1]) == 0:
|
| 366 |
return None
|
| 367 |
if frame_memory_masks is None or frame_memory_masks.get(stream_name) is None:
|
| 368 |
return None
|
| 369 |
return frame_memory_masks[stream_name].to(device=stream_hidden.device, dtype=torch.bool)
|
| 370 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 371 |
def _apply_frame_memory_reference_attention(self, x, c, frame_memory_segments, frame_memory_masks, frame_memory_geometry=None):
|
| 372 |
x_target, x_anchor, x_dynamic, x_revisit = self._split_frame_memory(x, frame_memory_segments)
|
| 373 |
if int(x_target.shape[1]) == 0:
|
|
@@ -412,6 +436,12 @@ class SpatioTemporalDiTBlock(nn.Module):
|
|
| 412 |
frame_memory_segments=None, frame_memory_masks=None, frame_memory_pose=None,
|
| 413 |
image_hw=None, frame_memory_geometry=None):
|
| 414 |
B, T, H, W, D = x.shape
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 415 |
|
| 416 |
# spatial block
|
| 417 |
|
|
|
|
| 356 |
anchor = int(frame_memory_segments.get("anchor", 0))
|
| 357 |
dynamic = int(frame_memory_segments.get("dynamic", 0))
|
| 358 |
revisit = int(frame_memory_segments.get("revisit", 0))
|
| 359 |
+
if min(target, anchor, dynamic, revisit) < 0:
|
| 360 |
+
raise ValueError(
|
| 361 |
+
"frame_memory_segments lengths must be nonnegative; "
|
| 362 |
+
f"got target={target}, anchor={anchor}, dynamic={dynamic}, revisit={revisit}"
|
| 363 |
+
)
|
| 364 |
+
total = target + anchor + dynamic + revisit
|
| 365 |
+
if total != int(x.shape[1]):
|
| 366 |
+
raise ValueError(
|
| 367 |
+
f"frame_memory_segments lengths sum to {total}, expected x.shape[1]={int(x.shape[1])}"
|
| 368 |
+
)
|
| 369 |
a0, a1 = target, target + anchor
|
| 370 |
d0, d1 = a1, a1 + dynamic
|
| 371 |
r0, r1 = d1, d1 + revisit
|
| 372 |
return x[:, :target], x[:, a0:a1], x[:, d0:d1], x[:, r0:r1]
|
| 373 |
|
| 374 |
def _frame_memory_stream_mask(self, frame_memory_masks, stream_name, stream_hidden):
|
| 375 |
+
self._check_frame_memory_stream_mask_shape(frame_memory_masks, stream_name, stream_hidden)
|
| 376 |
if stream_hidden is None or int(stream_hidden.shape[1]) == 0:
|
| 377 |
return None
|
| 378 |
if frame_memory_masks is None or frame_memory_masks.get(stream_name) is None:
|
| 379 |
return None
|
| 380 |
return frame_memory_masks[stream_name].to(device=stream_hidden.device, dtype=torch.bool)
|
| 381 |
|
| 382 |
+
def _check_frame_memory_stream_mask_shape(self, frame_memory_masks, stream_name, stream_hidden):
|
| 383 |
+
if stream_hidden is None:
|
| 384 |
+
return
|
| 385 |
+
if frame_memory_masks is None or frame_memory_masks.get(stream_name) is None:
|
| 386 |
+
return
|
| 387 |
+
stream_mask = frame_memory_masks[stream_name]
|
| 388 |
+
expected_shape = (int(stream_hidden.shape[0]), int(stream_hidden.shape[1]))
|
| 389 |
+
if tuple(stream_mask.shape) != expected_shape:
|
| 390 |
+
raise ValueError(
|
| 391 |
+
f"frame_memory_masks[{stream_name!r}] shape {tuple(stream_mask.shape)} "
|
| 392 |
+
f"must match {expected_shape}"
|
| 393 |
+
)
|
| 394 |
+
|
| 395 |
def _apply_frame_memory_reference_attention(self, x, c, frame_memory_segments, frame_memory_masks, frame_memory_geometry=None):
|
| 396 |
x_target, x_anchor, x_dynamic, x_revisit = self._split_frame_memory(x, frame_memory_segments)
|
| 397 |
if int(x_target.shape[1]) == 0:
|
|
|
|
| 436 |
frame_memory_segments=None, frame_memory_masks=None, frame_memory_pose=None,
|
| 437 |
image_hw=None, frame_memory_geometry=None):
|
| 438 |
B, T, H, W, D = x.shape
|
| 439 |
+
if frame_memory_segments is not None:
|
| 440 |
+
for stream_name, stream_hidden in zip(
|
| 441 |
+
("target", "anchor", "dynamic", "revisit"),
|
| 442 |
+
self._split_frame_memory(x, frame_memory_segments),
|
| 443 |
+
):
|
| 444 |
+
self._check_frame_memory_stream_mask_shape(frame_memory_masks, stream_name, stream_hidden)
|
| 445 |
|
| 446 |
# spatial block
|
| 447 |
|
tests/test_dememwm_temporal_attention.py
CHANGED
|
@@ -108,6 +108,70 @@ class DeMemWMTemporalAttentionTests(unittest.TestCase):
|
|
| 108 |
self.assertIs(spy.calls[0][0], segments)
|
| 109 |
self.assertIs(spy.calls[0][1], masks)
|
| 110 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 111 |
def test_dit_threads_frame_memory_metadata_to_blocks(self):
|
| 112 |
class SpyBlock(nn.Module):
|
| 113 |
def __init__(self):
|
|
|
|
| 108 |
self.assertIs(spy.calls[0][0], segments)
|
| 109 |
self.assertIs(spy.calls[0][1], masks)
|
| 110 |
|
| 111 |
+
def test_block_rejects_bad_frame_memory_segment_lengths_before_temporal_attention(self):
|
| 112 |
+
class ZeroAttention(nn.Module):
|
| 113 |
+
def forward(self, x):
|
| 114 |
+
return torch.zeros_like(x)
|
| 115 |
+
|
| 116 |
+
class UnexpectedTemporalAttention(nn.Module):
|
| 117 |
+
def forward(self, x, frame_memory_segments=None, frame_memory_masks=None):
|
| 118 |
+
raise AssertionError("temporal attention should not run before frame-memory validation")
|
| 119 |
+
|
| 120 |
+
block = SpatioTemporalDiTBlock(
|
| 121 |
+
hidden_size=4,
|
| 122 |
+
num_heads=1,
|
| 123 |
+
reference_length=0,
|
| 124 |
+
spatial_rotary_emb=None,
|
| 125 |
+
temporal_rotary_emb=None,
|
| 126 |
+
)
|
| 127 |
+
block.s_attn = ZeroAttention()
|
| 128 |
+
block.t_attn = UnexpectedTemporalAttention()
|
| 129 |
+
x = torch.zeros((1, 5, 1, 1, 4))
|
| 130 |
+
c = torch.zeros((1, 5, 4))
|
| 131 |
+
|
| 132 |
+
with self.assertRaisesRegex(ValueError, "lengths must be nonnegative"):
|
| 133 |
+
block(x, c, frame_memory_segments={"target": 2, "anchor": -1, "dynamic": 2, "revisit": 2})
|
| 134 |
+
|
| 135 |
+
with self.assertRaisesRegex(ValueError, r"sum to 6, expected x\.shape\[1\]=5"):
|
| 136 |
+
block(x, c, frame_memory_segments={"target": 2, "anchor": 1, "dynamic": 2, "revisit": 1})
|
| 137 |
+
|
| 138 |
+
def test_block_rejects_bad_frame_memory_stream_mask_shape_before_temporal_attention(self):
|
| 139 |
+
class ZeroAttention(nn.Module):
|
| 140 |
+
def forward(self, x):
|
| 141 |
+
return torch.zeros_like(x)
|
| 142 |
+
|
| 143 |
+
class UnexpectedTemporalAttention(nn.Module):
|
| 144 |
+
def forward(self, x, frame_memory_segments=None, frame_memory_masks=None):
|
| 145 |
+
raise AssertionError("temporal attention should not run before frame-memory validation")
|
| 146 |
+
|
| 147 |
+
block = SpatioTemporalDiTBlock(
|
| 148 |
+
hidden_size=4,
|
| 149 |
+
num_heads=1,
|
| 150 |
+
reference_length=0,
|
| 151 |
+
spatial_rotary_emb=None,
|
| 152 |
+
temporal_rotary_emb=None,
|
| 153 |
+
)
|
| 154 |
+
block.s_attn = ZeroAttention()
|
| 155 |
+
block.t_attn = UnexpectedTemporalAttention()
|
| 156 |
+
x = torch.zeros((1, 5, 1, 1, 4))
|
| 157 |
+
c = torch.zeros((1, 5, 4))
|
| 158 |
+
segments = {"target": 2, "anchor": 1, "dynamic": 1, "revisit": 1}
|
| 159 |
+
masks = {"dynamic": torch.ones((1, 2), dtype=torch.bool)}
|
| 160 |
+
|
| 161 |
+
with self.assertRaisesRegex(
|
| 162 |
+
ValueError,
|
| 163 |
+
r"frame_memory_masks\[\'dynamic\'\] shape \(1, 2\) must match \(1, 1\)",
|
| 164 |
+
):
|
| 165 |
+
block(x, c, frame_memory_segments=segments, frame_memory_masks=masks)
|
| 166 |
+
|
| 167 |
+
zero_dynamic_segments = {"target": 2, "anchor": 1, "dynamic": 0, "revisit": 2}
|
| 168 |
+
zero_dynamic_masks = {"dynamic": torch.ones((1, 1), dtype=torch.bool)}
|
| 169 |
+
with self.assertRaisesRegex(
|
| 170 |
+
ValueError,
|
| 171 |
+
r"frame_memory_masks\[\'dynamic\'\] shape \(1, 1\) must match \(1, 0\)",
|
| 172 |
+
):
|
| 173 |
+
block(x, c, frame_memory_segments=zero_dynamic_segments, frame_memory_masks=zero_dynamic_masks)
|
| 174 |
+
|
| 175 |
def test_dit_threads_frame_memory_metadata_to_blocks(self):
|
| 176 |
class SpyBlock(nn.Module):
|
| 177 |
def __init__(self):
|