farm_uf850_patch_policy / infer_patch_policy.py
noahnowac's picture
Patch Policy task_2 L1: ckpt + inference runner + training script
c7b77de verified
Raw
History Blame Contribute Delete
1.6 kB
"""Minimal inference for the FARM UF850 Patch Policy specialist (L1 head).
Inputs: base + wrist RGB frames (any size), 7-d proprio state.
Output: 50-step action chunk in dataset space (6 relative joint deltas rad + gripper [0,1])."""
import cv2
import numpy as np
import torch
from train_patch_policy import PatchPolicy
class PatchPolicyRunner:
def __init__(self, ckpt_path, device="cuda"):
ck = torch.load(ckpt_path, map_location="cpu", weights_only=False)
self.norm = ck["norm"]
cfg = ck["cfg"]
self.model = PatchPolicy(cfg["chunk"], head=cfg["head"]).to(device).eval()
self.model.load_state_dict(ck["ema"], strict=False) # EMA weights (incl. frozen ViT)
self.device = device
@torch.no_grad()
def predict(self, base_rgb, wrist_rgb, state7):
imgs = np.stack([cv2.resize(base_rgb, (224, 224), interpolation=cv2.INTER_AREA),
cv2.resize(wrist_rgb, (224, 224), interpolation=cv2.INTER_AREA)])[None]
st = ((np.asarray(state7, np.float32) - self.norm["s_mean"]) / self.norm["s_std"])[None]
with torch.autocast("cuda", torch.bfloat16):
chunk = self.model.predict(torch.from_numpy(imgs).to(self.device),
torch.from_numpy(st).float().to(self.device))
return chunk[0].float().cpu().numpy() * self.norm["a_std"] + self.norm["a_mean"] # (50, 7)
if __name__ == "__main__":
r = PatchPolicyRunner("patch_policy_task2_l1.pt")
dummy = np.zeros((480, 640, 3), np.uint8)
print("chunk:", r.predict(dummy, dummy, np.zeros(7)).shape)