File size: 17,914 Bytes
40fa6ec
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
f78bc70
40fa6ec
6e53006
df7bb3b
40fa6ec
 
 
 
 
 
 
 
 
f5a56d6
 
40fa6ec
df7bb3b
 
 
 
 
40fa6ec
 
 
 
f78bc70
 
 
 
40fa6ec
 
 
 
 
f5a56d6
 
 
 
 
 
 
 
40fa6ec
df7bb3b
 
 
 
 
 
 
 
 
 
 
 
 
 
40fa6ec
 
 
 
 
 
 
df7bb3b
 
40fa6ec
df7bb3b
40fa6ec
df7bb3b
 
 
 
 
40fa6ec
 
df7bb3b
40fa6ec
 
 
 
3a60a27
 
 
 
 
 
 
40fa6ec
df7bb3b
40fa6ec
 
 
 
 
 
 
 
 
 
 
 
df7bb3b
 
 
 
 
 
 
 
 
 
 
 
 
 
 
40fa6ec
 
 
 
 
6e53006
 
 
 
 
 
 
 
 
 
f5a56d6
 
 
 
 
 
 
 
 
 
 
40fa6ec
 
 
 
6e53006
 
 
 
f5a56d6
 
40fa6ec
 
 
 
 
 
 
 
 
04f1403
40fa6ec
 
df7bb3b
40fa6ec
 
04f1403
 
 
 
 
 
 
 
 
 
 
 
 
 
df7bb3b
 
 
 
 
 
 
 
 
 
 
 
 
 
26dfbe2
 
 
 
 
 
 
 
 
 
 
 
 
df7bb3b
 
26dfbe2
 
 
 
 
 
df7bb3b
26dfbe2
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
40fa6ec
df7bb3b
 
 
40fa6ec
 
 
f78bc70
 
 
40fa6ec
 
6e53006
 
 
 
f5a56d6
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
6e53006
 
 
f5a56d6
 
 
 
 
 
 
 
6e53006
 
 
 
 
 
 
 
 
f5a56d6
6e53006
f5a56d6
 
 
 
 
 
 
 
 
 
 
 
6e53006
 
 
 
 
 
 
 
 
 
40fa6ec
 
 
 
 
 
f5a56d6
 
 
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
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
"""
Nori cloud-inference server for MolmoAct2-SO100_101  (spike β€” task #38).

Runs on an AWS GPU instance (g5.xlarge / A10G 24GB is enough in bf16 <16GB).
Serves the robot rollout over plain JSON (NO pickle on the wire β€” avoids the
LeRobot PolicyServer CVE-2026-25874 class):

    POST /act  { images:[b64...], state:[6 floats], instruction:str, num_steps? }
      -> { actions: [[...DOF...], ...] }   # a 10-30 move chunk, ROBOT SCALE

The model is loaded once at startup. Inference is serialized behind a lock
(single GPU). Bearer-token auth (NORI_INFER_TOKEN) on every call.

The exact model API mirrors the allenai/MolmoAct2-SO100_101 model card:
    model.predict_action(processor=..., images=[...], task=..., state=...,
        norm_tag="so100_so101_molmoact2", inference_action_mode="continuous",
        num_steps=10, normalize_language=True, enable_cuda_graph=True).actions

Deploy + test: see README.md in this directory.
"""

import base64
import io
import os
import secrets
import threading
import time
from pathlib import Path
from typing import Optional

import numpy as np
import torch
from fastapi import FastAPI, Header, HTTPException
from PIL import Image
from pydantic import BaseModel
from transformers import AutoModelForImageTextToText, AutoProcessor

import rtc as rtcmod

