File size: 15,633 Bytes
0283577
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
diff --git a/latent_action_model/core/lam_lightinng.py b/latent_action_model/core/lam_lightinng.py
index cf1b883..347a691 100644
--- a/latent_action_model/core/lam_lightinng.py
+++ b/latent_action_model/core/lam_lightinng.py
@@ -898,6 +898,34 @@ class VJEPA_LAM(LightningModule):
         plt.savefig(f"{filename}.png", bbox_inches="tight", pad_inches=0.0)
         plt.close()
 
+    def on_save_checkpoint(self, checkpoint: Dict[str, Any]) -> None:
+        """Drop the frozen visual encoder from the checkpoint.
+
+        `lam.vision_encoder.*` is never trained and is re-loaded from
+        `vision_model_id` on every `LatentLAMModel.__init__`, so persisting it
+        only wastes disk. With VGGT-1B that is 909 M params -> ~3.6 GB per
+        checkpoint file.
+        """
+        state_dict = checkpoint.get("state_dict")
+        if not state_dict:
+            return
+        for key in [k for k in state_dict if k.startswith("lam.vision_encoder.")]:
+            del state_dict[key]
+
+    def on_load_checkpoint(self, checkpoint: Dict[str, Any]) -> None:
+        """Re-inject the frozen encoder weights stripped by `on_save_checkpoint`.
+
+        They are already loaded in `__init__`, so copying them back from the live
+        module lets Lightning keep loading strictly instead of silencing genuinely
+        missing trained weights.
+        """
+        state_dict = checkpoint.get("state_dict")
+        if state_dict is None:
+            return
+        for key, value in self.state_dict().items():
+            if key.startswith("lam.vision_encoder.") and key not in state_dict:
+                state_dict[key] = value
+
     def configure_optimizers(self) -> Any:
         if self.exclude_bias_norm_from_wd:
             param_groups, decay_names, no_decay_names = self._build_optimizer_param_groups()
diff --git a/latent_action_model/core/lam_model.py b/latent_action_model/core/lam_model.py
index c06dbd9..8bb4f0b 100644
--- a/latent_action_model/core/lam_model.py
+++ b/latent_action_model/core/lam_model.py
@@ -596,6 +596,10 @@ def load_latent_action_model(ckpt_path, yaml_path):
     ckpt_keys = set(new_ckpt.keys())
 
     missing_keys = sorted(list(model_keys - ckpt_keys))
+    # `VJEPA_LAM.on_save_checkpoint` strips the frozen visual encoder, which
+    # `LatentLAMModel.__init__` has already rebuilt from `vision_model_id`.
+    # Its absence is expected, not an error.
+    missing_keys = [k for k in missing_keys if not k.startswith("vision_encoder.")]
     unexpected_keys = sorted(list(ckpt_keys - model_keys))
     shape_mismatches = []
     for k in sorted(model_keys & ckpt_keys):
@@ -615,6 +619,12 @@ def load_latent_action_model(ckpt_path, yaml_path):
             error_lines += [f"  - {k}: model{ms} vs checkpoint{cs}" for k, ms, cs in shape_mismatches]
         raise RuntimeError("\n".join(error_lines))
 
+    # Keep strict=True meaningful for the trained weights by filling the stripped
+    # encoder entries from the freshly built (already pretrained) module.
+    for key, value in model_state.items():
+        if key.startswith("vision_encoder.") and key not in new_ckpt:
+            new_ckpt[key] = value
+
     latent_action_model.load_state_dict(new_ckpt, strict=True)
     for p in latent_action_model.parameters():
         p.requires_grad = False
diff --git a/latent_action_model/core/vjepa_encoder.py b/latent_action_model/core/vjepa_encoder.py
index 5d5eebe..e60d7ca 100644
--- a/latent_action_model/core/vjepa_encoder.py
+++ b/latent_action_model/core/vjepa_encoder.py
@@ -391,6 +391,244 @@ class DINOv3Encoder(nn.Module):
 
             return features[0].detach() if isinstance(n, int) else [f.detach() for f in features]
 
