BonanDing commited on
Commit
5a627dd
·
1 Parent(s): 63bd9b0

Use key-only DeMemWM pose geometry

Browse files

Disable DeMemWM timestamp embedding by default and route frame-memory pose geometry to memory keys while leaving query features unconditioned.

algorithms/dememwm/df_video.py CHANGED
@@ -901,7 +901,8 @@ class DeMemWMMinecraft(DiffusionForcingBase):
901
  self.relative_embedding = getattr(cfg, "relative_embedding", True)
902
  self.state_embed_only_on_qk = getattr(cfg, "state_embed_only_on_qk", True)
903
  self.use_memory_attention = getattr(cfg, "use_memory_attention", True)
904
- self.add_timestamp_embedding = getattr(cfg, "add_timestamp_embedding", True)
 
905
  self.ref_mode = getattr(cfg, "ref_mode", 'sequential')
906
  self.log_curve = getattr(cfg, "log_curve", False)
907
  self.focal_length = getattr(cfg, "focal_length", 0.35)
@@ -951,6 +952,7 @@ class DeMemWMMinecraft(DiffusionForcingBase):
951
  add_timestamp_embedding=self.add_timestamp_embedding,
952
  ref_mode=self.ref_mode,
953
  focal_length=self.focal_length,
 
954
  )
955
 
956
  self.validation_lpips_model = LearnedPerceptualImagePatchSimilarity(sync_on_compute=False)
 
901
  self.relative_embedding = getattr(cfg, "relative_embedding", True)
902
  self.state_embed_only_on_qk = getattr(cfg, "state_embed_only_on_qk", True)
903
  self.use_memory_attention = getattr(cfg, "use_memory_attention", True)
904
+ self.add_timestamp_embedding = getattr(cfg, "add_timestamp_embedding", False)
905
+ self.memory_attention_key_only_geometry = getattr(cfg, "memory_attention_key_only_geometry", True)
906
  self.ref_mode = getattr(cfg, "ref_mode", 'sequential')
907
  self.log_curve = getattr(cfg, "log_curve", False)
908
  self.focal_length = getattr(cfg, "focal_length", 0.35)
 
952
  add_timestamp_embedding=self.add_timestamp_embedding,
953
  ref_mode=self.ref_mode,
954
  focal_length=self.focal_length,
955
+ memory_attention_key_only_geometry=self.memory_attention_key_only_geometry,
956
  )
957
 
958
  self.validation_lpips_model = LearnedPerceptualImagePatchSimilarity(sync_on_compute=False)
algorithms/dememwm/models/diffusion.py CHANGED
@@ -31,6 +31,7 @@ class Diffusion(nn.Module):
31
  add_timestamp_embedding=False,
32
  ref_mode='sequential',
33
  focal_length=0.35,
 
34
  ):
35
  super().__init__()
36
  self.cfg = cfg
@@ -60,6 +61,7 @@ class Diffusion(nn.Module):
60
  self.add_timestamp_embedding = add_timestamp_embedding
61
  self.ref_mode = ref_mode
62
  self.focal_length = focal_length
 
63
 
64
  self._build_model()
65
  self._build_buffer()
@@ -75,7 +77,8 @@ class Diffusion(nn.Module):
75
  use_memory_attention=self.use_memory_attention,
76
  add_timestamp_embedding=self.add_timestamp_embedding,
77
  ref_mode=self.ref_mode,
78
- focal_length=self.focal_length)
 
79
  else:
80
  raise NotImplementedError
81
 
 
31
  add_timestamp_embedding=False,
32
  ref_mode='sequential',
33
  focal_length=0.35,
34
+ memory_attention_key_only_geometry=True,
35
  ):
36
  super().__init__()
37
  self.cfg = cfg
 
61
  self.add_timestamp_embedding = add_timestamp_embedding
62
  self.ref_mode = ref_mode
63
  self.focal_length = focal_length
64
+ self.memory_attention_key_only_geometry = memory_attention_key_only_geometry
65
 
66
  self._build_model()
67
  self._build_buffer()
 
77
  use_memory_attention=self.use_memory_attention,
78
  add_timestamp_embedding=self.add_timestamp_embedding,
79
  ref_mode=self.ref_mode,
80
+ focal_length=self.focal_length,
81
+ memory_attention_key_only_geometry=self.memory_attention_key_only_geometry)
82
  else:
83
  raise NotImplementedError
84
 
algorithms/dememwm/models/dit.py CHANGED
@@ -143,11 +143,12 @@ class FinalLayer(nn.Module):
143
 
144
 
145
  class FrameMemoryReferenceAttention(nn.Module):
146
- def __init__(self, hidden_size, num_heads):
147
  super().__init__()
148
  self.hidden_size = hidden_size
149
  self.num_heads = num_heads
150
  self.head_dim = hidden_size // num_heads
 
151
  self.to_q = nn.Linear(hidden_size, hidden_size, bias=False)
152
  self.to_k = nn.Linear(hidden_size, hidden_size, bias=False)
153
  self.to_v = nn.Linear(hidden_size, hidden_size, bias=False)
@@ -167,9 +168,13 @@ class FrameMemoryReferenceAttention(nn.Module):
167
  T_memory = memory_hidden.shape[1]
168
  P = H * W
169
  if geometry_cache is not None:
170
- query_rays = geometry_cache["query_rays"].to(device=target_hidden.device, dtype=target_hidden.dtype)
171
- relative_rays = geometry_cache["relative_rays"].to(device=target_hidden.device, dtype=target_hidden.dtype)
172
- target_time = geometry_cache.get("target_time")
 
 
 
 
173
  relative_time = geometry_cache.get("relative_time")
174
  else:
175
  query_rays = None
@@ -252,6 +257,7 @@ class SpatioTemporalDiTBlock(nn.Module):
252
  relative_embedding=False,
253
  state_embed_only_on_qk=False,
254
  use_memory_attention=False,
 
255
  ref_mode='sequential'
256
  ):
257
  super().__init__()
@@ -304,9 +310,9 @@ class SpatioTemporalDiTBlock(nn.Module):
304
  drop=0,
305
  )
306
  self.r_adaLN_modulation = nn.Sequential(nn.SiLU(), nn.Linear(hidden_size, 6 * hidden_size, bias=True))
307
- self.r_attn_anchor = FrameMemoryReferenceAttention(hidden_size, num_heads)
308
- self.r_attn_dynamic = FrameMemoryReferenceAttention(hidden_size, num_heads)
309
- self.r_attn_revisit = FrameMemoryReferenceAttention(hidden_size, num_heads)
310
 
311
  self.reference_length = reference_length
312
  self.relative_embedding = relative_embedding
@@ -473,6 +479,7 @@ class DiT(nn.Module):
473
  add_timestamp_embedding=False,
474
  ref_mode='sequential',
475
  focal_length=0.35,
 
476
  ):
477
  super().__init__()
478
  self.in_channels = in_channels
@@ -522,6 +529,7 @@ class DiT(nn.Module):
522
  relative_embedding=relative_embedding,
523
  state_embed_only_on_qk=state_embed_only_on_qk,
524
  use_memory_attention=use_memory_attention,
 
525
  ref_mode=ref_mode
526
  )
527
  for _ in range(depth)
@@ -792,7 +800,7 @@ class DiT(nn.Module):
792
  def DiT_S_2(action_cond_dim, pose_cond_dim, reference_length,
793
  use_plucker, relative_embedding,
794
  state_embed_only_on_qk, use_memory_attention, add_timestamp_embedding,
795
- ref_mode, focal_length=0.35):
796
  return DiT(
797
  patch_size=2,
798
  hidden_size=1024,
@@ -806,6 +814,7 @@ ref_mode, focal_length=0.35):
806
  state_embed_only_on_qk=state_embed_only_on_qk,
807
  use_memory_attention=use_memory_attention,
808
  add_timestamp_embedding=add_timestamp_embedding,
 
809
  ref_mode=ref_mode,
810
  focal_length=focal_length,
811
  )
 
143
 
144
 
145
  class FrameMemoryReferenceAttention(nn.Module):
146
+ def __init__(self, hidden_size, num_heads, key_only_geometry=True):
147
  super().__init__()
148
  self.hidden_size = hidden_size
149
  self.num_heads = num_heads
150
  self.head_dim = hidden_size // num_heads
151
+ self.key_only_geometry = key_only_geometry
152
  self.to_q = nn.Linear(hidden_size, hidden_size, bias=False)
153
  self.to_k = nn.Linear(hidden_size, hidden_size, bias=False)