REPO_ID = os.environ.get("MOLMOACT_REPO", "allenai/MolmoAct2-SO100_101")
# Inference Endpoints mount the endpoint's model repo at /repository (platform
# fast-path β€” no 21GB Hub download at boot). Load from there when present, else
# fall back to the Hub download so the SAME image still runs as a Docker Space
# during the transition. Override the probe location with MODEL_PATH.
MODEL_PATH = os.environ.get("MODEL_PATH", "/repository")
NORM_TAG = os.environ.get("MOLMOACT_NORM_TAG", "so100_so101_molmoact2")
AUTH_TOKEN = os.environ.get("NORI_INFER_TOKEN")  # REQUIRED β€” the rollout sends it
# bf16 fits <16GB (A10G/L4). Set MOLMOACT_BF16=0 to run fp32 (~26GB, needs L40S/48GB).
DTYPE = torch.bfloat16 if os.environ.get("MOLMOACT_BF16", "1") == "1" else torch.float32
# Defensive caps: a valid-token caller can't burn unbounded GPU via a huge solver
# step count or a flood of images (the endpoint is public, token-gated).
MAX_NUM_STEPS = 50
MAX_IMAGES = 6

app = FastAPI(title="nori-molmoact2")
_model = None
_processor = None
_lock = threading.Lock()  # single GPU: serialize predict_action calls
# RTC (Real-Time Chunking). The flow-loop patch closes over ONE state object, but
# the guidance target is per-CLIENT β€” two rollouts sharing a server would otherwise
# steer each other's arms. Inference is already serialized behind _lock, so we keep
# one state and swap the cached chunk in/out per session under that same lock.
ACTION_HORIZON = 30            # this checkpoint's chunk length
_rtc_state = rtcmod.RTCState()
_rtc_prev: dict = {}           # session id -> previous chunk (normalized, on-device)
_RTC_SESSION_CAP = 8           # bound the cache; robot sessions are few and long-lived
_load_error: Optional[str] = None  # set if the background load failed
_model_source: Optional[str] = None  # /repository mount or the Hub repo id


def _resolve_model_source() -> str:
    """Prefer the platform-mounted weights (Inference Endpoints: /repository);
    fall back to the Hub repo id (Docker Space / bare GPU box). A non-empty dir
    is treated as the mount β€” trust_remote_code loads the model code from it."""
    p = Path(MODEL_PATH)
    try:
        if p.is_dir() and any(p.iterdir()):
            return str(p)
    except OSError:
        pass
    return REPO_ID


def _load_model() -> None:
    """Load weights in a background thread so the HTTP port is up immediately.

    MolmoAct2 is ~21GB β€” a blocking startup event would keep the port dark for
    minutes and a HuggingFace Space health-probe would kill the container as
    unhealthy before the model ever finishes loading. /health reports progress;
    /ready gives probes the 503-until-loaded semantic (Endpoints health_route).
    """
    global _model, _processor, _load_error, _model_source
    try:
        _model_source = _resolve_model_source()
        print(f"[molmoact2] loading from {_model_source} "
              f"({'mounted /repository' if _model_source != REPO_ID else 'Hub download'})",
              flush=True)
        proc = AutoProcessor.from_pretrained(_model_source, trust_remote_code=True)
        model = (
            AutoModelForImageTextToText.from_pretrained(
                _model_source, trust_remote_code=True, dtype=DTYPE
            )
            .to("cuda")
            .eval()
        )
        # Inert until a session sets a target. Guarded: RTC is an optimisation, and
        # nothing here may be allowed to stop the model from loading.
        try:
            if rtcmod.install_rtc(model, _rtc_state) is None:
                print("[molmoact2] RTC: flow loop not found β€” serving un-guided", flush=True)
        except Exception as exc:
            print(f"[molmoact2] RTC install failed ({exc}) β€” serving un-guided", flush=True)
        _processor, _model = proc, model
        print(f"[molmoact2] loaded {_model_source} dtype={DTYPE} (RTC patch installed)", flush=True)
    except Exception as exc:  # surface load failures via /health instead of a dead port
        _load_error = f"{type(exc).__name__}: {exc}"
        print(f"[molmoact2] LOAD FAILED β€” {_load_error}", flush=True)


@app.on_event("startup")
def _startup() -> None:
    if not AUTH_TOKEN:
        raise RuntimeError("NORI_INFER_TOKEN must be set (bearer token for /act)")
    threading.Thread(target=_load_model, name="molmoact2-load", daemon=True).start()


