Lam_VGGT / code /changes_vs_upstream.patch
YuanhaoXD's picture
LaWAM Stage 1 with frozen VGGT-1B encoder: weights, code, report, figures
0283577
Raw
History Blame Contribute Delete
15.6 kB
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,