154
  self.to_v = nn.Linear(hidden_size, hidden_size, bias=False)
 
168
  T_memory = memory_hidden.shape[1]
169
  P = H * W
170
  if geometry_cache is not None:
171
+ query_rays = None if self.key_only_geometry else geometry_cache.get("query_rays")
172
+ if query_rays is not None:
173
+ query_rays = query_rays.to(device=target_hidden.device, dtype=target_hidden.dtype)
174
+ relative_rays = geometry_cache.get("relative_rays")
175
+ if relative_rays is not None:
176
+ relative_rays = relative_rays.to(device=target_hidden.device, dtype=target_hidden.dtype)
177
+ target_time = None if self.key_only_geometry else geometry_cache.get("target_time")
178
  relative_time = geometry_cache.get("relative_time")
179
  else:
180
  query_rays = None
 
257
  relative_embedding=False,
258
  state_embed_only_on_qk=False,
259
  use_memory_attention=False,
260
+ key_only_geometry=True,
261
  ref_mode='sequential'
262
  ):
263
  super().__init__()
 
310
  drop=0,
311
  )
312
  self.r_adaLN_modulation = nn.Sequential(nn.SiLU(), nn.Linear(hidden_size, 6 * hidden_size, bias=True))
313
+ self.r_attn_anchor = FrameMemoryReferenceAttention(hidden_size, num_heads, key_only_geometry=key_only_geometry)
314
+ self.r_attn_dynamic = FrameMemoryReferenceAttention(hidden_size, num_heads, key_only_geometry=key_only_geometry)
315
+ self.r_attn_revisit = FrameMemoryReferenceAttention(hidden_size, num_heads, key_only_geometry=key_only_geometry)
316
 
317
  self.reference_length = reference_length
318
  self.relative_embedding = relative_embedding
 
479
  add_timestamp_embedding=False,
480
  ref_mode='sequential',
481
  focal_length=0.35,
482
+ memory_attention_key_only_geometry=True,
483
  ):
484
  super().__init__()
485
  self.in_channels = in_channels
 
529
  relative_embedding=relative_embedding,
530
  state_embed_only_on_qk=state_embed_only_on_qk,
531
  use_memory_attention=use_memory_attention,
532
+ key_only_geometry=memory_attention_key_only_geometry,
533
  ref_mode=ref_mode
534
  )
535
  for _ in range(depth)
 
800
  def DiT_S_2(action_cond_dim, pose_cond_dim, reference_length,
801
  use_plucker, relative_embedding,
802
  state_embed_only_on_qk, use_memory_attention, add_timestamp_embedding,
803
+ ref_mode, focal_length=0.35, memory_attention_key_only_geometry=True):
804
  return DiT(
805
  patch_size=2,
806
  hidden_size=1024,
 
814
  state_embed_only_on_qk=state_embed_only_on_qk,
815
  use_memory_attention=use_memory_attention,
816
  add_timestamp_embedding=add_timestamp_embedding,
817
+ memory_attention_key_only_geometry=memory_attention_key_only_geometry,
818
  ref_mode=ref_mode,
819
  focal_length=focal_length,
820
  )
configurations/algorithm/dememwm_base.yaml CHANGED
@@ -36,7 +36,8 @@ use_plucker: true
36
  relative_embedding: true
37
  state_embed_only_on_qk: true
38
  use_memory_attention: true
39
- add_timestamp_embedding: true
 
40
  ref_mode: sequential
41
  focal_length: 0.35
42
  log_video: false
 
36
  relative_embedding: true
37
  state_embed_only_on_qk: true
38
  use_memory_attention: true
39
+ memory_attention_key_only_geometry: true
40
+ add_timestamp_embedding: false
41
  ref_mode: sequential
42
  focal_length: 0.35
43
  log_video: false
tests/test_dememwm_temporal_attention.py CHANGED
@@ -502,7 +502,7 @@ class DeMemWMTemporalAttentionTests(unittest.TestCase):
502
  self.assertTrue(torch.allclose(second_only, torch.full_like(second_only, 10.0)))
503
  self.assertTrue(torch.equal(empty, torch.zeros_like(empty)))
504
 
505
- def test_frame_memory_reference_attention_uses_geometry_qk(self):
506
  attn = FrameMemoryReferenceAttention(hidden_size=2, num_heads=1)
507
  with torch.no_grad():
508
  attn.to_q.weight.zero_()
