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