Add DeMemWM geometry-aware memory attention
Browse files
algorithms/dememwm/models/dit.py
CHANGED
|
@@ -150,11 +150,39 @@ class FrameMemoryReferenceAttention(nn.Module):
|
|
| 150 |
self.to_q = nn.Linear(hidden_size, hidden_size, bias=False)
|
| 151 |
self.to_k = nn.Linear(hidden_size, hidden_size, bias=False)
|
| 152 |
self.to_v = nn.Linear(hidden_size, hidden_size, bias=False)
|
|
|
|
|
|
|
| 153 |
self.out_proj = nn.Linear(hidden_size, hidden_size, bias=True)
|
| 154 |
|
| 155 |
def _split_heads(self, x):
|
| 156 |
return rearrange(x, "b n (h d) -> b h n d", h=self.num_heads)
|
| 157 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 158 |
def forward(self, target_hidden, memory_hidden, memory_mask=None, geometry_cache=None):
|
| 159 |
B, T_target, H, W, D = target_hidden.shape
|
| 160 |
if memory_hidden is None or int(memory_hidden.shape[1]) == 0:
|
|
@@ -162,18 +190,34 @@ class FrameMemoryReferenceAttention(nn.Module):
|
|
| 162 |
|
| 163 |
T_memory = memory_hidden.shape[1]
|
| 164 |
P = H * W
|
| 165 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 166 |
q = self._split_heads(self.to_q(q))
|
| 167 |
|
| 168 |
-
memory_tokens = rearrange(memory_hidden, "b m h w d -> b
|
| 169 |
-
|
| 170 |
-
|
| 171 |
-
|
| 172 |
-
|
| 173 |
-
|
| 174 |
-
|
| 175 |
-
|
| 176 |
-
|
|
|
|
|
|
|
| 177 |
|
| 178 |
if memory_mask is None:
|
| 179 |
frame_valid = torch.ones((B, T_memory), device=target_hidden.device, dtype=torch.bool)
|
|
@@ -184,13 +228,22 @@ class FrameMemoryReferenceAttention(nn.Module):
|
|
| 184 |
key_valid = torch.cat([key_valid, ~active[:, None]], dim=1)
|
| 185 |
|
| 186 |
row_batch = torch.arange(B, device=target_hidden.device).repeat_interleave(T_target)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 187 |
attn_bias = torch.zeros((B * T_target, 1, 1, key_valid.shape[1]), device=q.device, dtype=q.dtype)
|
| 188 |
attn_bias = attn_bias.masked_fill(~key_valid[row_batch][:, None, None], float("-inf"))
|
| 189 |
|
| 190 |
x = F.scaled_dot_product_attention(
|
| 191 |
query=q.contiguous(),
|
| 192 |
-
key=k.
|
| 193 |
-
value=v.
|
| 194 |
attn_mask=attn_bias,
|
| 195 |
)
|
| 196 |
x = rearrange(x, "n h p d -> n p (h d)")
|
|
@@ -658,16 +711,18 @@ class DiT(nn.Module):
|
|
| 658 |
else:
|
| 659 |
pc = None
|
| 660 |
|
| 661 |
-
frame_memory_geometry =
|
| 662 |
-
|
| 663 |
-
|
| 664 |
-
|
| 665 |
-
|
| 666 |
-
|
| 667 |
-
|
| 668 |
-
|
| 669 |
-
|
| 670 |
-
|
|
|
|
|
|
|
| 671 |
|
| 672 |
for i, block in enumerate(self.blocks):
|
| 673 |
x = block(x, c, current_frame=current_frame, timestep=t, is_last_block= (i+1 == len(self.blocks)),
|
|
|
|
| 150 |
self.to_q = nn.Linear(hidden_size, hidden_size, bias=False)
|
| 151 |
self.to_k = nn.Linear(hidden_size, hidden_size, bias=False)
|
| 152 |
self.to_v = nn.Linear(hidden_size, hidden_size, bias=False)
|
| 153 |
+
self.query_pose_proj = nn.Linear(6, hidden_size, bias=True)
|
| 154 |
+
self.key_pose_proj = nn.Linear(6, hidden_size, bias=True)
|
| 155 |
self.out_proj = nn.Linear(hidden_size, hidden_size, bias=True)
|
| 156 |
|
| 157 |
def _split_heads(self, x):
|
| 158 |
return rearrange(x, "b n (h d) -> b h n d", h=self.num_heads)
|
| 159 |
|
| 160 |
+
def _geometry_rays(
|
| 161 |
+
self,
|
| 162 |
+
geometry_cache,
|
| 163 |
+
B,
|
| 164 |
+
T_target,
|
| 165 |
+
T_memory,
|
| 166 |
+
H,
|
| 167 |
+
W,
|
| 168 |
+
device,
|
| 169 |
+
dtype,
|
| 170 |
+
):
|
| 171 |
+
if not isinstance(geometry_cache, dict):
|
| 172 |
+
return None, None
|
| 173 |
+
query_rays = geometry_cache.get("query_rays")
|
| 174 |
+
relative_rays = geometry_cache.get("relative_rays")
|
| 175 |
+
if not torch.is_tensor(query_rays) or not torch.is_tensor(relative_rays):
|
| 176 |
+
return None, None
|
| 177 |
+
if tuple(query_rays.shape) != (B, T_target, H, W, 6):
|
| 178 |
+
return None, None
|
| 179 |
+
if tuple(relative_rays.shape) != (B, T_target, T_memory, H, W, 6):
|
| 180 |
+
return None, None
|
| 181 |
+
return (
|
| 182 |
+
query_rays.to(device=device, dtype=dtype),
|
| 183 |
+
relative_rays.to(device=device, dtype=dtype),
|
| 184 |
+
)
|
| 185 |
+
|
| 186 |
def forward(self, target_hidden, memory_hidden, memory_mask=None, geometry_cache=None):
|
| 187 |
B, T_target, H, W, D = target_hidden.shape
|
| 188 |
if memory_hidden is None or int(memory_hidden.shape[1]) == 0:
|
|
|
|
| 190 |
|
| 191 |
T_memory = memory_hidden.shape[1]
|
| 192 |
P = H * W
|
| 193 |
+
query_rays, relative_rays = self._geometry_rays(
|
| 194 |
+
geometry_cache,
|
| 195 |
+
B,
|
| 196 |
+
T_target,
|
| 197 |
+
T_memory,
|
| 198 |
+
H,
|
| 199 |
+
W,
|
| 200 |
+
target_hidden.device,
|
| 201 |
+
target_hidden.dtype,
|
| 202 |
+
)
|
| 203 |
+
q_input = target_hidden
|
| 204 |
+
if query_rays is not None:
|
| 205 |
+
# Values stay visual-only; Plucker features only steer Q/K matching.
|
| 206 |
+
q_input = q_input + self.query_pose_proj(query_rays)
|
| 207 |
+
q = rearrange(q_input, "b t h w d -> (b t) (h w) d")
|
| 208 |
q = self._split_heads(self.to_q(q))
|
| 209 |
|
| 210 |
+
memory_tokens = rearrange(memory_hidden, "b m h w d -> b m (h w) d")
|
| 211 |
+
memory_flat = rearrange(memory_tokens, "b m p d -> b (m p) d")
|
| 212 |
+
if relative_rays is None:
|
| 213 |
+
k = self._split_heads(self.to_k(memory_flat))
|
| 214 |
+
else:
|
| 215 |
+
key_pose = self.key_pose_proj(
|
| 216 |
+
rearrange(relative_rays, "b t m h w r -> b t m (h w) r")
|
| 217 |
+
)
|
| 218 |
+
k = self.to_k(memory_tokens[:, None] + key_pose)
|
| 219 |
+
k = self._split_heads(rearrange(k, "b t m p d -> (b t) (m p) d"))
|
| 220 |
+
v = self._split_heads(self.to_v(memory_flat))
|
| 221 |
|
| 222 |
if memory_mask is None:
|
| 223 |
frame_valid = torch.ones((B, T_memory), device=target_hidden.device, dtype=torch.bool)
|
|
|
|
| 228 |
key_valid = torch.cat([key_valid, ~active[:, None]], dim=1)
|
| 229 |
|
| 230 |
row_batch = torch.arange(B, device=target_hidden.device).repeat_interleave(T_target)
|
| 231 |
+
if relative_rays is None:
|
| 232 |
+
k = k.index_select(0, row_batch)
|
| 233 |
+
# A zero dummy key keeps all-padded samples finite; the projected delta
|
| 234 |
+
# is masked back to zero below so padded memory cannot contribute.
|
| 235 |
+
dummy_k = k.new_zeros((k.shape[0], self.num_heads, 1, self.head_dim))
|
| 236 |
+
k = torch.cat([k, dummy_k], dim=2)
|
| 237 |
+
dummy_v = v.new_zeros((B, self.num_heads, 1, self.head_dim))
|
| 238 |
+
v = torch.cat([v, dummy_v], dim=2).index_select(0, row_batch)
|
| 239 |
+
|
| 240 |
attn_bias = torch.zeros((B * T_target, 1, 1, key_valid.shape[1]), device=q.device, dtype=q.dtype)
|
| 241 |
attn_bias = attn_bias.masked_fill(~key_valid[row_batch][:, None, None], float("-inf"))
|
| 242 |
|
| 243 |
x = F.scaled_dot_product_attention(
|
| 244 |
query=q.contiguous(),
|
| 245 |
+
key=k.contiguous(),
|
| 246 |
+
value=v.contiguous(),
|
| 247 |
attn_mask=attn_bias,
|
| 248 |
)
|
| 249 |
x = rearrange(x, "n h p d -> n p (h d)")
|
|
|
|
| 711 |
else:
|
| 712 |
pc = None
|
| 713 |
|
| 714 |
+
frame_memory_geometry = None
|
| 715 |
+
if self.use_memory_attention and self.use_plucker:
|
| 716 |
+
frame_memory_geometry = self._build_frame_memory_geometry(
|
| 717 |
+
frame_memory_segments,
|
| 718 |
+
frame_memory_masks,
|
| 719 |
+
frame_memory_pose,
|
| 720 |
+
image_hw,
|
| 721 |
+
x.shape[2],
|
| 722 |
+
x.shape[3],
|
| 723 |
+
x.device,
|
| 724 |
+
x.dtype,
|
| 725 |
+
)
|
| 726 |
|
| 727 |
for i, block in enumerate(self.blocks):
|
| 728 |
x = block(x, c, current_frame=current_frame, timestep=t, is_last_block= (i+1 == len(self.blocks)),
|
tests/test_dememwm_temporal_attention.py
CHANGED
|
@@ -264,6 +264,7 @@ class DeMemWMTemporalAttentionTests(unittest.TestCase):
|
|
| 264 |
action_cond_dim=3,
|
| 265 |
pose_cond_dim=0,
|
| 266 |
reference_length=0,
|
|
|
|
| 267 |
use_memory_attention=True,
|
| 268 |
)
|
| 269 |
spies = []
|
|
@@ -328,5 +329,48 @@ class DeMemWMTemporalAttentionTests(unittest.TestCase):
|
|
| 328 |
self.assertTrue(torch.allclose(second_only, torch.full_like(second_only, 10.0)))
|
| 329 |
self.assertTrue(torch.equal(empty, torch.zeros_like(empty)))
|
| 330 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 331 |
if __name__ == "__main__":
|
| 332 |
unittest.main()
|
|
|
|
| 264 |
action_cond_dim=3,
|
| 265 |
pose_cond_dim=0,
|
| 266 |
reference_length=0,
|
| 267 |
+
use_plucker=True,
|
| 268 |
use_memory_attention=True,
|
| 269 |
)
|
| 270 |
spies = []
|
|
|
|
| 329 |
self.assertTrue(torch.allclose(second_only, torch.full_like(second_only, 10.0)))
|
| 330 |
self.assertTrue(torch.equal(empty, torch.zeros_like(empty)))
|
| 331 |
|
| 332 |
+
def test_frame_memory_reference_attention_uses_geometry_qk(self):
|
| 333 |
+
attn = FrameMemoryReferenceAttention(hidden_size=2, num_heads=1)
|
| 334 |
+
with torch.no_grad():
|
| 335 |
+
attn.to_q.weight.zero_()
|
| 336 |
+
attn.to_q.weight[0, 0] = 1.0
|
| 337 |
+
attn.to_k.weight.zero_()
|
| 338 |
+
attn.to_k.weight[0, 0] = 1.0
|
| 339 |
+
attn.to_v.weight.zero_()
|
| 340 |
+
attn.to_v.weight[0, 1] = 1.0
|
| 341 |
+
attn.query_pose_proj.weight.zero_()
|
| 342 |
+
attn.query_pose_proj.bias.zero_()
|
| 343 |
+
attn.query_pose_proj.weight[0, 5] = 1.0
|
| 344 |
+
attn.key_pose_proj.weight.zero_()
|
| 345 |
+
attn.key_pose_proj.bias.zero_()
|
| 346 |
+
attn.key_pose_proj.weight[0, 5] = 1.0
|
| 347 |
+
attn.out_proj.weight.zero_()
|
| 348 |
+
attn.out_proj.weight[0, 0] = 1.0
|
| 349 |
+
attn.out_proj.bias.zero_()
|
| 350 |
+
|
| 351 |
+
target = torch.zeros((1, 1, 1, 1, 2))
|
| 352 |
+
memory = torch.zeros((1, 2, 1, 1, 2))
|
| 353 |
+
memory[0, :, 0, 0, 1] = torch.tensor([1.0, 10.0])
|
| 354 |
+
query_rays = torch.zeros((1, 1, 1, 1, 6))
|
| 355 |
+
query_rays[..., 5] = 1.0
|
| 356 |
+
relative_a = torch.zeros((1, 1, 2, 1, 1, 6))
|
| 357 |
+
relative_b = torch.zeros_like(relative_a)
|
| 358 |
+
relative_a[0, 0, 0, 0, 0, 5] = 2.0
|
| 359 |
+
relative_b[0, 0, 1, 0, 0, 5] = 2.0
|
| 360 |
+
|
| 361 |
+
out_a = attn(
|
| 362 |
+
target,
|
| 363 |
+
memory,
|
| 364 |
+
geometry_cache={"query_rays": query_rays, "relative_rays": relative_a},
|
| 365 |
+
)
|
| 366 |
+
out_b = attn(
|
| 367 |
+
target,
|
| 368 |
+
memory,
|
| 369 |
+
geometry_cache={"query_rays": query_rays, "relative_rays": relative_b},
|
| 370 |
+
)
|
| 371 |
+
|
| 372 |
+
self.assertLess(out_a[0, 0, 0, 0, 0].item(), out_b[0, 0, 0, 0, 0].item())
|
| 373 |
+
self.assertFalse(torch.allclose(out_a, out_b))
|
| 374 |
+
|
| 375 |
if __name__ == "__main__":
|
| 376 |
unittest.main()
|