def _require_auth(x_nori_token: Optional[str], authorization: Optional[str]) -> None:
    """App-level auth for /act and /point. `X-Nori-Token` is the PRIMARY
    credential: on a *protected* Inference Endpoint HF's edge consumes the
    `Authorization` header (it must carry an HF token to get past the proxy), so
    our own bearer can no longer ride it β€” custom headers pass through untouched.
    `Authorization: Bearer <token>` stays accepted for the Space-transition
    client (which sends BOTH). Each comparison is constant-time; checked BEFORE
    any model work so unauthenticated calls never touch the GPU."""
    if x_nori_token and secrets.compare_digest(x_nori_token, AUTH_TOKEN):
        return
    if authorization and secrets.compare_digest(authorization, f"Bearer {AUTH_TOKEN}"):
        return
    raise HTTPException(status_code=401, detail="bad or missing auth token")


class ActRequest(BaseModel):
    images: list[str]      # base64 JPEG/PNG (optionally a data: URL), 2+ camera views
    state: list[float]     # robot joint state (6 for a single SO-100/101 arm)
    instruction: str       # natural-language task, e.g. "pick up the red cup"
    num_steps: int = 10    # flow-matching integration steps (latency <-> quality)
    # RTC FEASIBILITY PROBE (see cloud_inference/rtc.py). RTC's per-step PiGDM
    # correction is a VJP, so it needs (a) autograd enabled and (b) the CUDA-graph
    # fast path off. Both are optimisations we currently rely on, and latency is
    # already the binding constraint: after the chunk-stride fix the queue covers
    # ~1s of motion against a 0.65-1.2s round-trip. So measure the cost BEFORE
    # building the integration β€” if compute doubles, RTC needs a latency
    # reduction (in-region GPU) to be viable at all.
    # Cost is still bounded by MAX_NUM_STEPS and the bearer token.
    enable_cuda_graph: bool = True
    enable_grad: bool = False
    # RTC: send {session, consumed, delay} to make this chunk continuous with the
    # previous one. `consumed` aligns the cached chunk to the new timeline;
    # `delay` is how many actions will execute WHILE this inference runs, and
    # becomes the frozen prefix. Omit the block entirely to run without RTC.
    rtc: Optional["RTCParams"] = None


class RTCParams(BaseModel):
    session: str
    consumed: int = 0
    delay: int = 0


class ActResponse(BaseModel):
    actions: list[list[float]]  # chunk: N moves x DOF, ROBOT SCALE (already de-normalized)
    # Server-side compute time, so the client can separate GPU cost from network
    # RTT. Additive + optional: existing clients that read only `actions` are
    # unaffected.
    compute_ms: Optional[float] = None
    # None when RTC wasn't requested; otherwise why it did or didn't apply.
    rtc: Optional[dict] = None


def _decode(b64: str) -> np.ndarray:
    if b64.lstrip().startswith("data:") and "," in b64[:64]:
        b64 = b64.split(",", 1)[1]
    img = Image.open(io.BytesIO(base64.b64decode(b64))).convert("RGB")
    return np.asarray(img)


def _status() -> dict:
    status = "ready" if _model is not None else ("error" if _load_error else "loading")
    return {"ok": _model is not None, "status": status, "error": _load_error,
            "repo": REPO_ID, "source": _model_source, "dtype": str(DTYPE)}


@app.get("/")
def root() -> dict:
    # HuggingFace Docker Spaces route external traffic only after their readiness
    # probe gets a 2xx on "/". Without this the app loads fine but the proxy 404s
    # every request (incl. /health) and the Space auto-sleeps unused. Harmless
    # elsewhere (AWS/Modal just get an extra liveness route).
    return _status()


@app.get("/health")
def health() -> dict:
    return _status()


@app.get("/ready")
def ready() -> dict:
    """Readiness with 503-until-loaded semantics β€” set this as the Inference
    Endpoint's `health_route` so the platform routes no traffic (and marks the
    replica initializing) until the model is actually servable. Kept SEPARATE
    from `/` and `/health`, which must stay 200-while-loading: a Docker Space
    routes external traffic only after a 2xx on `/`, so a 503 there would keep
    the Space dark for the whole model load."""
    if _model is None:
        detail = f"model load failed: {_load_error}" if _load_error else "model loading"
        raise HTTPException(status_code=503, detail=detail)
    return _status()


