| """VLANeXt inference server for RoboLab eval (ZMQ REQ/REP + msgpack). |
| |
| RoboLab runs in its own Isaac-Sim uv venv (Python 3.11); VLANeXt runs in the |
| fla_triton32 venv (torch 2.6 / triton 3.2). They cannot share a process, so the |
| model is served here and RoboLab connects via policies/vlanext/client.py. |
| |
| Wire protocol (mirrors policies/gr00t/client.py's self-contained msgpack codec, so |
| no openpi dependency is needed on either side): |
| |
| request = { |
| "exterior_image": uint8 HWC ndarray, # over-shoulder camera |
| "wrist_image": uint8 HWC ndarray, # wrist camera |
| "joint_position": float32 (7,) ndarray, # arm joints, rad |
| "gripper_position": float32 (1,) ndarray, # [0,1], 0=open 1=closed |
| "prompt": str, |
| "reset": bool, # optional, clears history |
| } |
| response = {"actions": float32 (horizon, 8) ndarray} # 7 joint (rad) + gripper [0,1] |
| |
| Actions are produced by the model in normalized [-1, 1] space and de-normalized |
| here with the MolmoAct2-DROID per-dim stats before being returned, so the client |
| can command RoboLab's DroidJointPositionActionCfg directly (7 joints in rad + |
| BinaryJointPositionZeroToOneAction gripper in [0, 1]). |
| """ |
|
|
| import io |
| import argparse |
|
|
| import numpy as np |
| import torch |
| import zmq |
| import msgpack |
|
|
| from src.evaluation.libero_bench.VLANeXt_utils import ( |
| get_vla as get_model, |
| get_processor, |
| get_vla_action, |
| ) |
| from src.datasets.molmoact_droid_act import get_molmoact_action_stats |
|
|
|
|
| |
| |
| |
| def _encode(obj): |
| if isinstance(obj, np.ndarray): |
| buf = io.BytesIO() |
| np.save(buf, obj, allow_pickle=False) |
| return {"__ndarray__": True, "as_npy": buf.getvalue()} |
| return obj |
|
|
|
|
| def _decode(obj): |
| if isinstance(obj, dict) and "__ndarray__" in obj: |
| return np.load(io.BytesIO(obj["as_npy"]), allow_pickle=False) |
| return obj |
|
|
|
|
| def to_bytes(data): |
| return msgpack.packb(data, default=_encode) |
|
|
|
|
| def from_bytes(data): |
| return msgpack.unpackb(data, object_hook=_decode, raw=False) |
|
|
|
|
| class _Cfg: |
| """Minimal dot-access cfg for get_vla/get_processor (mirrors libero eval DictConfig).""" |
|
|
| def __init__(self, d): |
| for k, v in d.items(): |
| setattr(self, k, _Cfg(v) if isinstance(v, dict) else v) |
|
|
|
|
| def build_cfg(checkpoint, image_size, diffusion_steps, use_incremental_gen, ttt_cuda): |
| return _Cfg({ |
| "eval": { |
| "finetuned_checkpoint": checkpoint, |
| "image_size": image_size, |
| "diffusion_steps": diffusion_steps, |
| "use_incremental_gen": use_incremental_gen, |
| "ttt_use_cuda_kernel": ttt_cuda, |
| }, |
| "model": { |
| "attn_implementation": "sdpa", |
| "diffusion_steps": diffusion_steps, |
| }, |
| }) |
|
|
|
|
| def main(): |
| p = argparse.ArgumentParser() |
| p.add_argument("--checkpoint", required=True, help="VLANeXt .pt checkpoint") |
| p.add_argument("--data-root", default="/mnt/afs-h200/NTU_slab/draven/data/MolmoAct2-DROID", |
| help="MolmoAct2-DROID root (for action denorm stats)") |
| p.add_argument("--port", type=int, default=5556) |
| p.add_argument("--image-size", type=int, default=256, |
| help="Eval image size (Qwen processor resizes; matches train resolution policy)") |
| p.add_argument("--diffusion-steps", type=int, default=10) |
| p.add_argument("--use-incremental-gen", action="store_true", |
| help="O(n) incremental image-token decode in predict_action") |
| p.add_argument("--ttt-cuda", action="store_true", help="Use TTT CUDA kernel at eval") |
| args = p.parse_args() |
|
|
| cfg = build_cfg(args.checkpoint, args.image_size, args.diffusion_steps, |
| args.use_incremental_gen, args.ttt_cuda) |
|
|
| print(f"[serve] loading model from {args.checkpoint}") |
| model = get_model(cfg) |
| |
| |
| |
| processor = getattr(model, "processor", None) |
| if processor is None: |
| processor = get_processor(cfg) |
| view_mode = model.train_config["data"].get("view_mode", "single") |
| print(f"[serve] view_mode={view_mode} num_history={getattr(model, 'num_history', 0)} " |
| f"action_dim={model.action_dim}") |
|
|
| |
| amin, amax = get_molmoact_action_stats(args.data_root) |
| amin = amin.astype(np.float32) |
| amax = amax.astype(np.float32) |
| print(f"[serve] denorm stats loaded: action dim={amin.shape[0]}") |
|
|
| def denorm(chunk): |
| |
| return ((chunk + 1.0) / 2.0) * (amax - amin) + amin |
|
|
| ctx = zmq.Context() |
| sock = ctx.socket(zmq.REP) |
| sock.bind(f"tcp://0.0.0.0:{args.port}") |
| print(f"[serve] listening on tcp://0.0.0.0:{args.port}") |
|
|
| while True: |
| msg = sock.recv() |
| try: |
| req = from_bytes(msg) |
|
|
| if req.get("reset", False): |
| |
| |
| sock.send(to_bytes({"ok": True})) |
| continue |
|
|
| ext = np.asarray(req["exterior_image"], dtype=np.uint8) |
| wrist = np.asarray(req["wrist_image"], dtype=np.uint8) |
| joint = np.asarray(req["joint_position"], dtype=np.float32).reshape(-1) |
| grip = np.asarray(req["gripper_position"], dtype=np.float32).reshape(-1) |
| prompt = str(req.get("prompt", "")) |
|
|
| |
| |
| state = np.concatenate([joint[:7], grip[:1]]).astype(np.float32) |
| state_norm = np.clip(2.0 * (state - amin) / np.where(amax - amin == 0, 1.0, amax - amin) - 1.0, |
| -1.0, 1.0) |
|
|
| |
| |
| obs = { |
| "full_image": ext, |
| "full_image_wrist": wrist if view_mode == "multi" else ext, |
| "image_history": [ext], |
| "image_history_wrist": [wrist] if view_mode == "multi" else [], |
| "state_history": [state_norm], |
| "action_history": [], |
| } |
|
|
| chunk_norm = get_vla_action(cfg, model, processor, obs, prompt) |
| if chunk_norm.ndim == 1: |
| chunk_norm = chunk_norm[None, :] |
| chunk = denorm(chunk_norm.astype(np.float32)) |
|
|
| sock.send(to_bytes({"actions": chunk.astype(np.float32)})) |
|
|
| except Exception as e: |
| import traceback |
| traceback.print_exc() |
| sock.send(to_bytes({"error": str(e)})) |
|
|
|
|
| if __name__ == "__main__": |
| main() |
|
|