Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
Show all changes
29 commits
Select commit Hold shift + click to select a range
1a4dafd
[Diffusion] Add FLUX 3 DiT, video VAE and text encoder
niehen6174 Sep 24, 2026
9ff432f
[Diffusion] Action endpoints: model-owned metadata and OpenPI respons…
niehen6174 Sep 24, 2026
743aa80
[Diffusion] Support FLUX 3 Action robot policies
niehen6174 Sep 24, 2026
6498e6a
[Diffusion] FLUX 3 Action: native FP8r packages
niehen6174 Sep 24, 2026
a352e2c
[Docs] Add FLUX 3 Action cookbook
niehen6174 Sep 24, 2026
6df1f2b
[Diffusion] FLUX 3 Action: drop RoboLab-specific request and response…
niehen6174 Sep 24, 2026
938faaa
[Diffusion] FLUX 3 Action: keep the sampling loop in the denoising st…
niehen6174 Sep 24, 2026
823b942
[Diffusion] FLUX 3 Action: register the UniPC scheduler as a pipeline…
niehen6174 Sep 24, 2026
da17feb
[Diffusion] FLUX 3 Action: build the scheduler in load_modules, drop …
niehen6174 Sep 24, 2026
72f0b52
[Diffusion] FLUX 3 Action: split conditioning into text and observati…
niehen6174 Sep 24, 2026
eef044d
[Diffusion] FLUX 3 Action: support layerwise offload
niehen6174 Sep 24, 2026
1b4d361
[Docs] FLUX 3 Action cookbook: sync with the current implementation
niehen6174 Sep 24, 2026
9d70cd0
[Diffusion] FLUX 3 Action: fused LN-modulate, SwiGLU, residual and QK…
niehen6174 Sep 24, 2026
77abae4
[Diffusion] FLUX 3 Action: accept integer pixel arrays from JSON requ…
niehen6174 Sep 24, 2026
601e2f8
[Diffusion] Rank-local TP loading: fall back when sources span merged…
niehen6174 Sep 24, 2026
a0fd827
[Diffusion] FLUX 3 Action: CFG parallel
niehen6174 Sep 24, 2026
843b117
[Diffusion] FLUX 3 Action: tensor and sequence parallel DiT
niehen6174 Sep 24, 2026
97fdc72
[Diffusion] FLUX 3 Action: drop the unservable so101 entry, add the V…
niehen6174 Sep 24, 2026
4a06af9
[Diffusion] FLUX 3 Action: fall back to FlexAttention when NATTEN is …
niehen6174 Sep 24, 2026
d5ffe8d
[Diffusion] FLUX 3 Action: add a server-side latency benchmark
niehen6174 Sep 24, 2026
7c05150
Merge origin/main into flux3-action
niehen6174 Sep 24, 2026
6ee4e62
[Docs] FLUX 3 Action cookbook: H200 latency and a shared latency prot…
niehen6174 Sep 24, 2026
7a8300d
[Diffusion] FLUX 3 Action: GPU CI case
niehen6174 Sep 24, 2026
7542397
[Diffusion] FLUX 3 Action CI: loosen action thresholds to 0.2 / 0.05
niehen6174 Sep 24, 2026
0915fa4
[Diffusion] CI: bump ci-data-diffusion revision for the FLUX 3 Action GT
niehen6174 Sep 24, 2026
5484e4f
[Diffusion] CI: pin ci-data-diffusion to the reverted GT
niehen6174 Sep 26, 2026
89af3b1
Merge origin/main into flux3-action
niehen6174 Sep 26, 2026
5e910e5
Merge remote-tracking branch 'origin/main' into flux3-action
niehen6174 Sep 27, 2026
24d7949
[Diffusion] CI: pin ci-data-diffusion to the H100 Qwen-Image 2.1 TP2 …
niehen6174 Sep 28, 2026
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
166 changes: 166 additions & 0 deletions docs/cookbook/vla/FLUX/FLUX-3-Action.mdx
Original file line number Diff line number Diff line change
@@ -0,0 +1,166 @@
---
title: FLUX 3 Action
metatags:
description: "Deploy Black Forest Labs FLUX 3 Action robot policies (joint video + action flow matching) with SGLang's native multimodal_gen runtime."
tag: dVLA
---

## 1. Model Introduction

FLUX 3 Action is a robot policy from Black Forest Labs built on a FLUX 3 video diffusion transformer. From the current camera frames, the robot state and a language instruction, it denoises a chunk of future video latents **jointly** with a chunk of continuous actions; only the actions are returned.

The model has three parts:

- **DiT** (`JointSingleSeq`, about 6.6B parameters). Each stream (text, `video`, `video_cond`, the action and state streams) runs through its own 5 mode blocks. The text and all streams then share 28 joint single-stream blocks with per-stream modulation and a 4-axis `(t, h, w, l)` RoPE.
- **Text encoder**: Qwen3-VL-4B. Eight hidden layers are stacked into a 20480-dim context.
- **Video VAE**: a Swin3D neighborhood-attention VAE (NATTEN), 96 latent channels, 32x spatial compression.

The sampler is Cosmos UniPC (order 2). Classifier-free guidance is applied per stream (DROID: `4.0` on video, `1.0` on actions).

SGLang runs it in the native `multimodal_gen` runtime. Work that does not depend on the noised streams is computed once:

- The text context, once per caption, across requests.
- The observation and state streams, once per request.
- The mode blocks of the noised streams, once per step, shared by the conditional and unconditional passes.

Supported checkpoints (`black-forest-labs/flux-3-action-droid`; three 360x640 cameras `wrist`, `left`, `right`; state and action dim 8; 32-action chunks):

| `--model-variant` | Recipe | Weights | Latency on 1x RTX 5090 | Latency on 1x H200 |
| --- | --- | --- | ---: | ---: |
| (default) / `base` | 4 steps, guidance 4.0 (video) / 1.0 (action) | BF16 | 1.14 s | 0.41 s |
| `fp8r` | same | FP8 rowwise | 0.93 s | 0.47 s |
| `gd` | guidance-distilled: 4 steps, no CFG | BF16 | 0.63 s | 0.24 s |
| `gd-fp8r` | same | FP8 rowwise | 0.52 s | 0.27 s |
| `sd` | step-distilled: 1 step | BF16 | 0.18 s | 0.09 s |
| `sd-fp8r` | same | FP8 rowwise | 0.16 s | 0.11 s |

Latency is the median server-side time per request (`server_timing.infer_ms`) over 50 sequential requests after 10 warmup requests, on one GPU with the default serve command, eager mode (no `torch.compile`). Requests go through the OpenPI WebSocket with msgpack numpy images, and the caption is cached. Measure it with `python -m sglang.multimodal_gen.benchmarks.bench_flux3_action --url ws://127.0.0.1:30000` against a running server. On H200 the FP8r packages are slower than BF16; they save memory. LayerNorm + modulation, SwiGLU and the gated residual run on bit-exact fused kernels; QK RMSNorm + RoPE runs on a fused kernel at bf16 rounding level (set `SGLANG_ENABLE_FUSED_QKNORM_ROPE=0` to use the eager path). FP8r packages load their native E4M3 weights (one scale per output row) and run `torch._scaled_mm` with per-token activation scales. This needs an SM89+ GPU.