class PointRequest(BaseModel):
    image: str             # base64 JPEG/PNG, one camera view
    query: str = "the red cup"
    max_new_tokens: int = 96


class PointResponse(BaseModel):
    raw: str                       # the VLM's verbatim generation
    points: list[list[float]]      # parsed [[x, y], ...] in PERCENT of image size
    compute_ms: Optional[float] = None


@app.post("/point", response_model=PointResponse)
def point(req: PointRequest, authorization: Optional[str] = Header(None),
          x_nori_token: Optional[str] = Header(None)) -> PointResponse:
    """Perception probe (diagnostic, not on the control path): ask the Molmo2-ER
    backbone β€” a pixel-accurate pointing model β€” to point at `query` in ONE
    frame. Separates "does the model SEE the target in our camera domain" from
    "does it act correctly": wrong/absent points on live robot frames = visual
    domain gap (no calibration work can fix it); correct points + wrong motion
    = the failure is downstream of perception."""
    _require_auth(x_nori_token, authorization)
    if _model is None:
        detail = f"model load failed: {_load_error}" if _load_error else "model not loaded yet"
        raise HTTPException(status_code=503, detail=detail)
    img = Image.fromarray(_decode(req.image))
    prompt = f"Point to {req.query}."
    t0 = time.time()
    try:
        with _lock, torch.inference_mode():
            # Preferred: the processor's chat template (Molmo2 family). Fallback:
            # the classic Molmo processor.process() API. Both produce tensors the
            # underlying ImageTextToText model can generate from.
            try:
                inputs = _processor.apply_chat_template(
                    [{"role": "user",
                      "content": [{"type": "image", "image": img},
                                  {"type": "text", "text": prompt}]}],
                    add_generation_prompt=True, tokenize=True,
                    return_dict=True, return_tensors="pt")
            except Exception:
                inputs = _processor.process(images=[img], text=prompt)
                inputs = {k: (v.unsqueeze(0) if hasattr(v, "dim") and v.dim() in (1, 3) else v)
                          for k, v in inputs.items()}
            inputs = {k: (v.to(_model.device) if hasattr(v, "to") else v)
                      for k, v in inputs.items()}
            out = _model.generate(**inputs, max_new_tokens=int(req.max_new_tokens))
            n_in = inputs["input_ids"].shape[1] if "input_ids" in inputs else 0
            text = _processor.tokenizer.decode(out[0][n_in:], skip_special_tokens=False)
    except Exception as exc:
        raise HTTPException(status_code=500, detail=f"pointing failed: {type(exc).__name__}: {exc}")
    # Parse Molmo point markup: <point x="53.1" y="42.2" ...> (single) and the
    # <points x1=".." y1=".." x2=".." ...> multi-point form. Percent coordinates.
    import re
    pts = [[float(x), float(y)] for x, y in
           re.findall(r'x\d*="([0-9.]+)"\s+y\d*="([0-9.]+)"', text)]
    return PointResponse(raw=text, points=pts,
                         compute_ms=round((time.time() - t0) * 1000.0, 1))


