Use key-only DeMemWM pose geometry
Browse filesDisable 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",
|
|
|
|
| 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
|
| 171 |
-
|
| 172 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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 |
-
|
|
|
|
| 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
|
| 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] =
|
| 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 |
-
|
|
|
|
| 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 |
|