#!/usr/bin/env python """DINO-WM style multi-view dynamics: predict future DINOv2 patch features, not pixels. {I_{t-8}, I_t}^1..V, a[t:t+8], s_t -> DINOv2 features of {I_{t+8}}^1..V Follows DINO-WM (arXiv:2411.04983): a frozen DINOv2 encoder, a small ViT predictor trained with plain MSE in feature space, and a decoder trained separately on detached features for visualisation only. No diffusion, no sampling -- the forecast is deterministic. Differences from the reference implementation, both deliberate: * multi-view: patch tokens of every camera share one sequence, each with a learned view embedding. * the target is read from explicit query tokens. The reference feeds frames 0..N-1 with full (non-causal) attention and scores them against frames 1..N, so every position except the last has its own target visible in the input; only the final position is a true forecast. Pixel metrics need a decoder: predicted DINO features -> SD3 VAE latent -> frozen SD3 VAE -> image. Reusing the SD3 VAE avoids training a pixel decoder from scratch. P=policy_learning/lerobot/.venv/bin/python PYTHONUNBUFFERED=1 $P dino_dynamics.py --data --steps 20000 """ import argparse import time from pathlib import Path import torch import torch.nn as nn import torch.nn.functional as F from sd3_dynamics import ( HORIZON, build_datasets, dataset_features, object_state_keys, per_camera_metrics, reconstruction_table, ) DINO_ID = "facebook/dinov2-small" # ViT-S/14, 384-dim, as in the reference config DINO_MEAN = (0.485, 0.456, 0.406) DINO_STD = (0.229, 0.224, 0.225) class DinoDynamics(nn.Module): """Frozen DINOv2 encoder + ViT predictor over [history x views] patch tokens.""" def __init__( self, dino, vae, n_views, action_dim, state_dim, history=2, depth=6, heads=6, mlp_dim=2048, horizon=HORIZON, dino_size=196, image_size=224, use_action=True, use_state=True, ): super().__init__() self.dino = dino.eval().requires_grad_(False) self.vae = vae.eval().requires_grad_(False) if vae is not None else None d = dino.config.hidden_size self.dim, self.n_views, self.history = d, n_views, history # ablations zero the embedding rather than drop the token, so token count and # model capacity are unchanged and the only variable is the information itself self.use_action, self.use_state = use_action, use_state self.patch = dino.config.patch_size self.dino_size = dino_size # multiple of patch size; 196 -> 14x14 = 196 tokens self.side = dino_size // self.patch self.n_patches = self.side**2 self.view_emb = nn.Parameter(torch.zeros(n_views, d)) self.frame_emb = nn.Parameter(torch.zeros(history + 1, d)) # last entry = target frame self.pos_emb = nn.Parameter(torch.zeros(self.n_patches, d)) self.query = nn.Parameter(torch.randn(self.n_patches, d) * 0.02) self.act_tok = nn.Linear(horizon * action_dim, d) self.state_tok = nn.Linear(state_dim, d) layer = nn.TransformerEncoderLayer( d, heads, mlp_dim, dropout=0.0, batch_first=True, norm_first=True, activation="gelu" ) self.predictor = nn.TransformerEncoder(layer, depth) self.norm = nn.LayerNorm(d) # feature -> SD3 latent head, for pixel metrics only; trained on detached features if vae is not None: c = vae.config.latent_channels # the head upsamples the patch grid 2x, so it must land exactly on the VAE latent grid lat_side = image_size // 2 ** (len(vae.config.block_out_channels) - 1) if self.side * 2 != lat_side: raise ValueError( f"patch grid {self.side}x2={self.side * 2} != VAE latent grid {lat_side} " f"(image_size={image_size}, dino_size={dino_size}); pick dino_size = " f"{lat_side // 2 * self.patch}" ) self.to_latent = nn.Sequential( nn.Conv2d(d, 4 * c, 3, padding=1), nn.GELU(), nn.PixelShuffle(2), nn.Conv2d(c, c, 3, padding=1) ) # -- features -- @torch.no_grad() def encode(self, imgs): """(N,3,H,W) in [0,1] -> (N, n_patches, d) DINOv2 patch tokens.""" x = F.interpolate(imgs, self.dino_size, mode="bilinear", antialias=True, align_corners=False) mean = torch.tensor(DINO_MEAN, device=x.device).view(1, 3, 1, 1) std = torch.tensor(DINO_STD, device=x.device).view(1, 3, 1, 1) out = self.dino(pixel_values=((x - mean) / std).to(self.dino.dtype)).last_hidden_state return out[:, 1:].float() # drop CLS def encode_views(self, imgs): """(B,R,V,3,H,W) -> (B,R,V,P,d).""" b, r, v = imgs.shape[:3] z = self.encode(imgs.flatten(0, 2)) return z.reshape(b, r, v, self.n_patches, self.dim) # -- prediction -- def forward(self, batch): """-> predicted target features (B,V,P,d).""" return self.predict_from_features( self.encode_views(batch["context"]), batch["action"], batch["state"] ) def predict_from_features(self, z, action, state): """Context features (B,R,V,P,d) -> predicted target features (B,V,P,d). Split out from forward so a planner can encode the (fixed) context once and reuse it across every CEM/gradient iteration; only this part is on the action-gradient path. """ b, r, v = z.shape[:3] z = z + self.view_emb[None, None, :, None] + self.pos_emb[None, None, None] z = z + self.frame_emb[:r][None, :, None, None] tokens = z.reshape(b, r * v * self.n_patches, self.dim) q = self.query[None, None] + self.view_emb[None, :, None] + self.pos_emb[None, None] q = q + self.frame_emb[-1][None, None, None] q = q.expand(b, v, self.n_patches, self.dim).reshape(b, v * self.n_patches, self.dim) act = self.act_tok(action.flatten(1)) * float(self.use_action) st = self.state_tok(state) * float(self.use_state) cond = torch.stack([act, st], 1) # (B,2,d) out = self.predictor(torch.cat([tokens, cond, q], 1)) return self.norm(out[:, -v * self.n_patches :]).reshape(b, v, self.n_patches, self.dim) def loss(self, batch): pred = self(batch) with torch.no_grad(): tgt = self.encode_views(batch["future"][:, None])[:, 0] # (B,V,P,d) return F.mse_loss(pred, tgt), tgt # -- pixels, for metrics and W&B only -- def features_to_image(self, feats): """(B,V,P,d) -> (B,V,3,H,W) in [0,1], through the frozen SD3 VAE.""" b, v = feats.shape[:2] x = feats.reshape(b * v, self.side, self.side, self.dim).permute(0, 3, 1, 2) lat = self.to_latent(x) z = lat.to(self.vae.dtype) / self.vae.config.scaling_factor + self.vae.config.shift_factor img = (self.vae.decode(z).sample / 2 + 0.5).clamp(0, 1).float() return img.reshape(b, v, *img.shape[1:]) def decoder_loss(self, batch, tgt_feats): """Train the feature->latent head against the true frames. Detached from the predictor.""" with torch.no_grad(): imgs = batch["future"].flatten(0, 1) z = self.vae.encode(imgs.to(self.vae.dtype) * 2 - 1).latent_dist.mode() z = ((z - self.vae.config.shift_factor) * self.vae.config.scaling_factor).float() b, v = batch["future"].shape[:2] x = tgt_feats.detach().reshape(b * v, self.side, self.side, self.dim).permute(0, 3, 1, 2) return F.mse_loss(self.to_latent(x), z) def main(): ap = argparse.ArgumentParser(description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter) ap.add_argument("--data", nargs="+", required=True) ap.add_argument("--dino", default=DINO_ID) ap.add_argument("--sd3", default="stabilityai/stable-diffusion-3.5-medium", help="VAE for decoding") ap.add_argument("--state-key", default="observation.state") ap.add_argument("--state-dim", type=int, default=9) ap.add_argument("--image-size", type=int, nargs=2, default=(224, 224)) ap.add_argument("--history", type=int, default=2, help="context frames incl. current") ap.add_argument("--history-stride", type=int, default=HORIZON, help="default = horizon: evenly spaced") ap.add_argument("--no-action", action="store_true", help="ablation: zero the action conditioning") ap.add_argument("--no-state", action="store_true", help="ablation: zero the state conditioning") ap.add_argument("--depth", type=int, default=6) ap.add_argument("--heads", type=int, default=6) ap.add_argument("--batch-size", type=int, default=32) ap.add_argument("--steps", type=int, default=20000) ap.add_argument("--lr", type=float, default=5e-4, help="reference predictor_lr") ap.add_argument("--decoder-lr", type=float, default=3e-4) ap.add_argument("--log-every", type=int, default=100) ap.add_argument("--val-episodes", type=int, default=20) ap.add_argument("--val-every", type=int, default=500) ap.add_argument("--val-batches", type=int, default=4) ap.add_argument("--log-images", type=int, default=4) ap.add_argument("--workers", type=int, default=8) ap.add_argument("--out", default=str(Path(__file__).parent / "outputs")) ap.add_argument("--wandb-project", default="sd3-dynamics") ap.add_argument("--run-name", default=None) args = ap.parse_args() dev = "cuda" if torch.cuda.is_available() else "cpu" torch.manual_seed(0) image_size = tuple(args.image_size) train_ds, val_ds, cameras = build_datasets( args.data, image_size, args.state_key, args.state_dim, args.val_episodes, args.history, args.history_stride, ) feats = dataset_features(args.data[0]) assert args.state_key not in object_state_keys(feats), "object state must not reach the model" action_dim = feats["action"]["shape"][0] print(f"cameras={cameras} action_dim={action_dim} state={args.state_key}[:{args.state_dim}]") print(f"excluded from the model: {object_state_keys(feats)}") print(f"train windows={len(train_ds)} val windows={len(val_ds) if val_ds else 0}") from diffusers import AutoencoderKL from transformers import AutoModel dino = AutoModel.from_pretrained(args.dino, torch_dtype=torch.float32) vae = AutoencoderKL.from_pretrained(args.sd3, subfolder="vae", torch_dtype=torch.bfloat16) model = DinoDynamics( dino, vae, len(cameras), action_dim, args.state_dim, args.history, args.depth, args.heads, image_size=image_size[0], use_action=not args.no_action, use_state=not args.no_state, ).to(dev) pred_params = [p for n, p in model.named_parameters() if p.requires_grad and not n.startswith("to_latent")] dec_params = list(model.to_latent.parameters()) print(f"predictor params: {sum(p.numel() for p in pred_params) / 1e6:.1f}M " f"decoder head: {sum(p.numel() for p in dec_params) / 1e6:.1f}M " f"tokens: {args.history * len(cameras) * model.n_patches} ctx + {len(cameras) * model.n_patches} query") opt = torch.optim.AdamW( [{"params": pred_params, "lr": args.lr}, {"params": dec_params, "lr": args.decoder_lr}] ) lpips = None if val_ds is not None: from torchmetrics.image.lpip import LearnedPerceptualImagePatchSimilarity lpips = LearnedPerceptualImagePatchSimilarity(net_type="alex", normalize=True).to(dev) run = None if args.wandb_project: import wandb run = wandb.init( project=args.wandb_project, name=args.run_name, config=vars(args) | {"cameras": cameras} ) def loader(ds, shuffle): return torch.utils.data.DataLoader( ds, batch_size=args.batch_size, shuffle=shuffle, num_workers=args.workers, drop_last=True, pin_memory=True, ) def to_dev(b): return {k: v.to(dev, non_blocking=True) for k, v in b.items()} def log(d, step): print(f"step {step}: " + " ".join(f"{k}={v:.4f}" for k, v in d.items() if isinstance(v, float))) if run: run.log(d, step=step) @torch.no_grad() def validate(step): model.eval() sums, n = {}, 0 for i, batch in enumerate(loader(val_ds, False)): if i >= args.val_batches: break batch = to_dev(batch) pred = model(batch) tgt = model.encode_views(batch["future"][:, None])[:, 0] cur = model.encode_views(batch["context"][:, -1:])[:, 0] sums["feat_mse"] = sums.get("feat_mse", 0.0) + F.mse_loss(pred, tgt).item() sums["copy_feat_mse"] = sums.get("copy_feat_mse", 0.0) + F.mse_loss(cur, tgt).item() img = model.features_to_image(pred) current = batch["context"][:, -1] for k, v in per_camera_metrics(img, batch["future"], cameras, lpips).items(): sums[k] = sums.get(k, 0.0) + v for k, v in per_camera_metrics(current, batch["future"], cameras, lpips).items(): sums[f"copy_{k}"] = sums.get(f"copy_{k}", 0.0) + v # decoder ceiling: what the feature->pixel head gives on GROUND-TRUTH features for k, v in per_camera_metrics( model.features_to_image(tgt), batch["future"], cameras, lpips ).items(): sums[f"oracle_{k}"] = sums.get(f"oracle_{k}", 0.0) + v n += 1 if i == 0 and run: k = min(args.log_images, img.shape[0]) run.log( {"val/reconstructions": reconstruction_table( current[:k], batch["future"][:k], img[:k], cameras)}, step=step, ) model.train() log({f"val/{k}": v / max(n, 1) for k, v in sums.items()}, step) out = Path(args.out) out.mkdir(parents=True, exist_ok=True) step, running, rdec, t0 = 0, 0.0, 0.0, time.time() model.train() while step < args.steps: for batch in loader(train_ds, True): batch = to_dev(batch) feat_loss, tgt = model.loss(batch) dec_loss = model.decoder_loss(batch, tgt) (feat_loss + dec_loss).backward() opt.step() opt.zero_grad(set_to_none=True) running += feat_loss.item() rdec += dec_loss.item() step += 1 if step % args.log_every == 0: dt = (time.time() - t0) / args.log_every log( { "train/feat_mse": running / args.log_every, "train/dec_mse": rdec / args.log_every, "train/peak_gb": torch.cuda.max_memory_allocated() / 1e9 if dev == "cuda" else 0.0, "train/s_per_step": dt, "train/samples_per_s": args.batch_size / dt, }, step, ) running, rdec, t0 = 0.0, 0.0, time.time() if val_ds is not None and step % args.val_every == 0: validate(step) if step % 5000 == 0 or step == args.steps: torch.save( {k: v for k, v in model.state_dict().items() if not k.startswith(("dino.", "vae."))}, out / f"dino_step_{step}.pt", ) if step >= args.steps: break if __name__ == "__main__": main()