@app.post("/act", response_model=ActResponse)
def act(req: ActRequest, authorization: Optional[str] = Header(None),
        x_nori_token: Optional[str] = Header(None)) -> ActResponse:
    _require_auth(x_nori_token, authorization)
    if _model is None:
        detail = f"model load failed: {_load_error}" if _load_error else "model not loaded yet"
        raise HTTPException(status_code=503, detail=detail)
    if not 1 <= len(req.images) <= MAX_IMAGES:
        raise HTTPException(status_code=422, detail=f"need 1..{MAX_IMAGES} camera images")
    num_steps = max(1, min(int(req.num_steps), MAX_NUM_STEPS))  # clamp GPU cost
    images = [_decode(b) for b in req.images]
    state = np.asarray(req.state, dtype=np.float32)
    # Autograd roughly doubles activation memory, and the weights alone are ~21GB
    # on a 24GB A10G β€” so the grad path can simply OOM. Report that as a clean 507
    # rather than letting it wedge the container: "RTC does not fit on this
    # hardware" is a legitimate measurement outcome, not a crash.
    # RTC decides the execution knobs: its per-step PiGDM correction is a VJP, so
    # autograd must be ON and the CUDA-graph fast path OFF (you cannot backprop a
    # captured graph). Measured cost of that swap: 303ms -> 827ms compute.
    rtc_req = req.rtc
    rtc_horizon = None
    rtc_note = None
    if rtc_req is not None:
        rtc_horizon = rtcmod.pick_execution_horizon(rtc_req.delay, ACTION_HORIZON)
        if rtc_horizon is None:
            # d > H/2: frozen prefix and free tail would overlap. Serve a normal
            # chunk rather than a silently ill-defined one.
            rtc_note = (f"skipped: delay {rtc_req.delay} too large for horizon "
                        f"{ACTION_HORIZON} (needs d <= H/2)")
    use_rtc = rtc_req is not None and rtc_horizon is not None
    cuda_graph = False if use_rtc else bool(req.enable_cuda_graph)
    grad_ctx = torch.enable_grad() if (use_rtc or req.enable_grad) else torch.no_grad()
    t0 = time.perf_counter()
    try:
        with _lock, grad_ctx:
            if use_rtc:
                # Swap this session's cached chunk in under the lock (see _rtc_prev).
                _rtc_state.prev = _rtc_prev.get(rtc_req.session)
                _rtc_state.enabled = _rtc_state.prev is not None
                _rtc_state.consumed = max(0, int(rtc_req.consumed))
                _rtc_state.delay = max(0, int(rtc_req.delay))
                _rtc_state.execution_horizon = rtc_horizon
                _rtc_state.applied = 0
            out = _model.predict_action(
                processor=_processor,
                images=images,
                task=req.instruction,
                state=state,
                norm_tag=NORM_TAG,
                inference_action_mode="continuous",
                num_steps=num_steps,
                normalize_language=True,
                enable_cuda_graph=cuda_graph,
            )
            if use_rtc:
                # the patched flow loop wrote the new chunk into state.prev
                if len(_rtc_prev) >= _RTC_SESSION_CAP and rtc_req.session not in _rtc_prev:
                    _rtc_prev.pop(next(iter(_rtc_prev)))
                _rtc_prev[rtc_req.session] = _rtc_state.prev
                rtc_note = {"guided_steps": _rtc_state.applied,
                            "execution_horizon": rtc_horizon,
                            "delay": _rtc_state.delay,
                            "consumed": _rtc_state.consumed,
                            "had_target": bool(_rtc_state.enabled)}
                _rtc_state.prev = None      # don't leak one session's chunk to the next
                _rtc_state.enabled = False
    except torch.cuda.OutOfMemoryError as e:
        torch.cuda.empty_cache()
        raise HTTPException(
            status_code=507,
            detail=f"CUDA OOM (enable_grad={req.enable_grad}, "
                   f"cuda_graph={req.enable_cuda_graph}): {e}",
        ) from e
    if torch.cuda.is_available():
        torch.cuda.synchronize()   # predict_action is async; time the real compute
    compute_ms = (time.perf_counter() - t0) * 1000.0
    acts = out.actions
    if torch.is_tensor(acts):  # predict_action returns a CUDA tensor β€” move to host first
        acts = acts.detach().float().cpu().numpy()
    acts = np.asarray(acts, dtype=np.float32)
    if acts.ndim == 3 and acts.shape[0] == 1:  # (1, chunk, DOF) -> (chunk, DOF)
        acts = acts[0]
    return ActResponse(actions=acts.tolist(), compute_ms=round(compute_ms, 1),
                       rtc=(rtc_note if isinstance(rtc_note, dict)
                            else ({"skipped": rtc_note} if rtc_note else None)))