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()