From 649ce5dd3ddb8ecc8196062c1292d74528ead8f2 Mon Sep 17 00:00:00 2001 From: Mick Date: Sat, 11 Jul 2026 07:50:58 +0800 Subject: [PATCH] model: support Pi0.5 (#30633) --- docs_new/cookbook/intro.mdx | 7 +- docs_new/cookbook/vla/OpenPI/Pi0.5.mdx | 474 +++++ docs_new/cookbook/vla/intro.mdx | 21 + docs_new/docs.json | 13 + python/pyproject.toml | 1 + python/sglang/cli/serve.py | 15 +- .../benchmarks/bench_pi05_openpi.py | 1243 +++++++++++++ .../configs/pipeline_configs/__init__.py | 2 + .../configs/pipeline_configs/base.py | 103 +- .../configs/pipeline_configs/pi05.py | 99 ++ .../multimodal_gen/configs/sample/__init__.py | 4 + .../multimodal_gen/configs/sample/pi05.py | 76 + .../configs/sample/sampling_params.py | 75 +- .../multimodal_gen/configs/sample/vla.py | 272 +++ python/sglang/multimodal_gen/registry.py | 16 + .../runtime/distributed/group_coordinator.py | 23 +- .../entrypoints/diffusion_generator.py | 34 +- .../runtime/entrypoints/http_server.py | 10 +- .../runtime/entrypoints/utils.py | 27 +- .../runtime/entrypoints/vla/__init__.py | 1 + .../runtime/entrypoints/vla/api.py | 89 + .../runtime/entrypoints/vla/openpi.py | 29 + .../runtime/entrypoints/vla/protocol.py | 443 +++++ .../runtime/entrypoints/vla/ws_utils.py | 57 + .../multimodal_gen/runtime/launch_server.py | 6 + .../runtime/layers/attention/layer.py | 3 +- .../multimodal_gen/runtime/loader/utils.py | 8 +- .../runtime/loader/weight_utils.py | 30 +- .../runtime/managers/gpu_worker.py | 10 +- .../runtime/managers/scheduler.py | 65 +- .../runtime/models/vlas/__init__.py | 5 + .../runtime/models/vlas/pi05_core.py | 1534 +++++++++++++++++ .../runtime/models/vlas/pi05_policy.py | 1107 ++++++++++++ .../multimodal_gen/runtime/pipelines/pi05.py | 126 ++ .../pipelines_core/composed_pipeline_base.py | 14 +- .../executors/pipeline_executor.py | 4 + .../runtime/pipelines_core/schedule_batch.py | 34 +- .../runtime/pipelines_core/stages/__init__.py | 10 + .../model_specific_stages/pi05_preprocess.py | 212 +++ .../runtime/pipelines_core/stages/vla.py | 496 ++++++ .../runtime/server_args/auto_tune.py | 2 + .../runtime/server_args/server_args.py | 35 +- .../multimodal_gen/runtime/server_warmup.py | 22 +- .../multimodal_gen/runtime/vla/__init__.py | 3 + .../runtime/vla/denoise_cuda_graph.py | 213 +++ .../multimodal_gen/runtime/vla/observation.py | 76 + .../multimodal_gen/runtime/vla/parallel.py | 161 ++ .../runtime/vla/prefix_cache.py | 171 ++ .../runtime/warmup_request_builder.py | 26 +- .../multimodal_gen/test/server/gpu_cases.py | 11 + .../test/server/test_server_common.py | 98 ++ .../test/server/test_server_utils.py | 102 ++ .../test/server/testcase_configs.py | 25 +- .../test/single_test_file/test_pi05_e2e.py | 212 +++ .../sglang/multimodal_gen/test/test_utils.py | 147 ++ .../test/unit/test_cfg_parallel_warmup.py | 88 +- .../test/unit/test_consistency_metrics.py | 31 + .../test/unit/test_pi05_action_api.py | 241 +++ .../test/unit/test_pi05_prefix_cache.py | 61 + .../test/unit/test_pi05_runtime_helpers.py | 174 ++ python/sglang/utils.py | 4 + 61 files changed, 8593 insertions(+), 108 deletions(-) create mode 100644 docs_new/cookbook/vla/OpenPI/Pi0.5.mdx create mode 100644 docs_new/cookbook/vla/intro.mdx create mode 100644 python/sglang/multimodal_gen/benchmarks/bench_pi05_openpi.py create mode 100644 python/sglang/multimodal_gen/configs/pipeline_configs/pi05.py create mode 100644 python/sglang/multimodal_gen/configs/sample/pi05.py create mode 100644 python/sglang/multimodal_gen/configs/sample/vla.py create mode 100644 python/sglang/multimodal_gen/runtime/entrypoints/vla/__init__.py create mode 100644 python/sglang/multimodal_gen/runtime/entrypoints/vla/api.py create mode 100644 python/sglang/multimodal_gen/runtime/entrypoints/vla/openpi.py create mode 100644 python/sglang/multimodal_gen/runtime/entrypoints/vla/protocol.py create mode 100644 python/sglang/multimodal_gen/runtime/entrypoints/vla/ws_utils.py create mode 100644 python/sglang/multimodal_gen/runtime/models/vlas/__init__.py create mode 100644 python/sglang/multimodal_gen/runtime/models/vlas/pi05_core.py create mode 100644 python/sglang/multimodal_gen/runtime/models/vlas/pi05_policy.py create mode 100644 python/sglang/multimodal_gen/runtime/pipelines/pi05.py create mode 100644 python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/pi05_preprocess.py create mode 100644 python/sglang/multimodal_gen/runtime/pipelines_core/stages/vla.py create mode 100644 python/sglang/multimodal_gen/runtime/vla/__init__.py create mode 100644 python/sglang/multimodal_gen/runtime/vla/denoise_cuda_graph.py create mode 100644 python/sglang/multimodal_gen/runtime/vla/observation.py create mode 100644 python/sglang/multimodal_gen/runtime/vla/parallel.py create mode 100644 python/sglang/multimodal_gen/runtime/vla/prefix_cache.py create mode 100644 python/sglang/multimodal_gen/test/single_test_file/test_pi05_e2e.py create mode 100644 python/sglang/multimodal_gen/test/unit/test_pi05_action_api.py create mode 100644 python/sglang/multimodal_gen/test/unit/test_pi05_prefix_cache.py create mode 100644 python/sglang/multimodal_gen/test/unit/test_pi05_runtime_helpers.py diff --git a/docs_new/cookbook/intro.mdx b/docs_new/cookbook/intro.mdx index 3cd1747ee..6914ade5f 100644 --- a/docs_new/cookbook/intro.mdx +++ b/docs_new/cookbook/intro.mdx @@ -9,7 +9,7 @@ A community-maintained repository of practical guides and recipes for deploying ## Guides - + + ## Benchmarks diff --git a/docs_new/cookbook/vla/OpenPI/Pi0.5.mdx b/docs_new/cookbook/vla/OpenPI/Pi0.5.mdx new file mode 100644 index 000000000..3b837f828 --- /dev/null +++ b/docs_new/cookbook/vla/OpenPI/Pi0.5.mdx @@ -0,0 +1,474 @@ +--- +title: Pi0.5 +metatags: + description: "Deploy OpenPI / LeRobot Pi0.5 diffusion Vision-Language-Action policies with SGLang's native multimodal_gen runtime." +tag: dVLA +--- + +
+ dVLA + OpenPI / LeRobot + flow matching action + robot edge +
+ +## 1. Model Introduction + +Pi0.5 is an OpenPI / LeRobot diffusion Vision-Language-Action (dVLA) policy. It consumes camera images, a language instruction, and robot state, then returns a continuous action chunk for robot control. + +SGLang serves Pi0.5 through the native `multimodal_gen` runtime. The implementation uses a SigLIP/PaliGemma prefix encoder and a Gemma action expert: the prefix is encoded once, then the action expert runs the flow-matching denoising loop. This is not a token decode workload, so the Pi0.5 path does not use the LLM sampler, logits processor, token streaming, paged decode KV cache, or a separate SRT serving engine. + +The prefix encoder covers both stages of observation encoding: SigLIP turns resized camera pixels into continuous patch embeddings, then the PaliGemma transformer jointly encodes those patches with tokenized task/state inputs and produces per-layer prefix K/V. At flow timestep `t`, the action expert projects the noisy continuous action chunk `x_t` into action embeddings. Its queries attend to both the fixed prefix K/V and the current action K/V, while the timestep follows a separate sinusoidal-MLP path and conditions every action-expert layer through AdaRMSNorm gates. + +Supported public checkpoints: + +| Checkpoint | Cameras | State Dim | Output Action Dim | Action Horizon | Denoise Steps | +| --- | --- | ---: | ---: | ---: | ---: | +| `lerobot/pi05_base` | `base_0_rgb`, `left_wrist_0_rgb`, `right_wrist_0_rgb` | 32 | 32 | 50 | 10 | +| `lerobot/pi05_libero_base` | `image`, `image2`, one empty camera | 8 | 7 | 50 | 10 | + +References: + +- [OpenPI](https://github.com/Physical-Intelligence/openpi) +- [lerobot/pi05_base](https://huggingface.co/lerobot/pi05_base) +- [LeRobot Pi0.5 docs](https://huggingface.co/docs/lerobot/en/pi05) + +## 2. Installation + +Install SGLang with the diffusion extra. Pi0.5 lives in `multimodal_gen`, and the extra includes the runtime dependencies used by the policy server. + +```bash Command +git clone https://github.com/sgl-project/sglang.git +cd sglang +pip install -e "python[diffusion]" +``` + +For general environment setup, see the [SGLang Diffusion installation guide](/docs/sglang-diffusion/installation). + +## 3. Model Deployment + +Serve the base Pi0.5 policy: + +```bash Command +sglang serve lerobot/pi05_base \ + --model-type diffusion \ + --host 127.0.0.1 \ + --port 30000 +``` + +Serve the LIBERO checkpoint: + +```bash Command +sglang serve lerobot/pi05_libero_base \ + --model-type diffusion \ + --host 127.0.0.1 \ + --port 30000 +``` + +These registered LeRobot checkpoints dispatch to the native `multimodal_gen` Pi0.5 pipeline automatically. The `--pipeline` / `--pipeline-class-name` flag is only an advanced override for local or private checkpoints that cannot be resolved from the model registry. Pi0.5 serving does not start a second SRT LLM serving engine. + +### 3.1 Action Request Schema + +| Field | Type | Description | +| --- | --- | --- | +| `model` | string, optional | Served model name. When omitted, the server uses the currently loaded policy. | +| `input.task` or `input.prompt` | string | Language instruction for the policy. | +| `input.observation.images` | object | Map from camera name to RGB image. Images can be nested here or sent as OpenPI observation keys such as `observation.images.base_0_rgb` through the websocket adapter. | +| `input.observation.state` | array or tensor object | Robot state vector. Use the same normalization convention as the OpenPI / LeRobot checkpoint. | +| `input.observation.noise` | array or tensor object, optional | Initial action noise with shape `[action_horizon, action_dim]`. Use this for deterministic debugging. | +| `parameters.num_inference_steps` | integer, optional | Flow-matching denoise steps. Defaults to `10`. | +| `parameters.action_horizon` | integer, optional | Output action horizon. Defaults to the checkpoint config. | +| `parameters.action_dim` | integer, optional | Internal padded action dimension. Defaults to the checkpoint config. | +| `runtime.return_timing` | boolean, optional | Return stage timing fields. Defaults to `true`. | +| `runtime.prefix_cache` | boolean or `"auto"`, optional | Enable exact full-prefix lookup for this request when the server has `enable_global_prefix_cache=true`. Defaults to `"auto"`. | +| `runtime.cuda_graph` | boolean or `"auto"`, optional | Enable the action denoise CUDA graph path for this request when a matching shape bucket is available. Defaults to `"auto"`. | +| `runtime.output_format` | `"list"` or `"numpy"`, optional | Use `"list"` for JSON compatibility. Use `"numpy"` with msgpack or Python clients to avoid Python-list materialization. Defaults to `"list"`. | +| `runtime.response_format` | `"envelope"` or `"raw"`, optional | HTTP-only response shape. `"envelope"` returns the generic action envelope. `"raw"` returns the policy payload directly. Defaults to `"envelope"`. | + +## 4. API Usage + +### 4.1 Generic Action HTTP API + +Use `/v1/actions/generations` for direct policy calls and debugging. Use `/v1/actions/metadata` to discover the camera keys, state size, action shape, defaults, and websocket capabilities of the currently served policy. + +```python Example +import numpy as np +import requests + +image = np.zeros((224, 224, 3), dtype=np.uint8) +payload = { + "model": "lerobot/pi05_base", + "input": { + "task": "pick up the block", + "observation": { + "images": { + "base_0_rgb": image.tolist(), + "left_wrist_0_rgb": image.tolist(), + "right_wrist_0_rgb": image.tolist(), + }, + "state": np.zeros(32, dtype=np.float32).tolist(), + }, + }, + "runtime": { + "return_timing": True, + "prefix_cache": "auto", + "cuda_graph": "auto", + }, +} + +response = requests.post( + "http://127.0.0.1:30000/v1/actions/generations", + json=payload, + timeout=60, +) +response.raise_for_status() +data = response.json() + +actions = data["data"][0]["action"]["values"] +print(len(actions), len(actions[0])) +print(data.get("timings")) +``` + +The same `/v1/actions/generations` endpoint also accepts `Content-Type: application/msgpack` and can return msgpack when `Accept: application/msgpack` is set. For msgpack clients, send numpy arrays directly using the `pack_numpy_payload` helper from the websocket example below. Msgpack requests default to numpy action output on the server side; set `runtime.output_format` to `"list"` only when a client explicitly needs nested Python lists. + +For the lowest-overhead generic HTTP path, use msgpack with the raw policy response: + +```python Example +payload["runtime"] = { + "return_timing": True, + "prefix_cache": False, + "cuda_graph": "auto", + "response_format": "raw", +} + +response = requests.post( + "http://127.0.0.1:30000/v1/actions/generations", + data=packb(payload), + headers={ + "Content-Type": "application/msgpack", + "Accept": "application/msgpack", + }, + timeout=60, +) +result = unpackb(response.content) +actions = result["actions"] +``` + +For `lerobot/pi05_libero_base`, use the LIBERO camera names and an 8-dimensional state vector: + +```python Example +payload = { + "input": { + "task": "pick up the object", + "observation": { + "images": { + "image": image.tolist(), + "image2": image.tolist(), + }, + "state": np.zeros(8, dtype=np.float32).tolist(), + }, + }, +} +``` + +### 4.2 Generic Realtime WebSocket + +`/v1/actions/realtime` is the generic msgpack websocket path. It sends `action.metadata` on connect and returns the same `action.generation` envelope as the HTTP API for each request. + +### 4.3 OpenPI-Compatible WebSocket + +Robot clients can use the OpenPI-compatible msgpack websocket endpoint at `/openpi/policy`. The server sends metadata immediately after connection, then each client message should contain one observation. + +```python Example +import asyncio + +import msgpack +import numpy as np +import websockets + + +def pack_array(obj): + if isinstance(obj, np.ndarray): + return { + b"__ndarray__": True, + b"data": obj.tobytes(), + b"dtype": obj.dtype.str, + b"shape": obj.shape, + } + if isinstance(obj, np.generic): + return { + b"__npgeneric__": True, + b"data": obj.item(), + b"dtype": obj.dtype.str, + } + return obj + + +def unpack_array(obj): + ndarray_marker = obj.get("__ndarray__") or obj.get(b"__ndarray__") + npgeneric_marker = obj.get("__npgeneric__") or obj.get(b"__npgeneric__") + data = obj.get("data", obj.get(b"data")) + dtype = obj.get("dtype", obj.get(b"dtype")) + shape = obj.get("shape", obj.get(b"shape")) + if ndarray_marker: + return np.ndarray( + buffer=data, + dtype=np.dtype(dtype), + shape=shape, + ) + if npgeneric_marker: + return np.dtype(dtype).type(data) + return obj + + +def packb(payload): + return msgpack.packb(payload, default=pack_array, use_bin_type=True) + + +def unpackb(payload): + return msgpack.unpackb(payload, object_hook=unpack_array, raw=False) + + +async def main(): + image = np.zeros((224, 224, 3), dtype=np.uint8) + observation = { + "task": "pick up the block", + "observation.images.base_0_rgb": image, + "observation.images.left_wrist_0_rgb": image, + "observation.images.right_wrist_0_rgb": image, + "observation.state": np.zeros(32, dtype=np.float32), + } + + async with websockets.connect( + "ws://127.0.0.1:30000/openpi/policy", + max_size=None, + ) as websocket: + metadata = unpackb(await websocket.recv()) + print(metadata) + + await websocket.send(packb(observation)) + result = unpackb(await websocket.recv()) + actions = result["actions"] + print(len(actions), len(actions[0])) + print(result.get("server_timing")) + + +asyncio.run(main()) +``` + +## 5. Configuration Tips + +- Request-local `PrefixContext` is always reused across all denoise steps in one request. The prefix K/V is not cloned per step. +- The optional global prefix cache is a bounded exact-match LRU. It is disabled by default because changing robot frames rarely hit it and enabling it prevents unrelated misses from entering grouped prefix execution. Set `enable_global_prefix_cache=true` for repeated observations, retries, or multiple policy calls over the same camera/state sample; `runtime.prefix_cache` can then disable lookup per request. +- Partial-prefix reuse is not supported because Pi0.5 combines image and tokenized task/state inputs under full attention. Changing any input can change every deeper-layer prefix K/V tensor. The exact key hashes resized and normalized pixels before SigLIP, plus effective token IDs, token masks, camera masks, model revision, dtype, and parallel layout. Tensor content hashing reuses SRT's CPU/CUDA implementation; hashing the pre-SigLIP input lets an exact hit skip both the vision encoder and prefix transformer. +- CUDA graph capture targets one action-denoise step by shape bucket, then replays it across the flow-matching loop. Shape buckets include batch size, prefix length, action horizon, action dim, dtype, and parallel layout. With action SP enabled, the graph bucket uses the local action shard length and rank-specific position offset. +- Cache-DiT is not used in the default Pi0.5 path. The current robot policy target is numerically lossless inference, while Cache-DiT-style reuse is an image/video DiT approximation that needs separate policy-quality validation before it can be recommended for action control. +- Do not use CFG parallelism to split the 10 Euler steps. Use it only for independent branches such as multiple candidate actions or future conditional/unconditional branches. +- Prefix TP uses native SGLang parallel linear layers for the PaliGemma language prefix model when model parallel TP is initialized and the VLA split broadcast group is not active. The action expert does not share that TP layout. The v1 split prefix/action path instead uses the SP group: prefix root computes/broadcasts `PrefixContext`, while action ranks run the SP action path. +- The split path uses the SP group as the action group: prefix root computes/broadcasts `PrefixContext`, all action ranks broadcast the initial action noise once, shard the action horizon, and run the action expert through Ulysses attention when the prefix is full-attention, ring degree is one, heads are divisible by SP size, and the horizon is evenly shardable. Otherwise it falls back to action-root execution. +- `lerobot/pi05_libero_base` returns 7 action dimensions even though the internal padded action tensor uses 32 dimensions. + +## 6. VRAM Tuning + +For current public Pi0.5 checkpoints, a stable 16GB discrete GPU target is a reasonable v1 deployment bar for robot workstations. In practice, leave headroom for the driver, camera middleware, robot process, and allocator fragmentation. Jetson/Orin unified-memory devices need extra caution because system RAM and GPU memory share the same pool. + +OpenPI inference is mixed precision, not full fp32: most weights and compute run in `bf16`, selected stability-sensitive weights stay in `fp32`, and returned actions are `float32`. SGLang mirrors that policy by default. The validated `pi05_aloha` OpenPI PyTorch checkpoint keeps `119,720,608` parameters in fp32; SGLang reports the same fp32 stability set, plus `3,233,713,264` bf16 runtime parameters after skipping unused LM heads for continuous action inference. The fp32 set includes SigLIP patch/position embeddings, Gemma layer norms/final norms, and the action/time projection heads. Keep `materialize_dtype` at `bf16` unless you are debugging numerical parity. + +Current H100 pressure validation shows that the bf16 model path fits a 16GB-free budget without layerwise offload. The Pi0.5 path batches all camera frames for a grouped request into one SigLIP forward before splitting the embeddings back by camera, which is important for multi-camera robot workloads. The latency numbers below use the Python grouped API with global prefix cache disabled and CUDA graph enabled on unconstrained H100; rerun HTTP/OpenPI websocket on your target server before using the policy in a closed-loop robot. + +| Mode | Command Shape | Steady VRAM Snapshot | Notes | +| --- | --- | --- | --- | +| Single GPU bf16 | `--num-gpus 1`, no offload | fit with `16383 MiB` free before load in H100 pressure test | Recommended first path for 16GB-class discrete GPUs. | +| ALOHA bf16 grouped | Python API, batch size 1 / 2 / 4 / 8 | unconstrained H100 | `52.4 ms` single; `65.1 ms / 2`; `91.9 ms / 4`; `147.4 ms / 8`. | +| ALOHA 16GB-free pressure | Python API, no offload | same 16GB-free pressure | Fit was validated after the bf16 correction; rerun latency on the target because older pressure latency was collected before the final OpenPI precision fix. | +| Global prefix cache disabled | `enable_global_prefix_cache=false` or request `runtime.prefix_cache=false` | avoids cache growth across changing frames | Recommended for tight edge budgets unless repeated exact frames are common. | +| Offload fallback | config in 6.2 | target-dependent | Use only if the default bf16 path OOMs on the target device. | + +Use these knobs first because they do not change action numerics for a fixed input/noise: + +### 6.1 16GB Edge Config + +Use this single-GPU config first for 16GB-class robot workstations. It keeps parameters resident on GPU, keeps CUDA graph enabled, and disables global prefix cache growth. + +```json File +{ + "materialize_dtype": "bf16", + "enable_global_prefix_cache": false, + "prefix_cache_max_entries": 0, + "enable_action_cuda_graph": true +} +``` + +Start a single-GPU server with the override: + +```bash Command +sglang serve lerobot/pi05_base \ + --model-type diffusion \ + --pipeline-config-path pi05_edge_16gb.json \ + --num-gpus 1 \ + --warmup-mode off \ + --host 127.0.0.1 \ + --port 30000 +``` + +Treat the H100 pressure result as a memory-budget validation, not a Jetson latency guarantee. Jetson/Orin devices use shared system memory and have much lower effective memory bandwidth than H100, so measure closed-loop control latency on the target device before deciding the action chunk cadence. + +### 6.2 Offload Fallback + +Use offload only when the default bf16 config still does not fit on the target. These modes are numerically lossless for fixed inputs/noise, but they move weights between CPU and GPU and can hurt latency substantially. + +For a moderate fallback, offload cache growth and selected stage-resident modules first: + +```json File +{ + "materialize_dtype": "bf16", + "enable_global_prefix_cache": false, + "prefix_cache_max_entries": 0, + "enable_action_cuda_graph": false, + "offload_prefix_image_encoder_after_embed": true, + "offload_prefix_token_embedding": true, + "offload_prefix_language_layer_count_after_prefix": 2, + "offload_action_expert_after_denoise": true, + "empty_cache_after_prefix": true +} +``` + +If that still does not fit, full prefix layerwise CPU offload keeps every PaliGemma language layer on CPU and moves one layer at a time to GPU during prefix compute: + +```json File +{ + "materialize_dtype": "bf16", + "enable_global_prefix_cache": false, + "prefix_cache_max_entries": 0, + "enable_action_cuda_graph": false, + "offload_prefix_image_encoder": true, + "offload_prefix_token_embedding": true, + "offload_prefix_language_layers": true, + "offload_prefix_language_layers_empty_cache": true, + "empty_cache_after_prefix": true +} +``` + +Offload validation should be repeated on the target hardware after any dtype or loader change. Earlier fp32-runtime offload numbers are not comparable to the current bf16 path. + +### 6.3 Per-Request Controls + +For HTTP calls, disable cache or CUDA graph without restarting the server: + +```python Example +payload = { + "input": { + "task": "pick up the block", + "observation": { + "images": images, + "state": state, + }, + }, + "runtime": { + "prefix_cache": False, + "cuda_graph": False, + }, +} +``` + +For OpenPI websocket clients, include equivalent `enable_prefix_cache` and `enable_cuda_graph` fields in each raw msgpack observation if you need per-request compatibility controls. + +### 6.4 Deployment Choices + +- Keep batch size to one control stream unless grouped robot streams are explicitly validated. More concurrent observations increase activation and PrefixContext residency. +- Keep `materialize_dtype` at the default `bf16`. `fp32` is useful only for debugging numerical issues and will increase memory and latency. +- Use split prefix/action only when you have multiple GPUs and have validated the robot control latency on that topology. On two GPUs with `--sp-degree 2 --ulysses-degree 2`, the action horizon can be sequence-sharded while the prefix root still computes and broadcasts `PrefixContext`. A single 16GB-class GPU should try the default bf16 config first. +- Reducing `num_inference_steps` lowers latency but is not a pure memory fix and can change policy behavior. Validate closed-loop task success before using fewer than the checkpoint default. +- CPU offload is a compatibility fallback, not the preferred 16GB path. Quantization and deeper prefix/vision sharding are the next steps for sub-16GB devices. + +### 6.5 Loader And Run:ai Model Streamer + +Run:ai Model Streamer can improve cold-start and checkpoint loading by reading safetensors concurrently and streaming tensors toward GPU memory. It is useful for local SSD, object storage, and cloud deployments where startup time is dominated by model file IO. + +The `python[diffusion]` extra includes `runai_model_streamer`. If the package is installed, SGLang enables it by default through `SGLANG_USE_RUNAI_MODEL_STREAMER=true`. Set it to `false` to force the plain safetensors loader: + +```bash Command +SGLANG_USE_RUNAI_MODEL_STREAMER=false \ +sglang serve lerobot/pi05_base \ + --model-type diffusion \ + --host 127.0.0.1 \ + --port 30000 +``` + +Pi0.5 uses the direct SSD-to-GPU Run:ai path when all load targets for the current process are GPU-resident. This check is rank-local: distributed action ranks can stream their action-expert subset directly to GPU, while any rank with CPU/offloaded target tensors stays on the CPU safetensors fallback. The single-GPU validation streamed `13.5 GiB` of safetensors directly to `cuda:0` in about `1.5 s`, then returned `[50, 32]` actions over both HTTP and OpenPI websocket. + +For mixed CPU offload and low-VRAM 16GB-class modes, Pi0.5 still uses the header-filtered safe loader for ranks whose target tensors are not fully CUDA-resident. Direct GPU streaming can increase the exact VRAM pressure the low-memory path is trying to avoid. When debugging distributed startup, compare with `SGLANG_USE_RUNAI_MODEL_STREAMER=false` to separate streamer behavior from model execution behavior. + +Run:ai Model Streamer does not reduce steady inference VRAM after parameters, caches, activations, CUDA contexts, and graph buffers are resident. For Pi0.5 low-VRAM work, prioritize component placement, prefix cache size, CUDA graph residency, and CPU/offload first. Model Streamer is a cold-start optimization after the steady-memory budget is correct. + +## 7. OpenPI Comparison Benchmark + +Use `bench_pi05_openpi.py` when you need a side-by-side latency and action-difference report against the OpenPI policy implementation. Start the SGLang Pi0.5 server first for HTTP or websocket modes, then run the benchmark from the SGLang repository root: + +```bash Command +python python/sglang/multimodal_gen/benchmarks/bench_pi05_openpi.py \ + --profile aloha \ + --sglang-url http://127.0.0.1:30000 \ + --openpi-checkpoint gs://openpi-assets/checkpoints/pi05_base \ + --num-inference-steps 10 \ + --batch-size 4 \ + --repeats 20 \ + --warmup 3 \ + --deterministic-noise \ + --output pi05_openpi_compare.json +``` + +The benchmark reports: + +- SGLang HTTP latency for `/v1/actions/generations`, including stage timings when the server returns them. +- SGLang msgpack HTTP latency when `--sglang-api http_msgpack` is set. This keeps the generic HTTP endpoint but avoids JSON image-array overhead. Use `--sglang-http-response-format raw` to benchmark the compact raw policy response. +- SGLang OpenPI-compatible websocket latency when `--sglang-api openpi_ws` is set. This uses persistent msgpack websocket connections and is closer to the robot client path than JSON-over-HTTP. +- SGLang Python in-process latency when `--sglang-api python` is set. This loads the native Pi0.5 pipeline in the benchmark process and avoids HTTP, websocket, scheduler, and serialization overhead. Use `--sglang-python-batch-mode grouped` to exercise the conservative native grouped-batch path. The Python path also reports actual SGLang module parameter dtype counts and example parameter names. +- OpenPI single-request latency through `Policy.infer`. +- Batch latency for grouped robot streams. SGLang uses concurrent HTTP requests in HTTP mode and persistent multi-connection msgpack calls in websocket mode. The Python backend can use true grouped model execution for fresh-prefix requests. OpenPI defaults to its internal direct model batch path because the public `Policy.infer` API is single-observation. +- Action difference on the common output prefix. Use `--deterministic-noise` for strict debugging when the SGLang and OpenPI horizons match. The `aloha` profile supports this directly; the LIBERO profile compares the common prefix because OpenPI's released LIBERO config uses a shorter output horizon than the LeRobot Pi0.5 checkpoint metadata. + +For one-sided 16GB-class checks, run each backend separately under the same VRAM pressure. The SGLang Python path accepts the same pipeline config override as serving: + +```bash Command +python python/sglang/multimodal_gen/benchmarks/bench_pi05_openpi.py \ + --profile aloha \ + --sglang-api python \ + --sglang-python-batch-mode grouped \ + --skip-openpi \ + --sglang-pipeline-config-path pi05_edge_16gb.json \ + --num-inference-steps 10 \ + --batch-size 4 \ + --repeats 5 \ + --warmup 2 \ + --deterministic-noise \ + --disable-prefix-cache \ + --disable-cuda-graph +``` + +Use `--skip-sglang --openpi-pytorch-compile-mode none` to measure an OpenPI eager baseline in the same environment. For a PyTorch-native OpenPI baseline, point `--openpi-checkpoint` at a checkpoint directory containing `model.safetensors` and the OpenPI `assets/` norm-stat tree. The `keep` mode preserves OpenPI's checkpoint default compile setting. The grouped Python path currently requires compatible fresh-prefix requests without split prefix/action workers or effective prefix-cache hits; other cases fall back to per-request execution. + +## 8. Validation Notes + +The following checks were run on H100 GPUs with the native SGLang Pi0.5 path: + +| Check | Result | +| --- | --- | +| `lerobot/pi05_base` direct end-to-end | Prefix length `968`, output shape `[1, 50, 32]`, peak allocated memory `12.817 GiB`. | +| LeRobot reference parity | One-step velocity max absolute difference `1.17e-6`; final 10-step action max absolute difference `1.17e-7`. | +| Action denoise CUDA graph | Eager 10-step denoise `125.4 ms`; steady graph replay `50.8 ms`; max output difference `0`. | +| Exact full-prefix cache | First prefix pass about `203 ms`; exact cache hit prefix stage about `0.2 ms`. | +| `lerobot/pi05_libero_base` direct end-to-end | Image keys `image`, `image2`, `empty_camera_0`; state dim `8`; output action dim `7`; output tensor shape `[1, 50, 32]`. | +| Python grouped execution | ALOHA batch=4 grouped path measured `91.9 ms / 4` on current mixed precision: prefix `18.4 ms`, action denoise `61.2 ms`, preprocess about `2.4 ms` per request. Sequential Python loop batch=4 measured `211.9 ms / 4`. | +| `sglang serve` HTTP | `/v1/actions/generations` returned action shape `[50, 32]`. JSON-over-HTTP remains compatible but image-array serialization dominates. Msgpack HTTP with prefix cache disabled measured `57.2 ms` single and `221.9 ms / 4` for the envelope response, and `56.4 ms` single and `219.7 ms / 4` for `runtime.response_format="raw"`. | +| OpenPI websocket | `/openpi/policy` returned action shape `[50, 32]`; with persistent msgpack connections and prefix cache disabled, ALOHA measured about `77.0 ms` single and `162.0 ms / 4`. | +| 2-GPU prefix/action split | Earlier split validation returned action shape `[50, 32]` and matched single-GPU HTTP with max absolute difference `0`. After the true action-SP change, rerun this check with `--num-gpus 2 --sp-degree 2 --ulysses-degree 2` and verify both action ranks enter denoise kernels. | +| OpenPI/SGLang precision | Official OpenPI JAX inference restores the public GCS checkpoint as bf16 with selected fp32 stability compute and returns float32 actions. The converted OpenPI PyTorch `pi05_aloha` checkpoint keeps `119,720,608` fp32 stability params; SGLang reports the same fp32 set and `3,233,713,264` bf16 runtime params after skipping unused LM heads. | +| Native attention dtype | Checkpoint source tensors may be fp32, but SGLang finalizes PiGemma and SigLIP compute dtype before native attention backend selection; backend logs showed `Using fa attention backend` for the PiGemma path in the prior run. | +| 16GB-free Python pressure | With an H100 artificially constrained to `16381 MiB` free before model load, single-GPU bf16 no-offload Python grouped path completed without OOM. Re-run latency after precision or loader changes before using pressure numbers for deployment sizing. | +| Low-VRAM switches | Disabling prefix cache prevents cache growth across changing robot frames. CUDA graph can stay enabled when the action expert remains resident; disable it only for offload fallback modes. | +| Offload fallback | CPU/offload modes are retained as numerically lossless compatibility fallbacks, but earlier fp32-runtime offload latency numbers are stale after the bf16 dtype correction and should be revalidated before deployment decisions. | +| Run:ai direct loader | Single-GPU serve streamed `13.5 GiB` safetensors to `cuda:0` in about `1.5 s` and returned `[50, 32]` actions. Distributed direct streaming is now rank-local and should be revalidated on the target split topology; offload ranks with CPU targets still use the safe loader. | +| OpenPI comparison status | Official OpenPI GCS `pi05_base` is a JAX checkpoint; converted PyTorch eager was validated without `torch.compile`. On 80GB H100, ALOHA OpenPI PyTorch eager was about `125-130 ms` single and about `164 ms / 4` in the direct-model batch path. Current SGLang Python grouped measured `52.4 ms` single and `91.9 ms / 4`; JAX OpenPI was `53.0 ms` single and `59.5 ms / 2` in a short check. | + +Run a full HTTP or websocket smoke test in your robot deployment environment before using the policy in a closed-loop controller. diff --git a/docs_new/cookbook/vla/intro.mdx b/docs_new/cookbook/vla/intro.mdx new file mode 100644 index 000000000..9c9869442 --- /dev/null +++ b/docs_new/cookbook/vla/intro.mdx @@ -0,0 +1,21 @@ +--- +title: Overview +mode: wide +description: Practical guides for deploying and using Vision-Language-Action policies with SGLang. +metatags: + description: "Explore SGLang Vision-Language-Action policy cookbooks for robot action serving, OpenPI compatibility, caching, and performance tuning." +--- + +Vision-Language-Action policies map camera observations, language instructions, and robot state into continuous action chunks. They share some runtime machinery with diffusion pipelines, but the user-facing workload is robot action inference rather than image or video generation. + +This section keeps VLA policies separate from the diffusion model cookbook so robot deployments can document their own input schemas, policy endpoints, cache behavior, and control-loop performance targets. + +## OpenPI + + + + diff --git a/docs_new/docs.json b/docs_new/docs.json index 15aed06d7..842e79d23 100644 --- a/docs_new/docs.json +++ b/docs_new/docs.json @@ -1265,6 +1265,19 @@ } ] }, + { + "group": "VLA (Vision-Language-Action) Models", + "pages": [ + "cookbook/vla/intro", + { + "group": "OpenPI", + "tag": "NEW", + "pages": [ + "cookbook/vla/OpenPI/Pi0.5" + ] + } + ] + }, { "group": "SpecBundle", "pages": [ diff --git a/python/pyproject.toml b/python/pyproject.toml index 0bcd942ff..7ab0c9df8 100755 --- a/python/pyproject.toml +++ b/python/pyproject.toml @@ -104,6 +104,7 @@ diffusion = [ "imageio==2.36.0", "imageio-ffmpeg==0.5.1", "moviepy>=2.0.0", + "msgpack", "nvidia-modelopt", "opencv-python-headless==4.10.0.84", "PyYAML==6.0.1", diff --git a/python/sglang/cli/serve.py b/python/sglang/cli/serve.py index 0268a1100..00dbd210b 100644 --- a/python/sglang/cli/serve.py +++ b/python/sglang/cli/serve.py @@ -46,13 +46,21 @@ def _extract_model_type_override(extra_argv): return model_type, filtered_argv +def _normalize_positional_model_path(extra_argv): + """Allow `sglang serve ` while preserving existing flag parsing.""" + if extra_argv and not extra_argv[0].startswith("-"): + return ["--model-path", extra_argv[0], *extra_argv[1:]], True + return extra_argv, False + + def serve(args, extra_argv): if any(h in extra_argv for h in ("-h", "--help")): # Since the server type is determined by the model, and we don't have a model path, # we can't show the exact help. Instead, we show a general help message and then # the help for both possible server types. print( - "Usage: sglang serve --model-path [additional-arguments]\n\n" + "Usage: sglang serve [additional-arguments]\n" + " or: sglang serve --model-path [additional-arguments]\n\n" "This command can launch either a standard language model server or a diffusion model server.\n" "The server type is determined by the --model-path.\n" "Optional override: --model-type {auto,llm,diffusion} " @@ -91,6 +99,9 @@ def serve(args, extra_argv): load_plugins() model_type, dispatch_argv = _extract_model_type_override(extra_argv) + dispatch_argv, positional_model_path = _normalize_positional_model_path( + dispatch_argv + ) model_path = get_model_path(dispatch_argv) try: if model_type == "auto": @@ -116,6 +127,8 @@ def serve(args, extra_argv): ) add_multimodal_gen_serve_args(parser) parsed_args, remaining_argv = parser.parse_known_args(dispatch_argv) + if positional_model_path: + parsed_args._sglang_explicit_arg_names = {"model_path"} execute_serve_cmd(parsed_args, remaining_argv) else: diff --git a/python/sglang/multimodal_gen/benchmarks/bench_pi05_openpi.py b/python/sglang/multimodal_gen/benchmarks/bench_pi05_openpi.py new file mode 100644 index 000000000..a200a59f0 --- /dev/null +++ b/python/sglang/multimodal_gen/benchmarks/bench_pi05_openpi.py @@ -0,0 +1,1243 @@ +# SPDX-License-Identifier: Apache-2.0 + +"""Manual Pi0.5 SGLang vs OpenPI benchmark. + +This script is intentionally outside the unit-test path. It needs GPU memory, +Pi0.5 checkpoints, and an OpenPI install for the baseline. +""" + +from __future__ import annotations + +import argparse +import asyncio +import dataclasses +import json +import sys +import time +from concurrent.futures import ThreadPoolExecutor +from dataclasses import dataclass +from pathlib import Path +from typing import Any + +import numpy as np +import requests + +SCRIPT_DIR = Path(__file__).resolve().parent +if sys.path and Path(sys.path[0]).resolve() == SCRIPT_DIR: + sys.path.pop(0) + +from sglang.multimodal_gen.runtime.entrypoints.vla.protocol import ( # noqa: E402 + pack_msgpack, + unpack_msgpack, +) + + +@dataclass(frozen=True) +class Pi05BenchProfile: + name: str + sglang_model: str + openpi_config: str + openpi_checkpoint: str + prompt: str + sglang_action_horizon: int + openpi_action_horizon: int + action_dim: int + output_action_dim: int + + +PROFILES = { + "libero": Pi05BenchProfile( + name="libero", + sglang_model="lerobot/pi05_libero_base", + openpi_config="pi05_libero", + openpi_checkpoint="gs://openpi-assets/checkpoints/pi05_libero", + prompt="pick up the object", + sglang_action_horizon=50, + openpi_action_horizon=10, + action_dim=32, + output_action_dim=7, + ), + "aloha": Pi05BenchProfile( + name="aloha", + sglang_model="lerobot/pi05_base", + openpi_config="pi05_aloha", + openpi_checkpoint="gs://openpi-assets/checkpoints/pi05_base", + prompt="pick up the block", + sglang_action_horizon=50, + openpi_action_horizon=50, + action_dim=32, + output_action_dim=14, + ), +} + + +def _stats_ms(samples: list[float]) -> dict[str, float]: + if not samples: + return {} + values = np.asarray(samples, dtype=np.float64) + return { + "count": float(len(samples)), + "mean_ms": float(np.mean(values)), + "std_ms": float(np.std(values)), + "p50_ms": float(np.quantile(values, 0.50)), + "p90_ms": float(np.quantile(values, 0.90)), + "p95_ms": float(np.quantile(values, 0.95)), + "min_ms": float(np.min(values)), + "max_ms": float(np.max(values)), + } + + +def _inc(mapping: dict[str, int], key: object, value: int) -> None: + key_str = str(key) + mapping[key_str] = mapping.get(key_str, 0) + value + + +def _summarize_torch_module(module) -> dict[str, Any]: + param_dtypes: dict[str, int] = {} + param_dtype_examples: dict[str, list[str]] = {} + buffer_dtypes: dict[str, int] = {} + buffer_dtype_examples: dict[str, list[str]] = {} + devices: dict[str, int] = {} + param_count = 0 + buffer_count = 0 + trainable_param_count = 0 + for name, param in module.named_parameters(recurse=True): + numel = int(param.numel()) + param_count += numel + if param.requires_grad: + trainable_param_count += numel + _inc(param_dtypes, param.dtype, numel) + _inc(devices, param.device, numel) + examples = param_dtype_examples.setdefault(str(param.dtype), []) + if len(examples) < 16: + examples.append(name) + for name, buffer in module.named_buffers(recurse=True): + numel = int(buffer.numel()) + buffer_count += numel + _inc(buffer_dtypes, buffer.dtype, numel) + _inc(devices, buffer.device, numel) + examples = buffer_dtype_examples.setdefault(str(buffer.dtype), []) + if len(examples) < 16: + examples.append(name) + return { + "class": module.__class__.__name__, + "param_dtypes": param_dtypes, + "param_dtype_examples": param_dtype_examples, + "buffer_dtypes": buffer_dtypes, + "buffer_dtype_examples": buffer_dtype_examples, + "devices": devices, + "param_count": param_count, + "trainable_param_count": trainable_param_count, + "buffer_count": buffer_count, + } + + +def _torch_autocast_dtype(torch_module, device: str) -> str: + get_autocast_dtype = getattr(torch_module, "get_autocast_dtype", None) + if get_autocast_dtype is not None: + return str(get_autocast_dtype(device)) + if device == "cuda": + return str(torch_module.get_autocast_gpu_dtype()) + return str(torch_module.get_autocast_cpu_dtype()) + + +def openpi_precision_metadata(policy) -> dict[str, Any]: + import torch + + metadata: dict[str, Any] = { + "policy_class": policy.__class__.__name__, + "is_pytorch_model": bool(getattr(policy, "_is_pytorch_model", False)), + "pytorch_device": str(getattr(policy, "_pytorch_device", "")), + "torch_default_dtype": str(torch.get_default_dtype()), + "torch_autocast_cpu_dtype": _torch_autocast_dtype(torch, "cpu"), + } + if torch.cuda.is_available(): + metadata["torch_autocast_cuda_dtype"] = _torch_autocast_dtype(torch, "cuda") + + modules = {} + for name, value in vars(policy).items(): + if isinstance(value, torch.nn.Module): + modules[name] = _summarize_torch_module(value) + metadata["torch_modules"] = modules + return metadata + + +def sglang_precision_metadata(pipeline) -> dict[str, Any]: + import torch + + metadata: dict[str, Any] = { + "torch_default_dtype": str(torch.get_default_dtype()), + "torch_autocast_cpu_dtype": _torch_autocast_dtype(torch, "cpu"), + } + if torch.cuda.is_available(): + metadata["torch_autocast_cuda_dtype"] = _torch_autocast_dtype(torch, "cuda") + + modules = {} + policy_model = pipeline.get_module("policy_model") + if isinstance(policy_model, torch.nn.Module): + modules["policy_model"] = _summarize_torch_module(policy_model) + core_model = getattr(policy_model, "core_model", None) + if isinstance(core_model, torch.nn.Module): + modules["core_model"] = _summarize_torch_module(core_model) + metadata["torch_modules"] = modules + return metadata + + +def _image(rng: np.random.Generator, *, chw: bool = False) -> np.ndarray: + image = rng.integers(0, 256, size=(224, 224, 3), dtype=np.uint8) + if chw: + return np.transpose(image, (2, 0, 1)) + return image + + +def _make_libero_observation( + rng: np.random.Generator, + prompt: str, +) -> tuple[dict[str, Any], dict[str, Any]]: + base_image = _image(rng) + wrist_image = _image(rng) + state = rng.random(8, dtype=np.float32) + openpi_obs = { + "observation/state": state, + "observation/image": base_image, + "observation/wrist_image": wrist_image, + "prompt": prompt, + } + sglang_observation = { + "images": { + "image": base_image, + "image2": wrist_image, + }, + "state": state, + } + return openpi_obs, sglang_observation + + +def _make_aloha_observation( + rng: np.random.Generator, + prompt: str, +) -> tuple[dict[str, Any], dict[str, Any]]: + cam_high = _image(rng, chw=True) + cam_left = _image(rng, chw=True) + cam_right = _image(rng, chw=True) + state = np.ones((14,), dtype=np.float32) + openpi_obs = { + "state": state, + "images": { + "cam_high": cam_high, + "cam_low": _image(rng, chw=True), + "cam_left_wrist": cam_left, + "cam_right_wrist": cam_right, + }, + "prompt": prompt, + } + sglang_state = np.zeros((32,), dtype=np.float32) + sglang_state[: state.shape[0]] = state + sglang_observation = { + "images": { + "base_0_rgb": np.transpose(cam_high, (1, 2, 0)), + "left_wrist_0_rgb": np.transpose(cam_left, (1, 2, 0)), + "right_wrist_0_rgb": np.transpose(cam_right, (1, 2, 0)), + }, + "state": sglang_state, + } + return openpi_obs, sglang_observation + + +def build_observations( + profile: Pi05BenchProfile, + count: int, + seed: int, +) -> tuple[list[dict[str, Any]], list[dict[str, Any]]]: + rng = np.random.default_rng(seed) + openpi_observations = [] + sglang_observations = [] + for _ in range(count): + if profile.name == "libero": + openpi_obs, sglang_obs = _make_libero_observation(rng, profile.prompt) + elif profile.name == "aloha": + openpi_obs, sglang_obs = _make_aloha_observation(rng, profile.prompt) + else: + raise ValueError(f"Unsupported Pi0.5 benchmark profile: {profile.name}") + openpi_observations.append(openpi_obs) + sglang_observations.append(sglang_obs) + return openpi_observations, sglang_observations + + +def _json_tensor(array: np.ndarray) -> dict[str, Any]: + return { + "dtype": str(array.dtype), + "shape": list(array.shape), + "values": array.tolist(), + } + + +def build_sglang_payload( + profile: Pi05BenchProfile, + observation: dict[str, Any], + *, + num_inference_steps: int, + prefix_cache: bool, + cuda_graph: bool, + noise: np.ndarray | None, + response_format: str = "envelope", +) -> dict[str, Any]: + encoded_images = { + key: _json_tensor(np.asarray(value)) + for key, value in observation["images"].items() + } + encoded_observation = { + "images": encoded_images, + "state": _json_tensor(np.asarray(observation["state"], dtype=np.float32)), + } + if noise is not None: + encoded_observation["noise"] = _json_tensor(noise.astype(np.float32)) + return { + "model": profile.sglang_model, + "input": { + "task": profile.prompt, + "observation": encoded_observation, + }, + "parameters": { + "num_inference_steps": num_inference_steps, + }, + "runtime": { + "return_timing": True, + "prefix_cache": prefix_cache, + "cuda_graph": cuda_graph, + "response_format": response_format, + }, + } + + +def build_sglang_python_payload( + profile: Pi05BenchProfile, + observation: dict[str, Any], + *, + num_inference_steps: int, + prefix_cache: bool, + cuda_graph: bool, + noise: np.ndarray | None, + response_format: str = "envelope", +) -> dict[str, Any]: + encoded_observation = { + "images": { + key: np.asarray(value) for key, value in observation["images"].items() + }, + "state": np.asarray(observation["state"], dtype=np.float32), + } + if noise is not None: + encoded_observation["noise"] = noise.astype(np.float32) + return { + "model": profile.sglang_model, + "input": { + "task": profile.prompt, + "observation": encoded_observation, + }, + "parameters": { + "num_inference_steps": num_inference_steps, + }, + "runtime": { + "return_timing": True, + "prefix_cache": prefix_cache, + "cuda_graph": cuda_graph, + "output_format": "numpy", + "response_format": response_format, + }, + } + + +def build_sglang_openpi_ws_payload( + profile: Pi05BenchProfile, + observation: dict[str, Any], + *, + num_inference_steps: int, + prefix_cache: bool, + cuda_graph: bool, + noise: np.ndarray | None, +) -> dict[str, Any]: + payload: dict[str, Any] = { + "task": profile.prompt, + "observation.state": np.asarray(observation["state"], dtype=np.float32), + "num_inference_steps": num_inference_steps, + "enable_pi_prefix_cache": prefix_cache, + "enable_pi_cuda_graph": cuda_graph, + "output_format": "numpy", + } + for key, value in observation["images"].items(): + payload[f"observation.images.{key}"] = np.asarray(value) + if noise is not None: + payload["observation.noise"] = noise.astype(np.float32) + return payload + + +def _post_action( + session: requests.Session, + url: str, + payload: dict[str, Any], + timeout_s: float, +) -> dict[str, Any]: + response = session.post(url, json=payload, timeout=timeout_s) + response.raise_for_status() + return response.json() + + +def _post_action_msgpack( + session: requests.Session, + url: str, + payload: dict[str, Any], + timeout_s: float, +) -> dict[str, Any]: + response = session.post( + url, + data=pack_msgpack(payload), + headers={ + "Content-Type": "application/msgpack", + "Accept": "application/msgpack", + }, + timeout=timeout_s, + ) + response.raise_for_status() + return unpack_msgpack(response.content) + + +def _get_action_metadata(session: requests.Session, url: str, timeout_s: float): + response = session.get( + url.rstrip("/") + "/v1/actions/metadata", + timeout=timeout_s, + ) + response.raise_for_status() + return response.json() + + +def run_sglang_http( + url: str, + payloads: list[dict[str, Any]], + *, + warmup: int, + repeats: int, + batch_size: int, + timeout_s: float, + msgpack: bool = False, +) -> dict[str, Any]: + endpoint = url.rstrip("/") + "/v1/actions/generations" + post_action = _post_action_msgpack if msgpack else _post_action + + single_latencies = [] + single_outputs = [] + batch_latencies = [] + with requests.Session() as session: + metadata = _get_action_metadata(session, url, timeout_s) + for idx in range(min(warmup, len(payloads))): + post_action(session, endpoint, payloads[idx], timeout_s) + + for idx in range(repeats): + payload = payloads[idx % len(payloads)] + start = time.perf_counter() + output = post_action(session, endpoint, payload, timeout_s) + single_latencies.append((time.perf_counter() - start) * 1000) + single_outputs.append(output) + + if batch_size > 1: + sessions = [requests.Session() for _ in range(batch_size)] + try: + with ThreadPoolExecutor(max_workers=batch_size) as pool: + + def post_item(item): + session, payload = item + return post_action(session, endpoint, payload, timeout_s) + + for warmup_idx in range(warmup): + batch = [ + payloads[(warmup_idx * batch_size + offset) % len(payloads)] + for offset in range(batch_size) + ] + list(pool.map(post_item, zip(sessions, batch))) + + for start_idx in range(repeats): + batch = [ + payloads[(start_idx * batch_size + offset) % len(payloads)] + for offset in range(batch_size) + ] + start = time.perf_counter() + list(pool.map(post_item, zip(sessions, batch))) + batch_latencies.append((time.perf_counter() - start) * 1000) + finally: + for session in sessions: + session.close() + + stage_timings = {} + for output in single_outputs: + for key, value in output.get("timings", {}).items(): + stage_timings.setdefault(key, []).append(float(value)) + + return { + "single": _stats_ms(single_latencies), + "batch": _stats_ms(batch_latencies), + "batch_size": batch_size, + "stage_timings": { + key: _stats_ms(values) for key, values in stage_timings.items() + }, + "first_output": single_outputs[0] if single_outputs else None, + "batch_mode": "concurrent_http_msgpack" if msgpack else "concurrent_http_json", + "metadata": metadata, + } + + +def _action_ws_url(url: str) -> str: + if url.startswith("https://"): + return "wss://" + url[len("https://") :].rstrip("/") + "/openpi/policy" + if url.startswith("http://"): + return "ws://" + url[len("http://") :].rstrip("/") + "/openpi/policy" + return url.rstrip("/") + "/openpi/policy" + + +async def _ws_send_recv(websocket, payload: dict[str, Any]) -> dict[str, Any]: + await websocket.send(pack_msgpack(payload)) + response = await websocket.recv() + if isinstance(response, str): + raise RuntimeError(response) + return unpack_msgpack(response) + + +async def _run_sglang_openpi_ws_async( + url: str, + payloads: list[dict[str, Any]], + *, + warmup: int, + repeats: int, + batch_size: int, +) -> dict[str, Any]: + import websockets + + endpoint = _action_ws_url(url) + async with websockets.connect(endpoint, max_size=None) as websocket: + metadata = unpack_msgpack(await websocket.recv()) + for idx in range(min(warmup, len(payloads))): + await _ws_send_recv(websocket, payloads[idx]) + + single_latencies = [] + single_outputs = [] + for idx in range(repeats): + payload = payloads[idx % len(payloads)] + start = time.perf_counter() + output = await _ws_send_recv(websocket, payload) + single_latencies.append((time.perf_counter() - start) * 1000) + single_outputs.append(output) + + batch_latencies = [] + if batch_size > 1: + websockets_list = [] + try: + for _ in range(batch_size): + websocket = await websockets.connect(endpoint, max_size=None) + await websocket.recv() + websockets_list.append(websocket) + + for warmup_idx in range(warmup): + batch = [ + payloads[(warmup_idx * batch_size + offset) % len(payloads)] + for offset in range(batch_size) + ] + await asyncio.gather( + *[ + _ws_send_recv(websocket, payload) + for websocket, payload in zip(websockets_list, batch) + ] + ) + + for start_idx in range(repeats): + batch = [ + payloads[(start_idx * batch_size + offset) % len(payloads)] + for offset in range(batch_size) + ] + start = time.perf_counter() + await asyncio.gather( + *[ + _ws_send_recv(websocket, payload) + for websocket, payload in zip(websockets_list, batch) + ] + ) + batch_latencies.append((time.perf_counter() - start) * 1000) + finally: + for websocket in websockets_list: + await websocket.close() + + stage_timings = {} + server_timings = {} + for output in single_outputs: + for key, value in output.get("timings", {}).items(): + stage_timings.setdefault(key, []).append(float(value)) + for key, value in output.get("server_timing", {}).items(): + server_timings.setdefault(key, []).append(float(value)) + + return { + "single": _stats_ms(single_latencies), + "batch": _stats_ms(batch_latencies), + "batch_size": batch_size, + "stage_timings": { + key: _stats_ms(values) for key, values in stage_timings.items() + }, + "server_timings": { + key: _stats_ms(values) for key, values in server_timings.items() + }, + "first_output": single_outputs[0] if single_outputs else None, + "batch_mode": "persistent_openpi_websocket", + "metadata": metadata, + } + + +def run_sglang_openpi_ws( + url: str, + payloads: list[dict[str, Any]], + *, + warmup: int, + repeats: int, + batch_size: int, +) -> dict[str, Any]: + return asyncio.run( + _run_sglang_openpi_ws_async( + url, + payloads, + warmup=warmup, + repeats=repeats, + batch_size=batch_size, + ) + ) + + +def create_sglang_python_pipeline( + model_path: str, + *, + pipeline_config_path: str | None, +): + from sglang.multimodal_gen.runtime.pipelines.pi05 import Pi05Pipeline + from sglang.multimodal_gen.runtime.pipelines_core.executors.sync_executor import ( + SyncExecutor, + ) + from sglang.multimodal_gen.runtime.server_args import ( + ServerArgs, + set_global_server_args, + ) + + kwargs: dict[str, Any] = { + "model_path": model_path, + "warmup_mode": "off", + "num_gpus": 1, + } + if pipeline_config_path: + kwargs["pipeline_config_path"] = pipeline_config_path + server_args = ServerArgs.from_kwargs(**kwargs) + set_global_server_args(server_args) + pipeline = Pi05Pipeline( + model_path, + server_args, + executor=SyncExecutor(server_args=server_args), + ) + return pipeline, server_args + + +def _make_sglang_python_req(server_args, payload: dict[str, Any]): + from sglang.multimodal_gen.runtime.entrypoints.utils import prepare_request + from sglang.multimodal_gen.runtime.entrypoints.vla.protocol import ( + build_action_sampling_params, + ) + + sampling_params = build_action_sampling_params(payload, server_args) + req = prepare_request(server_args, sampling_params) + req.suppress_logs = True + return req + + +def _run_sglang_python_once(pipeline, server_args, payload: dict[str, Any]): + req = _make_sglang_python_req(server_args, payload) + output_batch = pipeline.forward(req, server_args) + if output_batch.error: + raise RuntimeError(output_batch.error) + if not output_batch.output: + raise RuntimeError("SGLang Python policy returned no output") + return output_batch.output[0] + + +def _run_sglang_python_group(pipeline, server_args, payloads: list[dict[str, Any]]): + reqs = [_make_sglang_python_req(server_args, payload) for payload in payloads] + output_batches = pipeline.forward_batch(reqs, server_args) + outputs = [] + for output_batch in output_batches: + if output_batch.error: + raise RuntimeError(output_batch.error) + if not output_batch.output: + raise RuntimeError("SGLang Python grouped policy returned no output") + outputs.append(output_batch.output[0]) + return outputs + + +def run_sglang_python( + model_path: str, + payloads: list[dict[str, Any]], + *, + pipeline_config_path: str | None, + warmup: int, + repeats: int, + batch_size: int, + batch_mode: str, +) -> dict[str, Any]: + pipeline, server_args = create_sglang_python_pipeline( + model_path, + pipeline_config_path=pipeline_config_path, + ) + from sglang.multimodal_gen.runtime.entrypoints.vla.protocol import action_metadata + + metadata = action_metadata(server_args) + metadata["precision"] = sglang_precision_metadata(pipeline) + for idx in range(min(warmup, len(payloads))): + _run_sglang_python_once(pipeline, server_args, payloads[idx]) + + single_latencies = [] + single_outputs = [] + for idx in range(repeats): + payload = payloads[idx % len(payloads)] + start = time.perf_counter() + output = _run_sglang_python_once(pipeline, server_args, payload) + single_latencies.append((time.perf_counter() - start) * 1000) + single_outputs.append(output) + + batch_latencies = [] + batch_outputs = [] + if batch_size > 1: + for warmup_idx in range(warmup): + batch = [ + payloads[(warmup_idx * batch_size + offset) % len(payloads)] + for offset in range(batch_size) + ] + if batch_mode == "grouped": + _run_sglang_python_group(pipeline, server_args, batch) + else: + for payload in batch: + _run_sglang_python_once(pipeline, server_args, payload) + for start_idx in range(repeats): + batch = [ + payloads[(start_idx * batch_size + offset) % len(payloads)] + for offset in range(batch_size) + ] + start = time.perf_counter() + if batch_mode == "grouped": + outputs = _run_sglang_python_group(pipeline, server_args, batch) + else: + outputs = [] + for payload in batch: + outputs.append( + _run_sglang_python_once(pipeline, server_args, payload) + ) + batch_latencies.append((time.perf_counter() - start) * 1000) + batch_outputs.extend(outputs) + + stage_timings = {} + for output in single_outputs: + for key, value in output.get("timings", {}).items(): + stage_timings.setdefault(key, []).append(float(value)) + batch_stage_timings = {} + for output in batch_outputs: + for key, value in output.get("timings", {}).items(): + batch_stage_timings.setdefault(key, []).append(float(value)) + + return { + "single": _stats_ms(single_latencies), + "batch": _stats_ms(batch_latencies), + "batch_size": batch_size, + "stage_timings": { + key: _stats_ms(values) for key, values in stage_timings.items() + }, + "batch_stage_timings": { + key: _stats_ms(values) for key, values in batch_stage_timings.items() + }, + "first_output": single_outputs[0] if single_outputs else None, + "batch_mode": f"python_policy_{batch_mode}", + "metadata": metadata, + } + + +def create_openpi_policy( + config_name: str, + checkpoint_dir: str, + *, + pytorch_device: str, + num_inference_steps: int, + pytorch_compile_mode: str | None, +): + from openpi.policies import policy_config + from openpi.training import config as openpi_config + + train_config = openpi_config.get_config(config_name) + if pytorch_compile_mode != "keep": + train_config = dataclasses.replace( + train_config, + model=dataclasses.replace( + train_config.model, + pytorch_compile_mode=pytorch_compile_mode, + ), + ) + return policy_config.create_trained_policy( + train_config, + checkpoint_dir, + sample_kwargs={"num_steps": num_inference_steps}, + pytorch_device=pytorch_device, + ) + + +def _openpi_infer(policy, observation: dict[str, Any], noise: np.ndarray | None): + if noise is None: + return policy.infer(observation) + return policy.infer(observation, noise=noise) + + +def _openpi_direct_batch( + policy, + observations: list[dict[str, Any]], + noises: list[np.ndarray] | None, +): + import jax + import numpy as onp + from openpi.models import model as openpi_model + + inputs_list = [ + policy._input_transform(jax.tree.map(lambda value: value, observation)) + for observation in observations + ] + batched_inputs = jax.tree.map( + lambda *values: onp.stack(values, axis=0), + *inputs_list, + ) + sample_kwargs = dict(policy._sample_kwargs) + if noises is not None: + noise = onp.stack(noises, axis=0) + if policy._is_pytorch_model: + import torch + + sample_kwargs["noise"] = torch.from_numpy(noise).to(policy._pytorch_device) + else: + import jax.numpy as jnp + + sample_kwargs["noise"] = jnp.asarray(noise) + + if policy._is_pytorch_model: + import torch + + inputs = jax.tree.map( + lambda value: torch.from_numpy(onp.asarray(value)).to( + policy._pytorch_device + ), + batched_inputs, + ) + sample_key_or_device = policy._pytorch_device + else: + import jax.numpy as jnp + + inputs = jax.tree.map(lambda value: jnp.asarray(value), batched_inputs) + policy._rng, sample_key_or_device = jax.random.split(policy._rng) + + observation = openpi_model.Observation.from_dict(inputs) + actions = policy._sample_actions(sample_key_or_device, observation, **sample_kwargs) + if policy._is_pytorch_model: + actions_np = actions.detach().cpu().numpy() + states_np = inputs["state"].detach().cpu().numpy() + else: + actions_np = onp.asarray(actions) + states_np = onp.asarray(inputs["state"]) + + outputs = [] + for idx in range(actions_np.shape[0]): + outputs.append( + policy._output_transform( + { + "state": states_np[idx], + "actions": actions_np[idx], + } + ) + ) + return outputs + + +def run_openpi_policy( + policy, + observations: list[dict[str, Any]], + *, + warmup: int, + repeats: int, + batch_size: int, + noise: np.ndarray | None, + batch_mode: str, +) -> dict[str, Any]: + for idx in range(min(warmup, len(observations))): + _openpi_infer(policy, observations[idx], noise) + + single_latencies = [] + single_outputs = [] + for idx in range(repeats): + observation = observations[idx % len(observations)] + start = time.perf_counter() + output = _openpi_infer(policy, observation, noise) + single_latencies.append((time.perf_counter() - start) * 1000) + single_outputs.append(output) + + batch_latencies = [] + if batch_size > 1: + for warmup_idx in range(warmup): + batch = [ + observations[(warmup_idx * batch_size + offset) % len(observations)] + for offset in range(batch_size) + ] + noises = [noise] * len(batch) if noise is not None else None + if batch_mode == "direct_model": + _openpi_direct_batch(policy, batch, noises) + elif batch_mode == "policy_loop": + for obs in batch: + _openpi_infer(policy, obs, noise) + else: + raise ValueError(f"Unsupported OpenPI batch mode: {batch_mode}") + for start_idx in range(repeats): + batch = [ + observations[(start_idx * batch_size + offset) % len(observations)] + for offset in range(batch_size) + ] + noises = [noise] * len(batch) if noise is not None else None + start = time.perf_counter() + if batch_mode == "direct_model": + _openpi_direct_batch(policy, batch, noises) + elif batch_mode == "policy_loop": + for obs in batch: + _openpi_infer(policy, obs, noise) + else: + raise ValueError(f"Unsupported OpenPI batch mode: {batch_mode}") + batch_latencies.append((time.perf_counter() - start) * 1000) + + policy_timings = {} + for output in single_outputs: + for key, value in output.get("policy_timing", {}).items(): + policy_timings.setdefault(key, []).append(float(value)) + + precision = openpi_precision_metadata(policy) + first_actions = _openpi_actions(single_outputs[0]) if single_outputs else None + if first_actions is not None: + precision["output_action_dtype"] = str(first_actions.dtype) + precision["output_action_shape"] = list(first_actions.shape) + + return { + "single": _stats_ms(single_latencies), + "batch": _stats_ms(batch_latencies), + "batch_size": batch_size, + "policy_timings": { + key: _stats_ms(values) for key, values in policy_timings.items() + }, + "first_output": single_outputs[0] if single_outputs else None, + "batch_mode": batch_mode, + "precision": precision, + } + + +def _sglang_actions(output: dict[str, Any]) -> np.ndarray | None: + if output is None: + return None + if "actions" in output: + return np.asarray(output["actions"], dtype=np.float32) + return np.asarray(output["data"][0]["action"]["values"], dtype=np.float32) + + +def _openpi_actions(output: dict[str, Any]) -> np.ndarray | None: + if output is None: + return None + return np.asarray(output["actions"], dtype=np.float32) + + +def compare_first_actions( + sglang_output: dict[str, Any] | None, + openpi_output: dict[str, Any] | None, +) -> dict[str, Any]: + sglang_actions = _sglang_actions(sglang_output) + openpi_actions = _openpi_actions(openpi_output) + if sglang_actions is None or openpi_actions is None: + return {"available": False} + horizon = min(sglang_actions.shape[0], openpi_actions.shape[0]) + dim = min(sglang_actions.shape[1], openpi_actions.shape[1]) + if horizon == 0 or dim == 0: + return { + "available": False, + "sglang_shape": list(sglang_actions.shape), + "openpi_shape": list(openpi_actions.shape), + } + diff = np.abs(sglang_actions[:horizon, :dim] - openpi_actions[:horizon, :dim]) + return { + "available": True, + "common_shape": [int(horizon), int(dim)], + "sglang_shape": list(sglang_actions.shape), + "openpi_shape": list(openpi_actions.shape), + "max_abs_diff": float(np.max(diff)), + "mean_abs_diff": float(np.mean(diff)), + } + + +def print_summary(result: dict[str, Any]) -> None: + print(json.dumps(result, indent=2, sort_keys=True)) + sglang_result = result.get("sglang") or {} + openpi_result = result.get("openpi") or {} + sgl = sglang_result.get("single", {}).get("mean_ms") + opi = openpi_result.get("single", {}).get("mean_ms") + if sgl and opi: + print( + "\nSingle mean latency: " + f"SGLang={sgl:.2f} ms, OpenPI={opi:.2f} ms, " + f"speedup={opi / sgl:.2f}x" + ) + sgl_batch = sglang_result.get("batch", {}).get("mean_ms") + opi_batch = openpi_result.get("batch", {}).get("mean_ms") + batch_size = sglang_result.get("batch_size") or openpi_result.get("batch_size", 0) + if sgl_batch and opi_batch and batch_size: + print( + "Batch mean latency: " + f"SGLang={sgl_batch:.2f} ms/{batch_size} req, " + f"OpenPI={opi_batch:.2f} ms/{batch_size} req, " + f"speedup={opi_batch / sgl_batch:.2f}x" + ) + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--profile", choices=sorted(PROFILES), default="libero") + parser.add_argument("--sglang-url", default="http://127.0.0.1:30000") + parser.add_argument( + "--sglang-api", + choices=("http", "http_msgpack", "openpi_ws", "python"), + default="http", + ) + parser.add_argument("--sglang-model", default=None) + parser.add_argument("--sglang-pipeline-config-path", default=None) + parser.add_argument( + "--sglang-http-response-format", + choices=("envelope", "raw"), + default="envelope", + ) + parser.add_argument( + "--sglang-python-batch-mode", + choices=("loop", "grouped"), + default="loop", + ) + parser.add_argument("--openpi-config", default=None) + parser.add_argument("--openpi-checkpoint", default=None) + parser.add_argument("--openpi-device", default="cuda") + parser.add_argument( + "--openpi-pytorch-compile-mode", + choices=( + "keep", + "none", + "default", + "reduce-overhead", + "max-autotune", + "max-autotune-no-cudagraphs", + ), + default="keep", + ) + parser.add_argument( + "--openpi-batch-mode", + choices=("direct_model", "policy_loop"), + default="direct_model", + ) + parser.add_argument("--num-inference-steps", type=int, default=10) + parser.add_argument("--num-samples", type=int, default=16) + parser.add_argument("--repeats", type=int, default=20) + parser.add_argument("--warmup", type=int, default=3) + parser.add_argument("--batch-size", type=int, default=4) + parser.add_argument("--seed", type=int, default=42) + parser.add_argument("--timeout-s", type=float, default=180.0) + parser.add_argument("--disable-prefix-cache", action="store_true") + parser.add_argument("--disable-cuda-graph", action="store_true") + parser.add_argument("--deterministic-noise", action="store_true") + parser.add_argument("--skip-sglang", action="store_true") + parser.add_argument("--skip-openpi", action="store_true") + parser.add_argument("--output-file", default="") + return parser.parse_args() + + +def main() -> None: + args = parse_args() + profile = PROFILES[args.profile] + if args.sglang_model: + profile = Pi05BenchProfile( + name=profile.name, + sglang_model=args.sglang_model, + openpi_config=profile.openpi_config, + openpi_checkpoint=profile.openpi_checkpoint, + prompt=profile.prompt, + sglang_action_horizon=profile.sglang_action_horizon, + openpi_action_horizon=profile.openpi_action_horizon, + action_dim=profile.action_dim, + output_action_dim=profile.output_action_dim, + ) + openpi_config = args.openpi_config or profile.openpi_config + openpi_checkpoint = args.openpi_checkpoint or profile.openpi_checkpoint + + openpi_observations, sglang_observations = build_observations( + profile, + max(args.num_samples, args.batch_size), + args.seed, + ) + noise = None + if args.skip_sglang and args.skip_openpi: + raise ValueError("At least one backend must be enabled") + if args.deterministic_noise: + if profile.sglang_action_horizon != profile.openpi_action_horizon: + raise ValueError( + "--deterministic-noise requires matching SGLang and OpenPI " + f"action horizons; profile {profile.name!r} has " + f"{profile.sglang_action_horizon} vs {profile.openpi_action_horizon}" + ) + rng = np.random.default_rng(args.seed + 1) + noise = rng.standard_normal( + (profile.sglang_action_horizon, profile.action_dim), + dtype=np.float32, + ) + + payloads = [] + if args.skip_sglang: + pass + elif args.sglang_api in ("python", "http_msgpack"): + payloads = [ + build_sglang_python_payload( + profile, + observation, + num_inference_steps=args.num_inference_steps, + prefix_cache=not args.disable_prefix_cache, + cuda_graph=not args.disable_cuda_graph, + noise=noise, + response_format=args.sglang_http_response_format, + ) + for observation in sglang_observations + ] + elif args.sglang_api == "openpi_ws": + payloads = [ + build_sglang_openpi_ws_payload( + profile, + observation, + num_inference_steps=args.num_inference_steps, + prefix_cache=not args.disable_prefix_cache, + cuda_graph=not args.disable_cuda_graph, + noise=noise, + ) + for observation in sglang_observations + ] + else: + payloads = [ + build_sglang_payload( + profile, + observation, + num_inference_steps=args.num_inference_steps, + prefix_cache=not args.disable_prefix_cache, + cuda_graph=not args.disable_cuda_graph, + noise=noise, + response_format=args.sglang_http_response_format, + ) + for observation in sglang_observations + ] + + openpi_policy = None + if not args.skip_openpi: + openpi_policy = create_openpi_policy( + openpi_config, + openpi_checkpoint, + pytorch_device=args.openpi_device, + num_inference_steps=args.num_inference_steps, + pytorch_compile_mode=( + None + if args.openpi_pytorch_compile_mode == "none" + else args.openpi_pytorch_compile_mode + ), + ) + + sglang_result = None + if args.skip_sglang: + pass + elif args.sglang_api == "python": + sglang_result = run_sglang_python( + profile.sglang_model, + payloads, + pipeline_config_path=args.sglang_pipeline_config_path, + warmup=args.warmup, + repeats=args.repeats, + batch_size=args.batch_size, + batch_mode=args.sglang_python_batch_mode, + ) + elif args.sglang_api == "openpi_ws": + sglang_result = run_sglang_openpi_ws( + args.sglang_url, + payloads, + warmup=args.warmup, + repeats=args.repeats, + batch_size=args.batch_size, + ) + else: + sglang_result = run_sglang_http( + args.sglang_url, + payloads, + warmup=args.warmup, + repeats=args.repeats, + batch_size=args.batch_size, + timeout_s=args.timeout_s, + msgpack=args.sglang_api == "http_msgpack", + ) + openpi_result = None + if openpi_policy is not None: + openpi_result = run_openpi_policy( + openpi_policy, + openpi_observations, + warmup=args.warmup, + repeats=args.repeats, + batch_size=args.batch_size, + noise=noise, + batch_mode=args.openpi_batch_mode, + ) + + result = { + "profile": profile.name, + "sglang_model": profile.sglang_model, + "sglang_api": args.sglang_api, + "sglang_pipeline_config_path": args.sglang_pipeline_config_path, + "sglang_http_response_format": args.sglang_http_response_format, + "sglang_python_batch_mode": args.sglang_python_batch_mode, + "openpi_config": openpi_config, + "openpi_checkpoint": openpi_checkpoint, + "openpi_pytorch_compile_mode": args.openpi_pytorch_compile_mode, + "num_inference_steps": args.num_inference_steps, + "num_samples": args.num_samples, + "repeats": args.repeats, + "warmup": args.warmup, + "deterministic_noise": args.deterministic_noise, + "action_diff": compare_first_actions( + None if sglang_result is None else sglang_result.get("first_output"), + None if openpi_result is None else openpi_result.get("first_output"), + ), + "sglang": ( + None + if sglang_result is None + else { + key: value + for key, value in sglang_result.items() + if key not in ("first_output",) + } + ), + "openpi": ( + None + if openpi_result is None + else { + key: value + for key, value in openpi_result.items() + if key not in ("first_output",) + } + ), + } + if args.output_file: + with open(args.output_file, "w", encoding="utf-8") as f: + json.dump(result, f, indent=2, sort_keys=True) + print_summary(result) + + +if __name__ == "__main__": + main() diff --git a/python/sglang/multimodal_gen/configs/pipeline_configs/__init__.py b/python/sglang/multimodal_gen/configs/pipeline_configs/__init__.py index 2fa07b0b6..2801e2e4c 100644 --- a/python/sglang/multimodal_gen/configs/pipeline_configs/__init__.py +++ b/python/sglang/multimodal_gen/configs/pipeline_configs/__init__.py @@ -40,6 +40,7 @@ from sglang.multimodal_gen.configs.pipeline_configs.ltx_2 import ( LTX23PipelineConfig, ) from sglang.multimodal_gen.configs.pipeline_configs.mova import MOVAPipelineConfig +from sglang.multimodal_gen.configs.pipeline_configs.pi05 import Pi05PipelineConfig from sglang.multimodal_gen.configs.pipeline_configs.sana import SanaPipelineConfig from sglang.multimodal_gen.configs.pipeline_configs.stablediffusion3 import ( StableDiffusion3PipelineConfig, @@ -71,6 +72,7 @@ __all__ = [ "SanaPipelineConfig", "SlidingTileAttnConfig", "MOVAPipelineConfig", + "Pi05PipelineConfig", "StableDiffusion3PipelineConfig", "WanT2V480PConfig", "WanI2V480PConfig", diff --git a/python/sglang/multimodal_gen/configs/pipeline_configs/base.py b/python/sglang/multimodal_gen/configs/pipeline_configs/base.py index 09a2068ec..8def48f24 100644 --- a/python/sglang/multimodal_gen/configs/pipeline_configs/base.py +++ b/python/sglang/multimodal_gen/configs/pipeline_configs/base.py @@ -59,6 +59,7 @@ class ModelTaskType(Enum): I2I = auto() # Image to Image TI2I = auto() # Image to Image or Text-Image to Image I2M = auto() # Image to Mesh + VLA_ACTION = auto() # Vision-language-action policy output def is_image_gen(self) -> bool: return ( @@ -67,6 +68,22 @@ class ModelTaskType(Enum): or self == ModelTaskType.TI2I ) + def is_action_gen(self) -> bool: + return self == ModelTaskType.VLA_ACTION + + def is_mesh_gen(self) -> bool: + return self == ModelTaskType.I2M + + def is_video_gen(self) -> bool: + return ( + self == ModelTaskType.I2V + or self == ModelTaskType.T2V + or self == ModelTaskType.TI2V + ) + + def is_visual_gen(self) -> bool: + return self.is_image_gen() or self.is_video_gen() + def requires_image_input(self) -> bool: return ( self == ModelTaskType.I2V @@ -81,15 +98,17 @@ class ModelTaskType(Enum): or self == ModelTaskType.TI2I or self == ModelTaskType.TI2V or self == ModelTaskType.I2M + or self == ModelTaskType.VLA_ACTION ) def data_type(self) -> DataType: - if self == ModelTaskType.I2M: + if self.is_action_gen(): + return DataType.ACTION + if self.is_mesh_gen(): return DataType.MESH if self.is_image_gen(): return DataType.IMAGE - else: - return DataType.VIDEO + return DataType.VIDEO class STA_Mode(str, Enum): @@ -374,6 +393,10 @@ class PipelineConfig: """ return self.task_type in (ModelTaskType.T2I, ModelTaskType.T2V) + def supports_native_grouped_requests(self): + """Return whether dynamic batches should run as grouped Req lists.""" + return False + def estimate_request_cost(self, batch) -> float: """Return the relative cost used for batching admission caps. @@ -888,6 +911,8 @@ class PipelineConfig: pipeline_config_or_path: str | PipelineConfig | dict[str, Any] | None = ( kwargs.get(prefix_with_dot + "pipeline_config", None) or kwargs.get("pipeline_config") + or kwargs.get(prefix_with_dot + "pipeline_config_path", None) + or kwargs.get("pipeline_config_path") ) if model_path is None: raise ValueError("model_path is required in kwargs") @@ -943,36 +968,52 @@ class PipelineConfig: model_id=kwargs.get("model_id"), ) if model_info is None: - raise ValueError( - f"Could not get model info for '{model_path}'. " - f"If using a safetensors file, please specify pipeline_class_name" - ) - # 1.5. Adjust pipeline config for fine-tuned VAE if needed - pipeline_config_cls = model_info.pipeline_config_cls - # If an explicit pipeline_class_name refines the model-default config - # (e.g. SanaWMRealtimePipeline -> SanaWMRealtimeConfig, a subclass of - # the model-resolved SanaWMPipelineConfig), prefer the pipeline's own - # config so realtime-only wiring (the /v1/realtime_video adapter) is - # selected. Only applies when the explicit config strictly subclasses - # the model default, so non-realtime pipelines are unaffected. - if pipeline_class_name: - explicit_config_classes = get_pipeline_config_classes( - pipeline_class_name - ) - if explicit_config_classes is not None: - explicit_config_cls = explicit_config_classes[0] - if ( - isinstance(explicit_config_cls, type) - and isinstance(pipeline_config_cls, type) - and explicit_config_cls is not pipeline_config_cls - and issubclass(explicit_config_cls, pipeline_config_cls) - ): + if pipeline_class_name: + config_classes = get_pipeline_config_classes(pipeline_class_name) + if config_classes is not None: + pipeline_config_cls = config_classes[0] logger.info( - f"Refining pipeline config {pipeline_config_cls.__name__} " - f"-> {explicit_config_cls.__name__} for explicit " - f"pipeline_class_name={pipeline_class_name}" + "Using %s from explicit pipeline_class_name=%s", + pipeline_config_cls.__name__, + pipeline_class_name, ) - pipeline_config_cls = explicit_config_cls + else: + raise ValueError( + f"Could not get model info for '{model_path}'. " + "Please specify a valid model_id or pipeline_class_name." + ) + else: + raise ValueError( + f"Could not get model info for '{model_path}'. " + f"If using a safetensors file, please specify pipeline_class_name" + ) + else: + # 1.5. Adjust pipeline config for fine-tuned VAE if needed + pipeline_config_cls = model_info.pipeline_config_cls + # If an explicit pipeline_class_name refines the model-default config + # (e.g. SanaWMRealtimePipeline -> SanaWMRealtimeConfig, a subclass of + # the model-resolved SanaWMPipelineConfig), prefer the pipeline's own + # config so realtime-only wiring (the /v1/realtime_video adapter) is + # selected. Only applies when the explicit config strictly subclasses + # the model default, so non-realtime pipelines are unaffected. + if pipeline_class_name: + explicit_config_classes = get_pipeline_config_classes( + pipeline_class_name + ) + if explicit_config_classes is not None: + explicit_config_cls = explicit_config_classes[0] + if ( + isinstance(explicit_config_cls, type) + and isinstance(pipeline_config_cls, type) + and explicit_config_cls is not pipeline_config_cls + and issubclass(explicit_config_cls, pipeline_config_cls) + ): + logger.info( + f"Refining pipeline config {pipeline_config_cls.__name__} " + f"-> {explicit_config_cls.__name__} for explicit " + f"pipeline_class_name={pipeline_class_name}" + ) + pipeline_config_cls = explicit_config_cls vae_path = kwargs.get(prefix_with_dot + "vae_path") or kwargs.get("vae_path") if vae_path is None: component_paths = kwargs.get( diff --git a/python/sglang/multimodal_gen/configs/pipeline_configs/pi05.py b/python/sglang/multimodal_gen/configs/pipeline_configs/pi05.py new file mode 100644 index 000000000..0a72daede --- /dev/null +++ b/python/sglang/multimodal_gen/configs/pipeline_configs/pi05.py @@ -0,0 +1,99 @@ +# SPDX-License-Identifier: Apache-2.0 + +from dataclasses import dataclass, field + +from sglang.multimodal_gen.configs.pipeline_configs.base import ( + ModelTaskType, + PipelineConfig, +) +from sglang.multimodal_gen.configs.pipeline_configs.model_deployment_config import ( + ModelDeploymentConfig, +) + + +@dataclass +class Pi05PipelineConfig(PipelineConfig): + """Configuration for OpenPI / LeRobot Pi0.5 action policies.""" + + task_type: ModelTaskType = ModelTaskType.VLA_ACTION + should_use_guidance: bool = False + enable_autocast: bool = True + generator_device: str | None = None + + # OpenPI pi0.5 public checkpoint layout. + pi05: bool = True + paligemma_variant: str = "gemma_2b" + action_expert_variant: str = "gemma_300m" + max_token_len: int = 200 + action_horizon: int = 50 + action_dim: int = 32 + state_dim: int = 32 + output_action_dim: int = 32 + n_action_steps: int = 50 + default_num_inference_steps: int = 10 + time_embedding_min_period: float = 4e-3 + time_embedding_max_period: float = 4.0 + tokenizer_name: str = "google/paligemma-3b-pt-224" + + image_keys: tuple[str, ...] = ( + "base_0_rgb", + "left_wrist_0_rgb", + "right_wrist_0_rgb", + ) + empty_cameras: int = 0 + image_size: tuple[int, int] = (224, 224) + image_normalization_mean: tuple[float, float, float] = (0.5, 0.5, 0.5) + image_normalization_std: tuple[float, float, float] = (0.5, 0.5, 0.5) + + enable_global_prefix_cache: bool = False + enable_action_cuda_graph: bool = True + prefix_cache_max_entries: int = 1 + prefix_cache_layout_version: str = "pi05-prefix-v1" + offload_prefix_image_encoder: bool = False + offload_prefix_image_encoder_after_embed: bool = False + offload_prefix_token_embedding: bool = False + offload_prefix_language_layers: bool = False + offload_prefix_language_layers_after_prefix: bool = False + offload_prefix_language_layer_count_after_prefix: int = 0 + offload_prefix_language_layers_empty_cache: bool = True + offload_action_expert_after_denoise: bool = False + empty_cache_after_prefix: bool = False + + # Prefix VLM and action expert are separate logical groups. The concrete + # process-group construction lands with the native model parallel kernels. + prefix_parallel_strategy: str = "tp" + action_parallel_strategy: str = "sp" + parallel_layout_version: str = "pi05-split-prefix-action-v1" + + skip_unused_lm_head: bool = True + materialize_dtype: str = "bf16" + loader_component_map: dict[str, tuple[str, ...]] = field( + default_factory=lambda: { + "vision_tower": ("paligemma_with_expert.paligemma.model.vision_tower.",), + "paligemma": ("paligemma_with_expert.paligemma.model.language_model.",), + "multi_modal_projector": ( + "paligemma_with_expert.paligemma.model.multi_modal_projector.", + ), + "action_expert": ("paligemma_with_expert.gemma_expert.",), + "action_heads": ( + "action_in_proj.", + "action_out_proj.", + "time_mlp_in.", + "time_mlp_out.", + ), + } + ) + + def supports_dynamic_batching(self): + return True + + def supports_native_grouped_requests(self): + return True + + def estimate_request_cost(self, batch) -> float: + return float( + self.action_horizon * self.action_dim * self.default_num_inference_steps + ) + + def get_model_deployment_config(self) -> ModelDeploymentConfig: + return ModelDeploymentConfig() diff --git a/python/sglang/multimodal_gen/configs/sample/__init__.py b/python/sglang/multimodal_gen/configs/sample/__init__.py index 047622e33..3b20f860e 100644 --- a/python/sglang/multimodal_gen/configs/sample/__init__.py +++ b/python/sglang/multimodal_gen/configs/sample/__init__.py @@ -4,10 +4,14 @@ from sglang.multimodal_gen.configs.sample.diffusers_generic import ( DiffusersGenericSamplingParams, ) from sglang.multimodal_gen.configs.sample.ideogram import Ideogram4SamplingParams +from sglang.multimodal_gen.configs.sample.pi05 import Pi05SamplingParams from sglang.multimodal_gen.configs.sample.sampling_params import SamplingParams +from sglang.multimodal_gen.configs.sample.vla import VLASamplingParams __all__ = [ "SamplingParams", + "VLASamplingParams", "DiffusersGenericSamplingParams", "Ideogram4SamplingParams", + "Pi05SamplingParams", ] diff --git a/python/sglang/multimodal_gen/configs/sample/pi05.py b/python/sglang/multimodal_gen/configs/sample/pi05.py new file mode 100644 index 000000000..273b04d32 --- /dev/null +++ b/python/sglang/multimodal_gen/configs/sample/pi05.py @@ -0,0 +1,76 @@ +# SPDX-License-Identifier: Apache-2.0 + +from dataclasses import dataclass, field +from typing import Any + +from sglang.multimodal_gen.configs.sample.vla import VLASamplingParams + + +@dataclass +class Pi05SamplingParams(VLASamplingParams): + """Sampling parameters for Pi0.5 flow-matching action inference.""" + + num_inference_steps: int = 10 + + action_horizon: int = 50 + action_dim: int = 32 + output_format: str = "list" + return_timing: bool = True + enable_prefix_cache: bool = True + enable_cuda_graph: bool = True + + state: Any = field(default=None, metadata={"batch_sig_exclude": True}) + images: dict[str, Any] | None = field( + default=None, metadata={"batch_sig_exclude": True} + ) + image_masks: dict[str, bool] | None = field( + default=None, metadata={"batch_sig_exclude": True} + ) + camera_order: list[str] | tuple[str, ...] | None = field( + default=None, metadata={"batch_sig_exclude": True} + ) + noise: Any = field(default=None, metadata={"batch_sig_exclude": True}) + observation: dict[str, Any] | None = field( + default=None, metadata={"batch_sig_exclude": True} + ) + + def build_request_extra(self) -> dict[str, Any]: + extra = super().build_request_extra() + observation = dict(self.observation or {}) + if self.images is not None: + observation["images"] = self.images + if self.image_masks is not None: + observation["image_masks"] = self.image_masks + if self.state is not None: + observation["state"] = self.state + if self.camera_order is not None: + observation["camera_order"] = tuple(self.camera_order) + if self.prompt is not None: + observation["prompt"] = self.prompt + if self.noise is not None: + observation["noise"] = self.noise + + extra["vla"] = { + "observation": observation, + "options": { + "output_format": self.output_format, + "return_timing": self.return_timing, + "enable_prefix_cache": self.enable_prefix_cache, + "enable_cuda_graph": self.enable_cuda_graph, + }, + } + return extra + + def _validate(self): + super()._validate() + if self.action_horizon <= 0: + raise ValueError("action_horizon must be positive") + if self.action_dim <= 0: + raise ValueError("action_dim must be positive") + if self.output_format not in ("list", "numpy"): + raise ValueError("output_format must be 'list' or 'numpy'") + + def _set_output_file_name(self): + if self.output_file_name is None: + self.output_file_name = "pi05_action" + super()._set_output_file_name() diff --git a/python/sglang/multimodal_gen/configs/sample/sampling_params.py b/python/sglang/multimodal_gen/configs/sample/sampling_params.py index 3dc24ebc5..e280ba2e9 100644 --- a/python/sglang/multimodal_gen/configs/sample/sampling_params.py +++ b/python/sglang/multimodal_gen/configs/sample/sampling_params.py @@ -76,12 +76,15 @@ class DataType(Enum): IMAGE = auto() VIDEO = auto() MESH = auto() + ACTION = auto() def get_default_extension(self) -> str: if self == DataType.IMAGE: return "png" if self == DataType.VIDEO: return "mp4" + if self == DataType.ACTION: + return "json" return "glb" @@ -246,10 +249,8 @@ class SamplingParams: def _set_output_file_ext(self): # add extension if needed - if not any( - self.output_file_name.endswith(ext) - for ext in [".mp4", ".jpg", ".png", ".webp", ".obj", ".glb"] - ): + output_extensions = (".mp4", ".jpg", ".png", ".webp", ".obj", ".glb", ".json") + if not any(self.output_file_name.endswith(ext) for ext in output_extensions): self.output_file_name = ( f"{self.output_file_name}.{self.data_type.get_default_extension()}" ) @@ -329,6 +330,8 @@ class SamplingParams: def _adjust_output_quality(self, output_quality: str, data_type: DataType) -> int: """Convert output_quality string to compression level.""" + if data_type == DataType.ACTION: + return 0 output_quality_mapper = {"maximum": 100, "high": 90, "medium": 55, "low": 35} if output_quality == "default": return 50 if data_type == DataType.VIDEO else 75 @@ -469,18 +472,22 @@ class SamplingParams: """ check if the sampling params is compatible and valid with server_args """ - if pipeline_config.task_type.requires_image_input(): + task_type = pipeline_config.task_type + if task_type.is_action_gen(): + return + + if task_type.requires_image_input(): # requires image input if self.image_path is None: raise ValueError( - f"Served model with task type '{pipeline_config.task_type.name}' requires an 'image_path' input, but none was provided" + f"Served model with task type '{task_type.name}' requires an 'image_path' input, but none was provided" ) - if not pipeline_config.task_type.accepts_image_input(): + if not task_type.accepts_image_input(): # does not support image input if self.image_path is not None: raise ValueError( - f"input_reference is not supported for {pipeline_config.task_type.name} models." + f"input_reference is not supported for {task_type.name} models." ) def _adjust( @@ -494,7 +501,39 @@ class SamplingParams: # TODO: SamplingParams should not rely on ServerArgs pipeline_config = server_args.pipeline_config + task_type = pipeline_config.task_type + self.data_type = task_type.data_type() + self._adjust_output_path(server_args) + if task_type.is_action_gen(): + self._adjust_action_fields(server_args) + return + + if task_type.is_mesh_gen(): + self._adjust_mesh_fields(server_args, pipeline_config) + return + + if task_type.is_visual_gen(): + self._adjust_visual_fields(server_args, pipeline_config) + + def _adjust_output_path(self, server_args): + if self.output_path is None: + if server_args.output_path is not None: + self.output_path = server_args.output_path + logger.debug( + f"Overriding output_path with server configuration: {self.output_path}" + ) + else: + self.save_output = False + + def _adjust_action_fields(self, server_args): + self.return_file_paths_only = False + self.num_frames = 1 + self.adjust_frames = False + if self.save_output and not server_args.comfyui_mode: + self._set_output_file_name() + + def _adjust_mesh_fields(self, server_args, pipeline_config): if self.guidance_scale is None: try: from sglang.multimodal_gen.configs.pipeline_configs.hunyuan3d import ( @@ -507,17 +546,16 @@ class SamplingParams: self.guidance_scale = 1.0 except ImportError: self.guidance_scale = 1.0 + self.return_frames = False + self.return_video = False + self.num_frames = 1 + self.adjust_frames = False + if self.save_output and not server_args.comfyui_mode: + self._set_output_file_name() - self.data_type = server_args.pipeline_config.task_type.data_type() - - if self.output_path is None: - if server_args.output_path is not None: - self.output_path = server_args.output_path - logger.debug( - f"Overriding output_path with server configuration: {self.output_path}" - ) - else: - self.save_output = False + def _adjust_visual_fields(self, server_args, pipeline_config): + if self.guidance_scale is None: + self.guidance_scale = 1.0 # Process negative prompt if self.negative_prompt is not None and not self.negative_prompt.isspace(): @@ -576,7 +614,6 @@ class SamplingParams: if not server_args.pipeline_config.allow_set_num_frames(): logger.debug("Setting `num_frames` to 1 for image generation model") self.num_frames = 1 - else: # mandatory frame adjusting logic, mod # NOTE: We must apply adjust_num_frames BEFORE the SP alignment logic below. diff --git a/python/sglang/multimodal_gen/configs/sample/vla.py b/python/sglang/multimodal_gen/configs/sample/vla.py new file mode 100644 index 000000000..571f9428d --- /dev/null +++ b/python/sglang/multimodal_gen/configs/sample/vla.py @@ -0,0 +1,272 @@ +# SPDX-License-Identifier: Apache-2.0 + +import argparse +import dataclasses +import os +from dataclasses import dataclass, field +from typing import TYPE_CHECKING, Any + +from sglang.multimodal_gen.configs.sample.sampling_params import ( + DataType, + _sanitize_filename, +) +from sglang.multimodal_gen.utils import StoreBoolean, expand_path_fields + +if TYPE_CHECKING: + from sglang.multimodal_gen.runtime.server_args import ServerArgs + + +@dataclass +class VLASamplingParams: + """Sampling parameters for VLA/action-generation policies.""" + + data_type: DataType = DataType.ACTION + request_id: str | None = field(default=None, metadata={"batch_sig_exclude": True}) + prompt: str | list[str] | None = field( + default="", metadata={"batch_sig_exclude": True} + ) + num_outputs_per_prompt: int = 1 + seed: int | list[int] = field(default=42, metadata={"batch_sig_exclude": True}) + generator_device: str | None = None + num_inference_steps: int = 10 + + output_path: str | None = field(default=None, metadata={"batch_sig_exclude": True}) + output_file_name: str | None = field( + default=None, metadata={"batch_sig_exclude": True} + ) + save_output: bool = False + return_file_paths_only: bool = False + + profile: bool = field(default=False, metadata={"batch_sig_exclude": True}) + num_profiled_timesteps: int = field(default=5, metadata={"batch_sig_exclude": True}) + profile_all_stages: bool = field( + default=False, metadata={"batch_sig_exclude": True} + ) + debug: bool = field(default=False, metadata={"batch_sig_exclude": True}) + perf_dump_path: str | None = field( + default=None, metadata={"batch_sig_exclude": True} + ) + suppress_logs: bool = field(default=False, metadata={"batch_sig_exclude": True}) + + enable_sequence_shard: bool | None = None + max_sequence_length: int | None = None + no_override_protected_fields: bool = field( + default=False, metadata={"batch_sig_exclude": True} + ) + + def __post_init__(self) -> None: + self.data_type = DataType.ACTION + self._validate() + + env_steps = os.environ.get("SGLANG_TEST_NUM_INFERENCE_STEPS") + if env_steps is not None and self.num_inference_steps is not None: + self.num_inference_steps = int(env_steps) + + def build_request_extra(self) -> dict[str, Any]: + extra = {} + diffusers_kwargs = getattr(self, "diffusers_kwargs", None) + if diffusers_kwargs: + extra["diffusers_kwargs"] = diffusers_kwargs + explicit_fields = getattr(self, "_explicit_fields", None) + if explicit_fields is not None: + extra["explicit_fields"] = sorted(explicit_fields) + return extra + + def apply_request_extra(self, req: Any) -> None: + req.extra.update(self.build_request_extra()) + + def _validate(self): + if ( + not isinstance(self.num_outputs_per_prompt, int) + or self.num_outputs_per_prompt <= 0 + ): + raise ValueError( + "num_outputs_per_prompt must be a positive int, " + f"got {self.num_outputs_per_prompt!r}" + ) + + if isinstance(self.seed, list): + if not self.seed: + raise ValueError("seed list must not be empty") + for seed in self.seed: + if isinstance(seed, bool) or not isinstance(seed, int) or seed < 0: + raise ValueError( + f"seed list must contain non-negative ints, got {self.seed!r}" + ) + elif ( + isinstance(self.seed, bool) + or not isinstance(self.seed, int) + or self.seed < 0 + ): + raise ValueError( + f"seed must be a non-negative int or list of ints, got {self.seed!r}" + ) + + if ( + not isinstance(self.num_inference_steps, int) + or self.num_inference_steps <= 0 + ): + raise ValueError( + "num_inference_steps must be a positive int, " + f"got {self.num_inference_steps!r}" + ) + + if self.generator_device not in (None, "cuda", "musa", "cpu"): + raise ValueError( + "generator_device must be one of None, 'cuda', 'musa', or 'cpu', " + f"got {self.generator_device!r}" + ) + + def _validate_with_pipeline_config(self, pipeline_config): + if not pipeline_config.task_type.is_action_gen(): + raise ValueError( + f"VLASamplingParams requires an ACTION pipeline, got {pipeline_config.task_type.name}" + ) + + def _adjust(self, server_args: "ServerArgs"): + expand_path_fields(self) + self.data_type = DataType.ACTION + self.return_file_paths_only = False + if self.output_path is None and server_args.output_path is not None: + self.output_path = server_args.output_path + if self.output_path is None: + self.save_output = False + if self.save_output and not server_args.comfyui_mode: + self._set_output_file_name() + + def _set_output_file_ext(self): + if self.output_file_name and not self.output_file_name.endswith(".json"): + self.output_file_name = f"{self.output_file_name}.json" + + def _set_output_file_name(self): + if self.output_file_name is None: + self.output_file_name = "vla_action" + self.output_file_name = _sanitize_filename(self.output_file_name) + self._set_output_file_ext() + + def output_file_path(self): + if self.output_path is None or self.output_file_name is None: + return None + return os.path.join(self.output_path, self.output_file_name) + + def _merge_with_user_params( + self, + user_params: "VLASamplingParams", + explicit_fields: set[str] | None = None, + ): + if user_params is None: + return + + predefined_fields = set(type(self).__annotations__.keys()) + allow_override_protected = not user_params.no_override_protected_fields + for field_info in dataclasses.fields(user_params): + field_name = field_info.name + user_value = getattr(user_params, field_name) + if field_info.default is not dataclasses.MISSING: + default_class_value = field_info.default + elif field_info.default_factory is not dataclasses.MISSING: + default_class_value = field_info.default_factory() + else: + default_class_value = dataclasses.MISSING + + if explicit_fields is not None: + is_user_modified = field_name in explicit_fields + else: + is_user_modified = user_value != default_class_value + is_protected_field = field_name in predefined_fields + if is_user_modified and ( + allow_override_protected or not is_protected_field + ): + setattr(self, field_name, user_value) + + if explicit_fields is not None: + self._explicit_fields = set(explicit_fields) + self.__post_init__() + + @staticmethod + def add_cli_args(parser: Any) -> Any: + def add_argument(*name_or_flags, **kwargs): + kwargs.setdefault("default", argparse.SUPPRESS) + return parser.add_argument(*name_or_flags, **kwargs) + + add_argument( + "--prompt", + type=str, + nargs="+", + help="Language instruction(s) for the VLA policy.", + ) + add_argument( + "--num-inference-steps", + type=int, + help="Number of action denoising steps.", + ) + add_argument( + "--num-outputs-per-prompt", + type=int, + help="Number of candidate actions to generate per observation.", + ) + add_argument( + "--seed", + type=int, + nargs="+", + help="Random seed for action noise generation.", + ) + add_argument( + "--generator-device", + type=str, + choices=["cuda", "musa", "cpu"], + help="Device for random generator. Default: use the model-specific setting.", + ) + add_argument( + "--profile", + action="store_true", + help="Enable torch profiler for action denoising.", + ) + add_argument( + "--num-profiled-timesteps", + type=int, + help="Number of denoising timesteps to profile after warmup.", + ) + add_argument( + "--profile-all-stages", + action="store_true", + dest="profile_all_stages", + help="Used with --profile, profile all pipeline stages.", + ) + add_argument("--debug", action="store_true") + add_argument( + "--enable-sequence-shard", + action=StoreBoolean, + help="Enable sequence dimension shard with sequence parallelism.", + ) + add_argument( + "--max-sequence-length", + type=int, + help="Maximum prefix sequence length.", + ) + add_argument( + "--no-override-protected-fields", + action="store_true", + help="If set, disallow user params to override subclass-defined fields.", + ) + return parser + + @classmethod + def get_cli_args(cls, args: argparse.Namespace): + sampling_params_fields = {attr.name for attr in dataclasses.fields(cls)} + args_attrs = set(vars(args).keys()) + attrs = sampling_params_fields & args_attrs + cli_args = { + attr: getattr(args, attr) + for attr in attrs + if hasattr(args, attr) and getattr(args, attr) is not None + } + if isinstance(cli_args.get("seed"), list) and len(cli_args["seed"]) == 1: + cli_args["seed"] = cli_args["seed"][0] + return cli_args + + def output_size_str(self) -> str: + return "action" + + def seconds(self) -> float: + return 0.0 diff --git a/python/sglang/multimodal_gen/registry.py b/python/sglang/multimodal_gen/registry.py index cf751a86d..c5f06ef6f 100644 --- a/python/sglang/multimodal_gen/registry.py +++ b/python/sglang/multimodal_gen/registry.py @@ -76,6 +76,7 @@ from sglang.multimodal_gen.configs.pipeline_configs.mova import ( MOVA360PConfig, MOVA720PConfig, ) +from sglang.multimodal_gen.configs.pipeline_configs.pi05 import Pi05PipelineConfig from sglang.multimodal_gen.configs.pipeline_configs.qwen_image import ( QwenImageEditPipelineConfig, QwenImageEditPlus_2511_PipelineConfig, @@ -137,6 +138,7 @@ from sglang.multimodal_gen.configs.sample.mova import ( MOVA_360P_SamplingParams, MOVA_720P_SamplingParams, ) +from sglang.multimodal_gen.configs.sample.pi05 import Pi05SamplingParams from sglang.multimodal_gen.configs.sample.qwenimage import ( QwenImage2512SamplingParams, QwenImageEditPlusSamplingParams, @@ -629,6 +631,20 @@ def get_model_info( # Registration of model configs def _register_configs(): + # Pi0.5 / OpenPI / LeRobot action policies. + register_configs( + sampling_param_cls=Pi05SamplingParams, + pipeline_config_cls=Pi05PipelineConfig, + hf_model_paths=[ + "lerobot/pi05_base", + "lerobot/pi05_libero_base", + ], + model_detectors=[ + lambda hf_id: "pi05" in hf_id.lower(), + lambda hf_id: "pi0.5" in hf_id.lower(), + ], + ) + # LTX-2 register_configs( sampling_param_cls=LTX2SamplingParams, diff --git a/python/sglang/multimodal_gen/runtime/distributed/group_coordinator.py b/python/sglang/multimodal_gen/runtime/distributed/group_coordinator.py index ff398617b..4045e4a67 100644 --- a/python/sglang/multimodal_gen/runtime/distributed/group_coordinator.py +++ b/python/sglang/multimodal_gen/runtime/distributed/group_coordinator.py @@ -602,10 +602,11 @@ class GroupCoordinator: group = self.device_group metadata_group = self.cpu_group assert src < self.world_size, f"Invalid src rank ({src})" - src = self.ranks[src] + src_rank_in_group = src + src_global_rank = self.ranks[src_rank_in_group] rank = self.rank - if rank == src: + if rank == src_global_rank: metadata_list: List[Tuple[Any, Any]] = [] assert isinstance( tensor_dict, dict @@ -614,7 +615,7 @@ class GroupCoordinator: # `metadata_list` lives in CPU memory. # `broadcast_object_list` has serialization & deserialization, # all happening on CPU. Therefore, we can use the CPU group. - self.broadcast_object(metadata_list, src=src) + self.broadcast_object(metadata_list, src=src_rank_in_group) async_handles = [] for tensor in tensor_list: if tensor.numel() == 0: @@ -623,19 +624,22 @@ class GroupCoordinator: if tensor.is_cpu: # use metadata_group for CPU tensors handle = torch.distributed.broadcast( - tensor, src=src, group=metadata_group, async_op=True + tensor, + src=src_global_rank, + group=metadata_group, + async_op=True, ) else: # use group for GPU tensors handle = torch.distributed.broadcast( - tensor, src=src, group=group, async_op=True + tensor, src=src_global_rank, group=group, async_op=True ) async_handles.append(handle) for async_handle in async_handles: async_handle.wait() else: - metadata_list = self.broadcast_object(None, src=src) + metadata_list = self.broadcast_object(None, src=src_rank_in_group) tensor_dict = {} async_handles = [] for key, value in metadata_list: @@ -650,12 +654,15 @@ class GroupCoordinator: if tensor.is_cpu: # use metadata_group for CPU tensors handle = torch.distributed.broadcast( - tensor, src=src, group=metadata_group, async_op=True + tensor, + src=src_global_rank, + group=metadata_group, + async_op=True, ) else: # use group for GPU tensors handle = torch.distributed.broadcast( - tensor, src=src, group=group, async_op=True + tensor, src=src_global_rank, group=group, async_op=True ) async_handles.append(handle) _update_nested_dict(tensor_dict, key, tensor) diff --git a/python/sglang/multimodal_gen/runtime/entrypoints/diffusion_generator.py b/python/sglang/multimodal_gen/runtime/entrypoints/diffusion_generator.py index 2eaae8706..cd6e9504c 100644 --- a/python/sglang/multimodal_gen/runtime/entrypoints/diffusion_generator.py +++ b/python/sglang/multimodal_gen/runtime/entrypoints/diffusion_generator.py @@ -376,6 +376,34 @@ class DiffGenerator: return None return results[0] if len(results) == 1 else results + def generate_action( + self, + sampling_params_kwargs: dict | None = None, + external_trace_header: dict[str, str] | None = None, + ) -> dict[str, Any]: + sampling_params_kwargs = sampling_params_kwargs or {} + sampling_params = SamplingParams.from_user_sampling_params_args( + self.server_args.model_path, + server_args=self.server_args, + **sampling_params_kwargs, + ) + if sampling_params.data_type != DataType.ACTION: + raise ValueError( + f"generate_action requires an ACTION pipeline, got {sampling_params.data_type}" + ) + + req = prepare_request( + server_args=self.server_args, + sampling_params=sampling_params, + external_trace_header=external_trace_header, + ) + output_batch = self._send_to_scheduler_and_wait_for_response(req) + if output_batch.error: + raise RuntimeError(output_batch.error) + if output_batch.output is None: + raise RuntimeError("action policy returned no output") + return output_batch.output[0] + def _resolve_prompts( self, prompt: str | list[str] | None, @@ -430,9 +458,13 @@ class DiffGenerator: and output_index < len(output_batch.metrics_list) ): metrics = output_batch.metrics_list[output_index] + if req.data_type == DataType.ACTION: + size = ("action",) + else: + size = (req.height, req.width, req.num_frames) return dict( prompt=req.prompt, - size=(req.height, req.width, req.num_frames), + size=size, generation_time=generation_time, peak_memory_mb=output_batch.peak_memory_mb, metrics=metrics.to_dict() if metrics else {}, diff --git a/python/sglang/multimodal_gen/runtime/entrypoints/http_server.py b/python/sglang/multimodal_gen/runtime/entrypoints/http_server.py index 519dbd818..874b4e3f5 100644 --- a/python/sglang/multimodal_gen/runtime/entrypoints/http_server.py +++ b/python/sglang/multimodal_gen/runtime/entrypoints/http_server.py @@ -14,10 +14,7 @@ from fastapi import APIRouter, FastAPI, Request from fastapi.middleware.cors import CORSMiddleware from sglang.multimodal_gen.configs.sample.sampling_params import SamplingParams -from sglang.multimodal_gen.runtime.entrypoints.openai import ( - image_api, - video_api, -) +from sglang.multimodal_gen.runtime.entrypoints.openai import image_api, video_api from sglang.multimodal_gen.runtime.entrypoints.openai.protocol import ( VertexGenerateReqInput, ) @@ -33,6 +30,8 @@ from sglang.multimodal_gen.runtime.entrypoints.utils import ( prepare_request, save_outputs, ) +from sglang.multimodal_gen.runtime.entrypoints.vla import api as vla_api +from sglang.multimodal_gen.runtime.entrypoints.vla import openpi from sglang.multimodal_gen.runtime.scheduler_client import async_scheduler_client from sglang.multimodal_gen.runtime.server_args import ServerArgs, get_global_server_args from sglang.multimodal_gen.runtime.server_warmup import ( @@ -403,6 +402,9 @@ def create_app(server_args: ServerArgs): app.include_router(image_api.router) app.include_router(video_api.router) app.include_router(realtime_video_api.router) + if server_args.pipeline_config.task_type.is_action_gen(): + app.include_router(vla_api.router) + app.include_router(openpi.router) app.include_router(mesh_api.router) app.include_router(weights_api.router) app.include_router(rollout_api.router) diff --git a/python/sglang/multimodal_gen/runtime/entrypoints/utils.py b/python/sglang/multimodal_gen/runtime/entrypoints/utils.py index bca5c394c..d084bd9dc 100644 --- a/python/sglang/multimodal_gen/runtime/entrypoints/utils.py +++ b/python/sglang/multimodal_gen/runtime/entrypoints/utils.py @@ -8,6 +8,7 @@ This module provides a consolidated interface for generating videos using diffusion models. """ +import json import os import shutil import subprocess @@ -445,11 +446,13 @@ def prepare_request( if not isinstance(req.prompt, str): raise TypeError(f"`prompt` must be a string, but got {type(req.prompt)}") - if (req.width is not None and req.width <= 0) or ( - req.height is not None and req.height <= 0 + req_width = getattr(req, "width", None) + req_height = getattr(req, "height", None) + if (req_width is not None and req_width <= 0) or ( + req_height is not None and req_height <= 0 ): raise ValueError( - f"Height and width must be positive, got height={req.height}, width={req.width}" + f"Height and width must be positive, got height={req_height}, width={req_width}" ) if server_args.enable_trace: @@ -661,6 +664,21 @@ def save_outputs( output_paths: list[str] = [] for idx, sample in enumerate(outputs): save_file_path = build_output_path(idx) + if data_type == DataType.ACTION: + if samples_out is not None: + samples_out.append(sample) + if audios_out is not None: + audios_out.append(None) + if frames_out is not None: + frames_out.append([]) + if save_output and save_file_path: + os.makedirs(os.path.dirname(save_file_path) or ".", exist_ok=True) + with open(save_file_path, "w", encoding="utf-8") as f: + json.dump(sample, f, ensure_ascii=False) + logger.info(f"Output saved to {CYAN}{save_file_path}{RESET}") + output_paths.append(save_file_path) + continue + if data_type == DataType.VIDEO: sample = attach_audio_to_video_sample(sample, audio, idx) @@ -711,6 +729,9 @@ def post_process_sample( upscaling_scale: int = 4, ) -> list[Any]: """materialize frames and save outputs (optional)""" + if data_type == DataType.ACTION: + return [] + materialized = materialize_output_sample( sample, data_type, diff --git a/python/sglang/multimodal_gen/runtime/entrypoints/vla/__init__.py b/python/sglang/multimodal_gen/runtime/entrypoints/vla/__init__.py new file mode 100644 index 000000000..988131360 --- /dev/null +++ b/python/sglang/multimodal_gen/runtime/entrypoints/vla/__init__.py @@ -0,0 +1 @@ +# SPDX-License-Identifier: Apache-2.0 diff --git a/python/sglang/multimodal_gen/runtime/entrypoints/vla/api.py b/python/sglang/multimodal_gen/runtime/entrypoints/vla/api.py new file mode 100644 index 000000000..cc901ea83 --- /dev/null +++ b/python/sglang/multimodal_gen/runtime/entrypoints/vla/api.py @@ -0,0 +1,89 @@ +# SPDX-License-Identifier: Apache-2.0 + +from __future__ import annotations + +from fastapi import APIRouter, HTTPException, Request, Response, WebSocket + +from sglang.multimodal_gen.runtime.entrypoints.vla.protocol import ( + action_generation_response, + action_metadata, + action_raw_response, + infer_action, + pack_msgpack, + unpack_msgpack, +) +from sglang.multimodal_gen.runtime.entrypoints.vla.ws_utils import ( + run_action_msgpack_ws, +) +from sglang.multimodal_gen.runtime.server_args import ServerArgs +from sglang.srt.utils.json_response import orjson_response + +router = APIRouter(prefix="/v1/actions", tags=["actions"]) + + +def _wants_msgpack(request: Request) -> bool: + content_type = request.headers.get("content-type", "").lower() + accept = request.headers.get("accept", "").lower() + return "msgpack" in content_type or "msgpack" in accept + + +def _response_format(payload: dict) -> str: + runtime = payload.get("runtime") or {} + response_format = str(runtime.get("response_format", "envelope")).lower() + if response_format not in ("envelope", "raw"): + raise ValueError("runtime.response_format must be 'envelope' or 'raw'") + return response_format + + +def _prefer_numpy_output(payload: dict) -> None: + runtime = payload.setdefault("runtime", {}) + runtime.setdefault("output_format", "numpy") + + +@router.post("/generations") +async def create_action_generation(request: Request): + server_args: ServerArgs = request.app.state.server_args + try: + if "msgpack" in request.headers.get("content-type", "").lower(): + payload = unpack_msgpack(await request.body()) + else: + payload = await request.json() + wants_msgpack = _wants_msgpack(request) + if wants_msgpack: + _prefer_numpy_output(payload) + output = await infer_action(payload, server_args) + if _response_format(payload) == "raw": + response = action_raw_response(output, preserve_numpy=wants_msgpack) + else: + response = action_generation_response( + output, + server_args, + preserve_numpy=wants_msgpack, + ) + except ValueError as exc: + raise HTTPException(status_code=400, detail=str(exc)) from exc + if wants_msgpack: + return Response( + content=pack_msgpack(response), media_type="application/msgpack" + ) + return orjson_response(response) + + +@router.get("/metadata") +async def action_metadata_endpoint(request: Request): + return orjson_response(action_metadata(request.app.state.server_args)) + + +@router.websocket("/realtime") +async def action_realtime_ws(websocket: WebSocket): + server_args: ServerArgs = websocket.app.state.server_args + await run_action_msgpack_ws( + websocket, + server_args, + prepare_payload=_prefer_numpy_output, + build_response=lambda output: action_generation_response( + output, + server_args, + preserve_numpy=True, + ), + ) diff --git a/python/sglang/multimodal_gen/runtime/entrypoints/vla/openpi.py b/python/sglang/multimodal_gen/runtime/entrypoints/vla/openpi.py new file mode 100644 index 000000000..b12c3cbe4 --- /dev/null +++ b/python/sglang/multimodal_gen/runtime/entrypoints/vla/openpi.py @@ -0,0 +1,29 @@ +# SPDX-License-Identifier: Apache-2.0 + +from __future__ import annotations + +from typing import Any + +from fastapi import APIRouter, WebSocket + +from sglang.multimodal_gen.runtime.entrypoints.vla.ws_utils import ( + run_action_msgpack_ws, +) +from sglang.multimodal_gen.runtime.server_args import ServerArgs + +router = APIRouter() + + +def _prefer_numpy_output(observation: dict[str, Any]) -> None: + observation.setdefault("output_format", "numpy") + + +@router.websocket("/openpi/policy") +async def openpi_policy_ws(websocket: WebSocket): + server_args: ServerArgs = websocket.app.state.server_args + await run_action_msgpack_ws( + websocket, + server_args, + prepare_payload=_prefer_numpy_output, + build_response=lambda output: output, + ) diff --git a/python/sglang/multimodal_gen/runtime/entrypoints/vla/protocol.py b/python/sglang/multimodal_gen/runtime/entrypoints/vla/protocol.py new file mode 100644 index 000000000..16e8d2e8a --- /dev/null +++ b/python/sglang/multimodal_gen/runtime/entrypoints/vla/protocol.py @@ -0,0 +1,443 @@ +# SPDX-License-Identifier: Apache-2.0 + +from __future__ import annotations + +import base64 +import dataclasses +import io +import time +import uuid +from functools import lru_cache +from typing import Any + +import numpy as np +from PIL import Image + +from sglang.multimodal_gen.configs.sample.vla import VLASamplingParams +from sglang.multimodal_gen.runtime.entrypoints.utils import prepare_request +from sglang.multimodal_gen.runtime.scheduler_client import async_scheduler_client +from sglang.multimodal_gen.runtime.server_args import ServerArgs + + +def pack_numpy_payload(obj): + if isinstance(obj, (np.ndarray, np.generic)) and obj.dtype.kind in ("V", "O", "c"): + raise ValueError(f"Unsupported dtype: {obj.dtype}") + if isinstance(obj, np.ndarray): + return { + b"__ndarray__": True, + b"data": obj.tobytes(), + b"dtype": obj.dtype.str, + b"shape": obj.shape, + } + if isinstance(obj, np.generic): + return { + b"__npgeneric__": True, + b"data": obj.item(), + b"dtype": obj.dtype.str, + } + return obj + + +def unpack_numpy_payload(obj): + ndarray_marker = obj.get("__ndarray__") or obj.get(b"__ndarray__") + npgeneric_marker = obj.get("__npgeneric__") or obj.get(b"__npgeneric__") + data = obj.get("data", obj.get(b"data")) + dtype = obj.get("dtype", obj.get(b"dtype")) + shape = obj.get("shape", obj.get(b"shape")) + if ndarray_marker: + return np.ndarray( + buffer=data, + dtype=np.dtype(dtype), + shape=shape, + ) + if npgeneric_marker: + return np.dtype(dtype).type(data) + return obj + + +def pack_msgpack(payload: Any) -> bytes: + import msgpack + + return msgpack.packb(payload, default=pack_numpy_payload, use_bin_type=True) + + +def unpack_msgpack(payload: bytes) -> Any: + import msgpack + + return msgpack.unpackb(payload, object_hook=unpack_numpy_payload, raw=False) + + +def _decode_b64_image(payload: dict[str, Any]) -> Image.Image: + data = payload.get("b64_json") or payload.get("base64") + if not data: + raise ValueError("image payload requires b64_json") + if isinstance(data, str) and "," in data and data.startswith("data:"): + data = data.split(",", 1)[1] + return Image.open(io.BytesIO(base64.b64decode(data))).convert("RGB") + + +def _decode_tensor_payload(payload: dict[str, Any]) -> Any: + values = payload.get("values") + if values is None: + values = payload.get("data") + if values is None: + return payload + dtype = payload.get("dtype") + array = np.asarray(values, dtype=np.dtype(dtype) if dtype else None) + shape = payload.get("shape") + if shape is not None: + array = array.reshape(tuple(shape)) + return array + + +def _normalize_image_value(value: Any) -> Any: + if not isinstance(value, dict): + return value + if "b64_json" in value or "base64" in value: + return _decode_b64_image(value) + if "values" in value or "data" in value: + return _decode_tensor_payload(value) + return value + + +def _normalize_observation(observation: dict[str, Any]) -> dict[str, Any]: + normalized = dict(observation) + images = normalized.get("images") + if isinstance(images, dict): + normalized["images"] = { + name: _normalize_image_value(value) for name, value in images.items() + } + state = normalized.get("state") + if isinstance(state, dict): + normalized["state"] = _decode_tensor_payload(state) + observation_state = normalized.get("observation.state") + if isinstance(observation_state, dict): + normalized["observation.state"] = _decode_tensor_payload(observation_state) + noise = normalized.get("noise") + if isinstance(noise, dict): + normalized["noise"] = _decode_tensor_payload(noise) + observation_noise = normalized.get("observation.noise") + if isinstance(observation_noise, dict): + normalized["observation.noise"] = _decode_tensor_payload(observation_noise) + return normalized + + +def images_from_observation( + observation: dict[str, Any], + pipeline_config: Any, +) -> dict[str, Any]: + if isinstance(observation.get("images"), dict): + images = dict(observation["images"]) + else: + images = {} + for key in pipeline_config.image_keys: + if key in observation: + images[key] = observation[key] + full_key = f"observation.images.{key}" + if full_key in observation: + images[key] = observation[full_key] + return {name: _normalize_image_value(value) for name, value in images.items()} + + +def action_metadata(server_args: ServerArgs) -> dict[str, Any]: + pipeline_config = server_args.pipeline_config + policy_family = getattr( + pipeline_config, + "policy_family", + type(pipeline_config).__name__.removesuffix("PipelineConfig").lower(), + ) + return { + "object": "action.metadata", + "model": server_args.model_id or server_args.model_path, + "model_path": server_args.model_path, + "policy_family": policy_family, + "input": { + "image_keys": list(pipeline_config.image_keys), + "image_size": list(pipeline_config.image_size), + "state_dim": pipeline_config.state_dim, + }, + "output": { + "action_type": "continuous", + "action_horizon": pipeline_config.action_horizon, + "action_dim": pipeline_config.output_action_dim, + "padded_action_dim": pipeline_config.action_dim, + "dtype": "float32", + }, + "runtime": { + "materialize_dtype": pipeline_config.materialize_dtype, + "enable_autocast": pipeline_config.enable_autocast, + "parallelism": { + "num_gpus": server_args.num_gpus, + "tp_size": server_args.tp_size, + "sp_degree": server_args.sp_degree, + "ulysses_degree": server_args.ulysses_degree, + "ring_degree": server_args.ring_degree, + "prefix_strategy": pipeline_config.prefix_parallel_strategy, + "action_strategy": pipeline_config.action_parallel_strategy, + "layout_version": pipeline_config.parallel_layout_version, + }, + }, + "defaults": { + "num_inference_steps": pipeline_config.default_num_inference_steps, + "prefix_cache": ( + "auto" if pipeline_config.enable_global_prefix_cache else False + ), + "cuda_graph": "auto" if pipeline_config.enable_action_cuda_graph else False, + }, + "capabilities": { + "exact_prefix_cache": True, + "cuda_graph": pipeline_config.enable_action_cuda_graph, + "realtime_websocket": True, + "openpi_websocket": True, + "batch_inputs": False, + "multiple_candidates": False, + }, + } + + +def _runtime_bool(value: Any, default: bool) -> bool: + if value is None: + return default + if isinstance(value, str): + value = value.lower() + if value == "auto": + return default + if value in ("true", "1", "yes"): + return True + if value in ("false", "0", "no"): + return False + return bool(value) + + +def _action_request_to_observation(payload: dict[str, Any]) -> dict[str, Any]: + if "input" not in payload: + return _normalize_observation(payload) + + input_payload = payload.get("input") or {} + observation = dict(input_payload.get("observation") or {}) + if "task" in input_payload: + observation["prompt"] = input_payload["task"] + elif "prompt" in input_payload: + observation["prompt"] = input_payload["prompt"] + if "images" in input_payload: + observation["images"] = input_payload["images"] + if "state" in input_payload: + observation["state"] = input_payload["state"] + if "noise" in input_payload: + observation["noise"] = input_payload["noise"] + return _normalize_observation(observation) + + +@lru_cache(maxsize=32) +def _resolve_action_sampling_params_cls_cached( + model_path: str, + backend: str | None, + model_id: str | None, + pipeline_class_name: str | None, +) -> type[VLASamplingParams]: + if pipeline_class_name: + from sglang.multimodal_gen.registry import get_pipeline_config_classes + + config_classes = get_pipeline_config_classes(pipeline_class_name) + if config_classes is not None: + _, sampling_params_cls = config_classes + if issubclass(sampling_params_cls, VLASamplingParams): + return sampling_params_cls + + from sglang.multimodal_gen.registry import get_model_info + + model_info = get_model_info( + model_path, + backend=backend, + model_id=model_id, + ) + sampling_params_cls = model_info.sampling_param_cls + if not issubclass(sampling_params_cls, VLASamplingParams): + raise ValueError( + f"Action endpoint requires VLASamplingParams, got {sampling_params_cls.__name__}" + ) + return sampling_params_cls + + +def _resolve_action_sampling_params_cls( + server_args: ServerArgs, +) -> type[VLASamplingParams]: + return _resolve_action_sampling_params_cls_cached( + server_args.model_path, + getattr(server_args, "backend", None), + getattr(server_args, "model_id", None), + getattr(server_args, "pipeline_class_name", None), + ) + + +@lru_cache(maxsize=32) +def _sampling_params_field_names( + sampling_params_cls: type[VLASamplingParams], +) -> frozenset[str]: + return frozenset(field.name for field in dataclasses.fields(sampling_params_cls)) + + +def build_action_sampling_params( + payload: dict[str, Any], + server_args: ServerArgs, +) -> VLASamplingParams: + pipeline_config = server_args.pipeline_config + observation = _action_request_to_observation(payload) + parameters = dict(payload.get("parameters") or {}) + runtime = dict(payload.get("runtime") or {}) + if "return_timing" in payload and "return_timing" not in runtime: + runtime["return_timing"] = payload["return_timing"] + images = images_from_observation(observation, pipeline_config) + state = observation.get("state") + if state is None: + state = observation.get("observation.state") + noise = observation.get("noise") + if noise is None: + noise = observation.get("observation.noise") + prompt = observation.get("prompt") or observation.get("task") or "" + prefix_cache = runtime.get("prefix_cache") + if prefix_cache is None: + prefix_cache = observation.get("enable_prefix_cache") + if prefix_cache is None: + prefix_cache = observation.get("enable_pi_prefix_cache") + cuda_graph = runtime.get("cuda_graph") + if cuda_graph is None: + cuda_graph = observation.get("enable_cuda_graph") + if cuda_graph is None: + cuda_graph = observation.get("enable_pi_cuda_graph") + output_format = str( + runtime.get( + "output_format", + parameters.get( + "output_format", + observation.get("output_format", "list"), + ), + ) + ).lower() + if output_format not in ("list", "numpy"): + raise ValueError("output_format must be 'list' or 'numpy'") + + sampling_params_cls = _resolve_action_sampling_params_cls(server_args) + sampling_kwargs = { + "request_id": payload.get("request_id") or payload.get("id"), + "prompt": prompt, + "images": images, + "image_masks": observation.get("image_masks"), + "camera_order": observation.get("camera_order"), + "state": state, + "noise": noise, + "observation": observation, + "action_horizon": int( + parameters.get( + "action_horizon", + observation.get("action_horizon", pipeline_config.action_horizon), + ) + ), + "action_dim": int( + parameters.get( + "action_dim", + observation.get("action_dim", pipeline_config.action_dim), + ) + ), + "num_inference_steps": int( + parameters.get( + "num_inference_steps", + observation.get( + "num_inference_steps", + pipeline_config.default_num_inference_steps, + ), + ) + ), + "output_format": output_format, + "return_timing": _runtime_bool(runtime.get("return_timing"), True), + "enable_prefix_cache": _runtime_bool(prefix_cache, True), + "enable_cuda_graph": _runtime_bool(cuda_graph, True), + } + supported_fields = _sampling_params_field_names(sampling_params_cls) + sp = sampling_params_cls( + **{ + name: value + for name, value in sampling_kwargs.items() + if name in supported_fields + } + ) + sp._adjust(server_args) + return sp + + +async def infer_action( + payload: dict[str, Any], + server_args: ServerArgs, +) -> dict[str, Any]: + sp = build_action_sampling_params(payload, server_args) + req = prepare_request(server_args, sp) + response = await async_scheduler_client.forward(req) + if getattr(response, "error", None): + raise RuntimeError(response.error) + if response.output is None: + raise RuntimeError("action policy returned no output") + return response.output[0] + + +def action_generation_response( + output: dict[str, Any], + server_args: ServerArgs, + *, + preserve_numpy: bool = False, +) -> dict[str, Any]: + actions = output["actions"] + if isinstance(actions, np.ndarray): + action_shape = list(actions.shape) + action_values = actions if preserve_numpy else actions.tolist() + else: + horizon = len(actions) if isinstance(actions, list) else 0 + action_dim = len(actions[0]) if horizon and isinstance(actions[0], list) else 0 + action_shape = [horizon, action_dim] + action_values = actions + response = { + "id": output.get("request_id") or f"act_{uuid.uuid4().hex}", + "object": "action.generation", + "created": int(time.time()), + "model": server_args.model_id or server_args.model_path, + "data": [ + { + "index": 0, + "input_index": 0, + "candidate_index": 0, + "action": { + "type": "continuous", + "dtype": "float32", + "shape": action_shape, + "values": action_values, + }, + } + ], + "usage": { + "action_horizon": action_shape[0] if action_shape else 0, + "action_dim": action_shape[1] if len(action_shape) > 1 else 0, + "denoise_steps": output.get("parameters", {}).get( + "num_inference_steps", + server_args.pipeline_config.default_num_inference_steps, + ), + "prefix_cache_hit": bool(output.get("cache", {}).get("hit", False)), + }, + } + if "timings" in output: + response["timings"] = output["timings"] + if "cache" in output: + response["cache"] = output["cache"] + if "parallel" in output: + response["parallel"] = output["parallel"] + return response + + +def action_raw_response( + output: dict[str, Any], + *, + preserve_numpy: bool = False, +) -> dict[str, Any]: + response = dict(output) + actions = response.get("actions") + if isinstance(actions, np.ndarray) and not preserve_numpy: + response["actions"] = actions.tolist() + return response diff --git a/python/sglang/multimodal_gen/runtime/entrypoints/vla/ws_utils.py b/python/sglang/multimodal_gen/runtime/entrypoints/vla/ws_utils.py new file mode 100644 index 000000000..d6916c0da --- /dev/null +++ b/python/sglang/multimodal_gen/runtime/entrypoints/vla/ws_utils.py @@ -0,0 +1,57 @@ +# SPDX-License-Identifier: Apache-2.0 + +from __future__ import annotations + +import time +import traceback +from collections.abc import Callable +from typing import Any + +from fastapi import WebSocket, WebSocketDisconnect + +from sglang.multimodal_gen.runtime.entrypoints.vla.protocol import ( + action_metadata, + infer_action, + pack_msgpack, + unpack_msgpack, +) +from sglang.multimodal_gen.runtime.server_args import ServerArgs + + +async def run_action_msgpack_ws( + websocket: WebSocket, + server_args: ServerArgs, + *, + prepare_payload: Callable[[dict[str, Any]], None], + build_response: Callable[[dict[str, Any]], dict[str, Any]], +) -> None: + await websocket.accept() + await websocket.send_bytes(pack_msgpack(action_metadata(server_args))) + + prev_total_time = None + while True: + try: + start_time = time.monotonic() + payload = unpack_msgpack(await websocket.receive_bytes()) + prepare_payload(payload) + infer_start = time.monotonic() + output = await infer_action(payload, server_args) + response = build_response(output) + response.setdefault("server_timing", {})["infer_ms"] = ( + time.monotonic() - infer_start + ) * 1000 + if prev_total_time is not None: + response["server_timing"]["prev_total_ms"] = prev_total_time * 1000 + await websocket.send_bytes(pack_msgpack(response)) + prev_total_time = time.monotonic() - start_time + except WebSocketDisconnect: + break + except Exception: + try: + await websocket.send_bytes( + pack_msgpack({"error": traceback.format_exc()}) + ) + except Exception: + pass + await websocket.close(code=1011, reason="Internal server error") + raise diff --git a/python/sglang/multimodal_gen/runtime/launch_server.py b/python/sglang/multimodal_gen/runtime/launch_server.py index 0e4ab4938..e5b5ada7f 100644 --- a/python/sglang/multimodal_gen/runtime/launch_server.py +++ b/python/sglang/multimodal_gen/runtime/launch_server.py @@ -279,6 +279,12 @@ def launch_server(server_args: ServerArgs, launch_http_server: bool = True): logger.debug("All workers are ready") if launch_http_server: + if server_args.pipeline_config.task_type.is_action_gen(): + logger.info( + "VLA pipeline ready: model=%s; per-request details are " + "debug-only (use --log-level debug).", + server_args.model_id or server_args.model_path, + ) logger.info("Starting FastAPI server.") if server_args.webui: logger.info("Launch FastAPI server in another process because of webui.") diff --git a/python/sglang/multimodal_gen/runtime/layers/attention/layer.py b/python/sglang/multimodal_gen/runtime/layers/attention/layer.py index 8c0ced117..67904d1cc 100644 --- a/python/sglang/multimodal_gen/runtime/layers/attention/layer.py +++ b/python/sglang/multimodal_gen/runtime/layers/attention/layer.py @@ -420,6 +420,7 @@ class LocalAttention(nn.Module): softmax_scale: float | None = None, causal: bool = False, supported_attention_backends: set[AttentionBackendEnum] | None = None, + compute_dtype: torch.dtype | None = None, **extra_impl_args, ) -> None: super().__init__() @@ -430,7 +431,7 @@ class LocalAttention(nn.Module): if num_kv_heads is None: num_kv_heads = num_heads - dtype = get_compute_dtype() + dtype = compute_dtype or get_compute_dtype() attn_backend = get_attn_backend( head_size, dtype, supported_attention_backends=supported_attention_backends ) diff --git a/python/sglang/multimodal_gen/runtime/loader/utils.py b/python/sglang/multimodal_gen/runtime/loader/utils.py index 7f4bddcc8..b2ce831f1 100644 --- a/python/sglang/multimodal_gen/runtime/loader/utils.py +++ b/python/sglang/multimodal_gen/runtime/loader/utils.py @@ -179,14 +179,20 @@ class skip_init_modules: def __enter__(self): # Save originals self._orig_reset = {} - for cls in (nn.Linear, nn.Conv1d, nn.Conv2d, nn.Conv3d): + for cls in (nn.Linear, nn.Conv1d, nn.Conv2d, nn.Conv3d, nn.Embedding): self._orig_reset[cls] = cls.reset_parameters cls.reset_parameters = lambda self: None # skip init + from transformers.modeling_utils import PreTrainedModel + + self._pretrained_model_cls = PreTrainedModel + self._orig_post_init = PreTrainedModel.post_init + PreTrainedModel.post_init = lambda self: None def __exit__(self, exc_type, exc_value, traceback): # restore originals for cls, orig in self._orig_reset.items(): cls.reset_parameters = orig + self._pretrained_model_cls.post_init = self._orig_post_init def _normalize_component_type(module_type: str) -> str: diff --git a/python/sglang/multimodal_gen/runtime/loader/weight_utils.py b/python/sglang/multimodal_gen/runtime/loader/weight_utils.py index 5f786eeb8..2d172af6d 100644 --- a/python/sglang/multimodal_gen/runtime/loader/weight_utils.py +++ b/python/sglang/multimodal_gen/runtime/loader/weight_utils.py @@ -9,7 +9,7 @@ import json import os import tempfile from collections import defaultdict -from collections.abc import Generator, Iterable +from collections.abc import Callable, Generator, Iterable from pathlib import Path import filelock @@ -27,6 +27,7 @@ except ImportError: from sglang.multimodal_gen import envs from sglang.multimodal_gen.runtime.distributed import get_local_torch_device +from sglang.multimodal_gen.runtime.loader.weight_load_plan import WeightLoadPlan from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger logger = init_logger(__name__) @@ -183,12 +184,20 @@ def safetensors_weights_iterator( hf_weights_files: list[str], to_cpu: bool = True, use_runai_model_streamer: bool | None = None, + key_filter: Callable[[str], bool] | None = None, + clone_streamed_tensors: bool = True, + weight_load_plan: WeightLoadPlan | None = None, ) -> Generator[tuple[str, torch.Tensor], None, None]: """Iterate over the weights in the model safetensor files.""" enable_tqdm = ( not torch.distributed.is_initialized() or torch.distributed.get_rank() == 0 ) - device = "cpu" if to_cpu else str(get_local_torch_device()) + if weight_load_plan is not None: + checkpoint_device = torch.device(weight_load_plan.checkpoint_load_device) + to_cpu = checkpoint_device.type == "cpu" + device = str(checkpoint_device) + else: + device = "cpu" if to_cpu else str(get_local_torch_device()) if use_runai_model_streamer is None: use_runai_model_streamer = ( HAS_RUNAI_MODEL_STREAMER and envs.SGLANG_USE_RUNAI_MODEL_STREAMER @@ -233,13 +242,24 @@ def safetensors_weights_iterator( _raise_if_duplicate_safetensors_keys(hf_weights_files) if use_runai_model_streamer: + logger.info( + "Loading safetensors with Run:ai Model Streamer to %s", + "cpu" if to_cpu else device, + ) with SafetensorsStreamer() as streamer: - streamer.stream_files(hf_weights_files) + if to_cpu: + streamer.stream_files(hf_weights_files) + else: + streamer.stream_files(hf_weights_files, device=device) for name, tensor in streamer.get_tensors(): + if key_filter is not None and not key_filter(name): + continue if to_cpu: yield name, tensor.clone().detach() + elif clone_streamed_tensors: + yield name, tensor.clone().detach() else: - yield name, tensor.to(device) + yield name, tensor else: for st_file in tqdm( hf_weights_files, @@ -249,6 +269,8 @@ def safetensors_weights_iterator( ): with safe_open(st_file, framework="pt", device=device) as f: for name in f.keys(): # noqa: SIM118 + if key_filter is not None and not key_filter(name): + continue param = f.get_tensor(name) yield name, param diff --git a/python/sglang/multimodal_gen/runtime/managers/gpu_worker.py b/python/sglang/multimodal_gen/runtime/managers/gpu_worker.py index 3f943c92a..cc3168d2f 100644 --- a/python/sglang/multimodal_gen/runtime/managers/gpu_worker.py +++ b/python/sglang/multimodal_gen/runtime/managers/gpu_worker.py @@ -480,7 +480,7 @@ class GPUWorker(GPUWorkerPostTrainingMixin): self._materialize_raw_frame_transport(output_batch, req) elif req.save_output and req.return_file_paths_only: self._materialize_file_path_transport(output_batch, save_output_paths) - elif req.return_frames: + elif getattr(req, "return_frames", False): self._materialize_frame_outputs_for_return(output_batch, req) def _materialize_raw_frame_transport( @@ -518,7 +518,11 @@ class GPUWorker(GPUWorkerPostTrainingMixin): self, output_batch: OutputBatch, req: Req ) -> None: """materialize the output from tensor to numpy frames for faster serialization""" - if self.rank != 0 or output_batch.output is None or not req.return_frames: + if ( + self.rank != 0 + or output_batch.output is None + or not getattr(req, "return_frames", False) + ): return if ( @@ -692,7 +696,7 @@ class GPUWorker(GPUWorkerPostTrainingMixin): mismatched = [ field for field in shared_output_fields - if getattr(req, field) != getattr(first_req, field) + if getattr(req, field, None) != getattr(first_req, field, None) ] if mismatched: raise ValueError( diff --git a/python/sglang/multimodal_gen/runtime/managers/scheduler.py b/python/sglang/multimodal_gen/runtime/managers/scheduler.py index 9f631d3f3..cf64593e5 100644 --- a/python/sglang/multimodal_gen/runtime/managers/scheduler.py +++ b/python/sglang/multimodal_gen/runtime/managers/scheduler.py @@ -282,6 +282,9 @@ class Scheduler(SchedulerWarmupMixin, SchedulerPostTrainingMixin, SchedulerDisag if len(reqs) == 1 or not allow_dynamic_batching: return self.worker.execute_forward(reqs) + if self.server_args.pipeline_config.supports_native_grouped_requests(): + return self._execute_generation_grouped(reqs) + merged_req = self._try_merge_generation_reqs(reqs) if merged_req is None: return self._execute_generation_sequential(reqs) @@ -327,6 +330,48 @@ class Scheduler(SchedulerWarmupMixin, SchedulerPostTrainingMixin, SchedulerDisag error_msg=f"Dynamic batching failed: {e}", ) + def _execute_generation_grouped(self, reqs: List[Req]) -> List[OutputBatch]: + batch_size = len(reqs) + try: + output_batch = self.worker.execute_forward(reqs) + if output_batch.error: + logger.error( + "Native grouped execution returned error. Returning per-request errors: %s", + output_batch.error, + ) + return self._build_dynamic_batch_error_outputs( + reqs=reqs, + error_msg=output_batch.error, + ) + + split_outputs = self._split_batched_output(output_batch, reqs) + if split_outputs is None: + logger.error( + "Failed to split native grouped output cleanly. Returning per-request errors." + ) + return self._build_dynamic_batch_error_outputs( + reqs=reqs, + error_msg="Native grouped execution failed: could not split output.", + ) + + logger.info( + "Processed native grouped batch of %d/%d request(s) with max_delay=%.2fms", + batch_size, + self._batching_max_size, + self._batching_delay_s * 1000.0, + ) + return split_outputs + except Exception as e: + logger.error( + "Native grouped execution failed (%s). Returning per-request errors.", + e, + exc_info=True, + ) + return self._build_dynamic_batch_error_outputs( + reqs=reqs, + error_msg=f"Native grouped execution failed: {e}", + ) + def _execute_generation_sequential(self, reqs: List[Req]) -> List[OutputBatch]: return [self.worker.execute_forward([req]) for req in reqs] @@ -452,7 +497,10 @@ class Scheduler(SchedulerWarmupMixin, SchedulerPostTrainingMixin, SchedulerDisag candidate_req.prompt, str ): return "prompt_type" - if base_req.image_path is not None or candidate_req.image_path is not None: + if ( + getattr(base_req, "image_path", None) is not None + or getattr(candidate_req, "image_path", None) is not None + ): return "image_conditioning" if base_req.return_file_paths_only != candidate_req.return_file_paths_only: return "return_file_paths_only" @@ -486,7 +534,10 @@ class Scheduler(SchedulerWarmupMixin, SchedulerPostTrainingMixin, SchedulerDisag ): return False - if base_req.image_path is not None or candidate_req.image_path is not None: + if ( + getattr(base_req, "image_path", None) is not None + or getattr(candidate_req, "image_path", None) is not None + ): return False if base_req.return_file_paths_only != candidate_req.return_file_paths_only: return False @@ -722,8 +773,14 @@ class Scheduler(SchedulerWarmupMixin, SchedulerPostTrainingMixin, SchedulerDisag outputs: list[OutputBatch] = [] start = 0 - for req, req_count in zip(reqs, per_req_counts): + for req_index, (req, req_count) in enumerate(zip(reqs, per_req_counts)): end = start + req_count + metrics = ( + deepcopy(output_batch.metrics_list[req_index]) + if output_batch.metrics_list is not None + and req_index < len(output_batch.metrics_list) + else deepcopy(output_batch.metrics) + ) split = OutputBatch( output=self._slice_batched_value( output_batch.output, start, end, total_items @@ -748,7 +805,7 @@ class Scheduler(SchedulerWarmupMixin, SchedulerPostTrainingMixin, SchedulerDisag output_file_paths=self._slice_batched_value( output_batch.output_file_paths, start, end, total_items ), - metrics=deepcopy(output_batch.metrics), + metrics=metrics, noise_pred=self._slice_batched_value( output_batch.noise_pred, start, end, total_items ), diff --git a/python/sglang/multimodal_gen/runtime/models/vlas/__init__.py b/python/sglang/multimodal_gen/runtime/models/vlas/__init__.py new file mode 100644 index 000000000..5b0bef068 --- /dev/null +++ b/python/sglang/multimodal_gen/runtime/models/vlas/__init__.py @@ -0,0 +1,5 @@ +# SPDX-License-Identifier: Apache-2.0 + +from sglang.multimodal_gen.runtime.models.vlas.pi05_policy import Pi05PolicyModel + +__all__ = ["Pi05PolicyModel"] diff --git a/python/sglang/multimodal_gen/runtime/models/vlas/pi05_core.py b/python/sglang/multimodal_gen/runtime/models/vlas/pi05_core.py new file mode 100644 index 000000000..e158ce507 --- /dev/null +++ b/python/sglang/multimodal_gen/runtime/models/vlas/pi05_core.py @@ -0,0 +1,1534 @@ +# SPDX-License-Identifier: Apache-2.0 +# Adapted from OpenPI and LeRobot Pi0.5 PyTorch inference semantics. + +from __future__ import annotations + +import math +from dataclasses import dataclass +from typing import Any, Literal + +import torch +import torch.nn.functional as F +from torch import Tensor, nn +from transformers.modeling_outputs import BaseModelOutputWithPooling +from transformers.models.auto import CONFIG_MAPPING +from transformers.models.gemma.modeling_gemma import GemmaConfig +from transformers.models.paligemma.modeling_paligemma import PaliGemmaModel + +from sglang.multimodal_gen.configs.pipeline_configs.pi05 import Pi05PipelineConfig +from sglang.multimodal_gen.runtime.distributed.parallel_state import ( + get_ring_parallel_world_size, + get_sequence_parallel_world_size, + get_ulysses_parallel_world_size, + model_parallel_is_initialized, +) +from sglang.multimodal_gen.runtime.layers.attention import LocalAttention, USPAttention +from sglang.multimodal_gen.runtime.layers.linear import ( + MergedColumnParallelLinear, + QKVParallelLinear, + RowParallelLinear, +) +from sglang.multimodal_gen.runtime.layers.rotary_embedding import RotaryEmbedding +from sglang.multimodal_gen.runtime.managers.forward_context import set_forward_context +from sglang.multimodal_gen.runtime.platforms import AttentionBackendEnum +from sglang.multimodal_gen.runtime.vla.prefix_cache import VLADensePrefixCache +from sglang.srt.layers.activation import GeluAndMul +from sglang.srt.layers.rotary_embedding import ( + apply_rotary_pos_emb as native_apply_rotary_pos_emb, +) + + +def config_compute_dtype(config: GemmaConfig) -> torch.dtype | None: + dtype = getattr(config, "dtype", None) + if dtype is None or isinstance(dtype, torch.dtype): + return dtype + dtype_name = str(dtype).lower() + if dtype_name in ("bf16", "bfloat16", "torch.bfloat16"): + return torch.bfloat16 + if dtype_name in ("fp16", "float16", "half", "torch.float16"): + return torch.float16 + if dtype_name in ("fp32", "float32", "torch.float32"): + return torch.float32 + return None + + +@dataclass +class PiGemmaModelOutput: + last_hidden_state: torch.Tensor + past_key_values: object | None = None + hidden_states: tuple[torch.Tensor, ...] | None = None + attentions: tuple[torch.Tensor, ...] | None = None + + +def gated_residual( + x: torch.Tensor | None, + y: torch.Tensor | None, + gate: torch.Tensor | None, +) -> torch.Tensor | None: + if x is None and y is None: + return None + if x is None or y is None: + return x if x is not None else y + if gate is None: + return x + y + return x + y * gate + + +def layernorm_forward( + layernorm: nn.Module, + x: torch.Tensor, + cond: torch.Tensor | None = None, +) -> tuple[torch.Tensor, torch.Tensor | None]: + if cond is not None: + return layernorm(x, cond=cond) + return layernorm(x) + + +def linear_forward(module: nn.Module, x: torch.Tensor) -> torch.Tensor: + output = module(x) + return output[0] if isinstance(output, tuple) else output + + +def _use_ulysses_action_attention(num_heads: int) -> bool: + if not model_parallel_is_initialized(): + return False + try: + sp_world_size = get_sequence_parallel_world_size() + ulysses_world_size = get_ulysses_parallel_world_size() + ring_world_size = get_ring_parallel_world_size() + except AssertionError: + return False + return ( + sp_world_size > 1 + and ulysses_world_size > 1 + and ring_world_size == 1 + and num_heads % sp_world_size == 0 + ) + + +class Pi05SiglipAttention(nn.Module): + def __init__(self, attention: nn.Module): + super().__init__() + self.embed_dim = attention.embed_dim + self.num_heads = attention.num_heads + self.head_dim = attention.head_dim + self.scale = getattr(attention, "scale", self.head_dim**-0.5) + self.dropout = getattr(attention, "dropout", 0.0) + self.q_proj = attention.q_proj + self.k_proj = attention.k_proj + self.v_proj = attention.v_proj + self.out_proj = attention.out_proj + self.attn = LocalAttention( + num_heads=self.num_heads, + head_size=self.head_dim, + num_kv_heads=self.num_heads, + softmax_scale=self.scale, + causal=False, + supported_attention_backends={ + AttentionBackendEnum.FA, + AttentionBackendEnum.FA2, + AttentionBackendEnum.TORCH_SDPA, + }, + compute_dtype=self.q_proj.weight.dtype, + allow_cudnn_sdp=True, + ) + + def forward( + self, + hidden_states: torch.Tensor, + attention_mask: torch.Tensor | None = None, + output_attentions: bool = False, + **kwargs, + ) -> tuple[torch.Tensor, None]: + input_shape = hidden_states.shape[:-1] + query_states = self.q_proj(hidden_states).view( + *input_shape, + self.num_heads, + self.head_dim, + ) + key_states = self.k_proj(hidden_states).view( + *input_shape, + self.num_heads, + self.head_dim, + ) + value_states = self.v_proj(hidden_states).view( + *input_shape, + self.num_heads, + self.head_dim, + ) + attn_output = self.attn( + query_states, + key_states, + value_states, + attn_mask=attention_mask, + ) + attn_output = attn_output.reshape(*input_shape, self.embed_dim).contiguous() + return self.out_proj(attn_output), None + + +def patch_siglip_vision_attention_to_native(vision_model: nn.Module) -> None: + for layer in vision_model.encoder.layers: + if isinstance(layer.self_attn, Pi05SiglipAttention): + continue + layer.self_attn = Pi05SiglipAttention(layer.self_attn) + + +class PiGemmaRMSNorm(nn.Module): + def __init__(self, dim: int, eps: float = 1e-6, cond_dim: int | None = None): + super().__init__() + self.eps = eps + self.dim = dim + self.cond_dim = cond_dim + if cond_dim is None: + self.weight = nn.Parameter(torch.zeros(dim)) + self.dense = None + else: + self.dense = nn.Linear(cond_dim, dim * 3, bias=True) + nn.init.zeros_(self.dense.weight) + + def forward( + self, + x: torch.Tensor, + cond: torch.Tensor | None = None, + ) -> tuple[torch.Tensor, torch.Tensor | None]: + dtype = x.dtype + variance = torch.mean(torch.square(x.float()), dim=-1, keepdim=True) + normed = x * torch.rsqrt(variance + self.eps) + if cond is None: + if self.dense is not None: + return normed.type_as(x), None + normed = normed * (1.0 + self.weight.float()) + return normed.type_as(x), None + if self.dense is None: + normed = normed * (1.0 + self.weight.float()) + return normed.type_as(x), None + + modulation = self.dense(cond.to(dtype=self.dense.weight.dtype)) + if x.ndim == 3: + modulation = modulation.unsqueeze(1) + scale, shift, gate = modulation.chunk(3, dim=-1) + normed = normed * (1.0 + scale.float()) + shift.float() + return normed.to(dtype), gate.to(dtype) + + +class PiGemmaMLP(nn.Module): + def __init__(self, config: GemmaConfig, *, tensor_parallel: bool = False): + super().__init__() + self.hidden_size = config.hidden_size + self.intermediate_size = config.intermediate_size + self.tensor_parallel = tensor_parallel + if tensor_parallel: + self.gate_up_proj = MergedColumnParallelLinear( + input_size=self.hidden_size, + output_sizes=[self.intermediate_size] * 2, + bias=False, + ) + self.down_proj = RowParallelLinear( + input_size=self.intermediate_size, + output_size=self.hidden_size, + bias=False, + ) + self.gate_proj = None + self.up_proj = None + else: + self.gate_proj = nn.Linear( + self.hidden_size, self.intermediate_size, bias=False + ) + self.up_proj = nn.Linear( + self.hidden_size, self.intermediate_size, bias=False + ) + self.down_proj = nn.Linear( + self.intermediate_size, self.hidden_size, bias=False + ) + self.gate_up_proj = None + if config.hidden_act != "gelu_pytorch_tanh": + raise ValueError(f"Unsupported PiGemma activation: {config.hidden_act}") + self.act_fn = GeluAndMul(approximate="tanh") + + @property + def projection_dtype(self) -> torch.dtype: + if self.tensor_parallel: + return self.gate_up_proj.weight.dtype + return self.up_proj.weight.dtype + + def forward(self, x: torch.Tensor) -> torch.Tensor: + if self.tensor_parallel: + gate_up = linear_forward(self.gate_up_proj, x) + else: + gate_up = torch.cat([self.gate_proj(x), self.up_proj(x)], dim=-1) + return linear_forward(self.down_proj, self.act_fn(gate_up)) + + +class PiGemmaRotaryEmbedding(nn.Module): + def __init__(self, config: GemmaConfig): + super().__init__() + self.max_seq_len_cached = config.max_position_embeddings + self.head_dim = getattr( + config, + "head_dim", + config.hidden_size // config.num_attention_heads, + ) + rope_parameters = config.rope_parameters + if rope_parameters["rope_type"] != "default": + raise ValueError( + f"Unsupported PiGemma rope type: {rope_parameters['rope_type']}" + ) + self.rope = RotaryEmbedding( + head_size=self.head_dim, + rotary_dim=self.head_dim, + max_position_embeddings=self.max_seq_len_cached, + base=rope_parameters["rope_theta"], + is_neox_style=True, + dtype=torch.float32, + ) + + def _ensure_cache(self, x: torch.Tensor) -> None: + cache = self.rope.cos_sin_cache + if cache.device == x.device and cache.dtype == torch.float: + return + if cache.dtype != torch.float: + cache = self.rope._compute_cos_sin_cache() + self.rope.cos_sin_cache = cache.to(device=x.device, dtype=torch.float) + + def forward( + self, + x: torch.Tensor, + position_ids: torch.Tensor, + ) -> tuple[torch.Tensor, torch.Tensor]: + self._ensure_cache(x) + flat_positions = position_ids.reshape(-1) + cos_sin = self.rope.cos_sin_cache.index_select(0, flat_positions) + cos_half, sin_half = cos_sin.chunk(2, dim=-1) + cos = torch.cat((cos_half, cos_half), dim=-1) + sin = torch.cat((sin_half, sin_half), dim=-1) + output_shape = (*position_ids.shape, self.head_dim) + return cos.reshape(output_shape).to(x.dtype), sin.reshape(output_shape).to( + x.dtype + ) + + +class PiGemmaAttention(nn.Module): + def __init__( + self, + config: GemmaConfig, + layer_idx: int, + *, + tensor_parallel: bool = False, + ): + super().__init__() + self.config = config + self.layer_idx = layer_idx + self.tensor_parallel = tensor_parallel + self.head_dim = getattr( + config, + "head_dim", + config.hidden_size // config.num_attention_heads, + ) + self.scaling = self.head_dim**-0.5 + self.attention_dropout = config.attention_dropout + self.is_causal = not getattr(config, "use_bidirectional_attention", False) + + if tensor_parallel: + self.total_num_heads = config.num_attention_heads + self.total_num_key_value_heads = config.num_key_value_heads + self.qkv_proj = QKVParallelLinear( + hidden_size=config.hidden_size, + head_size=self.head_dim, + total_num_heads=self.total_num_heads, + total_num_kv_heads=self.total_num_key_value_heads, + bias=config.attention_bias, + ) + self.num_heads = self.qkv_proj.num_heads + self.num_key_value_heads = self.qkv_proj.num_kv_heads + self.q_size = self.num_heads * self.head_dim + self.kv_size = self.num_key_value_heads * self.head_dim + self.o_proj = RowParallelLinear( + input_size=self.total_num_heads * self.head_dim, + output_size=config.hidden_size, + bias=config.attention_bias, + ) + self.q_proj = None + self.k_proj = None + self.v_proj = None + else: + self.num_heads = config.num_attention_heads + self.num_key_value_heads = config.num_key_value_heads + self.q_proj = nn.Linear( + config.hidden_size, + self.num_heads * self.head_dim, + bias=config.attention_bias, + ) + self.k_proj = nn.Linear( + config.hidden_size, + self.num_key_value_heads * self.head_dim, + bias=config.attention_bias, + ) + self.v_proj = nn.Linear( + config.hidden_size, + self.num_key_value_heads * self.head_dim, + bias=config.attention_bias, + ) + self.o_proj = nn.Linear( + self.num_heads * self.head_dim, + config.hidden_size, + bias=config.attention_bias, + ) + self.qkv_proj = None + self.q_size = self.num_heads * self.head_dim + self.kv_size = self.num_key_value_heads * self.head_dim + self.num_key_value_groups = self.num_heads // self.num_key_value_heads + self.attn = LocalAttention( + num_heads=self.num_heads, + head_size=self.head_dim, + num_kv_heads=self.num_key_value_heads, + softmax_scale=self.scaling, + causal=self.is_causal, + supported_attention_backends={ + AttentionBackendEnum.FA, + AttentionBackendEnum.FA2, + AttentionBackendEnum.TORCH_SDPA, + }, + compute_dtype=config_compute_dtype(config), + allow_cudnn_sdp=True, + ) + self.sp_attn = ( + USPAttention( + num_heads=config.num_attention_heads, + head_size=self.head_dim, + num_kv_heads=config.num_attention_heads, + softmax_scale=self.scaling, + causal=self.is_causal, + supported_attention_backends={ + AttentionBackendEnum.FA, + AttentionBackendEnum.FA2, + AttentionBackendEnum.TORCH_SDPA, + }, + allow_cudnn_sdp=True, + ) + if _use_ulysses_action_attention(config.num_attention_heads) + and not tensor_parallel + else None + ) + + @property + def projection_dtype(self) -> torch.dtype: + if self.tensor_parallel: + return self.qkv_proj.weight.dtype + return self.q_proj.weight.dtype + + def project_qkv( + self, + hidden_states: torch.Tensor, + ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]: + input_shape = hidden_states.shape[:-1] + query_shape = (*input_shape, self.num_heads, self.head_dim) + kv_shape = (*input_shape, self.num_key_value_heads, self.head_dim) + + if self.tensor_parallel: + qkv = linear_forward(self.qkv_proj, hidden_states) + query_states, key_states, value_states = qkv.split( + [self.q_size, self.kv_size, self.kv_size], + dim=-1, + ) + else: + query_states = self.q_proj(hidden_states) + key_states = self.k_proj(hidden_states) + value_states = self.v_proj(hidden_states) + return ( + query_states.view(query_shape).transpose(1, 2), + key_states.view(kv_shape).transpose(1, 2), + value_states.view(kv_shape).transpose(1, 2), + ) + + def _repeat_kv_for_sequence_parallel( + self, + states: torch.Tensor, + ) -> torch.Tensor: + if self.num_key_value_groups == 1: + return states + return states.repeat_interleave(self.num_key_value_groups, dim=2) + + def forward( + self, + hidden_states: torch.Tensor, + position_embeddings: tuple[torch.Tensor, torch.Tensor] | None = None, + attention_mask: torch.Tensor | None = None, + past_key_values: VLADensePrefixCache | None = None, + **kwargs, + ) -> tuple[torch.Tensor, None]: + input_shape = hidden_states.shape[:-1] + query_states, key_states, value_states = self.project_qkv(hidden_states) + + cos, sin = position_embeddings + query_states, key_states = native_apply_rotary_pos_emb( + query_states, + key_states, + cos, + sin, + unsqueeze_dim=1, + ) + + if past_key_values is not None: + if ( + self.sp_attn is not None + and past_key_values.read_only + and attention_mask is None + ): + prefix_key_states, prefix_value_states = past_key_values.get_prefix( + self.layer_idx + ) + attn_output = self.sp_attn.forward_with_replicated_kv_prefix( + query_states.transpose(1, 2), + self._repeat_kv_for_sequence_parallel( + prefix_key_states.transpose(1, 2) + ), + self._repeat_kv_for_sequence_parallel( + prefix_value_states.transpose(1, 2) + ), + self._repeat_kv_for_sequence_parallel(key_states.transpose(1, 2)), + self._repeat_kv_for_sequence_parallel(value_states.transpose(1, 2)), + ) + attn_output = attn_output.reshape(*input_shape, -1).contiguous() + return linear_forward(self.o_proj, attn_output), None + + key_states, value_states = past_key_values.update( + key_states, + value_states, + self.layer_idx, + ) + + attn_output = self.attn( + query_states.transpose(1, 2), + key_states.transpose(1, 2), + value_states.transpose(1, 2), + attn_mask=attention_mask, + ) + attn_output = attn_output.reshape(*input_shape, -1).contiguous() + return linear_forward(self.o_proj, attn_output), None + + +class PiGemmaDecoderLayer(nn.Module): + def __init__( + self, + config: GemmaConfig, + layer_idx: int, + *, + tensor_parallel: bool = False, + ): + super().__init__() + self.hidden_size = config.hidden_size + self.self_attn = PiGemmaAttention( + config=config, + layer_idx=layer_idx, + tensor_parallel=tensor_parallel, + ) + self.mlp = PiGemmaMLP(config, tensor_parallel=tensor_parallel) + cond_dim = ( + getattr(config, "adarms_cond_dim", None) + if getattr(config, "use_adarms", False) + else None + ) + self.input_layernorm = PiGemmaRMSNorm( + config.hidden_size, eps=config.rms_norm_eps, cond_dim=cond_dim + ) + self.post_attention_layernorm = PiGemmaRMSNorm( + config.hidden_size, eps=config.rms_norm_eps, cond_dim=cond_dim + ) + + def forward( + self, + hidden_states: torch.Tensor, + attention_mask: torch.Tensor | None = None, + past_key_values=None, + position_embeddings: tuple[torch.Tensor, torch.Tensor] | None = None, + adarms_cond: torch.Tensor | None = None, + **kwargs, + ) -> torch.Tensor: + residual = hidden_states + hidden_states, gate = self.input_layernorm(hidden_states, cond=adarms_cond) + hidden_states, _ = self.self_attn( + hidden_states, + attention_mask=attention_mask, + past_key_values=past_key_values, + position_embeddings=position_embeddings, + **kwargs, + ) + hidden_states = gated_residual(residual, hidden_states, gate) + + residual = hidden_states + hidden_states, gate = self.post_attention_layernorm( + hidden_states, cond=adarms_cond + ) + hidden_states = self.mlp(hidden_states) + hidden_states = gated_residual(residual, hidden_states, gate) + return hidden_states + + +class PiGemmaModel(nn.Module): + def __init__( + self, + config: GemmaConfig, + *, + tensor_parallel: bool = False, + **kwargs, + ): + super().__init__() + self.config = config + self.tensor_parallel = tensor_parallel + self.padding_idx = config.pad_token_id + self.vocab_size = config.vocab_size + self.embed_tokens = nn.Embedding( + config.vocab_size, + config.hidden_size, + self.padding_idx, + ) + self.layerwise_cpu_offload_enabled = False + self.layerwise_cpu_offload_device: torch.device | None = None + self.layerwise_cpu_offload_empty_cache = True + cond_dim = getattr(config, "adarms_cond_dim", None) + self.layers = nn.ModuleList( + [ + PiGemmaDecoderLayer( + config, + layer_idx, + tensor_parallel=tensor_parallel, + ) + for layer_idx in range(config.num_hidden_layers) + ] + ) + self.norm = PiGemmaRMSNorm( + config.hidden_size, eps=config.rms_norm_eps, cond_dim=cond_dim + ) + self.rotary_emb = PiGemmaRotaryEmbedding(config=config) + self.gradient_checkpointing = False + + def get_input_embeddings(self) -> nn.Module: + return self.embed_tokens + + def configure_layerwise_cpu_offload( + self, + *, + compute_device: torch.device, + empty_cache: bool = True, + ) -> None: + self.layerwise_cpu_offload_enabled = True + self.layerwise_cpu_offload_device = torch.device(compute_device) + self.layerwise_cpu_offload_empty_cache = empty_cache + + def _maybe_move_inputs_for_layerwise_offload( + self, + inputs_embeds: torch.Tensor, + attention_mask: torch.Tensor | None, + position_ids: torch.LongTensor | None, + cache_position: torch.LongTensor | None, + adarms_cond: torch.Tensor | None, + ) -> tuple[ + torch.Tensor, + torch.Tensor | None, + torch.LongTensor | None, + torch.LongTensor | None, + torch.Tensor | None, + ]: + if not self.layerwise_cpu_offload_enabled: + return ( + inputs_embeds, + attention_mask, + position_ids, + cache_position, + adarms_cond, + ) + if self.layerwise_cpu_offload_device is None: + raise RuntimeError("PiGemma layerwise CPU offload compute device is unset") + + device = self.layerwise_cpu_offload_device + if inputs_embeds.device != device: + inputs_embeds = inputs_embeds.to(device=device) + if attention_mask is not None and attention_mask.device != device: + attention_mask = attention_mask.to(device=device) + if position_ids is not None and position_ids.device != device: + position_ids = position_ids.to(device=device) + if cache_position is not None and cache_position.device != device: + cache_position = cache_position.to(device=device) + if adarms_cond is not None and adarms_cond.device != device: + adarms_cond = adarms_cond.to(device=device) + return inputs_embeds, attention_mask, position_ids, cache_position, adarms_cond + + def _layer_to_compute_device(self, decoder_layer: nn.Module) -> None: + if not self.layerwise_cpu_offload_enabled: + return + decoder_layer.to(self.layerwise_cpu_offload_device) + + def _layer_to_cpu_after_compute(self, decoder_layer: nn.Module) -> None: + if not self.layerwise_cpu_offload_enabled: + return + decoder_layer.to("cpu") + device = self.layerwise_cpu_offload_device + if ( + self.layerwise_cpu_offload_empty_cache + and device is not None + and device.type == "cuda" + ): + torch.cuda.empty_cache() + + def forward( + self, + input_ids: torch.LongTensor | None = None, + attention_mask: torch.Tensor | None = None, + position_ids: torch.LongTensor | None = None, + past_key_values: Any | None = None, + inputs_embeds: torch.FloatTensor | None = None, + use_cache: bool | None = None, + output_attentions: bool | None = None, + output_hidden_states: bool | None = None, + cache_position: torch.LongTensor | None = None, + adarms_cond: torch.Tensor | None = None, + **kwargs, + ) -> PiGemmaModelOutput: + output_attentions = ( + output_attentions + if output_attentions is not None + else self.config.output_attentions + ) + output_hidden_states = ( + output_hidden_states + if output_hidden_states is not None + else self.config.output_hidden_states + ) + use_cache = use_cache if use_cache is not None else self.config.use_cache + + if (input_ids is None) == (inputs_embeds is None): + raise ValueError("Specify exactly one of input_ids or inputs_embeds") + + if inputs_embeds is None: + inputs_embeds = self.embed_tokens(input_ids) + ( + inputs_embeds, + attention_mask, + position_ids, + cache_position, + adarms_cond, + ) = self._maybe_move_inputs_for_layerwise_offload( + inputs_embeds, + attention_mask, + position_ids, + cache_position, + adarms_cond, + ) + + if use_cache and past_key_values is None: + past_key_values = VLADensePrefixCache() + + if cache_position is None: + past_seen_tokens = ( + past_key_values.get_seq_length() if past_key_values is not None else 0 + ) + cache_position = torch.arange( + past_seen_tokens, + past_seen_tokens + inputs_embeds.shape[1], + device=inputs_embeds.device, + ) + + if position_ids is None: + position_ids = cache_position.unsqueeze(0) + + causal_mask = attention_mask + + hidden_states = inputs_embeds + if ( + len(self.layers) > 0 + and self.layers[0].self_attn.projection_dtype == torch.bfloat16 + ): + hidden_states = hidden_states.to(torch.bfloat16) + + position_embeddings = self.rotary_emb(hidden_states, position_ids) + all_hidden_states = () if output_hidden_states else None + all_self_attns = () if output_attentions else None + + for decoder_layer in self.layers[: self.config.num_hidden_layers]: + if output_hidden_states: + all_hidden_states += (hidden_states,) + self._layer_to_compute_device(decoder_layer) + layer_outputs = decoder_layer( + hidden_states, + attention_mask=causal_mask, + past_key_values=past_key_values, + output_attentions=output_attentions, + position_embeddings=position_embeddings, + adarms_cond=adarms_cond, + **kwargs, + ) + hidden_states = layer_outputs + self._layer_to_cpu_after_compute(decoder_layer) + if output_attentions: + all_self_attns += (layer_outputs[1],) + + hidden_states, _ = self.norm(hidden_states, adarms_cond) + if output_hidden_states: + all_hidden_states += (hidden_states,) + + return PiGemmaModelOutput( + last_hidden_state=hidden_states, + past_key_values=past_key_values if use_cache else None, + hidden_states=all_hidden_states, + attentions=all_self_attns, + ) + + +class PiGemmaForCausalLM(nn.Module): + def __init__( + self, + config: GemmaConfig, + *, + tensor_parallel: bool = False, + **kwargs, + ): + super().__init__() + self.config = config + self.model = PiGemmaModel(config, tensor_parallel=tensor_parallel) + self.lm_head = None + + +class PaliGemmaModelWithPiGemma(PaliGemmaModel): + def __init__(self, config, *, tensor_parallel: bool = False): + super().__init__(config) + del self.language_model + self.language_model = PiGemmaModel( + config.text_config, + tensor_parallel=tensor_parallel, + ) + + +class PaliGemmaForConditionalGenerationWithPiGemma(nn.Module): + def __init__(self, config, *, tensor_parallel: bool = False): + super().__init__() + self.config = config + self.model = PaliGemmaModelWithPiGemma( + config, + tensor_parallel=tensor_parallel, + ) + self.lm_head = None + + @property + def language_model(self): + return self.model.language_model + + +OPENPI_ATTENTION_MASK_VALUE = -2.3819763e38 +PALIGEMMA_VOCAB_SIZE = 257_152 + + +@dataclass(frozen=True) +class GemmaVariantConfig: + width: int + depth: int + mlp_dim: int + num_heads: int + num_kv_heads: int + head_dim: int + + +def get_gemma_variant_config(variant: str) -> GemmaVariantConfig: + if variant == "gemma_300m": + return GemmaVariantConfig( + width=1024, + depth=18, + mlp_dim=4096, + num_heads=8, + num_kv_heads=1, + head_dim=256, + ) + if variant == "gemma_2b": + return GemmaVariantConfig( + width=2048, + depth=18, + mlp_dim=16_384, + num_heads=8, + num_kv_heads=1, + head_dim=256, + ) + raise ValueError(f"Unknown Pi05 Gemma variant: {variant}") + + +def create_sinusoidal_pos_embedding( + time: torch.Tensor, + dimension: int, + min_period: float, + max_period: float, +) -> Tensor: + if dimension % 2 != 0: + raise ValueError(f"dimension ({dimension}) must be divisible by 2") + if time.ndim != 1: + raise ValueError("time must have shape [batch]") + fraction = torch.linspace( + 0.0, + 1.0, + dimension // 2, + dtype=torch.float64, + device=time.device, + ) + period = min_period * (max_period / min_period) ** fraction + scaling = 1.0 / period * 2 * math.pi + sin_input = scaling[None, :] * time[:, None].to(torch.float64) + return torch.cat([torch.sin(sin_input), torch.cos(sin_input)], dim=1) + + +def make_att_2d_masks( + pad_masks: torch.Tensor, + att_masks: torch.Tensor, +) -> torch.Tensor: + if att_masks.ndim != 2 or pad_masks.ndim != 2: + raise ValueError("pad_masks and att_masks must be [batch, seq]") + cumsum = torch.cumsum(att_masks, dim=1) + att_2d_masks = cumsum[:, None, :] <= cumsum[:, :, None] + pad_2d_masks = pad_masks[:, None, :] * pad_masks[:, :, None] + return att_2d_masks & pad_2d_masks + + +def trim_trailing_padding_tokens( + tokens: torch.Tensor, + token_masks: torch.Tensor, +) -> tuple[torch.Tensor, torch.Tensor]: + token_len = int(token_masks.sum(dim=1).max().item()) + if token_len <= 0 or token_len >= tokens.shape[1]: + return tokens, token_masks + return tokens[:, :token_len], token_masks[:, :token_len] + + +def prepare_optional_full_attention_mask( + att_2d_masks: torch.Tensor, + *, + full_attention: bool | None = None, +) -> torch.Tensor | None: + if full_attention is None: + full_attention = bool(att_2d_masks.all().item()) + if full_attention: + return None + masks_4d = att_2d_masks[:, None, :, :] + return torch.where(masks_4d, 0.0, OPENPI_ATTENTION_MASK_VALUE) + + +def siglip_vision_forward_with_openpi_dtype( + self, + pixel_values, + interpolate_pos_encoding: bool | None = False, + **kwargs, +) -> BaseModelOutputWithPooling: + hidden_states = self.embeddings( + pixel_values, + interpolate_pos_encoding=interpolate_pos_encoding, + ) + if ( + len(self.encoder.layers) > 0 + and self.encoder.layers[0].self_attn.q_proj.weight.dtype == torch.bfloat16 + ): + hidden_states = hidden_states.to(torch.bfloat16) + + encoder_outputs = self.encoder(inputs_embeds=hidden_states, **kwargs) + last_hidden_state = encoder_outputs.last_hidden_state + last_hidden_state = self.post_layernorm(last_hidden_state) + pooler_output = self.head(last_hidden_state) if self.use_head else None + return BaseModelOutputWithPooling( + last_hidden_state=last_hidden_state, + pooler_output=pooler_output, + hidden_states=encoder_outputs.hidden_states, + attentions=encoder_outputs.attentions, + ) + + +def compute_layer_complete( + inputs_embeds, + attention_mask, + position_ids, + adarms_cond, + *, + layers, + rotary_emb, +): + query_states = [] + key_states = [] + value_states = [] + gates = [] + for i, hidden_states in enumerate(inputs_embeds): + layer = layers[i] + hidden_states, gate = layernorm_forward( + layer.input_layernorm, hidden_states, adarms_cond[i] + ) + gates.append(gate) + query_state, key_state, value_state = layer.self_attn.project_qkv(hidden_states) + query_states.append(query_state) + key_states.append(key_state) + value_states.append(value_state) + + query_states = torch.cat(query_states, dim=2) + key_states = torch.cat(key_states, dim=2) + value_states = torch.cat(value_states, dim=2) + dummy_tensor = torch.zeros( + query_states.shape[0], + query_states.shape[2], + query_states.shape[-1], + device=query_states.device, + dtype=query_states.dtype, + ) + cos, sin = rotary_emb(dummy_tensor, position_ids) + query_states, key_states = native_apply_rotary_pos_emb( + query_states, + key_states, + cos, + sin, + unsqueeze_dim=1, + ) + paligemma_layer = layers[0] + att_output = paligemma_layer.self_attn.attn( + query_states.transpose(1, 2), + key_states.transpose(1, 2), + value_states.transpose(1, 2), + attn_mask=attention_mask, + ) + batch_size = query_states.shape[0] + head_dim = paligemma_layer.self_attn.head_dim + hidden_size = paligemma_layer.self_attn.num_heads * head_dim + att_output = att_output.reshape(batch_size, -1, hidden_size) + + outputs_embeds = [] + start_pos = 0 + for i, hidden_states in enumerate(inputs_embeds): + layer = layers[i] + end_pos = start_pos + hidden_states.shape[1] + if att_output.dtype != layer.self_attn.projection_dtype: + att_output = att_output.to(layer.self_attn.projection_dtype) + out_emb = linear_forward( + layer.self_attn.o_proj, + att_output[:, start_pos:end_pos], + ) + out_emb = gated_residual(hidden_states, out_emb, gates[i]) + after_first_residual = out_emb.clone() + out_emb, gate = layernorm_forward( + layer.post_attention_layernorm, out_emb, adarms_cond[i] + ) + if layer.mlp.projection_dtype == torch.bfloat16: + out_emb = out_emb.to(dtype=torch.bfloat16) + out_emb = layer.mlp(out_emb) + out_emb = gated_residual(after_first_residual, out_emb, gate) + outputs_embeds.append(out_emb) + start_pos = end_pos + return outputs_embeds + + +class PaliGemmaWithExpertModel(nn.Module): + def __init__( + self, + vlm_config: GemmaVariantConfig, + action_expert_config: GemmaVariantConfig, + *, + use_adarms: list[bool], + precision: Literal["bfloat16", "float32"], + image_size: int, + runtime_role: Literal["all", "prefix", "action", "idle"] = "all", + prefix_tensor_parallel: bool = False, + ): + super().__init__() + self.paligemma = None + self.gemma_expert = None + self.prefix_output_device: torch.device | None = None + + if runtime_role in ("all", "prefix"): + vlm_config_hf = CONFIG_MAPPING["paligemma"]() + vlm_config_hf._vocab_size = PALIGEMMA_VOCAB_SIZE + vlm_config_hf.image_token_index = PALIGEMMA_VOCAB_SIZE + vlm_config_hf.text_config.hidden_size = vlm_config.width + vlm_config_hf.text_config.intermediate_size = vlm_config.mlp_dim + vlm_config_hf.text_config.num_attention_heads = vlm_config.num_heads + vlm_config_hf.text_config.head_dim = vlm_config.head_dim + vlm_config_hf.text_config.num_hidden_layers = vlm_config.depth + vlm_config_hf.text_config.num_key_value_heads = vlm_config.num_kv_heads + vlm_config_hf.text_config.hidden_activation = "gelu_pytorch_tanh" + vlm_config_hf.text_config.dtype = precision + vlm_config_hf.text_config.vocab_size = PALIGEMMA_VOCAB_SIZE + vlm_config_hf.text_config.use_adarms = use_adarms[0] + vlm_config_hf.text_config.is_causal = False + vlm_config_hf.text_config.use_bidirectional_attention = True + vlm_config_hf.text_config.adarms_cond_dim = ( + vlm_config.width if use_adarms[0] else None + ) + vlm_config_hf.vision_config.image_size = image_size + vlm_config_hf.vision_config.intermediate_size = 4304 + vlm_config_hf.vision_config.projection_dim = 2048 + vlm_config_hf.vision_config.projector_hidden_act = "gelu_fast" + vlm_config_hf.vision_config.dtype = "float32" + self.paligemma = PaliGemmaForConditionalGenerationWithPiGemma( + config=vlm_config_hf, + tensor_parallel=prefix_tensor_parallel, + ) + vision_tower = self.paligemma.model.vision_tower + vision_model = getattr(vision_tower, "vision_model", vision_tower) + vision_model.forward = siglip_vision_forward_with_openpi_dtype.__get__( + vision_model, + type(vision_model), + ) + self.paligemma.lm_head = None + + if runtime_role in ("all", "action"): + action_config_hf = CONFIG_MAPPING["gemma"]( + head_dim=action_expert_config.head_dim, + hidden_size=action_expert_config.width, + intermediate_size=action_expert_config.mlp_dim, + num_attention_heads=action_expert_config.num_heads, + num_hidden_layers=action_expert_config.depth, + num_key_value_heads=action_expert_config.num_kv_heads, + vocab_size=PALIGEMMA_VOCAB_SIZE, + hidden_activation="gelu_pytorch_tanh", + dtype=precision, + use_adarms=use_adarms[1], + is_causal=False, + use_bidirectional_attention=True, + adarms_cond_dim=(action_expert_config.width if use_adarms[1] else None), + ) + self.gemma_expert = PiGemmaForCausalLM( + config=action_config_hf, + tensor_parallel=False, + ) + self.gemma_expert.lm_head = None + self.gemma_expert.model.embed_tokens = None + self.to_selected_dtype(precision) + self.patch_native_attention_after_dtype_finalize() + + def patch_native_attention_after_dtype_finalize(self) -> None: + if self.paligemma is None: + return + vision_tower = self.paligemma.model.vision_tower + vision_model = getattr(vision_tower, "vision_model", vision_tower) + patch_siglip_vision_attention_to_native(vision_model) + + def to_selected_dtype( + self, precision: Literal["bfloat16", "float32"] = "bfloat16" + ) -> None: + if precision == "float32": + self.to(dtype=torch.float32) + return + if precision != "bfloat16": + raise ValueError(f"Invalid Pi05 precision: {precision}") + self.to(dtype=torch.bfloat16) + keep_fp32 = [ + "vision_tower.embeddings.patch_embedding.weight", + "vision_tower.embeddings.patch_embedding.bias", + "vision_tower.embeddings.position_embedding.weight", + "vision_tower.vision_model.embeddings.patch_embedding.weight", + "vision_tower.vision_model.embeddings.patch_embedding.bias", + "vision_tower.vision_model.embeddings.position_embedding.weight", + "input_layernorm", + "post_attention_layernorm", + "model.norm", + ] + for name, param in self.named_parameters(): + if any(selector in name for selector in keep_fp32): + param.data = param.data.to(dtype=torch.float32) + + def set_prefix_output_device(self, device: torch.device) -> None: + self.prefix_output_device = torch.device(device) + + @staticmethod + def _module_device(module: nn.Module) -> torch.device: + return next(module.parameters()).device + + def _prefix_transformer_device(self) -> torch.device: + if self.prefix_output_device is not None: + return self.prefix_output_device + language_model = self.paligemma.model.language_model + return self._module_device(language_model.layers[0]) + + def embed_image(self, image: torch.Tensor) -> torch.Tensor: + out_dtype = image.dtype + vision_device = self._module_device(self.paligemma.model.vision_tower) + output_device = self._prefix_transformer_device() + if image.device != vision_device or image.dtype != torch.float32: + image = image.to(device=vision_device, dtype=torch.float32) + with set_forward_context(current_timestep=0, attn_metadata=None): + image_outputs = self.paligemma.model.get_image_features(image) + features = image_outputs.pooler_output + if features.device != output_device or features.dtype != out_dtype: + features = features.to(device=output_device, dtype=out_dtype) + return features + + def embed_images(self, images: list[torch.Tensor]) -> list[torch.Tensor]: + if len(images) == 1: + return [self.embed_image(images[0])] + out_dtype = images[0].dtype + vision_device = self._module_device(self.paligemma.model.vision_tower) + output_device = self._prefix_transformer_device() + batch_sizes = [image.shape[0] for image in images] + batched_images = torch.cat( + [ + ( + image.to(device=vision_device, dtype=torch.float32) + if image.device != vision_device or image.dtype != torch.float32 + else image + ) + for image in images + ], + dim=0, + ) + with set_forward_context(current_timestep=0, attn_metadata=None): + image_outputs = self.paligemma.model.get_image_features(batched_images) + features = image_outputs.pooler_output + if features.device != output_device or features.dtype != out_dtype: + features = features.to(device=output_device, dtype=out_dtype) + return list(features.split(batch_sizes, dim=0)) + + def embed_language_tokens(self, tokens: torch.Tensor) -> torch.Tensor: + embedding = self.paligemma.model.language_model.get_input_embeddings() + embedding_device = self._module_device(embedding) + output_device = self._prefix_transformer_device() + if tokens.device != embedding_device: + tokens = tokens.to(device=embedding_device) + embeds = embedding(tokens) + if embeds.device != output_device: + embeds = embeds.to(device=output_device) + return embeds + + def forward( + self, + attention_mask: torch.Tensor | None, + position_ids: torch.LongTensor | None, + past_key_values, + inputs_embeds: list[torch.FloatTensor | None], + use_cache: bool | None, + adarms_cond: list[torch.Tensor | None] | None = None, + ): + if adarms_cond is None: + adarms_cond = [None, None] + if inputs_embeds[1] is None: + prefix_output = self.paligemma.model.language_model.forward( + inputs_embeds=inputs_embeds[0], + attention_mask=attention_mask, + position_ids=position_ids, + past_key_values=past_key_values, + use_cache=use_cache, + adarms_cond=adarms_cond[0], + ) + return [ + prefix_output.last_hidden_state, + None, + ], prefix_output.past_key_values + + if inputs_embeds[0] is None: + suffix_output = self.gemma_expert.model.forward( + inputs_embeds=inputs_embeds[1], + attention_mask=attention_mask, + position_ids=position_ids, + past_key_values=past_key_values, + use_cache=use_cache, + adarms_cond=adarms_cond[1], + ) + return [None, suffix_output.last_hidden_state], None + + paligemma_layers = self.paligemma.model.language_model.layers + expert_layers = self.gemma_expert.model.layers + rotary_emb = self.paligemma.model.language_model.rotary_emb + for layers in zip(paligemma_layers, expert_layers, strict=True): + inputs_embeds = compute_layer_complete( + inputs_embeds, + attention_mask, + position_ids, + adarms_cond, + layers=layers, + rotary_emb=rotary_emb, + ) + + final_norms = ( + self.paligemma.model.language_model.norm, + self.gemma_expert.model.norm, + ) + outputs = [] + for i, hidden_states in enumerate(inputs_embeds): + out_emb, _ = layernorm_forward( + final_norms[i], hidden_states, adarms_cond[i] + ) + outputs.append(out_emb) + return outputs, None + + +class Pi05CoreModel(nn.Module): + def __init__( + self, + config: Pi05PipelineConfig, + runtime_role: Literal["all", "prefix", "action", "idle"] = "all", + *, + prefix_tensor_parallel: bool = False, + ): + super().__init__() + self.config = config + vlm_config = get_gemma_variant_config(config.paligemma_variant) + action_config = get_gemma_variant_config(config.action_expert_variant) + precision = ( + "bfloat16" + if config.materialize_dtype in ("bf16", "bfloat16") + else "float32" + ) + self.paligemma_with_expert = PaliGemmaWithExpertModel( + vlm_config, + action_config, + use_adarms=[False, True], + precision=precision, + image_size=config.image_size[0], + runtime_role=runtime_role, + prefix_tensor_parallel=prefix_tensor_parallel, + ) + if runtime_role in ("all", "action"): + self.action_in_proj = nn.Linear(config.action_dim, action_config.width) + self.action_out_proj = nn.Linear(action_config.width, config.action_dim) + self.time_mlp_in = nn.Linear(action_config.width, action_config.width) + self.time_mlp_out = nn.Linear(action_config.width, action_config.width) + if precision == "bfloat16": + for module in ( + self.action_in_proj, + self.action_out_proj, + self.time_mlp_in, + self.time_mlp_out, + ): + module.to(dtype=torch.float32) + else: + self.action_in_proj = None + self.action_out_proj = None + self.time_mlp_in = None + self.time_mlp_out = None + + def retain_runtime_components( + self, + role: Literal["all", "prefix", "action", "idle"], + ) -> None: + if role == "all": + return + if role in ("prefix", "idle"): + self.paligemma_with_expert.gemma_expert = None + self.action_in_proj = None + self.action_out_proj = None + self.time_mlp_in = None + self.time_mlp_out = None + if role in ("action", "idle"): + self.paligemma_with_expert.paligemma = None + + def prepare_attention_masks_4d( + self, + att_2d_masks: torch.Tensor, + *, + full_attention: bool | None = None, + ) -> torch.Tensor | None: + return prepare_optional_full_attention_mask( + att_2d_masks, + full_attention=full_attention, + ) + + def embed_prefix( + self, + images: list[torch.Tensor], + image_masks: list[torch.Tensor], + tokens: torch.Tensor, + token_masks: torch.Tensor, + ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]: + embs = [] + pad_masks = [] + att_masks = [] + image_embs = self.paligemma_with_expert.embed_images(images) + for image_emb, image_mask in zip(image_embs, image_masks, strict=True): + batch_size, num_image_embs = image_emb.shape[:2] + embs.append(image_emb) + pad_masks.append(image_mask[:, None].expand(batch_size, num_image_embs)) + att_masks += [0] * num_image_embs + + lang_emb = self.paligemma_with_expert.embed_language_tokens(tokens) + embs.append(lang_emb) + pad_masks.append(token_masks) + att_masks += [0] * lang_emb.shape[1] + + embs = torch.cat(embs, dim=1) + pad_masks = torch.cat(pad_masks, dim=1) + att_masks_t = torch.tensor(att_masks, dtype=torch.bool, device=pad_masks.device) + att_masks_t = att_masks_t[None, :].expand(pad_masks.shape[0], len(att_masks)) + return embs, pad_masks, att_masks_t + + def embed_suffix( + self, + noisy_actions: torch.Tensor, + timestep: torch.Tensor, + ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]: + time_emb = create_sinusoidal_pos_embedding( + timestep, + self.action_in_proj.out_features, + min_period=self.config.time_embedding_min_period, + max_period=self.config.time_embedding_max_period, + ) + action_emb = self.action_in_proj( + noisy_actions.to(dtype=self.action_in_proj.weight.dtype) + ) + time_emb = time_emb.to(dtype=self.time_mlp_in.weight.dtype) + time_emb = self.time_mlp_in(time_emb) + time_emb = F.silu(time_emb) + time_emb = self.time_mlp_out(time_emb) + adarms_cond = F.silu(time_emb) + + batch_size, action_len = action_emb.shape[:2] + pad_masks = torch.ones( + batch_size, + action_len, + dtype=torch.bool, + device=noisy_actions.device, + ) + att_masks_t = torch.zeros( + batch_size, + action_len, + dtype=action_emb.dtype, + device=noisy_actions.device, + ) + att_masks_t[:, 0] = 1 + return action_emb, pad_masks, att_masks_t, adarms_cond + + def _move_prefix_image_encoder_to_device(self, device: torch.device) -> None: + paligemma = self.paligemma_with_expert.paligemma + paligemma.model.vision_tower.to(device) + paligemma.model.multi_modal_projector.to(device) + + def _prefix_language_phase_offload_layer_count(self) -> int: + language_model = self.paligemma_with_expert.paligemma.model.language_model + if self.config.offload_prefix_language_layers_after_prefix: + return len(language_model.layers) + return min( + max(self.config.offload_prefix_language_layer_count_after_prefix, 0), + len(language_model.layers), + ) + + def _move_prefix_language_layers_to_device(self, device: torch.device) -> None: + language_model = self.paligemma_with_expert.paligemma.model.language_model + layer_count = self._prefix_language_phase_offload_layer_count() + for layer in language_model.layers[:layer_count]: + layer.to(device) + + def _prepare_prefix_image_encoder_for_embed(self) -> None: + if not self.config.offload_prefix_image_encoder_after_embed: + return + self._move_prefix_image_encoder_to_device( + self.paligemma_with_expert._prefix_transformer_device() + ) + + def _offload_prefix_image_encoder_after_embed(self) -> None: + if not self.config.offload_prefix_image_encoder_after_embed: + return + self._move_prefix_image_encoder_to_device(torch.device("cpu")) + if torch.cuda.is_available(): + torch.cuda.empty_cache() + + def _prepare_prefix_language_layers_for_forward(self) -> None: + if ( + not self._prefix_language_phase_offload_layer_count() + or self.config.offload_prefix_language_layers + ): + return + self._move_prefix_language_layers_to_device( + self.paligemma_with_expert._prefix_transformer_device() + ) + + def _offload_prefix_language_layers_after_prefix(self) -> None: + if ( + not self._prefix_language_phase_offload_layer_count() + or self.config.offload_prefix_language_layers + ): + return + self._move_prefix_language_layers_to_device(torch.device("cpu")) + if torch.cuda.is_available(): + torch.cuda.empty_cache() + + @torch.no_grad() + def encode_prefix( + self, + images: list[torch.Tensor], + image_masks: list[torch.Tensor], + tokens: torch.Tensor, + token_masks: torch.Tensor, + prefix_full_attention_hint: bool | None = None, + tokens_trimmed: bool = False, + ): + if not tokens_trimmed: + tokens, token_masks = trim_trailing_padding_tokens(tokens, token_masks) + self._prepare_prefix_image_encoder_for_embed() + prefix_embs, prefix_pad_masks, prefix_att_masks = self.embed_prefix( + images, image_masks, tokens, token_masks + ) + self._offload_prefix_image_encoder_after_embed() + prefix_position_ids = torch.cumsum(prefix_pad_masks, dim=1) - 1 + prefix_full_attention = bool(prefix_full_attention_hint) + if prefix_full_attention: + attention_mask = None + else: + prefix_att_2d_masks = make_att_2d_masks(prefix_pad_masks, prefix_att_masks) + prefix_full_attention = bool(prefix_att_2d_masks.all().item()) + attention_mask = self.prepare_attention_masks_4d( + prefix_att_2d_masks, + full_attention=prefix_full_attention, + ) + self._prepare_prefix_language_layers_for_forward() + with set_forward_context(current_timestep=0, attn_metadata=None): + _, past_key_values = self.paligemma_with_expert.forward( + attention_mask=attention_mask, + position_ids=prefix_position_ids, + past_key_values=None, + inputs_embeds=[prefix_embs, None], + use_cache=True, + ) + self._offload_prefix_language_layers_after_prefix() + return ( + past_key_values, + prefix_pad_masks, + prefix_full_attention, + ) + + @torch.no_grad() + def denoise_step( + self, + prefix_pad_masks: torch.Tensor, + past_key_values, + x_t: torch.Tensor, + timestep: torch.Tensor, + prefix_full_attention: bool = False, + *, + action_position_offset: int = 0, + ) -> torch.Tensor: + suffix_embs, suffix_pad_masks, suffix_att_masks, adarms_cond = ( + self.embed_suffix(x_t, timestep) + ) + suffix_len = suffix_pad_masks.shape[1] + batch_size = prefix_pad_masks.shape[0] + prefix_len = prefix_pad_masks.shape[1] + if prefix_full_attention: + attention_mask = None + else: + prefix_pad_2d_masks = prefix_pad_masks[:, None, :].expand( + batch_size, suffix_len, prefix_len + ) + suffix_att_2d_masks = make_att_2d_masks(suffix_pad_masks, suffix_att_masks) + full_att_2d_masks = torch.cat( + [prefix_pad_2d_masks, suffix_att_2d_masks], + dim=2, + ) + attention_mask = self.prepare_attention_masks_4d(full_att_2d_masks) + prefix_offsets = torch.sum(prefix_pad_masks, dim=-1)[:, None] + position_ids = ( + prefix_offsets + + action_position_offset + + torch.cumsum(suffix_pad_masks, dim=1) + - 1 + ) + with set_forward_context(current_timestep=0, attn_metadata=None): + outputs_embeds, _ = self.paligemma_with_expert.forward( + attention_mask=attention_mask, + position_ids=position_ids, + past_key_values=VLADensePrefixCache( + past_key_values, + read_only=True, + ), + inputs_embeds=[None, suffix_embs], + use_cache=False, + adarms_cond=[None, adarms_cond], + ) + suffix_out = outputs_embeds[1][:, -x_t.shape[1] :] + return self.action_out_proj( + suffix_out.to(dtype=self.action_out_proj.weight.dtype) + ).to(dtype=torch.float32) diff --git a/python/sglang/multimodal_gen/runtime/models/vlas/pi05_policy.py b/python/sglang/multimodal_gen/runtime/models/vlas/pi05_policy.py new file mode 100644 index 000000000..5e523d447 --- /dev/null +++ b/python/sglang/multimodal_gen/runtime/models/vlas/pi05_policy.py @@ -0,0 +1,1107 @@ +# SPDX-License-Identifier: Apache-2.0 + +from __future__ import annotations + +import json +import os +from dataclasses import dataclass, field +from pathlib import Path +from typing import Any, Iterator + +import torch +from safetensors import safe_open +from torch import nn +from torch.distributed.fsdp import MixedPrecisionPolicy + +from sglang.multimodal_gen.configs.pipeline_configs.pi05 import Pi05PipelineConfig +from sglang.multimodal_gen.runtime.distributed.communication_op import ( + sequence_model_parallel_all_gather, + tensor_model_parallel_all_gather, +) +from sglang.multimodal_gen.runtime.distributed.parallel_state import ( + get_ring_parallel_world_size, + get_sequence_parallel_world_size, + get_sp_parallel_rank, + get_tp_world_size, + get_ulysses_parallel_world_size, + model_parallel_is_initialized, +) +from sglang.multimodal_gen.runtime.loader.utils import ( + set_default_torch_dtype, + skip_init_modules, +) +from sglang.multimodal_gen.runtime.loader.weight_load_plan import WeightLoadPlan +from sglang.multimodal_gen.runtime.loader.weight_utils import ( + safetensors_weights_iterator, +) +from sglang.multimodal_gen.runtime.models.vlas.pi05_core import Pi05CoreModel +from sglang.multimodal_gen.runtime.platforms import current_platform +from sglang.multimodal_gen.runtime.utils.hf_diffusers_utils import maybe_download_model +from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger +from sglang.multimodal_gen.runtime.vla.denoise_cuda_graph import ( + VLADenoiseGraphRunner, + VLADenoiseGraphSignature, +) +from sglang.multimodal_gen.runtime.vla.observation import ( + VLAObservationBatch, + tensor_fingerprint, +) +from sglang.multimodal_gen.runtime.vla.parallel import ( + broadcast_tensor_from_rank, + get_vla_split_group, +) +from sglang.multimodal_gen.runtime.vla.prefix_cache import ( + PrefixContext, + VLADensePrefixCache, + VLAPrefixCacheManager, +) +from sglang.multimodal_gen.utils import set_mixed_precision_policy + +logger = init_logger(__name__) + + +@dataclass +class Pi05CheckpointManifest: + model_path: str + safetensor_files: list[str] = field(default_factory=list) + component_keys: dict[str, list[str]] = field(default_factory=dict) + skipped_lm_head_keys: list[str] = field(default_factory=list) + + +class Pi05ActionExpert(nn.Module): + def __init__(self, config: Pi05PipelineConfig, core_model: Pi05CoreModel): + super().__init__() + self.config = config + self.core_model = core_model + + def forward( + self, + prefix_context: PrefixContext, + x_t: torch.Tensor, + timestep: torch.Tensor, + *, + action_position_offset: int = 0, + ) -> torch.Tensor: + return self.core_model.denoise_step( + prefix_context.prefix_pad_masks, + prefix_context.past_key_values, + x_t, + timestep, + bool(prefix_context.layout.get("full_attention", False)), + action_position_offset=action_position_offset, + ) + + +class Pi05PolicyModel(nn.Module): + _ROLE_COMPONENTS = { + "all": None, + "prefix": {"vision_tower", "paligemma", "multi_modal_projector"}, + "action": {"action_expert", "action_heads"}, + "idle": set(), + } + _FUSED_WEIGHT_MAPPINGS = ( + (".self_attn.qkv_proj.", ".self_attn.q_proj.", "q"), + (".self_attn.qkv_proj.", ".self_attn.k_proj.", "k"), + (".self_attn.qkv_proj.", ".self_attn.v_proj.", "v"), + (".mlp.gate_up_proj.", ".mlp.gate_proj.", 0), + (".mlp.gate_up_proj.", ".mlp.up_proj.", 1), + ) + + def __init__( + self, + config: Pi05PipelineConfig, + *, + model_path: str, + device: torch.device, + dtype: torch.dtype, + manifest: Pi05CheckpointManifest, + ): + super().__init__() + self.config = config + self.model_path = model_path + self.device = device + self.dtype = dtype + self.manifest = manifest + self.runtime_role = self._resolve_runtime_role() + mp_policy = MixedPrecisionPolicy( + dtype, + dtype, + dtype, + cast_forward_inputs=False, + ) + set_mixed_precision_policy( + param_dtype=dtype, + reduce_dtype=dtype, + output_dtype=dtype, + mp_policy=mp_policy, + ) + prefix_tensor_parallel = self._should_use_prefix_tensor_parallel() + with set_default_torch_dtype(dtype), skip_init_modules(): + self.core_model = Pi05CoreModel( + config, + runtime_role=self.runtime_role, + prefix_tensor_parallel=prefix_tensor_parallel, + ) + self.core_model.eval() + if device.type == "cuda": + if self._use_componentwise_empty_init(): + self._componentwise_empty_init() + else: + self._to_empty_preserve_buffers(self.core_model, device=device) + self._set_prefix_output_device() + self._move_offloaded_prefix_modules_to_empty_cpu() + if torch.cuda.is_available(): + torch.cuda.empty_cache() + self._load_weights() + torch.cuda.empty_cache() + else: + self._load_weights() + self.core_model.to(device) + self._set_prefix_output_device() + if self.runtime_role != "all": + logger.info("Pi05 split runtime role on rank: %s", self.runtime_role) + self.action_expert = Pi05ActionExpert(config, self.core_model) + self.graph_runner = VLADenoiseGraphRunner( + enabled=config.enable_action_cuda_graph + ) + + def _should_use_prefix_tensor_parallel(self) -> bool: + if self.runtime_role not in ("all", "prefix"): + return False + if self.config.prefix_parallel_strategy != "tp": + return False + if get_vla_split_group() is not None: + return False + if not model_parallel_is_initialized(): + return False + return get_tp_world_size() > 1 + + @staticmethod + def _to_empty_preserve_buffers(module: nn.Module, *, device: torch.device) -> None: + buffers = { + name: buffer.detach().clone().to(device=device) + for name, buffer in module.named_buffers(recurse=True) + } + module.to_empty(device=device) + for name, buffer in buffers.items(): + if "." in name: + parent_name, buffer_name = name.rsplit(".", 1) + parent = module.get_submodule(parent_name) + else: + parent = module + buffer_name = name + parent._buffers[buffer_name] = buffer + + def _use_componentwise_empty_init(self) -> bool: + prefix_offload = ( + self.config.offload_prefix_image_encoder + or self.config.offload_prefix_image_encoder_after_embed + or self.config.offload_prefix_token_embedding + or self.config.offload_prefix_language_layers + or self.config.offload_prefix_language_layers_after_prefix + or self.config.offload_prefix_language_layer_count_after_prefix > 0 + ) + return ( + self.runtime_role in ("all", "prefix") and prefix_offload + ) or self._offload_action_expert_between_requests() + + def _prefix_language_phase_offload_layer_count(self, layer_count: int) -> int: + if self.config.offload_prefix_language_layers_after_prefix: + return layer_count + return min( + max(self.config.offload_prefix_language_layer_count_after_prefix, 0), + layer_count, + ) + + def _componentwise_empty_init(self) -> None: + logger.info( + "Pi05 componentwise empty init enabled for runtime role %s", + self.runtime_role, + ) + self._to_empty_preserve_buffers(self.core_model, device=torch.device("cpu")) + paligemma = self.core_model.paligemma_with_expert.paligemma + if paligemma is not None: + language_model = paligemma.model.language_model + self._to_empty_preserve_buffers( + language_model.rotary_emb, + device=self.device, + ) + self._to_empty_preserve_buffers(language_model.norm, device=self.device) + + if ( + self.config.offload_prefix_image_encoder + or self.config.offload_prefix_image_encoder_after_embed + ): + self._to_empty_preserve_buffers( + paligemma.model.vision_tower, + device=torch.device("cpu"), + ) + self._to_empty_preserve_buffers( + paligemma.model.multi_modal_projector, + device=torch.device("cpu"), + ) + else: + self._to_empty_preserve_buffers( + paligemma.model.vision_tower, + device=self.device, + ) + self._to_empty_preserve_buffers( + paligemma.model.multi_modal_projector, + device=self.device, + ) + + if self.config.offload_prefix_token_embedding: + self._to_empty_preserve_buffers( + language_model.embed_tokens, + device=torch.device("cpu"), + ) + else: + self._to_empty_preserve_buffers( + language_model.embed_tokens, + device=self.device, + ) + + phase_layer_count = self._prefix_language_phase_offload_layer_count( + len(language_model.layers) + ) + if self.config.offload_prefix_language_layers: + self._to_empty_preserve_buffers( + language_model.layers, + device=torch.device("cpu"), + ) + language_model.configure_layerwise_cpu_offload( + compute_device=self.device, + empty_cache=self.config.offload_prefix_language_layers_empty_cache, + ) + elif phase_layer_count: + for i, layer in enumerate(language_model.layers): + self._to_empty_preserve_buffers( + layer, + device=( + torch.device("cpu") + if i < phase_layer_count + else self.device + ), + ) + else: + self._to_empty_preserve_buffers( + language_model.layers, + device=self.device, + ) + + self._set_prefix_output_device() + action_device = ( + torch.device("cpu") + if self._offload_action_expert_between_requests() + else self.device + ) + self._move_action_modules_to_empty_device(action_device) + if torch.cuda.is_available(): + torch.cuda.empty_cache() + + def _move_action_modules_to_empty_device(self, device: torch.device) -> None: + gemma_expert = self.core_model.paligemma_with_expert.gemma_expert + if gemma_expert is not None: + self._to_empty_preserve_buffers(gemma_expert, device=device) + for module in ( + self.core_model.action_in_proj, + self.core_model.action_out_proj, + self.core_model.time_mlp_in, + self.core_model.time_mlp_out, + ): + if module is not None: + self._to_empty_preserve_buffers(module, device=device) + + def _offload_action_expert_between_requests(self) -> bool: + return ( + self.runtime_role == "all" + and self.device.type == "cuda" + and self.config.offload_action_expert_after_denoise + ) + + def _move_action_modules_to_device(self, device: torch.device) -> None: + gemma_expert = self.core_model.paligemma_with_expert.gemma_expert + if gemma_expert is not None: + gemma_expert.to(device) + for module in ( + self.core_model.action_in_proj, + self.core_model.action_out_proj, + self.core_model.time_mlp_in, + self.core_model.time_mlp_out, + ): + if module is not None: + module.to(device) + + def _set_prefix_output_device(self) -> None: + paligemma_with_expert = self.core_model.paligemma_with_expert + if paligemma_with_expert.paligemma is not None: + paligemma_with_expert.set_prefix_output_device(self.device) + + def _move_offloaded_prefix_modules_to_empty_cpu(self) -> None: + paligemma = self.core_model.paligemma_with_expert.paligemma + if paligemma is None: + return + cpu = torch.device("cpu") + if ( + self.config.offload_prefix_image_encoder + or self.config.offload_prefix_image_encoder_after_embed + ): + self._to_empty_preserve_buffers( + paligemma.model.vision_tower, + device=cpu, + ) + self._to_empty_preserve_buffers( + paligemma.model.multi_modal_projector, + device=cpu, + ) + if self.config.offload_prefix_token_embedding: + self._to_empty_preserve_buffers( + paligemma.model.language_model.embed_tokens, + device=cpu, + ) + language_model = paligemma.model.language_model + phase_layer_count = self._prefix_language_phase_offload_layer_count( + len(language_model.layers) + ) + if self.config.offload_prefix_language_layers: + self._to_empty_preserve_buffers(language_model.layers, device=cpu) + language_model.configure_layerwise_cpu_offload( + compute_device=self.device, + empty_cache=self.config.offload_prefix_language_layers_empty_cache, + ) + elif phase_layer_count: + for layer in language_model.layers[:phase_layer_count]: + self._to_empty_preserve_buffers(layer, device=cpu) + + @staticmethod + def _resolve_runtime_role() -> str: + split = get_vla_split_group() + if split is None: + return "all" + if split.is_prefix_rank and split.is_action_rank: + return "all" + if split.is_prefix_rank: + return "prefix" + if split.is_action_rank: + return "action" + return "idle" + + @classmethod + def from_pretrained( + cls, + model_path: str, + config: Pi05PipelineConfig, + *, + dtype: torch.dtype | None = None, + ) -> Pi05PolicyModel: + local_path = maybe_download_model( + model_path, + force_diffusers_model=False, + allow_patterns=["*.json", "*.model", "*.safetensors", "*.txt"], + ) + cls._apply_checkpoint_config(local_path, config) + device = torch.device(current_platform.device_type) + dtype = dtype or cls._dtype_from_config(config.materialize_dtype) + manifest = cls._inspect_checkpoint(local_path, config) + logger.info( + "Loaded Pi05 checkpoint manifest from %s (%d safetensors, %d skipped lm_head tensors)", + local_path, + len(manifest.safetensor_files), + len(manifest.skipped_lm_head_keys), + ) + return cls( + config, + model_path=local_path, + device=device, + dtype=dtype, + manifest=manifest, + ) + + @staticmethod + def _dtype_from_config(dtype_name: str) -> torch.dtype: + name = (dtype_name or "bf16").lower() + if name in ("bf16", "bfloat16"): + return torch.bfloat16 + if name in ("fp16", "float16", "half"): + return torch.float16 + if name in ("fp32", "float32"): + return torch.float32 + raise ValueError(f"Unsupported Pi05 dtype: {dtype_name}") + + @staticmethod + def _apply_checkpoint_config( + model_path: str, + config: Pi05PipelineConfig, + ) -> None: + config_path = Path(model_path) / "config.json" + if not config_path.exists(): + return + with open(config_path, encoding="utf-8") as f: + payload = json.load(f) + + config.paligemma_variant = payload.get( + "paligemma_variant", config.paligemma_variant + ) + config.action_expert_variant = payload.get( + "action_expert_variant", config.action_expert_variant + ) + config.action_horizon = int(payload.get("chunk_size", config.action_horizon)) + config.n_action_steps = int( + payload.get("n_action_steps", config.n_action_steps) + ) + config.action_dim = int(payload.get("max_action_dim", config.action_dim)) + config.state_dim = int(payload.get("max_state_dim", config.state_dim)) + config.default_num_inference_steps = int( + payload.get("num_inference_steps", config.default_num_inference_steps) + ) + config.max_token_len = int( + payload.get("tokenizer_max_length", config.max_token_len) + ) + config.time_embedding_min_period = float( + payload.get("min_period", config.time_embedding_min_period) + ) + config.time_embedding_max_period = float( + payload.get("max_period", config.time_embedding_max_period) + ) + if "image_resolution" in payload: + resolution = tuple(payload["image_resolution"]) + config.image_size = (int(resolution[0]), int(resolution[1])) + input_features = payload.get("input_features") or {} + image_keys = [] + for key, feature in input_features.items(): + if feature.get("type") == "VISUAL": + image_keys.append(key.rsplit(".", 1)[-1]) + elif feature.get("type") == "STATE": + shape = feature.get("shape") or [] + if shape: + config.state_dim = int(shape[0]) + config.empty_cameras = int(payload.get("empty_cameras", 0) or 0) + image_keys.extend(f"empty_camera_{i}" for i in range(config.empty_cameras)) + if image_keys: + config.image_keys = tuple(image_keys) + + output_features = payload.get("output_features") or {} + action_feature = output_features.get("action") or {} + action_shape = action_feature.get("shape") or [] + if action_shape: + config.output_action_dim = int(action_shape[0]) + + @staticmethod + def _inspect_checkpoint( + model_path: str, + config: Pi05PipelineConfig, + ) -> Pi05CheckpointManifest: + path = Path(model_path) + safetensor_files = sorted(str(p) for p in path.glob("*.safetensors")) + component_keys = {name: [] for name in config.loader_component_map} + skipped_lm_head_keys: list[str] = [] + + try: + from safetensors import safe_open + except ImportError: + return Pi05CheckpointManifest( + model_path=model_path, + safetensor_files=safetensor_files, + component_keys=component_keys, + ) + + for filename in safetensor_files: + with safe_open(filename, framework="pt", device="cpu") as f: + for key in f.keys(): + if ( + config.skip_unused_lm_head + and key == "paligemma_with_expert.gemma_expert.lm_head.weight" + ): + skipped_lm_head_keys.append(key) + continue + for component, prefixes in config.loader_component_map.items(): + if any(key.startswith(prefix) for prefix in prefixes): + component_keys[component].append(key) + break + + return Pi05CheckpointManifest( + model_path=model_path, + safetensor_files=safetensor_files, + component_keys=component_keys, + skipped_lm_head_keys=skipped_lm_head_keys, + ) + + @staticmethod + def _candidate_weight_keys(key: str) -> list[str]: + if key.startswith("model."): + key = key[len("model.") :] + if key.startswith("PaligemmaWithExpert."): + key = key.replace("PaligemmaWithExpert.", "paligemma_with_expert.", 1) + if key.startswith("action_time_mlp_in."): + key = key.replace("action_time_mlp_in.", "time_mlp_in.", 1) + elif key.startswith("action_time_mlp_out."): + key = key.replace("action_time_mlp_out.", "time_mlp_out.", 1) + if key.startswith("state_proj."): + return [] + if key == "paligemma_with_expert.gemma_expert.lm_head.weight": + return [] + + candidates = [key] + replacements = { + ".vision_tower.vision_model.": ".vision_tower.", + ".paligemma.language_model.": ".paligemma.model.language_model.", + ".paligemma.vision_tower.": ".paligemma.model.vision_tower.", + ".paligemma.multi_modal_projector.": ( + ".paligemma.model.multi_modal_projector." + ), + } + for old, new in replacements.items(): + if old in key: + candidates.append(key.replace(old, new)) + + if key in { + "paligemma_with_expert.paligemma.lm_head.weight", + "paligemma_with_expert.paligemma.model.lm_head.weight", + }: + candidates.append( + "paligemma_with_expert.paligemma.model.language_model." + "embed_tokens.weight" + ) + return list(dict.fromkeys(candidates)) + + def _component_for_source_key(self, key: str) -> str | None: + if key in { + "paligemma_with_expert.paligemma.lm_head.weight", + "paligemma_with_expert.paligemma.model.lm_head.weight", + }: + return "paligemma" + candidates = self._candidate_weight_keys(key) + if not candidates: + return None + for component, prefixes in self.config.loader_component_map.items(): + if any( + candidate.startswith(prefix) + for candidate in candidates + for prefix in prefixes + ): + return component + return None + + def _should_load_source_key(self, key: str) -> bool: + components = self._ROLE_COMPONENTS[self.runtime_role] + if components is None: + return True + component = self._component_for_source_key(key) + return component in components + + def _should_read_source_key(self, key: str) -> bool: + return self._should_load_source_key(key) and bool( + self._candidate_weight_keys(key) + ) + + @classmethod + def _candidate_target_weights(cls, source_key: str) -> list[tuple[str, Any | None]]: + candidates = [] + for candidate in cls._candidate_weight_keys(source_key): + candidates.append((candidate, None)) + for target_name, weight_name, shard_id in cls._FUSED_WEIGHT_MAPPINGS: + if weight_name in candidate: + candidates.append( + (candidate.replace(weight_name, target_name), shard_id) + ) + return list(dict.fromkeys(candidates)) + + def _resolve_target_weight( + self, + source_key: str, + target_state: dict[str, torch.Tensor], + target_params: dict[str, nn.Parameter], + ) -> tuple[str, Any | None] | None: + candidates = self._candidate_target_weights(source_key) + if not candidates: + return None + for candidate, shard_id in candidates: + if candidate in target_params or candidate in target_state: + return candidate, shard_id + return None + + @staticmethod + def _target_tensor_for_key( + target_key: str, + target_state: dict[str, torch.Tensor], + target_params: dict[str, nn.Parameter], + ) -> torch.Tensor: + target = target_params.get(target_key) + if target is not None: + return target + return target_state[target_key] + + @staticmethod + def _load_tensor_to_target( + target: torch.Tensor, + tensor: torch.Tensor, + shard_id: Any | None, + ) -> bool: + if tensor.dtype != target.dtype: + tensor = tensor.to(dtype=target.dtype) + weight_loader = getattr(target, "weight_loader", None) + if weight_loader is not None: + if shard_id is None: + weight_loader(target, tensor) + else: + weight_loader(target, tensor, shard_id) + return True + if tuple(target.shape) != tuple(tensor.shape): + return False + target.copy_(tensor, non_blocking=target.device.type == "cuda") + return True + + def _should_stream_weights_to_gpu( + self, + target_state: dict[str, torch.Tensor], + target_params: dict[str, nn.Parameter], + ) -> bool: + if self.device.type != "cuda": + return False + + has_weight = False + for filename in self.manifest.safetensor_files: + with safe_open(filename, framework="pt", device="cpu") as f: + for source_key in f.keys(): + if not self._should_read_source_key(source_key): + continue + target_weight = self._resolve_target_weight( + source_key, + target_state, + target_params, + ) + if target_weight is None: + continue + target_key, _ = target_weight + has_weight = True + target = self._target_tensor_for_key( + target_key, + target_state, + target_params, + ) + if target.device.type != "cuda": + return False + return has_weight + + def _cpu_weights_iterator(self) -> Iterator[tuple[str, torch.Tensor]]: + for filename in self.manifest.safetensor_files: + with safe_open(filename, framework="pt", device="cpu") as f: + for source_key in f.keys(): + if self._should_read_source_key(source_key): + yield source_key, f.get_tensor(source_key) + + def _load_weights(self) -> None: + target_state = self.core_model.state_dict() + target_params = dict(self.core_model.named_parameters()) + loaded_keys: set[str] = set() + unexpected = 0 + mismatched = 0 + stream_to_gpu = self._should_stream_weights_to_gpu( + target_state, + target_params, + ) + checkpoint_load_device = self.device if stream_to_gpu else torch.device("cpu") + weight_load_plan = WeightLoadPlan(checkpoint_load_device=checkpoint_load_device) + if stream_to_gpu: + logger.info( + "Pi05 weight load streams safetensors directly to %s", + weight_load_plan.checkpoint_load_device, + ) + + with torch.no_grad(): + if stream_to_gpu: + weights = safetensors_weights_iterator( + self.manifest.safetensor_files, + weight_load_plan=weight_load_plan, + key_filter=self._should_read_source_key, + clone_streamed_tensors=False, + ) + else: + weights = self._cpu_weights_iterator() + + for source_key, tensor in weights: + target_weight = self._resolve_target_weight( + source_key, + target_state, + target_params, + ) + if target_weight is None: + unexpected += 1 + continue + target_key, shard_id = target_weight + target = self._target_tensor_for_key( + target_key, + target_state, + target_params, + ) + if not self._load_tensor_to_target(target, tensor, shard_id): + mismatched += 1 + continue + loaded_keys.add(target_key) + + missing = [key for key in target_state if key not in loaded_keys] + if missing or mismatched: + raise RuntimeError( + f"Pi05 weight load failed: {len(missing)} missing weights, " + f"{mismatched} mismatched weights. Running a robot policy with " + "uninitialized or partially loaded weights is unsafe." + ) + if unexpected: + logger.warning( + "Pi05 weight load: %d loaded, %d unexpected", + len(loaded_keys), + unexpected, + ) + else: + logger.info("Pi05 weights loaded successfully") + + def build_prefix_cache_key( + self, + observation: VLAObservationBatch, + ) -> str: + camera_order = tuple(observation.metadata.get("camera_order", ())) + image_hashes = { + name: tensor_fingerprint(observation.images[name]) for name in camera_order + } + masks = { + name: bool(mask.item()) for name, mask in observation.image_masks.items() + } + token_len = int(observation.token_masks.sum(dim=1).max().item()) + tokens = ( + observation.tokens[:, :token_len] if token_len > 0 else observation.tokens + ) + token_masks = ( + observation.token_masks[:, :token_len] + if token_len > 0 + else observation.token_masks + ) + model_revision = os.path.basename(os.path.normpath(self.model_path)) + return VLAPrefixCacheManager.make_key( + model_revision=model_revision, + tokenizer_id=f"{self.config.paligemma_variant}:{self.config.max_token_len}", + camera_order=camera_order, + image_hashes=image_hashes, + token_digest=tensor_fingerprint(tokens), + token_mask_digest=tensor_fingerprint(token_masks), + masks=masks, + positions_version=self.config.prefix_cache_layout_version, + dtype=str(self.dtype).replace("torch.", ""), + parallel_layout_version=self.config.parallel_layout_version, + cache_namespace="pi05", + ) + + def _prefix_language_model(self) -> nn.Module | None: + paligemma = self.core_model.paligemma_with_expert.paligemma + if paligemma is None: + return None + return paligemma.model.language_model + + def _prefix_kv_requires_tp_gather(self) -> bool: + language_model = self._prefix_language_model() + if language_model is None or not language_model.tensor_parallel: + return False + if not language_model.layers: + return False + attn = language_model.layers[0].self_attn + return attn.total_num_key_value_heads > attn.num_key_value_heads + + def _materialize_prefix_kv_for_action( + self, + past_key_values: VLADensePrefixCache, + ) -> VLADensePrefixCache: + if not self._prefix_kv_requires_tp_gather(): + return past_key_values + return VLADensePrefixCache( + tuple( + ( + tensor_model_parallel_all_gather(keys.contiguous(), dim=1), + tensor_model_parallel_all_gather(values.contiguous(), dim=1), + sliding_window, + ) + for keys, values, sliding_window in past_key_values + ) + ) + + def encode_prefix(self, observation: VLAObservationBatch) -> PrefixContext: + camera_order = tuple(observation.metadata.get("camera_order", ())) + images = [ + observation.images[name].to(self.device, dtype=torch.float32) + for name in camera_order + ] + image_masks = [ + observation.image_masks[name].to(self.device) for name in camera_order + ] + token_len = int(observation.token_masks.sum(dim=1).max().item()) + tokens_trimmed = token_len > 0 + if tokens_trimmed and token_len < observation.tokens.shape[1]: + tokens_cpu = observation.tokens[:, :token_len] + token_masks_cpu = observation.token_masks[:, :token_len] + else: + tokens_cpu = observation.tokens + token_masks_cpu = observation.token_masks + tokens = tokens_cpu.to(self.device) + token_masks = token_masks_cpu.to(self.device) + prefix_full_attention_hint = all( + bool(observation.image_masks[name].all().item()) for name in camera_order + ) and bool(token_masks_cpu.all().item()) + past_key_values, prefix_pad_masks, full_attention = ( + self.core_model.encode_prefix( + images, + image_masks, + tokens, + token_masks, + prefix_full_attention_hint=prefix_full_attention_hint, + tokens_trimmed=tokens_trimmed, + ) + ) + past_key_values = self._materialize_prefix_kv_for_action(past_key_values) + return PrefixContext( + past_key_values=past_key_values, + prefix_pad_masks=prefix_pad_masks, + prefix_len=prefix_pad_masks.shape[1], + layout={"full_attention": full_attention}, + ) + + def sample_noise( + self, + batch_size: int, + *, + generator: torch.Generator | None = None, + ) -> torch.Tensor: + return torch.randn( + batch_size, + self.config.action_horizon, + self.config.action_dim, + generator=generator, + device=self.device, + dtype=torch.float32, + ) + + def denoise_step( + self, + prefix_context: PrefixContext, + x_t: torch.Tensor, + timestep: torch.Tensor, + *, + use_cuda_graph: bool = True, + action_position_offset: int = 0, + action_sp_enabled: bool = False, + ) -> torch.Tensor: + if not bool(prefix_context.layout.get("full_attention", False)): + use_cuda_graph = False + if not use_cuda_graph: + return self.action_expert( + prefix_context, + x_t, + timestep, + action_position_offset=action_position_offset, + ) + parallel_layout = self.config.parallel_layout_version + if action_sp_enabled: + parallel_layout = ( + f"{parallel_layout}:action_sp" + f":rank{get_sp_parallel_rank()}:offset{action_position_offset}" + ) + signature = VLADenoiseGraphSignature( + batch_size=x_t.shape[0], + prefix_len=prefix_context.prefix_len, + action_horizon=x_t.shape[1], + action_dim=x_t.shape[2], + dtype=str(x_t.dtype).replace("torch.", ""), + parallel_layout=parallel_layout, + ) + + def step_fn( + current_prefix_context: PrefixContext, + current_x_t: torch.Tensor, + current_timestep: torch.Tensor, + ) -> torch.Tensor: + return self.action_expert( + current_prefix_context, + current_x_t, + current_timestep, + action_position_offset=action_position_offset, + ) + + return self.graph_runner.capture_or_run( + signature, + step_fn, + prefix_context, + x_t, + timestep, + ) + + def _can_use_action_sequence_parallel( + self, + prefix_context: PrefixContext | None, + action_horizon: int, + ) -> bool: + split = get_vla_split_group() + if split is None or not split.uses_action_sp: + return False + if self.runtime_role not in ("all", "action"): + return False + if self._offload_action_expert_between_requests(): + return False + if prefix_context is None or not bool( + prefix_context.layout.get("full_attention", False) + ): + return False + if not model_parallel_is_initialized(): + return False + try: + sp_world_size = get_sequence_parallel_world_size() + ulysses_world_size = get_ulysses_parallel_world_size() + ring_world_size = get_ring_parallel_world_size() + except AssertionError: + return False + if sp_world_size <= 1 or ulysses_world_size <= 1 or ring_world_size != 1: + return False + action_expert = self.core_model.paligemma_with_expert.gemma_expert + if action_expert is None: + return False + num_heads = action_expert.model.config.num_attention_heads + return action_horizon % sp_world_size == 0 and num_heads % sp_world_size == 0 + + def should_run_action_denoise( + self, + prefix_context: PrefixContext | None, + ) -> bool: + split = get_vla_split_group() + if split is None: + return True + if self._can_use_action_sequence_parallel( + prefix_context, + self.config.action_horizon, + ): + return split.is_action_rank + return split.rank == split.action_root + + def action_parallel_info( + self, + prefix_context: PrefixContext | None, + ) -> dict[str, Any]: + split = get_vla_split_group() + if split is None: + return { + "split_group": False, + "runtime_role": self.runtime_role, + "action_sequence_parallel": False, + } + return { + "split_group": True, + "runtime_role": self.runtime_role, + "world_size": split.group.world_size, + "prefix_root": split.prefix_root, + "action_root": split.action_root, + "action_ranks": list(split.action_ranks), + "action_sequence_parallel": self._can_use_action_sequence_parallel( + prefix_context, + self.config.action_horizon, + ), + } + + def _broadcast_initial_action_state( + self, + x_t: torch.Tensor | None, + ) -> torch.Tensor: + split = get_vla_split_group() + if split is None: + if x_t is None: + raise RuntimeError("Pi05 action state is missing on single-rank run") + return x_t + x_t = broadcast_tensor_from_rank( + x_t, + split, + src=split.action_root, + device=self.device, + ) + if x_t is None: + raise RuntimeError("Pi05 action state broadcast returned None") + return x_t + + def _shard_action_sequence(self, x_t: torch.Tensor) -> tuple[torch.Tensor, int]: + sp_world_size = get_sequence_parallel_world_size() + sp_rank = get_sp_parallel_rank() + local_len = x_t.shape[1] // sp_world_size + start = sp_rank * local_len + end = start + local_len + return x_t[:, start:end].contiguous(), start + + def sample_actions( + self, + observation: VLAObservationBatch, + prefix_context: PrefixContext, + *, + noise: torch.Tensor | None, + num_steps: int, + use_cuda_graph: bool = True, + generator: torch.Generator | None = None, + ) -> torch.Tensor: + offload_action = self._offload_action_expert_between_requests() + if offload_action: + self._move_action_modules_to_device(self.device) + use_cuda_graph = False + + split = get_vla_split_group() + action_sp_enabled = self._can_use_action_sequence_parallel( + prefix_context, + self.config.action_horizon, + ) + if split is None or split.rank == split.action_root: + x_t = noise + if x_t is None: + x_t = self.sample_noise(observation.batch_size, generator=generator) + else: + x_t = x_t.to(device=self.device, dtype=torch.float32).clone() + else: + x_t = None + if action_sp_enabled: + x_t = self._broadcast_initial_action_state(x_t) + elif x_t is None: + raise RuntimeError("Pi05 action fallback must run on the action root") + action_position_offset = 0 + if action_sp_enabled: + x_t, action_position_offset = self._shard_action_sequence(x_t) + + dt = -1.0 / num_steps + timesteps = torch.linspace( + 1.0, + 1.0 / num_steps, + num_steps, + dtype=torch.float32, + device=self.device, + ) + for timestep_value in timesteps: + timestep = timestep_value.expand(observation.batch_size) + velocity = self.denoise_step( + prefix_context, + x_t, + timestep, + use_cuda_graph=use_cuda_graph, + action_position_offset=action_position_offset, + action_sp_enabled=action_sp_enabled, + ) + x_t.add_(velocity, alpha=dt) + if action_sp_enabled: + x_t = sequence_model_parallel_all_gather(x_t.contiguous(), dim=1) + if offload_action: + self._move_action_modules_to_device(torch.device("cpu")) + torch.cuda.empty_cache() + return x_t + + def warmup_actions(self, batch_size: int = 1) -> torch.Tensor: + return torch.zeros( + batch_size, + self.config.action_horizon, + self.config.action_dim, + device=self.device, + dtype=torch.float32, + ) + + +__all__ = [ + "Pi05ActionExpert", + "Pi05CheckpointManifest", + "Pi05PolicyModel", +] diff --git a/python/sglang/multimodal_gen/runtime/pipelines/pi05.py b/python/sglang/multimodal_gen/runtime/pipelines/pi05.py new file mode 100644 index 000000000..f9353681d --- /dev/null +++ b/python/sglang/multimodal_gen/runtime/pipelines/pi05.py @@ -0,0 +1,126 @@ +# SPDX-License-Identifier: Apache-2.0 + +from __future__ import annotations + +import torch + +from sglang.multimodal_gen.configs.pipeline_configs.pi05 import Pi05PipelineConfig +from sglang.multimodal_gen.configs.sample.pi05 import Pi05SamplingParams +from sglang.multimodal_gen.runtime.disaggregation.roles import RoleType +from sglang.multimodal_gen.runtime.models.vlas import Pi05PolicyModel +from sglang.multimodal_gen.runtime.pipelines_core.composed_pipeline_base import ( + ComposedPipelineBase, +) +from sglang.multimodal_gen.runtime.pipelines_core.stages.model_specific_stages.pi05_preprocess import ( + Pi05Preprocessor, +) +from sglang.multimodal_gen.runtime.pipelines_core.stages.vla import ( + VLAActionDenoisingStage, + VLAActionPostprocessStage, + VLAObservationPreprocessStage, + VLAPrefixEncodingStage, +) +from sglang.multimodal_gen.runtime.server_args import ServerArgs +from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger +from sglang.multimodal_gen.runtime.vla.prefix_cache import VLAPrefixCacheManager + +logger = init_logger(__name__) + + +class Pi05Pipeline(ComposedPipelineBase): + pipeline_name = "Pi05Pipeline" + pipeline_config_cls = Pi05PipelineConfig + sampling_params_cls = Pi05SamplingParams + _required_config_modules: list[str] = [] + + def validate_disagg_role(self, role: RoleType) -> None: + if role != RoleType.MONOLITHIC: + raise ValueError( + "Pi05Pipeline v1 supports same-process execution only. " + "Use prefix/action logical groups inside one worker; cross-node " + "multimodal_gen disaggregation is a v2 target." + ) + + def load_modules( + self, + server_args: ServerArgs, + loaded_modules: dict[str, torch.nn.Module] | None = None, + ) -> dict[str, torch.nn.Module]: + if loaded_modules is not None: + return loaded_modules + + pipeline_config: Pi05PipelineConfig = server_args.pipeline_config + pipeline_config.offload_prefix_image_encoder = ( + pipeline_config.offload_prefix_image_encoder + or bool(server_args.image_encoder_cpu_offload) + ) + pipeline_config.offload_prefix_token_embedding = ( + pipeline_config.offload_prefix_token_embedding + or bool(server_args.text_encoder_cpu_offload) + ) + logger.info( + "Pi05 memory config: prefix_cache=%s/%s, action_cuda_graph=%s, " + "offload_image=%s, offload_image_after_embed=%s, " + "offload_tokens=%s, offload_language_layers=%s, " + "offload_language_after_prefix=%s/%s, " + "offload_action_after_denoise=%s, empty_cache_after_prefix=%s", + pipeline_config.enable_global_prefix_cache, + pipeline_config.prefix_cache_max_entries, + pipeline_config.enable_action_cuda_graph, + pipeline_config.offload_prefix_image_encoder, + pipeline_config.offload_prefix_image_encoder_after_embed, + pipeline_config.offload_prefix_token_embedding, + pipeline_config.offload_prefix_language_layers, + pipeline_config.offload_prefix_language_layers_after_prefix, + pipeline_config.offload_prefix_language_layer_count_after_prefix, + pipeline_config.offload_action_expert_after_denoise, + pipeline_config.empty_cache_after_prefix, + ) + policy_model = Pi05PolicyModel.from_pretrained( + self.model_path, + pipeline_config, + ) + if ( + pipeline_config.prefix_parallel_strategy + == pipeline_config.action_parallel_strategy + == "tp" + ): + raise ValueError( + "VLA action expert should not share the prefix TP layout. " + "Use SP, Ulysses, Ring, DP, or monolithic fallback for the " + "action path." + ) + return { + "policy_model": policy_model, + } + + def initialize_pipeline(self, server_args: ServerArgs) -> None: + pipeline_config: Pi05PipelineConfig = server_args.pipeline_config + self.preprocessor = Pi05Preprocessor(pipeline_config) + self.prefix_cache = VLAPrefixCacheManager( + max_entries=pipeline_config.prefix_cache_max_entries + ) + + def create_pipeline_stages(self, server_args: ServerArgs): + self.add_stage( + VLAObservationPreprocessStage(self.preprocessor), + "pi05_preprocess", + ) + self.add_stage( + VLAPrefixEncodingStage( + self.get_module("policy_model"), + self.prefix_cache, + ), + "pi05_prefix", + ) + self.add_stage( + VLAActionDenoisingStage(self.get_module("policy_model")), + "pi05_action_denoise", + ) + self.add_stage( + VLAActionPostprocessStage(), + "pi05_postprocess", + ) + + +EntryClass = Pi05Pipeline diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/composed_pipeline_base.py b/python/sglang/multimodal_gen/runtime/pipelines_core/composed_pipeline_base.py index efb88187b..d25fa3b28 100644 --- a/python/sglang/multimodal_gen/runtime/pipelines_core/composed_pipeline_base.py +++ b/python/sglang/multimodal_gen/runtime/pipelines_core/composed_pipeline_base.py @@ -979,7 +979,12 @@ class ComposedPipelineBase(ABC): # Execute each stage if not batch.is_warmup and not batch.suppress_logs: - logger.info( + stage_logger = ( + logger.debug + if server_args.pipeline_config.task_type.is_action_gen() + else logger.info + ) + stage_logger( "Running pipeline stages: %s", list(self._stage_name_mapping.keys()), main_process_only=True, @@ -1007,7 +1012,12 @@ class ComposedPipelineBase(ABC): ) if not batches[0].is_warmup and not batches[0].suppress_logs: - logger.info( + stage_logger = ( + logger.debug + if server_args.pipeline_config.task_type.is_action_gen() + else logger.info + ) + stage_logger( "Running grouped pipeline stages: %s", list(self._stage_name_mapping.keys()), main_process_only=True, diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/executors/pipeline_executor.py b/python/sglang/multimodal_gen/runtime/pipelines_core/executors/pipeline_executor.py index f8ecad2b5..cc5e22738 100644 --- a/python/sglang/multimodal_gen/runtime/pipelines_core/executors/pipeline_executor.py +++ b/python/sglang/multimodal_gen/runtime/pipelines_core/executors/pipeline_executor.py @@ -57,6 +57,10 @@ class PipelineExecutor(ABC): batch: Any, server_args: ServerArgs, ) -> None: + if isinstance(batch, list): + if not batch: + return + batch = batch[0] self.component_residency_manager.begin_request(stages, batch, server_args) def before_stage( diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/schedule_batch.py b/python/sglang/multimodal_gen/runtime/pipelines_core/schedule_batch.py index 066577df3..59166a632 100644 --- a/python/sglang/multimodal_gen/runtime/pipelines_core/schedule_batch.py +++ b/python/sglang/multimodal_gen/runtime/pipelines_core/schedule_batch.py @@ -22,7 +22,10 @@ from typing import Any, Optional, Sequence, Union import PIL.Image import torch -from sglang.multimodal_gen.configs.sample.sampling_params import SamplingParams +from sglang.multimodal_gen.configs.sample.sampling_params import ( + DataType, + SamplingParams, +) from sglang.multimodal_gen.runtime.post_training.rl_dataclasses import ( RolloutTrajectoryData, ) @@ -320,9 +323,11 @@ class Req: @property def resolution_key(self) -> str | None: """Return the batching config resolution key, e.g. "1024x1024".""" - if self.width is None or self.height is None: + width = getattr(self, "width", None) + height = getattr(self, "height", None) + if width is None or height is None: return None - return f"{int(self.width)}x{int(self.height)}" + return f"{int(width)}x{int(height)}" def set_as_warmup(self, warmup_steps: int = 1): self.is_warmup = True @@ -339,6 +344,13 @@ class Req: def validate(self): """Initialize dependent fields after dataclass initialization.""" + if getattr(self.sampling_params, "data_type", None) == DataType.ACTION: + self.do_classifier_free_guidance = False + if self.negative_prompt_embeds is None: + self.negative_prompt_embeds = [] + self.metrics = RequestMetrics(request_id=self.request_id) + return + # Prefer true_cfg_scale when it is explicitly provided. cfg_scale = ( self.true_cfg_scale @@ -360,6 +372,22 @@ class Req: def log(self, server_args: ServerArgs): if self.is_warmup or self.suppress_logs: return + if getattr(self.sampling_params, "data_type", None) == DataType.ACTION: + if not logger.isEnabledFor(logging.DEBUG): + return + logger.debug( + "VLA request: prompt=%s seed=%s steps=%s outputs=%s action=%sx%s " + "save_output=%s", + _sanitize_for_logging(self.prompt, key_hint="prompt"), + self.seed, + self.num_inference_steps, + self.num_outputs_per_prompt, + getattr(self, "action_horizon", None), + getattr(self, "action_dim", None), + self.save_output, + ) + return + # TODO: in some cases (e.g., TI2I), height and weight might be undecided at this moment if self.height: target_height = align_to(self.height, 16) diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/__init__.py b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/__init__.py index 7c124b21f..d0bd576d5 100644 --- a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/__init__.py +++ b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/__init__.py @@ -42,6 +42,12 @@ from sglang.multimodal_gen.runtime.pipelines_core.stages.timestep_preparation im DMDTimestepPreparationStage, TimestepPreparationStage, ) +from sglang.multimodal_gen.runtime.pipelines_core.stages.vla import ( + VLAActionDenoisingStage, + VLAActionPostprocessStage, + VLAObservationPreprocessStage, + VLAPrefixEncodingStage, +) __all__ = [ "PipelineStage", @@ -60,4 +66,8 @@ __all__ = [ "ImageEncodingStage", "ImageVAEEncodingStage", "TextEncodingStage", + "VLAObservationPreprocessStage", + "VLAPrefixEncodingStage", + "VLAActionDenoisingStage", + "VLAActionPostprocessStage", ] diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/pi05_preprocess.py b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/pi05_preprocess.py new file mode 100644 index 000000000..92d6ca2da --- /dev/null +++ b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/pi05_preprocess.py @@ -0,0 +1,212 @@ +# SPDX-License-Identifier: Apache-2.0 + +from __future__ import annotations + +from typing import Any + +import numpy as np +import torch +import torch.nn.functional as F +from PIL import Image +from transformers import AutoTokenizer + +from sglang.multimodal_gen.configs.pipeline_configs.pi05 import Pi05PipelineConfig +from sglang.multimodal_gen.runtime.vla.observation import VLAObservationBatch + + +def _tensor_from_image(value: Any) -> torch.Tensor: + if isinstance(value, torch.Tensor): + tensor = value.detach() + if tensor.ndim == 4: + if tensor.shape[0] != 1: + raise ValueError("Pi05 v1 expects one observation per request") + tensor = tensor[0] + if tensor.ndim != 3: + raise ValueError(f"Expected image tensor with 3 dims, got {tensor.shape}") + if tensor.shape[0] in (1, 3, 4): + tensor = tensor[:3] + elif tensor.shape[-1] in (1, 3, 4): + tensor = tensor[..., :3].permute(2, 0, 1) + else: + raise ValueError( + f"Could not infer image channels from shape {tensor.shape}" + ) + is_integer = not tensor.is_floating_point() + tensor = tensor.to(dtype=torch.float32) + if is_integer or tensor.max() > 2.0: + tensor = tensor / 255.0 + return tensor + + if isinstance(value, Image.Image): + image = value.convert("RGB") + arr = np.asarray(image, dtype=np.float32) / 255.0 + return torch.from_numpy(arr).permute(2, 0, 1) + + if isinstance(value, (np.ndarray, list)): + arr = np.asarray(value) + if arr.ndim != 3: + raise ValueError(f"Expected HWC image array, got shape {arr.shape}") + tensor = torch.from_numpy(np.ascontiguousarray(arr)) + if tensor.shape[0] in (1, 3, 4): + tensor = tensor[:3] + elif tensor.shape[-1] in (1, 3, 4): + tensor = tensor[..., :3].permute(2, 0, 1) + else: + raise ValueError( + f"Could not infer image channels from shape {tensor.shape}" + ) + is_integer = not tensor.is_floating_point() + tensor = tensor.to(dtype=torch.float32) + if is_integer or tensor.max() > 2.0: + tensor = tensor / 255.0 + return tensor + + raise TypeError(f"Unsupported Pi05 image type: {type(value)}") + + +def _resize_with_pad_image_tensor( + tensor: torch.Tensor, size: tuple[int, int] +) -> torch.Tensor: + height, width = size + if tensor.shape[-2:] == (height, width): + return tensor + _, cur_height, cur_width = tensor.shape + ratio = max(cur_width / width, cur_height / height) + resized_height = int(cur_height / ratio) + resized_width = int(cur_width / ratio) + tensor = F.interpolate( + tensor[None], + size=(resized_height, resized_width), + mode="bilinear", + align_corners=False, + )[0] + pad_h0, rem_h = divmod(height - resized_height, 2) + pad_w0, rem_w = divmod(width - resized_width, 2) + return F.pad( + tensor, + (pad_w0, pad_w0 + rem_w, pad_h0, pad_h0 + rem_h), + mode="constant", + value=0.0, + ) + + +class Pi05Preprocessor: + def __init__(self, config: Pi05PipelineConfig): + self.config = config + self.tokenizer = AutoTokenizer.from_pretrained(config.tokenizer_name) + self.tokenizer.padding_side = "right" + + def _tokenize(self, prompt: list[str], state: torch.Tensor | None): + if state is None: + state_for_prompt = torch.zeros(1, 0, dtype=torch.float32) + else: + state_for_prompt = state.detach().cpu().to(torch.float32) + bins = np.linspace(-1, 1, 256 + 1)[:-1] + state_np = state_for_prompt.numpy() + discretized = np.digitize(state_np, bins=bins) - 1 + full_prompts = [] + for idx, task in enumerate(prompt): + cleaned = task.strip().replace("_", " ").replace("\n", " ") + state_str = " ".join(map(str, discretized[idx])) + full_prompts.append(f"Task: {cleaned}, State: {state_str};\nAction: ") + + encoded = self.tokenizer( + full_prompts, + max_length=self.config.max_token_len, + padding="max_length", + truncation=True, + return_tensors="pt", + ) + return encoded["input_ids"].to(torch.long), encoded["attention_mask"].to( + torch.bool + ) + + def __call__(self, raw_observation: dict[str, Any]) -> VLAObservationBatch: + prompt_value = raw_observation.get("prompt", "") + if isinstance(prompt_value, list): + prompt = [str(x) for x in prompt_value] + else: + prompt = [str(prompt_value)] + if len(prompt) != 1: + raise ValueError("Pi05 v1 expects one prompt per action request") + + raw_images = raw_observation.get("images") or {} + image_masks_in = raw_observation.get("image_masks") or {} + camera_order = tuple( + raw_observation.get("camera_order") or self.config.image_keys + ) + + images: dict[str, torch.Tensor] = {} + image_masks: dict[str, torch.Tensor] = {} + for key in camera_order: + value = raw_images.get(key) + is_present = value is not None and bool(image_masks_in.get(key, True)) + if is_present: + tensor = _tensor_from_image(value) + tensor = _resize_with_pad_image_tensor(tensor, self.config.image_size) + tensor = tensor * 2.0 - 1.0 + else: + channels = 3 + height, width = self.config.image_size + tensor = torch.ones(channels, height, width, dtype=torch.float32) * -1.0 + + images[key] = tensor.unsqueeze(0) + image_masks[key] = torch.tensor([is_present], dtype=torch.bool) + + state = raw_observation.get("state") + state_tensor = None + if state is not None: + state_tensor = torch.as_tensor(state, dtype=torch.float32) + if state_tensor.ndim == 1: + state_tensor = state_tensor.unsqueeze(0) + if state_tensor.shape[0] != 1: + raise ValueError("Pi05 v1 expects one state vector per request") + if state_tensor.shape[-1] > self.config.state_dim: + raise ValueError( + f"Pi05 state dim must be <= {self.config.state_dim}, " + f"got {state_tensor.shape[-1]}" + ) + + noise = raw_observation.get("noise") + noise_tensor = None + if noise is not None: + noise_tensor = torch.as_tensor(noise, dtype=torch.float32) + if noise_tensor.ndim == 2: + noise_tensor = noise_tensor.unsqueeze(0) + expected = (1, self.config.action_horizon, self.config.action_dim) + if tuple(noise_tensor.shape) != expected: + raise ValueError( + f"Pi05 noise must have shape {expected}, " + f"got {tuple(noise_tensor.shape)}" + ) + + tokens = raw_observation.get("tokens") + if tokens is None: + tokens = raw_observation.get("tokenized_prompt") + token_masks = raw_observation.get("token_masks") + if token_masks is None: + token_masks = raw_observation.get("tokenized_prompt_mask") + if tokens is not None: + tokens_tensor = torch.as_tensor(tokens, dtype=torch.long) + if tokens_tensor.ndim == 1: + tokens_tensor = tokens_tensor.unsqueeze(0) + if token_masks is None: + token_masks_tensor = tokens_tensor != self.tokenizer.pad_token_id + else: + token_masks_tensor = torch.as_tensor(token_masks, dtype=torch.bool) + if token_masks_tensor.ndim == 1: + token_masks_tensor = token_masks_tensor.unsqueeze(0) + else: + tokens_tensor, token_masks_tensor = self._tokenize(prompt, state_tensor) + + return VLAObservationBatch( + prompt=prompt, + images=images, + image_masks=image_masks, + state=state_tensor, + noise=noise_tensor, + tokens=tokens_tensor, + token_masks=token_masks_tensor, + batch_size=1, + metadata={"camera_order": camera_order}, + ) diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/vla.py b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/vla.py new file mode 100644 index 000000000..da0d9a722 --- /dev/null +++ b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/vla.py @@ -0,0 +1,496 @@ +# SPDX-License-Identifier: Apache-2.0 + +from __future__ import annotations + +import time +from typing import Any + +import numpy as np +import torch + +from sglang.multimodal_gen.runtime.disaggregation.roles import RoleType +from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import ( + OutputBatch, + Req, +) +from sglang.multimodal_gen.runtime.pipelines_core.stages.base import PipelineStage +from sglang.multimodal_gen.runtime.server_args import ServerArgs +from sglang.multimodal_gen.runtime.vla.observation import ( + collate_vla_observation_batches, +) +from sglang.multimodal_gen.runtime.vla.parallel import ( + broadcast_prefix_context, + broadcast_tensor_from_rank, + get_vla_split_group, +) +from sglang.multimodal_gen.runtime.vla.prefix_cache import ( + PrefixContext, + VLAPrefixCacheManager, + slice_prefix_context, +) + + +def vla_state(batch: Req) -> dict[str, Any]: + """Per-request scratchpad shared by the VLA pipeline stages.""" + + return batch.extra["vla"] + + +def vla_timings(batch: Req) -> dict[str, float]: + return vla_state(batch).setdefault("timings", {}) + + +def vla_options(batch: Req) -> dict[str, Any]: + return vla_state(batch).get("options") or {} + + +def materialize_vla_action_batch( + actions: Any, + action_dim: int, + output_format: str, +) -> Any: + output_format = output_format.lower() + if isinstance(actions, torch.Tensor): + actions_out = actions[..., :action_dim].detach().float().cpu().numpy() + if output_format != "numpy": + actions_out = actions_out.tolist() + elif isinstance(actions, np.ndarray): + actions_out = actions[..., :action_dim].astype(np.float32, copy=False) + if output_format != "numpy": + actions_out = actions_out.tolist() + else: + actions_out = actions + if not isinstance(actions_out, list): + return actions_out + if not actions_out: + return [] + first = actions_out[0] + if isinstance(first, list) and first and isinstance(first[0], list): + actions_out = [[step[:action_dim] for step in sample] for sample in actions_out] + else: + actions_out = [[step[:action_dim] for step in actions_out]] + if output_format == "numpy": + return np.asarray(actions_out, dtype=np.float32) + return actions_out + + +def synchronize_vla_action_tensor(actions: torch.Tensor | None) -> None: + if actions is not None and actions.device.type == "cuda": + torch.cuda.synchronize(actions.device) + + +def _effective_prefix_cache_enabled( + batch: Req, + server_args: ServerArgs, +) -> bool: + options = vla_options(batch) + return bool(options.get("enable_prefix_cache", True)) and bool( + server_args.pipeline_config.enable_global_prefix_cache + ) + + +def _grouped_fingerprint( + batch: Req, + server_args: ServerArgs, +) -> tuple[Any, ...]: + if ( + batch.is_warmup + or get_vla_split_group() is not None + or _effective_prefix_cache_enabled(batch, server_args) + or batch.generator is not None + ): + return ("single", id(batch)) + + observation = vla_state(batch).get("observation_batch") + camera_order = tuple(observation.metadata.get("camera_order", ())) + image_shapes = tuple( + ( + name, + tuple(observation.images[name].shape), + bool(observation.image_masks[name].item()), + ) + for name in camera_order + ) + return ( + "grouped", + camera_order, + image_shapes, + None if observation.state is None else tuple(observation.state.shape), + None if observation.noise is None else tuple(observation.noise.shape), + tuple(observation.tokens.shape), + tuple(observation.token_masks.shape), + batch.action_horizon, + batch.action_dim, + batch.num_inference_steps, + ) + + +class VLAObservationPreprocessStage(PipelineStage): + def __init__(self, preprocessor: Any): + super().__init__() + self.preprocessor = preprocessor + + @property + def role_affinity(self) -> RoleType: + return RoleType.ENCODER + + def forward(self, batch: Req, server_args: ServerArgs) -> Req: + start = time.perf_counter() + state = vla_state(batch) + raw_observation = dict(state.get("observation") or {}) + raw_observation.setdefault("prompt", batch.prompt) + observation = self.preprocessor(raw_observation) + state["observation_batch"] = observation + vla_timings(batch)["preprocess_ms"] = (time.perf_counter() - start) * 1000 + return batch + + +class VLAPrefixEncodingStage(PipelineStage): + def __init__( + self, + policy_model: Any, + prefix_cache: VLAPrefixCacheManager, + ): + super().__init__() + self.policy_model = policy_model + self.prefix_cache = prefix_cache + + @property + def role_affinity(self) -> RoleType: + return RoleType.ENCODER + + def run_grouped_requests( + self, + batches: list[Req], + server_args: ServerArgs, + ) -> list[Req]: + results: list[Req | None] = [None] * len(batches) + for fingerprint, group in self._group_requests_by_fingerprint( + batches, + lambda batch: _grouped_fingerprint(batch, server_args), + ): + group_batches = [batch for _, batch in group] + if len(group_batches) == 1 or fingerprint[0] == "single": + for index, batch in group: + results[index] = self(batch, server_args) + continue + + prefix_start = time.perf_counter() + observations = [ + vla_state(batch)["observation_batch"] for batch in group_batches + ] + grouped_observation = collate_vla_observation_batches(observations) + prefix_context = self.policy_model.encode_prefix(grouped_observation) + prefix_ms = (time.perf_counter() - prefix_start) * 1000 + + for offset, (index, batch) in enumerate(group): + state = vla_state(batch) + state["observation_group"] = grouped_observation + state["prefix_context_group"] = prefix_context + state["prefix_context"] = slice_prefix_context( + prefix_context, + offset, + ) + state["cache"] = { + "hit": False, + "scope": "request", + "prefix_len": prefix_context.prefix_len, + "grouped": True, + "batch_size": len(group_batches), + } + timings = vla_timings(batch) + timings["cache_lookup_ms"] = 0.0 + timings["prefix_ms"] = prefix_ms + results[index] = batch + + if ( + server_args.pipeline_config.empty_cache_after_prefix + and torch.cuda.is_available() + ): + torch.cuda.empty_cache() + + return [result for result in results if result is not None] + + def _recv_prefix_result(self, batch: Req, split: Any) -> Req: + state = vla_state(batch) + state["prefix_context"] = broadcast_prefix_context( + None, + split, + src=split.prefix_root, + ) + state["cache"] = split.broadcast_object_from_rank( + None, + src=split.prefix_root, + ) + timings = split.broadcast_object_from_rank(None, src=split.prefix_root) + vla_timings(batch).update(timings) + return batch + + def _send_prefix_result( + self, + batch: Req, + split: Any, + prefix_context: Any, + ) -> None: + broadcast_prefix_context( + prefix_context, + split, + src=split.prefix_root, + ) + split.broadcast_object_from_rank( + vla_state(batch)["cache"], + src=split.prefix_root, + ) + split.broadcast_object_from_rank( + vla_timings(batch), + src=split.prefix_root, + ) + + def get_cached_context( + self, batch: Req, server_args: ServerArgs, observation: Any + ) -> tuple[str, PrefixContext]: + """try querying the cache for PrefixContext with prefix cache key built from observations and other keys""" + cache_enabled = _effective_prefix_cache_enabled(batch, server_args) + if cache_enabled: + cache_key = self.policy_model.build_prefix_cache_key(observation) + cached_context = self.prefix_cache.get(cache_key) + else: + cache_key = None + cached_context = None + return cache_key, cached_context + + def forward(self, batch: Req, server_args: ServerArgs) -> Req: + state = vla_state(batch) + if batch.is_warmup: + state["prefix_context"] = None + state["cache"] = {"hit": False, "warmup": True} + return batch + + split = get_vla_split_group() + if split is not None and not split.is_prefix_rank: + return self._recv_prefix_result(batch, split) + + observation = state["observation_batch"] + cache_start = time.perf_counter() + cache_enabled = _effective_prefix_cache_enabled(batch, server_args) + + # 1. try querying the per-request LRU prefix kv cache + cache_key, cached_context = self.get_cached_context( + batch, server_args, observation + ) + + vla_timings(batch)["cache_lookup_ms"] = ( + time.perf_counter() - cache_start + ) * 1000 + + # 2. prepare VLAState + if cached_context is not None: + state["prefix_context"] = cached_context + state["cache"] = { + "hit": True, + "scope": "global", + "mode": "exact", + "prefix_len": cached_context.prefix_len, + } + if split is not None: + self._send_prefix_result(batch, split, cached_context) + return batch + + prefix_start = time.perf_counter() + + # 3. run encoding + prefix_context = self.policy_model.encode_prefix(observation) + if cache_key is not None: + prefix_context.cache_key_digest = cache_key + + vla_timings(batch)["prefix_ms"] = (time.perf_counter() - prefix_start) * 1000 + state["prefix_context"] = prefix_context + state["cache"] = { + "hit": False, + "scope": "global" if cache_enabled else "request", + "mode": "exact" if cache_enabled else "disabled", + "prefix_len": prefix_context.prefix_len, + } + + # 4. update prefix kv cache + if cache_key is not None: + self.prefix_cache.put(cache_key, prefix_context) + if split is not None: + self._send_prefix_result(batch, split, prefix_context) + if ( + server_args.pipeline_config.empty_cache_after_prefix + and torch.cuda.is_available() + ): + torch.cuda.empty_cache() + return batch + + +class VLAActionDenoisingStage(PipelineStage): + def __init__(self, policy_model: Any): + super().__init__() + self.policy_model = policy_model + + @property + def role_affinity(self) -> RoleType: + return RoleType.DENOISER + + def run_grouped_requests( + self, + batches: list[Req], + server_args: ServerArgs, + ) -> list[Req]: + results: list[Req | None] = [None] * len(batches) + + def action_fingerprint(batch: Req) -> tuple[Any, ...]: + prefix_context = vla_state(batch).get("prefix_context_group") + if prefix_context is None: + return ("single", id(batch)) + return ( + "grouped", + id(prefix_context), + batch.num_inference_steps, + str(vla_options(batch).get("output_format") or "list"), + ) + + for _, group in self._group_requests_by_fingerprint( + batches, + action_fingerprint, + ): + group_batches = [batch for _, batch in group] + prefix_context = vla_state(group_batches[0]).get("prefix_context_group") + if len(group_batches) == 1 or prefix_context is None: + for index, batch in group: + results[index] = self(batch, server_args) + continue + + start = time.perf_counter() + options = vla_options(group_batches[0]) + observation = vla_state(group_batches[0])["observation_group"] + actions = self.policy_model.sample_actions( + observation, + prefix_context, + noise=observation.noise, + num_steps=group_batches[0].num_inference_steps, + use_cuda_graph=bool(options.get("enable_cuda_graph", True)), + generator=None, + ) + synchronize_vla_action_tensor(actions) + actions_out = materialize_vla_action_batch( + actions, + server_args.pipeline_config.output_action_dim, + str(options.get("output_format") or "list"), + ) + action_ms = (time.perf_counter() - start) * 1000 + parallel_info = self.policy_model.action_parallel_info(prefix_context) + + for offset, (index, batch) in enumerate(group): + state = vla_state(batch) + vla_timings(batch)["action_denoise_ms"] = action_ms + state["parallel"] = parallel_info + state["actions_output"] = actions_out[offset] + results[index] = batch + + return [result for result in results if result is not None] + + def forward(self, batch: Req, server_args: ServerArgs) -> Req: + start = time.perf_counter() + state = vla_state(batch) + observation = state.get("observation_batch") + split = get_vla_split_group() + prefix_context = state.get("prefix_context") + should_run_action = ( + split is None or self.policy_model.should_run_action_denoise(prefix_context) + ) + parallel_info = self.policy_model.action_parallel_info(prefix_context) + if batch.is_warmup: + actions = ( + self.policy_model.warmup_actions(batch_size=1) + if should_run_action + else None + ) + elif should_run_action: + # broadcast PrefixContext from action root rank to action ranks + options = vla_options(batch) + noise = observation.noise if observation is not None else None + actions = self.policy_model.sample_actions( + observation, + prefix_context, + noise=noise, + num_steps=batch.num_inference_steps, + use_cuda_graph=bool(options.get("enable_cuda_graph", True)), + generator=batch.generator, + ) + synchronize_vla_action_tensor(actions) + else: + actions = None + + if split is not None: + if should_run_action: + vla_timings(batch)["action_denoise_ms"] = ( + time.perf_counter() - start + ) * 1000 + actions = broadcast_tensor_from_rank( + actions, + split, + src=split.action_root, + device=self.policy_model.device, + ) + timings = split.broadcast_object_from_rank( + vla_timings(batch) if should_run_action else None, + src=split.action_root, + ) + vla_timings(batch).update(timings) + parallel_info = split.broadcast_object_from_rank( + parallel_info if should_run_action else None, + src=split.action_root, + ) + else: + vla_timings(batch)["action_denoise_ms"] = ( + time.perf_counter() - start + ) * 1000 + state["parallel"] = parallel_info + state["actions"] = actions + return batch + + +class VLAActionPostprocessStage(PipelineStage): + @property + def role_affinity(self) -> RoleType: + return RoleType.DENOISER + + def forward(self, batch: Req, server_args: ServerArgs) -> OutputBatch: + start = time.perf_counter() + state = vla_state(batch) + action_dim = server_args.pipeline_config.output_action_dim + options = vla_options(batch) + actions_out = state.get("actions_output") + if actions_out is None: + action_batch = materialize_vla_action_batch( + state["actions"], + action_dim, + str(options.get("output_format") or "list"), + ) + actions_out = ( + action_batch[0] if isinstance(action_batch, list) else action_batch + ) + if isinstance(action_batch, np.ndarray): + actions_out = action_batch[0] + + payload = { + "request_id": batch.request_id, + "actions": actions_out, + } + payload["parameters"] = {"num_inference_steps": batch.num_inference_steps} + if options.get("return_timing", True): + timings = dict(vla_timings(batch)) + timings["postprocess_ms"] = (time.perf_counter() - start) * 1000 + payload["timings"] = timings + if not batch.is_warmup: + payload["cache"] = state.get("cache", {}) + if state.get("parallel") is not None: + payload["parallel"] = state["parallel"] + + return OutputBatch( + output=[payload], + metrics=batch.metrics, + ) diff --git a/python/sglang/multimodal_gen/runtime/server_args/auto_tune.py b/python/sglang/multimodal_gen/runtime/server_args/auto_tune.py index df005777a..99504537a 100644 --- a/python/sglang/multimodal_gen/runtime/server_args/auto_tune.py +++ b/python/sglang/multimodal_gen/runtime/server_args/auto_tune.py @@ -398,6 +398,8 @@ class ServerArgsAutoTuner: def _default_layerwise_components_for_unset_placement(self) -> list[str]: args = self.server_args + if args.pipeline_config.task_type.is_action_gen(): + return [] if ( args.is_arg_explicitly_set("layerwise_offload_components") or args.dit_layerwise_offload is True diff --git a/python/sglang/multimodal_gen/runtime/server_args/server_args.py b/python/sglang/multimodal_gen/runtime/server_args/server_args.py index eae08df88..eb86fcdbd 100644 --- a/python/sglang/multimodal_gen/runtime/server_args/server_args.py +++ b/python/sglang/multimodal_gen/runtime/server_args/server_args.py @@ -615,6 +615,17 @@ class ServerArgs(DisaggServerArgsMixin): # CPU platform does not need offload return + if self.pipeline_config.task_type.is_action_gen(): + if self.dit_cpu_offload is None: + self.dit_cpu_offload = False + if self.text_encoder_cpu_offload is None: + self.text_encoder_cpu_offload = False + if self.image_encoder_cpu_offload is None: + self.image_encoder_cpu_offload = False + if self.vae_cpu_offload is None: + self.vae_cpu_offload = False + return + # TODO: to be handled by each platform if current_platform.get_device_total_memory() / BYTES_PER_GB < 30: logger.info( @@ -1106,6 +1117,20 @@ class ServerArgs(DisaggServerArgsMixin): self.use_fsdp_inference = False self.dit_layerwise_offload = False self.layerwise_offload_components = None + if ( + self.dit_cpu_offload + or self.text_encoder_cpu_offload + or self.image_encoder_cpu_offload + or self.vae_cpu_offload + ): + logger.warning( + "Disabling component CPU offload on MPS because CPU-to-MPS " + "module relocation can produce invalid diffusion outputs." + ) + self.dit_cpu_offload = False + self.text_encoder_cpu_offload = False + self.image_encoder_cpu_offload = False + self.vae_cpu_offload = False def is_arg_explicitly_set(self, arg_name: str) -> bool: return arg_name in self._explicit_arg_names @@ -1283,12 +1308,14 @@ class ServerArgs(DisaggServerArgsMixin): ), ) parser.add_argument( + "--pipeline", "--pipeline-class-name", + dest="pipeline_class_name", type=str, default=ServerArgs.pipeline_class_name, help=( - "Override pipeline class selection from model_index.json. " - "Must match a registered pipeline_name." + "Advanced override for pipeline class selection from the model registry " + "or model_index.json. Must match a registered pipeline_name." ), ) # attention @@ -2173,7 +2200,7 @@ class ServerArgs(DisaggServerArgsMixin): # Create a set of argument names that were present on the command line. # This handles both styles: '--arg=value' and '--arg value'. - provided_arg_names = set() + provided_arg_names = set(getattr(args, "_sglang_explicit_arg_names", ())) for arg in raw_argv: if arg.startswith("--"): # For '--arg=value', this gets 'arg'; for '--arg', this also gets 'arg'. @@ -2192,6 +2219,8 @@ class ServerArgs(DisaggServerArgsMixin): # Populate provided_args if the argument from the namespace was on the command line. for k, v in vars(args).items(): + if k.startswith("_sglang_"): + continue if k in provided_arg_names: provided_args[k] = v diff --git a/python/sglang/multimodal_gen/runtime/server_warmup.py b/python/sglang/multimodal_gen/runtime/server_warmup.py index e744665f5..1e5bacbb0 100644 --- a/python/sglang/multimodal_gen/runtime/server_warmup.py +++ b/python/sglang/multimodal_gen/runtime/server_warmup.py @@ -17,6 +17,7 @@ from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger from sglang.multimodal_gen.runtime.warmup_request_builder import ( build_warmup_reqs, should_include_warmup_image, + supports_synthetic_warmup, ) logger = init_logger(__name__) @@ -85,13 +86,19 @@ def is_realtime_serving(server_args: ServerArgs) -> bool: def should_run_synthetic_server_warmup(server_args: ServerArgs) -> bool: - return should_run_server_warmup(server_args) and not is_realtime_serving( - server_args + return ( + should_run_server_warmup(server_args) + and supports_synthetic_warmup(server_args) + and not is_realtime_serving(server_args) ) def should_run_explicit_client_warmup(server_args: ServerArgs) -> bool: - return server_args.warmup and server_args.warmup_resolutions is not None + return ( + server_args.warmup + and server_args.warmup_resolutions is not None + and supports_synthetic_warmup(server_args) + ) def format_warmup_req(req_or_group: Any) -> str: @@ -102,9 +109,12 @@ def format_warmup_req(req_or_group: Any) -> str: if req is None: return prefix - shape = f"{req.width}x{req.height}" - if req.num_frames is not None and req.num_frames > 1: - shape += f"x{req.num_frames}f" + width = getattr(req, "width", None) + height = getattr(req, "height", None) + shape = "action" if width is None or height is None else f"{width}x{height}" + num_frames = getattr(req, "num_frames", None) + if num_frames is not None and num_frames > 1: + shape += f"x{num_frames}f" default_steps = req.extra.get("cache_dit_num_inference_steps") if default_steps is not None and default_steps != req.num_inference_steps: diff --git a/python/sglang/multimodal_gen/runtime/vla/__init__.py b/python/sglang/multimodal_gen/runtime/vla/__init__.py new file mode 100644 index 000000000..f837e01ce --- /dev/null +++ b/python/sglang/multimodal_gen/runtime/vla/__init__.py @@ -0,0 +1,3 @@ +# SPDX-License-Identifier: Apache-2.0 + +"""Shared VLA runtime contracts and execution infrastructure.""" diff --git a/python/sglang/multimodal_gen/runtime/vla/denoise_cuda_graph.py b/python/sglang/multimodal_gen/runtime/vla/denoise_cuda_graph.py new file mode 100644 index 000000000..15913720f --- /dev/null +++ b/python/sglang/multimodal_gen/runtime/vla/denoise_cuda_graph.py @@ -0,0 +1,213 @@ +# SPDX-License-Identifier: Apache-2.0 + +from __future__ import annotations + +from dataclasses import dataclass +from typing import Any, Callable + +import torch + +from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger +from sglang.multimodal_gen.runtime.vla.prefix_cache import ( + PrefixContext, + VLADensePrefixCache, +) +from sglang.srt.distributed.device_communicators.pynccl_allocator import ( + set_graph_pool_id, +) +from sglang.srt.model_executor.runner_utils.pool import ( + get_or_create_global_graph_memory_pool, +) + +logger = init_logger(__name__) + + +@dataclass(frozen=True) +class VLADenoiseGraphSignature: + batch_size: int + prefix_len: int + action_horizon: int + action_dim: int + dtype: str + parallel_layout: str + + +@dataclass +class _CapturedDenoiseGraph: + graph: torch.cuda.CUDAGraph + static_prefix_context: PrefixContext + static_x_t: torch.Tensor + static_timestep: torch.Tensor + static_output: torch.Tensor + current_context_id: int | None = None + current_context_digest: str | None = None + + +def _clone_past_key_values(past_key_values: Any) -> Any: + return VLADensePrefixCache( + tuple( + (keys.detach().clone(), values.detach().clone(), sliding_window) + for keys, values, sliding_window in past_key_values + ) + ) + + +def _copy_past_key_values_(dst: Any, src: Any) -> None: + for (dst_keys, dst_values, _), (src_keys, src_values, _) in zip( + dst, src, strict=True + ): + dst_keys.copy_(src_keys) + dst_values.copy_(src_values) + + +def _clone_prefix_context(prefix_context: PrefixContext) -> PrefixContext: + return PrefixContext( + past_key_values=_clone_past_key_values(prefix_context.past_key_values), + prefix_pad_masks=prefix_context.prefix_pad_masks.detach().clone(), + prefix_len=prefix_context.prefix_len, + layout=dict(prefix_context.layout), + cache_key_digest=prefix_context.cache_key_digest, + ) + + +def _copy_prefix_context_(dst: PrefixContext, src: PrefixContext) -> None: + dst.prefix_pad_masks.copy_(src.prefix_pad_masks) + _copy_past_key_values_(dst.past_key_values, src.past_key_values) + dst.cache_key_digest = src.cache_key_digest + + +class VLADenoiseGraphRunner: + """Full CUDA graph runner for one VLA action-denoise step. + + Each signature owns fixed input and output buffers. This does not use + diffusion BCG and does not capture prefix encoding or token decode. + """ + + def __init__(self, enabled: bool = True): + self.enabled = enabled + self._captured: dict[VLADenoiseGraphSignature, _CapturedDenoiseGraph] = {} + self._disabled_signatures: set[VLADenoiseGraphSignature] = set() + self._capture_stream: torch.cuda.Stream | None = None + self._graph_pool: Any = None + + def _sync_context_if_needed( + self, + captured: _CapturedDenoiseGraph, + prefix_context: PrefixContext, + ) -> None: + context_id = id(prefix_context.past_key_values) + context_digest = prefix_context.cache_key_digest + if ( + context_digest is not None + and captured.current_context_digest == context_digest + ): + captured.current_context_id = context_id + return + if captured.current_context_id == context_id: + return + _copy_prefix_context_(captured.static_prefix_context, prefix_context) + captured.current_context_id = context_id + captured.current_context_digest = context_digest + + def _capture( + self, + signature: VLADenoiseGraphSignature, + step_fn: Callable[..., torch.Tensor], + prefix_context: PrefixContext, + x_t: torch.Tensor, + timestep: torch.Tensor, + ) -> _CapturedDenoiseGraph: + static_prefix_context = _clone_prefix_context(prefix_context) + static_x_t = x_t.detach().clone() + static_timestep = timestep.detach().clone() + + device_module = torch.get_device_module(x_t.device) + if self._capture_stream is None: + self._capture_stream = device_module.Stream(device=x_t.device) + if self._graph_pool is None: + self._graph_pool = get_or_create_global_graph_memory_pool(device_module) + set_graph_pool_id(self._graph_pool) + + # warm up lazy kernels and workspaces before capture + device_module.synchronize() + with device_module.stream(self._capture_stream), torch.inference_mode(): + step_fn( + static_prefix_context, + static_x_t, + static_timestep, + ) + self._capture_stream.synchronize() + + graph = torch.cuda.CUDAGraph() + with ( + device_module.graph( + cuda_graph=graph, + pool=self._graph_pool, + stream=self._capture_stream, + ), + torch.inference_mode(), + ): + static_output = step_fn( + static_prefix_context, + static_x_t, + static_timestep, + ) + self._capture_stream.synchronize() + + captured = _CapturedDenoiseGraph( + graph=graph, + static_prefix_context=static_prefix_context, + static_x_t=static_x_t, + static_timestep=static_timestep, + static_output=static_output, + current_context_id=id(prefix_context.past_key_values), + current_context_digest=prefix_context.cache_key_digest, + ) + self._captured[signature] = captured + logger.info( + "Captured VLA denoise CUDA graph: batch=%d prefix=%d action=%dx%d " + "dtype=%s", + signature.batch_size, + signature.prefix_len, + signature.action_horizon, + signature.action_dim, + signature.dtype, + ) + return captured + + def capture_or_run( + self, + signature: VLADenoiseGraphSignature, + step_fn: Callable[..., torch.Tensor], + prefix_context: PrefixContext, + x_t: torch.Tensor, + timestep: torch.Tensor, + ) -> torch.Tensor: + if not self.enabled or signature in self._disabled_signatures: + return step_fn(prefix_context, x_t, timestep) + + if x_t.device.type != "cuda": + return step_fn(prefix_context, x_t, timestep) + + captured = self._captured.get(signature) + try: + if captured is None: + captured = self._capture( + signature, step_fn, prefix_context, x_t, timestep + ) + captured.graph.replay() + else: + self._sync_context_if_needed(captured, prefix_context) + captured.static_x_t.copy_(x_t) + captured.static_timestep.copy_(timestep) + captured.graph.replay() + return captured.static_output + except Exception: + self._disabled_signatures.add(signature) + self._captured.pop(signature, None) + logger.warning( + "VLA denoise CUDA graph disabled for signature %s", + signature, + exc_info=True, + ) + return step_fn(prefix_context, x_t, timestep) diff --git a/python/sglang/multimodal_gen/runtime/vla/observation.py b/python/sglang/multimodal_gen/runtime/vla/observation.py new file mode 100644 index 000000000..d639ad753 --- /dev/null +++ b/python/sglang/multimodal_gen/runtime/vla/observation.py @@ -0,0 +1,76 @@ +# SPDX-License-Identifier: Apache-2.0 + +from __future__ import annotations + +from dataclasses import dataclass, field +from typing import Any + +import torch + +from sglang.srt.managers.mm_utils import tensor_hash + + +@dataclass +class VLAObservationBatch: + prompt: list[str] + images: dict[str, torch.Tensor] + image_masks: dict[str, torch.Tensor] + state: torch.Tensor | None + noise: torch.Tensor | None + tokens: torch.Tensor + token_masks: torch.Tensor + batch_size: int + metadata: dict[str, Any] = field(default_factory=dict) + + +def tensor_fingerprint(tensor: torch.Tensor) -> str: + """Hash tensor content with SRT's CPU/CUDA implementation.""" + + shape = ",".join(str(dim) for dim in tensor.shape) + return f"{tensor.dtype}:{shape}:{tensor_hash(tensor):016x}" + + +def collate_vla_observation_batches( + observations: list[VLAObservationBatch], +) -> VLAObservationBatch: + first = observations[0] + camera_order = tuple(first.metadata.get("camera_order", ())) + images = { + name: torch.cat([obs.images[name] for obs in observations], dim=0) + for name in camera_order + } + image_masks = { + name: torch.cat([obs.image_masks[name] for obs in observations], dim=0) + for name in camera_order + } + states = [obs.state for obs in observations] + noises = [obs.noise for obs in observations] + if any(item is None for item in states) and not all( + item is None for item in states + ): + raise ValueError("Cannot collate mixed VLA state presence") + if any(item is None for item in noises) and not all( + item is None for item in noises + ): + raise ValueError("Cannot collate mixed VLA noise presence") + state = ( + None + if states[0] is None + else torch.cat([item for item in states if item is not None], dim=0) + ) + noise = ( + None + if noises[0] is None + else torch.cat([item for item in noises if item is not None], dim=0) + ) + return VLAObservationBatch( + prompt=[prompt for obs in observations for prompt in obs.prompt], + images=images, + image_masks=image_masks, + state=state, + noise=noise, + tokens=torch.cat([obs.tokens for obs in observations], dim=0), + token_masks=torch.cat([obs.token_masks for obs in observations], dim=0), + batch_size=len(observations), + metadata={"camera_order": camera_order}, + ) diff --git a/python/sglang/multimodal_gen/runtime/vla/parallel.py b/python/sglang/multimodal_gen/runtime/vla/parallel.py new file mode 100644 index 000000000..108f37a53 --- /dev/null +++ b/python/sglang/multimodal_gen/runtime/vla/parallel.py @@ -0,0 +1,161 @@ +# SPDX-License-Identifier: Apache-2.0 + +from __future__ import annotations + +from dataclasses import dataclass + +import torch +import torch.distributed as dist + +from sglang.multimodal_gen.runtime.distributed import ( + get_sp_group, + model_parallel_is_initialized, +) +from sglang.multimodal_gen.runtime.distributed.group_coordinator import GroupCoordinator +from sglang.multimodal_gen.runtime.distributed.parallel_state import get_world_rank +from sglang.multimodal_gen.runtime.vla.prefix_cache import ( + PrefixContext, + VLADensePrefixCache, +) + + +@dataclass(frozen=True) +class VLASplitGroup: + """Runtime view for VLA prefix/action split execution. + + This reuses the existing SP group as the coordination group. It is not a + separate parallel topology: + 1. `prefix_root` computes/fetches PrefixContext and broadcasts it once. + 2. `action_root` owns fallback action denoise and initial noise broadcast. + 3. `action_ranks` may all participate in action SP when the policy allows it. + + All rank fields are global ranks; GroupCoordinator APIs take group-local + ranks, so call `group_rank_for` before collective helpers. + """ + + group: GroupCoordinator + prefix_root: int + action_root: int + action_ranks: tuple[int, ...] + rank: int + + @property + def is_prefix_rank(self) -> bool: + return self.rank == self.prefix_root + + @property + def is_action_rank(self) -> bool: + return self.rank in self.action_ranks + + @property + def uses_action_sp(self) -> bool: + return len(self.action_ranks) > 1 + + def group_rank_for(self, global_rank: int) -> int: + return self.group.ranks.index(global_rank) + + def broadcast_object_from_rank(self, obj, *, src: int): + return self.group.broadcast_object( + obj if self.rank == src else None, + src=self.group_rank_for(src), + ) + + +def get_vla_split_group() -> VLASplitGroup | None: + if not dist.is_available() or not dist.is_initialized(): + return None + if not model_parallel_is_initialized(): + return None + group = get_sp_group() + if group.world_size <= 1: + return None + # v1 maps the split view onto SP: first rank does prefix encode, last rank + # is the action fallback/root, and all SP ranks are eligible action ranks. + return VLASplitGroup( + group=group, + prefix_root=group.ranks[0], + action_root=group.ranks[-1], + action_ranks=tuple(group.ranks), + rank=get_world_rank(), + ) + + +def broadcast_tensor_from_rank( + tensor: torch.Tensor | None, + split: VLASplitGroup, + *, + src: int, + device: torch.device, +) -> torch.Tensor | None: + payload = ( + {"is_none": tensor is None, "tensor": tensor} if split.rank == src else None + ) + payload = split.group.broadcast_tensor_dict( + payload, + src=split.group_rank_for(src), + ) + if payload["is_none"]: + return None + output = payload["tensor"] + if output.device != device: + output = output.to(device) + return output + + +def broadcast_prefix_context( + context: PrefixContext | None, + split: VLASplitGroup, + *, + src: int, +) -> PrefixContext | None: + if split.rank == src and context is None: + payload = {"is_none": True} + elif split.rank == src: + prefix_pad_masks = context.prefix_pad_masks + prefix_pad_masks_is_bool = prefix_pad_masks.dtype == torch.bool + if prefix_pad_masks_is_bool: + prefix_pad_masks = prefix_pad_masks.to(torch.uint8) + payload = { + "is_none": False, + "prefix_pad_masks": prefix_pad_masks, + "prefix_pad_masks_is_bool": prefix_pad_masks_is_bool, + "prefix_len": context.prefix_len, + "layout": dict(context.layout), + "cache_key_digest": context.cache_key_digest, + "num_layers": len(context.past_key_values), + } + for i, (keys, values, sliding_window) in enumerate(context.past_key_values): + payload[f"layer_{i}_keys"] = keys + payload[f"layer_{i}_values"] = values + payload[f"layer_{i}_sliding_window"] = sliding_window + else: + payload = None + + payload = split.group.broadcast_tensor_dict( + payload, + src=split.group_rank_for(src), + ) + if payload["is_none"]: + return None + + kv_layers = [] + for i in range(int(payload["num_layers"])): + kv_layers.append( + ( + payload[f"layer_{i}_keys"], + payload[f"layer_{i}_values"], + payload[f"layer_{i}_sliding_window"], + ) + ) + + prefix_pad_masks = payload["prefix_pad_masks"] + if payload.get("prefix_pad_masks_is_bool"): + prefix_pad_masks = prefix_pad_masks.to(torch.bool) + + return PrefixContext( + past_key_values=VLADensePrefixCache(tuple(kv_layers)), + prefix_pad_masks=prefix_pad_masks, + prefix_len=int(payload["prefix_len"]), + layout=dict(payload["layout"]), + cache_key_digest=payload["cache_key_digest"], + ) diff --git a/python/sglang/multimodal_gen/runtime/vla/prefix_cache.py b/python/sglang/multimodal_gen/runtime/vla/prefix_cache.py new file mode 100644 index 000000000..f659c5386 --- /dev/null +++ b/python/sglang/multimodal_gen/runtime/vla/prefix_cache.py @@ -0,0 +1,171 @@ +# SPDX-License-Identifier: Apache-2.0 + +from __future__ import annotations + +import hashlib +import json +from collections import OrderedDict +from collections.abc import Iterable +from dataclasses import dataclass, field +from typing import Any + +import torch + + +@dataclass +class PrefixContext: + """Request-local observation K/V reused by every action denoise step. + + The optional digest identifies an exact server-level cache entry; suffix K/V + is step-dependent and never becomes part of this context. + """ + + past_key_values: Any + prefix_pad_masks: torch.Tensor + prefix_len: int + layout: dict[str, Any] = field(default_factory=dict) + cache_key_digest: str | None = None + + +class VLADensePrefixCache: + """a lightweight and naive dense per-layer K/V container for prefix fill and suffix attention. + + Mutable instances collect prefix K/V layer by layer. Read-only instances + prepend that fixed K/V to the current suffix K/V without changing storage. + """ + + def __init__( + self, + layers: Iterable[tuple[torch.Tensor, torch.Tensor, Any]] | None = None, + *, + read_only: bool = False, + ): + # cached_keys, cached_values, sliding_window + self.layers = list(layers or ()) + self.read_only = read_only + + def __iter__(self): + return iter(self.layers) + + def __len__(self) -> int: + return len(self.layers) + + def __getitem__(self, layer_idx: int): + return self.layers[layer_idx] + + def get_seq_length(self) -> int: + return 0 if not self.layers else int(self.layers[0][0].shape[-2]) + + def get_prefix(self, layer_idx: int) -> tuple[torch.Tensor, torch.Tensor]: + prefix_keys, prefix_values, _ = self.layers[layer_idx] + return prefix_keys, prefix_values + + def update( + self, + key_states: torch.Tensor, + value_states: torch.Tensor, + layer_idx: int, + ) -> tuple[torch.Tensor, torch.Tensor]: + """update the cache with fresh kv from each layer, return the appended prefix kv""" + if self.read_only: + prefix_keys, prefix_values = self.get_prefix(layer_idx) + return ( + torch.cat([prefix_keys, key_states], dim=-2), + torch.cat([prefix_values, value_states], dim=-2), + ) + + if layer_idx == len(self.layers): + self.layers.append((key_states, value_states, None)) + return key_states, value_states + if layer_idx > len(self.layers): + raise IndexError(f"Invalid VLA prefix cache layer: {layer_idx}") + cached_keys, cached_values, sliding_window = self.layers[layer_idx] + key_states = torch.cat([cached_keys, key_states], dim=-2) + value_states = torch.cat([cached_values, value_states], dim=-2) + self.layers[layer_idx] = (key_states, value_states, sliding_window) + return key_states, value_states + + +def slice_prefix_context(context: PrefixContext, index: int) -> PrefixContext: + return PrefixContext( + past_key_values=VLADensePrefixCache( + tuple( + ( + keys[index : index + 1], + values[index : index + 1], + sliding_window, + ) + for keys, values, sliding_window in context.past_key_values + ) + ), + prefix_pad_masks=context.prefix_pad_masks[index : index + 1], + prefix_len=context.prefix_len, + layout=dict(context.layout), + cache_key_digest=context.cache_key_digest, + ) + + +class VLAPrefixCacheManager: + """Bounded exact-match LRU for server-level VLA PrefixContext reuse. + + Partial-match prefix cache does not work well VLA scenario (with multiple combinations of keys). + + Request-local denoise reuse does not go through this cache. Partial-prefix + K/V reuse is invalid for VLA prefix blocks that use full attention. + """ + + def __init__(self, max_entries: int = 128): + self.max_entries = max(0, int(max_entries)) + self._cache: OrderedDict[str, PrefixContext] = OrderedDict() + + @staticmethod + def make_key( + *, + model_revision: str, + tokenizer_id: str, + camera_order: tuple[str, ...], + image_hashes: dict[str, str], + token_digest: str, + token_mask_digest: str, + masks: dict[str, bool], + positions_version: str, + dtype: str, + parallel_layout_version: str, + cache_namespace: str = "vla", + ) -> str: + # hash the effective prefix inputs plus runtime compatibility dimensions + payload = { + "cache_namespace": cache_namespace, + "model_revision": model_revision, + "tokenizer_id": tokenizer_id, + "camera_order": list(camera_order), + "image_hashes": image_hashes, + "token_digest": token_digest, + "token_mask_digest": token_mask_digest, + "masks": masks, + "positions_version": positions_version, + "dtype": dtype, + "parallel_layout_version": parallel_layout_version, + } + serialized = json.dumps( + payload, + sort_keys=True, + separators=(",", ":"), + default=str, + ) + return hashlib.sha256(serialized.encode("utf-8")).hexdigest() + + def get(self, key: str) -> PrefixContext | None: + context = self._cache.get(key) + if context is not None: + self._cache.move_to_end(key) + return context + + def put(self, key: str, context: PrefixContext) -> None: + if self.max_entries == 0: + return + if len(self._cache) >= self.max_entries and key not in self._cache: + self._cache.popitem(last=False) + context.cache_key_digest = key + self._cache[key] = context + self._cache.move_to_end(key) diff --git a/python/sglang/multimodal_gen/runtime/warmup_request_builder.py b/python/sglang/multimodal_gen/runtime/warmup_request_builder.py index 4cae0821f..44bd6ae73 100644 --- a/python/sglang/multimodal_gen/runtime/warmup_request_builder.py +++ b/python/sglang/multimodal_gen/runtime/warmup_request_builder.py @@ -1,5 +1,5 @@ # SPDX-License-Identifier: Apache-2.0 -"""Build synthetic diffusion warmup requests. +"""Build synthetic generation warmup requests. Default server warmup should cover a representative serving path before the first real request, without copying user traffic. It starts from the model's @@ -17,10 +17,7 @@ from copy import copy from typing import Any from sglang.multimodal_gen.configs.pipeline_configs.base import ModelTaskType -from sglang.multimodal_gen.configs.sample.sampling_params import ( - DataType, - SamplingParams, -) +from sglang.multimodal_gen.configs.sample.sampling_params import SamplingParams from sglang.multimodal_gen.registry import get_pipeline_config_classes from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import Req from sglang.multimodal_gen.runtime.server_args import ( @@ -146,7 +143,7 @@ def _fallback_warmup_resolution(server_args: ServerArgs) -> tuple[int, int]: def _is_video_warmup_task(server_args: ServerArgs) -> bool: - return server_args.pipeline_config.task_type.data_type() == DataType.VIDEO + return server_args.pipeline_config.task_type.is_video_gen() def _warmup_resolution_alignment(server_args: ServerArgs) -> int: @@ -234,7 +231,7 @@ def _resolve_warmup_num_frames( *, server_based_warmup: bool, ) -> int: - num_frames = sampling_defaults.num_frames + num_frames = getattr(sampling_defaults, "num_frames", 1) if ( not server_based_warmup or not _is_video_warmup_task(server_args) @@ -247,9 +244,9 @@ def _resolve_warmup_num_frames( def _effective_cfg_scale(sampling_defaults: SamplingParams) -> float | None: - if sampling_defaults.true_cfg_scale is not None: + if getattr(sampling_defaults, "true_cfg_scale", None) is not None: return sampling_defaults.true_cfg_scale - return sampling_defaults.guidance_scale + return getattr(sampling_defaults, "guidance_scale", None) def _resolve_warmup_steps( @@ -295,6 +292,8 @@ def should_include_warmup_image( server_args: ServerArgs, server_based_warmup: bool ) -> bool: task_type = server_args.pipeline_config.task_type + if not supports_synthetic_warmup(server_args): + return False if not task_type.accepts_image_input(): return False if task_type.requires_image_input(): @@ -306,6 +305,11 @@ def should_include_warmup_image( return True +def supports_synthetic_warmup(server_args: ServerArgs) -> bool: + task_type = server_args.pipeline_config.task_type + return task_type.is_visual_gen() or task_type.is_mesh_gen() + + def build_warmup_reqs( server_args: ServerArgs, *, @@ -315,6 +319,8 @@ def build_warmup_reqs( server_based_warmup: bool = False, ) -> list[Req]: task_type = server_args.pipeline_config.task_type + if not supports_synthetic_warmup(server_args): + return [] sampling_defaults = get_model_sampling_defaults(server_args) if warmup_resolutions is None: @@ -327,7 +333,7 @@ def build_warmup_reqs( else: resolutions = [parse_size(resolution) for resolution in warmup_resolutions] - negative_prompt: Any = sampling_defaults.negative_prompt + negative_prompt: Any = getattr(sampling_defaults, "negative_prompt", None) cfg_scale = _effective_cfg_scale(sampling_defaults) warmup_steps = _resolve_warmup_steps( server_args, diff --git a/python/sglang/multimodal_gen/test/server/gpu_cases.py b/python/sglang/multimodal_gen/test/server/gpu_cases.py index 754fafc39..a88baffc9 100644 --- a/python/sglang/multimodal_gen/test/server/gpu_cases.py +++ b/python/sglang/multimodal_gen/test/server/gpu_cases.py @@ -30,6 +30,7 @@ from sglang.multimodal_gen.test.server.testcase_configs import ( MULTI_FRAME_I2I_sampling_params, MULTI_IMAGE_TI2I_sampling_params, MULTI_IMAGE_TI2I_UPLOAD_sampling_params, + PI05_ACTION_CI_sampling_params, SANA_WM_TI2V_CI_sampling_params, T2I_sampling_params, T2V_sampling_params, @@ -102,6 +103,16 @@ ONE_GPU_CASES: list[DiffusionTestCase] = [ run_models_api_check=False, run_t2v_input_reference_check=False, ), + DiffusionTestCase( + "pi05_action_http", + DiffusionServerArgs( + model_path="lerobot/pi05_base", + ), + PI05_ACTION_CI_sampling_params, + run_perf_check=False, + run_component_accuracy_check=False, + run_t2v_input_reference_check=False, + ), DiffusionTestCase( "flux_image_t2i", DiffusionServerArgs(model_path=DEFAULT_FLUX_1_DEV_MODEL_NAME_FOR_TEST), diff --git a/python/sglang/multimodal_gen/test/server/test_server_common.py b/python/sglang/multimodal_gen/test/server/test_server_common.py index 0be5e5a03..c2b609403 100644 --- a/python/sglang/multimodal_gen/test/server/test_server_common.py +++ b/python/sglang/multimodal_gen/test/server/test_server_common.py @@ -7,6 +7,7 @@ Each collected request prints a performance log before validation. from __future__ import annotations +import json import os import queue import threading @@ -14,6 +15,7 @@ import time from pathlib import Path from typing import Any, Callable +import numpy as np import openai import pytest import requests @@ -47,8 +49,11 @@ from sglang.multimodal_gen.test.test_utils import ( SGL_TEST_FILES_CI_DATA_REVISION, _consistency_gt_filenames, _get_consistency_gt_dir, + action_gt_exists, compare_with_gt, extract_key_frames_from_video, + get_action_consistency_gt_candidates, + get_action_consistency_gt_remote_files, get_consistency_gt_candidates, get_consistency_gt_remote_files, get_consistency_threshold_path, @@ -56,6 +61,7 @@ from sglang.multimodal_gen.test.test_utils import ( get_dynamic_server_port, gt_exists, image_bytes_to_numpy, + load_action_consistency_gt, load_consistency_gt, save_consistency_failure_artifact, wait_for_req_perf_record, @@ -593,6 +599,10 @@ class DiffusionServerBase: ) return + if case.server_args.modality == "action": + self._validate_action_consistency(case, content) + return + num_gpus = case.server_args.num_gpus is_video = case.server_args.modality == "video" output_format = case.sampling_params.output_format @@ -727,6 +737,88 @@ Pinned revision used by this check: {SGL_TEST_FILES_CI_DATA_REVISION} f"max_mean_abs_diff={result.max_mean_abs_diff:.4f})" ) + def _extract_action_array( + self, + payload: dict[str, Any], + expected_horizon: int, + expected_dim: int, + ) -> np.ndarray: + action = payload["data"][0]["action"] + values = action["values"] + assert action["shape"] == [expected_horizon, expected_dim] + array = np.asarray(values, dtype=np.float32) + assert array.shape == (expected_horizon, expected_dim) + assert np.isfinite(array).all() + return array + + def _validate_action_consistency( + self, + case: DiffusionTestCase, + content: bytes, + ) -> None: + payload = json.loads(content.decode("utf-8")) + expected_horizon = int(case.sampling_params.extras.get("action_horizon", 50)) + expected_dim = int(case.sampling_params.extras.get("action_dim", 32)) + output = self._extract_action_array(payload, expected_horizon, expected_dim) + + num_gpus = case.server_args.num_gpus + if not action_gt_exists(case.id, num_gpus): + names = ", ".join(get_action_consistency_gt_candidates(case.id, num_gpus)) + logger.error(f""" +--- MISSING ACTION GROUND TRUTH DETECTED --- +GT action JSON not found for '{case.id}'. + +Add the expected file to sgl-project/ci-data in diffusion-ci/consistency_gt/sglang_generated/ with naming: + Action: {case.id}_{{n}}gpu.json + +For this case, expected file(s): {names} + +Repository: https://github.com/sgl-project/ci-data (path: diffusion-ci/consistency_gt/sglang_generated/, with optional platform subdirectories such as 5090/) +Pinned revision used by this check: {SGL_TEST_FILES_CI_DATA_REVISION} +""") + pytest.fail( + f"GT action JSON not found for {case.id}. See logs for instructions to add GT." + ) + + gt_payload = load_action_consistency_gt(case.id, num_gpus) + gt = self._extract_action_array(gt_payload, expected_horizon, expected_dim) + abs_diff = np.abs(output - gt) + max_abs_diff = float(abs_diff.max()) + mean_abs_diff = float(abs_diff.mean()) + max_abs_threshold = float( + case.sampling_params.extras.get("action_max_abs_diff_threshold", 0.05) + ) + mean_abs_threshold = float( + case.sampling_params.extras.get("action_mean_abs_diff_threshold", 0.005) + ) + + if max_abs_diff > max_abs_threshold or mean_abs_diff > mean_abs_threshold: + gt_remote_info = "\n".join( + f" - {filename}: {url}" + for filename, url in get_action_consistency_gt_remote_files( + case.id, + num_gpus, + ) + ) + pytest.fail( + f"Action consistency check failed for {case.id}:\n" + f" max_abs_diff={max_abs_diff:.6f} " + f"(threshold {max_abs_threshold:.6f})\n" + f" mean_abs_diff={mean_abs_diff:.6f} " + f"(threshold {mean_abs_threshold:.6f})\n" + f" Compared GT files and links:\n{gt_remote_info}" + ) + + logger.info( + "[Consistency] %s: PASSED action GT check " + "(shape=%sx%s, max_abs_diff=%.6f, mean_abs_diff=%.6f)", + case.id, + expected_horizon, + expected_dim, + max_abs_diff, + mean_abs_diff, + ) + def _save_gt_output( self, case: DiffusionTestCase, @@ -749,6 +841,12 @@ Pinned revision used by this check: {SGL_TEST_FILES_CI_DATA_REVISION} num_gpus = case.server_args.num_gpus is_video = case.server_args.modality == "video" + if case.server_args.modality == "action": + output_path = out_dir / f"{case.id}_{num_gpus}gpu.json" + output_path.write_bytes(content) + logger.info(f"Saved GT action JSON: {output_path}") + return + if is_video: # realtime consistency uses websocket raw frames to avoid lossy mp4 drift frames = pop_realtime_key_frames(case.id) diff --git a/python/sglang/multimodal_gen/test/server/test_server_utils.py b/python/sglang/multimodal_gen/test/server/test_server_utils.py index 7fbd13238..4b1b65dc4 100644 --- a/python/sglang/multimodal_gen/test/server/test_server_utils.py +++ b/python/sglang/multimodal_gen/test/server/test_server_utils.py @@ -854,6 +854,7 @@ VALIDATOR_REGISTRY = { "default": PerformanceValidator, "video": VideoPerformanceValidator, "mesh": MeshValidator, + "action": PerformanceValidator, } @@ -1501,8 +1502,109 @@ def get_generate_fn( pytest.fail(f"{case_id}: mesh generation timed out after {max_wait}s") + def generate_action(case_id, client) -> tuple[str, bytes]: + """VLA action generation using /v1/actions/generations.""" + import numpy as np + import requests as http_requests + + extra = dict(sampling_params.extras) + action_horizon = int(extra.get("action_horizon", 50)) + action_dim = int(extra.get("action_dim", 32)) + state_dim = int(extra.get("state_dim", action_dim)) + image_size = int(extra.get("image_size", 64)) + camera_order = tuple( + extra.get( + "camera_order", + ("base_0_rgb", "left_wrist_0_rgb", "right_wrist_0_rgb"), + ) + ) + + def tensor_payload(array): + return { + "dtype": str(array.dtype), + "shape": list(array.shape), + "values": array.tolist(), + } + + def image_payload(camera_index: int): + y = np.arange(image_size, dtype=np.uint16)[:, None] + x = np.arange(image_size, dtype=np.uint16)[None, :] + image = np.stack( + ( + (x + camera_index * 17) % 256 + np.zeros_like(y), + (y + camera_index * 29) % 256 + np.zeros_like(x), + (x + y + camera_index * 41) % 256, + ), + axis=-1, + ) + return tensor_payload(image.astype(np.uint8)) + + rng = np.random.default_rng(int(extra.get("seed", 0))) + request_id = f"{case_id}-{int(time.time() * 1000)}" + payload = { + "request_id": request_id, + "model": model_path, + "input": { + "task": sampling_params.prompt or "pick up the blue block", + "observation": { + "images": { + camera: image_payload(index) + for index, camera in enumerate(camera_order) + }, + "camera_order": list(camera_order), + "state": tensor_payload( + np.linspace(-0.5, 0.5, state_dim, dtype=np.float32) + ), + "noise": tensor_payload( + rng.standard_normal((action_horizon, action_dim)).astype( + np.float32 + ) + ), + }, + }, + "parameters": { + "action_horizon": action_horizon, + "action_dim": action_dim, + "num_inference_steps": int(extra.get("num_inference_steps", 2)), + }, + "runtime": { + "return_timing": True, + "prefix_cache": bool(extra.get("enable_prefix_cache", False)), + "cuda_graph": bool(extra.get("enable_cuda_graph", True)), + "output_format": "list", + }, + } + + base_url = str(client.base_url).rstrip("/") + endpoint = ( + f"{base_url}/actions/generations" + if base_url.endswith("/v1") + else f"{base_url}/v1/actions/generations" + ) + response = http_requests.post(endpoint, json=payload, timeout=600) + if response.status_code != 200: + pytest.fail(f"{case_id}: action generation failed: {response.text}") + + body = response.json() + action = body["data"][0]["action"] + if action["shape"] != [action_horizon, action_dim]: + pytest.fail( + f"{case_id}: action shape mismatch: {action['shape']} " + f"!= {[action_horizon, action_dim]}" + ) + values = action["values"] + if not all( + isinstance(value, (int, float)) and np.isfinite(value) + for row in values + for value in row + ): + pytest.fail(f"{case_id}: action response contains non-finite values") + return body["id"], response.content + if modality == "3d": fn = generate_mesh + elif modality == "action": + fn = generate_action elif modality == "video": if sampling_params.realtime_num_chunks is not None: fn = generate_realtime_video diff --git a/python/sglang/multimodal_gen/test/server/testcase_configs.py b/python/sglang/multimodal_gen/test/server/testcase_configs.py index 54cf3f589..bc5330431 100644 --- a/python/sglang/multimodal_gen/test/server/testcase_configs.py +++ b/python/sglang/multimodal_gen/test/server/testcase_configs.py @@ -167,7 +167,7 @@ class DiffusionServerArgs: """Configuration for a single model/scenario test case.""" model_path: str # HF repo or local path - modality: str | None = None # auto-inferred: "image" or "video" or "3d" + modality: str | None = None # auto-inferred: "image", "video", "3d", or "action" custom_validator: str | None = None # auto-derived unless explicitly overridden # resources @@ -208,6 +208,8 @@ class DiffusionServerArgs: self.custom_validator = "video" elif self.modality == "3d": self.custom_validator = "mesh" + elif self.modality == "action": + self.custom_validator = "action" @lru_cache(maxsize=None) @@ -219,6 +221,8 @@ def _infer_modality_from_model_path(model_path: str) -> str: task_type = model_info.pipeline_config_cls.task_type if task_type == ModelTaskType.I2M: return "3d" + if task_type.is_action_gen(): + return "action" if task_type.is_image_gen(): return "image" return "video" @@ -357,6 +361,23 @@ LINGBOT_WORLD_REALTIME_sampling_params = DiffusionSamplingParams( ) +PI05_ACTION_CI_sampling_params = DiffusionSamplingParams( + prompt="pick up the blue block", + extras={ + "action_horizon": 50, + "action_dim": 32, + "state_dim": 32, + "image_size": 64, + "num_inference_steps": 2, + "seed": 0, + "enable_prefix_cache": False, + "enable_cuda_graph": True, + "action_max_abs_diff_threshold": 0.05, + "action_mean_abs_diff_threshold": 0.005, + }, +) + + def sample_step_indices( step_map: dict[int, float], fractions: Sequence[float] ) -> list[int]: @@ -646,6 +667,8 @@ def get_default_sampling_params_for_model_task( return TI2V_sampling_params if task_type == ModelTaskType.I2M: return HUNYUAN3D_SHAPE_sampling_params + if task_type.is_action_gen(): + return PI05_ACTION_CI_sampling_params raise ValueError(f"No default sampling params for model task {task_type!r}") diff --git a/python/sglang/multimodal_gen/test/single_test_file/test_pi05_e2e.py b/python/sglang/multimodal_gen/test/single_test_file/test_pi05_e2e.py new file mode 100644 index 000000000..7f968e939 --- /dev/null +++ b/python/sglang/multimodal_gen/test/single_test_file/test_pi05_e2e.py @@ -0,0 +1,212 @@ +# SPDX-License-Identifier: Apache-2.0 + +from __future__ import annotations + +import json +import os +import statistics +import time +from pathlib import Path + +import numpy as np +import pytest + +from sglang.multimodal_gen.runtime.entrypoints.diffusion_generator import DiffGenerator + +pytestmark = pytest.mark.skipif( + os.getenv("SGLANG_RUN_PI05_E2E") != "1", + reason="set SGLANG_RUN_PI05_E2E=1 to run Pi0.5 GPU e2e tests", +) + +_MODEL_PATH = os.getenv("SGLANG_PI05_E2E_MODEL", "lerobot/pi05_base") +_CAMERA_ORDER = ("base_0_rgb", "left_wrist_0_rgb", "right_wrist_0_rgb") + + +def _env_int(name: str, default: int) -> int: + return int(os.getenv(name, str(default))) + + +def _env_float(name: str) -> float | None: + value = os.getenv(name) + return None if value is None else float(value) + + +def _image(camera_index: int) -> np.ndarray: + height = width = _env_int("SGLANG_PI05_E2E_IMAGE_SIZE", 224) + y = np.arange(height, dtype=np.uint16)[:, None] + x = np.arange(width, dtype=np.uint16)[None, :] + image = np.stack( + ( + (x + camera_index * 17) % 256 + np.zeros_like(y), + (y + camera_index * 29) % 256 + np.zeros_like(x), + (x + y + camera_index * 41) % 256, + ), + axis=-1, + ) + return image.astype(np.uint8) + + +def _action_request_kwargs(tag: str) -> dict: + action_horizon = _env_int("SGLANG_PI05_E2E_ACTION_HORIZON", 50) + action_dim = _env_int("SGLANG_PI05_E2E_ACTION_DIM", 32) + rng = np.random.default_rng(_env_int("SGLANG_PI05_E2E_NOISE_SEED", 0)) + prompt = os.getenv("SGLANG_PI05_E2E_PROMPT", "pick up the blue block") + return { + "prompt": f"{prompt} [{tag}]", + "images": {name: _image(idx) for idx, name in enumerate(_CAMERA_ORDER)}, + "camera_order": list(_CAMERA_ORDER), + "state": np.linspace( + -0.5, + 0.5, + _env_int("SGLANG_PI05_E2E_STATE_DIM", 32), + dtype=np.float32, + ), + "noise": rng.standard_normal((action_horizon, action_dim)).astype(np.float32), + "action_horizon": action_horizon, + "action_dim": action_dim, + "num_inference_steps": _env_int("SGLANG_PI05_E2E_NUM_STEPS", 2), + "return_timing": True, + "enable_prefix_cache": True, + "enable_cuda_graph": os.getenv("SGLANG_PI05_E2E_CUDA_GRAPH", "1") != "0", + } + + +@pytest.fixture(scope="module") +def pi05_generator(): + num_gpus = _env_int("SGLANG_PI05_E2E_NUM_GPUS", 1) + kwargs = { + "model_path": _MODEL_PATH, + "num_gpus": num_gpus, + "warmup": False, + "trust_remote_code": False, + } + if num_gpus > 1: + kwargs.update( + { + "sp_degree": _env_int("SGLANG_PI05_E2E_SP_DEGREE", num_gpus), + "ulysses_degree": _env_int( + "SGLANG_PI05_E2E_ULYSSES_DEGREE", + num_gpus, + ), + "ring_degree": _env_int("SGLANG_PI05_E2E_RING_DEGREE", 1), + } + ) + generator = DiffGenerator.from_pretrained(local_mode=True, **kwargs) + try: + yield generator + finally: + generator.shutdown() + + +def _actions(output: dict) -> np.ndarray: + return np.asarray(output["actions"], dtype=np.float32) + + +def _assert_action_output( + output: dict, *, expect_cache_hit: bool | None = None +) -> None: + actions = _actions(output) + assert actions.shape[0] == _env_int("SGLANG_PI05_E2E_ACTION_HORIZON", 50) + expected_output_dim = os.getenv("SGLANG_PI05_E2E_OUTPUT_ACTION_DIM") + if expected_output_dim is not None: + assert actions.shape[1] == int(expected_output_dim) + else: + assert 0 < actions.shape[1] <= _env_int("SGLANG_PI05_E2E_ACTION_DIM", 32) + assert np.isfinite(actions).all() + timings = output.get("timings") or {} + assert timings.get("preprocess_ms", 0.0) >= 0.0 + assert timings.get("prefix_ms", 0.0) >= 0.0 + assert timings.get("action_denoise_ms", 0.0) > 0.0 + assert timings.get("postprocess_ms", 0.0) >= 0.0 + + cache = output.get("cache") or {} + if expect_cache_hit is not None: + assert bool(cache.get("hit")) is expect_cache_hit + + parallel = output.get("parallel") or {} + num_gpus = _env_int("SGLANG_PI05_E2E_NUM_GPUS", 1) + assert bool(parallel.get("split_group", False)) is (num_gpus > 1) + if num_gpus > 1: + assert int(parallel["world_size"]) == num_gpus + assert parallel["prefix_root"] == 0 + assert parallel["action_root"] == num_gpus - 1 + assert bool(parallel.get("action_sequence_parallel")) is True + + +def test_pi05_python_action_e2e(pi05_generator): + output = pi05_generator.generate_action(_action_request_kwargs("e2e")) + _assert_action_output(output) + + +def test_pi05_python_action_consistency(pi05_generator): + first = pi05_generator.generate_action(_action_request_kwargs("consistency")) + second = pi05_generator.generate_action(_action_request_kwargs("consistency")) + _assert_action_output(first, expect_cache_hit=False) + _assert_action_output(second, expect_cache_hit=True) + + first_actions = _actions(first) + second_actions = _actions(second) + np.testing.assert_allclose( + first_actions, + second_actions, + rtol=_env_float("SGLANG_PI05_E2E_CONSISTENCY_RTOL") or 1e-3, + atol=_env_float("SGLANG_PI05_E2E_CONSISTENCY_ATOL") or 1e-3, + ) + + gt_path = os.getenv("SGLANG_PI05_E2E_CONSISTENCY_GT") + if gt_path is None: + return + path = Path(gt_path) + if os.getenv("SGLANG_PI05_E2E_UPDATE_GT") == "1": + path.parent.mkdir(parents=True, exist_ok=True) + np.save(path, first_actions) + else: + np.testing.assert_allclose( + first_actions, + np.load(path), + rtol=_env_float("SGLANG_PI05_E2E_GT_RTOL") or 1e-2, + atol=_env_float("SGLANG_PI05_E2E_GT_ATOL") or 1e-2, + ) + + +def test_pi05_python_action_perf(pi05_generator): + for _ in range(_env_int("SGLANG_PI05_E2E_PERF_WARMUP", 1)): + pi05_generator.generate_action(_action_request_kwargs("perf")) + + records = [] + for _ in range(_env_int("SGLANG_PI05_E2E_PERF_REPEAT", 3)): + start = time.perf_counter() + output = pi05_generator.generate_action(_action_request_kwargs("perf")) + wall_ms = (time.perf_counter() - start) * 1000 + _assert_action_output(output, expect_cache_hit=True) + records.append( + { + "wall_ms": wall_ms, + "timings": output.get("timings") or {}, + "parallel": output.get("parallel") or {}, + } + ) + + wall = [record["wall_ms"] for record in records] + denoise = [record["timings"].get("action_denoise_ms", 0.0) for record in records] + summary = { + "model": _MODEL_PATH, + "num_gpus": _env_int("SGLANG_PI05_E2E_NUM_GPUS", 1), + "num_steps": _env_int("SGLANG_PI05_E2E_NUM_STEPS", 2), + "median_wall_ms": statistics.median(wall), + "median_action_denoise_ms": statistics.median(denoise), + "records": records, + } + + dump_path = os.getenv("SGLANG_PI05_E2E_PERF_DUMP") + if dump_path: + path = Path(dump_path) + path.parent.mkdir(parents=True, exist_ok=True) + path.write_text(json.dumps(summary, indent=2), encoding="utf-8") + + max_wall_ms = _env_float("SGLANG_PI05_E2E_MAX_WALL_MS") + if max_wall_ms is not None: + assert summary["median_wall_ms"] <= max_wall_ms + max_denoise_ms = _env_float("SGLANG_PI05_E2E_MAX_ACTION_DENOISE_MS") + if max_denoise_ms is not None: + assert summary["median_action_denoise_ms"] <= max_denoise_ms diff --git a/python/sglang/multimodal_gen/test/test_utils.py b/python/sglang/multimodal_gen/test/test_utils.py index f59ec5093..063a96a6c 100644 --- a/python/sglang/multimodal_gen/test/test_utils.py +++ b/python/sglang/multimodal_gen/test/test_utils.py @@ -1082,6 +1082,30 @@ def get_consistency_gt_candidates( ] +def _action_consistency_gt_filenames(case_id: str, num_gpus: int) -> list[str]: + case_id = get_consistency_gt_case_id(case_id) + return [f"{case_id}_{num_gpus}gpu.json"] + + +def get_action_consistency_gt_candidate_sets( + case_id: str, + num_gpus: int, +) -> list[list[str]]: + candidates = _action_consistency_gt_filenames(case_id, num_gpus) + if _is_ascend_consistency_case(case_id) or current_platform.is_npu(): + return [candidates] + platform = get_consistency_platform() + return [[f"{platform}/{candidate}" for candidate in candidates], candidates] + + +def get_action_consistency_gt_candidates(case_id: str, num_gpus: int) -> list[str]: + return [ + candidate + for candidate_set in get_action_consistency_gt_candidate_sets(case_id, num_gpus) + for candidate in candidate_set + ] + + def get_consistency_gt_remote_files( case_id: str, num_gpus: int, is_video: bool, output_format: str | None = None ) -> list[tuple[str, str]]: @@ -1097,6 +1121,19 @@ def get_consistency_gt_remote_files( ) +def get_action_consistency_gt_remote_files( + case_id: str, num_gpus: int +) -> list[tuple[str, str]]: + files = _find_remote_action_consistency_gt_files(case_id, num_gpus) + if files: + return files + filenames = get_action_consistency_gt_candidates(case_id, num_gpus) + return [ + (filename, f"{SGL_TEST_FILES_CONSISTENCY_GT_BASE}/{filename}") + for filename in filenames + ] + + def _remote_consistency_gt_candidates( base_url: str, case_id: str, @@ -1300,6 +1337,35 @@ def _find_remote_consistency_gt_files( return [] +def _find_remote_action_consistency_gt_files( + case_id: str, + num_gpus: int, +) -> list[tuple[str, str]]: + for filenames in get_action_consistency_gt_candidate_sets(case_id, num_gpus): + for base_url in _remote_consistency_gt_base_urls(case_id): + candidates = [ + (filename, f"{base_url}/{filename}") for filename in filenames + ] + if _is_official_consistency_gt_base_url(base_url): + candidates = [ + (filename, url) + for filename, url in candidates + if _official_consistency_gt_candidate_is_declared(case_id, filename) + ] + if not candidates: + continue + uncertain_candidate = None + for filename, url in candidates: + exists = _remote_file_exists(url) + if exists is True: + return [(filename, url)] + if exists is None and uncertain_candidate is None: + uncertain_candidate = (filename, url) + if uncertain_candidate is not None: + return [uncertain_candidate] + return [] + + def _get_consistency_gt_dir() -> Path | None: """Return the local GT directory when configured.""" d = os.environ.get("SGLANG_CONSISTENCY_GT_DIR") @@ -1320,6 +1386,13 @@ def _get_consistency_gt_cache_key( return f"{platform}:{case_id}:{num_gpus}:{is_video}:{output_format or ''}:{source}" +def _get_action_consistency_gt_cache_key(case_id: str, num_gpus: int) -> str: + gt_dir = _get_consistency_gt_dir() + source = str(gt_dir) if gt_dir is not None else "remote" + platform = get_consistency_platform() + return f"{platform}:{case_id}:{num_gpus}:action:{source}" + + def load_consistency_gt( case_id: str, num_gpus: int, @@ -1398,6 +1471,60 @@ def load_consistency_gt( return loaded_gt +def _load_remote_gt_json(url: str) -> dict[str, Any]: + last_error: Exception | None = None + for _ in range(3): + try: + resp = requests.get(url, timeout=60) + try: + if resp.status_code == 200: + return resp.json() + last_error = FileNotFoundError(f"GT JSON not found: {url}") + if resp.status_code not in (403, 429) and resp.status_code < 500: + break + finally: + resp.close() + except (ValueError, requests.RequestException) as exc: + last_error = exc + raise FileNotFoundError(f"GT JSON not found: {url}") from last_error + + +def load_action_consistency_gt(case_id: str, num_gpus: int) -> dict[str, Any]: + cache_key = _get_action_consistency_gt_cache_key(case_id, num_gpus) + cached = _consistency_gt_cache.get(cache_key) + if cached is not None: + return cached + + gt_dir = _get_consistency_gt_dir() + if gt_dir is not None: + path = None + for fn in get_action_consistency_gt_candidates(case_id, num_gpus): + candidate = gt_dir / fn + if candidate.exists(): + path = candidate + break + if path is None: + candidates = get_action_consistency_gt_candidates(case_id, num_gpus) + raise FileNotFoundError( + f"GT action JSON not found in {gt_dir}. Tried: {', '.join(candidates)}" + ) + with path.open("r", encoding="utf-8") as f: + loaded_gt = json.load(f) + logger.info("Loaded action GT for %s from %s", case_id, path) + else: + remote_files = _find_remote_action_consistency_gt_files(case_id, num_gpus) + if not remote_files: + candidates = get_action_consistency_gt_candidates(case_id, num_gpus) + raise FileNotFoundError( + f"GT action JSON not found for {case_id}. Tried: {', '.join(candidates)}" + ) + loaded_gt = _load_remote_gt_json(remote_files[0][1]) + logger.info("Loaded action GT for %s from %s", case_id, remote_files[0][1]) + + _consistency_gt_cache[cache_key] = loaded_gt + return loaded_gt + + def load_gt_embeddings( case_id: str, num_gpus: int, @@ -1449,6 +1576,26 @@ def gt_exists( return found +def action_gt_exists(case_id: str, num_gpus: int) -> bool: + gt_dir = _get_consistency_gt_dir() + if gt_dir is not None: + return any( + (gt_dir / candidate).exists() + for candidate_set in get_action_consistency_gt_candidate_sets( + case_id, num_gpus + ) + for candidate in candidate_set + ) + + cache_key = _get_action_consistency_gt_cache_key(case_id, num_gpus) + if cache_key in _gt_exists_remote_cache: + return True + found = bool(_find_remote_action_consistency_gt_files(case_id, num_gpus)) + if found: + _gt_exists_remote_cache.add(cache_key) + return found + + def extract_key_frames_from_video( video_bytes: bytes, num_frames: int | None = None, diff --git a/python/sglang/multimodal_gen/test/unit/test_cfg_parallel_warmup.py b/python/sglang/multimodal_gen/test/unit/test_cfg_parallel_warmup.py index d61fb4fa9..61fb3c5ab 100644 --- a/python/sglang/multimodal_gen/test/unit/test_cfg_parallel_warmup.py +++ b/python/sglang/multimodal_gen/test/unit/test_cfg_parallel_warmup.py @@ -41,12 +41,17 @@ from sglang.multimodal_gen.runtime.pipelines_core.stages.image_encoding import ( from sglang.multimodal_gen.runtime.pipelines_core.stages.input_validation import ( InputValidationStage, ) -from sglang.multimodal_gen.runtime.server_warmup import format_warmup_req +from sglang.multimodal_gen.runtime.server_warmup import ( + format_warmup_req, + should_run_explicit_client_warmup, + should_run_synthetic_server_warmup, +) from sglang.multimodal_gen.runtime.warmup_request_builder import ( DEFAULT_PLACEHOLDER_PROMPT, SERVER_WARMUP_IMAGE_FALLBACK_RESOLUTION, build_warmup_reqs, should_include_warmup_image, + supports_synthetic_warmup, ) @@ -642,7 +647,12 @@ class TestWarmupReqCfgParallel(unittest.TestCase): ModelTaskType.I2I: True, ModelTaskType.I2V: True, ModelTaskType.I2M: True, + ModelTaskType.VLA_ACTION: False, } + request_based_expected = { + task_type: task_type.accepts_image_input() for task_type in ModelTaskType + } + request_based_expected[ModelTaskType.VLA_ACTION] = False for task_type in ModelTaskType: server_args = MagicMock() @@ -655,10 +665,84 @@ class TestWarmupReqCfgParallel(unittest.TestCase): ) self.assertEqual( should_include_warmup_image(server_args, server_based_warmup=False), - task_type.accepts_image_input(), + request_based_expected[task_type], task_type.name, ) + def test_action_pipeline_skips_synthetic_warmup_before_sampling_defaults(self): + server_args = MagicMock() + server_args.pipeline_config.task_type = ModelTaskType.VLA_ACTION + + with ( + patch( + "sglang.multimodal_gen.runtime.warmup_request_builder.get_model_sampling_defaults" + ) as get_defaults, + patch( + "sglang.multimodal_gen.runtime.warmup_request_builder._resolve_default_warmup_resolution" + ) as resolve_resolution, + ): + reqs = build_warmup_reqs( + server_args, + warmup_resolutions=None, + server_based_warmup=True, + ) + + self.assertEqual(reqs, []) + get_defaults.assert_not_called() + resolve_resolution.assert_not_called() + + def test_action_pipeline_disables_synthetic_warmup(self): + server_args = MagicMock() + server_args.warmup = True + server_args.server_warmup = True + server_args.warmup_resolutions = ["512x512"] + server_args.pipeline_config.task_type = ModelTaskType.VLA_ACTION + + self.assertFalse(supports_synthetic_warmup(server_args)) + self.assertFalse(should_run_synthetic_server_warmup(server_args)) + self.assertFalse(should_run_explicit_client_warmup(server_args)) + + def test_mesh_pipeline_builds_image_conditioned_warmup(self): + server_args = MagicMock() + server_args.warmup = True + server_args.server_warmup = True + server_args.warmup_steps = 1 + server_args.warmup_resolutions = None + server_args.enable_cfg_parallel = False + server_args.enable_torch_compile = False + server_args.enable_breakable_cuda_graph = False + server_args.backend = "native" + server_args.pipeline_class_name = None + server_args.is_arg_explicitly_set.return_value = False + server_args.pipeline_config = SimpleNamespace( + task_type=ModelTaskType.I2M, + vae_stride=None, + vae_scale_factor=None, + vae_config=None, + ) + + with ( + patch( + "sglang.multimodal_gen.runtime.warmup_request_builder.get_model_sampling_defaults", + return_value=SamplingParams(width=512, height=512), + ), + patch( + "sglang.multimodal_gen.runtime.server_warmup.is_realtime_serving", + return_value=False, + ), + ): + reqs = build_warmup_reqs( + server_args, + warmup_resolutions=None, + warmup_input_path="/tmp/warmup.png", + server_based_warmup=True, + ) + self.assertTrue(should_run_synthetic_server_warmup(server_args)) + + self.assertEqual(len(reqs), 1) + self.assertEqual(reqs[0].data_type, ModelTaskType.I2M.data_type()) + self.assertEqual(reqs[0].image_path, ["/tmp/warmup.png"]) + def test_server_based_warmup_keeps_ti2i_image_input(self): server_args = MagicMock() server_args.warmup_steps = 1 diff --git a/python/sglang/multimodal_gen/test/unit/test_consistency_metrics.py b/python/sglang/multimodal_gen/test/unit/test_consistency_metrics.py index 61b79e9fd..f35f7b70e 100644 --- a/python/sglang/multimodal_gen/test/unit/test_consistency_metrics.py +++ b/python/sglang/multimodal_gen/test/unit/test_consistency_metrics.py @@ -413,6 +413,37 @@ def test_consistency_gt_case_alias_reuses_canonical_filename(monkeypatch): ] +def test_action_gt_candidates_prefer_platform_then_default(monkeypatch): + monkeypatch.setenv(test_utils.CONSISTENCY_PLATFORM_ENV, "h100") + + assert test_utils.get_action_consistency_gt_candidates("unit_action", 1) == [ + "h100/unit_action_1gpu.json", + "unit_action_1gpu.json", + ] + + +def test_remote_action_gt_uses_sglang_generated(monkeypatch): + monkeypatch.setenv(test_utils.CONSISTENCY_PLATFORM_ENV, "h100") + sglang_prefix = test_utils.SGL_TEST_FILES_SGLANG_CONSISTENCY_GT_BASE + "/" + monkeypatch.setattr( + test_utils, + "_remote_file_exists", + lambda url: url.startswith(sglang_prefix), + ) + + files = test_utils._find_remote_action_consistency_gt_files("unit_action", 1) + + assert files == [ + ( + "h100/unit_action_1gpu.json", + ( + f"{test_utils.SGL_TEST_FILES_SGLANG_CONSISTENCY_GT_BASE}" + "/h100/unit_action_1gpu.json" + ), + ) + ] + + def test_threshold_metadata_merges_platform_override(): metadata = test_utils._merge_threshold_metadata( { diff --git a/python/sglang/multimodal_gen/test/unit/test_pi05_action_api.py b/python/sglang/multimodal_gen/test/unit/test_pi05_action_api.py new file mode 100644 index 000000000..c5482db05 --- /dev/null +++ b/python/sglang/multimodal_gen/test/unit/test_pi05_action_api.py @@ -0,0 +1,241 @@ +# SPDX-License-Identifier: Apache-2.0 + +import dataclasses +from types import SimpleNamespace + +import numpy as np + +from sglang.multimodal_gen.configs.pipeline_configs.pi05 import Pi05PipelineConfig +from sglang.multimodal_gen.configs.sample.pi05 import Pi05SamplingParams +from sglang.multimodal_gen.configs.sample.sampling_params import ( + DataType, + SamplingParams, +) +from sglang.multimodal_gen.configs.sample.vla import VLASamplingParams +from sglang.multimodal_gen.runtime.entrypoints.vla.protocol import ( + action_generation_response, + action_metadata, + action_raw_response, + build_action_sampling_params, + pack_msgpack, + unpack_msgpack, +) + + +def _server_args(config: Pi05PipelineConfig | None = None) -> SimpleNamespace: + return SimpleNamespace( + model_id=None, + model_path="lerobot/pi05_base", + output_path=None, + comfyui_mode=False, + num_gpus=1, + tp_size=1, + sp_degree=1, + ulysses_degree=1, + ring_degree=1, + pipeline_config=config or Pi05PipelineConfig(), + ) + + +def test_pi05_uses_vla_sampling_params_not_visual_sampling_params(): + params = Pi05SamplingParams() + field_names = {field.name for field in dataclasses.fields(params)} + + assert isinstance(params, VLASamplingParams) + assert not isinstance(params, SamplingParams) + assert "action_horizon" in field_names + assert "action_dim" in field_names + assert "height" not in field_names + assert "width" not in field_names + assert "fps" not in field_names + assert "negative_prompt" not in field_names + assert "return_frames" not in field_names + assert "diffusers_kwargs" not in field_names + + +def test_action_adjust_skips_visual_image_video_logic(): + params = SamplingParams() + params.num_frames = 0 + params.adjust_frames = True + params.return_file_paths_only = True + + params._adjust(_server_args()) + + assert params.data_type == DataType.ACTION + assert params.num_frames == 1 + assert params.adjust_frames is False + assert params.return_file_paths_only is False + + +def test_action_request_schema_builds_pi05_sampling_params(): + image = np.zeros((8, 8, 3), dtype=np.uint8) + payload = { + "request_id": "action-req-1", + "model": "lerobot/pi05_base", + "input": { + "task": "pick up the block", + "observation": { + "images": { + "base_0_rgb": { + "dtype": "uint8", + "shape": [8, 8, 3], + "values": image.tolist(), + }, + }, + "state": { + "dtype": "float32", + "shape": [32], + "values": np.arange(32, dtype=np.float32).tolist(), + }, + }, + }, + "parameters": { + "action_horizon": 25, + "action_dim": 32, + "num_inference_steps": 4, + }, + "runtime": { + "return_timing": "false", + "prefix_cache": False, + "cuda_graph": "0", + }, + } + + params = build_action_sampling_params(payload, _server_args()) + + assert params.prompt == "pick up the block" + assert params.request_id == "action-req-1" + assert params.action_horizon == 25 + assert params.action_dim == 32 + assert params.num_inference_steps == 4 + assert not params.return_timing + assert not params.enable_prefix_cache + assert not params.enable_cuda_graph + assert params.return_file_paths_only is False + assert params.save_output is False + assert set(params.images) == {"base_0_rgb"} + assert params.images["base_0_rgb"].shape == (8, 8, 3) + assert params.state.shape == (32,) + + vla_state = params.build_request_extra()["vla"] + assert vla_state["observation"]["prompt"] == "pick up the block" + assert not vla_state["options"]["enable_prefix_cache"] + + +def test_openpi_raw_observation_compatibility_fields_are_normalized(): + payload = { + "task": "push the cube", + "observation.images.base_0_rgb": np.ones((4, 4, 3), dtype=np.uint8), + "observation.state": { + "dtype": "float32", + "shape": [32], + "data": [0.25] * 32, + }, + "observation.noise": { + "dtype": "float32", + "shape": [50, 32], + "data": np.zeros((50, 32), dtype=np.float32).tolist(), + }, + "enable_pi_prefix_cache": False, + "enable_pi_cuda_graph": False, + } + + params = build_action_sampling_params(payload, _server_args()) + + assert params.prompt == "push the cube" + assert set(params.images) == {"base_0_rgb"} + assert params.state.shape == (32,) + assert params.noise.shape == (50, 32) + assert not params.enable_prefix_cache + assert not params.enable_cuda_graph + + +def test_action_metadata_reports_policy_shape_and_capabilities(): + config = Pi05PipelineConfig( + image_keys=("front", "wrist"), + image_size=(256, 256), + state_dim=8, + action_horizon=10, + action_dim=32, + output_action_dim=7, + enable_action_cuda_graph=True, + ) + + metadata = action_metadata(_server_args(config)) + + assert metadata["object"] == "action.metadata" + assert metadata["policy_family"] == "pi05" + assert metadata["input"]["image_keys"] == ["front", "wrist"] + assert metadata["input"]["image_size"] == [256, 256] + assert metadata["input"]["state_dim"] == 8 + assert metadata["output"]["action_horizon"] == 10 + assert metadata["output"]["action_dim"] == 7 + assert metadata["output"]["padded_action_dim"] == 32 + assert metadata["runtime"]["materialize_dtype"] == "bf16" + assert metadata["runtime"]["enable_autocast"] is True + assert metadata["runtime"]["parallelism"]["num_gpus"] == 1 + assert metadata["runtime"]["parallelism"]["prefix_strategy"] == "tp" + assert metadata["runtime"]["parallelism"]["action_strategy"] == "sp" + assert metadata["defaults"]["prefix_cache"] is False + assert metadata["capabilities"]["realtime_websocket"] + assert metadata["capabilities"]["openpi_websocket"] + + +def test_action_generation_response_uses_actual_output_parameters(): + output = { + "request_id": "action-response-1", + "actions": [[1.0, 2.0], [3.0, 4.0]], + "parameters": {"num_inference_steps": 3}, + "timings": {"preprocess_ms": 1.5}, + "cache": {"hit": True}, + "parallel": {"split_group": False}, + } + + response = action_generation_response(output, _server_args()) + + assert response["id"] == "action-response-1" + assert response["object"] == "action.generation" + assert response["data"][0]["action"]["shape"] == [2, 2] + assert response["data"][0]["action"]["values"] == output["actions"] + assert response["usage"]["action_horizon"] == 2 + assert response["usage"]["action_dim"] == 2 + assert response["usage"]["denoise_steps"] == 3 + assert response["usage"]["prefix_cache_hit"] is True + assert response["timings"] == output["timings"] + assert response["cache"] == output["cache"] + assert response["parallel"] == output["parallel"] + + +def test_action_raw_response_preserves_policy_payload_shape(): + actions = np.arange(6, dtype=np.float32).reshape(2, 3) + output = { + "actions": actions, + "timings": {"preprocess_ms": 1.5}, + } + + response = action_raw_response(output) + + assert response["actions"] == actions.tolist() + assert response["timings"] == output["timings"] + + +def test_action_raw_response_can_preserve_numpy_for_msgpack(): + actions = np.arange(6, dtype=np.float32).reshape(2, 3) + + response = action_raw_response({"actions": actions}, preserve_numpy=True) + + assert response["actions"] is actions + + +def test_msgpack_roundtrip_preserves_string_keys_and_numpy_payloads(): + payload = { + "task": "pick", + "array": np.arange(6, dtype=np.float32).reshape(2, 3), + "scalar": np.float32(1.25), + } + + decoded = unpack_msgpack(pack_msgpack(payload)) + + assert decoded["task"] == "pick" + np.testing.assert_array_equal(decoded["array"], payload["array"]) + assert decoded["scalar"] == payload["scalar"] diff --git a/python/sglang/multimodal_gen/test/unit/test_pi05_prefix_cache.py b/python/sglang/multimodal_gen/test/unit/test_pi05_prefix_cache.py new file mode 100644 index 000000000..d9a66b720 --- /dev/null +++ b/python/sglang/multimodal_gen/test/unit/test_pi05_prefix_cache.py @@ -0,0 +1,61 @@ +# SPDX-License-Identifier: Apache-2.0 + +import torch + +from sglang.multimodal_gen.runtime.vla.prefix_cache import ( + PrefixContext, + VLAPrefixCacheManager, +) + + +def _context() -> PrefixContext: + return PrefixContext( + past_key_values=("kv",), + prefix_pad_masks=torch.ones(1, 3, dtype=torch.bool), + prefix_len=3, + ) + + +def test_pi05_prefix_cache_full_hit_returns_context(): + manager = VLAPrefixCacheManager(max_entries=4) + key = "a" + context = _context() + + manager.put(key, context) + cached = manager.get(key) + + assert cached is context + assert context.cache_key_digest == key + + +def test_pi05_prefix_cache_different_key_misses(): + manager = VLAPrefixCacheManager(max_entries=4) + manager.put("a", _context()) + + assert manager.get("b") is None + + +def test_pi05_prefix_cache_zero_capacity_does_not_retain_context(): + manager = VLAPrefixCacheManager(max_entries=0) + context = _context() + + manager.put("a", context) + + assert manager.get("a") is None + assert context.cache_key_digest is None + + +def test_pi05_prefix_cache_evicts_least_recently_used_entry(): + manager = VLAPrefixCacheManager(max_entries=2) + context_a = _context() + context_b = _context() + context_c = _context() + manager.put("a", context_a) + manager.put("b", context_b) + assert manager.get("a") is context_a + + manager.put("c", context_c) + + assert manager.get("a") is context_a + assert manager.get("b") is None + assert manager.get("c") is context_c diff --git a/python/sglang/multimodal_gen/test/unit/test_pi05_runtime_helpers.py b/python/sglang/multimodal_gen/test/unit/test_pi05_runtime_helpers.py new file mode 100644 index 000000000..a09f58e4f --- /dev/null +++ b/python/sglang/multimodal_gen/test/unit/test_pi05_runtime_helpers.py @@ -0,0 +1,174 @@ +# SPDX-License-Identifier: Apache-2.0 + +from types import SimpleNamespace + +import torch +from torch import nn + +import sglang.multimodal_gen.runtime.models.vlas.pi05_policy as pi05_policy_module +from sglang.multimodal_gen.configs.pipeline_configs.pi05 import Pi05PipelineConfig +from sglang.multimodal_gen.runtime.models.vlas.pi05_core import ( + Pi05SiglipAttention, + patch_siglip_vision_attention_to_native, +) +from sglang.multimodal_gen.runtime.models.vlas.pi05_policy import ( + Pi05CheckpointManifest, + Pi05PolicyModel, +) +from sglang.multimodal_gen.runtime.vla.denoise_cuda_graph import ( + VLADenoiseGraphRunner, + _CapturedDenoiseGraph, +) +from sglang.multimodal_gen.runtime.vla.parallel import VLASplitGroup +from sglang.multimodal_gen.runtime.vla.prefix_cache import ( + PrefixContext, + VLADensePrefixCache, +) + + +def _prefix_context(value: float, digest: str | None) -> PrefixContext: + keys = torch.full((1, 1, 2, 4), value) + values = torch.full((1, 1, 2, 4), value) + return PrefixContext( + past_key_values=VLADensePrefixCache(((keys, values, None),)), + prefix_pad_masks=torch.ones(1, 2, dtype=torch.bool), + prefix_len=2, + cache_key_digest=digest, + ) + + +def test_vla_split_group_marks_all_action_ranks(): + split = VLASplitGroup( + group=SimpleNamespace(world_size=2), + prefix_root=0, + action_root=1, + action_ranks=(0, 1), + rank=0, + ) + + assert split.is_prefix_rank + assert split.is_action_rank + assert split.uses_action_sp + + +def test_denoise_graph_skips_prefix_copy_for_same_digest(monkeypatch): + runner = VLADenoiseGraphRunner(enabled=True) + static_context = _prefix_context(1.0, "same") + captured = _CapturedDenoiseGraph( + graph=object(), + static_prefix_context=static_context, + static_x_t=torch.empty(1, 2, 4), + static_timestep=torch.empty(1), + static_output=torch.empty(1, 2, 4), + current_context_id=123, + current_context_digest="same", + ) + + def fail_copy(*args, **kwargs): + raise AssertionError("PrefixContext should not be copied on digest hit") + + monkeypatch.setattr( + "sglang.multimodal_gen.runtime.vla.denoise_cuda_graph._copy_prefix_context_", + fail_copy, + ) + + runner._sync_context_if_needed(captured, _prefix_context(2.0, "same")) + + assert captured.static_prefix_context.past_key_values[0][0].eq(1.0).all() + + +def test_runai_direct_gpu_loader_does_not_reject_split_roles(monkeypatch): + class FakeSafeOpen: + def __enter__(self): + return self + + def __exit__(self, exc_type, exc, tb): + return False + + def keys(self): + return ["action.weight"] + + monkeypatch.setattr( + pi05_policy_module, + "safe_open", + lambda *args, **kwargs: FakeSafeOpen(), + ) + + model = Pi05PolicyModel.__new__(Pi05PolicyModel) + model.device = torch.device("cuda") + model.runtime_role = "action" + model.config = Pi05PipelineConfig() + model.manifest = Pi05CheckpointManifest( + model_path="fake", + safetensor_files=["fake.safetensors"], + ) + model._should_read_source_key = lambda key: True + target_state = { + "action.weight": SimpleNamespace(device=SimpleNamespace(type="cuda")), + } + + assert model._should_stream_weights_to_gpu(target_state, {}) + + model.runtime_role = "idle" + assert model._should_stream_weights_to_gpu(target_state, {}) + + +def test_pi05_loader_maps_unfused_prefix_weights_to_parallel_targets(): + q_key = ( + "paligemma_with_expert.paligemma.model.language_model.layers.0." + "self_attn.q_proj.weight" + ) + gate_key = ( + "paligemma_with_expert.paligemma.model.language_model.layers.0." + "mlp.gate_proj.weight" + ) + + assert ( + "paligemma_with_expert.paligemma.model.language_model.layers.0." + "self_attn.qkv_proj.weight", + "q", + ) in Pi05PolicyModel._candidate_target_weights(q_key) + assert ( + "paligemma_with_expert.paligemma.model.language_model.layers.0." + "mlp.gate_up_proj.weight", + 0, + ) in Pi05PolicyModel._candidate_target_weights(gate_key) + + +def test_action_parallel_info_reports_single_rank_without_process_group(): + model = Pi05PolicyModel.__new__(Pi05PolicyModel) + model.runtime_role = "all" + + info = model.action_parallel_info(prefix_context=None) + + assert info == { + "split_group": False, + "runtime_role": "all", + "action_sequence_parallel": False, + } + + +class _FakeSiglipAttention(nn.Module): + def __init__(self): + super().__init__() + self.embed_dim = 8 + self.num_heads = 2 + self.head_dim = 4 + self.scale = self.head_dim**-0.5 + self.dropout = 0.0 + self.q_proj = nn.Linear(self.embed_dim, self.embed_dim) + self.k_proj = nn.Linear(self.embed_dim, self.embed_dim) + self.v_proj = nn.Linear(self.embed_dim, self.embed_dim) + self.out_proj = nn.Linear(self.embed_dim, self.embed_dim) + + +def test_siglip_attention_patch_uses_native_wrapper_once(): + layer = SimpleNamespace(self_attn=_FakeSiglipAttention()) + vision_model = SimpleNamespace(encoder=SimpleNamespace(layers=[layer])) + + patch_siglip_vision_attention_to_native(vision_model) + first = layer.self_attn + patch_siglip_vision_attention_to_native(vision_model) + + assert isinstance(first, Pi05SiglipAttention) + assert layer.self_attn is first diff --git a/python/sglang/utils.py b/python/sglang/utils.py index c7dc1ae78..a04d184cd 100644 --- a/python/sglang/utils.py +++ b/python/sglang/utils.py @@ -32,6 +32,10 @@ from sglang.srt.environ import envs logger = logging.getLogger(__name__) KNOWN_NON_DIFFUSERS_DIFFUSION_MODEL_PATTERNS: dict[str, str] = { + "lerobot/pi05": "Pi05Pipeline", + "lerobot--pi05": "Pi05Pipeline", + "pi05": "Pi05Pipeline", + "pi0.5": "Pi05Pipeline", "hunyuan3d": "Hunyuan3D2Pipeline", "flux.2-dev-nvfp4": "Flux2NvfpPipeline", "comfy-org/ideogram-4": "Ideogram4Nvfp4Pipeline",