Download tests/test_dememwm_temporal_attention.py from BonanDing/DeMemWM: direct link, hf CLI and curl.
- Browser
- Download file 21.6 kB
-
https://huggingface.co/BonanDing/DeMemWM/resolve/main/tests/test_dememwm_temporal_attention.py
- Command line
-
hf download hf://BonanDing/DeMemWM/tests/test_dememwm_temporal_attention.py
-
curl -L -o test_dememwm_temporal_attention.py https://huggingface.co/BonanDing/DeMemWM/resolve/main/tests/test_dememwm_temporal_attention.py
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() | |