groot-n1.7-mlx / preprocess_lerobot.py
tokimoa's picture
groot-n1.7-mlx: MLX port (E2E parity cos 1.000000) + runtime + NVIDIA non-commercial license
e1ee88a verified
Raw
History Blame Contribute Delete
2.13 kB
"""GR00T N1.7の前処理をLeRobotパイプラインで実行し、MLXランタイム用の.ptを出力する。
要: pip install 'lerobot[groot] @ git+https://github.com/huggingface/lerobot' (Python 3.12+)
+ Hugging FaceでのCosmos-Reason2-2B利用同意(gated)
Usage: python preprocess_lerobot.py --ckpt <this_repo_dir> --image cam.png \
--state 0,0,... --task "pick up the cube" --embodiment-tag <tag> --out processed.pt
"""
import argparse, json
import numpy as np, torch
from PIL import Image
from lerobot.policies.groot.modeling_groot import GrootPolicy
from lerobot.policies.groot.processor_groot import make_groot_pre_post_processors_from_pretrained
ap = argparse.ArgumentParser()
ap.add_argument("--ckpt", required=True)
ap.add_argument("--image", required=True)
ap.add_argument("--state", required=True)
ap.add_argument("--task", required=True)
ap.add_argument("--embodiment-tag", default=None)
ap.add_argument("--out", default="processed.pt")
args = ap.parse_args()
policy = GrootPolicy.from_pretrained(args.ckpt)
policy.to("cpu").eval()
if args.embodiment_tag:
policy.config.embodiment_tag = args.embodiment_tag
pre, _ = make_groot_pre_post_processors_from_pretrained(policy.config, args.ckpt)
img = np.asarray(Image.open(args.image).convert("RGB").resize((224, 224))).astype("float32") / 255.0
state = np.array([float(x) for x in args.state.split(",")], dtype="float32")
state132 = np.zeros(132, dtype="float32"); state132[: len(state)] = state
batch = {
"observation.images.camera": torch.tensor(img).permute(2, 0, 1)[None],
"observation.state": torch.tensor(state132)[None],
"task": [args.task],
}
proc = pre(batch)
proc = {k: (v.cpu() if isinstance(v, torch.Tensor) else v) for k, v in proc.items()}
model = policy._groot_model.cpu()
with torch.no_grad():
backbone_in, action_in = model.prepare_input(proc)
tensors = {k: v for k, v in dict(backbone_in).items() if isinstance(v, torch.Tensor)}
tensors["state"] = action_in.state.cpu()
tensors["embodiment_id"] = action_in.embodiment_id.cpu()
torch.save(tensors, args.out)
print(f"saved {args.out}: {sorted(tensors)}")