BonanDing commited on
Commit
824571d
·
1 Parent(s): b285586

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