File size: 1,601 Bytes
c7b77de
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
"""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)