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 =
|
| 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 |
|