+
+# VGGT-1B aggregator output dim = 2 * embed_dim (frame-attention ‖ global-attention concat).
+VGGT_FEATURE_DIM = 2048
+VGGT_DEFAULT_INPUT_SIZE = 518
+
+
+class VGGTEncoder(nn.Module):
+    """Frozen VGGT-1B aggregator, drop-in compatible with :class:`DINOv3Encoder`.
+
+    VGGT is a 3D-geometry backbone rather than a 2D semantic one, so four
+    adaptations are needed to keep the rest of the LAM stack byte-for-byte
+    unchanged:
+
+    1. ``video_aug`` hands us ImageNet-normalized pixels, but the VGGT aggregator
+       applies its *own* ImageNet normalization internally and expects [0, 1].
+       We de-normalize before the forward pass.
+    2. VGGT is patch-14 at 518px -> a 37x37 grid, while the LAM decoder is wired
+       for ``LAM_PATCH_SIZE=16`` -> 16x16. We average-pool 37x37 down to 16x16 so
+       the token count matches and no downstream shape changes.
+    3. VGGT's global-attention blocks mix *all* frames stacked on the S axis, and
+       ``LatentLAMModel._run`` concatenates (enc_t, enc_T, dec_t, dec_T) into a
+       single ``encode`` call. Encoding with S>1 would leak the future frame into
+       the current one, so we force S=1 and put every frame on the batch axis.
+       That makes frames independent exactly like DINOv3 (verified bit-identical
+       against per-frame calls; S-stacking differs by max|delta|~20).
+    4. ``Aggregator.forward`` returns a length-``depth`` list where uncached
+       layers are ``None`` (only blocks {4, 11, 17, 23} are kept). Indexing it
+       with the shipped ``latent_layer_to_use=-2`` would yield ``None``, so we
+       index into the compacted list of cached layers instead: -1 -> block 23,
+       -2 -> block 17.
+    """
+
+    def __init__(
+        self,
+        model_id: str = "facebook/VGGT-1B",
+        num_latent_layers: int = 1,
+        norm_layer_type: str = "l2",
+        enable_norm: bool = False,
+        image_size: int = 256,
+        target_grid: int = 16,
+        vggt_input_size: int = VGGT_DEFAULT_INPUT_SIZE,
+    ):
+        super().__init__()
+        from vggt.models.vggt import VGGT
+
+        self.device = torch.device("cpu")
+        self.model_id = model_id
+        self.num_latent_layers = max(int(num_latent_layers), 1)
+        self.norm_layer_type = norm_layer_type
+        self.enable_norm = enable_norm
+        self.feature_dim = VGGT_FEATURE_DIM
+        # `image_size` is what the LAM pipeline feeds us (256); `vggt_input_size`
+        # is what we upsample to internally. Reporting 256 keeps LatentLAMModel's
+        # image_hw consistency check quiet.
+        self.image_size = int(image_size)
+        self.vggt_input_size = int(vggt_input_size)
+        self.target_grid = int(target_grid)
+        self.patch_size = self.image_size // self.target_grid
+
+        # Prediction heads are dead weight here: skip building them entirely
+        # (~500M params) and let strict=False drop their checkpoint entries.
+        model = VGGT(
+            enable_camera=False,
+            enable_point=False,
+            enable_depth=False,
+            enable_track=False,
+        )
+        state_dict = _load_vggt_state_dict(model_id)
+        missing, unexpected = model.load_state_dict(state_dict, strict=False)
+        aggregator_missing = [k for k in missing if k.startswith("aggregator.")]
+        if aggregator_missing:
+            raise RuntimeError(
+                f"VGGT aggregator weights are incomplete: {len(aggregator_missing)} missing keys, "
+                f"e.g. {aggregator_missing[:5]}"
+            )
+        print(
+            f"[VGGTEncoder] loaded `{model_id}` | missing={len(missing)} unexpected={len(unexpected)} "
+            f"| {self.image_size}px -> {self.vggt_input_size}px -> {self.target_grid}x{self.target_grid} "
+            f"tokens x {self.feature_dim}d"
+        )
+        model.eval()
+        self.model = model
+        for param in self.model.parameters():
+            param.requires_grad = False
+
+        mean = torch.tensor(IMAGENET_DEFAULT_MEAN).view(1, 3, 1, 1)
+        std = torch.tensor(IMAGENET_DEFAULT_STD).view(1, 3, 1, 1)
+        self.register_buffer("_imagenet_mean", mean, persistent=False)
+        self.register_buffer("_imagenet_std", std, persistent=False)
+
+        if self.norm_layer_type in ("bn", "ln"):
+            if self.norm_layer_type == "bn":
+                norm_builder = lambda: nn.SyncBatchNorm(self.feature_dim, affine=False)
+            else:
+                norm_builder = lambda: nn.LayerNorm(self.feature_dim, elementwise_affine=False)
+            self.latent_norms = nn.ModuleList([norm_builder() for _ in range(self.num_latent_layers)])
+        else:
+            self.latent_norms = None
+
+    def train(self, mode: bool = True):
+        # Keep the frozen backbone in eval mode. This is load-bearing beyond the
+        # usual dropout/BN reason: Aggregator gates `torch.utils.checkpoint` on
+        # `self.training`, which is pure waste under no_grad.
+        super().train(False)
+        self.model.eval()
+        return self
+
+    def _denormalize_to_unit(self, images: torch.Tensor) -> torch.Tensor:
+        """Undo `video_aug`'s ImageNet normalization; the aggregator redoes it."""
+        mean = self._imagenet_mean.to(device=images.device, dtype=images.dtype)
+        std = self._imagenet_std.to(device=images.device, dtype=images.dtype)
+        return (images * std + mean).clamp_(0.0, 1.0)
+
+    def _to_vggt_resolution(self, images: torch.Tensor) -> torch.Tensor:
+        if images.shape[-2:] == (self.vggt_input_size, self.vggt_input_size):
+            return images
+        return F.interpolate(
+            images,
+            size=(self.vggt_input_size, self.vggt_input_size),
+            mode="bilinear",
+            align_corners=False,
+        )
+
+    def _pool_to_target_grid(self, tokens: torch.Tensor) -> torch.Tensor:
+        """[N, side*side, D] patch tokens -> [N, target_grid**2, D]."""
+        n, num_patches, dim = tokens.shape
+        side = int(math.sqrt(num_patches))
+        if side * side != num_patches:
+            raise ValueError(f"VGGT patch count {num_patches} is not a square grid.")
+        if side == self.target_grid:
+            return tokens
+        grid = tokens.reshape(n, side, side, dim).permute(0, 3, 1, 2)
+        grid = F.adaptive_avg_pool2d(grid, (self.target_grid, self.target_grid))
+        return grid.flatten(2).transpose(1, 2).contiguous()
+
+    def _apply_norm(self, tokens: torch.Tensor, idx: int) -> torch.Tensor:
+        if not self.enable_norm:
+            return tokens
+        if self.norm_layer_type == "bn":
+            if self.latent_norms is None:
+                raise ValueError("VGGTEncoder has no BN layer initialized; set norm_layer_type to 'bn'.")
+            tokens_2d = tokens.reshape(-1, self.feature_dim)
+            tokens_2d = self.latent_norms[idx](tokens_2d)
+            return tokens_2d.view(tokens.shape[0], tokens.shape[1], self.feature_dim)
+        if self.norm_layer_type == "ln":
+            if self.latent_norms is None:
+                raise ValueError("VGGTEncoder has no LN layer initialized; set norm_layer_type to 'ln'.")
+            return self.latent_norms[idx](tokens)
+        if self.norm_layer_type == "l2":
+            return F.normalize(tokens, p=2, dim=-1)
+        return tokens
+
+    @torch.no_grad()
+    def encode(self, images: torch.Tensor, remove_cls: bool = True, n: Union[int, Sequence] = -1) -> torch.Tensor:
+        if not remove_cls:
+            raise NotImplementedError(
+                "VGGTEncoder always strips camera/register tokens: the remaining patch tokens "
+                "are pooled onto a square grid, which the special tokens cannot join."
+            )
+        if images.dim() == 5:
+            B, T = images.shape[0], images.shape[1]
+            flat = images.reshape(-1, images.shape[-3], images.shape[-2], images.shape[-1])
+        elif images.dim() == 4:
+            B, T = images.shape[0], 1
+            flat = images
+        else:
+            raise ValueError(f"Expected 4D or 5D input, got {tuple(images.shape)}")
+
+        pixels = self._to_vggt_resolution(self._denormalize_to_unit(flat.float()))
+        param_dtype = next(self.model.parameters()).dtype
+        # S=1 -> every frame is encoded independently; see class docstring (3).
+        token_list, patch_start_idx = self.model.aggregator(pixels.unsqueeze(1).to(dtype=param_dtype))
+
+        cached = [t for t in token_list if t is not None]
+        if not cached:
+            raise RuntimeError("VGGT aggregator returned no cached layers.")
+
+        list_n = [n] if isinstance(n, int) else list(n)
+        assert len(list_n) <= self.num_latent_layers, (
+            f"VGGTEncoder expected at most {self.num_latent_layers} normalization layers, "
+            f"but received {len(list_n)} feature layers. Ensure this matches len(latent_layer_to_use)."
+        )
+
+        features = []
+        for idx, layer in enumerate(list_n):
+            layer_tokens = cached[self._resolve_cached_index(layer, len(cached))]  # [N, 1, 5+P, D]
+            layer_tokens = layer_tokens[:, 0, patch_start_idx:, :]                 # [N, P, D]
+            layer_tokens = self._pool_to_target_grid(layer_tokens)                 # [N, K, D]
+            layer_tokens = self._apply_norm(layer_tokens, idx)
+            features.append(layer_tokens.reshape(B, T, -1, self.feature_dim))
+
+        return features[0].detach() if isinstance(n, int) else [f.detach() for f in features]
+
+    @staticmethod
+    def _resolve_cached_index(layer: int, num_cached: int) -> int:
+        """Map a config layer index onto the compacted list of cached layers.
+
+        The configs speak DINOv3 (`-2` = penultimate block). VGGT only keeps four
+        blocks, so `-1` -> block 23 and `-2` -> block 17. Indices beyond the
+        cached range are clamped rather than silently returning `None`.
+        """
+        layer = int(layer)
+        if layer < 0:
+            layer = max(layer, -num_cached)
+        else:
+            layer = min(layer, num_cached - 1)
+        return layer
+
+
+def _load_vggt_state_dict(model_id: str) -> dict:
+    """Load VGGT weights from a local .pt/.safetensors file, a directory, or the hub."""
+    path = _get_existing_path(model_id)
+    if path is not None:
+        if path.is_dir():
+            candidates = [
+                path / "model.pt",
+                path / "model.safetensors",
+                path / "pytorch_model.bin",
+            ]
+            found = next((c for c in candidates if c.exists()), None)
+            if found is None:
+                raise FileNotFoundError(
+                    f"No VGGT weight file (model.pt / model.safetensors / pytorch_model.bin) under {path}"
+                )
+            path = found
+        if path.suffix == ".safetensors":
+            from safetensors.torch import load_file
+
+            return load_file(str(path))
+        state_dict = _safe_torch_load(path)
+        return state_dict.get("model", state_dict)
+
+    # Not a local path: fall back to the hub via PyTorchModelHubMixin.
+    from vggt.models.vggt import VGGT
+
+    return VGGT.from_pretrained(model_id).state_dict()
+
+
 class CosmosAutoencoder(nn.Module):
 
     def __init__(
@@ -610,7 +848,15 @@ def build_vision_encoder(
 
     key = str(model_id).lower()
     vjepa_hub_name = _infer_vjepa_hub_name(model_id)
-    if "dinov3" in key:
+    if "vggt" in key:
+        encoder = VGGTEncoder(
+            model_id=model_id,
+            num_latent_layers=num_latent_layers,
+            norm_layer_type=norm_layer_type,
+            enable_norm=enable_norm,
+        )
+        return encoder, encoder.feature_dim
+    elif "dinov3" in key:
         encoder = DINOv3Encoder(
             model_id=model_id,
             num_latent_layers=num_latent_layers,