"""Patch Policy (arXiv 2607.18236) on FARM UF850 — single-task specialist. Frozen DINOv2-S/14 dense patch tokens (all 256/cam, no pooling) -> small transformer -> action-chunk head (L1 regression or DDPM/DDIM diffusion). No language conditioning (single-task, per-SKU analog). Self-contained: lerobot-v2 reader, RAM preload of the episode subset (FARM subsets are 13-111 eps -> <=20GB), per-subset normalization, EMA, held-out open-loop eval (triage screen only — rollouts on the arm are the verdict). Deviations from paper (noted): T=1 obs (no temporal context -> no block-causal mask); action readout via a learned [ACT] query token; diffusion head is a conditional MLP denoiser rather than their (unspecified) DP head. Run (openpi venv, 1 GPU): python train_patch_policy.py --data-root ~/data/NoahWeiss/farm_uf850_home_full \ --episodes-file ~/farm_task_eps/task_2.json --out ~/pp_runs/task2_l1 --head l1 """ import argparse import json import math import os import time from concurrent.futures import ThreadPoolExecutor import av import cv2 import numpy as np import pyarrow.parquet as pq import torch import torch.nn as nn import torch.nn.functional as F IMNET_MEAN = torch.tensor([0.485, 0.456, 0.406]).view(1, 3, 1, 1) IMNET_STD = torch.tensor([0.229, 0.224, 0.225]).view(1, 3, 1, 1) CAMS = ("observation.images.base", "observation.images.wrist") # ---------------------------------------------------------------- data def decode_video(path, size=224): frames = [] with av.open(path) as c: for fr in c.decode(c.streams.video[0]): img = fr.to_ndarray(format="rgb24") frames.append(cv2.resize(img, (size, size), interpolation=cv2.INTER_AREA)) return np.stack(frames) # (T, H, W, 3) uint8 class FarmSubset: """Preloads one episode subset fully into RAM.""" def __init__(self, root, episode_ids, chunk): self.chunk = chunk self.eps = [] t0 = time.time() def load_one(ei): ch = ei // 1000 pf = pq.read_table(f"{root}/data/chunk-{ch:03d}/episode_{ei:06d}.parquet").to_pydict() state = np.asarray(pf["observation.state"], np.float32) action = np.asarray(pf["action"], np.float32) vids = [decode_video(f"{root}/videos/chunk-{ch:03d}/{cam}/episode_{ei:06d}.mp4") for cam in CAMS] n = min(len(state), len(action), *[len(v) for v in vids]) return {"state": state[:n], "action": action[:n], "base": vids[0][:n], "wrist": vids[1][:n], "n": n} with ThreadPoolExecutor(max_workers=16) as ex: self.eps = list(ex.map(load_one, episode_ids)) frames = sum(e["n"] for e in self.eps) print(f"preloaded {len(self.eps)} eps / {frames} frames in {time.time()-t0:.0f}s " f"({sum(e['base'].nbytes + e['wrist'].nbytes for e in self.eps)/1e9:.1f} GB)", flush=True) acts = np.concatenate([e["action"] for e in self.eps]) sts = np.concatenate([e["state"] for e in self.eps]) self.a_mean, self.a_std = acts.mean(0), acts.std(0) + 1e-6 self.s_mean, self.s_std = sts.mean(0), sts.std(0) + 1e-6 self.index = [(i, t) for i, e in enumerate(self.eps) for t in range(e["n"])] def sample(self, rng, batch): idx = rng.choice(len(self.index), batch) imgs, states, chunks, masks = [], [], [], [] H = self.chunk for k in idx: i, t = self.index[k] e = self.eps[i] imgs.append(np.stack([e["base"][t], e["wrist"][t]])) states.append((e["state"][t] - self.s_mean) / self.s_std) n_valid = min(H, e["n"] - t) ch = np.repeat(e["action"][t + n_valid - 1][None], H, 0) ch[:n_valid] = e["action"][t:t + n_valid] chunks.append((ch - self.a_mean) / self.a_std) m = np.zeros(H, np.float32); m[:n_valid] = 1 masks.append(m) return (torch.from_numpy(np.stack(imgs)), # (B, 2, 224, 224, 3) u8 torch.from_numpy(np.stack(states)).float(), # (B, 7) torch.from_numpy(np.stack(chunks)).float(), # (B, H, 7) torch.from_numpy(np.stack(masks)).float()) # (B, H) # ---------------------------------------------------------------- model class DiffusionHead(nn.Module): """Conditional MLP denoiser over the flattened action chunk. DDPM train / DDIM sample.""" def __init__(self, cond_dim, chunk, adim, T=100): super().__init__() self.T, self.out = T, chunk * adim t = torch.linspace(0, 1, T + 1) abar = torch.cos((t + 0.008) / 1.008 * math.pi / 2) ** 2 self.register_buffer("abar", abar / abar[0]) self.temb = nn.Sequential(nn.Linear(128, 256), nn.GELU(), nn.Linear(256, 256)) self.net = nn.Sequential( nn.Linear(self.out + cond_dim + 256, 1024), nn.GELU(), nn.Linear(1024, 1024), nn.GELU(), nn.Linear(1024, 1024), nn.GELU(), nn.Linear(1024, self.out)) def t_embed(self, t): half = 64 freqs = torch.exp(-math.log(1000) * torch.arange(half, device=t.device) / half) ang = t[:, None].float() * freqs[None] return self.temb(torch.cat([ang.sin(), ang.cos()], -1)) def eps(self, x, t, cond): return self.net(torch.cat([x, cond, self.t_embed(t)], -1)) def loss(self, chunk, mask, cond): B = chunk.shape[0] x0 = chunk.flatten(1) t = torch.randint(1, self.T + 1, (B,), device=x0.device) ab = self.abar[t][:, None] noise = torch.randn_like(x0) xt = ab.sqrt() * x0 + (1 - ab).sqrt() * noise pred = self.eps(xt, t, cond) m = mask[:, :, None].expand(-1, -1, chunk.shape[-1]).flatten(1) return ((pred - noise) ** 2 * m).sum() / m.sum().clamp(min=1) @torch.no_grad() def sample(self, cond, steps=10): B = cond.shape[0] x = torch.randn(B, self.out, device=cond.device) ts = torch.linspace(self.T, 0, steps + 1).long().to(cond.device) for i in range(steps): t, tn = ts[i], ts[i + 1] ab, abn = self.abar[t], self.abar[tn] e = self.eps(x, t.expand(B), cond) x0 = ((x - (1 - ab).sqrt() * e) / ab.sqrt()).clamp(-4, 4) x = abn.sqrt() * x0 + (1 - abn).sqrt() * e return x class PatchPolicy(nn.Module): def __init__(self, chunk, adim=7, d=512, layers=8, heads=8, head="l1", n_patch=256): super().__init__() self.backbone = torch.hub.load("facebookresearch/dinov2", "dinov2_vits14") self.backbone.eval().requires_grad_(False) self.proj = nn.Linear(384, d) self.cam_emb = nn.Parameter(torch.zeros(2, 1, d)) self.pos = nn.Parameter(torch.zeros(2 * n_patch, d).normal_(std=0.02)) self.state_in = nn.Sequential(nn.Linear(7, d), nn.GELU(), nn.Linear(d, d)) self.act_query = nn.Parameter(torch.zeros(1, 1, d).normal_(std=0.02)) enc = nn.TransformerEncoderLayer(d, heads, 4 * d, batch_first=True, norm_first=True, activation="gelu", dropout=0.1) self.encoder = nn.TransformerEncoder(enc, layers) self.head_type = head self.chunk, self.adim = chunk, adim if head == "l1": self.head = nn.Sequential(nn.Linear(d, 1024), nn.GELU(), nn.Linear(1024, chunk * adim)) else: self.head = DiffusionHead(d, chunk, adim) def encode(self, imgs, states): # imgs (B, 2, H, W, 3) uint8 B = imgs.shape[0] x = imgs.flatten(0, 1).permute(0, 3, 1, 2).float() / 255.0 x = (x - IMNET_MEAN.to(x)) / IMNET_STD.to(x) with torch.no_grad(): tok = self.backbone.forward_features(x)["x_norm_patchtokens"] # (2B, 256, 384) tok = self.proj(tok) tok = tok + self.cam_emb.repeat_interleave(B, 0).to(tok) tok = tok.reshape(B, -1, tok.shape[-1]) + self.pos[None].to(tok) st = self.state_in(states)[:, None] seq = torch.cat([tok, st, self.act_query.expand(B, -1, -1).to(tok)], 1) return self.encoder(seq)[:, -1] # [ACT] token output def loss(self, imgs, states, chunk, mask): cond = self.encode(imgs, states) if self.head_type == "l1": pred = self.head(cond).view(-1, self.chunk, self.adim) return (F.l1_loss(pred, chunk, reduction="none") * mask[:, :, None]).sum() / \ (mask.sum() * self.adim).clamp(min=1) return self.head.loss(chunk, mask, cond) @torch.no_grad() def predict(self, imgs, states): cond = self.encode(imgs, states) if self.head_type == "l1": return self.head(cond).view(-1, self.chunk, self.adim) return self.head.sample(cond).view(-1, self.chunk, self.adim) # ---------------------------------------------------------------- train def main(): ap = argparse.ArgumentParser() ap.add_argument("--data-root", required=True) ap.add_argument("--episodes-file", required=True) ap.add_argument("--out", required=True) ap.add_argument("--head", choices=["l1", "diff"], default="l1") ap.add_argument("--chunk", type=int, default=50) ap.add_argument("--steps", type=int, default=30000) ap.add_argument("--bs", type=int, default=64) ap.add_argument("--lr", type=float, default=3e-4) ap.add_argument("--val-frac", type=int, default=10, help="every Nth episode held out") ap.add_argument("--eval-every", type=int, default=2000) ap.add_argument("--save-every", type=int, default=10000) args = ap.parse_args() os.makedirs(args.out, exist_ok=True) dev = "cuda" torch.manual_seed(0) rng = np.random.default_rng(0) ids = sorted(json.load(open(os.path.expanduser(args.episodes_file)))) val_ids = ids[:: args.val_frac] train_ids = [i for i in ids if i not in set(val_ids)] print(f"episodes: {len(train_ids)} train / {len(val_ids)} val", flush=True) root = os.path.expanduser(args.data_root) tr = FarmSubset(root, train_ids, args.chunk) va = FarmSubset(root, val_ids, args.chunk) va.a_mean, va.a_std, va.s_mean, va.s_std = tr.a_mean, tr.a_std, tr.s_mean, tr.s_std model = PatchPolicy(args.chunk, head=args.head).to(dev) trainable = [p for p in model.parameters() if p.requires_grad] print(f"trainable params: {sum(p.numel() for p in trainable)/1e6:.2f}M", flush=True) opt = torch.optim.AdamW(trainable, lr=args.lr, weight_decay=0.05) sched = torch.optim.lr_scheduler.CosineAnnealingLR(opt, args.steps, eta_min=1e-5) ema = {k: v.detach().clone() for k, v in model.state_dict().items() if v.dtype.is_floating_point} def run_eval(n=200): model.eval() maes, mae1 = [], [] for _ in range(n // 50): imgs, st, ch, m = va.sample(rng, 50) with torch.autocast("cuda", torch.bfloat16): pred = model.predict(imgs.to(dev), st.to(dev)) err = (pred.float().cpu() - ch).abs() * torch.from_numpy(tr.a_std) # raw units maes.append((err * m[:, :, None]).sum() / (m.sum() * 7)) mae1.append(err[:, 0].mean()) model.train() return float(np.mean(maes)), float(np.mean(mae1)) model.train() t0, losses = time.time(), [] for step in range(1, args.steps + 1): imgs, st, ch, m = tr.sample(rng, args.bs) with torch.autocast("cuda", torch.bfloat16): loss = model.loss(imgs.to(dev, non_blocking=True), st.to(dev), ch.to(dev), m.to(dev)) opt.zero_grad(set_to_none=True) loss.backward() torch.nn.utils.clip_grad_norm_(trainable, 1.0) opt.step(); sched.step() with torch.no_grad(): sd = model.state_dict() for k in ema: ema[k].mul_(0.999).add_(sd[k].detach(), alpha=0.001) losses.append(loss.item()) if step % 100 == 0: r = 100 / (time.time() - t0); t0 = time.time() print(f"step {step}/{args.steps} loss={np.mean(losses):.4f} {r:.1f} it/s", flush=True) losses = [] if step % args.eval_every == 0: bak = {k: v.detach().clone() for k, v in model.state_dict().items() if k in ema} model.load_state_dict(ema, strict=False) mae, m1 = run_eval() model.load_state_dict(bak, strict=False) print(f"EVAL step {step}: heldout chunk-MAE={mae:.4f} first-action-MAE={m1:.4f} (rad, EMA)", flush=True) if step % args.save_every == 0 or step == args.steps: torch.save({"ema": ema, "cfg": vars(args), "norm": {"a_mean": tr.a_mean, "a_std": tr.a_std, "s_mean": tr.s_mean, "s_std": tr.s_std}}, f"{args.out}/ckpt_{step}.pt") print("DONE", flush=True) if __name__ == "__main__": main()