BonanDing commited on
Commit
ca62cbd
·
1 Parent(s): 6e0a1a9

Use configured focal length in DeMemWM geometry

Browse files
algorithms/dememwm/df_video.py CHANGED
@@ -690,7 +690,8 @@ class DeMemWMMinecraft(DiffusionForcingBase):
690
  state_embed_only_on_qk=self.state_embed_only_on_qk,
691
  use_memory_attention=self.use_memory_attention,
692
  add_timestamp_embedding=self.add_timestamp_embedding,
693
- ref_mode=self.ref_mode
 
694
  )
695
 
696
  self.validation_lpips_model = LearnedPerceptualImagePatchSimilarity(sync_on_compute=False)
 
690
  state_embed_only_on_qk=self.state_embed_only_on_qk,
691
  use_memory_attention=self.use_memory_attention,
692
  add_timestamp_embedding=self.add_timestamp_embedding,
693
+ ref_mode=self.ref_mode,
694
+ focal_length=self.focal_length,
695
  )
696
 
697
  self.validation_lpips_model = LearnedPerceptualImagePatchSimilarity(sync_on_compute=False)
algorithms/dememwm/models/diffusion.py CHANGED
@@ -29,7 +29,8 @@ class Diffusion(nn.Module):
29
  state_embed_only_on_qk=False,
30
  use_memory_attention=False,
31
  add_timestamp_embedding=False,
32
- ref_mode='sequential'
 
33
  ):
34
  super().__init__()
35
  self.cfg = cfg
@@ -58,6 +59,7 @@ class Diffusion(nn.Module):
58
  self.use_memory_attention = use_memory_attention
59
  self.add_timestamp_embedding = add_timestamp_embedding
60
  self.ref_mode = ref_mode
 
61
 
62
  self._build_model()
63
  self._build_buffer()
@@ -72,7 +74,8 @@ class Diffusion(nn.Module):
72
  state_embed_only_on_qk=self.state_embed_only_on_qk,
73
  use_memory_attention=self.use_memory_attention,
74
  add_timestamp_embedding=self.add_timestamp_embedding,
75
- ref_mode=self.ref_mode)
 
76
  else:
77
  raise NotImplementedError
78
 
 
29
  state_embed_only_on_qk=False,
30
  use_memory_attention=False,
31
  add_timestamp_embedding=False,
32
+ ref_mode='sequential',
33
+ focal_length=0.35,
34
  ):
35
  super().__init__()
36
  self.cfg = cfg
 
59
  self.use_memory_attention = use_memory_attention
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()
 
74
  state_embed_only_on_qk=self.state_embed_only_on_qk,
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
 
algorithms/dememwm/models/dit.py CHANGED
@@ -459,7 +459,8 @@ class DiT(nn.Module):
459
  state_embed_only_on_qk=False,
460
  use_memory_attention=False,
461
  add_timestamp_embedding=False,
462
- ref_mode='sequential'
 
463
  ):
464
  super().__init__()
465
  self.in_channels = in_channels
@@ -467,6 +468,7 @@ class DiT(nn.Module):
467
  self.patch_size = patch_size
468
  self.num_heads = num_heads
469
  self.max_frames = max_frames
 
470
 
471
  self.x_embedder = PatchEmbed(input_h, input_w, patch_size, in_channels, hidden_size, flatten=False)
472
  self.t_embedder = TimestepEmbedder(hidden_size)
@@ -615,7 +617,7 @@ class DiT(nn.Module):
615
  # Token grid sets position count; image_hw keeps camera scaling in original image coordinates.
616
  y = (torch.arange(grid_h, device=device, dtype=torch.float32).view(1, grid_h, 1) + 0.5) * (image_h / grid_h)
617
  x = (torch.arange(grid_w, device=device, dtype=torch.float32).view(1, 1, grid_w) + 0.5) * (image_w / grid_w)
618
- fx, fy = 0.35 * image_w, 0.35 * image_h
619
  cx, cy = 0.5 * image_w, 0.5 * image_h
620
  directions = torch.stack(
621
  (
@@ -758,7 +760,7 @@ class DiT(nn.Module):
758
  def DiT_S_2(action_cond_dim, pose_cond_dim, reference_length,
759
  use_plucker, relative_embedding,
760
  state_embed_only_on_qk, use_memory_attention, add_timestamp_embedding,
761
- ref_mode):
762
  return DiT(
763
  patch_size=2,
764
  hidden_size=1024,
@@ -772,7 +774,8 @@ ref_mode):
772
  state_embed_only_on_qk=state_embed_only_on_qk,
773
  use_memory_attention=use_memory_attention,
774
  add_timestamp_embedding=add_timestamp_embedding,
775
- ref_mode=ref_mode
 
776
  )
777
 
778
 
 
459
  state_embed_only_on_qk=False,
460
  use_memory_attention=False,
461
  add_timestamp_embedding=False,
462
+ ref_mode='sequential',
463
+ focal_length=0.35,
464
  ):
465
  super().__init__()
466
  self.in_channels = in_channels
 
468
  self.patch_size = patch_size
469
  self.num_heads = num_heads
470
  self.max_frames = max_frames
471
+ self.focal_length = focal_length
472
 
473
  self.x_embedder = PatchEmbed(input_h, input_w, patch_size, in_channels, hidden_size, flatten=False)
474
  self.t_embedder = TimestepEmbedder(hidden_size)
 
617
  # Token grid sets position count; image_hw keeps camera scaling in original image coordinates.
618
  y = (torch.arange(grid_h, device=device, dtype=torch.float32).view(1, grid_h, 1) + 0.5) * (image_h / grid_h)
619
  x = (torch.arange(grid_w, device=device, dtype=torch.float32).view(1, 1, grid_w) + 0.5) * (image_w / grid_w)
620
+ fx, fy = self.focal_length * image_w, self.focal_length * image_h
621
  cx, cy = 0.5 * image_w, 0.5 * image_h
622
  directions = torch.stack(
623
  (
 
760
  def DiT_S_2(action_cond_dim, pose_cond_dim, reference_length,
761
  use_plucker, relative_embedding,
762
  state_embed_only_on_qk, use_memory_attention, add_timestamp_embedding,
763
+ ref_mode, focal_length=0.35):
764
  return DiT(
765
  patch_size=2,
766
  hidden_size=1024,
 
774
  state_embed_only_on_qk=state_embed_only_on_qk,
775
  use_memory_attention=use_memory_attention,
776
  add_timestamp_embedding=add_timestamp_embedding,
777
+ ref_mode=ref_mode,
778
+ focal_length=focal_length,
779
  )
780
 
781