DeMemWM / tests /test_dememwm_temporal_attention.py
BonanDing's picture
Use key-only DeMemWM pose geometry
5a627dd
Raw History Blame Contribute Delete
21.6 kB
import unittest
import torch
from torch import nn
from algorithms.dememwm.models.attention import TemporalAxialAttention
from algorithms.dememwm.models.dit import DiT, FrameMemoryReferenceAttention, SpatioTemporalDiTBlock
class IdentityRotary:
freqs = None
def rotate_queries_or_keys(self, x, freqs):
return x
def _averaging_temporal_attention(reference_length=0):
attn = TemporalAxialAttention(
dim=1,
heads=1,
dim_head=1,
reference_length=reference_length,
rotary_emb=IdentityRotary(),
)
with torch.no_grad():
attn.to_qkv.weight.zero_()
attn.to_qkv.weight[2, 0] = 1.0
attn.to_out.weight.fill_(1.0)
attn.to_out.bias.zero_()
return attn
class DeMemWMTemporalAttentionTests(unittest.TestCase):
def test_without_frame_memory_keeps_causal_reference_mask(self):
attn = _averaging_temporal_attention(reference_length=2)
x = torch.tensor([10.0, 20.0, 30.0, 100.0, 200.0]).view(1, 5, 1, 1, 1)
out = attn(x)
expected = torch.tensor([10.0, 15.0, 20.0, 100.0, 200.0])
self.assertTrue(torch.allclose(out.flatten(), expected))
def test_frame_memory_attention_bias_keeps_target_causal_and_streams_self_only(self):
attn = _averaging_temporal_attention()
segments = {"target": 3, "anchor": 2, "dynamic": 2, "revisit": 2}
total_frames = sum(segments.values())
bias = attn._frame_memory_attn_bias(
B=1,
T=total_frames,
H=1,
W=1,
dtype=torch.float32,
device=torch.device("cpu"),
frame_memory_segments=segments,
frame_memory_masks=None,
)[0, 0]
allow = torch.isfinite(bias)
target_expected = torch.tril(torch.ones((3, 3), dtype=torch.bool))
self.assertTrue(torch.equal(allow[:3, :3], target_expected))
self.assertFalse(allow[:3, 3:].any().item())
cursor = segments["target"]
for stream in ("anchor", "dynamic", "revisit"):
length = segments[stream]
rows = slice(cursor, cursor + length)
self.assertTrue(torch.equal(allow[rows, rows], torch.eye(length, dtype=torch.bool)))
self.assertFalse(allow[rows, :cursor].any().item())
self.assertFalse(allow[rows, cursor + length :].any().item())
cursor += length
def test_frame_memory_attention_bias_restores_invalid_memory_row_diagonal(self):
attn = _averaging_temporal_attention()
segments = {"target": 2, "anchor": 2, "dynamic": 2, "revisit": 2}
total_frames = sum(segments.values())
masks = {
"target": torch.ones((1, 2), dtype=torch.bool),
"anchor": torch.tensor([[False, True]]),
"dynamic": torch.tensor([[True, False]]),
"revisit": torch.tensor([[False, False]]),
}
bias = attn._frame_memory_attn_bias(
B=1,
T=total_frames,
H=1,
W=1,
dtype=torch.float32,
device=torch.device("cpu"),
frame_memory_segments=segments,
frame_memory_masks=masks,
)[0, 0]
allow = torch.isfinite(bias)
self.assertTrue(allow.any(dim=-1).all().item())
for row in (2, 5, 6, 7):
row_expected = torch.zeros(total_frames, dtype=torch.bool)
row_expected[row] = True
self.assertTrue(torch.equal(allow[row], row_expected))
def test_frame_memory_segments_mask_temporal_streams(self):
attn = _averaging_temporal_attention()
x = torch.tensor([10.0, 20.0, 30.0, 100.0, 1.0, 3.0, 200.0]).view(1, 7, 1, 1, 1)
segments = {"target": 3, "anchor": 1, "dynamic": 2, "revisit": 1}
out = attn(x, frame_memory_segments=segments)
expected = torch.tensor([10.0, 15.0, 20.0, 100.0, 1.0, 3.0, 200.0])
self.assertTrue(torch.allclose(out.flatten(), expected))
masks = {
"target": torch.ones((1, 3), dtype=torch.bool),
"anchor": torch.tensor([[False]]),
"dynamic": torch.tensor([[True, False]]),
"revisit": torch.ones((1, 1), dtype=torch.bool),
}
out = attn(x, frame_memory_segments=segments, frame_memory_masks=masks)
expected = torch.tensor([10.0, 15.0, 20.0, 100.0, 1.0, 3.0, 200.0])
self.assertTrue(torch.allclose(out.flatten(), expected))
def test_frame_memory_temporal_targets_do_not_depend_on_appended_memory(self):
attn = _averaging_temporal_attention()
segments = {"target": 3, "anchor": 1, "dynamic": 2, "revisit": 1}
x = torch.tensor([10.0, 20.0, 30.0, 100.0, 1.0, 3.0, 200.0]).view(1, 7, 1, 1, 1)
changed_memory = torch.tensor([10.0, 20.0, 30.0, -500.0, 700.0, 900.0, -300.0]).view(1, 7, 1, 1, 1)
out = attn(x, frame_memory_segments=segments)
changed_out = attn(changed_memory, frame_memory_segments=segments)
expected_target = torch.tensor([10.0, 15.0, 20.0]).view(1, 3, 1, 1, 1)
self.assertTrue(torch.allclose(out[:, :3], expected_target))
self.assertTrue(torch.allclose(out[:, :3], changed_out[:, :3]))
def test_block_threads_frame_memory_metadata_to_temporal_attention(self):
class ZeroAttention(nn.Module):
def forward(self, x):
return torch.zeros_like(x)
class SpyTemporalAttention(nn.Module):
def __init__(self):
super().__init__()
self.calls = []
def forward(self, x, frame_memory_segments=None, frame_memory_masks=None):
self.calls.append((frame_memory_segments, frame_memory_masks))
return torch.zeros_like(x)
block = SpatioTemporalDiTBlock(
hidden_size=4,
num_heads=1,
reference_length=0,
spatial_rotary_emb=None,
temporal_rotary_emb=None,
)
spy = SpyTemporalAttention()
block.s_attn = ZeroAttention()
block.t_attn = spy
x = torch.zeros((1, 5, 1, 1, 4))
c = torch.zeros((1, 5, 4))
segments = {"target": 2, "anchor": 1, "dynamic": 1, "revisit": 1}
masks = {"dynamic": torch.ones((1, 1), dtype=torch.bool)}
block(x, c, frame_memory_segments=segments, frame_memory_masks=masks)
self.assertIs(spy.calls[0][0], segments)
self.assertIs(spy.calls[0][1], masks)
def test_block_rejects_bad_frame_memory_segment_lengths_before_temporal_attention(self):
class ZeroAttention(nn.Module):
def forward(self, x):
return torch.zeros_like(x)
class UnexpectedTemporalAttention(nn.Module):
def forward(self, x, frame_memory_segments=None, frame_memory_masks=None):
raise AssertionError("temporal attention should not run before frame-memory validation")
block = SpatioTemporalDiTBlock(
hidden_size=4,
num_heads=1,
reference_length=0,
spatial_rotary_emb=None,
temporal_rotary_emb=None,
)
block.s_attn = ZeroAttention()
block.t_attn = UnexpectedTemporalAttention()
x = torch.zeros((1, 5, 1, 1, 4))
c = torch.zeros((1, 5, 4))
with self.assertRaisesRegex(ValueError, "lengths must be nonnegative"):
block(x, c, frame_memory_segments={"target": 2, "anchor": -1, "dynamic": 2, "revisit": 2})
with self.assertRaisesRegex(ValueError, r"sum to 6, expected x\.shape\[1\]=5"):
block(x, c, frame_memory_segments={"target": 2, "anchor": 1, "dynamic": 2, "revisit": 1})
def test_block_rejects_bad_frame_memory_stream_mask_shape_before_temporal_attention(self):
class ZeroAttention(nn.Module):
def forward(self, x):
return torch.zeros_like(x)
class UnexpectedTemporalAttention(nn.Module):
def forward(self, x, frame_memory_segments=None, frame_memory_masks=None):
raise AssertionError("temporal attention should not run before frame-memory validation")
block = SpatioTemporalDiTBlock(
hidden_size=4,
num_heads=1,
reference_length=0,
spatial_rotary_emb=None,
temporal_rotary_emb=None,
)
block.s_attn = ZeroAttention()
block.t_attn = UnexpectedTemporalAttention()
x = torch.zeros((1, 5, 1, 1, 4))
c = torch.zeros((1, 5, 4))
segments = {"target": 2, "anchor": 1, "dynamic": 1, "revisit": 1}
masks = {"dynamic": torch.ones((1, 2), dtype=torch.bool)}
with self.assertRaisesRegex(
ValueError,
r"frame_memory_masks\[\'dynamic\'\] shape \(1, 2\) must match \(1, 1\)",
):
block(x, c, frame_memory_segments=segments, frame_memory_masks=masks)
zero_dynamic_segments = {"target": 2, "anchor": 1, "dynamic": 0, "revisit": 2}
zero_dynamic_masks = {"dynamic": torch.ones((1, 1), dtype=torch.bool)}
with self.assertRaisesRegex(
ValueError,
r"frame_memory_masks\[\'dynamic\'\] shape \(1, 1\) must match \(1, 0\)",
):
block(x, c, frame_memory_segments=zero_dynamic_segments, frame_memory_masks=zero_dynamic_masks)
def test_dit_threads_frame_memory_metadata_to_blocks(self):
class SpyBlock(nn.Module):
def __init__(self):
super().__init__()
self.calls = []
def forward(self, x, c, **kwargs):
self.calls.append(kwargs)
return x
model = DiT(
input_h=2,
input_w=2,
patch_size=1,
in_channels=1,
hidden_size=8,
depth=1,
num_heads=1,
mlp_ratio=1.0,
action_cond_dim=3,
pose_cond_dim=0,
reference_length=0,
use_memory_attention=False,
)
spy = SpyBlock()
model.blocks[0] = spy
model.eval()
x = torch.zeros((1, 5, 1, 2, 2))
t = torch.zeros((1, 5), dtype=torch.long)
action_cond = torch.zeros((1, 5, 3))
segments = {"target": 2, "anchor": 1, "dynamic": 1, "revisit": 1}
masks = {"dynamic": torch.ones((1, 1), dtype=torch.bool)}
with torch.no_grad():
out = model(x, t, action_cond, frame_memory_segments=segments, frame_memory_masks=masks)
self.assertEqual(tuple(out.shape), (1, 2, 1, 2, 2))
self.assertIs(spy.calls[0]["frame_memory_segments"], segments)
self.assertIs(spy.calls[0]["frame_memory_masks"], masks)
def test_dit_baseline_and_packed_frame_memory_output_lengths_with_real_blocks(self):
torch.manual_seed(0)
model = DiT(
input_h=2,
input_w=2,
patch_size=1,
in_channels=1,
hidden_size=8,
depth=1,
num_heads=1,
mlp_ratio=1.0,
action_cond_dim=3,
pose_cond_dim=0,
reference_length=1,
use_memory_attention=True,
)
model.eval()
x = torch.randn((1, 3, 1, 2, 2))
t = torch.zeros((1, 3), dtype=torch.long)
action_cond = torch.zeros((1, 3, 3))
with torch.no_grad():
baseline = model(x, t, action_cond, reference_length=1)
self.assertEqual(tuple(baseline.shape), (1, 3, 1, 2, 2))
packed = torch.randn((1, 5, 1, 2, 2))
packed_t = torch.zeros((1, 5), dtype=torch.long)
packed_action = torch.zeros((1, 5, 3))
segments = {"target": 2, "anchor": 1, "dynamic": 1, "revisit": 1}
masks = {
"target": torch.ones((1, 2), dtype=torch.bool),
"anchor": torch.ones((1, 1), dtype=torch.bool),
"dynamic": torch.zeros((1, 1), dtype=torch.bool),
"revisit": torch.ones((1, 1), dtype=torch.bool),
}
with torch.no_grad():
packed_out = model(
packed,
packed_t,
packed_action,
reference_length=1,
frame_memory_segments=segments,
frame_memory_masks=masks,
)
self.assertEqual(tuple(packed_out.shape), (1, 2, 1, 2, 2))
def test_dit_has_no_worldmem_reference_attention_fallback(self):
model = DiT(
input_h=2,
input_w=2,
patch_size=1,
in_channels=1,
hidden_size=8,
depth=1,
num_heads=1,
mlp_ratio=1.0,
action_cond_dim=3,
pose_cond_dim=0,
reference_length=1,
use_memory_attention=True,
)
block = model.blocks[0]
self.assertFalse(hasattr(block, "r_attn"))
self.assertFalse(hasattr(block, "pose_cond_mlp"))
calls = []
class SpyReferenceAttention(nn.Module):
def forward(self, target_hidden, memory_hidden, memory_mask=None, geometry_cache=None):
calls.append(memory_hidden)
return torch.zeros_like(target_hidden)
block.r_attn_anchor = SpyReferenceAttention()
block.r_attn_dynamic = SpyReferenceAttention()
block.r_attn_revisit = SpyReferenceAttention()
x = torch.randn((1, 3, 1, 2, 2))
t = torch.zeros((1, 3), dtype=torch.long)
action_cond = torch.zeros((1, 3, 3))
with torch.no_grad():
out = model(x, t, action_cond, reference_length=1)
self.assertEqual(tuple(out.shape), (1, 3, 1, 2, 2))
self.assertEqual(calls, [])
def test_dit_builds_geometry_cache_once_and_threads_to_reference_attention(self):
class ZeroAttention(nn.Module):
def forward(self, x):
return torch.zeros_like(x)
class ZeroTemporalAttention(nn.Module):
def forward(self, x, frame_memory_segments=None, frame_memory_masks=None):
return torch.zeros_like(x)
class SpyReferenceAttention(nn.Module):
def __init__(self):
super().__init__()
self.geometry = []
def forward(self, target_hidden, memory_hidden, memory_mask=None, geometry_cache=None):
self.geometry.append(geometry_cache)
return torch.zeros_like(target_hidden)
model = DiT(
input_h=2,
input_w=2,
patch_size=1,
in_channels=1,
hidden_size=8,
depth=2,
num_heads=1,
mlp_ratio=1.0,
action_cond_dim=3,
pose_cond_dim=0,
reference_length=0,
use_plucker=True,
use_memory_attention=True,
)
spies = []
for block in model.blocks:
block.s_attn = ZeroAttention()
block.t_attn = ZeroTemporalAttention()
anchor = SpyReferenceAttention()
dynamic = SpyReferenceAttention()
revisit = SpyReferenceAttention()
block.r_attn_anchor = anchor
block.r_attn_dynamic = dynamic
block.r_attn_revisit = revisit
spies.append((anchor, dynamic, revisit))
model.eval()
x = torch.zeros((1, 5, 1, 2, 2))
t = torch.zeros((1, 5), dtype=torch.long)
action_cond = torch.zeros((1, 5, 3))
frame_memory_pose = torch.zeros((1, 5, 5))
frame_memory_pose[0, :, 0] = torch.arange(5, dtype=torch.float32)
image_hw = torch.tensor([[4, 8]], dtype=torch.long)
segments = {"target": 2, "anchor": 1, "dynamic": 1, "revisit": 1}
masks = {"dynamic": torch.tensor([[False]])}
with torch.no_grad():
model(
x,
t,
action_cond,
frame_memory_segments=segments,
frame_memory_masks=masks,
frame_memory_pose=frame_memory_pose,
image_hw=image_hw,
)
anchor_cache = spies[0][0].geometry[0]
self.assertIs(anchor_cache, spies[1][0].geometry[0])
self.assertEqual(tuple(anchor_cache["query_rays"].shape), (1, 2, 2, 2, 6))
self.assertEqual(tuple(anchor_cache["relative_rays"].shape), (1, 2, 1, 2, 2, 6))
self.assertTrue(torch.equal(anchor_cache["image_hw"], torch.tensor([[4.0, 8.0]])))
self.assertTrue(torch.allclose(anchor_cache["intrinsics"], torch.tensor([[2.8, 1.4, 4.0, 2.0]])))
self.assertTrue(torch.equal(anchor_cache["relative_pose"][0, :, 0, 0], torch.tensor([2.0, 1.0])))
self.assertTrue(torch.equal(spies[0][1].geometry[0]["mask"], masks["dynamic"]))
def test_frame_memory_geometry_uses_fp32_math_then_returns_activation_dtype(self):
model = DiT(
input_h=2,
input_w=2,
patch_size=1,
in_channels=1,
hidden_size=8,
depth=1,
num_heads=1,
mlp_ratio=1.0,
action_cond_dim=3,
pose_cond_dim=0,
reference_length=0,
use_plucker=True,
use_memory_attention=True,
)
device = torch.device("cpu")
segments = {"target": 1, "anchor": 1, "dynamic": 0, "revisit": 0}
frame_memory_pose = torch.zeros((1, 2, 5), device=device, dtype=torch.float32)
frame_memory_pose[0, 0, 0] = 10000.0
frame_memory_pose[0, 1, 0] = 10001.0
image_hw = torch.tensor([[4, 8]], device=device, dtype=torch.long)
with torch.autocast(device_type=device.type, dtype=torch.bfloat16):
cache = model._build_frame_memory_geometry(
segments,
None,
frame_memory_pose,
image_hw,
grid_h=2,
grid_w=2,
device=device,
dtype=torch.float16,
)
anchor_cache = cache["anchor"]
for tensor in (
cache["query_rays"],
cache["target_pose"],
cache["image_hw"],
cache["intrinsics"],
anchor_cache["relative_pose"],
anchor_cache["relative_c2w"],
anchor_cache["relative_rays"],
):
self.assertEqual(tensor.dtype, torch.float16)
self.assertTrue(torch.isfinite(tensor).all().item())
self.assertEqual(anchor_cache["relative_pose"][0, 0, 0, 0].item(), 1.0)
self.assertEqual(anchor_cache["relative_c2w"][0, 0, 0, 0, 3].item(), 1.0)
def test_frame_memory_reference_attention_masks_padded_keys(self):
attn = FrameMemoryReferenceAttention(hidden_size=1, num_heads=1)
with torch.no_grad():
attn.to_q.weight.zero_()
attn.to_k.weight.zero_()
attn.to_v.weight.fill_(1.0)
attn.out_proj.weight.fill_(1.0)
attn.out_proj.bias.zero_()
target = torch.zeros((1, 1, 1, 1, 1))
memory = torch.tensor([1.0, 10.0]).view(1, 2, 1, 1, 1)
first_only = attn(target, memory, torch.tensor([[True, False]]))
second_only = attn(target, memory, torch.tensor([[False, True]]))
empty = attn(target, memory, torch.zeros((1, 2), dtype=torch.bool))
self.assertTrue(torch.allclose(first_only, torch.ones_like(first_only)))
self.assertTrue(torch.allclose(second_only, torch.full_like(second_only, 10.0)))
self.assertTrue(torch.equal(empty, torch.zeros_like(empty)))
def test_frame_memory_reference_attention_uses_key_geometry_only(self):
attn = FrameMemoryReferenceAttention(hidden_size=2, num_heads=1)
with torch.no_grad():
attn.to_q.weight.zero_()
attn.to_q.weight[0, 0] = 1.0
attn.to_k.weight.zero_()
attn.to_k.weight[0, 0] = 1.0
attn.to_v.weight.zero_()
attn.to_v.weight[0, 1] = 1.0
attn.query_pose_proj.weight.zero_()
attn.query_pose_proj.bias.zero_()
attn.query_pose_proj.weight[0, 5] = 100.0
attn.key_pose_proj.weight.zero_()
attn.key_pose_proj.bias.zero_()
attn.key_pose_proj.weight[0, 5] = 1.0
attn.out_proj.weight.zero_()
attn.out_proj.weight[0, 0] = 1.0
attn.out_proj.bias.zero_()
target = torch.zeros((1, 1, 1, 1, 2))
target[0, 0, 0, 0, 0] = 1.0
memory = torch.zeros((1, 2, 1, 1, 2))
memory[0, :, 0, 0, 1] = torch.tensor([1.0, 10.0])
query_rays = torch.zeros((1, 1, 1, 1, 6))
query_rays_changed = query_rays.clone()
query_rays_changed[..., 5] = 10.0
relative_a = torch.zeros((1, 1, 2, 1, 1, 6))
relative_b = torch.zeros_like(relative_a)
relative_a[0, 0, 0, 0, 0, 5] = 2.0
relative_b[0, 0, 1, 0, 0, 5] = 2.0
out_a = attn(
target,
memory,
geometry_cache={"query_rays": query_rays, "relative_rays": relative_a},
)
out_a_query_changed = attn(
target,
memory,
geometry_cache={"query_rays": query_rays_changed, "relative_rays": relative_a},
)
out_b = attn(
target,
memory,
geometry_cache={"query_rays": query_rays, "relative_rays": relative_b},
)
self.assertTrue(torch.allclose(out_a, out_a_query_changed))
self.assertLess(out_a[0, 0, 0, 0, 0].item(), out_b[0, 0, 0, 0, 0].item())
self.assertFalse(torch.allclose(out_a, out_b))
if __name__ == "__main__":
unittest.main()