The frozen text encoder and video VAE are downloaded from the pinned revision of [`black-forest-labs/flux-3-action-base`](https://huggingface.co/black-forest-labs/flux-3-action-base) that the policy config references.

References:

- [FLUX Action](https://github.com/black-forest-labs/flux-action)
- [FLUX 3 Action collection](https://huggingface.co/collections/black-forest-labs/flux-3-action)

## 2. Installation

```bash Command
git clone https://github.com/sgl-project/sglang.git
cd sglang
pip install -e "python[diffusion]"
```

The video VAE runs its neighborhood attention on [NATTEN](https://natten.org) when it is installed. NATTEN is not a dependency of `sglang[diffusion]`; without it the VAE uses a compiled FlexAttention fallback with the same windows (actions stay at the bf16 rounding level). The fallback compiles on the first request (about 3 s on H200) and is then no slower than NATTEN for this single-frame encode (24 ms vs. 36 ms per request on H200). To use NATTEN, pick the wheel that matches your torch and CUDA versions from [whl.natten.org](https://whl.natten.org/).

## 3. Model Deployment

Serve the DROID policy:

```bash Command
sglang serve black-forest-labs/flux-3-action-droid \
--model-type diffusion \
--host 127.0.0.1 \
--port 30000
```

Serve a distilled variant:

```bash Command
sglang serve black-forest-labs/flux-3-action-droid \
--model-type diffusion \
--model-variant gd \
--port 30000
```

A local policy export (a directory containing `manifest.json`, `config.native.json` and `model.safetensors`) is detected automatically when passed as the model path. Pass `--revision` to pin a Hub revision. At startup the policy config and weights are checked against the SHA-256 hashes in `manifest.json`.

Peak GPU memory and per-request latency on one RTX 5090. The resident rows use the same median protocol as the table above. The layerwise-offload rows are from the earlier single-run measurement.

| Flags | Peak memory | Latency |
| --- | ---: | ---: |
| (none) | 23.3 GiB | 1.14 s |
| `--dit-layerwise-offload` | 14.5 GB | 2.51 s |
| `--layerwise-offload-components all` | 9.8 GB | 2.56 s |
| `--model-variant fp8r` | 17.9 GiB | 0.93 s |
| `--model-variant fp8r --dit-layerwise-offload` | 13.2 GB | 1.29 s |

Layerwise offload streams the DiT blocks (and, with `all`, the text encoder and VAE blocks) from host memory, so use it only when the resident configuration does not fit.

### 3.1 Multi-GPU

Three layouts split the DiT across GPUs. They compose as `--num-gpus` = TP size x SP degree x (2 with CFG parallel):

| Flags | What is split | Actions vs 1 GPU |
| --- | --- | --- |
| `--num-gpus 2 --enable-cfg-parallel` | The conditional and unconditional passes | Bit-identical |
| `--num-gpus 2 --tp-size 2` | DiT weights (heads and MLP channels), one all-reduce per block | bf16 rounding level |
| `--num-gpus 2 --sp-degree 2` | The joint sequence and the video mode blocks (K/V gather; add `--ulysses-degree 2` for Ulysses) | bf16 rounding level |

CFG parallel only helps recipes that use guidance (`base`, `fp8r`). The distilled `gd` and `sd` variants run one pass per step, so both GPUs compute the same pass. TP and SP accept the FP8r packages too. Ring attention is not supported.

```bash Command
sglang serve black-forest-labs/flux-3-action-droid \
--model-type diffusion \
--num-gpus 2 \
--enable-cfg-parallel \
--port 30000
```

### 3.2 Action Request Schema

| Field | Type | Description |
| --- | --- | --- |
| `input.task` | string | Language instruction. |
| `input.observation.images` | object | Camera name -> HWC RGB image, uint8 or float in `[0, 1]`. DROID: `wrist`, `left`, `right` (or the LeRobot names `wrist_image_left`, `exterior_image_1_left`, `exterior_image_2_left`). Alternatively send `composite`: the 540x640 image with the wrist camera on top and the two exterior cameras at half resolution below. |
| `input.observation.state` | array | Robot state in dataset units. DROID: 7 joint positions (rad) followed by the gripper position. |
| `parameters.num_inference_steps` | integer, optional | Defaults to the checkpoint recipe. |
| `parameters.guidance_scale` / `guidance_scale_action` | number, optional | Guidance scale on the video / action stream. Defaults to the checkpoint recipe. |
| `parameters.seed` | integer, optional | Noise seed. Defaults to the package's `inference_seed` (`0` for DROID), so a repeated observation returns the same actions. |
| `runtime.prefix_cache` | boolean, optional | `false` bypasses the per-caption text context cache for this request. Defaults to `true`. |
| `runtime.output_format` | `"list"` or `"numpy"`, optional | Use `"numpy"` with msgpack clients. |

The response returns absolute commands of shape `[32, 8]` in the dataset's conventions.

## 4. API Usage

### 4.1 Generic Action HTTP API

```python Example
import numpy as np
import requests

image = np.zeros((360, 640, 3), dtype=np.uint8)
payload = {
"input": {
"task": "put the marker in the cup",
"observation": {
"images": {"wrist": image.tolist(), "left": image.tolist(), "right": image.tolist()},
"state": np.zeros(8, dtype=np.float32).tolist(),
},
},
}
response = requests.post("http://127.0.0.1:30000/v1/actions/generations", json=payload)
action = response.json()["data"][0]["action"]
print(action["shape"]) # [32, 8]
```

`GET /v1/actions/metadata` reports the camera keys, action shape and sampler defaults of the served policy.

### 4.2 OpenPI-Compatible WebSocket

`/openpi/policy` takes one observation per message, with camera images as
`observation.images.<camera>` (`wrist`, `left`, `right`, or the LeRobot names
`wrist_image_left`, `exterior_image_1_left`, `exterior_image_2_left`), the state
as `observation.state` and the instruction as `task` or `prompt`. The response
carries the `[32, 8]` chunk as `actions`. Use the msgpack helpers from the
[Pi0.5 page](/cookbook/vla/OpenPI/Pi0.5) to pack numpy arrays.

## 5. Accuracy

With the same observation and seed, SGLang matches the FLUX Action reference implementation to the bf16 rounding level:

- The DiT agrees to 1e-6 (relative) in fp32.
- The VAE latents and text contexts are bit-identical.
- Across the full 4-step, CFG 4.0 sampling loop, actions differ by at most 0.015 rad (mean 0.003 rad). The reference's own eager and prepared paths differ from each other by 0.015 rad.
- FP8r packages differ from the reference FP8r path by at most 0.025 rad. The reference FP8r path itself differs from BF16 by 0.055 rad.
10 changes: 10 additions & 0 deletions docs/cookbook/vla/intro.mdx
Original file line number Diff line number Diff line change
Expand Up @@ -19,3 +19,13 @@ This section keeps VLA policies separate from the diffusion model cookbook so ro
href="/cookbook/vla/OpenPI/Pi0.5"
/>
</CardGroup>

## FLUX

<CardGroup cols={2}>
<Card
title="FLUX 3 Action"
mode="card"
href="/cookbook/vla/FLUX/FLUX-3-Action"
/>
</CardGroup>
7 changes: 7 additions & 0 deletions docs/docs.json
Original file line number Diff line number Diff line change
Expand Up @@ -1617,6 +1617,13 @@
"pages": [
"cookbook/vla/OpenPI/Pi0.5"
]
},
{
"group": "FLUX",
"tag": "NEW",
"pages": [
"cookbook/vla/FLUX/FLUX-3-Action"
]
}
]
},
Expand Down
5 changes: 5 additions & 0 deletions docs/src/snippets/diffusion/model-catalog.jsx
Original file line number Diff line number Diff line change
Expand Up @@ -254,6 +254,11 @@ export const DiffusionModelCatalog = ({ category }) => {
],
cookbook: "/cookbook/diffusion/Cosmos/Cosmos3",
},
{
name: "FLUX 3 Action",
modelIds: ["black-forest-labs/flux-3-action-droid"],
cookbook: "/cookbook/vla/FLUX/FLUX-3-Action",
},
{
name: "LingBotWorld",
modelIds: [
Expand Down
109 changes: 109 additions & 0 deletions python/sglang/multimodal_gen/benchmarks/bench_flux3_action.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,109 @@
# SPDX-License-Identifier: Apache-2.0
"""Latency of a running FLUX 3 Action server, as reported in the cookbook.

Protocol (keep the cookbook numbers comparable across GPUs):

- one GPU, the default ``sglang serve`` command of the variant, eager mode;
- the OpenPI WebSocket (``/openpi/policy``) with msgpack numpy payloads: three
uint8 360x640 cameras, an 8-dim float32 state and a fixed prompt, so the
caption context is cached after the first request;
- one request at a time; the first ``--warmup`` requests are discarded
(CUDA / JIT / FlexAttention compilation and the caption cache miss);
- the reported latency is the median (and p90) of the server-side
``server_timing.infer_ms`` over the next ``--requests`` requests, which
excludes network and client serialization.

Usage::

sglang serve black-forest-labs/flux-3-action-droid --model-type diffusion --port 30000
python -m sglang.multimodal_gen.benchmarks.bench_flux3_action --url ws://127.0.0.1:30000
"""

from __future__ import annotations

import argparse
import asyncio
import json
import statistics
import time

import numpy as np
import websockets

from sglang.multimodal_gen.runtime.entrypoints.action.protocol import (
pack_msgpack,
unpack_msgpack,
)

CAMERAS = ("wrist", "left", "right")


def _observation(seed: int) -> dict:
rng = np.random.default_rng(seed)
observation = {
f"observation.images.{name}": rng.integers(
0, 256, (360, 640, 3), dtype=np.uint8
)
for name in CAMERAS
}
observation["observation.state"] = rng.uniform(-1, 1, 8).astype(np.float32)
observation["prompt"] = "put the marker in the cup"
return observation


def _percentile(values: list[float], q: float) -> float:
return float(np.percentile(np.asarray(values), q))


async def _run(url: str, warmup: int, requests: int, seed: int) -> dict:
observation = pack_msgpack(_observation(seed))
async with websockets.connect(f"{url}/openpi/policy", max_size=None) as ws:
metadata = unpack_msgpack(await ws.recv())
infer, stages, round_trip = [], {}, []
for i in range(warmup + requests):
start = time.perf_counter()
await ws.send(observation)
response = unpack_msgpack(await ws.recv())
elapsed = (time.perf_counter() - start) * 1000
if i < warmup:
continue
round_trip.append(elapsed)
infer.append(float(response["server_timing"]["infer_ms"]))
for name, value in response["timings"].items():
stages.setdefault(name, []).append(float(value))
return {
"model": metadata.get("model"),
"variant": metadata.get("defaults", {}).get("variant"),
"warmup": warmup,
"requests": requests,
"infer_ms_median": statistics.median(infer),
"infer_ms_p90": _percentile(infer, 90),
"round_trip_ms_median": statistics.median(round_trip),
"stage_ms_median": {k: statistics.median(v) for k, v in stages.items()},
}


def main() -> None:
parser = argparse.ArgumentParser(description=__doc__.split("\n\n")[0])
parser.add_argument("--url", default="ws://127.0.0.1:30000")
parser.add_argument("--warmup", type=int, default=10)
parser.add_argument("--requests", type=int, default=50)
parser.add_argument("--seed", type=int, default=0)
parser.add_argument("--json", action="store_true", help="print the raw result")
args = parser.parse_args()
result = asyncio.run(_run(args.url, args.warmup, args.requests, args.seed))
if args.json:
print(json.dumps(result, indent=2))
return
stages = " ".join(f"{k}={v:.1f}" for k, v in result["stage_ms_median"].items())
print(
f"{result['model']} [{result['variant']}]: "
f"infer {result['infer_ms_median']:.0f} ms median, "
f"{result['infer_ms_p90']:.0f} ms p90 "
f"(round trip {result['round_trip_ms_median']:.0f} ms; "
f"{result['requests']} requests after {result['warmup']} warmup)\n {stages}"
)


if __name__ == "__main__":
main()
67 changes: 67 additions & 0 deletions python/sglang/multimodal_gen/configs/models/dits/flux3.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,67 @@
# SPDX-License-Identifier: Apache-2.0
"""FLUX 3 ``JointSingleSeq`` DiT architecture configuration.

The DiT routes every modality ("content stream") through its own mode blocks
before all streams and the text context share the joint single-stream blocks.
Streams are declared by ``in_channels``; ``sequence`` maps the model inputs
(``x_<name>``) to streams and fixes their order in the joint sequence.
"""

from dataclasses import dataclass, field

from sglang.multimodal_gen.configs.models.dits.base import DiTArchConfig, DiTConfig


@dataclass
class Flux3ArchConfig(DiTArchConfig):
in_channels: dict[str, int] = field(
default_factory=lambda: {"video": 96, "video_cond": 96}
)
sequence: dict[str, str] = field(
default_factory=lambda: {"x_video": "video", "x_video_cond": "video_cond"}
)
vec_in_dim: int | None = 768
context_in_dim: int = 20480
hidden_size: int = 3072
num_attention_heads: int = 24
depth: int = 5
depth_single_blocks: int = 28
axes_dim: tuple[int, ...] = (32, 32, 32, 32)
theta: int = 10000
mlp_ratio: float = 3.0

# Exported policies prefix the DiT tensors with ``dit.``; every block's
# q/k/v/mlp_in projections are fused into one ``qkv_mlp`` weight.
param_names_mapping: dict = field(
default_factory=lambda: {
r"^dit\.(.*)$": r"\1",
r"^(.*)\.q_proj\.(.*)$": (r"\1.qkv_mlp.\2", 0, 4),
r"^(.*)\.k_proj\.(.*)$": (r"\1.qkv_mlp.\2", 1, 4),
r"^(.*)\.v_proj\.(.*)$": (r"\1.qkv_mlp.\2", 2, 4),
r"^(.*)\.mlp_in\.(.*)$": (r"\1.qkv_mlp.\2", 3, 4),
}
)

def __post_init__(self) -> None:
super().__post_init__()
unknown = set(self.sequence.values()) - set(self.in_channels)
if unknown:
raise ValueError(f"sequence names undeclared streams: {sorted(unknown)}")
if self.hidden_size % self.num_attention_heads:
raise ValueError("hidden_size must be divisible by num_attention_heads")
if sum(self.axes_dim) != self.hidden_size // self.num_attention_heads:
raise ValueError("axes_dim must sum to the attention head dim")
self.num_channels_latents = self.in_channels.get("video", 0)

def with_streams(self, extra: dict[str, int]) -> "Flux3ArchConfig":
"""Declare additional streams ``name -> channels`` (appended to the sequence)."""
self.in_channels = {**self.in_channels, **extra}
self.sequence = {**self.sequence, **{f"x_{name}": name for name in extra}}
self.__post_init__()
return self


@dataclass
class Flux3DiTConfig(DiTConfig):
arch_config: DiTArchConfig = field(default_factory=Flux3ArchConfig)
prefix: str = "flux3"
Loading
Loading