File size: 7,358 Bytes
f6d03a4 | 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 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 | """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
# --------------------------------------------------------------------------
# msgpack ndarray codec (identical to policies/gr00t/client.py _MsgSerializer)
# --------------------------------------------------------------------------
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)
# The VLANeXt model builds its own processor internally; get_vla_action prefers
# model.processor. Reuse it instead of calling get_processor (which would re-read
# the full multi-GB checkpoint from disk a second time just for the lmm_path).
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}")
# De-normalization stats: normalized [-1,1] -> physical units (rad / [0,1] gripper).
amin, amax = get_molmoact_action_stats(args.data_root) # (8,), (8,)
amin = amin.astype(np.float32)
amax = amax.astype(np.float32)
print(f"[serve] denorm stats loaded: action dim={amin.shape[0]}")
def denorm(chunk):
# chunk: (horizon, 8) in [-1, 1]. Inverse of dataset _normalize.
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):
# No server-side session state to clear (history is sent each call),
# but acknowledge so the client can flush its chunk cache.
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", ""))
# Proprioception state = normalized [joint(7), gripper(1)] (8-dim), matching
# the training observation.state layout. Normalize with the same stats.
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)
# Single-step obs (history is rebuilt from the single frame inside
# get_vla_action via _take_last padding when no history is provided).
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) # (horizon, 8)
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()
|