@@ -513,7 +513,7 @@ class DeMemWMTemporalAttentionTests(unittest.TestCase):
513
  attn.to_v.weight[0, 1] = 1.0
514
  attn.query_pose_proj.weight.zero_()
515
  attn.query_pose_proj.bias.zero_()
516
- attn.query_pose_proj.weight[0, 5] = 1.0
517
  attn.key_pose_proj.weight.zero_()
518
  attn.key_pose_proj.bias.zero_()
519
  attn.key_pose_proj.weight[0, 5] = 1.0
@@ -522,10 +522,12 @@ class DeMemWMTemporalAttentionTests(unittest.TestCase):
522
  attn.out_proj.bias.zero_()
523
 
524
  target = torch.zeros((1, 1, 1, 1, 2))
 
525
  memory = torch.zeros((1, 2, 1, 1, 2))
526
  memory[0, :, 0, 0, 1] = torch.tensor([1.0, 10.0])
527
  query_rays = torch.zeros((1, 1, 1, 1, 6))
528
- query_rays[..., 5] = 1.0
 
529
  relative_a = torch.zeros((1, 1, 2, 1, 1, 6))
530
  relative_b = torch.zeros_like(relative_a)
531
  relative_a[0, 0, 0, 0, 0, 5] = 2.0
@@ -536,12 +538,18 @@ class DeMemWMTemporalAttentionTests(unittest.TestCase):
536
  memory,
537
  geometry_cache={"query_rays": query_rays, "relative_rays": relative_a},
538
  )
 
 
 
 
 
539
  out_b = attn(
540
  target,
541
  memory,
542
  geometry_cache={"query_rays": query_rays, "relative_rays": relative_b},
543
  )
544
 
 
545
  self.assertLess(out_a[0, 0, 0, 0, 0].item(), out_b[0, 0, 0, 0, 0].item())
546
  self.assertFalse(torch.allclose(out_a, out_b))
547
 
 
502
  self.assertTrue(torch.allclose(second_only, torch.full_like(second_only, 10.0)))
503
  self.assertTrue(torch.equal(empty, torch.zeros_like(empty)))
504
 
505
+ def test_frame_memory_reference_attention_uses_key_geometry_only(self):
506
  attn = FrameMemoryReferenceAttention(hidden_size=2, num_heads=1)
507
  with torch.no_grad():
508
  attn.to_q.weight.zero_()
 
513
  attn.to_v.weight[0, 1] = 1.0
514
  attn.query_pose_proj.weight.zero_()
515
  attn.query_pose_proj.bias.zero_()
516
+ attn.query_pose_proj.weight[0, 5] = 100.0
517
  attn.key_pose_proj.weight.zero_()
518
  attn.key_pose_proj.bias.zero_()
519
  attn.key_pose_proj.weight[0, 5] = 1.0
 
522
  attn.out_proj.bias.zero_()
523
 
524
  target = torch.zeros((1, 1, 1, 1, 2))
525
+ target[0, 0, 0, 0, 0] = 1.0
526
  memory = torch.zeros((1, 2, 1, 1, 2))
527
  memory[0, :, 0, 0, 1] = torch.tensor([1.0, 10.0])
528
  query_rays = torch.zeros((1, 1, 1, 1, 6))
529
+ query_rays_changed = query_rays.clone()
530
+ query_rays_changed[..., 5] = 10.0
531
  relative_a = torch.zeros((1, 1, 2, 1, 1, 6))
532
  relative_b = torch.zeros_like(relative_a)
533
  relative_a[0, 0, 0, 0, 0, 5] = 2.0
 
538
  memory,
539
  geometry_cache={"query_rays": query_rays, "relative_rays": relative_a},
540
  )
541
+ out_a_query_changed = attn(
542
+ target,
543
+ memory,
544
+ geometry_cache={"query_rays": query_rays_changed, "relative_rays": relative_a},
545
+ )
546
  out_b = attn(
547
  target,
548
  memory,
549
  geometry_cache={"query_rays": query_rays, "relative_rays": relative_b},
550
  )
551
 
552
+ self.assertTrue(torch.allclose(out_a, out_a_query_changed))
553
  self.assertLess(out_a[0, 0, 0, 0, 0].item(), out_b[0, 0, 0, 0, 0].item())
554
  self.assertFalse(torch.allclose(out_a, out_b))
555