farm_uf850_patch_policy / train_patch_policy.py
noahnowac's picture
Patch Policy task_2 L1: ckpt + inference runner + training script
c7b77de verified
Raw
History Blame Contribute Delete
12.9 kB
"""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()