diff --git a/docs/cookbook/vla/FLUX/FLUX-3-Action.mdx b/docs/cookbook/vla/FLUX/FLUX-3-Action.mdx new file mode 100644 index 000000000000..66125afe6e43 --- /dev/null +++ b/docs/cookbook/vla/FLUX/FLUX-3-Action.mdx @@ -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.` (`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. diff --git a/docs/cookbook/vla/intro.mdx b/docs/cookbook/vla/intro.mdx index 9c9869442cfb..69d5d45ef658 100644 --- a/docs/cookbook/vla/intro.mdx +++ b/docs/cookbook/vla/intro.mdx @@ -19,3 +19,13 @@ This section keeps VLA policies separate from the diffusion model cookbook so ro href="/cookbook/vla/OpenPI/Pi0.5" /> + +## FLUX + + + + diff --git a/docs/docs.json b/docs/docs.json index f993070abd25..bc6168a7cabb 100644 --- a/docs/docs.json +++ b/docs/docs.json @@ -1617,6 +1617,13 @@ "pages": [ "cookbook/vla/OpenPI/Pi0.5" ] + }, + { + "group": "FLUX", + "tag": "NEW", + "pages": [ + "cookbook/vla/FLUX/FLUX-3-Action" + ] } ] }, diff --git a/docs/src/snippets/diffusion/model-catalog.jsx b/docs/src/snippets/diffusion/model-catalog.jsx index 5ed174b4b569..fea6a805b726 100644 --- a/docs/src/snippets/diffusion/model-catalog.jsx +++ b/docs/src/snippets/diffusion/model-catalog.jsx @@ -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: [ diff --git a/python/sglang/multimodal_gen/benchmarks/bench_flux3_action.py b/python/sglang/multimodal_gen/benchmarks/bench_flux3_action.py new file mode 100644 index 000000000000..93f87286f523 --- /dev/null +++ b/python/sglang/multimodal_gen/benchmarks/bench_flux3_action.py @@ -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() diff --git a/python/sglang/multimodal_gen/configs/models/dits/flux3.py b/python/sglang/multimodal_gen/configs/models/dits/flux3.py new file mode 100644 index 000000000000..6de6f80815ed --- /dev/null +++ b/python/sglang/multimodal_gen/configs/models/dits/flux3.py @@ -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_``) 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" diff --git a/python/sglang/multimodal_gen/configs/models/vaes/flux3_video.py b/python/sglang/multimodal_gen/configs/models/vaes/flux3_video.py new file mode 100644 index 000000000000..5ab2a0c8ba23 --- /dev/null +++ b/python/sglang/multimodal_gen/configs/models/vaes/flux3_video.py @@ -0,0 +1,35 @@ +# SPDX-License-Identifier: Apache-2.0 +"""FLUX 3 video VAE (Swin3D neighborhood-attention "ViTNorm") configuration.""" + +from dataclasses import dataclass, field + +from sglang.multimodal_gen.configs.models.vaes.base import VAEArchConfig, VAEConfig + + +@dataclass +class Flux3VideoVAEArchConfig(VAEArchConfig): + z_dim: int = 96 + embed_dim: int = 256 + patch_size: tuple[int, int, int] = (1, 4, 4) + window_size: tuple[int, int, int] = (5, 5, 5) + enc_depths: tuple[int, ...] = (1, 4, 8, 8) + dec_depths: tuple[int, ...] = (1, 4, 8, 8) + num_heads: tuple[int, ...] = (4, 8, 16, 32) + temporal: tuple[bool, ...] = (False, False, True, True) + enc_causal: bool = True + dec_causal: bool = False + qk_norm: bool = True + patch_norm: bool = False + temporal_compression_ratio: int = 4 + spatial_compression_ratio: int = 32 + # Encoding chunk length in frames; consecutive chunks overlap by one frame. + chunk_size_frames: int = 45 + # Looped decode: every decoder block attends over temporal windows of at + # most this many latent frames (plus the attention halo). ``None`` decodes + # the whole latent at once. + decoder_max_t: int | None = 8 + + +@dataclass +class Flux3VideoVAEConfig(VAEConfig): + arch_config: VAEArchConfig = field(default_factory=Flux3VideoVAEArchConfig) diff --git a/python/sglang/multimodal_gen/configs/pipeline_configs/base.py b/python/sglang/multimodal_gen/configs/pipeline_configs/base.py index 24e5fed0c5e8..d1072ea23689 100644 --- a/python/sglang/multimodal_gen/configs/pipeline_configs/base.py +++ b/python/sglang/multimodal_gen/configs/pipeline_configs/base.py @@ -299,6 +299,12 @@ def supports_openpi_endpoint(self) -> bool: return False + def action_metadata(self, server_args: Any) -> dict[str, Any] | None: + """Model-owned ``GET /v1/actions/metadata`` payload; None uses the generic one.""" + + del server_args + return None + # Wan2.2 TI2V parameters boundary_ratio: float | None = None diff --git a/python/sglang/multimodal_gen/configs/pipeline_configs/flux3_action.py b/python/sglang/multimodal_gen/configs/pipeline_configs/flux3_action.py new file mode 100644 index 000000000000..03f02277501e --- /dev/null +++ b/python/sglang/multimodal_gen/configs/pipeline_configs/flux3_action.py @@ -0,0 +1,446 @@ +# SPDX-License-Identifier: Apache-2.0 +"""FLUX 3 Action policy pipeline configuration. + +A FLUX 3 Action policy package (e.g. ``black-forest-labs/flux-3-action-droid``) +holds ``config.native.json`` (or ``config.json``), ``manifest.json`` and +``model.safetensors`` (the DiT with the embodiment's action streams). The +frozen text encoder and video VAE live in the shared base repository and are +referenced from the config as ``repo_id[:filename][@revision]``. + +Most fields below are filled from the package config when the server starts +(:meth:`Flux3ActionPipelineConfig.validate_server_args`), so the checkpoint +defines its camera layout, action space, sampler and guidance. +""" + +from __future__ import annotations + +import hashlib +import json +from dataclasses import dataclass, field +from pathlib import Path +from typing import Any + +from sglang.multimodal_gen.configs.models import DiTConfig, VAEConfig +from sglang.multimodal_gen.configs.models.dits.flux3 import ( + Flux3ArchConfig, + Flux3DiTConfig, +) +from sglang.multimodal_gen.configs.models.vaes.flux3_video import Flux3VideoVAEConfig +from sglang.multimodal_gen.configs.pipeline_configs.base import ( + ModelTaskType, + PipelineConfig, +) + +FLUX3_ACTION_POLICY_FILES = ( + "config.json", + "config.native.json", + "manifest.json", + "model.safetensors", +) +FLUX3_ACTION_VARIANTS = ("base", "fp8r", "gd", "gd-fp8r", "sd", "sd-fp8r") +_CONFIG_ONLY_FILES = ("config.json", "config.native.json", "manifest.json") + + +def flux3_action_variant_subfolder(variant: str | None) -> str: + """``--model-variant`` -> package subfolder (``base`` / ``None`` is the repository root).""" + if variant in (None, "", "base"): + return "" + if variant not in FLUX3_ACTION_VARIANTS: + raise ValueError( + f"unknown FLUX 3 Action variant {variant!r}; choose from {FLUX3_ACTION_VARIANTS}" + ) + return f"variants/{variant}" + + +def resolve_flux3_action_package( + model_path: str, + variant: str | None = None, + *, + config_only: bool = False, + revision: str | None = None, +) -> Path: + """Local directory of a policy package, downloading it from the Hub if needed.""" + subfolder = flux3_action_variant_subfolder(variant) + from sglang.multimodal_gen.runtime.utils.hf_diffusers_utils import ( + maybe_download_model, + ) + + prefix = f"{subfolder}/" if subfolder else "" + names = _CONFIG_ONLY_FILES if config_only else FLUX3_ACTION_POLICY_FILES + root = maybe_download_model( + str(Path(model_path).expanduser()), + allow_patterns=[prefix + name for name in names], + revision=revision, + ) + package = Path(root) / subfolder + if not package.is_dir(): + raise FileNotFoundError(f"FLUX 3 Action package {package} does not exist") + return package + + +def _policy_config_path(package: Path) -> Path: + """``config.native.json`` when present, else ``config.json`` (FP8r packages).""" + for name in ("config.native.json", "config.json"): + path = package / name + if path.is_file(): + config = json.loads(path.read_text()) + if isinstance(config, dict) and "action_modality" in config: + return path + raise FileNotFoundError(f"{package} holds no FLUX 3 Action policy config") + + +def read_flux3_action_config(package: Path) -> dict[str, Any]: + return json.loads(_policy_config_path(package).read_text()) + + +def _sha256(path: Path) -> str: + digest = hashlib.sha256() + with path.open("rb") as f: + for block in iter(lambda: f.read(1 << 24), b""): + digest.update(block) + return digest.hexdigest() + + +def verify_flux3_action_manifest(package: Path, *, include_weights: bool) -> None: + """Check the policy config (and weights) against the package manifest hashes.""" + manifest = json.loads((package / "manifest.json").read_text()) + if not isinstance(manifest, dict) or manifest.get("kind") != "policy_export": + raise ValueError(f"{package}/manifest.json is not a policy export manifest") + hashes = manifest.get("sha256") or {} + names = [_policy_config_path(package).name] + if include_weights: + names.append("model.safetensors") + for name in names: + if _sha256(package / name) != hashes.get(name): + raise ValueError(f"{package}/{name} does not match its manifest checksum") + + +# Reference ``JointSingleSeqParams`` fields -> Flux3ArchConfig fields. +_DIT_CONFIG_RENAMES = {"num_heads": "num_attention_heads"} +_DIT_CONFIG_PASSTHROUGH = ( + "vec_in_dim", + "context_in_dim", + "hidden_size", + "depth", + "depth_single_blocks", + "axes_dim", + "theta", + "mlp_ratio", +) +# Training-only knobs; the released trunk uses none of the non-default values. +_DIT_CONFIG_UNSUPPORTED = { + "depth_late_blocks": 0, + "qkv_bias": False, + "gate_type": None, +} +_DIT_CONFIG_IGNORED = ("attn_mode",) + + +def flux3_arch_config( + dit_config: dict[str, Any], streams: tuple[str, ...] +) -> Flux3ArchConfig: + """A reference ``dit_config`` restricted to ``streams`` -> :class:`Flux3ArchConfig`.""" + kwargs: dict[str, Any] = {} + for key, value in dit_config.items(): + if key in _DIT_CONFIG_UNSUPPORTED: + if value != _DIT_CONFIG_UNSUPPORTED[key]: + raise NotImplementedError( + f"dit_config.{key}={value!r} is not supported" + ) + elif key in _DIT_CONFIG_RENAMES: + kwargs[_DIT_CONFIG_RENAMES[key]] = value + elif key in _DIT_CONFIG_PASSTHROUGH: + kwargs[key] = tuple(value) if key == "axes_dim" else value + elif key not in _DIT_CONFIG_IGNORED + ("in_channels", "sequence"): + raise ValueError(f"unknown dit_config field {key!r}") + arch = Flux3ArchConfig(**kwargs) + in_channels = dit_config.get("in_channels", arch.in_channels) + missing = [s for s in streams if s not in in_channels] + if missing: + raise ValueError(f"dit_config.in_channels lacks content streams {missing}") + # Streams of the trunk the policy does not feed (e.g. image) are not built. + arch.in_channels = {s: in_channels[s] for s in streams} + arch.sequence = {f"x_{s}": s for s in streams} + arch.__post_init__() + return arch + + +def is_flux3_action_package(model_path: str) -> bool: + """Whether a local directory is a FLUX 3 Action policy export.""" + package = Path(model_path).expanduser() + manifest = package / "manifest.json" + if not manifest.is_file(): + return False + try: + manifest_data = json.loads(manifest.read_text()) + _policy_config_path(package) + except (OSError, ValueError): + return False + return ( + isinstance(manifest_data, dict) and manifest_data.get("kind") == "policy_export" + ) + + +def _validate_parallelism(server_args: Any, num_heads: int) -> None: + """Multi-GPU serving shards the DiT (TP, Ulysses SP) and/or splits the CFG passes.""" + if server_args.num_gpus == 1: + return + cfg_degree = ( + server_args.cfg_parallel_degree if server_args.enable_cfg_parallel else 1 + ) + supported = ( + cfg_degree in (1, 2) + and server_args.ring_degree == 1 + and server_args.num_gpus + == server_args.tp_size * server_args.sp_degree * cfg_degree + ) + if not supported: + raise NotImplementedError( + "FLUX 3 Action runs on num_gpus = tp_size x sp_degree x (2 with " + "--enable-cfg-parallel, else 1); ring attention is not supported" + ) + if num_heads % (server_args.tp_size * server_args.ulysses_degree): + raise ValueError( + f"tp_size x ulysses_degree must divide the {num_heads} attention heads" + ) + + +@dataclass +class Flux3ActionPipelineConfig(PipelineConfig): + """FLUX 3 Action: joint video + action flow matching, returns action chunks.""" + + task_type: ModelTaskType = ModelTaskType.VLA_ACTION + should_use_guidance: bool = True + enable_autocast: bool = False + dit_precision: str = "bf16" + vae_precision: str = "bf16" + + dit_config: DiTConfig = field(default_factory=Flux3DiTConfig) + # Only the encoder serves the policy (the predicted video is never decoded). + vae_config: VAEConfig = field( + default_factory=lambda: Flux3VideoVAEConfig(load_decoder=False) + ) + + # --- filled from the policy package (validate_server_args) --- + policy_family: str = "flux3_action" + policy_variant: str = "base" + policy_revision: str | None = None + action_modality: str = "action_prediction_droid" + action_dim: int = 8 + state_dim: int = 8 + output_action_dim: int = 8 + action_horizon: int = 32 + camera_layout: str = "droid" + image_keys: tuple[str, ...] = ("wrist", "left", "right") + # Alternative request names of the cameras (LeRobot / RoboArena DROID keys). + camera_aliases: dict[str, str] = field( + default_factory=lambda: { + "wrist_image_left": "wrist", + "exterior_image_1_left": "left", + "exterior_image_2_left": "right", + } + ) + canvas_hw: tuple[int, int] = (544, 736) + fps: float = 15.0 + action_scale: float = 2.0 + gripper_flip_dims: tuple[int, ...] = (-1,) + action_parameterization: str = "absolute" + absolute_action_dims: tuple[int, ...] = () + action_normalization: dict[str, list[float]] | None = None + state_normalization: dict[str, list[float]] | None = None + normalization_clip: float = 6.0 + inference_profile: str = "default" + sampler: str = "cosmos_unipc" + default_num_inference_steps: int = 4 + guidance_scale: float = 4.0 + # None: the action stream follows ``guidance_scale`` (also for request overrides). + guidance_scale_action: float | None = 1.0 + sampler_shift: float = 5.0 + inference_seed: int = 0 + quantization: str | None = None + video_vae_id: str | None = None + text_encoder_id: str | None = None + + # --- runtime --- + text_pad_multiple: int = 80 + text_max_length: int = 8192 + text_output_layers: tuple[int, ...] = (4, 8, 12, 16, 20, 24, 28, 32) + # Text contexts cached per caption, bounded by their total token count. + caption_cache_max_tokens: int = 1 << 16 + + def validate_server_args(self, server_args: Any) -> None: + super().validate_server_args(server_args) + variant = server_args.model_variant or "base" + package = resolve_flux3_action_package( + server_args.model_path, + variant, + config_only=True, + revision=server_args.revision, + ) + verify_flux3_action_manifest(package, include_weights=False) + self.load_policy_config(read_flux3_action_config(package)) + _validate_parallelism( + server_args, num_heads=self.dit_config.arch_config.num_attention_heads + ) + self.policy_variant = variant + self.policy_revision = server_args.revision + + def load_policy_config(self, config: dict[str, Any]) -> None: + """Adopt the embodiment, camera layout and inference recipe of a policy config.""" + profile = config.get("inference_profile", "default") + if profile != "default": + raise NotImplementedError( + f"FLUX 3 Action inference profile {profile!r} is not supported yet" + ) + if config.get("action_parameterization", "absolute") not in ( + "absolute", + "joint_delta", + ): + raise ValueError( + f"unknown action parameterization {config['action_parameterization']!r}" + ) + streams = tuple(config.get("content_streams") or ("video", "video_cond")) + if streams != ("video", "video_cond"): + raise NotImplementedError( + f"content streams {streams} are not supported yet" + ) + + self.action_modality = config["action_modality"] + self.action_dim = int(config["action_dim"]) + self.state_dim = self.action_dim + self.output_action_dim = self.action_dim + self.action_horizon = int(config["chunk_size"]) + self.camera_layout = config.get("camera_layout", "droid") + self.image_keys = tuple( + key.removeprefix("images.") for key in config.get("camera_keys", ()) + ) + self.canvas_hw = tuple(config.get("canvas_hw", (544, 736))) + self.fps = float(config.get("fps", 15.0)) + self.action_scale = float(config.get("action_scale", 2.0)) + self.gripper_flip_dims = tuple(config.get("gripper_flip_dims", ())) + self.action_parameterization = config.get("action_parameterization", "absolute") + self.absolute_action_dims = tuple(config.get("absolute_action_dims", ())) + self.action_normalization = config.get("action_normalization") + self.state_normalization = config.get("state_normalization") + self.normalization_clip = float(config.get("normalization_clip", 6.0)) + self.inference_profile = profile + self.sampler = config.get("sampler") or "cosmos_unipc" + if self.sampler != "cosmos_unipc": + raise NotImplementedError( + f"FLUX 3 Action sampler {self.sampler!r} is not supported" + ) + self.default_num_inference_steps = int(config.get("num_inference_steps") or 4) + self.guidance_scale = float( + 1.0 if config.get("guidance_scale") is None else config["guidance_scale"] + ) + action_guidance = config.get("guidance_scale_action") + self.guidance_scale_action = ( + None if action_guidance is None else float(action_guidance) + ) + self.sampler_shift = float(config.get("sampler_shift") or 1.0) + self.inference_seed = int(config.get("inference_seed", 0)) + self.quantization = config.get("quantization") + self.video_vae_id = config.get("video_vae_id") + self.text_encoder_id = config.get("text_encoder_id") + + conditioning_channels = config.get("conditioning_channels") or self.action_dim + self.dit_config = Flux3DiTConfig( + arch_config=flux3_arch_config( + config.get("dit_config") or {}, streams + ).with_streams( + { + self.action_modality: self.action_dim, + f"{self.action_modality}_cond": int(conditioning_channels), + } + ) + ) + + @property + def latent_hw(self) -> tuple[int, int]: + """Latent grid kept after the VAE: the content region of the canvas / 32 (ceil).""" + if self.camera_layout == "droid": + content = (540, 640) + else: + content = tuple(self.canvas_hw) + return tuple(-(-size // 32) for size in content) + + def resolve_guidance( + self, video: float | None, action: float | None + ) -> dict[str, float]: + """Per-stream guidance: request overrides, then the package recipe.""" + video_scale = self.guidance_scale if video is None else float(video) + if action is None: + action = self.guidance_scale_action + return { + "video": video_scale, + self.action_modality: video_scale if action is None else float(action), + } + + def supports_openpi_endpoint(self) -> bool: + return True + + def action_metadata(self, server_args: Any) -> dict[str, Any]: + return { + "object": "action.metadata", + "model": server_args.served_model_name, + "model_path": server_args.model_path, + "policy_family": self.policy_family, + "input": { + "image_keys": list(self.image_keys), + "camera_aliases": dict(self.camera_aliases), + "camera_layout": self.camera_layout, + "state_dim": self.state_dim, + }, + "output": { + "action_type": "continuous", + "action_horizon": self.action_horizon, + "action_dim": self.output_action_dim, + "padded_action_dim": self.action_dim, + "dtype": "float32", + }, + "defaults": { + "num_inference_steps": self.default_num_inference_steps, + "guidance_scale": self.guidance_scale, + "guidance_scale_action": self.guidance_scale_action, + "sampler": self.sampler, + "variant": self.policy_variant, + }, + "capabilities": { + "exact_prefix_cache": True, + "realtime_websocket": True, + "openpi_websocket": True, + "batch_inputs": False, + "multiple_candidates": False, + }, + } + + def estimate_request_cost(self, batch) -> float: + options = batch.extra.get("vla", {}).get("options", {}) + guidance = self.resolve_guidance( + options.get("guidance_scale"), options.get("guidance_scale_action") + ) + passes = 1 if all(g == 1.0 for g in guidance.values()) else 2 + steps = batch.num_inference_steps or self.default_num_inference_steps + return float(steps * passes) + + +# Policy exports (manifest.json + config.native.json) this pipeline serves. +FLUX3_ACTION_HF_PATHS = [ + "black-forest-labs/flux-3-action-droid", +] + + +def register(): + from sglang.multimodal_gen.configs.sample.flux3_action import ( + Flux3ActionSamplingParams, + ) + from sglang.multimodal_gen.registry import register_configs + + register_configs( + sampling_param_cls=Flux3ActionSamplingParams, + pipeline_config_cls=Flux3ActionPipelineConfig, + hf_model_paths=FLUX3_ACTION_HF_PATHS, + model_detectors=[ + lambda hf_id: "flux-3-action" in hf_id.lower(), + ], + ) diff --git a/python/sglang/multimodal_gen/configs/sample/flux3_action.py b/python/sglang/multimodal_gen/configs/sample/flux3_action.py new file mode 100644 index 000000000000..b810ac753a4c --- /dev/null +++ b/python/sglang/multimodal_gen/configs/sample/flux3_action.py @@ -0,0 +1,99 @@ +# SPDX-License-Identifier: Apache-2.0 +from dataclasses import dataclass, field +from typing import Any + +from sglang.multimodal_gen.configs.sample.action import ActionSamplingParams + + +@dataclass +class Flux3ActionSamplingParams(ActionSamplingParams): + """Sampling parameters for FLUX 3 Action policies. + + ``num_inference_steps``, ``guidance_scale``, ``guidance_scale_action`` and + ``seed`` default to the recipe of the served policy package (its + ``inference_seed`` for the noise) when left unset. + """ + + num_inference_steps: int | None = None + guidance_scale: float | None = None + guidance_scale_action: float | None = None + seed: int | list[int] | None = field( + default=None, metadata={"batch_sig_exclude": True} + ) + action_horizon: int | None = None + action_dim: int | None = None + output_format: str = "list" + return_timing: bool = True + # False bypasses the per-caption text context cache for this request. + enable_prefix_cache: bool = True + # Camera frames keyed by camera name (``wrist``, ``left``, ...), or a + # prebuilt ``composite``; HWC uint8 arrays, PIL images or tensors. + images: dict[str, Any] | None = field( + default=None, metadata={"batch_sig_exclude": True} + ) + state: 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.state is not None: + observation["state"] = self.state + if self.prompt is not None: + observation["prompt"] = self.prompt + extra["vla"] = { + "observation": observation, + "options": { + "output_format": self.output_format, + "return_timing": self.return_timing, + "guidance_scale": self.guidance_scale, + "guidance_scale_action": self.guidance_scale_action, + "enable_prefix_cache": self.enable_prefix_cache, + }, + } + return extra + + def _adjust(self, server_args): + super()._adjust(server_args) + config = server_args.pipeline_config + if self.num_inference_steps is None: + self.num_inference_steps = config.default_num_inference_steps + if self.seed is None: + self.seed = config.inference_seed + + def _validate(self): + steps, seed = self.num_inference_steps, self.seed + # None selects the package defaults; the base checks need ints. + self.num_inference_steps = 1 if steps is None else steps + self.seed = 0 if seed is None else seed + try: + super()._validate() + finally: + self.num_inference_steps, self.seed = steps, seed + if isinstance(seed, list) and len(seed) != 1: + raise ValueError("FLUX 3 Action takes one seed per request") + if self.num_outputs_per_prompt != 1: + raise ValueError("FLUX 3 Action returns one action chunk per request") + if self.action_horizon is not None and self.action_horizon <= 0: + raise ValueError("action_horizon must be positive") + if self.output_format not in ("list", "numpy"): + raise ValueError("output_format must be 'list' or 'numpy'") + for name, value in ( + ("guidance_scale", self.guidance_scale), + ("guidance_scale_action", self.guidance_scale_action), + ): + if value is None: + continue + if isinstance(value, bool) or not isinstance(value, (int, float)): + raise ValueError(f"{name} must be a number, got {value!r}") + if value < 0: + raise ValueError(f"{name} must be non-negative") + + def _set_output_file_name(self): + if self.output_file_name is None: + self.output_file_name = "flux3_action" + super()._set_output_file_name() diff --git a/python/sglang/multimodal_gen/envs.py b/python/sglang/multimodal_gen/envs.py index 67d06c10e50a..0a0ee5b0e92f 100644 --- a/python/sglang/multimodal_gen/envs.py +++ b/python/sglang/multimodal_gen/envs.py @@ -48,6 +48,7 @@ SGLANG_DIFFUSION_MINIMAX_H3_ADALN_GPU_PLANS: int = 64 SGLANG_DIFFUSION_MINIMAX_H3_ADALN_FP32: bool = False SGLANG_DIFFUSION_MINIMAX_H3_PDD_HEADS: str | None = None + SGLANG_DIFFUSION_FLUX3_NATTEN_BACKEND: str | None = None SGLANG_DIFFUSION_CFG_GATE_STEP: float = 1.0 # cache-dit env vars (primary transformer) # on by default; engages only on 2 ranks with peer-to-peer access and falls @@ -334,6 +335,11 @@ def _getter(): "SGLANG_DIFFUSION_MINIMAX_H3_PDD_HEADS": _lazy_str( "SGLANG_DIFFUSION_MINIMAX_H3_PDD_HEADS" ), + # NATTEN backend of the FLUX 3 video VAE (blackwell-fna, hopper-fna, + # cutlass-fna or flex-fna); probed per GPU when unset. + "SGLANG_DIFFUSION_FLUX3_NATTEN_BACKEND": _lazy_str( + "SGLANG_DIFFUSION_FLUX3_NATTEN_BACKEND" + ), "SGLANG_DIFFUSION_VAE_CHANNELS_LAST_3D": _lazy_str( "SGLANG_DIFFUSION_VAE_CHANNELS_LAST_3D", "auto" ), diff --git a/python/sglang/multimodal_gen/registry.py b/python/sglang/multimodal_gen/registry.py index cdb3af842228..d2b9b31a4920 100644 --- a/python/sglang/multimodal_gen/registry.py +++ b/python/sglang/multimodal_gen/registry.py @@ -28,6 +28,10 @@ from sglang.multimodal_gen.runtime.server_args import Backend from sglang.multimodal_gen.configs.pipeline_configs.base import PipelineConfig +from sglang.multimodal_gen.configs.pipeline_configs.flux3_action import ( + FLUX3_ACTION_HF_PATHS, + is_flux3_action_package, +) from sglang.multimodal_gen.configs.sensenova_u1 import ( SENSENOVA_U1_MODEL_IDS, is_sensenova_u1_adapter_only_model, @@ -177,6 +181,7 @@ class ConfigInfo: "lerobot/pi05": "Pi05Pipeline", "pi05": "Pi05Pipeline", "pi0.5": "Pi05Pipeline", + "flux-3-action": "Flux3ActionPipeline", "hunyuan3d": "Hunyuan3D2Pipeline", "flux.2-dev-nvfp4": "Flux2NvfpPipeline", "fal/ideogram-v4-fast": "Ideogram4FastPipeline", @@ -409,6 +414,10 @@ def _get_config_info( if registered_hf_id.lower() in SENSENOVA_U1_MODEL_IDS: return _CONFIG_REGISTRY.get(_MODEL_HF_PATH_TO_NAME[registered_hf_id]) + # Local FLUX 3 Action exports are identified by their manifest, not their name. + if is_flux3_action_package(model_path): + return _CONFIG_REGISTRY.get(_MODEL_HF_PATH_TO_NAME[FLUX3_ACTION_HF_PATHS[0]]) + # 1. Exact match if model_path in _MODEL_HF_PATH_TO_NAME: model_id = _MODEL_HF_PATH_TO_NAME[model_path] @@ -693,6 +702,8 @@ def get_non_diffusers_pipeline_name(model_path: str) -> Optional[str]: """Get the pipeline name for a known non-diffusers model.""" if is_sensenova_u1_model(model_path): return "SenseNovaU1Pipeline" + if is_flux3_action_package(model_path): + return "Flux3ActionPipeline" normalized_model_path = _normalize_hf_cache_path(model_path) model_short_name = get_model_short_name(normalized_model_path) diff --git a/python/sglang/multimodal_gen/runtime/entrypoints/action/protocol.py b/python/sglang/multimodal_gen/runtime/entrypoints/action/protocol.py index 50ca7603e2fa..0ec9945548a2 100644 --- a/python/sglang/multimodal_gen/runtime/entrypoints/action/protocol.py +++ b/python/sglang/multimodal_gen/runtime/entrypoints/action/protocol.py @@ -158,6 +158,9 @@ def action_metadata(server_args: ServerArgs) -> dict[str, Any]: pipeline_config = server_args.pipeline_config if isinstance(pipeline_config, Cosmos3Config): return cosmos3_action_metadata(server_args) + metadata = pipeline_config.action_metadata(server_args) + if metadata is not None: + return metadata policy_family = getattr( pipeline_config, @@ -387,6 +390,11 @@ def _build_action_model_sampling_params( "enable_prefix_cache": _runtime_bool(prefix_cache, True), "enable_cuda_graph": _runtime_bool(cuda_graph, True), } + # Optional per-request overrides; absent keys keep the model defaults. + for name in ("seed", "guidance_scale", "guidance_scale_action"): + value = parameters.get(name, observation.get(name)) + if value is not None: + sampling_kwargs[name] = value supported_fields = _sampling_params_field_names(sampling_params_cls) sp = sampling_params_cls( **{ diff --git a/python/sglang/multimodal_gen/runtime/models/dits/flux3.py b/python/sglang/multimodal_gen/runtime/models/dits/flux3.py new file mode 100644 index 000000000000..db08b85f0ec3 --- /dev/null +++ b/python/sglang/multimodal_gen/runtime/models/dits/flux3.py @@ -0,0 +1,940 @@ +# Copyright 2026 Black Forest Labs. All rights reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# SPDX-License-Identifier: Apache-2.0 +"""FLUX 3 ``JointSingleSeq`` DiT. + +Adapted from the FLUX Action reference implementation +(https://github.com/black-forest-labs/flux-action, +``flux_action/models/transformer.py``). The module tree matches the released +checkpoints, so tensors load without renaming. + +Every content stream (``video``, ``video_cond``, an action modality, ...) and +the text context first run through their own ``depth`` mode blocks +(self-attention within the stream). The text context and all active streams are +then concatenated and processed by ``depth_single_blocks`` joint blocks with +shared weights and per-stream modulation. Each stream carries its own +timesteps, so conditioning streams sit at ``t = 0`` while targets are noised. + +Besides the reference-compatible :meth:`Flux3Transformer.forward`, the model +exposes the pieces separately (:meth:`encode_context`, :meth:`encode_stream`, +:meth:`denoise`) so pipelines can cache everything that does not depend on the +noised streams: the text context per caption, and the conditioning streams per +request. +""" + +from __future__ import annotations + +import math +import os +from typing import Any + +import msgspec +import torch +import torch.nn as nn +import torch.nn.functional as F + +from sglang.kernels.ops.diffusion import ( + BitExactFusionGate, + can_use_fused_inplace_qknorm_rope, + can_use_fused_layernorm_modulate, + fused_inplace_qknorm_rope, + fused_layernorm_modulate_raw, + fused_packed_silu_mul_bitexact, + is_plain_layer_norm, + residual_gate_add, +) +from sglang.multimodal_gen.configs.models.dits.flux3 import ( + Flux3ArchConfig, + Flux3DiTConfig, +) +from sglang.multimodal_gen.runtime.distributed import ( + divide, + get_sp_world_size, + get_tp_rank, + get_tp_world_size, + tensor_model_parallel_all_reduce, +) +from sglang.multimodal_gen.runtime.distributed.sp_shard_utils import ( + SpShard, + build_shard_plan, + gather_seq, + shard_like, + tail_attn_meta, +) +from sglang.multimodal_gen.runtime.layers.attention import USPAttention +from sglang.multimodal_gen.runtime.layers.linear import ( + LinearBase, + MergedColumnParallelLinear, + RowParallelLinear, +) +from sglang.multimodal_gen.runtime.layers.quantization.configs.base_config import ( + QuantizationConfig, +) +from sglang.multimodal_gen.runtime.managers.memory_managers.layerwise_offload import ( + LayerwiseOffloadableModuleMixin, +) +from sglang.multimodal_gen.runtime.models.dits.base import BaseDiT +from sglang.multimodal_gen.runtime.platforms import AttentionBackendEnum +from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger + +logger = init_logger(__name__) + +Modulation = tuple[torch.Tensor, torch.Tensor, torch.Tensor] +TEXT_STREAM = "txt" + + +def rope_cos_sin( + ids: torch.Tensor, axes_dim: tuple[int, ...], theta: int +) -> torch.Tensor: + """Position ids ``(B, L, n_axes)`` -> fp32 ``(B, L, head_dim)`` rows ``[cos | sin]``. + + Pair ``i`` (channels ``2i, 2i + 1``) is rotated by angle ``i``; the angles of + the axes are concatenated in order (FLUX-style interleaved RoPE). + """ + angles = [] + for axis, dim in enumerate(axes_dim): + scale = torch.arange(0, dim, 2, dtype=torch.float64, device=ids.device) / dim + omega = 1.0 / (theta**scale) + angles.append(torch.einsum("...n,d->...nd", ids[..., axis], omega)) + angles = torch.cat(angles, dim=-1) + return torch.cat((torch.cos(angles), torch.sin(angles)), dim=-1).float() + + +def apply_rope( + q: torch.Tensor, k: torch.Tensor, cos_sin: torch.Tensor +) -> tuple[torch.Tensor, torch.Tensor]: + """Rotate adjacent channel pairs of ``q``/``k`` ``(B, L, H, D)`` in fp32.""" + cos, sin = cos_sin[:, :, None].chunk(2, dim=-1) + + def rotate(x: torch.Tensor) -> torch.Tensor: + pairs = x.float().reshape(*x.shape[:-1], -1, 2) + even = cos * pairs[..., 0] + (-sin) * pairs[..., 1] + odd = sin * pairs[..., 0] + cos * pairs[..., 1] + return torch.stack((even, odd), dim=-1).reshape_as(x).to(x.dtype) + + return rotate(q), rotate(k) + + +# Fused fast paths. The first two are bit-exact against the eager chain and +# verified per signature; QK-norm + RoPE is fused at bf16 rounding level. +_LN_MODULATE = BitExactFusionGate("FLUX 3 fused LN+modulate", per_signature=True) +_SWIGLU = BitExactFusionGate("FLUX 3 fused SwiGLU", per_signature=True) + + +def _eager_fast_path_allowed(x: torch.Tensor) -> bool: + return ( + x.is_cuda + and x.dtype is torch.bfloat16 + and not torch.compiler.is_compiling() + and not torch.cuda.is_current_stream_capturing() + ) + + +def _norm_modulate( + norm: nn.LayerNorm, x: torch.Tensor, shift: torch.Tensor, scale: torch.Tensor +) -> torch.Tensor: + """``(1 + scale) * LN(x) + shift`` with ``(B, 1, D)`` modulation.""" + scale_row, shift_row = scale[:, 0], shift[:, 0] + if ( + _LN_MODULATE.disabled + or not _eager_fast_path_allowed(x) + or not is_plain_layer_norm(norm, x.shape[-1]) + or not can_use_fused_layernorm_modulate(x, scale_row, shift_row) + ): + return (1 + scale) * norm(x) + shift + sig = (x.device, x.shape[0], x.shape[-1], norm.eps) + try: + out = fused_layernorm_modulate_raw(x, scale_row, shift_row, norm.eps) + except Exception as exc: + _LN_MODULATE.on_exception(exc, logger=logger) + return (1 + scale) * norm(x) + shift + if _LN_MODULATE.is_verified(sig): + return out + return _LN_MODULATE.accept_or_fallback( + out, + (1 + scale) * norm(x) + shift, + sig=sig, + logger=logger, + mismatch_msg="FLUX 3 fused LN+modulate is not bit-exact here; using eager", + ) + + +def _swiglu(packed: torch.Tensor) -> torch.Tensor: + """``silu(gate) * value`` of a packed ``[gate | value]`` projection.""" + gate, value = packed.chunk(2, dim=-1) + if _SWIGLU.disabled or not _eager_fast_path_allowed(packed): + return F.silu(gate) * value + sig = (packed.device, packed.shape[-1], packed.stride(-2)) + try: + out = fused_packed_silu_mul_bitexact(packed) + except Exception as exc: + _SWIGLU.on_exception(exc, logger=logger) + return F.silu(gate) * value + if _SWIGLU.is_verified(sig): + return out + return _SWIGLU.accept_or_fallback( + out, + F.silu(gate) * value, + sig=sig, + logger=logger, + mismatch_msg="FLUX 3 fused SwiGLU is not bit-exact here; using eager", + ) + + +def _fused_qknorm_rope_enabled(q: torch.Tensor, head_dim: int) -> bool: + return ( + _eager_fast_path_allowed(q) + and os.getenv("SGLANG_ENABLE_FUSED_QKNORM_ROPE", "1").lower() + not in ("0", "false", "off", "no") + and can_use_fused_inplace_qknorm_rope( + head_dim=head_dim, + rope_dim=head_dim, + is_neox=False, + dtype=q.dtype, + cache_dtype=torch.float32, + round_norm_before_rope=False, + ) + ) + + +def timestep_embedding(t: torch.Tensor, dim: int = 256) -> torch.Tensor: + """Sinusoidal embedding of ``t`` in ``[0, 1]`` (scaled by 1000), fp32.""" + half = dim // 2 + freqs = torch.exp( + -math.log(10000) + * torch.arange(half, device=t.device, dtype=torch.float32) + / half + ) + args = (1000.0 * t)[..., None].float() * freqs + return torch.cat((torch.cos(args), torch.sin(args)), dim=-1) + + +class Flux3MLPEmbedder(nn.Module): + def __init__(self, in_dim: int, hidden_dim: int): + super().__init__() + self.in_layer = nn.Linear(in_dim, hidden_dim, bias=False) + self.silu = nn.SiLU() + self.out_layer = nn.Linear(hidden_dim, hidden_dim, bias=False) + + def forward(self, x: torch.Tensor) -> torch.Tensor: + return self.out_layer(self.silu(self.in_layer(x))) + + +class Flux3RMSNorm(nn.Module): + def __init__(self, dim: int): + super().__init__() + self.scale = nn.Parameter(torch.ones(dim)) + + def forward(self, x: torch.Tensor) -> torch.Tensor: + dtype = x.dtype + x_float = x.float() + rrms = torch.rsqrt(torch.mean(x_float**2, dim=-1, keepdim=True) + 1e-6) + return (x_float * rrms).to(dtype) * self.scale + + +class Flux3QKNorm(nn.Module): + def __init__(self, dim: int): + super().__init__() + self.query_norm = Flux3RMSNorm(dim) + self.key_norm = Flux3RMSNorm(dim) + + def forward( + self, q: torch.Tensor, k: torch.Tensor, v: torch.Tensor + ) -> tuple[torch.Tensor, torch.Tensor]: + return self.query_norm(q).to(v), self.key_norm(k).to(v) + + +class Flux3Modulation(nn.Module): + """shift / scale / gate from the timestep vector (shared by all blocks of a phase).""" + + def __init__(self, hidden_size: int): + super().__init__() + self.lin = nn.Linear(hidden_size, 3 * hidden_size, bias=False) + + def forward(self, vec: torch.Tensor) -> Modulation: + out = self.lin(F.silu(vec)) + if out.ndim == 2: + out = out[:, None, :] + shift, scale, gate = out.chunk(3, dim=-1) + return shift, scale, gate + + +FP8_E4M3_MAX = 448.0 + + +def quantize_fp8_rowwise(x: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]: + """``(M, K)`` -> E4M3 values and one fp32 scale per row (amax / 448).""" + rows = x.float() + scale = (rows.abs().amax(dim=1) / FP8_E4M3_MAX).clamp(min=1e-12) + quantized = (rows / scale[:, None]).clamp(-FP8_E4M3_MAX, FP8_E4M3_MAX) + return quantized.to(torch.float8_e4m3fn), scale + + +class Flux3Fp8RowwiseLinear(nn.Module): + """FP8 "rowwise" linear of the FLUX 3 FP8r checkpoints. + + E4M3 weights with one fp32 scale per output row; activations are quantized + per token on the fly and multiplied with ``torch._scaled_mm`` (fp32 fast + accumulation, bf16 output), as in the reference FP8r inference path. + ``tuple_output`` mirrors the ``(output, bias)`` convention of SGLang's + parallel linears so the module can replace either kind. + """ + + # _scaled_mm needs M padded to a multiple of 16. + ROW_ALIGNMENT = 16 + + def __init__( + self, weight: torch.Tensor, weight_scale: torch.Tensor, tuple_output: bool + ): + super().__init__() + if weight.dtype != torch.float8_e4m3fn or weight_scale.shape != ( + weight.shape[0], + ): + raise ValueError( + "expected an E4M3 weight with one fp32 scale per output row" + ) + self.out_features, self.in_features = weight.shape + self.tuple_output = tuple_output + # Parameters (not buffers) so that layerwise offload streams them. + self.weight = nn.Parameter(weight.contiguous(), requires_grad=False) + self.weight_scale = nn.Parameter( + weight_scale.float().contiguous(), requires_grad=False + ) + + def forward(self, x: torch.Tensor): + leading = x.shape[:-1] + flat = x.reshape(-1, self.in_features).contiguous() + rows = flat.shape[0] + pad = -rows % self.ROW_ALIGNMENT + if pad: + flat = F.pad(flat, (0, 0, 0, pad)) + activation, activation_scale = quantize_fp8_rowwise(flat) + out = torch._scaled_mm( + activation, + self.weight.T, + activation_scale[:, None], + self.weight_scale[None, :], + out_dtype=torch.bfloat16, + use_fast_accum=True, + )[:rows] + out = out.reshape(*leading, self.out_features) + return (out, None) if self.tuple_output else out + + +class Flux3LastLayer(nn.Module): + def __init__(self, hidden_size: int, out_channels: int): + super().__init__() + self.norm_final = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6) + self.linear = nn.Linear(hidden_size, out_channels, bias=False) + self.adaLN_modulation = nn.Sequential( + nn.SiLU(), nn.Linear(hidden_size, 2 * hidden_size, bias=False) + ) + + def forward(self, x: torch.Tensor, vec: torch.Tensor) -> torch.Tensor: + if vec.ndim == 2: + vec = vec[:, None, :] + x = self.norm_final(x) + projection = self.adaLN_modulation[1] + if isinstance(projection, Flux3Fp8RowwiseLinear): + shift, scale = self.adaLN_modulation(vec).chunk(2, dim=-1) + return self.linear(x * (scale + 1) + shift) + # BF16: two half projections avoid materializing the (L, 2 * hidden) product. + activated = self.adaLN_modulation[0](vec) + shift_w, scale_w = projection.weight.chunk(2) + x.mul_(F.linear(activated, scale_w).add_(1)) + x.add_(F.linear(activated, shift_w)) + return self.linear(x) + + +# Streams shorter than this keep their mode blocks replicated under sequence +# parallelism (text, action and state tokens); arbitrary, not tuned. +SP_MIN_STREAM_TOKENS = 256 + + +class Flux3SequenceShard(msgspec.Struct, frozen=True): + """This rank's slice of a sequence-parallel sequence (Ulysses or K/V gather).""" + + shard: SpShard + # Tail-pad varlen meta for USPAttention (None when the split is even). + attn_mask_meta: dict[str, Any] | None + + +def _sequence_shard( + length: int, like: torch.Tensor, min_tokens: int = 0 +) -> Flux3SequenceShard | None: + """The SP split of a ``length``-token sequence, or None to keep it replicated.""" + if get_sp_world_size() == 1 or length < max(min_tokens, get_sp_world_size()): + return None + shard = build_shard_plan(length) + return Flux3SequenceShard( + shard=shard, + attn_mask_meta=tail_attn_meta(shard, like.shape[0], like.device), + ) + + +def _shard_mod(mod: Modulation, shard: SpShard) -> Modulation: + """Per-token modulation ``(B, L, D)`` follows the tokens; ``(B, 1, D)`` broadcasts.""" + return tuple(m if m.shape[1] == 1 else shard_like(m, shard, dim=1) for m in mod) + + +def _local_segments( + lengths: list[int], mods: list[Modulation], shard: SpShard +) -> tuple[list[int], list[Modulation]]: + """Segments of the joint sequence inside this rank's shard. + + The tail pad of the last rank becomes a segment with zero modulation, so + its rows pass through unchanged (attention masks them out). + """ + start = shard.sp_rank * shard.local_len + end = start + shard.local_real_len + local_lengths, local_mods = [], [] + pos = 0 + for length, mod in zip(lengths, mods): + lo, hi = max(pos, start), min(pos + length, end) + if lo < hi: + local_lengths.append(hi - lo) + local_mods.append( + tuple(m if m.shape[1] == 1 else m[:, lo - pos : hi - pos] for m in mod) + ) + pos += length + if shard.local_pad: + zero = mods[-1][0].new_zeros(mods[-1][0].shape[0], 1, mods[-1][0].shape[-1]) + local_lengths.append(shard.local_pad) + local_mods.append((zero, zero, zero)) + return local_lengths, local_mods + + +class Flux3Block(nn.Module): + """Parallel attention + SwiGLU MLP block with QK-RMSNorm and 4-axis RoPE. + + Used both as a per-stream mode block (one modulation for the whole input) + and as a joint block (one modulation per segment of the joint sequence). + """ + + def __init__( + self, + hidden_size: int, + num_heads: int, + mlp_ratio: float, + quant_config: QuantizationConfig | None = None, + supported_attention_backends: set[AttentionBackendEnum] | None = None, + prefix: str = "", + ): + super().__init__() + self.hidden_size = hidden_size + self.num_heads = num_heads + self.head_dim = hidden_size // num_heads + self.mlp_hidden_dim = int(hidden_size * mlp_ratio) + # Tensor parallel as in the FLUX 3 native DiT: each rank owns a slice of + # the heads and of the MLP channels; one all-reduce per block. + self.tp_size = get_tp_world_size() + self.local_heads = divide(num_heads, self.tp_size) + self.local_hidden = self.local_heads * self.head_dim + self.local_mlp_hidden = divide(self.mlp_hidden_dim, self.tp_size) + + # q, k, v and the MLP input share one GEMM; the checkpoint's separate + # tensors are concatenated at load time (see Flux3ArchConfig). Gate and + # value are separate partitions so each rank keeps matching halves. + self.qkv_mlp = MergedColumnParallelLinear( + hidden_size, + [hidden_size] * 3 + [self.mlp_hidden_dim] * 2, + bias=False, + gather_output=False, + quant_config=quant_config, + prefix=f"{prefix}.qkv_mlp", + ) + + def row_linear(name: str, in_features: int) -> RowParallelLinear: + return RowParallelLinear( + in_features, + hidden_size, + bias=False, + input_is_parallel=True, + reduce_results=False, + quant_config=quant_config, + prefix=f"{prefix}.{name}", + ) + + self.attn_out = row_linear("attn_out", hidden_size) + self.mlp_out = row_linear("mlp_out", self.mlp_hidden_dim) + self.norm = Flux3QKNorm(self.head_dim) + self.pre_norm = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6) + self.attn = USPAttention( + num_heads=self.local_heads, + head_size=self.head_dim, + causal=False, + supported_attention_backends=supported_attention_backends, + prefix=f"{prefix}.attn", + ) + + def _mix( + self, + modulated: torch.Tensor, + rope: torch.Tensor, + sp: Flux3SequenceShard | None, + ) -> torch.Tensor: + batch, length, _ = modulated.shape + heads = self.local_heads + q, k, v, mlp = self.qkv_mlp(modulated)[0].split( + ( + self.local_hidden, + self.local_hidden, + self.local_hidden, + 2 * self.local_mlp_hidden, + ), + dim=-1, + ) + q = q.view(batch, length, heads, self.head_dim) + k = k.view(batch, length, heads, self.head_dim) + v = v.view(batch, length, heads, self.head_dim) + if _fused_qknorm_rope_enabled(q, self.head_dim): + # In place on the fused projection: k follows q's heads in each row. + fused_inplace_qknorm_rope( + q=q.view(-1, heads, self.head_dim), + k=k.view(-1, heads, self.head_dim), + q_weight=self.norm.query_norm.scale, + k_weight=self.norm.key_norm.scale, + cos_sin_cache=rope.reshape(-1, self.head_dim), + positions=torch.arange(batch * length, device=q.device), + is_neox=False, + eps=1e-6, + round_norm_before_rope=False, + ) + else: + q, k = self.norm(q, k, v) + q, k = apply_rope(q, k, rope) + if sp is None: + # Replicated input: every rank already holds the whole stream. + attended = self.attn(q, k, v, skip_sequence_parallel_override=True) + else: + # q/k/v are strided views of the fused projection; the SP exchanges + # (all-to-all, K/V all-gather) need dense tensors. + q, k, v = q.contiguous(), k.contiguous(), v.contiguous() + attended = self.attn(q, k, v, attn_mask_meta=sp.attn_mask_meta) + attended = attended.reshape(batch, length, self.local_hidden) + out = self.attn_out(attended)[0] + self.mlp_out(_swiglu(mlp))[0] + if self.tp_size > 1: + # One reduce of the summed partials, not one per branch (the order + # the FLUX 3 reference TP uses). + out = tensor_model_parallel_all_reduce(out) + return out + + def forward( + self, + x: torch.Tensor, + rope: torch.Tensor, + mods: list[Modulation], + lengths: list[int] | None = None, + sp: Flux3SequenceShard | None = None, + ) -> torch.Tensor: + """Segment ``i`` of ``x`` (``lengths[i]`` tokens) is modulated by ``mods[i]``. + + A mode block has a single segment (``lengths=None``); a joint block has + one per stream. With ``sp``, ``x`` is this rank's sequence shard. + Blocks are always entered through ``__call__`` so that forward hooks + (layerwise offload) see every block. + """ + if lengths is None: + shift, scale, gate = mods[0] + modulated = _norm_modulate(self.pre_norm, x, shift=shift, scale=scale) + return residual_gate_add(x, self._mix(modulated, rope, sp), gate) + segments = torch.split(x, lengths, dim=1) + modulated = torch.cat( + [ + _norm_modulate(self.pre_norm, seg, shift=m[0], scale=m[1]) + for seg, m in zip(segments, mods) + ], + dim=1, + ) + output = torch.split(self._mix(modulated, rope, sp), lengths, dim=1) + return torch.cat( + [ + residual_gate_add(seg, out.contiguous(), m[2]) + for seg, out, m in zip(segments, output, mods) + ], + dim=1, + ) + + +class Flux3SegmentState(msgspec.Struct, frozen=True): + """A stream (or the text context) after its mode blocks, ready for the joint blocks.""" + + name: str + hidden: torch.Tensor # (B, L, hidden) + rope: torch.Tensor # (B, L, head_dim) fp32 [cos | sin] + vec: torch.Tensor # (B, 1 | L, hidden): timestep vector of the stream + joint_mod: Modulation # modulation of the stream in the joint blocks + + @property + def length(self) -> int: + return self.hidden.shape[1] + + +class Flux3Transformer(BaseDiT, LayerwiseOffloadableModuleMixin): + _fsdp_shard_conditions = [ + lambda name, module: isinstance(module, Flux3Block), + ] + _compile_conditions = _fsdp_shard_conditions + _fsdp_forward_methods = ("encode_context", "encode_stream", "denoise") + param_names_mapping = Flux3ArchConfig().param_names_mapping + reverse_param_names_mapping = {} + _supported_attention_backends = { + AttentionBackendEnum.FA, + AttentionBackendEnum.TORCH_SDPA, + AttentionBackendEnum.SAGE_ATTN, + AttentionBackendEnum.SAGE_ATTN_3, + AttentionBackendEnum.AITER, + } + + def __init__( + self, + config: Flux3DiTConfig, + hf_config: dict[str, Any] | None = None, + quant_config: QuantizationConfig | None = None, + **kwargs, + ) -> None: + super().__init__(config=config, hf_config=hf_config or {}, **kwargs) + arch: Flux3ArchConfig = config.arch_config + self.arch = arch + self.param_names_mapping = arch.param_names_mapping + self.hidden_size = arch.hidden_size + self.num_attention_heads = arch.num_attention_heads + self.num_channels_latents = arch.num_channels_latents + self.in_channels = dict(arch.in_channels) + self.sequence = dict(arch.sequence) + self.depth = arch.depth + self.axes_dim = tuple(arch.axes_dim) + self.theta = arch.theta + hidden = arch.hidden_size + + self.emb_in = nn.ModuleDict( + {m: nn.Linear(c, hidden, bias=False) for m, c in self.in_channels.items()} + ) + self.txt_in = nn.Linear(arch.context_in_dim, hidden, bias=False) + self.time_in = Flux3MLPEmbedder(256, hidden) + self.vector_in = ( + Flux3MLPEmbedder(arch.vec_in_dim, hidden) + if arch.vec_in_dim is not None + else None + ) + streams = sorted(self.in_channels) + self.early_stream_modulations = nn.ModuleDict( + {m: Flux3Modulation(hidden) for m in (*streams, TEXT_STREAM)} + ) + self.single_stream_modulations = nn.ModuleDict( + {m: Flux3Modulation(hidden) for m in (*streams, TEXT_STREAM)} + ) + + def block(prefix: str) -> Flux3Block: + return Flux3Block( + hidden, + arch.num_attention_heads, + arch.mlp_ratio, + quant_config=quant_config, + supported_attention_backends=self._supported_attention_backends, + prefix=prefix, + ) + + self.content_mode_blocks = nn.ModuleDict( + { + m: nn.ModuleList( + block(f"content_mode_blocks.{m}.{i}") for i in range(arch.depth) + ) + for m in streams + } + ) + self.txt_mode_blocks = nn.ModuleList( + block(f"txt_mode_blocks.{i}") for i in range(arch.depth) + ) + self.single_blocks = nn.ModuleList( + block(f"single_blocks.{i}") for i in range(arch.depth_single_blocks) + ) + self.final_layer = nn.ModuleDict( + {m: Flux3LastLayer(hidden, c) for m, c in self.in_channels.items()} + ) + self.layer_names = [ + "txt_mode_blocks", + *(f"content_mode_blocks.{m}" for m in streams), + "single_blocks", + ] + self.__post_init__() + + # ------------------------------------------------------------------ pieces + @property + def dtype(self) -> torch.dtype: + return _compute_dtype(self.txt_in) + + def rope(self, ids: torch.Tensor) -> torch.Tensor: + return rope_cos_sin(ids, self.axes_dim, self.theta) + + def _timestep_vector( + self, timesteps: torch.Tensor, vector: torch.Tensor | None = None + ) -> torch.Tensor: + """``timesteps`` ``(B,)`` or ``(B, L)`` in ``[0, 1]`` -> ``(B, 1 | L, hidden)``.""" + if timesteps.ndim == 1: + timesteps = timesteps[:, None] + weight_dtype = _compute_dtype(self.time_in.in_layer) + with torch.autocast(device_type=timesteps.device.type, enabled=False): + vec = self.time_in(timestep_embedding(timesteps).to(weight_dtype)) + vec = vec.to(self.dtype) + if self.vector_in is not None: + if vector is None: + vector = torch.zeros( + timesteps.shape[0], + self.vector_in.in_layer.in_features, + device=timesteps.device, + dtype=self.dtype, + ) + vec = vec + self.vector_in(vector.to(self.dtype))[:, None, :] + return vec + + def encode_context( + self, + ctx: torch.Tensor, + ctx_ids: torch.Tensor, + timesteps: torch.Tensor | None = None, + vector: torch.Tensor | None = None, + ) -> Flux3SegmentState: + """Text context ``(B, L, context_in_dim)`` through ``txt_in`` and the text mode blocks.""" + if timesteps is None: + timesteps = torch.zeros(ctx.shape[0], device=ctx.device) + vec = self._timestep_vector(timesteps=timesteps, vector=vector) + rope = self.rope(ctx_ids) + early = self.early_stream_modulations[TEXT_STREAM](vec) + hidden = self.txt_in(ctx.to(self.dtype)) + for block in self.txt_mode_blocks: + hidden = block(hidden, rope, [early]) + return Flux3SegmentState( + name=TEXT_STREAM, + hidden=hidden, + rope=rope, + vec=vec, + joint_mod=self.single_stream_modulations[TEXT_STREAM](vec), + ) + + def encode_stream( + self, + name: str, + x: torch.Tensor, + ids: torch.Tensor, + timesteps: torch.Tensor, + vector: torch.Tensor | None = None, + rope: torch.Tensor | None = None, + ) -> Flux3SegmentState: + """Stream tokens ``(B, L, in_channels[name])`` through ``emb_in`` and the stream's mode blocks.""" + vec = self._timestep_vector(timesteps=timesteps, vector=vector) + rope = self.rope(ids) if rope is None else rope + early = self.early_stream_modulations[name](vec) + hidden = self.emb_in[name](x.to(self.dtype)) + # Long streams run their mode blocks sequence-parallel; the state keeps + # the full stream so cached conditioning is layout independent. + sp = _sequence_shard(hidden.shape[1], hidden, min_tokens=SP_MIN_STREAM_TOKENS) + if sp is None: + for block in self.content_mode_blocks[name]: + hidden = block(hidden, rope, [early]) + else: + local = shard_like(hidden, sp.shard) + local_rope = shard_like(rope, sp.shard) + local_early = _shard_mod(early, sp.shard) + for block in self.content_mode_blocks[name]: + local = block(local, local_rope, [local_early], sp=sp) + hidden = gather_seq(local, sp.shard.orig_len) + return Flux3SegmentState( + name=name, + hidden=hidden, + rope=rope, + vec=vec, + joint_mod=self.single_stream_modulations[name](vec), + ) + + def _joint( + self, context: Flux3SegmentState, streams: list[Flux3SegmentState] + ) -> list[torch.Tensor]: + """Run the joint blocks over ``[context, *streams]``; returns the hidden state of each stream.""" + segments = [context, *streams] + lengths = [s.length for s in segments] + mods = [s.joint_mod for s in segments] + rope = torch.cat([s.rope for s in segments], dim=1) + x = torch.cat([s.hidden for s in segments], dim=1) + sp = _sequence_shard(x.shape[1], x) + if sp is None: + for block in self.single_blocks: + x = block(x, rope, mods, lengths) + else: + local = shard_like(x, sp.shard) + local_rope = shard_like(rope, sp.shard) + local_lengths, local_mods = _local_segments(lengths, mods, sp.shard) + for block in self.single_blocks: + local = block(local, local_rope, local_mods, local_lengths, sp=sp) + x = gather_seq(local, sp.shard.orig_len) + return list(torch.split(x, lengths, dim=1)[1:]) + + def denoise( + self, + context: Flux3SegmentState, + streams: list[Flux3SegmentState], + targets: list[str], + ) -> dict[str, torch.Tensor]: + """Joint blocks + output heads of the ``targets`` streams (by name).""" + hidden = self._joint(context, streams) + by_name = {s.name: (h, s) for h, s in zip(hidden, streams)} + return { + name: self.final_layer[name](by_name[name][0], by_name[name][1].vec) + for name in targets + } + + # ------------------------------------------------------------------ reference API + def forward( + self, + ctx: torch.Tensor, + ctx_ids: torch.Tensor, + vector: torch.Tensor | None = None, + timesteps_ctx: torch.Tensor | None = None, + **kwargs: torch.Tensor, + ) -> dict[str, torch.Tensor]: + """Predict every stream in ``kwargs`` (``x_``, ``x__ids``, ``x__timesteps``). + + Matches the reference ``JointSingleSeq.forward`` for dense batches: + streams absent from ``kwargs`` are skipped entirely (in the reference + they only feed their own discarded heads). + """ + context = self.encode_context( + ctx=ctx, ctx_ids=ctx_ids, timesteps=_uniform(timesteps_ctx), vector=vector + ) + streams, names = [], [] + for key, stream in self.sequence.items(): + if key not in kwargs: + continue + streams.append( + self.encode_stream( + name=stream, + x=kwargs[key], + ids=kwargs[f"{key}_ids"], + timesteps=_uniform(kwargs[f"{key}_timesteps"]), + vector=vector, + ) + ) + names.append(key) + outputs = self.denoise( + context=context, streams=streams, targets=[s.name for s in streams] + ) + return {key: outputs[s.name] for key, s in zip(names, streams)} + + +def _compute_dtype(linear: nn.Module) -> torch.dtype: + """Activation dtype of a linear: its weight dtype, bf16 for FP8 weights.""" + if isinstance(linear, Flux3Fp8RowwiseLinear): + return torch.bfloat16 + return linear.weight.dtype + + +def load_fp8r_checkpoint( + model: Flux3Transformer, state_dict: dict[str, torch.Tensor] +) -> None: + """Load a native FP8r checkpoint into ``model`` (built on the meta device). + + Every linear whose checkpoint weight is E4M3 (with ``.weight_scale``) + becomes a :class:`Flux3Fp8RowwiseLinear`; BF16 tensors (the embodiment's + action boundary layers and all norms) load as they are. + """ + from sglang.multimodal_gen.runtime.loader.utils import get_param_names_mapping + + mapping = get_param_names_mapping(model.param_names_mapping) + weights: dict[str, torch.Tensor] = {} + scales: dict[str, torch.Tensor] = {} + fused: dict[str, dict[int, tuple[torch.Tensor, torch.Tensor | None]]] = {} + fused_counts: dict[str, int] = {} + for name, tensor in state_dict.items(): + is_scale = name.endswith(".weight_scale") + weight_name = name.removesuffix("_scale") if is_scale else name + target, index, count = mapping(weight_name) + if index is None: + (scales if is_scale else weights)[target] = tensor + continue + weight, scale = fused.setdefault(target, {}).get(index, (None, None)) + fused[target][index] = (weight, tensor) if is_scale else (tensor, scale) + fused_counts[target] = count + for target, parts in fused.items(): + if sorted(parts) != list(range(fused_counts[target])) or any( + w is None for w, _ in parts.values() + ): + raise ValueError(f"{target}: checkpoint lacks fused parts {sorted(parts)}") + ordered = [parts[i] for i in sorted(parts)] + weights[target] = torch.cat([w for w, _ in ordered]) + part_scales = [sc for _, sc in ordered if sc is not None] + if part_scales: + if len(part_scales) != len(ordered): + raise ValueError(f"{target}: mixed FP8 and BF16 parts cannot be fused") + scales[target] = torch.cat(part_scales) + for name, scale in scales.items(): + module_path = name.removesuffix(".weight") + parent_path, _, child = module_path.rpartition(".") + parent = model.get_submodule(parent_path) + original = getattr(parent, child) + weight = _tp_partition(original, weights.pop(name), is_scale=False) + if weight.shape != original.weight.shape: + raise ValueError( + f"{name}: checkpoint shape {tuple(weight.shape)} does not match " + f"{tuple(original.weight.shape)}" + ) + setattr( + parent, + child, + Flux3Fp8RowwiseLinear( + weight, + _tp_partition(original, scale, is_scale=True), + tuple_output=isinstance(original, LinearBase), + ), + ) + for name, tensor in weights.items(): + owner = model.get_submodule(name.rpartition(".")[0]) + weights[name] = _tp_partition(owner, tensor, is_scale=False) + missing, unexpected = model.load_state_dict(weights, strict=False, assign=True) + fp8_params = { + n for n, _ in model.named_parameters() if n.removesuffix("_scale") in scales + } + missing = [n for n in missing if n not in fp8_params] + if missing or unexpected: + raise ValueError( + f"FP8r checkpoint mismatch: missing {missing}, unexpected {unexpected}" + ) + + +def _tp_partition( + module: nn.Module, tensor: torch.Tensor, *, is_scale: bool +) -> torch.Tensor: + """This rank's slice of a full checkpoint weight (or per-row scale) of ``module``.""" + tp_size = get_tp_world_size() + if tp_size == 1: + return tensor + rank = get_tp_rank() + if isinstance(module, MergedColumnParallelLinear): + parts = tensor.split(module.output_sizes) + return torch.cat([part.chunk(tp_size)[rank] for part in parts]) + if isinstance(module, RowParallelLinear) and not is_scale: + # Row scales stay whole: they scale output rows, which are not split. + return tensor.chunk(tp_size, dim=1)[rank].contiguous() + return tensor + + +def _uniform(timesteps: torch.Tensor | None) -> torch.Tensor | None: + """Collapse per-token timesteps ``(B, L)`` that are constant along ``L`` to ``(B,)``.""" + if timesteps is None or timesteps.ndim == 1: + return timesteps + if bool((timesteps == timesteps[:, :1]).all()): + return timesteps[:, 0] + return timesteps + + +EntryClass = Flux3Transformer diff --git a/python/sglang/multimodal_gen/runtime/models/encoders/flux3_text_encoder.py b/python/sglang/multimodal_gen/runtime/models/encoders/flux3_text_encoder.py new file mode 100644 index 000000000000..6147959496d6 --- /dev/null +++ b/python/sglang/multimodal_gen/runtime/models/encoders/flux3_text_encoder.py @@ -0,0 +1,124 @@ +# SPDX-License-Identifier: Apache-2.0 +"""FLUX 3 text context: stacked Qwen3-VL-4B hidden states. + +The DiT context is ``hidden_states[k]`` for ``k`` in ``output_layers`` (eight +layers of width 2560) concatenated along channels (20480). Each prompt is +wrapped in the chat template, right-padded to the next multiple of +``pad_multiple`` tokens and encoded with an attention mask; the padded +positions stay in the context (the DiT attends to them unmasked, as in the +reference implementation). +""" + +from __future__ import annotations + +import math +import os + +import torch +import torch.nn as nn + +from sglang.multimodal_gen.runtime.managers.memory_managers.layerwise_offload import ( + LayerwiseOffloadableModuleMixin, +) + + +def parse_weight_spec(spec: str) -> tuple[str, str | None, str | None]: + """``repo_id[:subfolder_or_file][@revision]`` -> ``(repo_id, subpath, revision)``.""" + body, _, revision = spec.partition("@") + repo_id, _, subpath = body.partition(":") + if not repo_id: + raise ValueError(f"invalid weight spec {spec!r}") + return repo_id, subpath or None, revision or None + + +class Flux3TextEncoder(nn.Module, LayerwiseOffloadableModuleMixin): + layerwise_offload_dit_group_enabled = False + layer_names = ["model.language_model.layers"] + + def __init__( + self, + spec: str, + *, + output_layers: tuple[int, ...] = (4, 8, 12, 16, 20, 24, 28, 32), + pad_multiple: int = 80, + max_length: int = 8192, + dtype: torch.dtype = torch.bfloat16, + ): + super().__init__() + from transformers import AutoProcessor, Qwen3VLForConditionalGeneration + + hub_kwargs = {} + if not os.path.exists(spec): + spec, subfolder, revision = parse_weight_spec(spec) + hub_kwargs = {"subfolder": subfolder or "", "revision": revision} + model = Qwen3VLForConditionalGeneration.from_pretrained( + spec, torch_dtype=dtype, **hub_kwargs + ) + # Text-only use of the multimodal backbone (keeps its M-RoPE position + # handling for padded prompts); the LM head and vision tower are dropped. + self.model = model.model + del self.model.visual + # hidden_states[k] is the input of layer k, so layers >= max(output_layers) + # and the final norm never reach the context. + language_model = self.model.language_model + del language_model.layers[max(output_layers) :] + language_model.norm = nn.Identity() + self.embed_dtype = dtype + self.processor = AutoProcessor.from_pretrained(spec, **hub_kwargs) + if self.processor.tokenizer.padding_side != "right": + raise ValueError("the FLUX 3 text encoder needs a right-padding tokenizer") + self.output_layers = tuple(output_layers) + self.pad_multiple = pad_multiple + self.max_length = max_length + + @property + def device(self) -> torch.device: + return next(self.model.parameters()).device + + def _bucket(self, prompt: str) -> int: + length = self.processor.tokenizer( + prompt, padding=False, truncation=True, max_length=self.max_length + )["input_ids"] + return min( + math.ceil(len(length) / self.pad_multiple) * self.pad_multiple, + self.max_length, + ) + + @torch.no_grad() + def encode(self, texts: list[str]) -> list[torch.Tensor]: + """Prompts -> contexts ``(1, L_i, len(output_layers) * hidden)``; one forward per length bucket.""" + formatted = [ + self.processor.apply_chat_template( + [{"role": "user", "content": text}], + tokenize=False, + add_generation_prompt=True, + ) + for text in texts + ] + buckets: dict[int, list[int]] = {} + for i, prompt in enumerate(formatted): + buckets.setdefault(self._bucket(prompt), []).append(i) + results: list[torch.Tensor | None] = [None] * len(texts) + tokenizer = self.processor.tokenizer + for length, members in sorted(buckets.items()): + toks = tokenizer( + [formatted[i] for i in members], + return_tensors="pt", + padding="max_length", + truncation=True, + max_length=length, + padding_side="right", + ) + out = self.model( + input_ids=toks["input_ids"].to(self.device), + attention_mask=toks["attention_mask"].to(self.device), + output_hidden_states=True, + use_cache=False, + ) + stacked = torch.cat( + [out.hidden_states[k] for k in self.output_layers], dim=-1 + ) + stacked = stacked.to(self.embed_dtype) + for j, i in enumerate(members): + results[i] = stacked[j : j + 1] + return results diff --git a/python/sglang/multimodal_gen/runtime/models/vaes/flux3_neighborhood_attention.py b/python/sglang/multimodal_gen/runtime/models/vaes/flux3_neighborhood_attention.py new file mode 100644 index 000000000000..98949bab48c6 --- /dev/null +++ b/python/sglang/multimodal_gen/runtime/models/vaes/flux3_neighborhood_attention.py @@ -0,0 +1,116 @@ +# SPDX-License-Identifier: Apache-2.0 +"""Neighborhood attention of the FLUX 3 video VAE without NATTEN. + +NATTEN is an optional dependency; when it is missing, the VAE runs the same +windows through a compiled FlexAttention block mask. The windows follow +NATTEN's ``na2d`` / ``na3d`` (stride 1, no dilation): on a non-causal axis the +window of query ``i`` is ``[start, start + k)`` with +``start = clamp(i - k // 2, 0, L - k)`` (shifted inward at the borders); on a +causal axis it is ``[max(0, i - k + 1), i]``. +""" + +from __future__ import annotations + +import importlib.util +from functools import lru_cache + +import torch + +from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger + +logger = init_logger(__name__) + +# One block mask per (grid, kernel, causal, device); a server sees a handful. +_BLOCK_MASK_CACHE_MAX = 16 +_block_masks: dict[tuple, object] = {} + + +@lru_cache(maxsize=1) +def natten_available() -> bool: + available = importlib.util.find_spec("natten") is not None + if not available: + logger.warning( + "NATTEN is not installed; the FLUX 3 video VAE uses the FlexAttention " + "fallback, compiled on the first request. Install a wheel matching " + "your torch/CUDA build from https://whl.natten.org/ to use NATTEN." + ) + return available + + +@lru_cache(maxsize=1) +def _compiled_flex_attention(): + from torch.nn.attention.flex_attention import flex_attention + + # Uncompiled, flex_attention materializes the dense score matrix. + return torch.compile(flex_attention, dynamic=False) + + +def _neighborhood_block_mask( + grid: tuple[int, ...], + kernel: tuple[int, ...], + causal: tuple[bool, ...], + device: torch.device, +): + from torch.nn.attention.flex_attention import create_block_mask + + key = (grid, kernel, causal, str(device)) + mask = _block_masks.get(key) + if mask is not None: + return mask + kernel = tuple(min(k, n) for k, n in zip(kernel, grid)) + strides = [1] * len(grid) + for axis in range(len(grid) - 2, -1, -1): + strides[axis] = strides[axis + 1] * grid[axis + 1] + + def mask_mod(batch_idx, head_idx, q_idx, kv_idx): + inside = None + for n, k, is_causal, stride in zip(grid, kernel, causal, strides): + q_pos = (q_idx // stride) % n + k_pos = (kv_idx // stride) % n + if is_causal: + axis_ok = (k_pos <= q_pos) & (k_pos > q_pos - k) + else: + start = torch.clamp(q_pos - k // 2, 0, n - k) + axis_ok = (k_pos >= start) & (k_pos < start + k) + inside = axis_ok if inside is None else inside & axis_ok + return inside + + seq_len = 1 + for n in grid: + seq_len *= n + # _compile=True keeps the mask construction sparse (eager is O(S^2) memory). + mask = create_block_mask( + mask_mod, + B=None, + H=None, + Q_LEN=seq_len, + KV_LEN=seq_len, + device=device, + _compile=True, + ) + if len(_block_masks) >= _BLOCK_MASK_CACHE_MAX: + _block_masks.pop(next(iter(_block_masks))) + _block_masks[key] = mask + return mask + + +def neighborhood_attention( + q: torch.Tensor, + k: torch.Tensor, + v: torch.Tensor, + *, + kernel_size: list[int], + is_causal: list[bool] | None = None, +) -> torch.Tensor: + """NATTEN-layout ``(B, *grid, heads, head_dim)`` in and out, like ``na2d`` / ``na3d``.""" + batch, *grid, heads, head_dim = q.shape + causal = tuple(is_causal or [False] * len(grid)) + mask = _neighborhood_block_mask(tuple(grid), tuple(kernel_size), causal, q.device) + + def to_flex(x: torch.Tensor) -> torch.Tensor: + return x.reshape(batch, -1, heads, head_dim).transpose(1, 2) + + out = _compiled_flex_attention()( + to_flex(q), to_flex(k), to_flex(v), block_mask=mask + ) + return out.transpose(1, 2).reshape(q.shape) diff --git a/python/sglang/multimodal_gen/runtime/models/vaes/flux3_video_vae.py b/python/sglang/multimodal_gen/runtime/models/vaes/flux3_video_vae.py new file mode 100644 index 000000000000..f3b3238255b6 --- /dev/null +++ b/python/sglang/multimodal_gen/runtime/models/vaes/flux3_video_vae.py @@ -0,0 +1,664 @@ +# Copyright 2026 Black Forest Labs. All rights reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# SPDX-License-Identifier: Apache-2.0 +"""FLUX 3 video VAE: Swin3D encoder/decoder with NATTEN neighborhood attention. + +Adapted from the FLUX Action reference implementation +(https://github.com/black-forest-labs/flux-action, +``flux_action/models/video_vae.py``). The module tree matches +``video_vae.safetensors`` so the checkpoint loads without renaming. + +Latents are ``(B, 96, 1 + (T - 1) // 4, H // 32, W // 32)`` and normalized by +the running statistics stored in the checkpoint. Neighborhood attention runs +on NATTEN (https://natten.org) when installed, else on a FlexAttention fallback +(``flux3_neighborhood_attention``). +""" + +from __future__ import annotations + +import math +from functools import partial + +import torch +import torch.nn as nn +import torch.nn.functional as F + +from sglang.multimodal_gen import envs +from sglang.multimodal_gen.configs.models.vaes.flux3_video import ( + Flux3VideoVAEArchConfig, + Flux3VideoVAEConfig, +) +from sglang.multimodal_gen.runtime.managers.memory_managers.layerwise_offload import ( + LayerwiseOffloadableModuleMixin, +) +from sglang.multimodal_gen.runtime.models.vaes.flux3_neighborhood_attention import ( + natten_available, + neighborhood_attention, +) +from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger + +logger = init_logger(__name__) + +_norm_layer = partial(nn.LayerNorm, eps=1e-5) + +# NATTEN backend, probed once per token rank (2D/3D): the fused CUTLASS +# kernels only exist for some architectures; flex-attention runs everywhere. +_PERSISTENT_KERNEL_BACKENDS = ("blackwell", "hopper") +_natten_backends: dict[int, str] = {} + + +def _natten_attention_kwargs( + q: torch.Tensor, + k: torch.Tensor, + v: torch.Tensor, + *, + kernel_size: list[int], + is_causal: list[bool] | None = None, +) -> dict: + backend = _natten_backends.get(q.ndim) + if backend is None: + backend = envs.SGLANG_DIFFUSION_FLUX3_NATTEN_BACKEND + if not backend: + from natten import backends as nb + + if nb.can_run_cutlass_blackwell_fna(q, k, v): + backend = "blackwell-fna" + elif nb.can_run_cutlass_hopper_fna(q, k, v): + backend = "hopper-fna" + elif nb.can_run_cutlass_fna(q, k, v): + backend = "cutlass-fna" + else: + backend = "flex-fna" + logger.info("FLUX 3 video VAE natten backend (%dD): %s", q.ndim - 3, backend) + _natten_backends[q.ndim] = backend + # A window covering every non-causal axis is dense attention; NATTEN then + # dispatches to its "*-fmha" kernels. + na_dim = q.ndim - 3 + causal = is_causal or [False] * na_dim + if all( + kk == s and not c + for kk, s, c in zip(kernel_size, q.shape[1 : 1 + na_dim], causal) + ): + backend = backend.replace("-fna", "-fmha") + kwargs = {"backend": backend} + if backend.split("-")[0] in _PERSISTENT_KERNEL_BACKENDS: + kwargs["run_persistent_kernel"] = True + return kwargs + + +class DistributedRunningStats(nn.Module): + """Latent mean / variance of the trained VAE; normalizes ``(B, C, ...)``.""" + + def __init__(self, num_channels: int): + super().__init__() + self.register_buffer("running_mean", torch.zeros(num_channels)) + self.register_buffer("running_var", torch.ones(num_channels)) + self.register_buffer("initialized", torch.tensor(False)) + + def _shape(self, x: torch.Tensor) -> tuple: + return (1, -1) + (1,) * (x.dim() - 2) + + def normalize(self, x: torch.Tensor) -> torch.Tensor: + s = self._shape(x) + return (x - self.running_mean.view(s)) / self.running_var.sqrt().view(s) + + def denormalize(self, x: torch.Tensor) -> torch.Tensor: + s = self._shape(x) + return x * self.running_var.sqrt().view(s) + self.running_mean.view(s) + + +class PatchMerging(nn.Module): + def __init__(self, dim: int, out_dim: int): + super().__init__() + self.norm = _norm_layer(4 * dim) + self.reduction = nn.Linear(4 * dim, out_dim, bias=False) + + def forward(self, x: torch.Tensor) -> torch.Tensor: + _, _, h, w, _ = x.shape + if h % 2 == 1 or w % 2 == 1: + x = F.pad(x, (0, 0, 0, w % 2, 0, h % 2)) + b, d, h, w, c = x.shape + x = x.reshape(b, d, h // 2, 2, w // 2, 2, c) + x = x.permute(0, 1, 2, 4, 3, 5, 6).flatten(4) + return self.reduction(self.norm(x)) + + +class TemporalMerging(nn.Module): + def __init__(self, dim: int, out_dim: int): + super().__init__() + self.norm = _norm_layer(2 * dim) + self.reduction = nn.Linear(2 * dim, out_dim, bias=False) + + def forward(self, x: torch.Tensor) -> torch.Tensor: + if x.shape[1] % 2 == 1: + x = torch.concat([x[:, :1], x], dim=1) + b, d, h, w, c = x.shape + x = x.reshape(b, d // 2, 2, h, w, c) + skip = x.mean(2) + x = x.permute(0, 1, 3, 4, 2, 5).reshape(b, d // 2, h, w, 2 * c) + return self.reduction(self.norm(x)) + skip + + +class PatchExpansion(nn.Module): + def __init__(self, dim: int, out_dim: int): + super().__init__() + self.dim = dim + self.out_dim = out_dim + self.norm = _norm_layer(dim) + self.expansion = nn.Linear(dim, 4 * out_dim, bias=False) + + def forward(self, x: torch.Tensor) -> torch.Tensor: + b, d, h, w, _ = x.shape + x = self.expansion(self.norm(x)) + x = x.view(b, d, h, w, 2, 2, self.out_dim) + x = x.permute(0, 1, 2, 4, 3, 5, 6).contiguous() + return x.view(b, d, h * 2, w * 2, self.out_dim) + + +class TemporalExpansion(nn.Module): + def __init__(self, dim: int, out_dim: int): + super().__init__() + self.dim = dim + self.out_dim = out_dim + self.norm = _norm_layer(dim) + self.expansion = nn.Linear(dim, 2 * out_dim, bias=False) + + def forward(self, x: torch.Tensor) -> torch.Tensor: + b, d, h, w, _ = x.shape + x = self.expansion(self.norm(x)) + torch.concat([x, x], -1) + x = x.view(b, d, h, w, 2, self.out_dim) + x = x.permute(0, 1, 4, 2, 3, 5).contiguous() + x = x.view(b, d * 2, h, w, self.out_dim) + return x[:, 1:] + + +def _temporal_core_halo( + core_start: int, core_end: int, length: int, kernel: int, causal: bool +) -> tuple[int, int]: + """Frames ``[halo_start, halo_end)`` a looped block must attend over to + reproduce frames ``[core_start, core_end)`` of the full-sequence output. + + NATTEN gives frame ``i`` the keys ``[start(i), start(i) + kernel)`` with + ``start(i) = clamp(i - kernel // 2, 0, length - kernel)`` (windows shift + inward at the edges), or ``[max(0, i - kernel + 1), i]`` when causal. + """ + if causal: + return max(0, core_start - kernel + 1), core_end + + def window_start(index: int) -> int: + return min(max(index - kernel // 2, 0), length - kernel) + + return window_start(core_start), window_start(core_end - 1) + kernel + + +class RotaryPositionEmbedding3D(nn.Module): + def __init__(self, head_dim: int, base: float = 256.0): + super().__init__() + assert head_dim % 8 == 0, "head dimension must be divisible by 8" + self.head_dim = head_dim + self.chunk_dim = head_dim // 4 + axis_inv_freq = 1.0 / ( + base ** (torch.arange(0, self.chunk_dim, 2).float() / self.chunk_dim) + ) + # (t, h, w, unused) axes; the checkpoint stores this buffer. + inv_freq = torch.stack( + [ + axis_inv_freq, + axis_inv_freq, + axis_inv_freq, + torch.zeros(self.chunk_dim // 2), + ] + ) + self.register_buffer("inv_freq", inv_freq) + + def forward( + self, + q: torch.Tensor, + k: torch.Tensor, + *, + temporal_offset: int | torch.Tensor = 0, + ) -> tuple[torch.Tensor, torch.Tensor]: + """Rotate ``(B, T, H, W, heads, head_dim)``; ``temporal_offset`` is the time of frame 0.""" + _, t, h, w, _, _ = q.shape + device, dtype = q.device, q.dtype + grids = torch.meshgrid( + torch.arange(t, device=device, dtype=torch.float32) + temporal_offset, + torch.arange(h, device=device, dtype=torch.float32), + torch.arange(w, device=device, dtype=torch.float32), + indexing="ij", + ) + pos = torch.stack(grids + (torch.zeros_like(grids[0]),), dim=-1) + freqs = torch.einsum("...a,af->...af", pos, self.inv_freq.float()) + freqs = freqs.reshape(1, t, h, w, 1, -1) + freqs = torch.cat([freqs, freqs], dim=-1) + cos = freqs.cos().to(dtype) + sin = freqs.sin().to(dtype) + q = q * cos + self._rotate_half(q) * sin + k = k * cos + self._rotate_half(k) * sin + return q, k + + @staticmethod + def _rotate_half(x: torch.Tensor) -> torch.Tensor: + x1, x2 = x.chunk(2, dim=-1) + return torch.cat([-x2, x1], dim=-1) + + +class Natten3D(nn.Module): + def __init__( + self, + dim: int, + window_size: list[int], + num_heads: int, + causal: bool = True, + qk_norm: bool = False, + ): + super().__init__() + self.window_size = list(window_size) + self.num_heads = num_heads + self.head_dim = dim // num_heads + self.causal = causal + self.qk_norm = qk_norm + self.qkv = nn.Linear(dim, dim * 3) + self.proj = nn.Linear(dim, dim) + self.rope = RotaryPositionEmbedding3D(self.head_dim) + if qk_norm: + self.q_norm = nn.RMSNorm(self.head_dim, eps=1e-5, elementwise_affine=False) + self.k_norm = nn.RMSNorm(self.head_dim, eps=1e-5, elementwise_affine=False) + + def _qkv(self, x: torch.Tensor): + b, t, h, w, _ = x.shape + q, k, v = ( + self.qkv(x).reshape(b, t, h, w, 3, self.num_heads, self.head_dim).unbind(4) + ) + if self.qk_norm: + q = self.q_norm(q) + k = self.k_norm(k) + return q, k, v + + def forward(self, x: torch.Tensor) -> torch.Tensor: + b, t, h, w, c = x.shape + q, k, v = self._qkv(x) + q, k = self.rope(q, k) + return self.proj(self._attend(q, k, v).reshape(b, t, h, w, c)) + + def forward_region( + self, + x: torch.Tensor, + *, + temporal_offset: torch.Tensor, + output_start: int, + output_end: int, + ) -> torch.Tensor: + """Attention over a temporal halo of a longer sequence; returns frames ``[output_start, output_end)``.""" + b, t, h, w, c = x.shape + q, k, v = self._qkv(x) + q, k = self.rope(q, k, temporal_offset=temporal_offset) + out = self._attend(q, k, v)[:, output_start:output_end] + return self.proj(out.reshape(b, output_end - output_start, h, w, c)) + + def _attend( + self, q: torch.Tensor, k: torch.Tensor, v: torch.Tensor + ) -> torch.Tensor: + if q.shape[1] == 1: # a single frame uses the 2-D kernel + q2, k2, v2 = q.squeeze(1), k.squeeze(1), v.squeeze(1) + kernel = self.window_size[1:] + if not natten_available(): + out = neighborhood_attention(q2, k2, v2, kernel_size=kernel) + return out.unsqueeze(1) + from natten.functional import na2d + + kwargs = _natten_attention_kwargs(q2, k2, v2, kernel_size=kernel) + return na2d( + q2, k2, v2, kernel_size=kernel, attention_kwargs=kwargs + ).unsqueeze(1) + causal = [self.causal, False, False] + if not natten_available(): + return neighborhood_attention( + q, k, v, kernel_size=self.window_size, is_causal=causal + ) + from natten.functional import na3d + + kwargs = _natten_attention_kwargs( + q, k, v, kernel_size=self.window_size, is_causal=causal + ) + return na3d( + q, + k, + v, + is_causal=causal, + kernel_size=self.window_size, + attention_kwargs=kwargs, + ) + + +class GLUMLP(nn.Module): + def __init__(self, dim: int, align_to: int = 64): + super().__init__() + hidden_dim = align_to * ((int(dim * 8 / 3) + align_to - 1) // align_to) + self.gate_up_proj = nn.Linear(dim, 2 * hidden_dim, bias=False) + self.down_proj = nn.Linear(hidden_dim, dim, bias=False) + + def forward(self, x: torch.Tensor) -> torch.Tensor: + gate, up = self.gate_up_proj(x).chunk(2, dim=-1) + return self.down_proj(F.silu(gate) * up) + + +class SwinTransformerBlock(nn.Module): + """Pre-norm neighborhood-attention block. + + With ``max_t`` set the block runs "looped": attention and MLP are applied + to temporal windows of at most ``max_t`` frames (each window attends over + its halo), writing the output into the input in place. The result equals + the full forward frame for frame at a fraction of the activation memory. + """ + + def __init__( + self, + dim: int, + num_heads: int, + window_size: list[int], + causal: bool = False, + qk_norm: bool = False, + max_t: int | None = None, + ): + super().__init__() + if max_t is not None and max_t < window_size[0]: + raise ValueError( + f"max_t={max_t} must be at least the temporal kernel {window_size[0]}" + ) + self.norm1 = _norm_layer(dim) + self.attn = Natten3D( + dim, window_size, num_heads, causal=causal, qk_norm=qk_norm + ) + self.norm2 = _norm_layer(dim) + self.mlp = GLUMLP(dim) + self.max_t = max_t + + def forward(self, x: torch.Tensor) -> torch.Tensor: + if self.max_t is None or x.shape[1] <= self.max_t: + x = x + self.attn(self.norm1(x)) + return x + self.mlp(self.norm2(x)) + return self._forward_looped(x) + + def _forward_looped(self, x: torch.Tensor) -> torch.Tensor: + if torch.is_grad_enabled(): + raise RuntimeError("looped decoding writes in place; run it under no_grad") + length = x.shape[1] + kernel = self.attn.window_size[0] + # A halo reaches at most kernel - 1 < max_t frames back into the + # previous window, so each window's output is written one iteration late. + pending: tuple[int, int, torch.Tensor] | None = None + for core_start in range(0, length, self.max_t): + core_end = min(core_start + self.max_t, length) + halo_start, halo_end = _temporal_core_halo( + core_start, core_end, length, kernel, self.attn.causal + ) + normalized_halo = self.norm1(x[:, halo_start:halo_end]) + if pending is not None: + x[:, pending[0] : pending[1]].copy_(pending[2]) + out = x[:, core_start:core_end] + self.attn.forward_region( + normalized_halo, + temporal_offset=torch.tensor(float(halo_start), device=x.device), + output_start=core_start - halo_start, + output_end=core_end - halo_start, + ) + pending = (core_start, core_end, out + self.mlp(self.norm2(out))) + assert pending is not None + x[:, pending[0] : pending[1]].copy_(pending[2]) + return x + + +class PatchEmbed3d(nn.Module): + def __init__( + self, patch_size, in_channels: int = 3, embed_dim: int = 96, norm: bool = True + ): + super().__init__() + self.patch = tuple(patch_size) + self.proj = nn.Conv3d( + in_channels, embed_dim, kernel_size=self.patch, stride=self.patch + ) + self.norm = _norm_layer(embed_dim) if norm else nn.Identity() + + def forward(self, x: torch.Tensor) -> torch.Tensor: + _, _, t, h, w = x.shape + pad = [(p - s % p) % p for s, p in zip((t, h, w), self.patch)] + x = F.pad(x, (0, pad[2], 0, pad[1], 0, pad[0])) + x = self.proj(x).permute(0, 2, 3, 4, 1) + return self.norm(x) + + +class EncoderSwin3D(nn.Module): + def __init__( + self, + z_ch: int, + patch_size, + embed_dim: int, + depths, + temporal, + num_heads, + window_size, + causal: bool = False, + qk_norm: bool = False, + patch_norm: bool = True, + ): + super().__init__() + self.proj = nn.Linear(embed_dim * 2 ** (len(depths) - 1), z_ch) + self.patch_embed = PatchEmbed3d( + patch_size=patch_size, embed_dim=embed_dim, norm=patch_norm + ) + layers: list[nn.Module] = [] + for i_stage in range(len(depths)): + dim = embed_dim * 2**i_stage + layers.append( + nn.Sequential( + *[ + SwinTransformerBlock( + dim, + num_heads[i_stage], + window_size, + causal=causal, + qk_norm=qk_norm, + ) + for _ in range(depths[i_stage]) + ] + ) + ) + downsampled = False + if i_stage < len(depths) - 1: + layers.append(PatchMerging(dim, 2 * dim)) + downsampled = True + if temporal[i_stage]: + heads = num_heads[i_stage + 1] if downsampled else num_heads[i_stage] + dim = 2 * dim if downsampled else dim + layers.append( + SwinTransformerBlock( + dim, heads, window_size, causal=causal, qk_norm=qk_norm + ) + ) + layers.append(TemporalMerging(dim, dim)) + self.features = nn.Sequential(*layers) + + def forward(self, x: torch.Tensor) -> torch.Tensor: + x = self.patch_embed(x) + x = self.features(x) + x = self.proj(x) + return x.permute(0, 4, 1, 2, 3).contiguous() + + +class DecoderSwin3D(nn.Module): + def __init__( + self, + z_ch: int, + patch_size, + embed_dim: int, + depths, + temporal, + num_heads, + window_size, + causal: bool = False, + qk_norm: bool = False, + max_t: int | None = None, + ): + super().__init__() + self.ps = list(patch_size) + self.proj_in = nn.Linear(z_ch, embed_dim * 2 ** (len(depths) - 1)) + self.proj_out = nn.Linear(embed_dim, math.prod(patch_size) * 3) + + def block(dim: int, heads: int) -> SwinTransformerBlock: + return SwinTransformerBlock( + dim, heads, window_size, causal=causal, qk_norm=qk_norm, max_t=max_t + ) + + layers: list[nn.Module] = [] + for i_stage in reversed(range(len(depths))): + dim = embed_dim * 2**i_stage + layers.append( + nn.Sequential( + *[block(dim, num_heads[i_stage]) for _ in range(depths[i_stage])] + ) + ) + if temporal[i_stage]: + layers.append(TemporalExpansion(dim, dim)) + layers.append(block(dim, num_heads[i_stage])) + if i_stage > 0: + layers.append(PatchExpansion(dim, dim // 2)) + self.features = nn.Sequential(*layers) + + def forward(self, x: torch.Tensor) -> torch.Tensor: + x = self.proj_in(x.permute(0, 2, 3, 4, 1).contiguous()) + x = self.proj_out(self.features(x)) + b, t, h, w, _ = x.shape + x = x.view(b, t, h, w, self.ps[0], self.ps[1], self.ps[2], 3) + x = x.permute(0, 1, 4, 2, 5, 3, 6, 7).contiguous() + x = x.view(b, t * self.ps[0], h * self.ps[1], w * self.ps[2], 3) + return x.permute(0, 4, 1, 2, 3).contiguous() + + +class ViTNorm(nn.Module): + def __init__(self, arch: Flux3VideoVAEArchConfig, *, load_decoder: bool = True): + super().__init__() + self.z_dim = arch.z_dim + common = dict( + patch_size=arch.patch_size, + window_size=arch.window_size, + embed_dim=arch.embed_dim, + num_heads=arch.num_heads, + temporal=arch.temporal, + qk_norm=arch.qk_norm, + ) + self.encoder = EncoderSwin3D( + z_ch=2 * arch.z_dim, + depths=arch.enc_depths, + causal=arch.enc_causal, + patch_norm=arch.patch_norm, + **common, + ) + self.decoder = ( + DecoderSwin3D( + z_ch=arch.z_dim, + depths=arch.dec_depths, + causal=arch.dec_causal, + max_t=arch.decoder_max_t, + **common, + ) + if load_decoder + else None + ) + self.z_normalizer = DistributedRunningStats(arch.z_dim) + + def encode(self, x: torch.Tensor) -> torch.Tensor: + mu, _ = self.encoder(x).chunk(2, dim=-4) + return self.z_normalizer.normalize(mu) + + def decode(self, z: torch.Tensor) -> torch.Tensor: + if self.decoder is None: + raise RuntimeError("this video VAE was built without its decoder") + return self.decoder(self.z_normalizer.denormalize(z)) + + +class Flux3VideoVAE(nn.Module, LayerwiseOffloadableModuleMixin): + """Frozen FLUX 3 video VAE: ``encode`` / ``encode_frame`` / ``decode``.""" + + layerwise_offload_dit_group_enabled = False + layer_names = ["model.encoder.features", "model.decoder.features"] + + def __init__(self, config: Flux3VideoVAEConfig, **kwargs): + super().__init__() + arch: Flux3VideoVAEArchConfig = config.arch_config + if tuple(arch.temporal) != (False, False, True, True) or tuple( + arch.patch_size + ) != (1, 4, 4): + raise ValueError( + "compression ratios assume temporal=(F, F, T, T) and patch (1, 4, 4)" + ) + self.config = config + self.arch = arch + self.model = ViTNorm(arch, load_decoder=config.load_decoder) + + @property + def temporal_compression_ratio(self) -> int: + return self.arch.temporal_compression_ratio + + @property + def spatial_compression_ratio(self) -> int: + return self.arch.spatial_compression_ratio + + @torch.no_grad() + def encode(self, video: torch.Tensor) -> torch.Tensor: + """``(B, 3, T, H, W)`` in ``[-1, 1]`` with ``T = 1 (mod 4)`` -> normalized latents. + + Clips longer than one chunk are encoded in ``chunk_size_frames`` chunks + overlapping by one frame (the repeated first latent of each later chunk + is dropped); shorter clips are padded by repeating the last frame. + """ + num_frames = video.shape[2] + if (num_frames - 1) % self.temporal_compression_ratio: + raise ValueError( + f"video VAE encode expects T = 1 (mod 4) frames, got {num_frames}" + ) + chunk = self.arch.chunk_size_frames + stride = chunk - 1 + padded = chunk + max(0, -(-(num_frames - chunk) // stride)) * stride + if padded > num_frames: + tail = video[:, :, -1:].expand(-1, -1, padded - num_frames, -1, -1) + video = torch.cat([video, tail], dim=2) + pieces = [] + for start in range(0, padded - chunk + 1, stride): + z = self.model.encode(video[:, :, start : start + chunk]) + pieces.append(z if start == 0 else z[:, :, 1:]) + latent = torch.cat(pieces, dim=2) + return latent[:, :, : 1 + (num_frames - 1) // self.temporal_compression_ratio] + + @torch.no_grad() + def encode_frame(self, frame: torch.Tensor) -> torch.Tensor: + """``(B, 3, H, W)`` -> ``(B, 96, 1, H // 32, W // 32)``: the frame alone, no chunk padding.""" + return self.model.encode(frame[:, :, None]) + + @torch.no_grad() + def decode(self, latents: torch.Tensor) -> torch.Tensor: + """``(B, 96, T_lat, h, w)`` -> pixels ``(B, 3, 4 * T_lat - 3, 32 h, 32 w)`` in ``[-1, 1]``.""" + return self.model.decode(latents).clamp(-1, 1) + + def load_checkpoint(self, state_dict: dict[str, torch.Tensor]) -> None: + """Strict load of ``video_vae.safetensors`` (decoder tensors skipped when not built).""" + if self.model.decoder is None: + state_dict = { + k: v + for k, v in state_dict.items() + if not k.startswith("model.decoder.") + } + self.load_state_dict(state_dict, strict=True, assign=True) + + +EntryClass = Flux3VideoVAE diff --git a/python/sglang/multimodal_gen/runtime/pipelines/flux3_action.py b/python/sglang/multimodal_gen/runtime/pipelines/flux3_action.py new file mode 100644 index 000000000000..c7f374ccea63 --- /dev/null +++ b/python/sglang/multimodal_gen/runtime/pipelines/flux3_action.py @@ -0,0 +1,202 @@ +# SPDX-License-Identifier: Apache-2.0 +"""FLUX 3 Action robot policy pipeline (``black-forest-labs/flux-3-action-*``).""" + +from __future__ import annotations + +import os + +import torch +from safetensors.torch import load_file + +from sglang.multimodal_gen.configs.pipeline_configs.flux3_action import ( + Flux3ActionPipelineConfig, + resolve_flux3_action_package, + verify_flux3_action_manifest, +) +from sglang.multimodal_gen.configs.sample.flux3_action import Flux3ActionSamplingParams +from sglang.multimodal_gen.runtime.disaggregation.roles import RoleType +from sglang.multimodal_gen.runtime.distributed import get_local_torch_device +from sglang.multimodal_gen.runtime.loader.fsdp_load import maybe_load_fsdp_model +from sglang.multimodal_gen.runtime.loader.utils import ( + get_memory_usage_of_component, + set_default_torch_dtype, +) +from sglang.multimodal_gen.runtime.models.dits.flux3 import ( + Flux3Transformer, + load_fp8r_checkpoint, +) +from sglang.multimodal_gen.runtime.models.encoders.flux3_text_encoder import ( + Flux3TextEncoder, + parse_weight_spec, +) +from sglang.multimodal_gen.runtime.models.schedulers.scheduling_flow_unipc_multistep import ( + FlowUniPCMultistepScheduler, +) +from sglang.multimodal_gen.runtime.models.vaes.flux3_video_vae import Flux3VideoVAE +from sglang.multimodal_gen.runtime.pipelines_core.composed_pipeline_base import ( + ComposedPipelineBase, +) +from sglang.multimodal_gen.runtime.pipelines_core.stages.model_specific_stages.flux3_action import ( + Flux3ActionDenoisingStage, + Flux3ActionObservationEncodingStage, + Flux3ActionPreprocessStage, + Flux3ActionTextEncodingStage, +) +from sglang.multimodal_gen.runtime.pipelines_core.stages.vla import ( + VLAActionPostprocessStage, +) +from sglang.multimodal_gen.runtime.server_args import ServerArgs +from sglang.multimodal_gen.runtime.utils.hf_diffusers_utils import hf_hub_download +from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger +from sglang.multimodal_gen.runtime.utils.precision import resolve_component_precision + +logger = init_logger(__name__) + +VIDEO_VAE_FILENAME = "video_vae.safetensors" + + +def _resolve_file(spec: str, default_filename: str) -> str: + """A local file as is, otherwise ``repo_id[:filename][@revision]`` from the Hub.""" + if os.path.exists(spec): + return spec + repo_id, filename, revision = parse_weight_spec(spec) + return hf_hub_download(repo_id, filename or default_filename, revision=revision) + + +class Flux3ActionPipeline(ComposedPipelineBase): + pipeline_name = "Flux3ActionPipeline" + pipeline_config_cls = Flux3ActionPipelineConfig + sampling_params_cls = Flux3ActionSamplingParams + _required_config_modules: list[str] = [] + + def validate_disagg_role(self, role: RoleType) -> None: + if role != RoleType.MONOLITHIC: + raise ValueError("Flux3ActionPipeline supports same-process execution only") + + 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 + config: Flux3ActionPipelineConfig = server_args.pipeline_config + if config.quantization not in (None, "fp8r"): + raise NotImplementedError( + f"FLUX 3 Action {config.quantization} packages are not supported" + ) + package = resolve_flux3_action_package( + self.model_path, config.policy_variant, revision=config.policy_revision + ) + verify_flux3_action_manifest(package, include_weights=True) + modules = { + "transformer": self._load_transformer(server_args, config, package), + "vae": self._load_vae(server_args, config), + "text_encoder": self._load_text_encoder(server_args, config), + } + for name, module in modules.items(): + module.requires_grad_(False).eval() + self.memory_usages[name] = get_memory_usage_of_component(module) + # Cosmos UniPC of the reference; the policy's shift is applied per request + # in set_timesteps, so the scheduler itself is unshifted. + modules["scheduler"] = FlowUniPCMultistepScheduler( + solver_order=2, + solver_type="bh2", + predict_x0=True, + lower_order_final=True, + final_sigmas_type="zero", + shift=1.0, + ) + return modules + + @staticmethod + def _load_transformer( + server_args: ServerArgs, config: Flux3ActionPipelineConfig, package + ) -> Flux3Transformer: + logger.info("Loading FLUX 3 Action DiT from %s", package) + weights = str(package / "model.safetensors") + device = get_local_torch_device() + if config.quantization == "fp8r": + # Native FP8r payloads load as they are (no requantization). + with torch.device("meta"), set_default_torch_dtype(torch.bfloat16): + transformer = Flux3Transformer(config=config.dit_config, hf_config={}) + # Offloaded (e.g. layerwise) components start on the host. + on_cpu = server_args.should_start_component_on_cpu("transformer") + load_fp8r_checkpoint( + transformer, load_file(weights, device="cpu" if on_cpu else str(device)) + ) + return transformer + return maybe_load_fsdp_model( + model_cls=Flux3Transformer, + init_params={"config": config.dit_config, "hf_config": {}}, + weight_dir_list=[weights], + device=device, + hsdp_replicate_dim=server_args.hsdp_replicate_dim, + hsdp_shard_dim=server_args.hsdp_shard_dim, + param_dtype=resolve_component_precision(server_args, "transformer"), + reduce_dtype=torch.float32, + component_starts_on_cpu=server_args.should_start_component_on_cpu( + "transformer" + ), + fsdp_inference=server_args.should_use_fsdp_for_component("transformer"), + pin_cpu_memory=server_args.pin_cpu_memory, + strict=True, + ) + + @staticmethod + def _load_vae( + server_args: ServerArgs, config: Flux3ActionPipelineConfig + ) -> Flux3VideoVAE: + spec = server_args.component_paths.get("vae") or config.video_vae_id + logger.info("Loading FLUX 3 video VAE from %s", spec) + device = get_local_torch_device() + with torch.device("meta"): + vae = Flux3VideoVAE(config.vae_config) + vae.load_checkpoint( + load_file(_resolve_file(spec, VIDEO_VAE_FILENAME), device=str(device)) + ) + return vae.to(resolve_component_precision(server_args, "vae")) + + @staticmethod + def _load_text_encoder( + server_args: ServerArgs, config: Flux3ActionPipelineConfig + ) -> Flux3TextEncoder: + spec = server_args.component_paths.get("text_encoder") or config.text_encoder_id + logger.info("Loading FLUX 3 text encoder from %s", spec) + text_encoder = Flux3TextEncoder( + spec, + output_layers=config.text_output_layers, + pad_multiple=config.text_pad_multiple, + max_length=config.text_max_length, + ) + on_cpu = server_args.should_start_component_on_cpu("text_encoder") + return text_encoder.to("cpu" if on_cpu else get_local_torch_device()) + + def create_pipeline_stages(self, server_args: ServerArgs): + config: Flux3ActionPipelineConfig = server_args.pipeline_config + transformer = self.get_module("transformer") + self.add_stage(Flux3ActionPreprocessStage(config), "flux3_action_preprocess") + self.add_stage( + Flux3ActionTextEncodingStage( + config, + transformer=transformer, + text_encoder=self.get_module("text_encoder"), + ), + "flux3_action_text_encoding", + ) + self.add_stage( + Flux3ActionObservationEncodingStage( + config, transformer=transformer, vae=self.get_module("vae") + ), + "flux3_action_observation_encoding", + ) + self.add_stage( + Flux3ActionDenoisingStage( + config, transformer, scheduler=self.get_module("scheduler") + ), + "flux3_action_denoise", + ) + self.add_stage(VLAActionPostprocessStage(), "flux3_action_postprocess") + + +EntryClass = Flux3ActionPipeline diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/flux3_action.py b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/flux3_action.py new file mode 100644 index 000000000000..2d6a7ee370a9 --- /dev/null +++ b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/flux3_action.py @@ -0,0 +1,751 @@ +# SPDX-License-Identifier: Apache-2.0 +"""FLUX 3 Action stages: observation -> conditioning -> joint video/action denoising. + +The policy follows the FLUX Action reference implementation +(https://github.com/black-forest-labs/flux-action): the camera frames are +composed onto one canvas and VAE-encoded as the ``video_cond`` stream, the +normalized state is the ``_cond`` token, and future video latents are +denoised jointly with the action chunk. Only the actions are returned. + +Everything that does not depend on the noised streams is computed once: the +text context per caption (cached across requests), the conditioning streams +per request, and the target streams' mode blocks per step (shared by the +conditional and unconditional CFG passes, since mode blocks never see text). +""" + +from __future__ import annotations + +import copy +import time +from collections import OrderedDict +from collections.abc import Callable +from typing import Any + +import msgspec +import numpy as np +import torch +import torch.nn.functional as F +from einops import rearrange + +from sglang.multimodal_gen.configs.pipeline_configs.flux3_action import ( + Flux3ActionPipelineConfig, +) +from sglang.multimodal_gen.runtime.distributed import get_local_torch_device +from sglang.multimodal_gen.runtime.distributed.cfg_parallel_utils import ( + run_cfg_parallel, +) +from sglang.multimodal_gen.runtime.distributed.cfg_policy import CFGBranch, CFGPolicy +from sglang.multimodal_gen.runtime.managers.forward_context import set_forward_context +from sglang.multimodal_gen.runtime.managers.memory_managers.component_manager import ( + ComponentUse, +) +from sglang.multimodal_gen.runtime.models.dits.flux3 import ( + Flux3SegmentState, + Flux3Transformer, +) +from sglang.multimodal_gen.runtime.models.encoders.flux3_text_encoder import ( + Flux3TextEncoder, +) +from sglang.multimodal_gen.runtime.models.schedulers.scheduling_flow_unipc_multistep import ( + FlowUniPCMultistepScheduler, +) +from sglang.multimodal_gen.runtime.models.vaes.flux3_video_vae import Flux3VideoVAE +from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import Req +from sglang.multimodal_gen.runtime.pipelines_core.stages.base import PipelineStage +from sglang.multimodal_gen.runtime.pipelines_core.stages.vla import ( + synchronize_vla_action_tensor, + vla_options, + vla_state, + vla_timings, +) +from sglang.multimodal_gen.runtime.server_args import ServerArgs +from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger + +logger = init_logger(__name__) + +DROID_CAMERA_HW = (360, 640) +DROID_COMPOSITE_HW = (540, 640) +LATENT_CHANNELS = 96 +TEMPORAL_DOWNSAMPLE = 4 +GRAY_LEVEL = 128 + + +class Flux3ActionObservation(msgspec.Struct, frozen=True): + canvas: torch.Tensor # (3, Hc, Wc) in [-1, 1] + state: torch.Tensor # (D,) fp32, dataset units (gripper as stored) + prompt: str + + +# ---------------------------------------------------------------- position ids +def _times_to_ids(seconds: torch.Tensor) -> torch.Tensor: + """Seconds -> ids on the shared 10 ms clock.""" + return (seconds * 1000 // 10).to(torch.int64) + + +def _cartesian_ids( + t: torch.Tensor, h: int, w: int, l_coord: torch.Tensor +) -> torch.Tensor: + return torch.cartesian_prod(t, torch.arange(h), torch.arange(w), l_coord) + + +def pack_video(latent: torch.Tensor, first_frame: int, fps: float): + """Latent ``(1, C, T, h, w)`` -> tokens ``(1, T*h*w, C)`` and ids; frame ``i`` at ``i * 4 / fps`` s.""" + _, _, t, h, w = latent.shape + seconds = ( + torch.arange(first_frame, first_frame + t).float() * TEMPORAL_DOWNSAMPLE / fps + ) + ids = _cartesian_ids(t=_times_to_ids(seconds), h=h, w=w, l_coord=torch.arange(1)) + return rearrange(latent, "b c t h w -> b (t h w) c"), ids[None] + + +def pack_action(values: torch.Tensor, seconds: torch.Tensor): + """``(1, D, K)`` values at ``seconds (K,)`` -> tokens ``(1, K, D)`` with ids ``(t, 0, 0, 0)``.""" + ids = _cartesian_ids(t=_times_to_ids(seconds), h=1, w=1, l_coord=torch.arange(1)) + return values.transpose(1, 2), ids[None] + + +def text_ids(length: int) -> torch.Tensor: + return _cartesian_ids(t=torch.arange(1), h=1, w=1, l_coord=torch.arange(length))[ + None + ] + + +# ---------------------------------------------------------------- observation +def _as_image(value: Any) -> torch.Tensor: + """A request image -> float ``(3, H, W)`` in ``[0, 1]``. + + Arrays and PIL images are HWC: uint8 pixels, or floats already in + ``[0, 1]``. Tensors are uint8 ``(H, W, 3)`` or float ``(3, H, W)`` in + ``[0, 1]`` (the reference policy's conventions). + """ + if isinstance(value, torch.Tensor): + if value.dtype == torch.uint8 and value.ndim == 3 and value.shape[-1] == 3: + return value.permute(2, 0, 1).float().div_(255.0) + if value.is_floating_point() and value.ndim == 3 and value.shape[0] == 3: + return _check_unit_range(value.float()) + raise ValueError( + "image tensors must be uint8 (H, W, 3) or float (3, H, W), " + f"got {value.dtype} {tuple(value.shape)}" + ) + array = np.asarray(value) + if array.ndim != 3 or array.shape[-1] != 3: + raise ValueError(f"expected an HWC RGB image, got shape {array.shape}") + if np.issubdtype(array.dtype, np.integer) and array.dtype != np.uint8: + # JSON pixel lists decode to int64. + if array.size and (array.min() < 0 or array.max() > 255): + raise ValueError("integer images must hold pixel values in [0, 255]") + array = array.astype(np.uint8) + tensor = torch.from_numpy(np.require(array, requirements=["C", "W"])) + if array.dtype == np.uint8: + return tensor.permute(2, 0, 1).float().div_(255.0) + if np.issubdtype(array.dtype, np.floating): + return _check_unit_range(tensor.permute(2, 0, 1).float()) + raise ValueError(f"unsupported image dtype {array.dtype}") + + +def _check_unit_range(image: torch.Tensor) -> torch.Tensor: + if image.numel() and (image.min() < 0.0 or image.max() > 1.0): + raise ValueError("float images must be in [0, 1]") + return image + + +def _canonical_camera(name: str, aliases: dict[str, str]) -> str: + name = name.removeprefix("observation.images.").removeprefix("images.") + return aliases.get(name, name) + + +def _pad_composite(composite: torch.Tensor, canvas_hw: tuple[int, int]) -> torch.Tensor: + """DROID composite ``(3, 540, 640)`` in ``[0, 1]`` -> canvas ``(3, Hc, Wc)`` in ``[-1, 1]`` (reflect pad).""" + if tuple(composite.shape) != (3, *DROID_COMPOSITE_HW): + raise ValueError(f"composite must be 3x540x640, got {tuple(composite.shape)}") + pad_right = canvas_hw[1] - composite.shape[-1] + pad_bottom = canvas_hw[0] - composite.shape[-2] + if pad_right < 0 or pad_bottom < 0: + raise ValueError(f"canvas {canvas_hw} is smaller than the DROID composite") + canvas = F.pad(composite[None], (0, pad_right, 0, pad_bottom), mode="reflect")[0] + return canvas.mul_(2.0).sub_(1.0) + + +def _grid_canvas(cams: list[torch.Tensor], canvas_hw: tuple[int, int]) -> torch.Tensor: + cols = int(np.ceil(np.sqrt(len(cams)))) + rows = int(np.ceil(len(cams) / cols)) + cell_h, cell_w = canvas_hw[0] // rows, canvas_hw[1] // cols + canvas = cams[0].new_zeros(3, *canvas_hw) + for i, cam in enumerate(cams): + r, c = divmod(i, cols) + canvas[:, r * cell_h : (r + 1) * cell_h, c * cell_w : (c + 1) * cell_w] = ( + F.interpolate( + cam[None], + size=(cell_h, cell_w), + mode="bilinear", + align_corners=False, + antialias=True, + )[0] + ) + return canvas + + +def _compose_canvas( + cams: list[torch.Tensor], layout: str, canvas_hw: tuple[int, int] +) -> torch.Tensor: + """Camera frames ``(3, H, W)`` in ``[0, 1]`` (layout order) -> canvas in ``[-1, 1]``.""" + if layout == "droid": + if len(cams) != 3 or any(tuple(c.shape[-2:]) != DROID_CAMERA_HW for c in cams): + raise ValueError( + "droid layout needs three 360x640 cameras [wrist, left, right]" + ) + wrist, left, right = cams + half = (DROID_CAMERA_HW[0] // 2, DROID_CAMERA_HW[1] // 2) + bottom = torch.cat( + [ + F.interpolate( + cam[None], size=half, mode="bilinear", align_corners=False + )[0] + for cam in (left, right) + ], + dim=-1, + ) + return _pad_composite(torch.cat([wrist, bottom], dim=-2), canvas_hw=canvas_hw) + if layout == "single": + if len(cams) != 1: + raise ValueError(f"single layout needs exactly one camera, got {len(cams)}") + canvas = F.interpolate( + cams[0][None], + size=canvas_hw, + mode="bilinear", + align_corners=False, + antialias=True, + )[0] + elif layout == "side_by_side": + if len(cams) != 2 or canvas_hw[1] % 2: + raise ValueError( + "side_by_side layout needs two cameras and an even canvas width" + ) + canvas = torch.cat( + [ + F.interpolate( + cam[None], + size=(canvas_hw[0], canvas_hw[1] // 2), + mode="bilinear", + align_corners=False, + antialias=True, + )[0] + for cam in cams + ], + dim=-1, + ) + elif layout == "grid": + canvas = _grid_canvas(cams, canvas_hw) + else: + raise ValueError(f"unknown camera layout {layout!r}") + return canvas.mul_(2.0).sub_(1.0) + + +def _observation_canvas( + observation: dict[str, Any], config: Flux3ActionPipelineConfig +) -> torch.Tensor: + # OpenPI clients may also send cameras as top-level "observation.images.". + named = { + k: v for k, v in observation.items() if k.startswith("observation.images.") + } + named.update(observation.get("images") or {}) + images = { + _canonical_camera(name, config.camera_aliases): value + for name, value in named.items() + } + if "composite" in images: + if config.camera_layout != "droid": + raise ValueError("a composite image requires the droid camera layout") + return _pad_composite( + _as_image(images["composite"]), canvas_hw=config.canvas_hw + ) + missing = [key for key in config.image_keys if key not in images] + if missing: + raise KeyError( + f"observation lacks cameras {missing}; expected {list(config.image_keys)} " + "or a droid 'composite'" + ) + cams = [_as_image(images[key]) for key in config.image_keys] + if config.camera_layout == "grid": + hw = (max(c.shape[-2] for c in cams), max(c.shape[-1] for c in cams)) + cams = [ + ( + F.interpolate( + c[None], + size=hw, + mode="bilinear", + align_corners=False, + antialias=True, + )[0] + if tuple(c.shape[-2:]) != hw + else c + ) + for c in cams + ] + return _compose_canvas( + cams, layout=config.camera_layout, canvas_hw=config.canvas_hw + ) + + +def _observation_state(observation: dict[str, Any], state_dim: int) -> torch.Tensor: + state = observation.get("state") + if state is None: + state = observation.get("observation.state") + if state is None: + raise KeyError("observation lacks 'state'") + state = torch.as_tensor(np.asarray(state, dtype=np.float32)).reshape(-1) + if state.shape != (state_dim,) or not torch.isfinite(state).all(): + raise ValueError( + f"state must be {state_dim} finite values, got {tuple(state.shape)}" + ) + return state + + +def parse_observation( + observation: dict[str, Any], config: Flux3ActionPipelineConfig +) -> Flux3ActionObservation: + prompt = observation.get("prompt") or observation.get("task") or "" + if isinstance(prompt, (list, tuple)): + if len(prompt) != 1: + raise ValueError("FLUX 3 Action serves one observation per request") + prompt = prompt[0] + return Flux3ActionObservation( + canvas=_observation_canvas(observation, config), + state=_observation_state(observation, state_dim=config.state_dim), + prompt=str(prompt), + ) + + +def _warmup_observation(config: Flux3ActionPipelineConfig) -> Flux3ActionObservation: + canvas = torch.full((3, *config.canvas_hw), GRAY_LEVEL / 255.0 * 2.0 - 1.0) + return Flux3ActionObservation( + canvas=canvas, state=torch.zeros(config.state_dim), prompt="" + ) + + +# ---------------------------------------------------------------- action space +def _flip_gripper(x: torch.Tensor, dims: tuple[int, ...]) -> torch.Tensor: + """``x -> 1 - x`` on the gripper dims (self-inverse).""" + if not dims: + return x + x = x.clone() + x[..., list(dims)] = 1.0 - x[..., list(dims)] + return x + + +def _bounds(stats: dict[str, list[float]], like: torch.Tensor): + q01 = torch.as_tensor(stats["q01"], dtype=like.dtype, device=like.device) + q99 = torch.as_tensor(stats["q99"], dtype=like.dtype, device=like.device) + span = q99 - q01 + return q01, torch.where(span > 1e-6, span, torch.ones_like(span)) + + +def normalize(x: torch.Tensor, stats, clip: float) -> torch.Tensor: + if stats is None: + return x + q01, span = _bounds(stats, x) + return (2.0 * (x - q01) / span - 1.0).clamp_(-clip, clip) + + +def denormalize(x: torch.Tensor, stats) -> torch.Tensor: + if stats is None: + return x + q01, span = _bounds(stats, x) + return (x + 1.0) * span / 2.0 + q01 + + +def targets_to_actions( + targets: torch.Tensor, state: torch.Tensor, config: Flux3ActionPipelineConfig +) -> torch.Tensor: + """Normalized targets ``(K, D)`` and the observed state ``(D,)`` -> absolute commands (dataset units).""" + flipped_state = _flip_gripper(state, dims=config.gripper_flip_dims) + actions = denormalize(targets, stats=config.action_normalization) + if config.action_parameterization == "joint_delta": + integrated = flipped_state[None] + torch.cumsum(actions, dim=0) + if config.absolute_action_dims: + dims = list(config.absolute_action_dims) + integrated[..., dims] = actions[..., dims] + actions = integrated + return _flip_gripper(actions, dims=config.gripper_flip_dims) + + +# ---------------------------------------------------------------- samplers +# Flow matching over a dict of streams: x_t = t * eps + (1 - t) * x0, velocity +# eps - x0; the solver state stays fp32. +Samples = dict[str, torch.Tensor] +Predictor = Callable[[Samples, float], Samples] +NUM_TRAIN_TIMESTEPS = 1000 + + +def cosmos_unipc( + samples: Samples, + predict: Predictor, + *, + scheduler: FlowUniPCMultistepScheduler, + n_steps: int, + shift: float, +) -> Samples: + """Solve with a private copy of ``scheduler`` per stream (UniPC keeps per-stream history).""" + device = next(iter(samples.values())).device + schedulers = {k: copy.deepcopy(scheduler) for k in samples} + for stream_scheduler in schedulers.values(): + stream_scheduler.set_timesteps(n_steps, device=device, shift=shift) + # Ticks can repeat at high step counts; index by step, not by tick. + stream_scheduler.set_begin_index(0) + for tick in next(iter(schedulers.values())).timesteps: + # The reference feeds float32(tick) / 1000 to the model. + t = torch.tensor(float(tick), dtype=torch.float32) / NUM_TRAIN_TIMESTEPS + velocity = predict(samples, t.item()) + samples = { + k: schedulers[k].step(velocity[k], tick, samples[k], return_dict=False)[0] + for k in samples + } + return samples + + +def _cfg_parallel_policy( + contexts: list[Flux3SegmentState], server_args: ServerArgs +) -> CFGPolicy | None: + """CFG-parallel branches of the denoising passes; None runs every context locally.""" + if not server_args.enable_cfg_parallel: + return None + if len(contexts) == 1: + logger.warning_once( + "CFG parallel is enabled but the request has no CFG; " + "every rank runs the same single pass" + ) + return None + cond, uncond = contexts + return CFGPolicy( + branches=[ + CFGBranch("conditional", True, {"context": cond}), + CFGBranch("unconditional", False, {"context": uncond}), + ], + parallel_uses_serial_arithmetic=True, + ) + + +# ---------------------------------------------------------------- stages +class Flux3ActionPreprocessStage(PipelineStage): + """Observation dict -> canvas, state and prompt.""" + + def __init__(self, config: Flux3ActionPipelineConfig): + super().__init__() + self.config = config + + def forward(self, batch: Req, server_args: ServerArgs) -> Req: + start = time.perf_counter() + state = vla_state(batch) + if batch.is_warmup: + state["flux3_observation"] = _warmup_observation(self.config) + else: + observation = dict(state.get("observation") or {}) + observation.setdefault("prompt", batch.prompt) + state["flux3_observation"] = parse_observation(observation, self.config) + vla_timings(batch)["preprocess_ms"] = (time.perf_counter() - start) * 1000 + return batch + + +class Flux3ActionTextEncodingStage(PipelineStage): + """Prompt (and CFG negative) -> DiT-encoded text contexts, cached per caption.""" + + def __init__( + self, + config: Flux3ActionPipelineConfig, + transformer: Flux3Transformer, + text_encoder: Flux3TextEncoder, + ): + super().__init__() + self.config = config + self.transformer = transformer + self.text_encoder = text_encoder + self._contexts: OrderedDict[str, Flux3SegmentState] = OrderedDict() + self._cached_tokens = 0 + + def component_uses( + self, server_args: ServerArgs, stage_name: str | None = None + ) -> list[ComponentUse]: + stage_name = self._component_stage_name(stage_name) + # Contexts are cached per caption: both components run only on a miss. + return [ + ComponentUse( + stage_name=stage_name, + component_name=name, + allow_prefetch=False, + start_at_stage_entry=False, + ) + for name in ("text_encoder", "transformer") + ] + + def _encode(self, captions: list[str], device: torch.device) -> dict: + with self.use_declared_component( + component_name="text_encoder", module=self.text_encoder + ): + # One caption per forward: batching changes bf16 GEMM results. + encoded = [self.text_encoder.encode([c])[0] for c in captions] + contexts = {} + with ( + self.use_declared_component( + component_name="transformer", module=self.transformer + ), + set_forward_context(current_timestep=0, attn_metadata=None), + ): + for caption, ctx in zip(captions, encoded): + contexts[caption] = self.transformer.encode_context( + ctx=ctx.to(device), ctx_ids=text_ids(ctx.shape[1]).to(device) + ) + return contexts + + def _contexts_for( + self, captions: list[str], device: torch.device, *, use_cache: bool + ) -> tuple[list[Flux3SegmentState], bool]: + """DiT-encoded text contexts of ``captions`` and whether all were cached.""" + cached = self._contexts if use_cache else {} + todo = [c for c in dict.fromkeys(captions) if c not in cached] + fresh = self._encode(todo, device) if todo else {} + contexts = [fresh.get(c) or cached[c] for c in captions] + if use_cache: + self._remember(captions, fresh) + return contexts, not todo + + def _remember( + self, captions: list[str], fresh: dict[str, Flux3SegmentState] + ) -> None: + for caption, context in fresh.items(): + self._contexts[caption] = context + self._cached_tokens += context.length + for caption in captions: + self._contexts.move_to_end(caption) + while ( + self._cached_tokens > self.config.caption_cache_max_tokens + and self._contexts + ): + _, evicted = self._contexts.popitem(last=False) + self._cached_tokens -= evicted.length + + def forward(self, batch: Req, server_args: ServerArgs) -> Req: + start = time.perf_counter() + state = vla_state(batch) + observation: Flux3ActionObservation = state["flux3_observation"] + options = vla_options(batch) + guidance = self.config.resolve_guidance( + options.get("guidance_scale"), options.get("guidance_scale_action") + ) + captions = [observation.prompt] + if any(g != 1.0 for g in guidance.values()): + captions.append("") + use_cache = bool(options.get("enable_prefix_cache", True)) + state["flux3_contexts"], hit = self._contexts_for( + captions, get_local_torch_device(), use_cache=use_cache + ) + state["flux3_guidance"] = guidance + state["cache"] = { + "enabled": use_cache, + "hit": hit, + "scope": "caption", + "mode": "exact", + } + vla_timings(batch)["text_ms"] = (time.perf_counter() - start) * 1000 + return batch + + +class Flux3ActionObservationEncodingStage(PipelineStage): + """Observed frame (VAE) and robot state -> DiT-encoded conditioning streams.""" + + def __init__( + self, + config: Flux3ActionPipelineConfig, + transformer: Flux3Transformer, + vae: Flux3VideoVAE, + ): + super().__init__() + self.config = config + self.transformer = transformer + self.vae = vae + + def component_uses( + self, server_args: ServerArgs, stage_name: str | None = None + ) -> list[ComponentUse]: + stage_name = self._component_stage_name(stage_name) + return [ + ComponentUse(stage_name=stage_name, component_name="vae"), + ComponentUse( + stage_name=stage_name, + component_name="transformer", + start_at_stage_entry=False, + ), + ] + + def _state_tokens(self, state: torch.Tensor, device: torch.device): + cfg = self.config + flipped = _flip_gripper(state, dims=cfg.gripper_flip_dims) + token = normalize( + flipped, stats=cfg.state_normalization, clip=cfg.normalization_clip + ) + values = (token[None, :, None] * cfg.action_scale).to(device) + return pack_action(values, seconds=torch.zeros(1)) + + def forward(self, batch: Req, server_args: ServerArgs) -> Req: + start = time.perf_counter() + cfg = self.config + device = get_local_torch_device() + state = vla_state(batch) + observation: Flux3ActionObservation = state["flux3_observation"] + h, w = cfg.latent_hw + with self.use_declared_component(component_name="vae", module=self.vae): + frame = observation.canvas.to(device, torch.bfloat16)[None] + latent = self.vae.encode_frame(frame)[..., :h, :w] + video, video_ids = pack_video(latent, first_frame=0, fps=cfg.fps) + action, action_ids = self._state_tokens(observation.state, device) + zero = torch.zeros(1, device=device) + with ( + self.use_declared_component( + component_name="transformer", module=self.transformer + ), + set_forward_context(current_timestep=0, attn_metadata=None), + ): + state["flux3_conditioning"] = [ + self.transformer.encode_stream( + name="video_cond", x=video, ids=video_ids.to(device), timesteps=zero + ), + self.transformer.encode_stream( + name=f"{cfg.action_modality}_cond", + x=action, + ids=action_ids.to(device), + timesteps=zero, + ), + ] + vla_timings(batch)["observation_ms"] = (time.perf_counter() - start) * 1000 + return batch + + +class Flux3ActionDenoisingStage(PipelineStage): + """Joint video + action flow matching from noise; stores the action chunk.""" + + def __init__( + self, + config: Flux3ActionPipelineConfig, + transformer: Flux3Transformer, + scheduler: FlowUniPCMultistepScheduler, + ): + super().__init__() + self.config = config + self.transformer = transformer + self.scheduler = scheduler + + def component_uses( + self, server_args: ServerArgs, stage_name: str | None = None + ) -> list[ComponentUse]: + return [ + ComponentUse( + stage_name=self._component_stage_name(stage_name), + component_name="transformer", + phase="denoise", + preferred_ready_after_request=True, + memory_intensive=True, + ) + ] + + def _noised_streams(self, seed: int): + cfg = self.config + h, w = cfg.latent_hw + # Latent frames 1.. of the (chunk + 1)-frame window. + n_pred = cfg.action_horizon // TEMPORAL_DOWNSAMPLE + # Draw order (video, then action) on a CPU generator matches the reference. + rng = torch.Generator().manual_seed(seed) + video_noise = torch.randn(1, LATENT_CHANNELS, n_pred, h, w, generator=rng) + action_noise = torch.randn(1, cfg.action_dim, cfg.action_horizon, generator=rng) + video, video_ids = pack_video(video_noise, first_frame=1, fps=cfg.fps) + seconds = (torch.arange(cfg.action_horizon).float() + 1) / cfg.fps + action, action_ids = pack_action(action_noise, seconds=seconds) + return {"video": video, cfg.action_modality: action}, { + "video": video_ids, + cfg.action_modality: action_ids, + } + + def forward(self, batch: Req, server_args: ServerArgs) -> Req: + cfg = self.config + device = get_local_torch_device() + state = vla_state(batch) + observation: Flux3ActionObservation = state["flux3_observation"] + contexts: list[Flux3SegmentState] = state["flux3_contexts"] + conditioning: list[Flux3SegmentState] = state["flux3_conditioning"] + guidance: dict[str, float] = state["flux3_guidance"] + steps = batch.num_inference_steps + seed = batch.seed[0] if isinstance(batch.seed, list) else batch.seed + horizon = batch.action_horizon or cfg.action_horizon + if horizon > cfg.action_horizon: + raise ValueError( + f"action_horizon {horizon} exceeds the policy chunk {cfg.action_horizon}" + ) + + start = time.perf_counter() + samples, ids = self._noised_streams(seed) + samples = {k: v.to(device) for k, v in samples.items()} + ropes = {k: self.transformer.rope(v.to(device)) for k, v in ids.items()} + video_cond, action_cond = conditioning + order = list(samples) # joint sequence: video, video_cond, action, action_cond + cfg_policy = _cfg_parallel_policy(contexts, server_args) + step = 0 + + def predict( + current: dict[str, torch.Tensor], t: float + ) -> dict[str, torch.Tensor]: + nonlocal step + timestep = torch.full((1,), t, device=device, dtype=torch.float32) + with set_forward_context( + current_timestep=step, attn_metadata=None, forward_batch=batch + ): + targets = { + name: self.transformer.encode_stream( + name=name, + x=current[name].to(torch.bfloat16), + ids=None, + timesteps=timestep, + rope=ropes[name], + ) + for name in order + } + streams = [targets["video"], video_cond, targets[order[1]], action_cond] + + def denoise(ctx: Flux3SegmentState) -> tuple[torch.Tensor, ...]: + out = self.transformer.denoise( + context=ctx, streams=streams, targets=order + ) + return tuple(out[k] for k in order) + + if cfg_policy is not None: + # Each rank runs one branch; all ranks get both predictions. + raw = run_cfg_parallel( + cfg_policy, lambda branch: denoise(branch.kwargs["context"]) + ) + else: + raw = [denoise(ctx) for ctx in contexts] + preds = [dict(zip(order, p)) for p in raw] + step += 1 + if len(preds) == 1: + return {k: v.float() for k, v in preds[0].items()} + cond, uncond = preds + return { + k: (uncond[k] + guidance[k] * (cond[k] - uncond[k])).float() + for k in order + } + + with self.use_declared_component( + component_name="transformer", module=self.transformer + ): + result = cosmos_unipc( + samples, + predict, + scheduler=self.scheduler, + n_steps=steps, + shift=cfg.sampler_shift, + ) + targets = result[cfg.action_modality][0].float() / cfg.action_scale + actions = targets_to_actions( + targets, state=observation.state.to(device), config=cfg + ) + state["actions"] = actions[None, :horizon] + synchronize_vla_action_tensor(actions) + vla_timings(batch)["denoise_ms"] = (time.perf_counter() - start) * 1000 + return batch diff --git a/python/sglang/multimodal_gen/test/server/gpu_cases.py b/python/sglang/multimodal_gen/test/server/gpu_cases.py index 102ec78aab67..83c62dcd36fc 100644 --- a/python/sglang/multimodal_gen/test/server/gpu_cases.py +++ b/python/sglang/multimodal_gen/test/server/gpu_cases.py @@ -20,6 +20,7 @@ DiffusionSamplingParams, DiffusionServerArgs, DiffusionTestCase, + FLUX3_ACTION_CI_sampling_params, IDEOGRAM4_CI_sampling_params, JOY_ECHO_T2V_CI_sampling_params, LINGBOT_VIDEO_T2V_CI_sampling_params, @@ -120,6 +121,18 @@ run_component_accuracy_check=False, run_t2v_input_reference_check=False, ), + DiffusionTestCase( + "flux3_action_http", + DiffusionServerArgs( + model_path="black-forest-labs/flux-3-action-droid", + ), + FLUX3_ACTION_CI_sampling_params, + run_perf_check=False, + perf_warmup_requests=1, + # No Diffusers counterpart to compare components against. + 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/perf_baselines/h100.json b/python/sglang/multimodal_gen/test/server/perf_baselines/h100.json index 926267e76c38..6e236784f981 100644 --- a/python/sglang/multimodal_gen/test/server/perf_baselines/h100.json +++ b/python/sglang/multimodal_gen/test/server/perf_baselines/h100.json @@ -63,6 +63,14 @@ "expected_avg_denoise_ms": 0.0, "expected_median_denoise_ms": 0.0 }, + "flux3_action_http": { + "expected_load_ms": 46252.88, + "stages_ms": {}, + "denoise_step_ms": {}, + "expected_e2e_ms": 442.53, + "expected_avg_denoise_ms": 0.0, + "expected_median_denoise_ms": 0.0 + }, "qwen_image_t2i": { "expected_load_ms": 40799.72, "stages_ms": { 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 9b87e5adad29..29efa8bee96d 100644 --- a/python/sglang/multimodal_gen/test/server/test_server_utils.py +++ b/python/sglang/multimodal_gen/test/server/test_server_utils.py @@ -1720,6 +1720,9 @@ def generate_action(case_id, client) -> tuple[str, bytes]: 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)) + # Policies with a fixed camera resolution (e.g. DROID 360x640) set both. + image_height = int(extra.get("image_height", image_size)) + image_width = int(extra.get("image_width", image_size)) camera_order = tuple( extra.get( "camera_order", @@ -1735,8 +1738,8 @@ def tensor_payload(array): } def image_payload(camera_index: int): - y = np.arange(image_size, dtype=np.uint16)[:, None] - x = np.arange(image_size, dtype=np.uint16)[None, :] + y = np.arange(image_height, dtype=np.uint16)[:, None] + x = np.arange(image_width, dtype=np.uint16)[None, :] image = np.stack( ( (x + camera_index * 17) % 256 + np.zeros_like(y), diff --git a/python/sglang/multimodal_gen/test/server/testcase_configs.py b/python/sglang/multimodal_gen/test/server/testcase_configs.py index 1840f84becd3..059b5eb40af8 100644 --- a/python/sglang/multimodal_gen/test/server/testcase_configs.py +++ b/python/sglang/multimodal_gen/test/server/testcase_configs.py @@ -451,6 +451,28 @@ def __post_init__(self) -> None: ) +# DROID policy: three fixed-name 360x640 cameras, 8-dim state and actions, +# the package recipe (4 steps, CFG on video). Noise comes from the seed. +FLUX3_ACTION_CI_sampling_params = DiffusionSamplingParams( + prompt="put the marker in the cup", + extras={ + "action_horizon": 32, + "action_dim": 8, + "state_dim": 8, + "image_height": 360, + "image_width": 640, + "camera_order": ("wrist", "left", "right"), + "num_inference_steps": 4, + "seed": 0, + "enable_prefix_cache": False, + # Same path is bit-exact across runs and GPUs. Kernel swaps move actions + # by up to max 0.064 / mean 0.020 (eager QK-norm+RoPE in every block). + "action_max_abs_diff_threshold": 0.2, + "action_mean_abs_diff_threshold": 0.05, + }, +) + + def sample_step_indices( step_map: dict[int, float], fractions: Sequence[float] ) -> list[int]: diff --git a/python/sglang/multimodal_gen/test/test_utils.py b/python/sglang/multimodal_gen/test/test_utils.py index b400b771c10c..9152e5512b3a 100644 --- a/python/sglang/multimodal_gen/test/test_utils.py +++ b/python/sglang/multimodal_gen/test/test_utils.py @@ -40,7 +40,7 @@ # NPU/ascend) is read from sgl-project/ci-data-diffusion, where the GT-gen workflows # publish. SGL_TEST_FILES_CI_DATA_REPO = "sgl-project/ci-data-diffusion" -SGL_TEST_FILES_CI_DATA_REVISION = "38ba32bd812b2dfb0eccc83ef063096c089e3389" +SGL_TEST_FILES_CI_DATA_REVISION = "dbb70135da54b1f2b172886b28e813d2bcdd16c1" # The NPU pin is kept as a separate branch so ascend GT can be bumped independently # when it's regenerated on its own cadence. diff --git a/python/sglang/multimodal_gen/test/unit/test_flux3_action.py b/python/sglang/multimodal_gen/test/unit/test_flux3_action.py new file mode 100644 index 000000000000..397be88129e1 --- /dev/null +++ b/python/sglang/multimodal_gen/test/unit/test_flux3_action.py @@ -0,0 +1,828 @@ +# SPDX-License-Identifier: Apache-2.0 + +import json +from types import SimpleNamespace + +import numpy as np +import pytest +import torch + +from sglang.multimodal_gen.configs.models.dits.flux3 import ( + Flux3ArchConfig, + Flux3DiTConfig, +) +from sglang.multimodal_gen.configs.pipeline_configs.flux3_action import ( + Flux3ActionPipelineConfig, + _validate_parallelism, + flux3_action_variant_subfolder, + is_flux3_action_package, + read_flux3_action_config, +) +from sglang.multimodal_gen.configs.sample.flux3_action import Flux3ActionSamplingParams +from sglang.multimodal_gen.runtime.entrypoints.action.protocol import ( + action_metadata, + build_action_sampling_params, +) +from sglang.multimodal_gen.runtime.models.schedulers.scheduling_flow_unipc_multistep import ( + FlowUniPCMultistepScheduler, +) +from sglang.multimodal_gen.runtime.pipelines_core.stages.model_specific_stages.flux3_action import ( + _cfg_parallel_policy, + cosmos_unipc, + denormalize, + normalize, + pack_action, + pack_video, + parse_observation, + targets_to_actions, + text_ids, +) + +DROID_CONFIG = { + "action_dim": 8, + "action_modality": "action_prediction_droid", + "camera_layout": "droid", + "camera_keys": ["images.wrist", "images.left", "images.right"], + "canvas_hw": [544, 736], + "chunk_size": 32, + "n_action_steps": 32, + "fps": 15.0, + "action_scale": 2.0, + "gripper_flip_dims": [-1], + "action_parameterization": "absolute", + "absolute_action_dims": [], + "action_normalization": None, + "state_normalization": None, + "normalization_clip": 6.0, + "video_vae_id": "black-forest-labs/flux-3-action-base:video_vae.safetensors@rev", + "text_encoder_id": "black-forest-labs/flux-3-action-base:text_encoder@rev", + "dit_config": {}, + "content_streams": ["video", "video_cond"], + "torch_dtype": "bfloat16", + "quantization": None, + "inference_profile": "default", + "sampler": "cosmos_unipc", + "num_inference_steps": 4, + "guidance_scale": 4.0, + "guidance_scale_action": 1.0, + "sampler_shift": 5.0, + "inference_seed": 0, +} + + +def _droid_config() -> Flux3ActionPipelineConfig: + config = Flux3ActionPipelineConfig() + config.load_policy_config(dict(DROID_CONFIG)) + return config + + +def _write_package(root) -> None: + (root / "config.native.json").write_text(json.dumps(DROID_CONFIG)) + (root / "manifest.json").write_text(json.dumps({"kind": "policy_export"})) + + +# ---------------------------------------------------------------- config / registry +def test_policy_config_adopts_droid_recipe(): + config = _droid_config() + assert config.image_keys == ("wrist", "left", "right") + assert config.latent_hw == (17, 20) + assert config.default_num_inference_steps == 4 + assert (config.guidance_scale, config.guidance_scale_action) == (4.0, 1.0) + arch = config.dit_config.arch_config + assert arch.in_channels == { + "video": 96, + "video_cond": 96, + "action_prediction_droid": 8, + "action_prediction_droid_cond": 8, + } + assert list(arch.sequence) == [ + "x_video", + "x_video_cond", + "x_action_prediction_droid", + "x_action_prediction_droid_cond", + ] + + +def test_policy_config_rejects_unsupported_profiles(): + config = Flux3ActionPipelineConfig() + with pytest.raises(NotImplementedError): + config.load_policy_config({**DROID_CONFIG, "inference_profile": "history"}) + + +def _parallel_args(**overrides): + args = dict( + num_gpus=2, + enable_cfg_parallel=True, + cfg_parallel_degree=2, + tp_size=1, + sp_degree=1, + ulysses_degree=1, + ring_degree=1, + ) + return SimpleNamespace(**{**args, **overrides}) + + +def test_multi_gpu_layouts_are_tp_times_sp_times_cfg(): + """Layouts the DiT does not shard (ring, uneven heads) must be refused, not run.""" + no_cfg = dict(enable_cfg_parallel=False, cfg_parallel_degree=1) + for overrides in ( + dict(num_gpus=1, **no_cfg), + dict(), + dict(tp_size=2, **no_cfg), + dict(sp_degree=2, ulysses_degree=2, **no_cfg), + dict(num_gpus=8, tp_size=2, sp_degree=2, ulysses_degree=2), + ): + _validate_parallelism(_parallel_args(**overrides), num_heads=24) + for overrides in ( + dict(sp_degree=2, ring_degree=2, **no_cfg), + dict(num_gpus=4, cfg_parallel_degree=4), + dict(num_gpus=4), + ): + with pytest.raises(NotImplementedError): + _validate_parallelism(_parallel_args(**overrides), num_heads=24) + with pytest.raises(ValueError): + _validate_parallelism( + _parallel_args(num_gpus=5, sp_degree=5, ulysses_degree=5, **no_cfg), + num_heads=24, + ) + + +def test_cfg_parallel_branches_are_conditional_then_unconditional(): + """The combine reads ``cond, uncond = preds``; a swapped order inverts guidance.""" + cond, uncond = object(), object() + policy = _cfg_parallel_policy([cond, uncond], _parallel_args()) + assert [b.kwargs["context"] for b in policy.branches] == [cond, uncond] + assert [b.is_conditional for b in policy.branches] == [True, False] + assert policy.parallel_uses_serial_arithmetic + assert _cfg_parallel_policy([cond], _parallel_args()) is None + assert ( + _cfg_parallel_policy([cond, uncond], _parallel_args(enable_cfg_parallel=False)) + is None + ) + + +def test_variant_subfolders(): + assert flux3_action_variant_subfolder(None) == "" + assert flux3_action_variant_subfolder("base") == "" + assert flux3_action_variant_subfolder("gd-fp8r") == "variants/gd-fp8r" + with pytest.raises(ValueError): + flux3_action_variant_subfolder("nope") + + +def test_local_package_detection_and_registry(tmp_path): + from sglang.multimodal_gen.registry import ( + get_model_info, + get_non_diffusers_pipeline_name, + ) + + assert not is_flux3_action_package(str(tmp_path)) + _write_package(tmp_path) + assert is_flux3_action_package(str(tmp_path)) + assert ( + read_flux3_action_config(tmp_path)["action_modality"] + == "action_prediction_droid" + ) + assert get_non_diffusers_pipeline_name(str(tmp_path)) == "Flux3ActionPipeline" + for path in (str(tmp_path), "black-forest-labs/flux-3-action-droid"): + get_model_info.cache_clear() + info = get_model_info(path) + assert info.pipeline_cls.__name__ == "Flux3ActionPipeline" + assert info.sampling_param_cls is Flux3ActionSamplingParams + get_model_info.cache_clear() + + +# ---------------------------------------------------------------- protocol +def _server_args(config: Flux3ActionPipelineConfig) -> SimpleNamespace: + return SimpleNamespace( + model_id=None, + model_path="black-forest-labs/flux-3-action-droid", + served_model_name="flux3-action", + output_path=None, + comfyui_mode=False, + backend=None, + pipeline_class_name=None, + pipeline_config=config, + ) + + +def test_action_request_builds_flux3_sampling_params(): + image = np.zeros((360, 640, 3), dtype=np.uint8) + payload = { + "input": { + "task": "pour the cup", + "observation": { + "images": {"wrist": image, "left": image, "right": image}, + "state": [0.0] * 8, + }, + }, + "parameters": {"guidance_scale": 2.0, "seed": 7}, + } + params = build_action_sampling_params(payload, _server_args(_droid_config())) + assert isinstance(params, Flux3ActionSamplingParams) + assert params.prompt == "pour the cup" + assert params.num_inference_steps == 4 + assert params.guidance_scale == 2.0 and params.guidance_scale_action is None + assert params.seed == 7 + extra = params.build_request_extra()["vla"] + assert set(extra["observation"]["images"]) == {"wrist", "left", "right"} + assert extra["options"]["guidance_scale"] == 2.0 + + +def test_action_metadata_reports_policy_recipe(): + metadata = action_metadata(_server_args(_droid_config())) + assert metadata["policy_family"] == "flux3_action" + assert metadata["input"]["image_keys"] == ["wrist", "left", "right"] + assert metadata["output"]["action_horizon"] == 32 + assert metadata["defaults"]["guidance_scale"] == 4.0 + + +def test_sampling_params_reject_unsupported_requests(): + """Unsupported request fields must fail loudly instead of being ignored.""" + for kwargs in ( + {"guidance_scale": -1.0}, + {"guidance_scale": "4"}, + {"seed": [1, 2]}, + {"num_outputs_per_prompt": 2}, + {"action_horizon": 0}, + ): + with pytest.raises(ValueError): + Flux3ActionSamplingParams(**kwargs) + + +def test_guidance_resolution_follows_video_scale_when_action_is_unset(): + """Reference rule: an unset action scale follows the (possibly overridden) video scale.""" + config = Flux3ActionPipelineConfig() + config.load_policy_config({**DROID_CONFIG, "guidance_scale_action": None}) + assert config.resolve_guidance(None, None) == { + "video": 4.0, + "action_prediction_droid": 4.0, + } + assert config.resolve_guidance(2.0, None) == { + "video": 2.0, + "action_prediction_droid": 2.0, + } + droid = _droid_config() # explicit action scale in the package + assert droid.resolve_guidance(2.0, None) == { + "video": 2.0, + "action_prediction_droid": 1.0, + } + + +def test_reference_dit_config_fields_are_translated(): + """Exports carry the reference JointSingleSeqParams names and the full trunk's streams.""" + config = Flux3ActionPipelineConfig() + config.load_policy_config( + { + **DROID_CONFIG, + "dit_config": { + "num_heads": 24, + "depth": 5, + "attn_mode": "torch", + "in_channels": {"video": 96, "video_cond": 96, "image": 128}, + }, + } + ) + arch = config.dit_config.arch_config + assert arch.num_attention_heads == 24 + assert "image" not in arch.in_channels + with pytest.raises(NotImplementedError): + config.load_policy_config( + {**DROID_CONFIG, "dit_config": {"depth_late_blocks": 2}} + ) + + +def test_manifest_checksums_are_verified(tmp_path): + import hashlib + + from sglang.multimodal_gen.configs.pipeline_configs.flux3_action import ( + verify_flux3_action_manifest, + ) + + config_bytes = json.dumps(DROID_CONFIG).encode() + (tmp_path / "config.native.json").write_bytes(config_bytes) + manifest = { + "kind": "policy_export", + "sha256": {"config.native.json": hashlib.sha256(config_bytes).hexdigest()}, + } + (tmp_path / "manifest.json").write_text(json.dumps(manifest)) + verify_flux3_action_manifest(tmp_path, include_weights=False) + (tmp_path / "config.native.json").write_bytes(config_bytes + b" ") + with pytest.raises(ValueError, match="checksum"): + verify_flux3_action_manifest(tmp_path, include_weights=False) + + +# ---------------------------------------------------------------- observation / action space +def _views(seed: int = 0) -> dict[str, np.ndarray]: + rng = np.random.default_rng(seed) + return { + name: rng.integers(0, 256, (360, 640, 3), dtype=np.uint8) + for name in ("wrist", "left", "right") + } + + +def test_droid_canvas_from_cameras(): + config = _droid_config() + views = _views() + obs = parse_observation( + { + "images": {f"images.{k}": v for k, v in views.items()}, + "state": np.zeros(8), + "prompt": "x", + }, + config, + ) + assert tuple(obs.canvas.shape) == (3, 544, 736) + assert obs.canvas.min() >= -1.0 and obs.canvas.max() <= 1.0 + wrist = torch.from_numpy(views["wrist"]).permute(2, 0, 1).float() / 255 * 2 - 1 + torch.testing.assert_close(obs.canvas[:, :360, :640], wrist) + # reflect padding to the right of the 640-wide composite + torch.testing.assert_close(obs.canvas[:, :, 640], obs.canvas[:, :, 638]) + + +def test_droid_composite_matches_three_camera_canvas(): + """A client-built 540x640 composite and the three cameras give the same canvas.""" + config = _droid_config() + views = _views(1) + from_cameras = parse_observation({"images": views, "state": np.zeros(8)}, config) + composite = (from_cameras.canvas[:, :540, :640] + 1) / 2 + from_composite = parse_observation( + {"images": {"composite": composite}, "state": np.zeros(8)}, config + ) + torch.testing.assert_close(from_composite.canvas, from_cameras.canvas) + + +def test_lerobot_camera_names_and_float_images(): + config = _droid_config() + views = _views() + lerobot = { + "observation.images.wrist_image_left": views["wrist"], + "observation.images.exterior_image_1_left": views["left"], + "observation.images.exterior_image_2_left": views["right"], + } + by_alias = parse_observation({**lerobot, "observation.state": np.zeros(8)}, config) + floats = {k: v.astype(np.float32) / 255 for k, v in views.items()} + by_float = parse_observation({"images": floats, "state": np.zeros(8)}, config) + torch.testing.assert_close(by_alias.canvas, by_float.canvas) + with pytest.raises(ValueError, match=r"\[0, 1\]"): + pixels = {k: v.astype(np.float32) for k, v in views.items()} + parse_observation({"images": pixels, "state": np.zeros(8)}, config) + + +def test_json_decoded_pixel_lists_are_accepted(): + """Cookbook JSON requests send image.tolist(); the lists decode to int64 arrays.""" + config = _droid_config() + views = _views(2) + as_json = {k: np.asarray(v.tolist()) for k, v in views.items()} + assert next(iter(as_json.values())).dtype == np.int64 + from_json = parse_observation({"images": as_json, "state": [0.0] * 8}, config) + from_uint8 = parse_observation({"images": views, "state": np.zeros(8)}, config) + assert torch.equal(from_json.canvas, from_uint8.canvas) + with pytest.raises(ValueError, match="0, 255"): + parse_observation( + { + "images": {**as_json, "wrist": as_json["wrist"] + 256}, + "state": [0.0] * 8, + }, + config, + ) + + +def test_missing_camera_is_reported(): + views = _views() + del views["left"] + with pytest.raises(KeyError, match="left"): + parse_observation({"images": views, "state": np.zeros(8)}, _droid_config()) + + +def test_action_space_roundtrip_and_joint_delta(): + stats = {"q01": [-1.0, 0.0, 2.0], "q99": [1.0, 4.0, 2.0]} + x = torch.tensor([[0.5, 1.0, 2.0]]) + roundtrip = denormalize(normalize(x, stats=stats, clip=6.0), stats=stats) + torch.testing.assert_close(roundtrip, x) + + config = _droid_config() + config.action_parameterization = "joint_delta" + config.gripper_flip_dims = (-1,) + config.absolute_action_dims = (2,) + state = torch.tensor([1.0, 2.0, 0.25]) + deltas = torch.tensor([[0.1, 0.0, 0.2], [0.1, -1.0, 0.4]]) + actions = targets_to_actions(deltas, state=state, config=config) + expected = torch.tensor([[1.1, 2.0, 0.8], [1.2, 1.0, 0.6]]) + torch.testing.assert_close(actions, expected) + + +def test_position_ids_follow_the_10ms_clock(): + _, ids = pack_video(torch.zeros(1, 96, 2, 2, 3), first_frame=1, fps=15.0) + assert ids.shape == (1, 12, 4) + assert ids[0, 0].tolist() == [26, 0, 0, 0] # frame 1 at 4 / 15 s + assert ids[0, -1].tolist() == [53, 1, 2, 0] + seconds = (torch.arange(3).float() + 1) / 15.0 + tokens, ids = pack_action(torch.zeros(1, 8, 3), seconds=seconds) + assert tokens.shape == (1, 3, 8) + assert ids[0, :, 0].tolist() == [6, 13, 20] + assert text_ids(3)[0, :, 3].tolist() == [0, 1, 2] + + +# ---------------------------------------------------------------- samplers +def _unipc_scheduler() -> FlowUniPCMultistepScheduler: + # Same configuration as Flux3ActionPipeline.load_modules. + return FlowUniPCMultistepScheduler( + solver_order=2, + solver_type="bh2", + predict_x0=True, + lower_order_final=True, + final_sigmas_type="zero", + shift=1.0, + ) + + +def test_cosmos_unipc_feeds_reference_ticks_and_reuses_the_scheduler(): + """DROID recipe (4 steps, shift 5): the model sees ticks 999, 937, 833, 624, + and the pipeline's scheduler module is not consumed by a request.""" + scheduler = _unipc_scheduler() + for _ in range(2): + seen = [] + cosmos_unipc( + {"a": torch.zeros(1, 3, 2)}, + lambda samples, t: ( + seen.append(round(t * 1000)), + {"a": torch.zeros(1, 3, 2)}, + )[1], + scheduler=scheduler, + n_steps=4, + shift=5.0, + ) + assert seen == [999, 937, 833, 624] + + +def test_cosmos_unipc_recovers_x0_on_a_straight_flow(): + """Exact velocities on a straight flow must land on x0 (solver bookkeeping is consistent).""" + generator = torch.Generator().manual_seed(0) + x0 = { + "a": torch.randn(1, 2, 5, generator=generator), + "b": torch.randn(1, 3, 8, generator=generator), + } + eps = {k: torch.randn(v.shape, generator=generator) for k, v in x0.items()} + # UniPC starts at the shifted 0.999: 5 * 0.999 / (1 + 4 * 0.999) + sigma_max = 5 * 0.999 / (1 + 4 * 0.999) + start = {k: sigma_max * eps[k] + (1 - sigma_max) * x0[k] for k in x0} + result = cosmos_unipc( + start, + lambda samples, t: {k: eps[k] - x0[k] for k in samples}, + scheduler=_unipc_scheduler(), + n_steps=4, + shift=5.0, + ) + for k in x0: + torch.testing.assert_close(result[k], x0[k], atol=1e-4, rtol=0) + + +# ---------------------------------------------------------------- DiT +def _init_single_process_parallel() -> None: + from sglang.multimodal_gen.runtime.distributed.parallel_state import ( + maybe_init_distributed_environment_and_model_parallel, + model_parallel_is_initialized, + ) + from sglang.multimodal_gen.test.single_test_file.component_accuracy.utils import ( + ensure_distributed_env_defaults, + ) + + if not model_parallel_is_initialized(): + ensure_distributed_env_defaults() + maybe_init_distributed_environment_and_model_parallel(tp_size=1, sp_size=1) + + +def _tiny_arch() -> Flux3ArchConfig: + return Flux3ArchConfig( + hidden_size=64, + num_attention_heads=4, + depth=2, + depth_single_blocks=2, + axes_dim=(4, 4, 4, 4), + context_in_dim=32, + vec_in_dim=16, + ).with_streams({"act": 3, "act_cond": 3}) + + +@pytest.mark.skipif( + not torch.cuda.is_available(), reason="needs the CUDA parallel runtime" +) +def test_dit_loads_released_checkpoint_names(): + """Released checkpoints load by name; a renamed module or mapping breaks them.""" + from sglang.multimodal_gen.runtime.loader.utils import get_param_names_mapping + from sglang.multimodal_gen.runtime.models.dits.flux3 import Flux3Transformer + + _init_single_process_parallel() + with torch.device("meta"): + model = Flux3Transformer(Flux3DiTConfig(arch_config=_tiny_arch())) + params = set(model.state_dict()) + assert all("bias" not in n for n in params) + # Layerwise offload streams exactly these block lists; a stale name is skipped silently. + modules = dict(model.named_modules()) + blocks = {name: modules.get(name) for name in model.layer_names} + assert all(isinstance(m, torch.nn.ModuleList) and len(m) for m in blocks.values()) + assert set(blocks) == { + "txt_mode_blocks", + "single_blocks", + *( + f"content_mode_blocks.{m}" + for m in ("video", "video_cond", "act", "act_cond") + ), + } + mapping = get_param_names_mapping(model.param_names_mapping) + fused = {} + for name in ( + "dit.emb_in.act.weight", + "dit.txt_in.weight", + "dit.time_in.in_layer.weight", + "dit.vector_in.out_layer.weight", + "dit.early_stream_modulations.txt.lin.weight", + "dit.single_stream_modulations.act_cond.lin.weight", + "dit.content_mode_blocks.act.0.norm.query_norm.scale", + "dit.txt_mode_blocks.1.attn_out.weight", + "dit.single_blocks.1.mlp_out.weight", + "dit.final_layer.act.adaLN_modulation.1.weight", + "dit.final_layer.video.linear.weight", + "dit.content_mode_blocks.video.1.q_proj.weight", + "dit.content_mode_blocks.video.1.k_proj.weight", + "dit.content_mode_blocks.video.1.v_proj.weight", + "dit.content_mode_blocks.video.1.mlp_in.weight", + ): + target, index, count = mapping(name) + assert target in params, (name, target) + if index is not None: + fused.setdefault(target, set()).add((index, count)) + # q, k, v, mlp_in concatenate in this order into the fused projection + assert fused == { + "content_mode_blocks.video.1.qkv_mlp.weight": {(0, 4), (1, 4), (2, 4), (3, 4)} + } + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="needs CUDA attention") +def test_dit_cached_streams_match_full_forward(): + """The pipeline's per-caption / per-request caching must not change predictions.""" + from sglang.multimodal_gen.runtime.managers.forward_context import ( + set_forward_context, + ) + from sglang.multimodal_gen.runtime.models.dits.flux3 import Flux3Transformer + + _init_single_process_parallel() + torch.manual_seed(0) + model = Flux3Transformer(Flux3DiTConfig(arch_config=_tiny_arch())).cuda().bfloat16() + for p in model.parameters(): + torch.nn.init.normal_(p, std=0.05) + + video, video_ids = pack_video(torch.randn(1, 96, 2, 2, 3), first_frame=1, fps=15.0) + cond, cond_ids = pack_video(torch.randn(1, 96, 1, 2, 3), first_frame=0, fps=15.0) + act, act_ids = pack_action( + torch.randn(1, 3, 4), seconds=(torch.arange(4).float() + 1) / 15 + ) + state, state_ids = pack_action(torch.randn(1, 3, 1), seconds=torch.zeros(1)) + ctx = torch.randn(1, 5, 32).cuda() + t = 0.7 + streams = { + "x_video": (video, video_ids, t), + "x_video_cond": (cond, cond_ids, 0.0), + "x_act": (act, act_ids, t), + "x_act_cond": (state, state_ids, 0.0), + } + kwargs = {} + for key, (x, ids, ts) in streams.items(): + kwargs[key] = x.cuda() + kwargs[f"{key}_ids"] = ids.cuda() + kwargs[f"{key}_timesteps"] = torch.full(x.shape[:2], ts).cuda() + with torch.no_grad(), set_forward_context(current_timestep=0, attn_metadata=None): + full = model(ctx=ctx, ctx_ids=text_ids(5).cuda(), **kwargs) + context = model.encode_context(ctx, text_ids(5).cuda()) + encoded = [ + model.encode_stream( + name=key[2:], + x=x.cuda(), + ids=ids.cuda(), + timesteps=torch.full((1,), ts).cuda(), + ) + for key, (x, ids, ts) in streams.items() + ] + cached = model.denoise( + context=context, streams=encoded, targets=["video", "act"] + ) + # Same kernels in the same order: the cached path is bit-identical. + assert torch.equal(cached["video"], full["x_video"]) + assert torch.equal(cached["act"], full["x_act"]) + assert full["x_act_cond"].shape == (1, 1, 3) + + +@pytest.mark.skipif( + not torch.cuda.is_available() or torch.cuda.get_device_capability() < (8, 9), + reason="needs FP8 scaled_mm", +) +def test_fp8r_checkpoint_loads_fused_rowwise_linears(): + """Native FP8r payloads (E4M3 + per-row scales) must fuse and dequantize consistently.""" + from sglang.multimodal_gen.runtime.models.dits.flux3 import ( + Flux3Fp8RowwiseLinear, + Flux3Transformer, + load_fp8r_checkpoint, + quantize_fp8_rowwise, + ) + + _init_single_process_parallel() + with torch.device("meta"): + model = Flux3Transformer(Flux3DiTConfig(arch_config=_tiny_arch())) + reference = {} + state = {} + generator = torch.Generator().manual_seed(0) + for name, tensor in model.state_dict().items(): + for part in ( + ("q_proj", "k_proj", "v_proj", "mlp_in") if ".qkv_mlp." in name else (None,) + ): + source = name if part is None else name.replace("qkv_mlp", part) + rows = tensor.shape[0] if part is None else {"mlp_in": 384}.get(part, 64) + value = torch.randn(rows, *tensor.shape[1:], generator=generator) * 0.05 + key = f"dit.{source}" + # the action boundary layers stay BF16, as in the released packages + if value.ndim == 2 and ".act" not in source: + q, scale = quantize_fp8_rowwise(value) + state[key], state[f"{key}_scale"] = q.cuda(), scale.cuda() + reference[source] = q.float() * scale[:, None] + else: + state[key] = value.to(torch.bfloat16).cuda() + load_fp8r_checkpoint(model, state) + + block = model.single_blocks[0] + assert isinstance(block.qkv_mlp, Flux3Fp8RowwiseLinear) + assert isinstance(model.early_stream_modulations["txt"].lin, Flux3Fp8RowwiseLinear) + assert isinstance(model.emb_in["act"], torch.nn.Linear) + assert model.dtype == torch.bfloat16 + x = torch.randn(5, 64, generator=generator).cuda().to(torch.bfloat16) + expected = torch.cat( + [ + reference[f"single_blocks.0.{p}.weight"] + for p in ("q_proj", "k_proj", "v_proj", "mlp_in") + ] + ).cuda() + out, _ = block.qkv_mlp(x) + torch.testing.assert_close( + out.float(), x.float() @ expected.T, atol=5e-2, rtol=5e-2 + ) + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="needs the CUDA fused kernel") +def test_fp8r_tp_partition_keeps_gate_and_value_paired(monkeypatch): + """Each rank's fused ``qkv_mlp`` rows must be matching q/k/v/gate/value slices. + + A contiguous split of the fused weight would give one rank all gate rows and + the other all value rows; row-parallel scales cover whole output rows. + """ + from torch import nn + + from sglang.multimodal_gen.runtime.layers.linear import ( + MergedColumnParallelLinear, + RowParallelLinear, + ) + from sglang.multimodal_gen.runtime.models.dits import flux3 + + monkeypatch.setattr(flux3, "get_tp_world_size", lambda: 2) + monkeypatch.setattr(flux3, "get_tp_rank", lambda: 1) + merged = MergedColumnParallelLinear.__new__(MergedColumnParallelLinear) + nn.Module.__init__(merged) + merged.output_sizes = [2, 2, 2, 4, 4] + rows = torch.arange(14.0) + local = flux3._tp_partition(merged, rows, is_scale=False) + assert local.tolist() == [1, 3, 5, 8, 9, 12, 13] + assert flux3._tp_partition(merged, rows, is_scale=True).tolist() == local.tolist() + + row = RowParallelLinear.__new__(RowParallelLinear) + nn.Module.__init__(row) + weight = torch.arange(8.0).reshape(2, 4) + assert flux3._tp_partition(row, weight, is_scale=False).tolist() == [ + [2, 3], + [6, 7], + ] + scale = torch.tensor([1.0, 2.0]) + assert flux3._tp_partition(row, scale, is_scale=True) is scale + + +def test_sp_joint_segments_follow_the_rank_shard(): + """Each rank modulates exactly its slice of every segment; the tail pad is inert. + + Joint blocks modulate per segment, so an SP rank must split its shard where + the global segments end, slice per-token modulation, and give the pad + rows a zero gate. + """ + from sglang.multimodal_gen.runtime.distributed.sp_shard_utils import SpShard + from sglang.multimodal_gen.runtime.models.dits.flux3 import _local_segments + + def mod(tokens: int, value: float): + return tuple(torch.full((1, tokens, 2), value) for _ in range(3)) + + lengths = [3, 6, 1] # text, a stream, a single-token stream: 10 tokens + mods = [mod(1, 1.0), mod(6, 2.0), mod(1, 3.0)] + mods[1][0][:, :, 0] = torch.arange(6.0) # per-token shift + # 3 ranks of 4 tokens: [0, 4), [4, 8), [8, 10) + 2 pad rows. + first = _local_segments(lengths, mods, SpShard(10, 4, 2, 3, 0)) + middle = _local_segments(lengths, mods, SpShard(10, 4, 2, 3, 1)) + last = _local_segments(lengths, mods, SpShard(10, 4, 2, 3, 2)) + assert first[0] == [3, 1] and middle[0] == [4] and last[0] == [1, 1, 2] + assert first[1][1][0][0, :, 0].tolist() == [0.0] + assert middle[1][0][0][0, :, 0].tolist() == [1.0, 2.0, 3.0, 4.0] + assert last[1][0][0][0, :, 0].tolist() == [5.0] + assert last[1][1][2].eq(3.0).all() + assert all(not m.any() for m in last[1][2]) + + +def _dense_neighborhood(q, k, v, kernel, causal): + """NATTEN windows spelled out as a dense mask over ``(B, *grid, heads, D)``.""" + grid = q.shape[1:-2] + coords = torch.cartesian_prod(*[torch.arange(n) for n in grid]).reshape( + -1, len(grid) + ) + allowed = torch.ones(coords.shape[0], coords.shape[0], dtype=torch.bool) + for axis, (n, size, is_causal) in enumerate(zip(grid, kernel, causal)): + qi, ki = coords[:, axis, None], coords[None, :, axis] + if is_causal: + allowed &= (ki <= qi) & (ki > qi - size) + else: + start = (qi - size // 2).clamp(0, n - size) + allowed &= (ki >= start) & (ki < start + size) + + def flat(x): + return x.reshape(x.shape[0], -1, *x.shape[-2:]).transpose(1, 2) + + out = torch.nn.functional.scaled_dot_product_attention( + flat(q).float(), + flat(k).float(), + flat(v).float(), + attn_mask=allowed.to(q.device), + ) + return out.transpose(1, 2).reshape(q.shape) + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="needs CUDA") +@pytest.mark.parametrize( + "grid, kernel, causal", + [ + ((9, 11), (5, 5), (False, False)), + ((6, 7, 9), (5, 5, 5), (True, False, False)), + ((5, 6, 7), (5, 5, 5), (False, False, False)), + ], +) +def test_vae_attention_fallback_uses_natten_windows(grid, kernel, causal): + """Without NATTEN the VAE must attend over NATTEN's windows. + + Non-causal windows shift inward at the borders instead of shrinking; a + causal axis takes the ``k`` latest positions. + """ + from sglang.multimodal_gen.runtime.models.vaes.flux3_neighborhood_attention import ( + natten_available, + neighborhood_attention, + ) + + torch.manual_seed(0) + q, k, v = ( + torch.randn(1, *grid, 2, 64, device="cuda", dtype=torch.bfloat16) + for _ in range(3) + ) + is_causal = list(causal) if any(causal) else None + out = neighborhood_attention(q, k, v, kernel_size=list(kernel), is_causal=is_causal) + reference = _dense_neighborhood(q, k, v, kernel, causal) + torch.testing.assert_close(out.float(), reference, atol=2e-2, rtol=2e-2) + if natten_available(): + from natten.functional import na2d, na3d + + kwargs = {} if is_causal is None else {"is_causal": is_causal} + natten_out = (na2d if len(grid) == 2 else na3d)( + q, k, v, kernel_size=list(kernel), **kwargs + ) + torch.testing.assert_close(out, natten_out, atol=2e-2, rtol=2e-2) + + +def test_fused_qknorm_rope_matches_the_eager_rotation(): + """The fused kernel must use the same [cos | sin] cache layout and interleaved pairs.""" + from sglang.kernels.ops.diffusion import fused_inplace_qknorm_rope + from sglang.multimodal_gen.runtime.models.dits.flux3 import ( + Flux3QKNorm, + apply_rope, + rope_cos_sin, + ) + + torch.manual_seed(0) + norm = Flux3QKNorm(128).cuda().bfloat16() + torch.nn.init.normal_(norm.query_norm.scale, mean=1.0, std=0.1) + torch.nn.init.normal_(norm.key_norm.scale, mean=1.0, std=0.1) + ids = torch.randint(0, 50, (1, 7, 4)).cuda() + rope = rope_cos_sin(ids, (32, 32, 32, 32), 10000) + qkv = torch.randn(1, 7, 3, 2, 128, device="cuda", dtype=torch.bfloat16) + q, k, v = qkv.unbind(2) + expected_q, expected_k = apply_rope(*norm(q, k, v), rope) + fused_inplace_qknorm_rope( + q=q.view(-1, 2, 128), + k=k.view(-1, 2, 128), + q_weight=norm.query_norm.scale, + k_weight=norm.key_norm.scale, + cos_sin_cache=rope.reshape(-1, 128), + positions=torch.arange(7, device="cuda"), + is_neox=False, + eps=1e-6, + ) + torch.testing.assert_close(q, expected_q, atol=3e-2, rtol=2e-2) + torch.testing.assert_close(k, expected_k, atol=3e-2, rtol=2e-2)