| """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") |
|
|
|
|
| |
| 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) |
|
|
|
|
| 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)), |
| torch.from_numpy(np.stack(states)).float(), |
| torch.from_numpy(np.stack(chunks)).float(), |
| torch.from_numpy(np.stack(masks)).float()) |
|
|
|
|
| |
| 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): |
| |
| 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"] |
| 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] |
|
|
| 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) |
|
|
|
|
| |
| 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) |
| 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() |
|
|