diff --git a/docs/cookbook/diffusion/LTX/LTX2.5.mdx b/docs/cookbook/diffusion/LTX/LTX2.5.mdx
new file mode 100644
index 000000000..d1ced3fdd
--- /dev/null
+++ b/docs/cookbook/diffusion/LTX/LTX2.5.mdx
@@ -0,0 +1,317 @@
+---
+title: LTX2.5
+description: Run LTX-2.5 video + audio generation with SGLang Diffusion.
+metatags:
+ description: "Deploy and use the LTX-2.5 video and audio generation model with SGLang Diffusion, including one-stage, two-stage, image-to-video, auto-duration, and diffusion-decoder examples."
+---
+
+import { DiffusionModelTags } from '/src/snippets/diffusion/model-tags.jsx';
+import { LTX25Deployment } from '/src/snippets/diffusion/ltx25-deployment.jsx';
+
+
+
+## 1. Model Introduction
+
+[LTX-2.5](https://huggingface.co/Lightricks/LTX-2.5) is an open world model from
+Lightricks, built for local execution and fine-tuning. Its established use is
+generating synchronized, high-fidelity video and audio from text, image and
+video inputs.
+
+It is a 22B DiT paired with a Gemma-4-12B text encoder, separate video and audio
+VAEs, and a vocoder that outputs 48 kHz stereo. Video and audio are denoised
+jointly in one pass rather than dubbed afterwards, so they stay in sync.
+
+Use **`Lightricks/LTX-2.5-Diffusers`** as `--model-path`.
+
+
+**License notice:** LTX-2.5 is released under the LTX-2.x Community License
+Agreement, not Apache 2.0. The license includes commercial-use restrictions for
+some entities. Review the [official Lightricks license](https://github.com/Lightricks/LTX-2/blob/main/LICENSE.md)
+before production or commercial use; SGLang support does not grant additional
+model usage rights.
+
+
+### 1.1 New in LTX-2.5
+
+Two capabilities have no equivalent in LTX-2 / LTX-2.3:
+
+
+
+ A duration head predicts how long the shot the caption implies should run,
+ and picks the frame count for you. Pass `--auto-duration` instead of
+ `--num-frames`.
+
+
+ A diffusion model replaces the convolutional VAE decoder for the
+ latent-to-pixel step. Enable with `--use-diffusion-decoder`.
+
+
+
+Both are optional and off by default.
+
+### 1.2 Components
+
+| Path | Component | Used by |
+| --- | --- | --- |
+| `transformer/` | Distilled DiT (the default) | always |
+| `transformer_full/` | Full / SFT DiT | `--model-variant dev` |
+| `vae/` | Convolutional video VAE | encode always; decode by default |
+| `diffusion_decoder/` | Diffusion video decoder, decoder-only | `--use-diffusion-decoder` |
+| `latent_upsampler/` | Spatial x2 latent upsampler | `LTX2TwoStagePipeline` |
+| `duration_head/` | Predicts clip length from the caption | `--auto-duration` |
+| `audio_vae/`, `vocoder/`, `connectors/`, `text_encoder/`, `tokenizer/`, `scheduler/` | Shared | always |
+
+Encoding always uses `vae/`, and both decoders consume the same latents, so the
+decoder choice does not change anything upstream of it.
+
+## 2. SGLang-diffusion Installation
+
+```bash
+uv pip install "sglang[diffusion]" --prerelease=allow
+```
+
+For platform-specific setup, see the [SGLang Diffusion installation guide](/docs/sglang-diffusion/installation).
+
+NATTEN is an optional extra, worth installing only if you plan to use the
+[diffusion decoder](#4-6-diffusion-decoder) — see that section for why.
+
+## 3. Model Deployment
+
+### 3.1 Basic Configuration
+
+```bash
+sglang serve \
+ --model-path Lightricks/LTX-2.5-Diffusers \
+ --pipeline-class-name LTX2Pipeline
+```
+
+On a single high-VRAM GPU no extra flags are needed.
+
+**Interactive Command Generator**: pick a target and the features you want; the
+command updates below. Server-side choices (pipeline class, weights variant,
+parallelism) go on `sglang serve`, while per-request choices (auto-duration,
+diffusion decoder, resolution) are listed separately, since they belong on the
+`sglang generate` call or the request body.
+
+
+
+### 3.2 Configuration Tips
+
+Choose the pipeline class based on the quality and latency target:
+
+| Use case | Pipeline class | Notes |
+| --- | --- | --- |
+| One-stage generation | `LTX2Pipeline` | Fastest path. Supports T2V and TI2V, auto-duration and the diffusion decoder. |
+| Two-stage generation | `LTX2TwoStagePipeline` | Half-resolution base stage, x2 latent upsample, then a short refinement. Pass the **final** resolution. |
+
+There is no HQ pipeline class for LTX-2.5, and no `--distilled-lora-path` for
+either weights variant: LTX-2.5 distils the weights themselves rather than
+merging a LoRA per stage, so `--ltx2-two-stage-device-mode` (which governs that
+swap) does not apply either.
+
+Every feature on this page — text-to-video, image conditioning, auto-duration,
+the diffusion decoder, and either weights variant — works with both pipeline
+classes.
+
+Selecting weights:
+
+- `--model-variant dev` serves the full / SFT DiT from `transformer_full/`; the
+ default is the distilled one. See [section 4.5](#4-5-the-dev-transformer).
+
+### 3.3 Multi-GPU presets
+
+| Target | Recommended server flags | Notes |
+| --- | --- | --- |
+| 1 high-VRAM GPU | *(no extra flags)* | 960×544 fits comfortably on an H200. |
+| 1 tight-VRAM GPU | `--quantization fp8` | Halves the DiT and cuts peak memory ~18 GB at unchanged speed. See [section 3.4](#3-4-fp8-quantization). |
+| 1 very tight GPU | `--dit-layerwise-offload` | Cuts peak memory by roughly 10 GB, at about 4x the wall clock. |
+| 2 GPUs, long sequences | `--num-gpus 2 --ulysses-degree 2` | Sequence parallel; the memory/long-sequence tool. |
+| 2 GPUs, large DiT | `--num-gpus 2 --tp-size 2` | Tensor parallel across attention heads. |
+| 2 GPUs, dev weights | `--num-gpus 2 --enable-cfg-parallel` | Splits the guided and unguided branches across GPUs. Measured 1.77x on denoising (15.1s to 8.5s, 960×544 / 57 frames / 30 steps). |
+
+
+**CFG parallelism does not apply on the default (distilled) path.** That DiT
+runs unguided, so there is no negative branch to split across GPUs and
+`--enable-cfg-parallel` buys nothing — the CFG-parallel presets on the
+LTX-2 / LTX-2.3 page do not carry over. It *is* worth using with
+`--model-variant dev`, which runs with guidance.
+
+
+### 3.4 fp8 quantization
+
+`--quantization fp8` quantizes the DiT's linear layers as it loads them, so it
+needs no pre-quantized checkpoint:
+
+```bash
+sglang serve \
+ --model-path Lightricks/LTX-2.5-Diffusers \
+ --pipeline-class-name LTX2Pipeline \
+ --quantization fp8
+```
+
+At 960×544 / 49 frames the transformer loads in 18.11 GB against 35.37 GB for
+bf16, and the run peaks at 53.5 GB against 71.1 GB. Denoising time is
+unchanged: the distilled 8-step path at this size is bound by memory traffic
+rather than matmul throughput, so fp8 buys headroom rather than speed.
+
+Expect a different sample for a given seed. Quantization nudges the denoising
+trajectory and diffusion amplifies that, so the result differs from bf16
+without being worse.
+
+## 4. Model Invocation
+
+### 4.1 Text-to-video with audio
+
+```bash
+sglang generate \
+ --model-path Lightricks/LTX-2.5-Diffusers \
+ --pipeline-class-name LTX2Pipeline \
+ --prompt "A cinematic shot of a red fox walking through a snowy forest at dawn, the camera tracking alongside, snow crunching underfoot." \
+ --save-output
+```
+
+Defaults: 960×544, 121 frames, 24 fps. Video and audio are generated jointly and
+muxed into one MP4.
+
+The default DiT is distilled and runs off a fixed 8-sigma schedule rather than a
+step count, so `--num-inference-steps` and `--guidance-scale` have no effect
+here. Use [`--model-variant dev`](#4-5-the-dev-transformer) when you want
+control over either.
+
+### 4.2 Image-to-video
+
+```bash
+sglang generate \
+ --model-path Lightricks/LTX-2.5-Diffusers \
+ --pipeline-class-name LTX2Pipeline \
+ --image-path ./inputs/start.png \
+ --prompt "The camera pushes forward as the subject turns toward the light." \
+ --save-output
+```
+
+The conditioning image is re-compressed to match the compression the model was
+trained against — CRF 18 for LTX-2.5, where LTX-2 / 2.3 use 33. SGLang picks the
+right one from the checkpoint, so nothing needs to be passed.
+
+### 4.3 Auto-duration
+
+NEW
+
+LTX-2.5 ships a duration head — a small module that reads the encoded caption
+and regresses the natural length of the shot it describes. Use it when the
+prompt implies a duration ("a quick glance" vs "a slow pan across the valley")
+and you would rather not guess a frame count:
+
+```bash
+sglang generate \
+ --model-path Lightricks/LTX-2.5-Diffusers \
+ --pipeline-class-name LTX2Pipeline \
+ --prompt "A red fox walking through a snowy forest at dawn." \
+ --auto-duration \
+ --save-output
+```
+
+The prediction is clamped to `--auto-duration-min-seconds` /
+`--auto-duration-max-seconds` (default 1–20 s) and snapped to the VAE's temporal
+grid, so the result is always a valid frame count. It overrides `--num-frames`.
+
+### 4.4 Two-stage (higher quality)
+
+Stage 1 runs at half the requested resolution, the latents are upsampled 2x, and
+a short sigma tail refines at full resolution. Pass the **final** size:
+
+```bash
+sglang generate \
+ --model-path Lightricks/LTX-2.5-Diffusers \
+ --pipeline-class-name LTX2TwoStagePipeline \
+ --prompt "A cinematic shot of a red fox walking through a snowy forest at dawn." \
+ --height 1088 --width 1920 \
+ --save-output
+```
+
+Resolution must be divisible by 64. Unlike LTX-2.3, no `--distilled-lora-path`
+is needed: the LTX-2.5 transformer is already distilled.
+
+### 4.5 The dev transformer
+
+LTX-2.5 ships two DiTs. `model_index.json` points at the distilled one; the
+full / SFT weights live in `transformer_full/` and are deliberately left out of
+the index. Select them with `--model-variant dev`:
+
+```bash
+sglang generate \
+ --model-path Lightricks/LTX-2.5-Diffusers \
+ --pipeline-class-name LTX2Pipeline \
+ --model-variant dev \
+ --prompt "A cinematic shot of a red fox walking through a snowy forest at dawn." \
+ --num-inference-steps 30 --guidance-scale 3.0 \
+ --save-output
+```
+
+The dev variant is not distilled, so SGLang automatically drops the pinned
+distilled sigma schedule and re-enables the dynamic shifting that `scheduler/`
+turns off for the distilled DiT. Unlike the distilled path it *is* driven by a
+step count and *does* want CFG, so pass `--num-inference-steps` and
+`--guidance-scale` yourself.
+
+Note that `from_pretrained` only fetches what `model_index.json` lists, so a
+partial snapshot download will not include `transformer_full/` (another 38 GB).
+
+### 4.6 Diffusion decoder
+
+NEW
+
+LTX-2.5 adds a diffusion-based video decoder as an alternative to the
+convolutional VAE decoder. Rather than deconvolving the latent it denoises
+pixels conditioned on a context volume built from it, which recovers detail a
+convolutional decoder tends to smooth away:
+
+```bash
+sglang generate \
+ --model-path Lightricks/LTX-2.5-Diffusers \
+ --pipeline-class-name LTX2Pipeline \
+ --prompt "A red fox walking through a snowy forest at dawn." \
+ --use-diffusion-decoder \
+ --save-output
+```
+
+It is a diffusion model in its own right and decodes more slowly than the VAE
+decoder, so it is off by default — matching upstream, where `LTX2Pipeline` also
+decodes with the VAE. The offline `generate` command loads the optional decoder
+automatically when `--use-diffusion-decoder` is present.
+
+For an online server, opt into loading the decoder at startup, then select it per
+request with `use_diffusion_decoder: true`:
+
+```bash
+sglang serve \
+ --model-path Lightricks/LTX-2.5-Diffusers \
+ --pipeline-class-name LTX2Pipeline \
+ --load-diffusion-decoder
+```
+
+This keeps the default server footprint unchanged while still allowing VAE and
+diffusion-decoder requests to share one server. When GPU memory is constrained,
+`--cpu-offload-components diffusion_decoder` keeps the optional decoder on CPU
+between uses.
+
+
+**Install NATTEN for this decoder.** Its stages run 3D neighborhood attention,
+and SGLang uses NATTEN's fused `na3d` kernel for it when the package is present.
+NATTEN is *not* a dependency of `sglang[diffusion]`: without it the decoder
+falls back to a compiled FlexAttention block mask. The two agree to bf16
+rounding, but the fallback is roughly **5x slower** on the decoder's largest
+attention grid, and has to build the mask on top of that.
+
+NATTEN ships prebuilt wheels pinned to a specific torch and CUDA build, so
+install the one matching your environment rather than a bare version — check
+your combination at [natten.org](https://natten.org). For torch 2.11 / CUDA
+13.0, for example:
+
+```bash
+uv pip install natten==0.21.6+torch2110cu130 -f https://whl.natten.org/
+```
+
+Nothing else changes if you skip it: the decoder still produces the same video,
+just slower.
+
diff --git a/docs/docs.json b/docs/docs.json
index 60707fd09..20c5c537e 100644
--- a/docs/docs.json
+++ b/docs/docs.json
@@ -1469,8 +1469,10 @@
},
{
"group": "LTX",
+ "tag": "NEW",
"pages": [
- "cookbook/diffusion/LTX/LTX2 & LTX2.3"
+ "cookbook/diffusion/LTX/LTX2 & LTX2.3",
+ "cookbook/diffusion/LTX/LTX2.5"
]
},
{
diff --git a/docs/docs/sglang-diffusion/compatibility_matrix.mdx b/docs/docs/sglang-diffusion/compatibility_matrix.mdx
index 1fcad4a8d..34393e9d4 100644
--- a/docs/docs/sglang-diffusion/compatibility_matrix.mdx
+++ b/docs/docs/sglang-diffusion/compatibility_matrix.mdx
@@ -145,6 +145,12 @@ Rows are grouped when a family shares the same runtime path or optimization supp
One-stage, two-stage, TI2V, HQ
No dedicated optimization listed
+
+ LTX-2.5
+ Lightricks/LTX-2.5-Diffusers
+ One-stage, two-stage, TI2V, auto-duration, diffusion decode
+ No dedicated optimization listed
+
Cosmos3
nvidia/Cosmos3-Nanonvidia/Cosmos3-Supernvidia/Cosmos3-Super-Text2Imagenvidia/Cosmos3-Super-Image2Video
@@ -605,6 +611,21 @@ Optimization columns are abbreviated to keep the matrix readable:
❌
❌
+
+ LTX-2.5 (one/two-stage/TI2V)
+ Lightricks/LTX-2.5-Diffusers
+ 960×544 (default) 1920×1088 (two-stage)
+ ❌
+ ❌
+ ❌
+ ❌
+ ❌
+ ❌
+ ❌
+ ❌
+ ❌
+ ❌
+
Cosmos3-Nano (T2V / I2V / T2I)
nvidia/Cosmos3-Nano
diff --git a/docs/src/snippets/diffusion/ltx25-deployment.jsx b/docs/src/snippets/diffusion/ltx25-deployment.jsx
new file mode 100644
index 000000000..8f182a48d
--- /dev/null
+++ b/docs/src/snippets/diffusion/ltx25-deployment.jsx
@@ -0,0 +1,182 @@
+export const LTX25Deployment = () => {
+ const options = {
+ hardware: {
+ name: 'hardware',
+ title: 'Deployment Target',
+ items: [
+ { id: 'h200', label: '1x H200', subtitle: 'no extra flags', default: true },
+ { id: 'tight', label: '1 GPU, tight VRAM', subtitle: 'layerwise offload', default: false },
+ { id: 'sp2', label: '2 GPUs', subtitle: 'sequence parallel', default: false },
+ { id: 'tp2', label: '2 GPUs', subtitle: 'tensor parallel', default: false },
+ { id: 'cfg2', label: '2 GPUs', subtitle: 'CFG parallel', default: false },
+ ],
+ },
+ precision: {
+ name: 'precision',
+ title: 'Precision',
+ items: [
+ { id: 'bf16', label: 'bf16', subtitle: 'default', default: true },
+ { id: 'fp8', label: 'fp8', subtitle: 'online, -18 GB', default: false },
+ ],
+ },
+ weights: {
+ name: 'weights',
+ title: 'Weights',
+ items: [
+ { id: 'distilled', label: 'Distilled', subtitle: '8 steps, unguided', default: true },
+ { id: 'dev', label: 'Dev / SFT', subtitle: 'steps + CFG', default: false },
+ ],
+ },
+ pipeline: {
+ name: 'pipeline',
+ title: 'Pipeline',
+ items: [
+ { id: 'one-stage', label: 'One Stage', subtitle: '960x544', default: true },
+ { id: 'two-stage', label: 'Two Stage', subtitle: '1920x1088', default: false },
+ ],
+ },
+ decoder: {
+ name: 'decoder',
+ title: 'Decoder',
+ items: [
+ { id: 'vae', label: 'VAE', subtitle: 'default, fast', default: true },
+ { id: 'diffusion', label: 'Diffusion', subtitle: 'slower, more detail', default: false },
+ ],
+ },
+ duration: {
+ name: 'duration',
+ title: 'Clip Length',
+ items: [
+ { id: 'fixed', label: 'Fixed', subtitle: '--num-frames', default: true },
+ { id: 'auto', label: 'Auto', subtitle: 'duration head', default: false },
+ ],
+ },
+ };
+
+ const REPO_ID = 'Lightricks/LTX-2.5-Diffusers';
+ const PIPELINE_CLASSES = {
+ 'one-stage': 'LTX2Pipeline',
+ 'two-stage': 'LTX2TwoStagePipeline',
+ };
+
+ const [values, setValues] = useState({
+ hardware: 'h200',
+ precision: 'bf16',
+ weights: 'distilled',
+ pipeline: 'one-stage',
+ decoder: 'vae',
+ duration: 'fixed',
+ });
+ const [isDark, setIsDark] = useState(false);
+
+ useEffect(() => {
+ const checkDarkMode = () => {
+ const html = document.documentElement;
+ const isDarkMode = html.classList.contains('dark') ||
+ html.getAttribute('data-theme') === 'dark' ||
+ html.style.colorScheme === 'dark';
+ setIsDark(isDarkMode);
+ };
+ checkDarkMode();
+ const observer = new MutationObserver(checkDarkMode);
+ observer.observe(document.documentElement, { attributes: true, attributeFilter: ['class', 'data-theme', 'style'] });
+ return () => observer.disconnect();
+ }, []);
+
+ const handleRadioChange = (key, id) => {
+ setValues((prev) => ({ ...prev, [key]: id }));
+ };
+
+ const getParallelFlags = () => {
+ const map = {
+ tight: ` \\\n --dit-layerwise-offload`,
+ sp2: ` \\\n --num-gpus 2 \\\n --ulysses-degree 2`,
+ tp2: ` \\\n --num-gpus 2 \\\n --tp-size 2`,
+ cfg2: ` \\\n --num-gpus 2 \\\n --enable-cfg-parallel`,
+ };
+ return map[values.hardware] || '';
+ };
+
+ const generateCommand = () => {
+ let command = `sglang serve \\\n --model-path ${REPO_ID}`;
+ command += ` \\\n --pipeline-class-name ${PIPELINE_CLASSES[values.pipeline]}`;
+ if (values.weights === 'dev') {
+ command += ` \\\n --model-variant dev`;
+ }
+ if (values.precision === 'fp8') {
+ command += ` \\\n --quantization fp8`;
+ }
+ command += getParallelFlags();
+ command += ` \\\n --port 30000`;
+
+ // The distilled DiT runs unguided, so there is no negative branch to split.
+ if (values.hardware === 'cfg2' && values.weights !== 'dev') {
+ command += `\n\n# Note: CFG parallel does nothing on the distilled weights (they run\n# unguided). Pick "Dev / SFT" above, or use sequence/tensor parallel.`;
+ }
+
+ // Per-request flags belong on the generate call, not the server.
+ const requestFlags = [];
+ if (values.pipeline === 'two-stage') {
+ requestFlags.push('--height 1088 --width 1920');
+ }
+ if (values.weights === 'dev') {
+ requestFlags.push('--num-inference-steps 30 --guidance-scale 3.0');
+ }
+ if (values.duration === 'auto') {
+ requestFlags.push('--auto-duration');
+ }
+ if (values.decoder === 'diffusion') {
+ requestFlags.push('--use-diffusion-decoder');
+ }
+ if (requestFlags.length > 0) {
+ command += `\n\n# Per-request flags (pass these to \`sglang generate\`, or as request fields):\n# ${requestFlags.join(' ')}`;
+ }
+ return command;
+ };
+
+ const containerStyle = { maxWidth: '900px', margin: '0 auto', display: 'flex', flexDirection: 'column', gap: '4px' };
+ const cardStyle = { padding: '8px 12px', border: `1px solid ${isDark ? '#374151' : '#e5e7eb'}`, borderLeft: `3px solid ${isDark ? '#E85D4D' : '#D45D44'}`, borderRadius: '4px', display: 'flex', alignItems: 'center', gap: '12px', background: isDark ? '#1f2937' : '#fff' };
+ const titleStyle = { fontSize: '13px', fontWeight: '600', minWidth: '140px', flexShrink: 0, color: isDark ? '#e5e7eb' : 'inherit' };
+ const itemsStyle = { display: 'flex', rowGap: '2px', columnGap: '6px', flexWrap: 'wrap', alignItems: 'center', flex: 1 };
+ const labelBaseStyle = { padding: '4px 10px', border: `1px solid ${isDark ? '#9ca3af' : '#d1d5db'}`, borderRadius: '3px', cursor: 'pointer', display: 'inline-flex', flexDirection: 'column', alignItems: 'center', justifyContent: 'center', fontWeight: '500', fontSize: '13px', transition: 'all 0.2s', userSelect: 'none', minWidth: '45px', textAlign: 'center', flex: 1, background: isDark ? '#374151' : '#fff', color: isDark ? '#e5e7eb' : 'inherit' };
+ const checkedStyle = { background: '#D45D44', color: 'white', borderColor: '#D45D44' };
+ const subtitleStyle = { display: 'block', fontSize: '9px', marginTop: '1px', lineHeight: '1.1', opacity: 0.7 };
+ const commandDisplayStyle = { flex: 1, padding: '12px 16px', background: isDark ? '#111827' : '#f5f5f5', borderRadius: '6px', fontFamily: "'Menlo', 'Monaco', 'Courier New', monospace", fontSize: '12px', lineHeight: '1.5', color: isDark ? '#e5e7eb' : '#374151', whiteSpace: 'pre-wrap', overflowX: 'auto', margin: 0, border: `1px solid ${isDark ? '#374151' : '#e5e7eb'}` };
+
+ return (
+
+ {Object.entries(options).map(([key, option]) => (
+
+
{option.title}
+
+ {option.items.map((item) => {
+ const isChecked = values[option.name] === item.id;
+ return (
+
+ handleRadioChange(key, item.id)}
+ style={{ display: 'none' }}
+ />
+ {item.label}
+ {item.subtitle && (
+
+ {item.subtitle}
+
+ )}
+
+ );
+ })}
+
+
+ ))}
+
+
+
Run this Command:
+
{generateCommand()}
+
+
+ );
+};
diff --git a/python/sglang/multimodal_gen/configs/models/adapter/ltx_2_connector.py b/python/sglang/multimodal_gen/configs/models/adapter/ltx_2_connector.py
index 03e2bcf9b..6caefb45e 100644
--- a/python/sglang/multimodal_gen/configs/models/adapter/ltx_2_connector.py
+++ b/python/sglang/multimodal_gen/configs/models/adapter/ltx_2_connector.py
@@ -5,9 +5,21 @@ from sglang.multimodal_gen.configs.models.adapter.base import (
AdapterConfig,
)
+# Diffusers names the per-modality projections `video_text_proj_in` /
+# `audio_text_proj_in`; SGLang follows ltx-core (`*_aggregate_embed`). Every
+# other connector weight already matches.
+LTX2_CONNECTOR_PARAM_NAMES_MAPPING: dict[str, str] = {
+ r"^video_text_proj_in\.(.*)$": r"video_aggregate_embed.\1",
+ r"^audio_text_proj_in\.(.*)$": r"audio_aggregate_embed.\1",
+}
+
@dataclass
class LTX2ConnectorArchConfig(AdapterArchConfig):
+ param_names_mapping: dict = field(
+ default_factory=lambda: dict(LTX2_CONNECTOR_PARAM_NAMES_MAPPING)
+ )
+
audio_connector_attention_head_dim: int = 128
audio_connector_num_attention_heads: int = 30
audio_connector_num_layers: int = 2
@@ -28,6 +40,30 @@ class LTX2ConnectorArchConfig(AdapterArchConfig):
video_connector_num_layers: int = 2
video_connector_num_learnable_registers: int = 128
+ # `update_model_arch` copies `connectors/config.json` verbatim onto this
+ # object, so declare its names here and derive the SGLang-side fields in
+ # `__post_init__`. LTX-2.0 leaves `per_modality_projections` false and keeps
+ # one shared `text_proj_in`; LTX-2.3 / 2.5 set it.
+ per_modality_projections: bool = False
+ video_hidden_dim: int = 4096
+ audio_hidden_dim: int = 2048
+ video_gated_attn: bool = False
+ audio_gated_attn: bool = False
+
+ def __post_init__(self) -> None:
+ super().__post_init__()
+
+ if self.per_modality_projections:
+ self.feature_extractor_in_features = (
+ self.caption_channels * self.text_proj_in_factor
+ )
+ self.video_feature_extractor_out_features = self.video_hidden_dim
+ self.audio_feature_extractor_out_features = self.audio_hidden_dim
+
+ # Upstream gates these separately; released checkpoints always pair them.
+ if self.video_gated_attn or self.audio_gated_attn:
+ self.connector_apply_gated_attention = True
+
@dataclass
class LTX2ConnectorConfig(AdapterConfig):
diff --git a/python/sglang/multimodal_gen/configs/models/adapter/ltx_2_duration_head.py b/python/sglang/multimodal_gen/configs/models/adapter/ltx_2_duration_head.py
new file mode 100644
index 000000000..f3fd4beb0
--- /dev/null
+++ b/python/sglang/multimodal_gen/configs/models/adapter/ltx_2_duration_head.py
@@ -0,0 +1,30 @@
+# SPDX-License-Identifier: Apache-2.0
+from dataclasses import dataclass, field
+
+from sglang.multimodal_gen.configs.models.adapter.base import (
+ AdapterArchConfig,
+ AdapterConfig,
+)
+
+
+@dataclass
+class LTX2DurationHeadArchConfig(AdapterArchConfig):
+ """LTX-2.5 duration head.
+
+ Field names match `duration_head/config.json` verbatim, so
+ `update_model_arch` populates this directly from the checkpoint.
+ """
+
+ video_cross_attention_dim: int = 4096
+ audio_cross_attention_dim: int = 2048
+ pooler_hidden_dim: int = 256
+ num_queries: int = 1
+ num_pooler_heads: int = 4
+ mlp_hidden_dim: int = 256
+
+
+@dataclass
+class LTX2DurationHeadConfig(AdapterConfig):
+ arch_config: AdapterArchConfig = field(default_factory=LTX2DurationHeadArchConfig)
+
+ prefix: str = "LTX2DurationHead"
diff --git a/python/sglang/multimodal_gen/configs/models/decoders/__init__.py b/python/sglang/multimodal_gen/configs/models/decoders/__init__.py
new file mode 100644
index 000000000..ac7ab008f
--- /dev/null
+++ b/python/sglang/multimodal_gen/configs/models/decoders/__init__.py
@@ -0,0 +1,11 @@
+# SPDX-License-Identifier: Apache-2.0
+
+from sglang.multimodal_gen.configs.models.decoders.ltx_2_5_diffusion_decoder import (
+ LTX25DiffusionDecoderArchConfig,
+ LTX25DiffusionDecoderConfig,
+)
+
+__all__ = [
+ "LTX25DiffusionDecoderArchConfig",
+ "LTX25DiffusionDecoderConfig",
+]
diff --git a/python/sglang/multimodal_gen/configs/models/decoders/ltx_2_5_diffusion_decoder.py b/python/sglang/multimodal_gen/configs/models/decoders/ltx_2_5_diffusion_decoder.py
new file mode 100644
index 000000000..09f7ab923
--- /dev/null
+++ b/python/sglang/multimodal_gen/configs/models/decoders/ltx_2_5_diffusion_decoder.py
@@ -0,0 +1,48 @@
+# SPDX-License-Identifier: Apache-2.0
+from dataclasses import dataclass, field
+
+from sglang.multimodal_gen.configs.models.base import ArchConfig, ModelConfig
+
+
+@dataclass
+class LTX25DiffusionDecoderArchConfig(ArchConfig):
+ """LTX-2.5 diffusion video decoder.
+
+ Field names match `diffusion_decoder/config.json` verbatim so
+ `update_model_arch` populates this straight from the checkpoint.
+ """
+
+ latent_channels: int = 128
+ out_channels: int = 3
+ patch_size: int = 4
+ spatial_compression_ratio: int = 32
+ temporal_compression_ratio: int = 8
+ scaling_factor: float = 1.0
+
+ decoder_head_dim: int = 64
+ decoder_t_emb_dim: int = 384
+ decoder_model_output_type: str = "x0"
+ decoder_num_inference_steps: int = 1
+ decoder_timestep_scale_multiplier: float = 1000.0
+
+ decoder_stage_channels: list[int] = field(
+ default_factory=lambda: [2048, 1024, 512, 512, 256]
+ )
+ decoder_stage_depths: list[int] = field(default_factory=lambda: [4, 6, 4, 2, 8])
+ decoder_stage_kernels: list[list[int]] = field(
+ default_factory=lambda: [[3, 7, 7], [3, 7, 7], [3, 5, 5], [3, 5, 5]]
+ )
+ decoder_stage5_kernel: list[int] = field(default_factory=lambda: [11, 11, 11])
+ decoder_upsample_strides: list[list[int]] = field(
+ default_factory=lambda: [[1, 2, 2], [2, 1, 1], [2, 2, 2], [2, 2, 2]]
+ )
+ decoder_upsample_channel_reductions: list[int] = field(
+ default_factory=lambda: [2, 2, 1, 2]
+ )
+
+
+@dataclass
+class LTX25DiffusionDecoderConfig(ModelConfig):
+ arch_config: LTX25DiffusionDecoderArchConfig = field(
+ default_factory=LTX25DiffusionDecoderArchConfig
+ )
diff --git a/python/sglang/multimodal_gen/configs/models/dits/ltx_2.py b/python/sglang/multimodal_gen/configs/models/dits/ltx_2.py
index 11bc750a9..42ed057cd 100644
--- a/python/sglang/multimodal_gen/configs/models/dits/ltx_2.py
+++ b/python/sglang/multimodal_gen/configs/models/dits/ltx_2.py
@@ -47,62 +47,62 @@ class LTX2AttentionFunction(str, Enum):
DEFAULT = "default"
+# HF checkpoint key -> SGLang module name (SGLang follows upstream naming).
+LTX2_PARAM_NAMES_MAPPING: dict[str, str] = {
+ r"^model\.diffusion_model\.(.*)$": r"\1",
+ r"^proj_in\.(.*)$": r"patchify_proj.\1",
+ r"^time_embed\.(.*)$": r"adaln_single.\1",
+ r"^audio_proj_in\.(.*)$": r"audio_patchify_proj.\1",
+ r"^audio_time_embed\.(.*)$": r"audio_adaln_single.\1",
+ # FeedForward
+ r"(.*)ff\.net\.0\.proj\.(.*)$": r"\1ff.proj_in.\2",
+ r"(.*)ff\.net\.2\.(.*)$": r"\1ff.proj_out.\2",
+ # Attention Norms
+ r"(.*)\.norm_q\.(.*)$": r"\1.q_norm.\2",
+ r"(.*)\.norm_k\.(.*)$": r"\1.k_norm.\2",
+ # Scale Shift Tables (Global)
+ r"^av_cross_attn_video_scale_shift\.(.*)$": r"av_ca_video_scale_shift_adaln_single.\1",
+ r"^av_cross_attn_audio_scale_shift\.(.*)$": r"av_ca_audio_scale_shift_adaln_single.\1",
+ r"^av_cross_attn_video_a2v_gate\.(.*)$": r"av_ca_a2v_gate_adaln_single.\1",
+ r"^av_cross_attn_audio_v2a_gate\.(.*)$": r"av_ca_v2a_gate_adaln_single.\1",
+ # Scale Shift Tables (Block Level)
+ r"(.*)scale_shift_table_a2v_ca_video": r"\1video_a2v_cross_attn_scale_shift_table",
+ r"(.*)scale_shift_table_a2v_ca_audio": r"\1audio_a2v_cross_attn_scale_shift_table",
+}
+
+# Reverse mapping: SGLang module names -> HF checkpoint keys (for saving).
+LTX2_REVERSE_PARAM_NAMES_MAPPING: dict[str, str] = {
+ r"^patchify_proj\.(.*)$": r"proj_in.\1",
+ r"^adaln_single\.(.*)$": r"time_embed.\1",
+ r"^audio_patchify_proj\.(.*)$": r"audio_proj_in.\1",
+ r"^audio_adaln_single\.(.*)$": r"audio_time_embed.\1",
+ # FeedForward
+ r"(.*)ff\.proj_in\.(.*)$": r"\1ff.net.0.proj.\2",
+ r"(.*)ff\.proj_out\.(.*)$": r"\1ff.net.2.\2",
+ # Attention Norms
+ r"(.*)\.q_norm\.(.*)$": r"\1.norm_q.\2",
+ r"(.*)\.k_norm\.(.*)$": r"\1.norm_k.\2",
+ # Scale Shift Tables (Global)
+ r"^av_ca_video_scale_shift_adaln_single\.(.*)$": r"av_cross_attn_video_scale_shift.\1",
+ r"^av_ca_audio_scale_shift_adaln_single\.(.*)$": r"av_cross_attn_audio_scale_shift.\1",
+ r"^av_ca_a2v_gate_adaln_single\.(.*)$": r"av_cross_attn_video_a2v_gate.\1",
+ r"^av_ca_v2a_gate_adaln_single\.(.*)$": r"av_cross_attn_audio_v2a_gate.\1",
+ # Scale Shift Tables (Block Level)
+ r"(.*)video_a2v_cross_attn_scale_shift_table": r"\1scale_shift_table_a2v_ca_video",
+ r"(.*)audio_a2v_cross_attn_scale_shift_table": r"\1scale_shift_table_a2v_ca_audio",
+}
+
+
@dataclass
class LTX2ArchConfig(DiTArchConfig):
"""Architecture configuration for LTX-2 Video Transformer."""
param_names_mapping: dict = field(
- default_factory=lambda: {
- # Parameter name mappings from HuggingFace checkpoint keys to SGLang module names.
- # We use upstream variable names (patchify_proj, adaln_single) but HF uses different keys.
- #
- # HF key -> SGLang key (upstream naming)
- r"^model\.diffusion_model\.(.*)$": r"\1",
- r"^proj_in\.(.*)$": r"patchify_proj.\1",
- r"^time_embed\.(.*)$": r"adaln_single.\1",
- r"^audio_proj_in\.(.*)$": r"audio_patchify_proj.\1",
- r"^audio_time_embed\.(.*)$": r"audio_adaln_single.\1",
- # FeedForward
- r"(.*)ff\.net\.0\.proj\.(.*)$": r"\1ff.proj_in.\2",
- r"(.*)ff\.net\.2\.(.*)$": r"\1ff.proj_out.\2",
- # Attention Norms
- r"(.*)\.norm_q\.(.*)$": r"\1.q_norm.\2",
- r"(.*)\.norm_k\.(.*)$": r"\1.k_norm.\2",
- # Scale Shift Tables (Global)
- r"^av_cross_attn_video_scale_shift\.(.*)$": r"av_ca_video_scale_shift_adaln_single.\1",
- r"^av_cross_attn_audio_scale_shift\.(.*)$": r"av_ca_audio_scale_shift_adaln_single.\1",
- r"^av_cross_attn_video_a2v_gate\.(.*)$": r"av_ca_a2v_gate_adaln_single.\1",
- r"^av_cross_attn_audio_v2a_gate\.(.*)$": r"av_ca_v2a_gate_adaln_single.\1",
- # Scale Shift Tables (Block Level)
- # HF: scale_shift_table_a2v_ca_video -> SGLang: video_a2v_cross_attn_scale_shift_table
- r"(.*)scale_shift_table_a2v_ca_video": r"\1video_a2v_cross_attn_scale_shift_table",
- r"(.*)scale_shift_table_a2v_ca_audio": r"\1audio_a2v_cross_attn_scale_shift_table",
- }
+ default_factory=lambda: dict(LTX2_PARAM_NAMES_MAPPING)
)
reverse_param_names_mapping: dict = field(
- default_factory=lambda: {
- # Reverse mapping: SGLang module names -> HF checkpoint keys (for saving).
- r"^patchify_proj\.(.*)$": r"proj_in.\1",
- r"^adaln_single\.(.*)$": r"time_embed.\1",
- r"^audio_patchify_proj\.(.*)$": r"audio_proj_in.\1",
- r"^audio_adaln_single\.(.*)$": r"audio_time_embed.\1",
- # FeedForward
- r"(.*)ff\.proj_in\.(.*)$": r"\1ff.net.0.proj.\2",
- r"(.*)ff\.proj_out\.(.*)$": r"\1ff.net.2.\2",
- # Attention Norms
- r"(.*)\.q_norm\.(.*)$": r"\1.norm_q.\2",
- r"(.*)\.k_norm\.(.*)$": r"\1.norm_k.\2",
- # Scale Shift Tables (Global)
- r"^av_ca_video_scale_shift_adaln_single\.(.*)$": r"av_cross_attn_video_scale_shift.\1",
- r"^av_ca_audio_scale_shift_adaln_single\.(.*)$": r"av_cross_attn_audio_scale_shift.\1",
- r"^av_ca_a2v_gate_adaln_single\.(.*)$": r"av_cross_attn_video_a2v_gate.\1",
- r"^av_ca_v2a_gate_adaln_single\.(.*)$": r"av_cross_attn_audio_v2a_gate.\1",
- # Scale Shift Tables (Block Level)
- # SGLang: video_a2v_cross_attn_scale_shift_table -> HF: scale_shift_table_a2v_ca_video
- r"(.*)video_a2v_cross_attn_scale_shift_table": r"\1scale_shift_table_a2v_ca_video",
- r"(.*)audio_a2v_cross_attn_scale_shift_table": r"\1scale_shift_table_a2v_ca_audio",
- }
+ default_factory=lambda: dict(LTX2_REVERSE_PARAM_NAMES_MAPPING)
)
lora_param_names_mapping: dict = field(
@@ -123,6 +123,13 @@ class LTX2ArchConfig(DiTArchConfig):
cross_attention_adaln: bool = False
caption_proj_before_connector: bool = False
+ # LTX-2.5 drops the video feed-forward bias but keeps the audio one.
+ # `use_keyframes_abs_pos_embedding` only allocates the parameter so the
+ # checkpoint round-trips; the forward does not consume it.
+ ff_bias: bool = True
+ audio_ff_bias: bool = True
+ use_keyframes_abs_pos_embedding: bool = False
+
# Video parameters
num_attention_heads: int = 32
attention_head_dim: int = 128
diff --git a/python/sglang/multimodal_gen/configs/models/dits/ltx_2_5.py b/python/sglang/multimodal_gen/configs/models/dits/ltx_2_5.py
new file mode 100644
index 000000000..60d5dd7d0
--- /dev/null
+++ b/python/sglang/multimodal_gen/configs/models/dits/ltx_2_5.py
@@ -0,0 +1,74 @@
+# SPDX-License-Identifier: Apache-2.0
+from dataclasses import dataclass, field
+
+from sglang.multimodal_gen.configs.models.dits.ltx_2 import (
+ LTX2_PARAM_NAMES_MAPPING,
+ LTX2_REVERSE_PARAM_NAMES_MAPPING,
+ LTX2ArchConfig,
+ LTX2Config,
+ LTX2RopeType,
+)
+
+# LTX-2.5 renames the prompt adaLN modules and adds the keyframe position
+# embedding; the shared LTX-2 mapping covers everything else.
+LTX25_EXTRA_PARAM_NAMES_MAPPING: dict[str, str] = {
+ r"^prompt_adaln\.(.*)$": r"prompt_adaln_single.\1",
+ r"^audio_prompt_adaln\.(.*)$": r"audio_prompt_adaln_single.\1",
+}
+
+LTX25_EXTRA_REVERSE_PARAM_NAMES_MAPPING: dict[str, str] = {
+ r"^prompt_adaln_single\.(.*)$": r"prompt_adaln.\1",
+ r"^audio_prompt_adaln_single\.(.*)$": r"audio_prompt_adaln.\1",
+}
+
+
+@dataclass
+class LTX25ArchConfig(LTX2ArchConfig):
+ """LTX-2.5 DiT architecture config.
+
+ LTX-2.5 reuses the LTX-2.3 audio-video transformer: gated attention,
+ cross-attention adaLN modulation, split RoPE in double precision, and
+ per-modality caption projections that live in the connector rather than the
+ DiT. On top of that it drops the video feed-forward bias and carries a
+ keyframe absolute-position embedding.
+ """
+
+ param_names_mapping: dict = field(
+ default_factory=lambda: {
+ **LTX2_PARAM_NAMES_MAPPING,
+ **LTX25_EXTRA_PARAM_NAMES_MAPPING,
+ }
+ )
+ reverse_param_names_mapping: dict = field(
+ default_factory=lambda: {
+ **LTX2_REVERSE_PARAM_NAMES_MAPPING,
+ **LTX25_EXTRA_REVERSE_PARAM_NAMES_MAPPING,
+ }
+ )
+
+ # LTX-2.3 audio-video base (`gated_attn` / `cross_attn_mod` /
+ # `use_prompt_embeddings: false` in transformer/config.json).
+ apply_gated_attention: bool = True
+ cross_attention_adaln: bool = True
+ caption_proj_before_connector: bool = True
+ rope_type: LTX2RopeType = LTX2RopeType.SPLIT
+ double_precision_rope: bool = True
+
+ # LTX-2.5 specific.
+ ff_bias: bool = False
+ audio_ff_bias: bool = True
+ use_keyframes_abs_pos_embedding: bool = True
+
+ # Mirrored here because these also appear in transformer/config.json.
+ connector_num_attention_heads: int = 32
+ connector_num_layers: int = 8
+ audio_connector_attention_head_dim: int = 64
+ audio_connector_num_attention_heads: int = 32
+ audio_connector_num_layers: int = 8
+
+
+@dataclass
+class LTX25Config(LTX2Config):
+ arch_config: LTX25ArchConfig = field(default_factory=LTX25ArchConfig)
+
+ prefix: str = "ltx2_5"
diff --git a/python/sglang/multimodal_gen/configs/models/encoders/__init__.py b/python/sglang/multimodal_gen/configs/models/encoders/__init__.py
index 58c00f45c..2fbc9ccf3 100644
--- a/python/sglang/multimodal_gen/configs/models/encoders/__init__.py
+++ b/python/sglang/multimodal_gen/configs/models/encoders/__init__.py
@@ -17,6 +17,9 @@ from sglang.multimodal_gen.configs.models.encoders.flux_2 import (
)
from sglang.multimodal_gen.configs.models.encoders.gemma2 import Gemma2Config
from sglang.multimodal_gen.configs.models.encoders.gemma_3 import Gemma3Config
+from sglang.multimodal_gen.configs.models.encoders.gemma_4_unified import (
+ Gemma4UnifiedConfig,
+)
from sglang.multimodal_gen.configs.models.encoders.ideogram import (
Ideogram4TextEncoderConfig,
)
@@ -47,5 +50,6 @@ __all__ = [
"T5Config",
"Gemma2Config",
"Gemma3Config",
+ "Gemma4UnifiedConfig",
"Ideogram4TextEncoderConfig",
]
diff --git a/python/sglang/multimodal_gen/configs/models/encoders/gemma_4_unified.py b/python/sglang/multimodal_gen/configs/models/encoders/gemma_4_unified.py
new file mode 100644
index 000000000..d281a3d4c
--- /dev/null
+++ b/python/sglang/multimodal_gen/configs/models/encoders/gemma_4_unified.py
@@ -0,0 +1,70 @@
+# SPDX-License-Identifier: Apache-2.0
+
+from dataclasses import dataclass, field
+
+from sglang.multimodal_gen.configs.models.encoders.base import (
+ TextEncoderArchConfig,
+ TextEncoderConfig,
+)
+from sglang.multimodal_gen.configs.models.fsdp import (
+ is_embed_tokens,
+ is_final_norm,
+ is_layer,
+)
+
+
+@dataclass
+class Gemma4UnifiedArchConfig(TextEncoderArchConfig):
+ """Gemma-4-Unified text encoder used by LTX-2.5.
+
+ Like `Gemma3ArchConfig`, the actual module is instantiated by transformers
+ from the repo's `text_encoder/config.json`
+ (`Gemma4UnifiedForConditionalGeneration`); this config carries tokenization
+ and sharding metadata.
+
+ LTX-2.5 consumes all 48 hidden layers plus the embedding output, which is why
+ the connector's `text_proj_in_factor` is 49 and `caption_channels` is 3840.
+
+ No `param_names_mapping` here: this encoder currently loads through
+ `TextEncoderLoader.load_native`, i.e. `transformers.from_pretrained`, which
+ never consults SGLang's mapping. If a customized (FSDP/TP) implementation is
+ added the way LTX-2/2.3 have `FSDPGemma3ForConditionalGeneration`, it will
+ need one, because 10 keys drift between the checkpoint and the installed
+ transformers:
+
+ model.vision_embedder.* -> model.embed_vision.*
+ model.embed_vision.embedding_projection.* ->
+ model.embed_vision.multimodal_embedder.embedding_projection.*
+
+ plus a tied `lm_head.weight`. All of it is on the vision path, which
+ text-to-video never runs; `from_pretrained` tolerates the drift today.
+ """
+
+ hidden_size: int = 3840
+ num_hidden_layers: int = 48
+ rms_norm_eps: float = 1e-6
+ rope_theta: float = 10000.0
+ max_position_embeddings: int = 262144
+ hidden_state_skip_layer: int = 2
+ text_len: int = 1024
+
+ stacked_params_mapping: list[tuple[str, str, str]] = field(
+ default_factory=lambda: [
+ # (param_name, shard_name, shard_id)
+ (".qkv_proj", ".q_proj", "q"),
+ (".qkv_proj", ".k_proj", "k"),
+ (".qkv_proj", ".v_proj", "v"),
+ (".gate_up_proj", ".gate_proj", "0"), # type: ignore
+ (".gate_up_proj", ".up_proj", "1"), # type: ignore
+ ]
+ )
+ _fsdp_shard_conditions: list = field(
+ default_factory=lambda: [is_layer, is_embed_tokens, is_final_norm]
+ )
+
+
+@dataclass
+class Gemma4UnifiedConfig(TextEncoderConfig):
+ arch_config: TextEncoderArchConfig = field(default_factory=Gemma4UnifiedArchConfig)
+
+ prefix: str = "gemma_4_unified"
diff --git a/python/sglang/multimodal_gen/configs/models/vaes/ltx_2_5_video.py b/python/sglang/multimodal_gen/configs/models/vaes/ltx_2_5_video.py
new file mode 100644
index 000000000..36a59392a
--- /dev/null
+++ b/python/sglang/multimodal_gen/configs/models/vaes/ltx_2_5_video.py
@@ -0,0 +1,55 @@
+# SPDX-License-Identifier: Apache-2.0
+from dataclasses import dataclass, field
+from typing import List
+
+from sglang.multimodal_gen.configs.models.vaes.ltx_video import (
+ LTXVideoVAEArchConfig,
+ LTXVideoVAEConfig,
+)
+
+
+@dataclass
+class LTX25VideoVAEArchConfig(LTXVideoVAEArchConfig):
+ """LTX-2.5 video VAE.
+
+ The encoder is unchanged from LTX-2. The decoder gains a fourth up block and
+ no longer upsamples every stage in all three dimensions -- `upsample_type`
+ makes the last two stages temporal-only and spatial-only respectively.
+ """
+
+ block_out_channels: List[int] = field(
+ default_factory=lambda: [256, 512, 1024, 1024]
+ )
+ layers_per_block: List[int] = field(default_factory=lambda: [4, 6, 4, 2, 2])
+
+ decoder_block_out_channels: List[int] = field(
+ default_factory=lambda: [256, 512, 512, 1024]
+ )
+ decoder_spatio_temporal_scaling: List[bool] = field(
+ default_factory=lambda: [True, True, True, True]
+ )
+ decoder_layers_per_block: List[int] = field(default_factory=lambda: [4, 6, 4, 2, 2])
+ decoder_inject_noise: List[bool] = field(
+ default_factory=lambda: [False, False, False, False, False]
+ )
+ decoder_spatial_padding_mode: str = "zeros"
+
+ upsample_residual: List[bool] = field(
+ default_factory=lambda: [False, False, False, False]
+ )
+ upsample_factor: List[int] = field(default_factory=lambda: [2, 2, 1, 2])
+ upsample_type: List[str] | None = field(
+ default_factory=lambda: [
+ "spatiotemporal",
+ "spatiotemporal",
+ "temporal",
+ "spatial",
+ ]
+ )
+
+
+@dataclass
+class LTX25VideoVAEConfig(LTXVideoVAEConfig):
+ arch_config: LTX25VideoVAEArchConfig = field(
+ default_factory=LTX25VideoVAEArchConfig
+ )
diff --git a/python/sglang/multimodal_gen/configs/models/vaes/ltx_video.py b/python/sglang/multimodal_gen/configs/models/vaes/ltx_video.py
index 70a59226b..8e3661728 100644
--- a/python/sglang/multimodal_gen/configs/models/vaes/ltx_video.py
+++ b/python/sglang/multimodal_gen/configs/models/vaes/ltx_video.py
@@ -51,6 +51,15 @@ class LTXVideoVAEArchConfig(VAEArchConfig):
decoder_layers_per_block: List[int] = field(default_factory=lambda: [5, 5, 5, 5])
decoder_causal: bool = False
decoder_spatial_padding_mode: str = "reflect"
+ decoder_inject_noise: List[bool] = field(
+ default_factory=lambda: [False, False, False, False]
+ )
+ upsample_residual: List[bool] = field(default_factory=lambda: [True, True, True])
+ upsample_factor: List[int] = field(default_factory=lambda: [2, 2, 2])
+ # Per-decoder-stage upsampling axis: "spatial", "temporal" or
+ # "spatiotemporal". `None` keeps every stage spatiotemporal (LTX-2).
+ upsample_type: List[str] | None = None
+ timestep_conditioning: bool = False
# Native LTX variant metadata.
ltx_variant: str = "ltx_2"
diff --git a/python/sglang/multimodal_gen/configs/models/vocoder/ltx_vocoder.py b/python/sglang/multimodal_gen/configs/models/vocoder/ltx_vocoder.py
index 8feea43cc..addefbca8 100644
--- a/python/sglang/multimodal_gen/configs/models/vocoder/ltx_vocoder.py
+++ b/python/sglang/multimodal_gen/configs/models/vocoder/ltx_vocoder.py
@@ -1,15 +1,33 @@
# SPDX-License-Identifier: Apache-2.0
from dataclasses import dataclass, field
-from typing import List
+from typing import Any, List
from sglang.multimodal_gen.configs.models.vocoder.base import (
VocoderArchConfig,
VocoderConfig,
)
+# `LTX2VocoderWithBWE` stores both stacks with diffusers module names; SGLang
+# follows ltx-core naming. The `vocoder.` / `bwe_generator.` prefixes match.
+LTX_VOCODER_PARAM_NAMES_MAPPING: dict[str, str] = {
+ r"^(vocoder|bwe_generator)\.conv_in\.(.*)$": r"\1.conv_pre.\2",
+ r"^(vocoder|bwe_generator)\.conv_out\.(.*)$": r"\1.conv_post.\2",
+ r"^(vocoder|bwe_generator)\.act_out\.(.*)$": r"\1.act_post.\2",
+ r"^(vocoder|bwe_generator)\.upsamplers\.(.*)$": r"\1.ups.\2",
+ r"^(vocoder|bwe_generator)\.resnets\.(.*)$": r"\1.resblocks.\2",
+ # DownSample1d holds its kernel on a LowPassFilter1d submodule; UpSample1d
+ # registers it directly. Must run after the renames above, so the rules are
+ # evaluated in order rather than first-match-wins.
+ r"^(vocoder|bwe_generator)\.(.*)downsample\.filter$": r"\1.\2downsample.lowpass.filter",
+}
+
@dataclass
class LTXVocoderArchConfig(VocoderArchConfig):
+ param_names_mapping: dict = field(
+ default_factory=lambda: dict(LTX_VOCODER_PARAM_NAMES_MAPPING)
+ )
+
# Architecture params
in_channels: int = 128
hidden_channels: int = 1024
@@ -23,6 +41,74 @@ class LTXVocoderArchConfig(VocoderArchConfig):
leaky_relu_negative_slope: float = 0.1
sample_rate: int = 24000
+ # --- LTX-2.5 `LTX2VocoderWithBWE` fields -------------------------------
+ # The base stack synthesises at `input_sampling_rate`, a mel STFT
+ # re-analyses it, and the BWE stack resynthesises at `output_sampling_rate`.
+ act_fn: str = "snake"
+ final_act_fn: str | None = None
+ final_bias: bool = True
+ antialias: bool = False
+ input_sampling_rate: int = 16000
+ output_sampling_rate: int = 24000
+ # Mel analysis feeding the BWE stack.
+ filter_length: int = 512
+ window_length: int = 512
+ hop_length: int = 80
+ num_mel_channels: int = 64
+ # `bwe_upsample_factors` being non-empty is what marks a BWE checkpoint.
+ bwe_act_fn: str = "snake"
+ bwe_final_act_fn: str | None = None
+ bwe_final_bias: bool = True
+ bwe_hidden_channels: int = 512
+ bwe_in_channels: int = 128
+ bwe_out_channels: int = 2
+ bwe_upsample_factors: List[int] = field(default_factory=list)
+ bwe_upsample_kernel_sizes: List[int] = field(default_factory=list)
+ bwe_resnet_kernel_sizes: List[int] = field(default_factory=list)
+ bwe_resnet_dilations: List[List[int]] = field(default_factory=list)
+
+ # `LTX2Vocoder` takes its BWE branch when this carries a "bwe" entry.
+ vocoder: dict[str, Any] | None = None
+
+ def __post_init__(self) -> None:
+ if self.bwe_upsample_factors and self.vocoder is None:
+ self.vocoder = self._build_nested_bwe_config()
+
+ def _build_nested_bwe_config(self) -> dict[str, Any]:
+ """Translate the flat diffusers fields into the nested ltx-core shape."""
+ return {
+ "vocoder": {
+ "resblock": "AMP1",
+ "activation": self.act_fn,
+ "resblock_kernel_sizes": self.resnet_kernel_sizes,
+ "resblock_dilation_sizes": self.resnet_dilations,
+ "upsample_rates": self.upsample_factors,
+ "upsample_kernel_sizes": self.upsample_kernel_sizes,
+ "upsample_initial_channel": self.hidden_channels,
+ "apply_final_activation": self.final_act_fn is not None,
+ "use_tanh_at_final": self.final_act_fn == "tanh",
+ "use_bias_at_final": self.final_bias,
+ },
+ "bwe": {
+ "resblock": "AMP1",
+ "activation": self.bwe_act_fn,
+ "resblock_kernel_sizes": self.bwe_resnet_kernel_sizes,
+ "resblock_dilation_sizes": self.bwe_resnet_dilations,
+ "upsample_rates": self.bwe_upsample_factors,
+ "upsample_kernel_sizes": self.bwe_upsample_kernel_sizes,
+ "upsample_initial_channel": self.bwe_hidden_channels,
+ "apply_final_activation": self.bwe_final_act_fn is not None,
+ "use_tanh_at_final": self.bwe_final_act_fn == "tanh",
+ "use_bias_at_final": self.bwe_final_bias,
+ "input_sampling_rate": self.input_sampling_rate,
+ "output_sampling_rate": self.output_sampling_rate,
+ "n_fft": self.filter_length,
+ "win_size": self.window_length,
+ "hop_length": self.hop_length,
+ "num_mels": self.num_mel_channels,
+ },
+ }
+
@dataclass
class LTXVocoderConfig(VocoderConfig):
diff --git a/python/sglang/multimodal_gen/configs/pipeline_configs/base.py b/python/sglang/multimodal_gen/configs/pipeline_configs/base.py
index 9ad6c558e..8931d3d08 100644
--- a/python/sglang/multimodal_gen/configs/pipeline_configs/base.py
+++ b/python/sglang/multimodal_gen/configs/pipeline_configs/base.py
@@ -226,6 +226,9 @@ class PipelineConfig:
vae_precision: str = "fp32"
vae_decode_precision: str | None = None
vae_tiling: bool = True
+ # Bounds the attention grid the diffusion decoder's stages see, which is
+ # what makes a full-length decode tractable.
+ diffusion_decoder_tiling: bool = True
vae_slicing: bool = False
vae_sp: bool = True
@@ -845,6 +848,13 @@ class PipelineConfig:
default=PipelineConfig.vae_tiling,
help="Enable VAE tiling",
)
+ parser.add_argument(
+ f"--{prefix_with_dot}diffusion-decoder-tiling",
+ action=StoreBoolean,
+ dest=f"{prefix_with_dot.replace('-', '_')}diffusion_decoder_tiling",
+ default=PipelineConfig.diffusion_decoder_tiling,
+ help="Enable tiling for the LTX-2.5 diffusion decoder",
+ )
parser.add_argument(
f"--{prefix_with_dot}vae-slicing",
action=StoreBoolean,
diff --git a/python/sglang/multimodal_gen/configs/pipeline_configs/ltx_2.py b/python/sglang/multimodal_gen/configs/pipeline_configs/ltx_2.py
index dd2caa7e6..ccc4d2715 100644
--- a/python/sglang/multimodal_gen/configs/pipeline_configs/ltx_2.py
+++ b/python/sglang/multimodal_gen/configs/pipeline_configs/ltx_2.py
@@ -178,6 +178,10 @@ class LTX2PipelineConfig(PipelineConfig):
generator_device: str = "cpu"
dit_config: LTX2Config = field(default_factory=LTX2Config)
+ # Distilled checkpoints are trained against one fixed sigma schedule rather
+ # than a step count. When set, it replaces the derived schedule.
+ default_sigmas: tuple[float, ...] | None = None
+
# Model architecture
in_channels: int = 128
out_channels: int = 128
@@ -309,6 +313,9 @@ class LTX2PipelineConfig(PipelineConfig):
self.patch_size,
)
latents = latents.permute(0, 2, 4, 6, 1, 3, 5, 7).flatten(4, 7).flatten(1, 3)
+ # Deliberately left non-contiguous: both flattens are views, so this
+ # keeps the permuted strides. Normalising here would change which GEMM
+ # kernel runs and move bf16 output. The fp8 path makes its own copy.
return latents
def _infer_video_latent_frames_and_tokens_per_frame(
diff --git a/python/sglang/multimodal_gen/configs/pipeline_configs/ltx_2_5.py b/python/sglang/multimodal_gen/configs/pipeline_configs/ltx_2_5.py
new file mode 100644
index 000000000..646c055dc
--- /dev/null
+++ b/python/sglang/multimodal_gen/configs/pipeline_configs/ltx_2_5.py
@@ -0,0 +1,57 @@
+# SPDX-License-Identifier: Apache-2.0
+import dataclasses
+from dataclasses import field
+
+from sglang.multimodal_gen.configs.models.dits.ltx_2_5 import LTX25Config
+from sglang.multimodal_gen.configs.models.encoders import (
+ EncoderConfig,
+ Gemma4UnifiedConfig,
+)
+from sglang.multimodal_gen.configs.models.vaes.ltx_2_5_video import LTX25VideoVAEConfig
+from sglang.multimodal_gen.configs.pipeline_configs.base import ModelTaskType
+from sglang.multimodal_gen.configs.pipeline_configs.ltx_2 import LTX2PipelineConfig
+
+# Explicit sigma schedule the distilled LTX-2.5 DiT was trained against. Upstream
+# exposes it as `diffusers.pipelines.ltx2.utils.DISTILLED_SIGMA_VALUES`.
+LTX25_DISTILLED_SIGMA_VALUES: tuple[float, ...] = (
+ 1.0,
+ 0.99375,
+ 0.9875,
+ 0.98125,
+ 0.975,
+ 0.909375,
+ 0.725,
+ 0.421875,
+)
+
+
+@dataclasses.dataclass
+class LTX25PipelineConfig(LTX2PipelineConfig):
+ """Pipeline configuration for LTX-2.5.
+
+ LTX-2.5 reuses the LTX-2 pipeline class (`model_index.json` still declares
+ `LTX2Pipeline`) and the LTX-2 *sigma* path -- upstream builds
+ `np.linspace(1.0, 1/steps, steps)` and lets the scheduler's
+ `use_dynamic_shifting: false` turn the shift into a no-op. So this must stay
+ an `ltx_2` variant; do not mark it as an LTX-2.3 native variant.
+
+ What differs from LTX-2 is the component geometry (DiT / VAE / connectors /
+ text encoder), and that the shipped DiT is distilled, hence the pinned
+ `default_sigmas`.
+ """
+
+ # One checkpoint drives both T2V and image-conditioned generation, so this
+ # must stay TI2V -- T2V rejects `--image-path` outright.
+ task_type: ModelTaskType = ModelTaskType.TI2V
+ native_only_components = ("diffusion_decoder",)
+
+ dit_config: LTX25Config = field(default_factory=LTX25Config)
+ vae_config: LTX25VideoVAEConfig = field(default_factory=LTX25VideoVAEConfig)
+
+ text_encoder_configs: tuple[EncoderConfig, ...] = field(
+ default_factory=lambda: (Gemma4UnifiedConfig(),)
+ )
+
+ default_sigmas: tuple[float, ...] | None = field(
+ default_factory=lambda: LTX25_DISTILLED_SIGMA_VALUES
+ )
diff --git a/python/sglang/multimodal_gen/configs/sample/ltx_2_5.py b/python/sglang/multimodal_gen/configs/sample/ltx_2_5.py
new file mode 100644
index 000000000..8d7bad1b5
--- /dev/null
+++ b/python/sglang/multimodal_gen/configs/sample/ltx_2_5.py
@@ -0,0 +1,33 @@
+import dataclasses
+
+from sglang.multimodal_gen.configs.sample.ltx_2 import LTX2SamplingParams
+
+
+@dataclasses.dataclass
+class LTX25SamplingParams(LTX2SamplingParams):
+ """Sampling defaults for the LTX-2.5 distilled transformer.
+
+ `model_index.json` points at the distilled DiT, which runs **unguided** off an
+ explicit sigma schedule (see `LTX25PipelineConfig.default_sigmas`) rather than
+ a step count. `guidance_scale=1.0` disables CFG; STG and modality guidance
+ stay off. Feeding it a generic linear schedule instead costs quality.
+
+ Reference: the "Quick start — distilled, convolutional decode" recipe in the
+ `Lightricks/LTX-2.5-Diffusers` model card.
+ """
+
+ seed: int = 42
+ generator_device: str = "cuda"
+
+ height: int = 544
+ width: int = 960
+ num_frames: int = 121
+ fps: int = 24
+
+ guidance_scale: float = 1.0
+
+ # `auto_duration` on the base class has the duration head predict this
+ # instead, overriding `num_frames`.
+ # The schedule is pinned by the pipeline config; this only keeps the
+ # reported step count honest.
+ num_inference_steps: int = 8
diff --git a/python/sglang/multimodal_gen/configs/sample/sampling_params.py b/python/sglang/multimodal_gen/configs/sample/sampling_params.py
index 33ab21614..32f713450 100644
--- a/python/sglang/multimodal_gen/configs/sample/sampling_params.py
+++ b/python/sglang/multimodal_gen/configs/sample/sampling_params.py
@@ -180,6 +180,16 @@ class SamplingParams:
width: int | None = None
fps: int = 24
+ # LTX-2.5 duration head. Ignored by other models, so the flags stay
+ # universally accepted.
+ # Decode with the diffusion decoder instead of the VAE one. Ignored by
+ # models that ship no such decoder.
+ use_diffusion_decoder: bool = False
+
+ auto_duration: bool = False
+ auto_duration_min_seconds: float = 1.0
+ auto_duration_max_seconds: float = 20.0
+
# Resolution validation
supported_resolutions: list[tuple[int, int]] | None = field(
default=None, metadata={"batch_sig_exclude": True}
@@ -885,6 +895,11 @@ class SamplingParams:
return parser.add_argument(*name_or_flags, **kwargs)
add_argument("--data-type", type=str, nargs="+")
+ # Predict the shot length from the caption, overriding `--num-frames`.
+ add_argument("--use-diffusion-decoder", action="store_true")
+ add_argument("--auto-duration", action="store_true")
+ add_argument("--auto-duration-min-seconds", type=float)
+ add_argument("--auto-duration-max-seconds", type=float)
add_argument(
"--num-frames-round-down",
action="store_true",
diff --git a/python/sglang/multimodal_gen/csrc/attn/vmoba_attn/tests/test_vmoba_attn.py b/python/sglang/multimodal_gen/csrc/attn/vmoba_attn/tests/test_vmoba_attn.py
index 9350f3c9e..55b969f39 100644
--- a/python/sglang/multimodal_gen/csrc/attn/vmoba_attn/tests/test_vmoba_attn.py
+++ b/python/sglang/multimodal_gen/csrc/attn/vmoba_attn/tests/test_vmoba_attn.py
@@ -6,6 +6,7 @@ import pytest
import torch
from sglang.multimodal_gen.csrc.attn.vmoba_attn.vmoba import moba_attn_varlen
+
def generate_test_data(
batch_size, total_seqlen, num_heads, head_dim, dtype, device="cuda"
):
diff --git a/python/sglang/multimodal_gen/registry.py b/python/sglang/multimodal_gen/registry.py
index 8e9f360da..72fcc21ed 100644
--- a/python/sglang/multimodal_gen/registry.py
+++ b/python/sglang/multimodal_gen/registry.py
@@ -81,6 +81,7 @@ from sglang.multimodal_gen.configs.pipeline_configs.ltx_2 import (
LTX2PipelineConfig,
LTX23PipelineConfig,
)
+from sglang.multimodal_gen.configs.pipeline_configs.ltx_2_5 import LTX25PipelineConfig
from sglang.multimodal_gen.configs.pipeline_configs.mova import (
MOVA360PConfig,
MOVA720PConfig,
@@ -157,6 +158,7 @@ from sglang.multimodal_gen.configs.sample.ltx_2 import (
LTX23HQSamplingParams,
LTX23SamplingParams,
)
+from sglang.multimodal_gen.configs.sample.ltx_2_5 import LTX25SamplingParams
from sglang.multimodal_gen.configs.sample.minimax_h3 import MiniMaxH3SamplingParams
from sglang.multimodal_gen.configs.sample.mova import (
MOVA_360P_SamplingParams,
@@ -692,7 +694,9 @@ def _register_configs():
hf_model_paths=["Lightricks/LTX-2"],
model_detectors=[
lambda path: "ltx" in path.lower() and "video" in path.lower(),
- lambda path: "ltx-2" in path.lower() and "ltx-2.3" not in path.lower(),
+ lambda path: "ltx-2" in path.lower()
+ and "ltx-2.3" not in path.lower()
+ and "ltx-2.5" not in path.lower(),
],
)
register_configs(
@@ -703,6 +707,18 @@ def _register_configs():
lambda path: "ltx-2.3" in path.lower(),
],
)
+ # Keeps the LTX-2 pipeline class; only component geometry and the pinned
+ # distilled schedule differ. Only the `-Diffusers` repo is listed --
+ # `Lightricks/LTX-2.5` is a split pack of bare `.safetensors` and would need
+ # a model overlay first.
+ register_configs(
+ sampling_param_cls=LTX25SamplingParams,
+ pipeline_config_cls=LTX25PipelineConfig,
+ hf_model_paths=["Lightricks/LTX-2.5-Diffusers"],
+ model_detectors=[
+ lambda path: "ltx-2.5" in path.lower(),
+ ],
+ )
# register dedicated sampling params for LTX2TwoStageHQPipeline
_PIPELINE_CONFIG_REGISTRY.setdefault(
"LTX2TwoStageHQPipeline",
diff --git a/python/sglang/multimodal_gen/runtime/disaggregation/roles.py b/python/sglang/multimodal_gen/runtime/disaggregation/roles.py
index c7253490b..42ed1e8f0 100644
--- a/python/sglang/multimodal_gen/runtime/disaggregation/roles.py
+++ b/python/sglang/multimodal_gen/runtime/disaggregation/roles.py
@@ -37,6 +37,7 @@ def get_module_role(module_name: str) -> "RoleType | None":
"image_processor",
"processor",
"connectors",
+ "duration_head",
"vision_language_encoder",
)
if any(
@@ -62,7 +63,13 @@ def get_module_role(module_name: str) -> "RoleType | None":
if module_name == "hy3dshape_model":
return RoleType.DENOISER
- decoder_prefixes = ("vae", "audio_vae", "video_vae", "vocoder")
+ decoder_prefixes = (
+ "vae",
+ "audio_vae",
+ "video_vae",
+ "vocoder",
+ "diffusion_decoder",
+ )
if any(
module_name == p or module_name.startswith(p + "_") for p in decoder_prefixes
):
diff --git a/python/sglang/multimodal_gen/runtime/entrypoints/cli/generate.py b/python/sglang/multimodal_gen/runtime/entrypoints/cli/generate.py
index 3ef45b65c..82b99b3f0 100644
--- a/python/sglang/multimodal_gen/runtime/entrypoints/cli/generate.py
+++ b/python/sglang/multimodal_gen/runtime/entrypoints/cli/generate.py
@@ -177,6 +177,8 @@ def generate_cmd(args: argparse.Namespace, unknown_args: list[str] | None = None
sampling_params_kwargs.update(sampling_params_cls.get_cli_args(args))
_apply_output_file_path_override(args, sampling_params_kwargs)
sampling_params_kwargs["request_id"] = generate_request_id()
+ if sampling_params_kwargs.get("use_diffusion_decoder", False):
+ server_args.load_diffusion_decoder = True
# Handle diffusers-specific kwargs passed via CLI
if hasattr(args, "diffusers_kwargs") and args.diffusers_kwargs:
diff --git a/python/sglang/multimodal_gen/runtime/layers/quantization/fp8.py b/python/sglang/multimodal_gen/runtime/layers/quantization/fp8.py
index 79817ba62..e6a95a628 100644
--- a/python/sglang/multimodal_gen/runtime/layers/quantization/fp8.py
+++ b/python/sglang/multimodal_gen/runtime/layers/quantization/fp8.py
@@ -455,6 +455,13 @@ class Fp8LinearMethod(LinearMethodBase):
x: torch.Tensor,
bias: Optional[torch.Tensor] = None,
) -> torch.Tensor:
+ # The activation quantization kernels assert on row-major input, and
+ # diffusion backbones routinely pass a permuted view. Normalising at the
+ # producer instead would also move the unquantized path's output, by
+ # changing which GEMM kernel it picks. No-op when already contiguous.
+ if not x.is_contiguous():
+ x = x.contiguous()
+
if self.use_marlin:
return apply_fp8_marlin_linear(
input=x,
diff --git a/python/sglang/multimodal_gen/runtime/layers/usp.py b/python/sglang/multimodal_gen/runtime/layers/usp.py
index e40313cc2..1ca6619bc 100644
--- a/python/sglang/multimodal_gen/runtime/layers/usp.py
+++ b/python/sglang/multimodal_gen/runtime/layers/usp.py
@@ -173,6 +173,11 @@ def _ipc_input_a2a_qkv(q, k, v):
None when unavailable."""
if get_ulysses_parallel_world_size() != 2:
return None
+ # One staging slot is sized from `q` and reused for all three, so unequal
+ # k/v lengths would copy mismatched extents into it. The general exchange
+ # guards the same way and handles them.
+ if q.shape != k.shape or q.shape != v.shape:
+ return None
group = _ipc_ready_group()
if group is None:
return None
diff --git a/python/sglang/multimodal_gen/runtime/loader/component_loaders/adapter_loader.py b/python/sglang/multimodal_gen/runtime/loader/component_loaders/adapter_loader.py
index af18bfc23..f45a6d510 100644
--- a/python/sglang/multimodal_gen/runtime/loader/component_loaders/adapter_loader.py
+++ b/python/sglang/multimodal_gen/runtime/loader/component_loaders/adapter_loader.py
@@ -1,13 +1,16 @@
-from safetensors.torch import load_file as safetensors_load_file
+import re
from sglang.multimodal_gen.configs.models.adapter.ltx_2_connector import (
LTX2ConnectorConfig,
)
+from sglang.multimodal_gen.configs.models.adapter.ltx_2_duration_head import (
+ LTX2DurationHeadConfig,
+)
from sglang.multimodal_gen.runtime.loader.component_loaders.component_loader import (
ComponentLoader,
)
from sglang.multimodal_gen.runtime.loader.utils import (
- _list_safetensors_files,
+ load_safetensors_state_dict,
set_default_torch_dtype,
skip_init_modules,
)
@@ -24,14 +27,24 @@ class AdapterLoader(ComponentLoader):
This loader intentionally avoids FSDP sharding and just:
1) Instantiates the module from `config.json`.
- 2) Loads a single safetensors state_dict.
+ 2) Loads the safetensors state_dict (single-file or sharded).
"""
- component_names = ["connectors"]
+ component_names = ["connectors", "duration_head"]
expected_library = "diffusers"
+ # `update_model_arch` fills each from the component's `config.json`.
+ _CONFIG_CLASSES = {
+ "connectors": LTX2ConnectorConfig,
+ "duration_head": LTX2DurationHeadConfig,
+ }
+
def load_customized(
- self, component_model_path: str, server_args: ServerArgs, *args
+ self,
+ component_model_path: str,
+ server_args: ServerArgs,
+ component_name: str = "connectors",
+ *args,
):
config = get_diffusers_component_config(component_path=component_model_path)
@@ -45,33 +58,47 @@ class AdapterLoader(ComponentLoader):
config.pop("_diffusers_version", None)
config.pop("_name_or_path", None)
- server_args.model_paths["connectors"] = component_model_path
+ server_args.model_paths[component_name] = component_model_path
model_cls, _ = ModelRegistry.resolve_model_cls(cls_name)
+ # Not a fixed name: connectors follow DiT offload, while the duration
+ # head stays resident unless selected explicitly.
target_device = self.target_device(
- server_args.should_cpu_offload_component("connectors")
+ server_args.should_cpu_offload_component(component_name)
)
default_dtype = resolve_precision(
- server_args, "connectors", precision_attr="dit_precision"
+ server_args, component_name, precision_attr="dit_precision"
)
+ config_cls = self._CONFIG_CLASSES[component_name]
with set_default_torch_dtype(default_dtype), skip_init_modules():
- connector_cfg = LTX2ConnectorConfig()
- connector_cfg.update_model_arch(config)
- model = model_cls(connector_cfg).to(
- device=target_device, dtype=default_dtype
- )
+ adapter_cfg = config_cls()
+ adapter_cfg.update_model_arch(config)
+ model = model_cls(adapter_cfg).to(device=target_device, dtype=default_dtype)
- safetensors_list = _list_safetensors_files(component_model_path)
- if not safetensors_list:
- raise ValueError(f"No safetensors files found in {component_model_path}")
- if len(safetensors_list) != 1:
+ loaded = load_safetensors_state_dict(component_model_path)
+ mapping = adapter_cfg.arch_config.param_names_mapping
+ loaded = {_remap_connector_key(k, mapping): v for k, v in loaded.items()}
+
+ missing, unexpected = model.load_state_dict(loaded, strict=False)
+ # `strict=False` because a checkpoint carries either the shared
+ # `text_proj_in` or the per-modality projections, never both. Anything
+ # else uninitialized would surface later as garbage embeddings.
+ if missing or unexpected:
raise ValueError(
- f"Found {len(safetensors_list)} safetensors files in {component_model_path}, expected 1"
+ f"Adapter weights at '{component_model_path}' do not match the "
+ f"instantiated {cls_name}. Missing: {sorted(missing)}. "
+ f"Unexpected: {sorted(unexpected)}. This usually means the "
+ "adapter config or its weight-name mapping is wrong."
)
- loaded = safetensors_load_file(safetensors_list[0])
- model.load_state_dict(loaded, strict=False)
-
return model
+
+
+def _remap_connector_key(key: str, param_names_mapping: dict[str, str]) -> str:
+ for pattern, replacement in param_names_mapping.items():
+ key, replaced = re.subn(pattern, replacement, key)
+ if replaced:
+ break
+ return key
diff --git a/python/sglang/multimodal_gen/runtime/loader/component_loaders/component_loader.py b/python/sglang/multimodal_gen/runtime/loader/component_loaders/component_loader.py
index 21f9ff7bc..bc3b29fb3 100644
--- a/python/sglang/multimodal_gen/runtime/loader/component_loaders/component_loader.py
+++ b/python/sglang/multimodal_gen/runtime/loader/component_loaders/component_loader.py
@@ -416,7 +416,14 @@ class ComponentLoader(ABC):
self, transformers_or_diffusers: str, component_name: str
) -> str:
# NOTE(FlamingoPg): special for LTX-2 models
- if component_name == "vocoder" or component_name == "connectors":
+ # `model_index.json` records these under an `ltx2` library that is not a
+ # real importable package; SGLang implements them natively.
+ if component_name in (
+ "vocoder",
+ "connectors",
+ "duration_head",
+ "diffusion_decoder",
+ ):
transformers_or_diffusers = "diffusers"
# NOTE(CloudRipple): special for MOVA models
diff --git a/python/sglang/multimodal_gen/runtime/loader/component_loaders/diffusion_decoder_loader.py b/python/sglang/multimodal_gen/runtime/loader/component_loaders/diffusion_decoder_loader.py
new file mode 100644
index 000000000..444845c25
--- /dev/null
+++ b/python/sglang/multimodal_gen/runtime/loader/component_loaders/diffusion_decoder_loader.py
@@ -0,0 +1,62 @@
+# SPDX-License-Identifier: Apache-2.0
+
+from sglang.multimodal_gen.configs.models.decoders.ltx_2_5_diffusion_decoder import (
+ LTX25DiffusionDecoderConfig,
+)
+from sglang.multimodal_gen.runtime.loader.component_loaders.component_loader import (
+ ComponentLoader,
+)
+from sglang.multimodal_gen.runtime.loader.utils import (
+ load_safetensors_state_dict,
+ set_default_torch_dtype,
+ skip_init_modules,
+)
+from sglang.multimodal_gen.runtime.models.registry import ModelRegistry
+from sglang.multimodal_gen.runtime.server_args import ServerArgs
+from sglang.multimodal_gen.runtime.utils.hf_diffusers_utils import (
+ get_diffusers_component_config,
+)
+from sglang.multimodal_gen.runtime.utils.precision import resolve_precision
+
+
+class DiffusionDecoderLoader(ComponentLoader):
+ """Loader for the standalone, replicated LTX-2.5 diffusion decoder."""
+
+ component_names = ["diffusion_decoder"]
+ expected_library = "diffusers"
+
+ def load_customized(
+ self,
+ component_model_path: str,
+ server_args: ServerArgs,
+ component_name: str = "diffusion_decoder",
+ *args,
+ ):
+ config = get_diffusers_component_config(component_path=component_model_path)
+ class_name = config.pop("_class_name", None)
+ if class_name is None:
+ raise ValueError(
+ "Model config does not contain a _class_name attribute. "
+ "Only diffusers format is supported."
+ )
+ config.pop("_diffusers_version", None)
+ config.pop("_name_or_path", None)
+
+ server_args.model_paths[component_name] = component_model_path
+ model_cls, _ = ModelRegistry.resolve_model_cls(class_name)
+ target_device = self.target_device(
+ server_args.should_cpu_offload_component(component_name)
+ )
+ dtype = resolve_precision(
+ server_args, component_name, precision_attr="vae_precision"
+ )
+
+ decoder_config = LTX25DiffusionDecoderConfig()
+ decoder_config.update_model_arch(config)
+ with set_default_torch_dtype(dtype), skip_init_modules():
+ model = model_cls(decoder_config).to(device=target_device, dtype=dtype)
+
+ model.load_state_dict(
+ load_safetensors_state_dict(component_model_path), strict=True
+ )
+ return model
diff --git a/python/sglang/multimodal_gen/runtime/loader/component_loaders/upsampler_loader.py b/python/sglang/multimodal_gen/runtime/loader/component_loaders/upsampler_loader.py
index 1fa6751e3..778bd5c2c 100644
--- a/python/sglang/multimodal_gen/runtime/loader/component_loaders/upsampler_loader.py
+++ b/python/sglang/multimodal_gen/runtime/loader/component_loaders/upsampler_loader.py
@@ -98,7 +98,13 @@ def _normalize_config(raw: dict) -> dict:
# diffusers uses rational_spatial_scale instead of rational_resampler + spatial_scale
if "rational_spatial_scale" in raw and "rational_resampler" not in config:
- config["rational_resampler"] = True
+ # LTX-2.5 states this explicitly and turns it off, so the scale alone
+ # no longer implies it. Assuming True builds the wrong module (3 missing
+ # / 2 unexpected tensors).
+ if "use_rational_resampler" in raw:
+ config["rational_resampler"] = bool(raw["use_rational_resampler"])
+ else:
+ config["rational_resampler"] = True
config.setdefault("spatial_scale", raw["rational_spatial_scale"])
return config
diff --git a/python/sglang/multimodal_gen/runtime/loader/component_loaders/vocoder_loader.py b/python/sglang/multimodal_gen/runtime/loader/component_loaders/vocoder_loader.py
index 80d18675b..5f5292501 100644
--- a/python/sglang/multimodal_gen/runtime/loader/component_loaders/vocoder_loader.py
+++ b/python/sglang/multimodal_gen/runtime/loader/component_loaders/vocoder_loader.py
@@ -1,3 +1,5 @@
+import re
+
from safetensors.torch import load_file as safetensors_load_file
from sglang.multimodal_gen.runtime.loader.component_loaders.component_loader import (
@@ -61,24 +63,22 @@ class VocoderLoader(ComponentLoader):
len(safetensors_list) == 1
), f"Found {len(safetensors_list)} safetensors files in {component_model_path}"
loaded = safetensors_load_file(safetensors_list[0])
- incompatible = vocoder.load_state_dict(loaded, strict=False)
- missing_keys = []
- unexpected_keys = []
- try:
- missing_keys = incompatible.missing_keys
- unexpected_keys = incompatible.unexpected_keys
- except AttributeError:
- # Best-effort fallback in case older torch returns a tuple-like.
- try:
- missing_keys = incompatible[0]
- unexpected_keys = incompatible[1]
- except Exception:
- pass
+ mapping = vocoder_config.arch_config.param_names_mapping
+ loaded = {_remap_vocoder_key(k, mapping): v for k, v in loaded.items()}
+ missing_keys, unexpected_keys = vocoder.load_state_dict(loaded, strict=False)
+ # A half-loaded vocoder produces plausible but wrong audio.
if missing_keys or unexpected_keys:
- logger.warning(
- "Loaded vocoder with missing_keys=%d unexpected_keys=%d",
- len(missing_keys),
- len(unexpected_keys),
+ raise ValueError(
+ f"Vocoder weights at '{component_model_path}' do not match the "
+ f"instantiated {class_name}. Missing: {sorted(missing_keys)}. "
+ f"Unexpected: {sorted(unexpected_keys)}."
)
return vocoder
+
+
+def _remap_vocoder_key(key: str, param_names_mapping: dict[str, str]) -> str:
+ # Applied in order, not first-match: one key can need several rules.
+ for pattern, replacement in param_names_mapping.items():
+ key = re.sub(pattern, replacement, key)
+ return key
diff --git a/python/sglang/multimodal_gen/runtime/loader/utils.py b/python/sglang/multimodal_gen/runtime/loader/utils.py
index b2ce831f1..5fc3bac23 100644
--- a/python/sglang/multimodal_gen/runtime/loader/utils.py
+++ b/python/sglang/multimodal_gen/runtime/loader/utils.py
@@ -5,6 +5,7 @@
import contextlib
import glob
+import json
import os
import re
from collections import defaultdict
@@ -12,6 +13,7 @@ from collections.abc import Callable, Iterator
from typing import Any, Dict, Type
import torch
+from safetensors.torch import load_file as safetensors_load_file
from torch import nn
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
@@ -260,8 +262,6 @@ def _list_safetensors_files(model_path: str) -> list[str]:
str(model_path), "diffusion_pytorch_model.safetensors.index.json"
)
if os.path.exists(index_path):
- import json
-
with open(index_path) as f:
index = json.load(f)
expected_shards = sorted(set(index.get("weight_map", {}).values()))
@@ -284,6 +284,33 @@ def _list_safetensors_files(model_path: str) -> list[str]:
return found
+def load_safetensors_state_dict(model_path: str) -> dict[str, torch.Tensor]:
+ """Load one safetensors checkpoint, including an indexed sharded set."""
+ index_path = os.path.join(
+ str(model_path), "diffusion_pytorch_model.safetensors.index.json"
+ )
+ safetensors_files = _list_safetensors_files(model_path)
+ if os.path.exists(index_path):
+ with open(index_path) as f:
+ index = json.load(f)
+ shard_names = sorted(set(index.get("weight_map", {}).values()))
+ state_dict: dict[str, torch.Tensor] = {}
+ for shard_name in shard_names:
+ state_dict.update(
+ safetensors_load_file(os.path.join(str(model_path), shard_name))
+ )
+ return state_dict
+
+ if not safetensors_files:
+ raise ValueError(f"No safetensors files found in {model_path}")
+ if len(safetensors_files) != 1:
+ raise ValueError(
+ f"Found {len(safetensors_files)} safetensors files in {model_path} "
+ "and no index to disambiguate them."
+ )
+ return safetensors_load_file(safetensors_files[0])
+
+
BYTES_PER_GB = 1024**3
diff --git a/python/sglang/multimodal_gen/runtime/managers/memory_managers/layerwise_offload_components.py b/python/sglang/multimodal_gen/runtime/managers/memory_managers/layerwise_offload_components.py
index 4e803c8f5..80bbac252 100644
--- a/python/sglang/multimodal_gen/runtime/managers/memory_managers/layerwise_offload_components.py
+++ b/python/sglang/multimodal_gen/runtime/managers/memory_managers/layerwise_offload_components.py
@@ -39,6 +39,7 @@ VAE_COMPONENT_NAMES = frozenset(
"vocoder",
"spatial_upsampler",
"condition_image_encoder",
+ "diffusion_decoder",
}
)
DEFAULT_LAYERWISE_VAE_COMPONENT_NAMES = frozenset(
diff --git a/python/sglang/multimodal_gen/runtime/models/adapter/ltx_2_duration_head.py b/python/sglang/multimodal_gen/runtime/models/adapter/ltx_2_duration_head.py
new file mode 100644
index 000000000..cb0f57366
--- /dev/null
+++ b/python/sglang/multimodal_gen/runtime/models/adapter/ltx_2_duration_head.py
@@ -0,0 +1,193 @@
+# SPDX-License-Identifier: Apache-2.0
+"""LTX-2.5 duration head.
+
+Predicts the shot length a caption implies from the text connector outputs.
+Used only when the caller omits `num_frames`.
+"""
+
+import torch
+import torch.nn.functional as F
+from torch import nn
+
+from sglang.multimodal_gen.configs.models.adapter.ltx_2_duration_head import (
+ LTX2DurationHeadConfig,
+)
+from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
+
+logger = init_logger(__name__)
+
+
+class LTX2DurationAttentionPooler(nn.Module):
+ """Cross-attends `num_queries` learnable tokens against the caption tokens.
+
+ Produces a fixed `(batch, num_queries, hidden_dim)` output regardless of
+ input length. No attention mask: the connectors already replaced padded
+ positions with learnable registers and marked everything attendable.
+ """
+
+ def __init__(
+ self, hidden_dim: int = 256, num_queries: int = 1, num_heads: int = 4
+ ) -> None:
+ super().__init__()
+ self.heads = num_heads
+ self.query_tokens = nn.Parameter(torch.randn(num_queries, hidden_dim) * 0.02)
+ self.to_q = nn.Linear(hidden_dim, hidden_dim)
+ self.to_k = nn.Linear(hidden_dim, hidden_dim)
+ self.to_v = nn.Linear(hidden_dim, hidden_dim)
+ self.to_out = nn.Linear(hidden_dim, hidden_dim)
+
+ def forward(self, tokens: torch.Tensor) -> torch.Tensor:
+ queries = self.query_tokens.unsqueeze(0).expand(tokens.shape[0], -1, -1)
+
+ query = self.to_q(queries).unflatten(2, (self.heads, -1)).transpose(1, 2)
+ key = self.to_k(tokens).unflatten(2, (self.heads, -1)).transpose(1, 2)
+ value = self.to_v(tokens).unflatten(2, (self.heads, -1)).transpose(1, 2)
+
+ hidden_states = F.scaled_dot_product_attention(query, key, value)
+ hidden_states = hidden_states.transpose(1, 2).flatten(2, 3)
+ return self.to_out(hidden_states)
+
+
+class LTX2DurationHead(nn.Module):
+ """Modality-agnostic duration regressor over the connector outputs.
+
+ Per-modality input projections map each stream into a shared pooler width,
+ learnable modality embeddings tag the streams, and a small MLP turns the
+ pooled vector into a log-duration. The target is trained in log-seconds, so
+ `forward` exponentiates and callers always get seconds.
+ """
+
+ def __init__(self, config: LTX2DurationHeadConfig) -> None:
+ super().__init__()
+ arch = config.arch_config
+ pooler_hidden_dim = arch.pooler_hidden_dim
+
+ self.video_input_proj = nn.Linear(
+ arch.video_cross_attention_dim, pooler_hidden_dim
+ )
+ self.video_modality_emb = nn.Parameter(torch.randn(pooler_hidden_dim) * 0.02)
+
+ self.audio_input_proj = nn.Linear(
+ arch.audio_cross_attention_dim, pooler_hidden_dim
+ )
+ self.audio_modality_emb = nn.Parameter(torch.randn(pooler_hidden_dim) * 0.02)
+
+ self.attention_pooler = LTX2DurationAttentionPooler(
+ hidden_dim=pooler_hidden_dim,
+ num_queries=arch.num_queries,
+ num_heads=arch.num_pooler_heads,
+ )
+ self.mlp_hidden = nn.Linear(
+ pooler_hidden_dim * arch.num_queries, arch.mlp_hidden_dim
+ )
+ self.mlp_out = nn.Linear(arch.mlp_hidden_dim, 1)
+
+ def forward(
+ self,
+ video_tokens: torch.Tensor | None = None,
+ audio_tokens: torch.Tensor | None = None,
+ ) -> torch.Tensor:
+ """Returns predicted duration in seconds, shape `(batch,)`."""
+ if video_tokens is None and audio_tokens is None:
+ raise ValueError(
+ "LTX2DurationHead requires at least one of video_tokens / audio_tokens."
+ )
+
+ # The connector output can arrive in a different dtype than the head.
+ head_dtype = self.mlp_out.weight.dtype
+
+ token_groups = []
+ if video_tokens is not None:
+ token_groups.append(
+ self.video_input_proj(video_tokens.to(head_dtype))
+ + self.video_modality_emb
+ )
+ if audio_tokens is not None:
+ token_groups.append(
+ self.audio_input_proj(audio_tokens.to(head_dtype))
+ + self.audio_modality_emb
+ )
+
+ tokens = torch.cat(token_groups, dim=1)
+ pooled = self.attention_pooler(tokens).flatten(1)
+
+ # tanh-approximated GELU matches the JAX-trained head; exact GELU does not.
+ hidden_states = F.gelu(self.mlp_hidden(pooled), approximate="tanh")
+ log_duration = self.mlp_out(hidden_states).squeeze(-1)
+ return log_duration.exp()
+
+ def predict_num_frames(
+ self,
+ video_tokens: torch.Tensor | None = None,
+ audio_tokens: torch.Tensor | None = None,
+ *,
+ frame_rate: float,
+ temporal_compression_ratio: int,
+ min_seconds: float = 1.0,
+ max_seconds: float = 20.0,
+ ) -> int:
+ """Predict a frame count on the VAE's causal temporal grid.
+
+ Clamp first, then snap: a clamped count is not necessarily grid-aligned,
+ so snapping first would give a different answer.
+ """
+ predicted_seconds = self(video_tokens, audio_tokens)
+ if predicted_seconds.numel() != 1:
+ raise ValueError(
+ "predict_num_frames supports a single prediction only, got shape "
+ f"{tuple(predicted_seconds.shape)}. One frame count cannot serve "
+ "prompts with different natural durations."
+ )
+ seconds = predicted_seconds.item()
+
+ # Floor at 1 so the grid arithmetic cannot go negative.
+ min_frames = max(1, round(min_seconds * frame_rate))
+ max_frames = round(max_seconds * frame_rate)
+ clamped_frames = max(min_frames, min(round(seconds * frame_rate), max_frames))
+
+ num_frames = (
+ (clamped_frames - 1) // temporal_compression_ratio
+ ) * temporal_compression_ratio + 1
+
+ if num_frames < min_frames:
+ # Flooring undershot the lower bound; take the next grid point up.
+ snapped_up = num_frames + temporal_compression_ratio
+ if snapped_up <= max_frames:
+ num_frames = snapped_up
+ else:
+ # No grid point fits the bounds; overshooting by under a step
+ # beats refusing to generate.
+ if abs(snapped_up - clamped_frames) < abs(num_frames - clamped_frames):
+ num_frames = snapped_up
+ logger.warning(
+ "Duration bounds [%.2fs, %.2fs] at %.2f fps admit no frame count "
+ "on the VAE temporal grid (k * %d + 1); using nearest: %d frames",
+ min_seconds,
+ max_seconds,
+ frame_rate,
+ temporal_compression_ratio,
+ num_frames,
+ )
+
+ if seconds < min_seconds or seconds > max_seconds:
+ logger.warning(
+ "Duration prediction clamped: raw %.2fs outside [%.2fs, %.2fs], "
+ "using %.2fs (%d frames) @ %.2f fps",
+ seconds,
+ min_seconds,
+ max_seconds,
+ num_frames / frame_rate,
+ num_frames,
+ frame_rate,
+ )
+ else:
+ logger.info(
+ "Predicted duration %.2fs (%d frames @ %.2f fps)",
+ seconds,
+ num_frames,
+ frame_rate,
+ )
+ return num_frames
+
+
+EntryClass = LTX2DurationHead
diff --git a/python/sglang/multimodal_gen/runtime/models/decoders/ltx_2_5_diffusion_decoder.py b/python/sglang/multimodal_gen/runtime/models/decoders/ltx_2_5_diffusion_decoder.py
new file mode 100644
index 000000000..b7b5fb4a5
--- /dev/null
+++ b/python/sglang/multimodal_gen/runtime/models/decoders/ltx_2_5_diffusion_decoder.py
@@ -0,0 +1,1011 @@
+# SPDX-License-Identifier: Apache-2.0
+"""LTX-2.5 diffusion video decoder.
+
+An alternative to the convolutional VAE decoder: stages 1-4 deterministically
+upsample the latent into a context volume, and stage 5 denoises patchified
+pixels conditioned on it. As shipped (`model_output_type="x0"`, one step) that
+single prediction *is* the output.
+
+Attention is 3D *neighborhood* attention -- each query attends to a fixed window
+that shifts inward at the grid borders. Prefers NATTEN's fused `na3d` like
+upstream, falling back to a FlexAttention block mask when NATTEN is missing.
+"""
+
+import math
+
+import torch
+import torch.nn.functional as F
+from torch import nn
+
+from sglang.multimodal_gen.configs.models.decoders.ltx_2_5_diffusion_decoder import (
+ LTX25DiffusionDecoderConfig,
+)
+from sglang.multimodal_gen.runtime.layers.visual_embedding import (
+ timestep_embedding,
+)
+from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
+
+logger = init_logger(__name__)
+
+# Bounds the three hidden-width temporaries the SwiGLU holds live. Exact: the
+# MLP is pointwise across tokens.
+_SWIGLU_TILE_SIZE = 16384
+
+_na3d_fn: object = None
+_NA3D_UNAVAILABLE = object()
+
+
+def _na3d():
+ """NATTEN's fused 3D neighborhood attention, or `None` if unavailable.
+
+ ~4.8x the compiled `flex_attention` fallback, and needs no block mask.
+ """
+ global _na3d_fn
+ if _na3d_fn is None:
+ try:
+ from natten.functional import na3d
+
+ _na3d_fn = na3d
+ except ImportError:
+ _na3d_fn = _NA3D_UNAVAILABLE
+ return None if _na3d_fn is _NA3D_UNAVAILABLE else _na3d_fn
+
+
+_compiled_flex_attention = None
+
+
+class LTX2VideoDecoderTimestepEmbedder(nn.Module):
+ """Replicated native timestep MLP used by the standalone decoder.
+
+ The decoder is replicated across TP ranks, so these projections must remain
+ ordinary linear layers. Reusing the DiT's tensor-parallel embedder shards
+ their parameters and makes the unsharded decoder checkpoint unloadable.
+ """
+
+ def __init__(self, embedding_dim: int, in_channels: int = 256) -> None:
+ super().__init__()
+ self.linear_1 = nn.Linear(in_channels, embedding_dim, bias=True)
+ self.linear_2 = nn.Linear(embedding_dim, embedding_dim, bias=True)
+
+ def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
+ hidden_states = self.linear_1(hidden_states)
+ hidden_states = F.silu(hidden_states)
+ return self.linear_2(hidden_states)
+
+
+class LTX2VideoDecoderCombinedTimestepEmbeddings(nn.Module):
+ def __init__(self, embedding_dim: int) -> None:
+ super().__init__()
+ self.timestep_embedder = LTX2VideoDecoderTimestepEmbedder(embedding_dim)
+
+ def forward(
+ self, timestep: torch.Tensor, hidden_dtype: torch.dtype | None = None
+ ) -> torch.Tensor:
+ timestep = timestep.reshape(-1).to(dtype=torch.float32)
+ hidden_states = timestep_embedding(
+ timestep, dim=256, max_period=10000, dtype=torch.float32
+ )
+ if hidden_dtype is not None:
+ hidden_states = hidden_states.to(dtype=hidden_dtype)
+ return self.timestep_embedder(hidden_states)
+
+
+def _flex_attention_fn():
+ """`flex_attention`, compiled.
+
+ Uncompiled it falls back to materializing the full `S x S` score matrix,
+ which is tens of GiB at these grids -- compiling is what makes the
+ neighborhood window actually sparse. Compiled once and cached; the handful
+ of distinct decoder-stage shapes each trigger one recompile.
+ """
+ global _compiled_flex_attention
+ if _compiled_flex_attention is None:
+ from torch.nn.attention.flex_attention import flex_attention
+
+ _compiled_flex_attention = torch.compile(flex_attention, dynamic=False)
+ return _compiled_flex_attention
+
+
+def _patchify(x: torch.Tensor, patch_size: int) -> torch.Tensor:
+ """Space-to-depth on H/W only: `(B,C,F,H,W)` -> `(B, C*p**2, F, H//p, W//p)`.
+
+ Channel packing order is `(channel, width_offset, height_offset)`.
+ """
+ batch_size, num_channels, num_frames, height, width = x.shape
+ x = x.reshape(
+ batch_size,
+ num_channels,
+ num_frames,
+ height // patch_size,
+ patch_size,
+ width // patch_size,
+ patch_size,
+ )
+ x = x.permute(0, 1, 6, 4, 2, 3, 5)
+ return x.reshape(
+ batch_size,
+ num_channels * patch_size * patch_size,
+ num_frames,
+ height // patch_size,
+ width // patch_size,
+ )
+
+
+def _unpatchify(x: torch.Tensor, patch_size: int) -> torch.Tensor:
+ """Depth-to-space on H/W only; the exact inverse of `_patchify`."""
+ batch_size, num_channels, num_frames, height, width = x.shape
+ num_channels = num_channels // (patch_size * patch_size)
+ x = x.reshape(
+ batch_size, num_channels, patch_size, patch_size, num_frames, height, width
+ )
+ x = x.permute(0, 1, 4, 5, 3, 6, 2)
+ return x.reshape(
+ batch_size, num_channels, num_frames, height * patch_size, width * patch_size
+ )
+
+
+# O(S^2) to build (17.5 s at 1.08M tokens) and a pure function of grid and
+# kernel, so a server at one resolution pays it once.
+_BLOCK_MASK_CACHE: dict = {}
+_BLOCK_MASK_CACHE_MAX = 16
+
+
+def _neighborhood_block_mask(
+ num_frames: int,
+ height: int,
+ width: int,
+ kernel_size: tuple[int, int, int],
+ device: torch.device,
+):
+ """FlexAttention `BlockMask` for a 3D neighborhood window.
+
+ The window is centered where possible and shifted inward at the borders so
+ it always holds exactly `kernel_size` positions -- that inward shift, rather
+ than truncation, is what NATTEN's `na3d` does.
+ """
+ from torch.nn.attention.flex_attention import create_block_mask
+
+ cache_key = (num_frames, height, width, tuple(kernel_size), str(device))
+ cached = _BLOCK_MASK_CACHE.get(cache_key)
+ if cached is not None:
+ return cached
+
+ kernel_t, kernel_h, kernel_w = kernel_size
+ kernel_t = min(kernel_t, num_frames)
+ kernel_h = min(kernel_h, height)
+ kernel_w = min(kernel_w, width)
+ hw = height * width
+
+ def mask_mod(batch_idx, head_idx, q_idx, kv_idx):
+ q_t, q_rem = q_idx // hw, q_idx % hw
+ q_h, q_w = q_rem // width, q_rem % width
+ k_t, k_rem = kv_idx // hw, kv_idx % hw
+ k_h, k_w = k_rem // width, k_rem % width
+
+ start_t = torch.clamp(q_t - kernel_t // 2, 0, num_frames - kernel_t)
+ start_h = torch.clamp(q_h - kernel_h // 2, 0, height - kernel_h)
+ start_w = torch.clamp(q_w - kernel_w // 2, 0, width - kernel_w)
+ window_t = (k_t >= start_t) & (k_t < start_t + kernel_t)
+ window_h = (k_h >= start_h) & (k_h < start_h + kernel_h)
+ window_w = (k_w >= start_w) & (k_w < start_w + kernel_w)
+ return window_t & window_h & window_w
+
+ seq_len = num_frames * hw
+ # `_compile=True` is required, not an optimisation: the eager path
+ # materialises O(S^2) booleans, tens of GiB at production grids.
+ block_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_MASK_CACHE) >= _BLOCK_MASK_CACHE_MAX:
+ _BLOCK_MASK_CACHE.pop(next(iter(_BLOCK_MASK_CACHE)))
+ _BLOCK_MASK_CACHE[cache_key] = block_mask
+ return block_mask
+
+
+class LTX2VideoVaeRotaryPosEmbed3D(nn.Module):
+ """Absolute 3D rotary embedding over the (T, H, W) grid.
+
+ `head_dim` splits into (T, H, W) chunks, each rotated by its own axis
+ position.
+ """
+
+ def __init__(self, head_dim: int, base: float = 10000.0) -> None:
+ super().__init__()
+ if head_dim % 8 != 0:
+ raise ValueError(f"head_dim must be a multiple of 8, got {head_dim}.")
+ # A quarter to T, the rest split H/W, both kept even for whole
+ # rotation pairs.
+ dim_t = (head_dim // 4) // 2 * 2
+ dim_hw = (head_dim - dim_t) // 2
+ if dim_hw % 2 != 0:
+ dim_t -= 2
+ dim_hw = (head_dim - dim_t) // 2
+ self.rope_dim_split = (dim_t, dim_hw, dim_hw)
+ self.base = base
+
+ def _inv_freqs(self, dim: int, device: torch.device) -> torch.Tensor:
+ exponents = torch.arange(0, dim, 2, dtype=torch.float64, device=device) / dim
+ return (1.0 / self.base**exponents).to(torch.float32)
+
+ def _rotate_axis(
+ self,
+ x: torch.Tensor,
+ positions: torch.Tensor,
+ inv_freqs: torch.Tensor,
+ axis: int,
+ ) -> torch.Tensor:
+ out_dtype = x.dtype
+ pairs = x.reshape(*x.shape[:-1], x.shape[-1] // 2, 2)
+ even = pairs[..., 0].float()
+ odd = pairs[..., 1].float()
+ # Broadcast over (B, T, H, W, heads, dim // 2), varying only along `axis`.
+ shape = [1, 1, 1, 1, 1, inv_freqs.shape[0]]
+ shape[axis] = positions.shape[0]
+ angles = (positions[:, None] * inv_freqs[None, :]).reshape(shape)
+ cos, sin = angles.cos(), angles.sin()
+ rotated = torch.stack([even * cos - odd * sin, even * sin + odd * cos], dim=-1)
+ return rotated.reshape(x.shape).to(out_dtype)
+
+ def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
+ """`hidden_states`: `(B, T, H, W, heads, head_dim)`."""
+ dim_t, dim_h, _ = self.rope_dim_split
+ num_frames, height, width = hidden_states.shape[1:4]
+ device = hidden_states.device
+ inv_t, inv_h, inv_w = (
+ self._inv_freqs(dim, device) for dim in self.rope_dim_split
+ )
+
+ positions_t = torch.arange(num_frames, dtype=torch.float32, device=device)
+ positions_h = torch.arange(height, dtype=torch.float32, device=device)
+ positions_w = torch.arange(width, dtype=torch.float32, device=device)
+ rotated_t = self._rotate_axis(
+ hidden_states[..., :dim_t], positions_t, inv_t, axis=1
+ )
+ rotated_h = self._rotate_axis(
+ hidden_states[..., dim_t : dim_t + dim_h], positions_h, inv_h, axis=2
+ )
+ rotated_w = self._rotate_axis(
+ hidden_states[..., dim_t + dim_h :], positions_w, inv_w, axis=3
+ )
+ return torch.cat([rotated_t, rotated_h, rotated_w], dim=-1)
+
+
+class LTX2VideoVaeNeighborhoodAttention(nn.Module):
+ """3D neighborhood attention over a channels-last `(B, T, H, W, C)` volume."""
+
+ def __init__(
+ self,
+ dim: int,
+ kernel_size: tuple[int, int, int],
+ head_dim: int = 64,
+ rope_base: float = 10000.0,
+ ) -> None:
+ super().__init__()
+ if dim % head_dim != 0:
+ raise ValueError(f"dim {dim} must be divisible by head_dim {head_dim}.")
+ self.heads = dim // head_dim
+ self.head_dim = head_dim
+ self.kernel_size = tuple(kernel_size)
+ self.scale = head_dim**-0.5
+
+ self.to_q = nn.Linear(dim, dim, bias=True)
+ self.to_k = nn.Linear(dim, dim, bias=True)
+ self.to_v = nn.Linear(dim, dim, bias=True)
+ self.to_out = nn.ModuleList([nn.Linear(dim, dim, bias=True), nn.Dropout(0.0)])
+ self.norm_q = nn.RMSNorm(head_dim, eps=1e-6)
+ self.norm_k = nn.RMSNorm(head_dim, eps=1e-6)
+ self.rope = LTX2VideoVaeRotaryPosEmbed3D(head_dim, base=rope_base)
+
+ def project_qkv(
+ self, hidden_states: torch.Tensor
+ ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
+ """Q/K/V as `(B, T, H, W, heads, head_dim)`: normed, query pre-scaled, rotated.
+
+ The query carries the `1/sqrt(head_dim)` factor so attention runs with
+ `scale=1.0` -- upstream's order is norm, scale, then rotate.
+ """
+ batch_size, num_frames, height, width, _ = hidden_states.shape
+ shape = (batch_size, num_frames, height, width, self.heads, self.head_dim)
+ query = self.to_q(hidden_states).view(shape)
+ key = self.to_k(hidden_states).view(shape)
+ value = self.to_v(hidden_states).view(shape)
+
+ query = self.norm_q(query)
+ key = self.norm_k(key)
+ query = query * self.scale
+ return self.rope(query), self.rope(key), value
+
+ def build_block_mask(self, hidden_states: torch.Tensor):
+ """The window mask for this grid, or `None` when NATTEN handles it.
+
+ Fixed within a stage, so built once.
+ """
+ if _na3d() is not None:
+ return None
+ num_frames, height, width = hidden_states.shape[1:4]
+ return _neighborhood_block_mask(
+ num_frames, height, width, self.kernel_size, hidden_states.device
+ )
+
+ def forward(self, hidden_states: torch.Tensor, block_mask=None) -> torch.Tensor:
+ batch_size, num_frames, height, width, _ = hidden_states.shape
+ kernel_t, kernel_h, kernel_w = self.kernel_size
+ if num_frames < kernel_t or height < kernel_h or width < kernel_w:
+ raise ValueError(
+ "Neighborhood attention requires each dim to be at least its "
+ f"kernel size; got (T, H, W) = ({num_frames}, {height}, {width}) "
+ f"with kernel_size {self.kernel_size}."
+ )
+
+ query, key, value = self.project_qkv(hidden_states)
+
+ na3d = _na3d()
+ if na3d is not None:
+ # `project_qkv` already yields NATTEN's layout. scale=1.0: the
+ # query is pre-scaled there.
+ hidden_states = na3d(
+ query, key, value, kernel_size=self.kernel_size, scale=1.0
+ )
+ hidden_states = hidden_states.reshape(
+ batch_size, num_frames, height, width, self.heads * self.head_dim
+ )
+ return self.to_out[0](hidden_states)
+
+ seq_len = num_frames * height * width
+ # flex_attention wants (B, heads, S, head_dim).
+ query = query.reshape(batch_size, seq_len, self.heads, self.head_dim).transpose(
+ 1, 2
+ )
+ key = key.reshape(batch_size, seq_len, self.heads, self.head_dim).transpose(
+ 1, 2
+ )
+ value = value.reshape(batch_size, seq_len, self.heads, self.head_dim).transpose(
+ 1, 2
+ )
+
+ if block_mask is None:
+ block_mask = self.build_block_mask(hidden_states)
+
+ # scale=1.0: the query is already scaled in project_qkv.
+ hidden_states = _flex_attention_fn()(
+ query, key, value, block_mask=block_mask, scale=1.0
+ )
+ hidden_states = hidden_states.transpose(1, 2).reshape(
+ batch_size, num_frames, height, width, self.heads * self.head_dim
+ )
+ return self.to_out[0](hidden_states)
+
+
+class LTX2VideoVaeSwiGLU(nn.Module):
+ """`w_down(silu(w_gate(x)) * w_up(x))`, evaluated in token tiles."""
+
+ def __init__(self, dim: int, hidden_dim: int) -> None:
+ super().__init__()
+ self.w_up = nn.Linear(dim, hidden_dim, bias=False)
+ self.w_gate = nn.Linear(dim, hidden_dim, bias=False)
+ self.w_down = nn.Linear(hidden_dim, dim, bias=False)
+
+ def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
+ batch_size, *token_dims, channels = hidden_states.shape
+ num_tokens = math.prod(token_dims)
+ if num_tokens <= _SWIGLU_TILE_SIZE:
+ return self.w_down(
+ F.silu(self.w_gate(hidden_states)) * self.w_up(hidden_states)
+ )
+
+ flat = hidden_states.reshape(batch_size, num_tokens, channels)
+ out = torch.empty_like(flat)
+ for start in range(0, num_tokens, _SWIGLU_TILE_SIZE):
+ tile = flat[:, start : start + _SWIGLU_TILE_SIZE]
+ out[:, start : start + _SWIGLU_TILE_SIZE] = self.w_down(
+ F.silu(self.w_gate(tile)) * self.w_up(tile)
+ )
+ return out.reshape(hidden_states.shape)
+
+
+def _swiglu_hidden_dim(dim: int, mlp_ratio: float) -> int:
+ return (int(dim * mlp_ratio) + 15) // 16 * 16
+
+
+class LTX2VideoVaeNABlock(nn.Module):
+ """Pre-norm neighborhood-attention block used by the deterministic stages."""
+
+ def __init__(
+ self,
+ dim: int,
+ kernel_size: tuple[int, int, int],
+ head_dim: int = 64,
+ mlp_ratio: float = 4.0,
+ ) -> None:
+ super().__init__()
+ self.norm1 = nn.RMSNorm(dim, eps=1e-6)
+ self.attn = LTX2VideoVaeNeighborhoodAttention(
+ dim, kernel_size, head_dim=head_dim
+ )
+ self.norm2 = nn.RMSNorm(dim, eps=1e-6)
+ self.mlp = LTX2VideoVaeSwiGLU(dim, _swiglu_hidden_dim(dim, mlp_ratio))
+
+ def forward(self, hidden_states: torch.Tensor, block_mask=None) -> torch.Tensor:
+ hidden_states = hidden_states + self.attn(self.norm1(hidden_states), block_mask)
+ hidden_states = hidden_states + self.mlp(self.norm2(hidden_states))
+ return hidden_states
+
+
+class LTX2VideoVaeAdaLNZero(nn.Module):
+ """Timestep embedding to seven `(B, 1, 1, 1, C)` modulation chunks.
+
+ Seven is upstream's shape (scale/shift/gate for attention and MLP, plus a
+ context gate); only the four scale/shift chunks are consumed, since the
+ residuals here are ungated.
+ """
+
+ def __init__(self, dim: int, t_emb_dim: int, num_chunks: int = 7) -> None:
+ super().__init__()
+ self.num_chunks = num_chunks
+ self.proj = nn.Linear(t_emb_dim, num_chunks * dim, bias=True)
+
+ def forward(self, t_emb: torch.Tensor) -> tuple[torch.Tensor, ...]:
+ chunks = self.proj(F.silu(t_emb)).chunk(self.num_chunks, dim=-1)
+ return tuple(chunk[:, None, None, None, :] for chunk in chunks)
+
+
+class LTX2VideoVaeDiffusionNABlock(nn.Module):
+ """Stage-5 block: neighborhood attention + SwiGLU under AdaLN-Zero modulation."""
+
+ def __init__(
+ self,
+ dim: int,
+ kernel_size: tuple[int, int, int],
+ context_channels: int,
+ head_dim: int = 64,
+ mlp_ratio: float = 4.0,
+ num_mod_params: int = 7,
+ ) -> None:
+ super().__init__()
+ self.context_channels = context_channels
+ self.num_mod_params = num_mod_params
+ self.context_proj = nn.Linear(context_channels, dim, bias=True)
+ self.scale_shift_table = nn.Parameter(torch.zeros(num_mod_params, dim))
+
+ self.norm1 = nn.RMSNorm(dim, eps=1e-6)
+ self.attn = LTX2VideoVaeNeighborhoodAttention(
+ dim, kernel_size, head_dim=head_dim
+ )
+ self.norm2 = nn.RMSNorm(dim, eps=1e-6)
+ self.mlp = LTX2VideoVaeSwiGLU(dim, _swiglu_hidden_dim(dim, mlp_ratio))
+
+ def forward(
+ self,
+ hidden_states: torch.Tensor,
+ latent_context: torch.Tensor,
+ modulation: tuple[torch.Tensor, ...],
+ block_mask=None,
+ ) -> torch.Tensor:
+ scale_msa, shift_msa, _, scale_mlp, shift_mlp, _, _ = [
+ modulation[i] + self.scale_shift_table[i].view(1, 1, 1, 1, -1)
+ for i in range(self.num_mod_params)
+ ]
+
+ hidden_states = hidden_states + self.context_proj(latent_context)
+ hidden_states = hidden_states + self.attn(
+ self.norm1(hidden_states) * (1 + scale_msa) + shift_msa, block_mask
+ )
+ hidden_states = hidden_states + self.mlp(
+ self.norm2(hidden_states) * (1 + scale_mlp) + shift_mlp
+ )
+ return hidden_states
+
+
+class LTX2VideoVaePixelShuffleUpsampler(nn.Module):
+ """Linear channel expansion then a channels-last pixel shuffle.
+
+ A temporal stride of 2 produces a duplicate leading frame, dropped to keep
+ the causal 1:2 (composed 1:8) frame mapping.
+ """
+
+ def __init__(
+ self,
+ in_channels: int,
+ stride: tuple[int, int, int],
+ out_channels_reduction_factor: int = 1,
+ ) -> None:
+ super().__init__()
+ self.stride = tuple(stride)
+ proj_out_channels = (
+ math.prod(self.stride) * in_channels // out_channels_reduction_factor
+ )
+ self.out_channels = proj_out_channels // math.prod(self.stride)
+ self.proj = nn.Linear(in_channels, proj_out_channels, bias=True)
+
+ def forward(
+ self, hidden_states: torch.Tensor, drop_leading_frame: bool = True
+ ) -> torch.Tensor:
+ batch_size, num_frames, height, width, _ = hidden_states.shape
+ stride_t, stride_h, stride_w = self.stride
+ hidden_states = self.proj(hidden_states)
+ hidden_states = hidden_states.reshape(
+ batch_size,
+ num_frames,
+ height,
+ width,
+ self.out_channels,
+ stride_t,
+ stride_h,
+ stride_w,
+ )
+ hidden_states = hidden_states.permute(0, 1, 5, 2, 6, 3, 7, 4)
+ hidden_states = hidden_states.reshape(
+ batch_size,
+ num_frames * stride_t,
+ height * stride_h,
+ width * stride_w,
+ self.out_channels,
+ )
+ if stride_t == 2 and drop_leading_frame:
+ hidden_states = hidden_states[:, 1:]
+ return hidden_states
+
+
+class LTX2VideoDiffusionDecoder3d(nn.Module):
+ """Stages 1-4 upsample the latent into a context volume; stage 5 denoises
+ patchified pixels conditioned on it."""
+
+ def __init__(self, config: LTX25DiffusionDecoderConfig) -> None:
+ super().__init__()
+ arch = config.arch_config
+ stage_channels = tuple(arch.decoder_stage_channels)
+ stage_depths = tuple(arch.decoder_stage_depths)
+ stage_kernels = tuple(tuple(k) for k in arch.decoder_stage_kernels)
+ upsample_strides = tuple(tuple(s) for s in arch.decoder_upsample_strides)
+ reductions = tuple(arch.decoder_upsample_channel_reductions)
+
+ if arch.decoder_model_output_type not in ("x0", "v"):
+ raise ValueError(
+ "decoder_model_output_type must be 'x0' or 'v', got "
+ f"{arch.decoder_model_output_type!r}."
+ )
+ # An inconsistent pair would only fail deep inside the first block.
+ for stage_idx, reduction in enumerate(reductions):
+ expected = stage_channels[stage_idx] // reduction
+ if stage_channels[stage_idx + 1] != expected:
+ raise ValueError(
+ f"decoder_stage_channels[{stage_idx + 1}] must be "
+ f"{expected}, got {stage_channels[stage_idx + 1]}."
+ )
+
+ self.patch_size = arch.patch_size
+ self.out_channels = arch.out_channels
+ self.timestep_scale_multiplier = arch.decoder_timestep_scale_multiplier
+ self.model_output_type = arch.decoder_model_output_type
+ self.default_num_inference_steps = arch.decoder_num_inference_steps
+ self.temporal_compression_ratio = arch.temporal_compression_ratio
+ self.context_channels = stage_channels[-1]
+ # Replicated through stages 1-4 and cropped before stage 5, moving the
+ # border effect past the frames that are kept.
+ self.trailing_pad_latent_frames = (stage_kernels[0][0] // 2) * 2
+
+ self.conv_in = nn.Linear(arch.latent_channels, stage_channels[0], bias=True)
+
+ self.det_stages = nn.ModuleList()
+ self.upsamples = nn.ModuleList()
+ for stage_idx, stride in enumerate(upsample_strides):
+ channels = stage_channels[stage_idx]
+ self.det_stages.append(
+ nn.ModuleList(
+ [
+ LTX2VideoVaeNABlock(
+ dim=channels,
+ kernel_size=stage_kernels[stage_idx],
+ head_dim=arch.decoder_head_dim,
+ )
+ for _ in range(stage_depths[stage_idx])
+ ]
+ )
+ )
+ self.upsamples.append(
+ LTX2VideoVaePixelShuffleUpsampler(
+ in_channels=channels,
+ stride=stride,
+ out_channels_reduction_factor=reductions[stage_idx],
+ )
+ )
+
+ self.t_embedder = LTX2VideoDecoderCombinedTimestepEmbeddings(
+ embedding_dim=arch.decoder_t_emb_dim
+ )
+
+ stage5_channels = stage_channels[-1]
+ noised_pixel_channels = arch.out_channels * arch.patch_size**2
+ self.conv_in_x_t = nn.Linear(noised_pixel_channels, stage5_channels, bias=True)
+ self.shared_adaln = LTX2VideoVaeAdaLNZero(
+ dim=stage5_channels, t_emb_dim=arch.decoder_t_emb_dim
+ )
+ self.diff_blocks = nn.ModuleList(
+ [
+ LTX2VideoVaeDiffusionNABlock(
+ dim=stage5_channels,
+ kernel_size=tuple(arch.decoder_stage5_kernel),
+ context_channels=self.context_channels,
+ head_dim=arch.decoder_head_dim,
+ num_mod_params=self.shared_adaln.num_chunks,
+ )
+ for _ in range(stage_depths[-1])
+ ]
+ )
+ self.norm_out = nn.RMSNorm(stage5_channels, eps=1e-6)
+ self.conv_out = nn.Linear(stage5_channels, noised_pixel_channels, bias=True)
+
+ def forward_stages_1_to_3(self, hidden_states: torch.Tensor) -> torch.Tensor:
+ """Latent `(B, C, T, H, W)` to a channels-last feature volume."""
+ num_pad = self.trailing_pad_latent_frames
+ if num_pad > 0:
+ trailing = hidden_states[:, :, -1:].expand(-1, -1, num_pad, -1, -1)
+ hidden_states = torch.cat([hidden_states, trailing], dim=2)
+
+ hidden_states = hidden_states.permute(0, 2, 3, 4, 1)
+ hidden_states = self.conv_in(hidden_states)
+ for blocks, upsample in zip(self.det_stages[:-1], self.upsamples[:-1]):
+ # Fixed within a stage, so one mask serves all of its blocks.
+ block_mask = blocks[0].attn.build_block_mask(hidden_states)
+ for block in blocks:
+ hidden_states = block(hidden_states, block_mask)
+ hidden_states = upsample(hidden_states)
+ return hidden_states
+
+ def forward_stage_4(
+ self,
+ hidden_states: torch.Tensor,
+ drop_leading_frame: bool = True,
+ crop_trailing_ghost: bool = True,
+ ) -> torch.Tensor:
+ """Last deterministic stage -> context `(B, T5, H5, W5, C5)`."""
+ blocks = self.det_stages[-1]
+ block_mask = blocks[0].attn.build_block_mask(hidden_states)
+ for block in blocks:
+ hidden_states = block(hidden_states, block_mask)
+ hidden_states = self.upsamples[-1](
+ hidden_states, drop_leading_frame=drop_leading_frame
+ )
+
+ num_pad = self.trailing_pad_latent_frames
+ if crop_trailing_ghost and num_pad > 0:
+ hidden_states = hidden_states[
+ :, : -num_pad * self.temporal_compression_ratio
+ ]
+ return hidden_states
+
+ def forward_diffusion_step(
+ self, latent_context: torch.Tensor, x_t: torch.Tensor, timestep: torch.Tensor
+ ) -> torch.Tensor:
+ """One stage-5 step; returns a pixel-space prediction `(B, C, F, H, W)`."""
+ t_emb = self.t_embedder(
+ self.timestep_scale_multiplier * timestep,
+ hidden_dtype=latent_context.dtype,
+ )
+ modulation = self.shared_adaln(t_emb)
+
+ hidden_states = _patchify(x_t, self.patch_size).permute(0, 2, 3, 4, 1)
+ hidden_states = self.conv_in_x_t(hidden_states)
+ block_mask = self.diff_blocks[0].attn.build_block_mask(hidden_states)
+ for block in self.diff_blocks:
+ hidden_states = block(hidden_states, latent_context, modulation, block_mask)
+
+ hidden_states = self.norm_out(hidden_states)
+ hidden_states = self.conv_out(hidden_states)
+ hidden_states = hidden_states.permute(0, 4, 1, 2, 3).contiguous()
+ return _unpatchify(hidden_states, self.patch_size)
+
+ def denoise(
+ self,
+ latent_context: torch.Tensor,
+ x_t: torch.Tensor,
+ num_inference_steps: int,
+ ) -> torch.Tensor:
+ batch_size = latent_context.shape[0]
+ timesteps = torch.linspace(
+ 1.0,
+ 1.0 / num_inference_steps,
+ num_inference_steps,
+ device=latent_context.device,
+ dtype=torch.float32,
+ )
+
+ # How LTX-2.5 ships: one step whose x0 prediction is the output.
+ if num_inference_steps == 1 and self.model_output_type == "x0":
+ return self.forward_diffusion_step(
+ latent_context, x_t, timesteps[:1].expand(batch_size)
+ )
+
+ for step_idx in range(num_inference_steps):
+ t_now = timesteps[step_idx].expand(batch_size)
+ t_next = (
+ timesteps[step_idx + 1]
+ if step_idx + 1 < num_inference_steps
+ else torch.zeros_like(t_now)
+ )
+ model_out = self.forward_diffusion_step(latent_context, x_t, t_now).float()
+ x_t_fp32 = x_t.float()
+ if self.model_output_type == "x0":
+ sigma = t_now.view(-1, *([1] * (x_t.ndim - 1)))
+ model_out = (x_t_fp32 - model_out) / sigma
+ dt = (t_now - t_next).view(-1, *([1] * (x_t.ndim - 1)))
+ x_t = (x_t_fp32 - dt * model_out).to(x_t.dtype)
+ return x_t
+
+ def forward(
+ self,
+ hidden_states: torch.Tensor,
+ generator: torch.Generator | None = None,
+ num_inference_steps: int | None = None,
+ ) -> torch.Tensor:
+ num_inference_steps = num_inference_steps or self.default_num_inference_steps
+ latent_context = self.forward_stage_4(self.forward_stages_1_to_3(hidden_states))
+ # Pixel canvas = stage-5 token grid times the patch size.
+ pixel_shape = (
+ hidden_states.shape[0],
+ self.out_channels,
+ latent_context.shape[1],
+ latent_context.shape[2] * self.patch_size,
+ latent_context.shape[3] * self.patch_size,
+ )
+ x_t = torch.randn(
+ pixel_shape,
+ generator=generator,
+ device=hidden_states.device,
+ dtype=hidden_states.dtype,
+ )
+ return self.denoise(latent_context, x_t, num_inference_steps)
+
+
+def _tile_intervals(
+ length: int, tile_size: int, stride: int, min_size: int
+) -> list[tuple[int, int]]:
+ """Overlapping `[start, end)` tiles covering `[0, length)`.
+
+ A trailing remnant shorter than `min_size` is merged into the previous tile
+ rather than decoded alone: neighborhood attention rejects any grid smaller
+ than its kernel, so a short remnant cannot always stand on its own.
+ """
+ if length <= tile_size:
+ return [(0, length)]
+ starts = list(range(0, length, stride))
+ while len(starts) > 1 and length - starts[-1] < min_size:
+ starts.pop()
+ return [(start, min(start + tile_size, length)) for start in starts[:-1]] + [
+ (starts[-1], length)
+ ]
+
+
+class LTX2VideoDiffusionDecoderModel(nn.Module):
+ """Checkpoint-level wrapper: the decoder plus the latent statistics.
+
+ `diffusion_decoder/` stores `latents_mean` / `latents_std` alongside a
+ `decoder.` submodule, so this mirrors that layout rather than flattening it.
+ """
+
+ def __init__(self, config: LTX25DiffusionDecoderConfig) -> None:
+ super().__init__()
+ self.config = config
+ latent_channels = config.arch_config.latent_channels
+ self.decoder = LTX2VideoDiffusionDecoder3d(config)
+ self.register_buffer(
+ "latents_mean", torch.zeros(latent_channels), persistent=True
+ )
+ self.register_buffer(
+ "latents_std", torch.ones(latent_channels), persistent=True
+ )
+
+ # Tiles the last deterministic stage and the diffusion blocks, so the
+ # output only moves near tile borders. Set by the decoding stage from
+ # `--diffusion-decoder-tiling`; tile sizes match upstream.
+ self.use_tiling = False
+ self.tile_sample_min_height = 768
+ self.tile_sample_min_width = 768
+ self.tile_sample_min_num_frames = 32
+ self.tile_sample_stride_height = 512
+ self.tile_sample_stride_width = 512
+ self.tile_sample_stride_num_frames = 16
+
+ @staticmethod
+ def _blend(a: torch.Tensor, b: torch.Tensor, extent: int, dim: int) -> torch.Tensor:
+ """Linear cross-fade of `a`'s tail into `b`'s head along `dim`."""
+ extent = min(a.shape[dim], b.shape[dim], extent)
+ if extent <= 0:
+ return b
+ ramp = torch.arange(extent, device=b.device, dtype=torch.float32) / extent
+ shape = [1] * b.ndim
+ shape[dim] = extent
+ ramp = ramp.reshape(shape).to(b.dtype)
+ a_tail = a.narrow(dim, a.shape[dim] - extent, extent)
+ b_head = b.narrow(dim, 0, extent)
+ b_head.copy_(a_tail * (1 - ramp) + b_head * ramp)
+ return b
+
+ def _should_tile(self, hidden_states: torch.Tensor) -> bool:
+ if not self.use_tiling:
+ return False
+ arch = self.config.arch_config
+ return (
+ hidden_states.shape[2]
+ > self.tile_sample_min_num_frames // arch.temporal_compression_ratio
+ or hidden_states.shape[3]
+ > self.tile_sample_min_height // arch.spatial_compression_ratio
+ or hidden_states.shape[4]
+ > self.tile_sample_min_width // arch.spatial_compression_ratio
+ )
+
+ def tiled_decode(
+ self,
+ hidden_states: torch.Tensor,
+ generator: torch.Generator | None = None,
+ num_inference_steps: int | None = None,
+ ) -> torch.Tensor:
+ """Decode with stage 4 and the diffusion stage running per tile.
+
+ Tiles live on the grid entering the last deterministic stage, where one
+ cell maps to a fixed block of output pixels. Temporal tiles follow the
+ causal frame mapping: only the tile holding t=0 drops the temporal
+ upsample's duplicate leading frame, and only the tile holding the end of
+ the video carries the border padding to crop.
+ """
+ decoder = self.decoder
+ arch = self.config.arch_config
+ num_inference_steps = num_inference_steps or decoder.default_num_inference_steps
+ batch_size = hidden_states.shape[0]
+ patch_size = decoder.patch_size
+
+ # Pixels per tiling-grid cell: the last upsample's stride times the patch.
+ stride_up = decoder.upsamples[-1].stride
+ scale_t, scale_h, scale_w = (
+ stride_up[0],
+ stride_up[1] * patch_size,
+ stride_up[2] * patch_size,
+ )
+ tile_t = self.tile_sample_min_num_frames // scale_t
+ step_t = self.tile_sample_stride_num_frames // scale_t
+ tile_h = self.tile_sample_min_height // scale_h
+ step_h = self.tile_sample_stride_height // scale_h
+ tile_w = self.tile_sample_min_width // scale_w
+ step_w = self.tile_sample_stride_width // scale_w
+ # Stage 4 sees the tile as-is, stage 5 sees it scaled by the stride.
+ min_sizes = [
+ max(k4, -(-k5 // stride))
+ for k4, k5, stride in zip(
+ arch.decoder_stage_kernels[-1], arch.decoder_stage5_kernel, stride_up
+ )
+ ]
+
+ features = decoder.forward_stages_1_to_3(hidden_states)
+ # Trailing ghost frames replicate through the earlier temporal
+ # upsamples; the composed mapping is affine in their stride product.
+ ghost = decoder.trailing_pad_latent_frames * math.prod(
+ up.stride[0] for up in decoder.upsamples[:-1]
+ )
+ num_frames = features.shape[1] - ghost
+ height, width = features.shape[2], features.shape[3]
+
+ temporal_tiles = _tile_intervals(num_frames, tile_t, step_t, min_sizes[0])
+ height_tiles = _tile_intervals(height, tile_h, step_h, min_sizes[1])
+ width_tiles = _tile_intervals(width, tile_w, step_w, min_sizes[2])
+ blend_frames = (tile_t - step_t) * scale_t
+ blend_height = (tile_h - step_h) * scale_h
+ blend_width = (tile_w - step_w) * scale_w
+
+ # Single-step x0 predicts from pure noise, so tiles may draw their own.
+ # Multi-step integrates across steps and needs one shared canvas.
+ single_step_x0 = num_inference_steps == 1 and decoder.model_output_type == "x0"
+ x_t_full = None
+ if not single_step_x0:
+ pixel_frames = num_frames * scale_t - (1 if scale_t == 2 else 0)
+ x_t_full = torch.randn(
+ (
+ batch_size,
+ decoder.out_channels,
+ pixel_frames,
+ height * scale_h,
+ width * scale_w,
+ ),
+ generator=generator,
+ device=hidden_states.device,
+ dtype=hidden_states.dtype,
+ )
+
+ frame_groups = []
+ for t0, t1 in temporal_tiles:
+ is_origin = t0 == 0
+ is_trailing = t1 == num_frames
+ feature_t1 = features.shape[1] if is_trailing else t1
+ rows = []
+ for h0, h1 in height_tiles:
+ row = []
+ for w0, w1 in width_tiles:
+ context = decoder.forward_stage_4(
+ features[:, t0:feature_t1, h0:h1, w0:w1],
+ drop_leading_frame=is_origin,
+ crop_trailing_ghost=is_trailing,
+ )
+ tile_shape = (
+ batch_size,
+ decoder.out_channels,
+ context.shape[1],
+ context.shape[2] * patch_size,
+ context.shape[3] * patch_size,
+ )
+ if single_step_x0:
+ x_t = torch.randn(
+ tile_shape,
+ generator=generator,
+ device=hidden_states.device,
+ dtype=hidden_states.dtype,
+ )
+ else:
+ # A non-origin tile keeps its duplicate leading frame, so
+ # it starts one pixel frame earlier than t0 * scale_t.
+ pixel_t0 = t0 * scale_t - (
+ 1 if not is_origin and scale_t == 2 else 0
+ )
+ x_t = x_t_full[
+ :,
+ :,
+ pixel_t0 : pixel_t0 + tile_shape[2],
+ h0 * scale_h : h0 * scale_h + tile_shape[3],
+ w0 * scale_w : w0 * scale_w + tile_shape[4],
+ ]
+ row.append(decoder.denoise(context, x_t, num_inference_steps))
+ rows.append(row)
+
+ result_rows = []
+ for i, row in enumerate(rows):
+ result_row = []
+ for j, tile in enumerate(row):
+ if i > 0:
+ tile = self._blend(rows[i - 1][j], tile, blend_height, dim=3)
+ if j > 0:
+ tile = self._blend(row[j - 1], tile, blend_width, dim=4)
+ # The last tile can run past the stride grid, since a short
+ # remnant is merged into it, so it keeps its full extent.
+ keep_h = step_h * scale_h if i < len(rows) - 1 else tile.shape[3]
+ keep_w = step_w * scale_w if j < len(row) - 1 else tile.shape[4]
+ result_row.append(tile[:, :, :, :keep_h, :keep_w])
+ result_rows.append(torch.cat(result_row, dim=4))
+ frame_groups.append(torch.cat(result_rows, dim=3))
+
+ result = []
+ for k, group in enumerate(frame_groups):
+ if k > 0:
+ group = self._blend(frame_groups[k - 1], group, blend_frames, dim=2)
+ if k < len(frame_groups) - 1:
+ # The origin group is one frame short of stride * scale: its
+ # first cell decodes to a single pixel frame under the causal
+ # mapping.
+ keep_frames = step_t * scale_t - (1 if k == 0 and scale_t == 2 else 0)
+ group = group[:, :, :keep_frames]
+ result.append(group)
+ return torch.cat(result, dim=2)
+
+ def forward(
+ self,
+ hidden_states: torch.Tensor,
+ generator: torch.Generator | None = None,
+ num_inference_steps: int | None = None,
+ ) -> torch.Tensor:
+ if self._should_tile(hidden_states):
+ return self.tiled_decode(hidden_states, generator, num_inference_steps)
+ return self.decoder(hidden_states, generator, num_inference_steps)
+
+ def decode(
+ self,
+ hidden_states: torch.Tensor,
+ generator: torch.Generator | None = None,
+ num_inference_steps: int | None = None,
+ ) -> torch.Tensor:
+ return self.forward(hidden_states, generator, num_inference_steps)
+
+
+EntryClass = LTX2VideoDiffusionDecoderModel
diff --git a/python/sglang/multimodal_gen/runtime/models/dits/ltx_2.py b/python/sglang/multimodal_gen/runtime/models/dits/ltx_2.py
index 0c31ac6af..819269c98 100644
--- a/python/sglang/multimodal_gen/runtime/models/dits/ltx_2.py
+++ b/python/sglang/multimodal_gen/runtime/models/dits/ltx_2.py
@@ -1050,6 +1050,7 @@ class LTX2FeedForward(nn.Module):
dim: int,
dim_out: int | None = None,
mult: int = 4,
+ bias: bool = True,
quant_config: QuantizationConfig | None = None,
) -> None:
super().__init__()
@@ -1058,13 +1059,13 @@ class LTX2FeedForward(nn.Module):
inner_dim = int(dim * mult)
self.proj_in = ColumnParallelLinear(
- dim, inner_dim, bias=True, gather_output=False, quant_config=quant_config
+ dim, inner_dim, bias=bias, gather_output=False, quant_config=quant_config
)
self.act = nn.GELU(approximate="tanh")
self.proj_out = RowParallelLinear(
inner_dim,
dim_out,
- bias=True,
+ bias=bias,
input_is_parallel=True,
quant_config=quant_config,
)
@@ -1096,6 +1097,8 @@ class LTX2TransformerBlock(nn.Module):
norm_eps: float = 1e-6,
apply_gated_attention: bool = False,
cross_attention_adaln: bool = False,
+ ff_bias: bool = True,
+ audio_ff_bias: bool = True,
use_local_av_cross_attention: bool = False,
force_sdpa_v2a_cross_attention: bool = False,
enable_packed_qkv_input_a2a: bool = False,
@@ -1202,10 +1205,13 @@ class LTX2TransformerBlock(nn.Module):
)
# 4. Feedforward layers
- self.ff = LTX2FeedForward(dim, dim_out=dim, quant_config=quant_config)
+ # LTX-2.5: `ff_bias: false`, `audio_ff_bias: true`.
+ self.ff = LTX2FeedForward(
+ dim, dim_out=dim, bias=ff_bias, quant_config=quant_config
+ )
mark_ltx2_rms_norm_modulate_site(self)
self.audio_ff = LTX2FeedForward(
- audio_dim, dim_out=audio_dim, quant_config=quant_config
+ audio_dim, dim_out=audio_dim, bias=audio_ff_bias, quant_config=quant_config
)
# 5. Modulation Parameters
@@ -1562,6 +1568,8 @@ class LTX2TransformerBlock(nn.Module):
class LTX2VideoTransformer3DModel(CachableDiT, LayerwiseOffloadableModuleMixin):
_fsdp_shard_conditions = [is_blocks_or_transformer_blocks]
_compile_conditions = [is_blocks_or_transformer_blocks]
+ # Class-level defaults satisfy BaseDiT's `__init_subclass__` contract;
+ # `__init__` overrides them per instance so variants can extend the mapping.
param_names_mapping = LTX2ArchConfig().param_names_mapping
reverse_param_names_mapping = LTX2ArchConfig().reverse_param_names_mapping
lora_param_names_mapping = LTX2ArchConfig().lora_param_names_mapping
@@ -1639,6 +1647,10 @@ class LTX2VideoTransformer3DModel(CachableDiT, LayerwiseOffloadableModuleMixin):
super().__init__(config=config, hf_config=hf_config)
arch = self.config
+ # Checkpoint naming is arch-config metadata, not a runtime capability.
+ self.param_names_mapping = arch.param_names_mapping
+ self.reverse_param_names_mapping = arch.reverse_param_names_mapping
+ self.lora_param_names_mapping = arch.lora_param_names_mapping
self.hidden_size = arch.hidden_size
self.num_attention_heads = arch.num_attention_heads
self.audio_hidden_size = arch.audio_hidden_size
@@ -1665,6 +1677,15 @@ class LTX2VideoTransformer3DModel(CachableDiT, LayerwiseOffloadableModuleMixin):
quant_config=quant_config,
)
+ # Marks single-pixel-frame keyframe tokens. Zero-initialized upstream
+ # and unused by the denoising forward; held so the checkpoint
+ # round-trips.
+ self.keyframes_abs_pos_embedding: nn.Parameter | None = None
+ if arch.use_keyframes_abs_pos_embedding:
+ self.keyframes_abs_pos_embedding = nn.Parameter(
+ torch.zeros(1, self.hidden_size)
+ )
+
# 2. Prompt embeddings
self.caption_projection: LTX2TextProjection | None = None
self.audio_caption_projection: LTX2TextProjection | None = None
@@ -1841,6 +1862,8 @@ class LTX2VideoTransformer3DModel(CachableDiT, LayerwiseOffloadableModuleMixin):
qk_norm=True, # Always True in LTX2
apply_gated_attention=arch.apply_gated_attention,
cross_attention_adaln=arch.cross_attention_adaln,
+ ff_bias=arch.ff_bias,
+ audio_ff_bias=arch.audio_ff_bias,
use_local_av_cross_attention=bool(
getattr(arch, "use_local_av_cross_attention", False)
),
diff --git a/python/sglang/multimodal_gen/runtime/models/registry.py b/python/sglang/multimodal_gen/runtime/models/registry.py
index 9eb0a0a32..3af49cb6d 100644
--- a/python/sglang/multimodal_gen/runtime/models/registry.py
+++ b/python/sglang/multimodal_gen/runtime/models/registry.py
@@ -372,6 +372,13 @@ class _ModelRegistry:
normalized_arch = []
for arch in architectures:
if arch not in self.registered_models:
+ # A checkpoint may name a class that is only a rename of one we
+ # already implement (e.g. LTX-2.5's `LTX2VocoderWithBWE` is
+ # `LTX2Vocoder`); `_aliases` declares those equivalences.
+ canonical = _ALIAS_TO_MODEL.get(arch)
+ if canonical is not None and canonical in self.registered_models:
+ normalized_arch.append(canonical)
+ continue
registered_models = list(self.registered_models.keys())
raise Exception(
f"Unsupported model architecture: {arch}. Registered architectures: {registered_models}"
diff --git a/python/sglang/multimodal_gen/runtime/models/vaes/ltx_2_vae.py b/python/sglang/multimodal_gen/runtime/models/vaes/ltx_2_vae.py
index e1da51bbf..45493e533 100644
--- a/python/sglang/multimodal_gen/runtime/models/vaes/ltx_2_vae.py
+++ b/python/sglang/multimodal_gen/runtime/models/vaes/ltx_2_vae.py
@@ -866,6 +866,15 @@ class LTX23VideoMidBlock3d(nn.Module):
# Like LTXVideoUpBlock3d but with no conv_in and the updated LTX2VideoResnetBlock3d
+# Per-stage upsampling strides, selected by the decoder's `upsample_type`.
+# LTX-2 upsamples every stage in 3D; LTX-2.5 mixes spatial- and temporal-only.
+_UPSAMPLE_STRIDES: dict[str, tuple[int, int, int]] = {
+ "spatial": (1, 2, 2),
+ "temporal": (2, 1, 1),
+ "spatiotemporal": (2, 2, 2),
+}
+
+
class LTX2VideoUpBlock3d(nn.Module):
r"""
Up block used in the LTXVideo model.
@@ -901,6 +910,7 @@ class LTX2VideoUpBlock3d(nn.Module):
resnet_eps: float = 1e-6,
resnet_act_fn: str = "swish",
spatio_temporal_scale: bool = True,
+ upsample_type: str = "spatiotemporal",
inject_noise: bool = False,
timestep_conditioning: bool = False,
upsample_residual: bool = False,
@@ -936,7 +946,7 @@ class LTX2VideoUpBlock3d(nn.Module):
[
LTXVideoUpsampler3d(
out_channels * upscale_factor,
- stride=(2, 2, 2),
+ stride=_UPSAMPLE_STRIDES[upsample_type],
residual=upsample_residual,
upscale_factor=upscale_factor,
spatial_padding_mode=spatial_padding_mode,
@@ -1236,6 +1246,7 @@ class LTX2VideoDecoder3d(nn.Module):
timestep_conditioning: bool = False,
upsample_residual: Tuple[bool, ...] = (True, True, True),
upsample_factor: Tuple[bool, ...] = (2, 2, 2),
+ upsample_type: Tuple[str, ...] | None = None,
spatial_padding_mode: str = "reflect",
) -> None:
super().__init__()
@@ -1245,12 +1256,17 @@ class LTX2VideoDecoder3d(nn.Module):
self.out_channels = out_channels * patch_size**2
self.is_causal = is_causal
+ if upsample_type is None:
+ upsample_type = ("spatiotemporal",) * len(block_out_channels)
+
block_out_channels = tuple(reversed(block_out_channels))
spatio_temporal_scaling = tuple(reversed(spatio_temporal_scaling))
layers_per_block = tuple(reversed(layers_per_block))
inject_noise = tuple(reversed(inject_noise))
upsample_residual = tuple(reversed(upsample_residual))
upsample_factor = tuple(reversed(upsample_factor))
+ # Deliberately not reversed: upstream indexes `upsample_type` in
+ # decoder order, the sibling lists in encoder order.
output_channel = block_out_channels[0]
self.conv_in = LTX2VideoCausalConv3d(
@@ -1283,6 +1299,7 @@ class LTX2VideoDecoder3d(nn.Module):
num_layers=layers_per_block[i + 1],
resnet_eps=resnet_norm_eps,
spatio_temporal_scale=spatio_temporal_scaling[i],
+ upsample_type=upsample_type[i],
inject_noise=inject_noise[i + 1],
timestep_conditioning=timestep_conditioning,
upsample_residual=upsample_residual[i],
@@ -1624,28 +1641,25 @@ class AutoencoderKLLTX2Video(ParallelTiledVAE):
config.arch_config.decoder_spatio_temporal_scaling
)
decoder_layers_per_block = config.arch_config.decoder_layers_per_block
- decoder_inject_noise = getattr(
- config.arch_config, "decoder_inject_noise", (False, False, False, False)
- )
+ decoder_inject_noise = config.arch_config.decoder_inject_noise
if isinstance(decoder_inject_noise, bool):
decoder_inject_noise = (decoder_inject_noise,) * 4
else:
decoder_inject_noise = tuple(decoder_inject_noise)
- upsample_residual = getattr(
- config.arch_config, "upsample_residual", (True, True, True)
- )
+ upsample_residual = config.arch_config.upsample_residual
if isinstance(upsample_residual, bool):
upsample_residual = (upsample_residual,) * 3
else:
upsample_residual = tuple(upsample_residual)
- upsample_factor = getattr(config.arch_config, "upsample_factor", (2, 2, 2))
+ upsample_factor = config.arch_config.upsample_factor
if isinstance(upsample_factor, int):
upsample_factor = (upsample_factor,) * 3
else:
upsample_factor = tuple(upsample_factor)
- timestep_conditioning = getattr(
- config.arch_config, "timestep_conditioning", False
- )
+ upsample_type = config.arch_config.upsample_type
+ if upsample_type is not None:
+ upsample_type = tuple(upsample_type)
+ timestep_conditioning = config.arch_config.timestep_conditioning
use_ltx23_video_decoder = (
str(config.arch_config.video_decoder_variant) == "ltx_2_3"
)
@@ -1732,6 +1746,7 @@ class AutoencoderKLLTX2Video(ParallelTiledVAE):
timestep_conditioning=timestep_conditioning,
upsample_residual=upsample_residual,
upsample_factor=upsample_factor,
+ upsample_type=upsample_type,
spatial_padding_mode=decoder_spatial_padding_mode,
)
diff --git a/python/sglang/multimodal_gen/runtime/models/vocoder/ltx_2_vocoder.py b/python/sglang/multimodal_gen/runtime/models/vocoder/ltx_2_vocoder.py
index 6432db263..7e604ca5d 100644
--- a/python/sglang/multimodal_gen/runtime/models/vocoder/ltx_2_vocoder.py
+++ b/python/sglang/multimodal_gen/runtime/models/vocoder/ltx_2_vocoder.py
@@ -537,8 +537,14 @@ class LTX23VocoderCore(nn.Module):
class LTX2Vocoder(ABC, nn.Module, LayerwiseOffloadableModuleMixin):
r"""
LTX 2.0 vocoder for converting generated mel spectrograms back to audio waveforms.
+
+ Also serves LTX-2.5, whose `LTX2VocoderWithBWE` adds a bandwidth-extension
+ stage on top of the same generator: the base stack synthesises at 16 kHz and
+ the BWE stack resynthesises at 48 kHz from a mel re-analysis.
"""
+ _aliases = ["LTX2VocoderWithBWE"]
+
layerwise_offload_dit_group_enabled = False
layer_names = [
"upsamplers",
diff --git a/python/sglang/multimodal_gen/runtime/pipelines/ltx_2_pipeline.py b/python/sglang/multimodal_gen/runtime/pipelines/ltx_2_pipeline.py
index da6cdd65f..8dd667981 100644
--- a/python/sglang/multimodal_gen/runtime/pipelines/ltx_2_pipeline.py
+++ b/python/sglang/multimodal_gen/runtime/pipelines/ltx_2_pipeline.py
@@ -1,3 +1,4 @@
+import json
import math
import os
@@ -42,6 +43,7 @@ from sglang.multimodal_gen.runtime.pipelines_core.stages.model_specific_stages.l
LTX2AVDecodingStage,
LTX2AVDenoisingStage,
LTX2AVLatentPreparationStage,
+ LTX2DurationStage,
LTX2HalveResolutionStage,
LTX2LoRASwitchStage,
LTX2RefinementStage,
@@ -168,6 +170,20 @@ class LTX2SigmaPreparationStage(PipelineStage):
def forward(self, batch: Req, server_args: ServerArgs) -> Req:
batch.extra["ltx2_phase"] = "stage1"
+ pinned_sigmas = server_args.pipeline_config.default_sigmas
+ if pinned_sigmas:
+ # Distilled checkpoints ship an explicit schedule; a generic linear
+ # one silently costs quality.
+ if int(batch.num_inference_steps) != len(pinned_sigmas):
+ logger.info(
+ "Overriding num_inference_steps=%d with the pinned distilled "
+ "sigma schedule (%d steps).",
+ int(batch.num_inference_steps),
+ len(pinned_sigmas),
+ )
+ batch.sigmas = list(pinned_sigmas)
+ batch.num_inference_steps = len(pinned_sigmas)
+ return batch
if is_ltx23_native_variant(server_args.pipeline_config.vae_config.arch_config):
# Gate on pipeline class to mirror the three official entry points:
# - HQ (`ti2vid_two_stages_hq.py:164`) calls
@@ -220,6 +236,11 @@ def _add_ltx2_front_stages(pipeline: ComposedPipelineBase):
LTX2TextConnectorStage(connectors=pipeline.get_module("connectors")),
]
)
+ # Must run before latent preparation, which derives shapes from
+ # `num_frames`. A no-op unless the request sets `auto_duration`.
+ duration_head = pipeline.get_module("duration_head", None)
+ if duration_head is not None:
+ pipeline.add_stage(LTX2DurationStage(duration_head=duration_head))
def _add_ltx2_stage1_generation_stages(
@@ -260,6 +281,8 @@ def _add_ltx2_decoding_stage(pipeline: ComposedPipelineBase):
audio_vae=pipeline.get_module("audio_vae"),
vocoder=pipeline.get_module("vocoder"),
pipeline=pipeline,
+ # LTX-2.5 only; None elsewhere, which keeps the VAE decode path.
+ diffusion_decoder=pipeline.get_module("diffusion_decoder", None),
)
)
@@ -314,9 +337,88 @@ class _BaseLTX2Pipeline(LoRAPipeline):
"connectors",
]
+ # `model_index.json` points at the distilled DiT; the full / SFT weights in
+ # `transformer_full/` are deliberately omitted from it.
+ _DEV_VARIANTS = frozenset({"dev", "full", "sft"})
+ _DEV_TRANSFORMER_SUBFOLDER = "transformer_full"
+
+ def __init__(self, model_path, server_args, required_config_modules=None, **kwargs):
+ self._maybe_route_dev_transformer(model_path, server_args)
+ # LTX-2 / 2.3 ship neither. The small duration head is always available
+ # when declared; the much larger decoder is loaded only on request.
+ modules = list(required_config_modules or self._required_config_modules)
+ if "duration_head" not in modules and self._declares_component(
+ model_path, "duration_head"
+ ):
+ modules.append("duration_head")
+ if server_args.load_diffusion_decoder:
+ if not self._declares_component(model_path, "diffusion_decoder"):
+ raise ValueError(
+ "--load-diffusion-decoder was requested, but this checkpoint "
+ "does not declare a diffusion_decoder component."
+ )
+ if "diffusion_decoder" not in modules:
+ modules.append("diffusion_decoder")
+ super().__init__(
+ model_path, server_args, required_config_modules=modules, **kwargs
+ )
+
+ @classmethod
+ def _is_dev_variant(cls, server_args: ServerArgs) -> bool:
+ return str(server_args.model_variant or "").lower() in cls._DEV_VARIANTS
+
+ @classmethod
+ def _maybe_route_dev_transformer(cls, model_path: str, server_args: ServerArgs):
+ """Point the transformer at `transformer_full/` for the dev variant."""
+ if not cls._is_dev_variant(server_args):
+ return
+ if server_args.component_paths.get("transformer"):
+ return
+ full_path = os.path.join(str(model_path), cls._DEV_TRANSFORMER_SUBFOLDER)
+ if not os.path.isdir(full_path):
+ raise ValueError(
+ f"--model-variant {server_args.model_variant} requires "
+ f"'{cls._DEV_TRANSFORMER_SUBFOLDER}' in the checkpoint, but "
+ f"{full_path} does not exist. It is excluded from "
+ "`model_index.json`, so a partial snapshot download may have "
+ "skipped it."
+ )
+ server_args.component_paths["transformer"] = full_path
+ logger.info("Serving the LTX-2.5 dev transformer from %s", full_path)
+
+ @staticmethod
+ def _declares_component(model_path: str, component_name: str) -> bool:
+ index_path = os.path.join(str(model_path), "model_index.json")
+ if not os.path.exists(index_path):
+ return False
+ try:
+ with open(index_path) as f:
+ model_index = json.load(f)
+ except (OSError, ValueError):
+ return False
+ entry = model_index.get(component_name)
+ # model_index.json records absent optional components as [null, null].
+ return bool(entry) and entry[0] is not None
+
def initialize_pipeline(self, server_args: ServerArgs):
orig = self.get_module("scheduler")
- self.modules["scheduler"] = LTX2FlowMatchScheduler.from_config(orig.config)
+ scheduler_overrides: dict = {}
+ if self._is_dev_variant(server_args):
+ # `scheduler/` is configured for the distilled DiT; the full DiT
+ # needs the shifting back.
+ scheduler_overrides = {
+ "use_dynamic_shifting": True,
+ "shift_terminal": 0.1,
+ }
+ # It is also driven by a step count, not the distilled sigma list.
+ server_args.pipeline_config.default_sigmas = None
+ logger.info(
+ "LTX-2.5 dev variant: re-enabled dynamic shifting and dropped the "
+ "pinned distilled sigma schedule."
+ )
+ self.modules["scheduler"] = LTX2FlowMatchScheduler.from_config(
+ orig.config, **scheduler_overrides
+ )
sync_ltx23_runtime_vae_markers(
server_args.pipeline_config.vae_config.arch_config,
getattr(self.get_module("vae"), "config", None),
@@ -582,8 +684,12 @@ class LTX2TwoStagePipeline(_BaseLTX2Pipeline):
self.modules["spatial_upsampler"] = module
self.memory_usages["spatial_upsampler"] = memory_usage
+ # LTX-2 / 2.3 merge a distilled LoRA per stage; LTX-2.5's transformer
+ # is already distilled, so the LoRA is optional there.
distilled_lora_path = server_args.component_paths.get("distilled_lora")
- if not distilled_lora_path:
+ if not distilled_lora_path and not self._transformer_is_predistilled(
+ server_args
+ ):
raise ValueError(
f"{self.pipeline_name} requires --distilled-lora-path "
"(component_paths['distilled_lora'])."
@@ -599,6 +705,15 @@ class LTX2TwoStagePipeline(_BaseLTX2Pipeline):
self._stage1_distilled_in_base = False
self._stage1_distilled_base_strength: float | None = None
+ @staticmethod
+ def _transformer_is_predistilled(server_args: ServerArgs) -> bool:
+ """Whether the checkpoint's own transformer is already distilled.
+
+ True for LTX-2.5, whose `model_index.json` points at the distilled DiT
+ and which pins the distilled sigma schedule rather than shipping a LoRA.
+ """
+ return bool(server_args.pipeline_config.default_sigmas)
+
def _initialize_premerged_stage2_transformer(self, server_args: ServerArgs) -> None:
transformer_path = self._resolve_component_path(
server_args, "transformer", "transformer"
@@ -740,6 +855,10 @@ class LTX2TwoStagePipeline(_BaseLTX2Pipeline):
return False
def should_skip_ltx2_lora_switch_stage(self) -> bool:
+ # Nothing to switch when the DiT is already distilled (LTX-2.5): there
+ # is no distilled LoRA, and both stages run the same weights.
+ if self._distilled_lora_path is None:
+ return True
return (
self._use_premerged_stage2_transformer
and self._ltx2_residency.mode == "resident"
@@ -820,6 +939,11 @@ class LTX2TwoStagePipeline(_BaseLTX2Pipeline):
return lora_nicknames, lora_paths, lora_strengths, lora_targets
def switch_lora_phase(self, phase: str, batch: Req | None = None) -> None:
+ # A pre-distilled DiT has no LoRA to switch to and runs the same
+ # weights in both stages. Guarding here covers every caller.
+ if self._distilled_lora_path is None:
+ self._active_lora_phase = phase
+ return
distilled_lora_strength = self._get_stage_distilled_lora_strength(phase, batch)
phase_signature = (phase, distilled_lora_strength)
if phase_signature == self._active_lora_signature:
diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/image_encoding.py b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/image_encoding.py
index 05a667d24..4ac1fe22b 100644
--- a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/image_encoding.py
+++ b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/image_encoding.py
@@ -535,6 +535,21 @@ class LTX2ImageEncodingStage(PipelineStage):
# -- image preprocessing ---------------------------------------------
+ # Conditioning images are re-compressed to match training: CRF 33 for
+ # LTX-2 / 2.3, 18 for LTX-2.5. Like upstream, keyed off the text-encoder
+ # generation -- the only signal that separates them.
+ _DEFAULT_IMAGE_CRF = 33
+ _LTX_2_5_IMAGE_CRF = 18
+ _GEMMA_4_MODEL_TYPES = ("gemma4_unified", "gemma4")
+
+ @classmethod
+ def _resolve_image_conditioning_crf(cls, server_args: ServerArgs) -> int:
+ text_encoder_configs = server_args.pipeline_config.text_encoder_configs
+ for encoder_config in text_encoder_configs:
+ if encoder_config.prefix in ("gemma_4_unified", "gemma_4"):
+ return cls._LTX_2_5_IMAGE_CRF
+ return cls._DEFAULT_IMAGE_CRF
+
@staticmethod
def _apply_video_codec_compression(
img_array: np.ndarray, crf: int = 33
@@ -704,11 +719,12 @@ class LTX2ImageEncodingStage(PipelineStage):
from sglang.multimodal_gen.runtime.utils.vision import load_image
# 1. Load images, apply codec compression, resize for condition_image
+ crf = self._resolve_image_conditioning_crf(server_args)
conditioned_imgs = []
for image_path in image_paths:
img = load_image(image_path)
arr = np.array(img).astype(np.uint8)[..., :3]
- arr = self._apply_video_codec_compression(arr, crf=33)
+ arr = self._apply_video_codec_compression(arr, crf=crf)
conditioned_img = PIL.Image.fromarray(arr)
conditioned_imgs.append(conditioned_img)
batch.condition_image = [
diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/ltx_2/__init__.py b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/ltx_2/__init__.py
index 8853051a8..3dac02908 100644
--- a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/ltx_2/__init__.py
+++ b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/ltx_2/__init__.py
@@ -12,6 +12,9 @@ from sglang.multimodal_gen.runtime.pipelines_core.stages.model_specific_stages.l
LTX2AVDenoisingStage,
LTX2RefinementStage,
)
+from sglang.multimodal_gen.runtime.pipelines_core.stages.model_specific_stages.ltx_2.duration import (
+ LTX2DurationStage,
+)
from sglang.multimodal_gen.runtime.pipelines_core.stages.model_specific_stages.ltx_2.latent_preparation_av import (
LTX2AVLatentPreparationStage,
)
@@ -26,6 +29,7 @@ from sglang.multimodal_gen.runtime.pipelines_core.stages.model_specific_stages.l
__all__ = [
"LTX2AVDecodingStage",
+ "LTX2DurationStage",
"LTX2AVDenoisingStage",
"LTX2AVLatentPreparationStage",
"LTX2DenoisingStage",
diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/ltx_2/decoding_av.py b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/ltx_2/decoding_av.py
index 2c6f9ce7a..e5f58189f 100644
--- a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/ltx_2/decoding_av.py
+++ b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/ltx_2/decoding_av.py
@@ -12,6 +12,7 @@ from sglang.multimodal_gen.runtime.utils.precision import (
align_tensor_to_module_dtype,
autocast_context,
autocast_enabled,
+ resolve_decode_precision,
resolve_precision,
temporary_module_dtype,
)
@@ -24,10 +25,13 @@ class LTX2AVDecodingStage(DecodingStage):
LTX-2 specific decoding stage that handles both video and audio decoding.
"""
- def __init__(self, vae, audio_vae, vocoder, pipeline=None):
+ def __init__(self, vae, audio_vae, vocoder, pipeline=None, diffusion_decoder=None):
super().__init__(vae, pipeline)
self.audio_vae = audio_vae
self.vocoder = vocoder
+ # Replaces the convolutional decoder; latents and denormalization are
+ # identical either way.
+ self.diffusion_decoder = diffusion_decoder
# Add video processor for postprocessing
from diffusers.video_processor import VideoProcessor
@@ -37,69 +41,117 @@ class LTX2AVDecodingStage(DecodingStage):
self, server_args: ServerArgs, stage_name: str | None = None
) -> list[ComponentUse]:
stage_name = self._component_stage_name(stage_name)
- vae_dtype = resolve_precision(
- server_args, "vae", precision_attr="vae_precision"
- )
+ vae_dtype = resolve_decode_precision(server_args, "vae")
audio_vae_dtype = resolve_precision(
server_args, "audio_vae", precision_attr="audio_vae_precision"
)
- return [
- ComponentUse(stage_name, "vae", target_dtype=vae_dtype),
- ComponentUse(stage_name, "audio_vae", target_dtype=audio_vae_dtype),
- ComponentUse(stage_name, "vocoder"),
- ]
+ uses = [ComponentUse(stage_name, "vae", target_dtype=vae_dtype)]
+ if self.diffusion_decoder is not None:
+ uses.append(
+ ComponentUse(
+ stage_name,
+ "diffusion_decoder",
+ target_dtype=vae_dtype,
+ allow_prefetch=False,
+ )
+ )
+ uses.extend(
+ [
+ ComponentUse(stage_name, "audio_vae", target_dtype=audio_vae_dtype),
+ ComponentUse(stage_name, "vocoder"),
+ ]
+ )
+ return uses
@staticmethod
def _ltx2_should_externally_denorm_video_latents(server_args: ServerArgs) -> bool:
arch_config = server_args.pipeline_config.vae_config.arch_config
- return str(getattr(arch_config, "video_decoder_variant", "ltx_2")) != "ltx_2_3"
+ return str(arch_config.video_decoder_variant) != "ltx_2_3"
+
+ def _decode_with_diffusion_decoder(
+ self, decoder, latents, batch, server_args: ServerArgs
+ ):
+ """Decode with the LTX-2.5 diffusion decoder.
+
+ It is a diffusion model in its own right, so it needs a generator; the
+ request's seed keeps a decode reproducible.
+ """
+ # Untiled, every stage attends over the whole volume -- minutes at a
+ # full-length 121-frame grid.
+ decoder.use_tiling = bool(server_args.pipeline_config.diffusion_decoder_tiling)
+ generator = torch.Generator(device=latents.device).manual_seed(int(batch.seed))
+ return decoder(latents, generator=generator)
+
+ def _prepare_video_latents(self, batch: Req, module, server_args: ServerArgs):
+ latents = batch.latents.to(get_local_torch_device())
+ if self._ltx2_should_externally_denorm_video_latents(server_args):
+ std = module.latents_std.view(1, -1, 1, 1, 1).to(latents)
+ mean = module.latents_mean.view(1, -1, 1, 1, 1).to(latents)
+ latents = latents * std + mean
+ return server_args.pipeline_config.preprocess_decoding(
+ latents, server_args, vae=module
+ )
def forward(self, batch: Req, server_args: ServerArgs) -> OutputBatch:
self.load_model()
- vae_dtype = resolve_precision(
- server_args,
- "vae",
- precision_attr="vae_precision",
- )
+ vae_dtype = resolve_decode_precision(server_args, "vae")
vae_autocast_enabled = autocast_enabled(vae_dtype, server_args.disable_autocast)
- with self.use_declared_component(component_name="vae", module=self.vae) as vae:
- assert vae is not None
- self.vae = vae
- self.vae.eval()
- latents = batch.latents.to(get_local_torch_device())
- if self._ltx2_should_externally_denorm_video_latents(server_args):
- std = self.vae.latents_std.view(1, -1, 1, 1, 1).to(latents)
- mean = self.vae.latents_mean.view(1, -1, 1, 1, 1).to(latents)
- latents = latents * std + mean
- latents = server_args.pipeline_config.preprocess_decoding(
- latents, server_args, vae=self.vae
- )
+ if batch.use_diffusion_decoder:
+ if self.diffusion_decoder is None:
+ raise ValueError(
+ "use_diffusion_decoder was requested, but the decoder is not "
+ "loaded. Start the server with --load-diffusion-decoder."
+ )
+ with self.use_declared_component(
+ component_name="diffusion_decoder", module=self.diffusion_decoder
+ ) as decoder:
+ assert decoder is not None
+ decoder.eval()
+ latents = self._prepare_video_latents(batch, decoder, server_args)
+ with autocast_context(
+ dtype=vae_dtype,
+ disable_autocast=server_args.disable_autocast,
+ enabled=vae_autocast_enabled,
+ ):
+ if not vae_autocast_enabled:
+ latents = latents.to(vae_dtype)
+ decode_output = self._decode_with_diffusion_decoder(
+ decoder, latents, batch, server_args
+ )
+ else:
+ with self.use_declared_component(
+ component_name="vae", module=self.vae
+ ) as vae:
+ assert vae is not None
+ self.vae = vae
+ self.vae.eval()
+ latents = self._prepare_video_latents(batch, self.vae, server_args)
+ with autocast_context(
+ dtype=vae_dtype,
+ disable_autocast=server_args.disable_autocast,
+ enabled=vae_autocast_enabled,
+ ):
+ try:
+ if server_args.pipeline_config.vae_tiling:
+ self.vae.enable_tiling()
+ except Exception:
+ pass
+ should_cast_vae = not vae_autocast_enabled
+ if not vae_autocast_enabled:
+ latents = latents.to(vae_dtype)
+ with temporary_module_dtype(
+ self.vae, vae_dtype, enabled=should_cast_vae
+ ) as vae:
+ decode_output = vae.decode(latents)
- with autocast_context(
- dtype=vae_dtype,
- disable_autocast=server_args.disable_autocast,
- enabled=vae_autocast_enabled,
- ):
- try:
- if server_args.pipeline_config.vae_tiling:
- self.vae.enable_tiling()
- except Exception:
- pass
- should_cast_vae = not vae_autocast_enabled
- if not vae_autocast_enabled:
- latents = latents.to(vae_dtype)
- with temporary_module_dtype(
- self.vae, vae_dtype, enabled=should_cast_vae
- ) as vae:
- decode_output = vae.decode(latents)
- if isinstance(decode_output, tuple):
- video = decode_output[0]
- elif hasattr(decode_output, "sample"):
- video = decode_output.sample
- else:
- video = decode_output
+ if isinstance(decode_output, tuple):
+ video = decode_output[0]
+ elif isinstance(decode_output, torch.Tensor):
+ video = decode_output
+ else:
+ video = decode_output.sample
video = self.video_processor.postprocess_video(video, output_type="np")
output_batch = OutputBatch(
diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/ltx_2/duration.py b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/ltx_2/duration.py
new file mode 100644
index 000000000..64e0c9f44
--- /dev/null
+++ b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/ltx_2/duration.py
@@ -0,0 +1,94 @@
+# SPDX-License-Identifier: Apache-2.0
+"""Auto-duration stage for LTX-2.5.
+
+Runs between the text connectors and latent preparation, so it can rewrite
+`batch.num_frames` before any latent shape is derived from it.
+"""
+
+import torch
+
+from sglang.multimodal_gen.runtime.managers.memory_managers.component_manager import (
+ ComponentUse,
+)
+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.server_args import ServerArgs
+from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
+from sglang.multimodal_gen.runtime.utils.precision import resolve_precision
+
+logger = init_logger(__name__)
+
+
+class LTX2DurationStage(PipelineStage):
+ """Predict `num_frames` from the caption when auto-duration is requested.
+
+ Upstream expresses this by omitting `num_frames` on a pipeline that has a
+ duration head. SGLang's sampling params always carry a frame count, so the
+ request opts in explicitly via `auto_duration`.
+ """
+
+ def __init__(self, duration_head) -> None:
+ super().__init__()
+ self.duration_head = duration_head
+
+ def component_uses(
+ self, server_args: ServerArgs, stage_name: str | None = None
+ ) -> list[ComponentUse]:
+ if self.duration_head is None:
+ return []
+ dtype = resolve_precision(
+ server_args, "duration_head", precision_attr="dit_precision"
+ )
+ return [
+ ComponentUse(
+ self._component_stage_name(stage_name),
+ "duration_head",
+ target_dtype=dtype,
+ )
+ ]
+
+ def forward(self, batch: Req, server_args: ServerArgs) -> Req:
+ if not batch.auto_duration:
+ return batch
+
+ if self.duration_head is None:
+ raise ValueError(
+ "auto_duration was requested but this checkpoint has no duration "
+ "head. It ships from LTX-2.5 onward."
+ )
+
+ video_tokens = batch.prompt_embeds
+ audio_tokens = batch.audio_prompt_embeds
+ if isinstance(video_tokens, list):
+ video_tokens = video_tokens[0]
+ if isinstance(audio_tokens, list):
+ audio_tokens = audio_tokens[0]
+
+ # A CFG batch carries [negative, positive] with duplicated rows, so
+ # predict from the first positive row only.
+ with (
+ self.use_declared_component(
+ component_name="duration_head", module=self.duration_head
+ ) as duration_head,
+ torch.no_grad(),
+ ):
+ assert duration_head is not None
+ num_frames = duration_head.predict_num_frames(
+ video_tokens[:1],
+ audio_tokens[:1],
+ frame_rate=float(batch.fps),
+ temporal_compression_ratio=int(
+ server_args.pipeline_config.vae_temporal_compression
+ ),
+ min_seconds=float(batch.auto_duration_min_seconds),
+ max_seconds=float(batch.auto_duration_max_seconds),
+ )
+
+ logger.info(
+ "Auto-duration: %d frames (requested %d) @ %.2f fps",
+ num_frames,
+ int(batch.num_frames),
+ float(batch.fps),
+ )
+ batch.num_frames = int(num_frames)
+ return batch
diff --git a/python/sglang/multimodal_gen/runtime/server_args/server_args.py b/python/sglang/multimodal_gen/runtime/server_args/server_args.py
index 2675b805f..1f95b352d 100644
--- a/python/sglang/multimodal_gen/runtime/server_args/server_args.py
+++ b/python/sglang/multimodal_gen/runtime/server_args/server_args.py
@@ -290,6 +290,8 @@ class ServerArgs(DisaggServerArgsMixin):
# Component path overrides (key = model_index.json component name, value = path)
component_paths: dict[str, str] = field(default_factory=dict)
+ # Optional LTX-2.5 decoder is large enough to load only when requested.
+ load_diffusion_decoder: bool = False
# path to pre-quantized transformer weights (single .safetensors or directory).
transformer_weights_path: str | None = None
@@ -1565,6 +1567,16 @@ class ServerArgs(DisaggServerArgsMixin):
"or model_index.json. Must match a registered pipeline_name."
),
)
+ parser.add_argument(
+ "--load-diffusion-decoder",
+ action=StoreBoolean,
+ default=ServerArgs.load_diffusion_decoder,
+ help=(
+ "Load the optional LTX-2.5 diffusion decoder so requests may set "
+ "use_diffusion_decoder. Offline generate enables this automatically "
+ "when --use-diffusion-decoder is passed."
+ ),
+ )
# attention
parser.add_argument(
"--attention-backend",
diff --git a/python/sglang/multimodal_gen/runtime/utils/precision.py b/python/sglang/multimodal_gen/runtime/utils/precision.py
index 6ac31c787..28db6d4b5 100644
--- a/python/sglang/multimodal_gen/runtime/utils/precision.py
+++ b/python/sglang/multimodal_gen/runtime/utils/precision.py
@@ -55,7 +55,7 @@ def resolve_component_precision(server_args, module_name: str) -> Optional[torch
if module_name in ("audio_vae", "vocoder"):
precision_attr = "audio_vae_precision"
- elif module_name in ("vae", "video_vae"):
+ elif module_name in ("vae", "video_vae", "diffusion_decoder"):
precision_attr = "vae_precision"
elif module_name in (
"transformer",
diff --git a/python/sglang/multimodal_gen/test/unit/test_adapter_loader_offload_target.py b/python/sglang/multimodal_gen/test/unit/test_adapter_loader_offload_target.py
new file mode 100644
index 000000000..b7258a3b6
--- /dev/null
+++ b/python/sglang/multimodal_gen/test/unit/test_adapter_loader_offload_target.py
@@ -0,0 +1,60 @@
+# SPDX-License-Identifier: Apache-2.0
+"""Custom component loaders place each component by its own offload policy.
+
+`connectors` follows `dit_cpu_offload`; the duration head and standalone
+diffusion decoder stay resident by default.
+"""
+
+import unittest
+
+from sglang.multimodal_gen.runtime.loader.component_loaders.adapter_loader import (
+ AdapterLoader,
+)
+from sglang.multimodal_gen.runtime.loader.component_loaders.diffusion_decoder_loader import (
+ DiffusionDecoderLoader,
+)
+from sglang.multimodal_gen.runtime.server_args.server_args import ServerArgs
+
+
+class TestAdapterLoaderOffloadTarget(unittest.TestCase):
+ def _server_args(self, **overrides):
+ server_args = ServerArgs(model_path="x")
+ server_args.cpu_offload_components = None
+ server_args.dit_cpu_offload = False
+ server_args.vae_cpu_offload = False
+ for key, value in overrides.items():
+ setattr(server_args, key, value)
+ return server_args
+
+ def test_every_component_the_loader_serves_has_a_policy_answer(self):
+ server_args = self._server_args()
+ for component_name in AdapterLoader.component_names:
+ # Must not raise: the loader asks the policy about each of these.
+ self.assertIsInstance(
+ server_args.should_cpu_offload_component(component_name), bool
+ )
+
+ def test_dit_offload_moves_connectors_but_not_optional_modules(self):
+ server_args = self._server_args(dit_cpu_offload=True)
+ self.assertTrue(server_args.should_cpu_offload_component("connectors"))
+ self.assertFalse(server_args.should_cpu_offload_component("duration_head"))
+ self.assertFalse(server_args.should_cpu_offload_component("diffusion_decoder"))
+
+ def test_diffusion_decoder_has_a_dedicated_loader(self):
+ self.assertNotIn("diffusion_decoder", AdapterLoader.component_names)
+ self.assertEqual(DiffusionDecoderLoader.component_names, ["diffusion_decoder"])
+
+ def test_explicit_selection_reaches_the_diffusion_decoder(self):
+ server_args = self._server_args(
+ cpu_offload_components=["diffusion_decoder"],
+ )
+ self.assertTrue(server_args.should_cpu_offload_component("diffusion_decoder"))
+ self.assertFalse(server_args.should_cpu_offload_component("connectors"))
+
+ def test_vae_group_includes_the_diffusion_decoder(self):
+ server_args = self._server_args(cpu_offload_components=["vae"])
+ self.assertTrue(server_args.should_cpu_offload_component("diffusion_decoder"))
+
+
+if __name__ == "__main__":
+ unittest.main()
diff --git a/python/sglang/multimodal_gen/test/unit/test_disagg_roles.py b/python/sglang/multimodal_gen/test/unit/test_disagg_roles.py
index f4b689977..2c0b6fc5c 100644
--- a/python/sglang/multimodal_gen/test/unit/test_disagg_roles.py
+++ b/python/sglang/multimodal_gen/test/unit/test_disagg_roles.py
@@ -139,12 +139,16 @@ class TestGetModuleRole(unittest.TestCase):
self.assertEqual(get_module_role("audio_vae"), RoleType.DECODER)
self.assertEqual(get_module_role("video_vae"), RoleType.DECODER)
self.assertEqual(get_module_role("vocoder"), RoleType.DECODER)
+ self.assertEqual(get_module_role("diffusion_decoder"), RoleType.DECODER)
self.assertEqual(get_module_role("hy3dshape_vae"), RoleType.DECODER)
def test_shared_modules(self):
self.assertIsNone(get_module_role("scheduler"))
self.assertIsNone(get_module_role("hy3dshape_scheduler"))
+ def test_ltx25_optional_modules(self):
+ self.assertEqual(get_module_role("duration_head"), RoleType.ENCODER)
+
class TestFilterModulesForRole(unittest.TestCase):
WAN_MODULES = ["text_encoder", "tokenizer", "vae", "transformer", "scheduler"]
diff --git a/python/sglang/multimodal_gen/test/unit/test_ltx2_5_config.py b/python/sglang/multimodal_gen/test/unit/test_ltx2_5_config.py
new file mode 100644
index 000000000..01a2d81c7
--- /dev/null
+++ b/python/sglang/multimodal_gen/test/unit/test_ltx2_5_config.py
@@ -0,0 +1,613 @@
+# SPDX-License-Identifier: Apache-2.0
+"""LTX-2.5 config wiring.
+
+These pin the handful of places where LTX-2.5 diverges from LTX-2 and where a
+silent regression would produce wrong output rather than an error. Everything
+here is CPU/meta-device only -- no weights, no GPU.
+"""
+
+import json
+import tempfile
+import unittest
+from types import SimpleNamespace
+from unittest import mock
+
+from sglang.multimodal_gen.configs.models.adapter.ltx_2_connector import (
+ LTX2ConnectorArchConfig,
+)
+from sglang.multimodal_gen.configs.models.dits.ltx_2 import LTX2ArchConfig
+from sglang.multimodal_gen.configs.models.dits.ltx_2_5 import LTX25ArchConfig
+from sglang.multimodal_gen.configs.models.vaes.ltx_2_5_video import (
+ LTX25VideoVAEArchConfig,
+)
+from sglang.multimodal_gen.configs.models.vaes.ltx_video import LTXVideoVAEArchConfig
+from sglang.multimodal_gen.configs.models.vocoder.ltx_vocoder import LTXVocoderConfig
+from sglang.multimodal_gen.configs.pipeline_configs.ltx_2 import LTX2PipelineConfig
+from sglang.multimodal_gen.configs.pipeline_configs.ltx_2_5 import (
+ LTX25_DISTILLED_SIGMA_VALUES,
+ LTX25PipelineConfig,
+)
+
+
+class TestLTX25DiTConfig(unittest.TestCase):
+ def test_inherits_ltx23_audio_video_base(self):
+ arch = LTX25ArchConfig()
+ self.assertTrue(arch.apply_gated_attention)
+ self.assertTrue(arch.cross_attention_adaln)
+ # `use_prompt_embeddings: false` upstream -- the caption projection
+ # lives in the connector, not the DiT.
+ self.assertTrue(arch.caption_proj_before_connector)
+ self.assertEqual(arch.rope_type.value, "split")
+ self.assertTrue(arch.double_precision_rope)
+
+ def test_feed_forward_bias_is_video_only(self):
+ # LTX-2.5 checkpoints have no `ff.net.*.bias` for the video branch but
+ # do for the audio one.
+ arch = LTX25ArchConfig()
+ self.assertFalse(arch.ff_bias)
+ self.assertTrue(arch.audio_ff_bias)
+
+ def test_param_names_mapping_extends_ltx2(self):
+ # Regression: read off a class attribute pinned to LTX2ArchConfig, the
+ # LTX-2.5 renames never reached the loader.
+ arch = LTX25ArchConfig()
+ for rule in LTX2ArchConfig().param_names_mapping:
+ self.assertIn(rule, arch.param_names_mapping)
+ self.assertIn(r"^prompt_adaln\.(.*)$", arch.param_names_mapping)
+ self.assertIn(r"^audio_prompt_adaln\.(.*)$", arch.param_names_mapping)
+
+ def test_prompt_adaln_rename_round_trips(self):
+ from sglang.multimodal_gen.runtime.loader.utils import get_param_names_mapping
+
+ arch = LTX25ArchConfig()
+ forward = get_param_names_mapping(arch.param_names_mapping)
+ reverse = get_param_names_mapping(arch.reverse_param_names_mapping)
+ for key in ("prompt_adaln.linear.weight", "audio_prompt_adaln.linear.bias"):
+ mapped = forward(key)[0]
+ self.assertTrue(mapped.startswith(key.split(".")[0] + "_single."), mapped)
+ self.assertEqual(reverse(mapped)[0], key)
+
+ def test_ltx2_defaults_unchanged(self):
+ # The shared LTX-2 config must keep its original behaviour.
+ arch = LTX2ArchConfig()
+ self.assertTrue(arch.ff_bias)
+ self.assertTrue(arch.audio_ff_bias)
+ self.assertFalse(arch.use_keyframes_abs_pos_embedding)
+
+
+class TestLTX25VAEConfig(unittest.TestCase):
+ # Reversed like its sibling lists, this would give the wrong strides.
+ EXPECTED_STRIDES = {
+ "spatiotemporal": (2, 2, 2),
+ "temporal": (2, 1, 1),
+ "spatial": (1, 2, 2),
+ }
+
+ def test_upsample_type_order_is_decoder_order(self):
+ arch = LTX25VideoVAEArchConfig()
+ self.assertEqual(
+ list(arch.upsample_type),
+ ["spatiotemporal", "spatiotemporal", "temporal", "spatial"],
+ )
+
+ def test_ltx2_defaults_to_all_spatiotemporal(self):
+ # `None` must keep LTX-2 bit-identical.
+ self.assertIsNone(LTXVideoVAEArchConfig().upsample_type)
+
+ def test_decoder_builds_expected_upsampler_strides(self):
+ import torch
+
+ from sglang.multimodal_gen.runtime.models.vaes.ltx_2_vae import (
+ LTX2VideoDecoder3d,
+ )
+
+ arch = LTX25VideoVAEArchConfig()
+ with torch.device("meta"):
+ decoder = LTX2VideoDecoder3d(
+ in_channels=arch.latent_channels,
+ out_channels=arch.out_channels,
+ block_out_channels=arch.decoder_block_out_channels,
+ spatio_temporal_scaling=arch.decoder_spatio_temporal_scaling,
+ layers_per_block=arch.decoder_layers_per_block,
+ patch_size=arch.patch_size,
+ patch_size_t=arch.patch_size_t,
+ inject_noise=arch.decoder_inject_noise,
+ upsample_residual=arch.upsample_residual,
+ upsample_factor=arch.upsample_factor,
+ upsample_type=arch.upsample_type,
+ spatial_padding_mode=arch.decoder_spatial_padding_mode,
+ )
+
+ actual = [tuple(b.upsamplers[0].stride) for b in decoder.up_blocks]
+ expected = [self.EXPECTED_STRIDES[t] for t in arch.upsample_type]
+ self.assertEqual(actual, expected)
+
+
+class TestLTX25ConnectorConfig(unittest.TestCase):
+ """The connector configures itself from the checkpoint's own config.json.
+
+ Field names there are diffusers'; SGLang's module reads different ones. If
+ the derivation breaks, LTX-2.5 silently falls back to the LTX-2.0 branch
+ (one shared `text_proj_in`) and produces garbage embeddings instead of
+ failing.
+ """
+
+ LTX25_CONNECTOR_CONFIG = {
+ "caption_channels": 3840,
+ "text_proj_in_factor": 49,
+ "per_modality_projections": True,
+ "video_hidden_dim": 4096,
+ "audio_hidden_dim": 2048,
+ "video_gated_attn": True,
+ "audio_gated_attn": True,
+ "video_connector_num_layers": 8,
+ "audio_connector_num_layers": 8,
+ "audio_connector_attention_head_dim": 64,
+ }
+
+ def test_derives_per_modality_projection_dims(self):
+ arch = LTX2ConnectorArchConfig(**self.LTX25_CONNECTOR_CONFIG)
+ self.assertEqual(arch.feature_extractor_in_features, 3840 * 49)
+ self.assertEqual(arch.video_feature_extractor_out_features, 4096)
+ self.assertEqual(arch.audio_feature_extractor_out_features, 2048)
+ self.assertTrue(arch.connector_apply_gated_attention)
+
+ def test_ltx2_keeps_shared_projection(self):
+ arch = LTX2ConnectorArchConfig()
+ self.assertFalse(arch.per_modality_projections)
+ self.assertEqual(arch.feature_extractor_in_features, 0)
+ self.assertFalse(arch.connector_apply_gated_attention)
+
+ def test_diffusers_projection_names_are_mapped(self):
+ arch = LTX2ConnectorArchConfig()
+ self.assertIn(r"^video_text_proj_in\.(.*)$", arch.param_names_mapping)
+ self.assertIn(r"^audio_text_proj_in\.(.*)$", arch.param_names_mapping)
+
+
+class TestLTX25VocoderConfig(unittest.TestCase):
+ """LTX-2.5 ships `LTX2VocoderWithBWE` with a flat diffusers config, while
+ SGLang's BWE implementation expects the nested ltx-core shape."""
+
+ LTX25_VOCODER_CONFIG = {
+ "hidden_channels": 1536,
+ "upsample_factors": [5, 2, 2, 2, 2, 2],
+ "upsample_kernel_sizes": [11, 4, 4, 4, 4, 4],
+ "resnet_kernel_sizes": [3, 7, 11],
+ "act_fn": "snakebeta",
+ "input_sampling_rate": 16000,
+ "output_sampling_rate": 48000,
+ "filter_length": 512,
+ "window_length": 512,
+ "hop_length": 80,
+ "num_mel_channels": 64,
+ "bwe_hidden_channels": 512,
+ "bwe_upsample_factors": [6, 5, 2, 2, 2],
+ "bwe_upsample_kernel_sizes": [12, 11, 4, 4, 4],
+ "bwe_resnet_kernel_sizes": [3, 7, 11],
+ "bwe_act_fn": "snakebeta",
+ }
+
+ def test_builds_nested_bwe_config(self):
+ config = LTXVocoderConfig()
+ config.update_model_arch(dict(self.LTX25_VOCODER_CONFIG))
+ nested = config.arch_config.vocoder
+
+ self.assertIsNotNone(nested)
+ self.assertIn("bwe", nested)
+ self.assertEqual(nested["vocoder"]["upsample_initial_channel"], 1536)
+ self.assertEqual(nested["bwe"]["upsample_initial_channel"], 512)
+ # The base stack synthesises at the BWE's input rate, not the final one.
+ self.assertEqual(nested["bwe"]["input_sampling_rate"], 16000)
+ self.assertEqual(nested["bwe"]["output_sampling_rate"], 48000)
+ self.assertEqual(nested["bwe"]["num_mels"], 64)
+
+ def test_ltx2_stays_on_the_non_bwe_branch(self):
+ # No `bwe_upsample_factors` -> no nested config -> original code path.
+ self.assertIsNone(LTXVocoderConfig().arch_config.vocoder)
+
+ def test_vocoder_with_bwe_class_name_resolves(self):
+ from sglang.multimodal_gen.runtime.models.registry import ModelRegistry
+
+ cls, _ = ModelRegistry.resolve_model_cls("LTX2VocoderWithBWE")
+ self.assertEqual(cls.__name__, "LTX2Vocoder")
+
+
+class TestLTX25PipelineConfig(unittest.TestCase):
+ def test_pins_the_distilled_sigma_schedule(self):
+ # The distilled DiT is driven by this schedule, not by a step count.
+ config = LTX25PipelineConfig()
+ self.assertEqual(config.default_sigmas, LTX25_DISTILLED_SIGMA_VALUES)
+ self.assertEqual(len(LTX25_DISTILLED_SIGMA_VALUES), 8)
+ self.assertEqual(LTX25_DISTILLED_SIGMA_VALUES[0], 1.0)
+ self.assertTrue(
+ all(
+ a > b
+ for a, b in zip(
+ LTX25_DISTILLED_SIGMA_VALUES, LTX25_DISTILLED_SIGMA_VALUES[1:]
+ )
+ )
+ )
+
+ def test_ltx2_has_no_pinned_schedule(self):
+ self.assertIsNone(LTX2PipelineConfig().default_sigmas)
+
+ def test_stays_an_ltx2_variant(self):
+ # LTX-2.5 uses the LTX-2 linspace sigma path, so it must NOT be marked
+ # as an LTX-2.3 native variant even though it shares 2.3's architecture.
+ from sglang.multimodal_gen.configs.pipeline_configs.ltx_2 import (
+ is_ltx23_native_variant,
+ )
+
+ config = LTX25PipelineConfig()
+ self.assertFalse(is_ltx23_native_variant(config.vae_config.arch_config))
+
+ def test_registry_resolves_ltx_variants_apart(self):
+ # Not `get_model_info`: it also reads `model_index.json` from the Hub,
+ # and offline that silently resolves to the generic diffusers config.
+ from sglang.multimodal_gen.registry import _get_config_info
+
+ self.assertIs(
+ _get_config_info("Lightricks/LTX-2.5-Diffusers").pipeline_config_cls,
+ LTX25PipelineConfig,
+ )
+ self.assertIs(
+ _get_config_info("Lightricks/LTX-2").pipeline_config_cls,
+ LTX2PipelineConfig,
+ )
+ self.assertEqual(
+ _get_config_info("Lightricks/LTX-2.3").pipeline_config_cls.__name__,
+ "LTX23PipelineConfig",
+ )
+
+ def test_derived_repos_keep_the_point_releases_apart(self):
+ """Forks and local copies resolve by longest registered path stem.
+
+ Resolution tries exact match, then the longest registered path that is
+ a substring of the request, and only then the detectors. So a derived
+ repo lands on the right config as long as it keeps the registered stem
+ -- `LTX-2.5-Diffusers` is longer than `LTX-2` and wins.
+ """
+ from sglang.multimodal_gen.registry import _get_config_info
+
+ self.assertIs(
+ _get_config_info("myorg/LTX-2.5-Diffusers-fp8").pipeline_config_cls,
+ LTX25PipelineConfig,
+ )
+ self.assertIs(
+ _get_config_info("myorg/LTX-2-custom").pipeline_config_cls,
+ LTX2PipelineConfig,
+ )
+ self.assertEqual(
+ _get_config_info("myorg/LTX-2.3-tuned").pipeline_config_cls.__name__,
+ "LTX23PipelineConfig",
+ )
+
+
+class TestLTX25ImageConditioningCRF(unittest.TestCase):
+ """LTX-2.5 trained image conditioning at CRF 18; LTX-2 / 2.3 at 33.
+
+ Getting this wrong does not raise -- it just feeds the model conditioning
+ images from the wrong compression distribution.
+ """
+
+ def _crf_for_config(self, pipeline_config):
+ """CRF for an already-resolved pipeline config.
+
+ Driving this from a model path would route through `ServerArgs`, which
+ reads `model_index.json` from the Hub; offline that falls back to the
+ generic config and the assertion becomes meaningless. The resolver only
+ reads `pipeline_config.text_encoder_configs`, so hand it the config
+ directly.
+ """
+ from types import SimpleNamespace
+
+ from sglang.multimodal_gen.runtime.pipelines_core.stages.image_encoding import (
+ LTX2ImageEncodingStage,
+ )
+
+ return LTX2ImageEncodingStage._resolve_image_conditioning_crf(
+ SimpleNamespace(pipeline_config=pipeline_config)
+ )
+
+ def test_ltx_2_5_uses_crf_18(self):
+ self.assertEqual(self._crf_for_config(LTX25PipelineConfig()), 18)
+
+ def test_earlier_ltx_generations_use_crf_33(self):
+ self.assertEqual(self._crf_for_config(LTX2PipelineConfig()), 33)
+
+
+class TestLTX25DurationHead(unittest.TestCase):
+ """Frame counts must land on the VAE's causal temporal grid (8k + 1)."""
+
+ def _head(self):
+ import torch
+
+ from sglang.multimodal_gen.configs.models.adapter.ltx_2_duration_head import (
+ LTX2DurationHeadConfig,
+ )
+ from sglang.multimodal_gen.runtime.models.adapter.ltx_2_duration_head import (
+ LTX2DurationHead,
+ )
+
+ with torch.device("meta"):
+ return LTX2DurationHead(LTX2DurationHeadConfig())
+
+ def test_predicted_frames_land_on_the_temporal_grid(self):
+ from unittest import mock
+
+ import torch
+
+ head = self._head()
+ for seconds in (1.0, 2.7, 3.28125, 7.5, 19.9):
+ with mock.patch.object(
+ head, "forward", return_value=torch.tensor([seconds])
+ ):
+ n = head.predict_num_frames(
+ frame_rate=24.0, temporal_compression_ratio=8
+ )
+ self.assertEqual((n - 1) % 8, 0, f"{n} frames is off-grid for {seconds}s")
+ self.assertGreaterEqual(n, 1)
+
+ def test_prediction_is_clamped_to_bounds(self):
+ from unittest import mock
+
+ import torch
+
+ head = self._head()
+ with mock.patch.object(head, "forward", return_value=torch.tensor([100.0])):
+ n = head.predict_num_frames(
+ frame_rate=24.0, temporal_compression_ratio=8, max_seconds=5.0
+ )
+ self.assertLessEqual(n / 24.0, 5.0)
+ self.assertEqual((n - 1) % 8, 0)
+
+ def test_requires_at_least_one_modality(self):
+ with self.assertRaises(ValueError):
+ self._head()(None, None)
+
+
+class TestLTX25DiffusionDecoder(unittest.TestCase):
+ """The 2.5 diffusion decoder: config shape and the geometry it implies."""
+
+ def _config(self):
+ from sglang.multimodal_gen.configs.models.decoders.ltx_2_5_diffusion_decoder import (
+ LTX25DiffusionDecoderConfig,
+ )
+
+ return LTX25DiffusionDecoderConfig()
+
+ def test_stage_channels_match_upsample_reductions(self):
+ # Two views of the same thing; an inconsistent pair would only fail deep
+ # inside the first block.
+ arch = self._config().arch_config
+ for i, reduction in enumerate(arch.decoder_upsample_channel_reductions):
+ self.assertEqual(
+ arch.decoder_stage_channels[i + 1],
+ arch.decoder_stage_channels[i] // reduction,
+ )
+
+ def test_upsample_strides_compose_to_the_vae_ratios(self):
+ arch = self._config().arch_config
+ temporal = 1
+ spatial = 1
+ for stride_t, stride_h, _ in arch.decoder_upsample_strides:
+ temporal *= stride_t
+ spatial *= stride_h
+ self.assertEqual(temporal, arch.temporal_compression_ratio)
+ # The remaining spatial factor is the pixel patch size.
+ self.assertEqual(spatial * arch.patch_size, arch.spatial_compression_ratio)
+
+ def test_ships_as_a_single_step_x0_decoder(self):
+ arch = self._config().arch_config
+ self.assertEqual(arch.decoder_num_inference_steps, 1)
+ self.assertEqual(arch.decoder_model_output_type, "x0")
+
+ def test_builds_and_reports_expected_context_width(self):
+ import torch
+
+ from sglang.multimodal_gen.runtime.models.decoders.ltx_2_5_diffusion_decoder import (
+ LTX2VideoDiffusionDecoderModel,
+ )
+
+ config = self._config()
+ with torch.device("meta"):
+ model = LTX2VideoDiffusionDecoderModel(config)
+ self.assertEqual(
+ model.decoder.context_channels,
+ config.arch_config.decoder_stage_channels[-1],
+ )
+ # The window shifts inward at the border, so stages 1-4 carry replicated
+ # trailing frames that stage 4 crops.
+ self.assertEqual(model.decoder.trailing_pad_latent_frames, 2)
+
+ def test_timestep_embedder_is_replicated_and_checkpoint_compatible(self):
+ import torch
+ from torch import nn
+
+ from sglang.multimodal_gen.runtime.models.decoders.ltx_2_5_diffusion_decoder import (
+ LTX2VideoDiffusionDecoderModel,
+ )
+
+ with torch.device("meta"):
+ model = LTX2VideoDiffusionDecoderModel(self._config())
+ timestep_embedder = model.decoder.t_embedder.timestep_embedder
+ self.assertIsInstance(timestep_embedder.linear_1, nn.Linear)
+ self.assertIsInstance(timestep_embedder.linear_2, nn.Linear)
+ self.assertEqual(
+ tuple(timestep_embedder.linear_1.weight.shape),
+ (self._config().arch_config.decoder_t_emb_dim, 256),
+ )
+ self.assertIn(
+ "decoder.t_embedder.timestep_embedder.linear_1.weight",
+ model.state_dict(),
+ )
+
+ def test_class_name_resolves(self):
+ from sglang.multimodal_gen.runtime.models.registry import ModelRegistry
+
+ cls, _ = ModelRegistry.resolve_model_cls("LTX2VideoDiffusionDecoderModel")
+ self.assertEqual(cls.__name__, "LTX2VideoDiffusionDecoderModel")
+
+
+class TestLTX25OptionalDecoderLoading(unittest.TestCase):
+ @staticmethod
+ def _server_args(load_diffusion_decoder: bool):
+ return SimpleNamespace(
+ load_diffusion_decoder=load_diffusion_decoder,
+ model_variant=None,
+ component_paths={},
+ )
+
+ @staticmethod
+ def _write_model_index(model_path: str, *, include_decoder: bool = True):
+ model_index = {
+ "_class_name": "LTX2Pipeline",
+ "duration_head": ["ltx2", "LTX2DurationHeadModel"],
+ }
+ if include_decoder:
+ model_index["diffusion_decoder"] = [
+ "ltx2",
+ "LTX2VideoDiffusionDecoderModel",
+ ]
+ with open(f"{model_path}/model_index.json", "w") as f:
+ json.dump(model_index, f)
+
+ def test_decoder_is_not_loaded_by_default(self):
+ from sglang.multimodal_gen.runtime.pipelines.ltx_2_pipeline import LTX2Pipeline
+ from sglang.multimodal_gen.runtime.pipelines_core.lora_pipeline import (
+ LoRAPipeline,
+ )
+
+ with tempfile.TemporaryDirectory() as model_path:
+ self._write_model_index(model_path)
+ with mock.patch.object(LoRAPipeline, "__init__", return_value=None) as init:
+ LTX2Pipeline(model_path, self._server_args(False))
+ modules = init.call_args.kwargs["required_config_modules"]
+ self.assertIn("duration_head", modules)
+ self.assertNotIn("diffusion_decoder", modules)
+
+ def test_decoder_load_is_explicit_and_validated(self):
+ from sglang.multimodal_gen.runtime.pipelines.ltx_2_pipeline import LTX2Pipeline
+ from sglang.multimodal_gen.runtime.pipelines_core.lora_pipeline import (
+ LoRAPipeline,
+ )
+
+ with tempfile.TemporaryDirectory() as model_path:
+ self._write_model_index(model_path)
+ with mock.patch.object(LoRAPipeline, "__init__", return_value=None) as init:
+ LTX2Pipeline(model_path, self._server_args(True))
+ modules = init.call_args.kwargs["required_config_modules"]
+ self.assertIn("diffusion_decoder", modules)
+
+ self._write_model_index(model_path, include_decoder=False)
+ with self.assertRaisesRegex(ValueError, "does not declare"):
+ LTX2Pipeline(model_path, self._server_args(True))
+
+
+class TestLTX25LatentUpsampler(unittest.TestCase):
+ """LTX-2.5 turns the rational resampler off explicitly.
+
+ Earlier LTX configs only carry `rational_spatial_scale`, so the loader
+ inferred the resampler from its presence. LTX-2.5 states the choice, and
+ assuming True there builds a different module than the checkpoint holds.
+ """
+
+ LTX25_UPSAMPLER_CONFIG = {
+ "dims": 3,
+ "in_channels": 128,
+ "mid_channels": 1024,
+ "num_blocks_per_stage": 4,
+ "rational_spatial_scale": 2.0,
+ "spatial_upsample": True,
+ "temporal_upsample": False,
+ "use_rational_resampler": False,
+ }
+
+ def _normalize(self, raw):
+ from sglang.multimodal_gen.runtime.loader.component_loaders.upsampler_loader import (
+ _normalize_config,
+ )
+
+ return _normalize_config(raw)
+
+ def test_explicit_flag_is_honoured(self):
+ config = self._normalize(dict(self.LTX25_UPSAMPLER_CONFIG))
+ self.assertFalse(config["rational_resampler"])
+ self.assertEqual(config["spatial_scale"], 2.0)
+
+ def test_absent_flag_keeps_legacy_behaviour(self):
+ raw = {
+ k: v
+ for k, v in self.LTX25_UPSAMPLER_CONFIG.items()
+ if k != "use_rational_resampler"
+ }
+ self.assertTrue(self._normalize(raw)["rational_resampler"])
+
+ def test_flag_changes_the_module_it_builds(self):
+ # Guards the fix: the two settings are not interchangeable.
+ import torch
+
+ from sglang.multimodal_gen.runtime.models.upsampler.latent_upsampler import (
+ LatentUpsampler,
+ )
+
+ kwargs = dict(
+ in_channels=128,
+ mid_channels=1024,
+ num_blocks_per_stage=4,
+ dims=3,
+ spatial_upsample=True,
+ temporal_upsample=False,
+ spatial_scale=2.0,
+ )
+ with torch.device("meta"):
+ without = set(
+ LatentUpsampler(**kwargs, rational_resampler=False).state_dict()
+ )
+ with_rr = set(
+ LatentUpsampler(**kwargs, rational_resampler=True).state_dict()
+ )
+ self.assertNotEqual(without, with_rr)
+
+
+class TestLTX25DevVariant(unittest.TestCase):
+ """`--model-variant dev` serves `transformer_full/`, which the index omits."""
+
+ def _pipeline_cls(self):
+ from sglang.multimodal_gen.runtime.pipelines.ltx_2_pipeline import (
+ _BaseLTX2Pipeline,
+ )
+
+ return _BaseLTX2Pipeline
+
+ def _args(self, variant):
+ class _Args:
+ model_variant = variant
+ component_paths: dict = {}
+
+ return _Args()
+
+ def test_variant_aliases(self):
+ cls = self._pipeline_cls()
+ for variant in ("dev", "full", "sft", "DEV"):
+ self.assertTrue(cls._is_dev_variant(self._args(variant)), variant)
+ for variant in (None, "", "distilled"):
+ self.assertFalse(cls._is_dev_variant(self._args(variant)), variant)
+
+ def test_missing_weights_raises_a_clear_error(self):
+ cls = self._pipeline_cls()
+ with self.assertRaises(ValueError) as ctx:
+ cls._maybe_route_dev_transformer("/nonexistent/model", self._args("dev"))
+ self.assertIn("transformer_full", str(ctx.exception))
+
+ def test_explicit_component_path_wins(self):
+ cls = self._pipeline_cls()
+ args = self._args("dev")
+ args.component_paths = {"transformer": "/some/other/transformer"}
+ # Must not raise, and must not overwrite the caller's choice.
+ cls._maybe_route_dev_transformer("/nonexistent/model", args)
+ self.assertEqual(args.component_paths["transformer"], "/some/other/transformer")
+
+
+if __name__ == "__main__":
+ unittest.main()
diff --git a/python/sglang/multimodal_gen/test/unit/test_usp_ipc_a2a_guard.py b/python/sglang/multimodal_gen/test/unit/test_usp_ipc_a2a_guard.py
new file mode 100644
index 000000000..e2553c650
--- /dev/null
+++ b/python/sglang/multimodal_gen/test/unit/test_usp_ipc_a2a_guard.py
@@ -0,0 +1,54 @@
+# SPDX-License-Identifier: Apache-2.0
+"""`_ipc_input_a2a_qkv` must decline cross-attention shapes.
+
+It sizes one staging slot from `q` and reuses it for q, k and v, which only
+holds when all three share a sequence length. Cross-attention with unequal
+query and key/value lengths -- LTX-2's video-to-audio blocks, say -- has to fall
+back to the general exchange, which handles them.
+"""
+
+import unittest
+from unittest import mock
+
+import torch
+
+from sglang.multimodal_gen.runtime.layers import usp
+
+
+class TestIpcInputA2AQkvGuard(unittest.TestCase):
+ def _call(self, q, k, v):
+ # Pretend ulysses degree 2 so the guard, not the degree check, decides.
+ with mock.patch.object(usp, "get_ulysses_parallel_world_size", lambda: 2):
+ return usp._ipc_input_a2a_qkv(q, k, v)
+
+ def test_declines_when_kv_length_differs(self):
+ q = torch.zeros(1, 1530, 8, 64)
+ kv = torch.zeros(1, 43, 8, 64)
+ self.assertIsNone(self._call(q, kv, kv))
+
+ def test_declines_when_only_v_differs(self):
+ q = torch.zeros(1, 128, 8, 64)
+ k = torch.zeros(1, 128, 8, 64)
+ v = torch.zeros(1, 64, 8, 64)
+ self.assertIsNone(self._call(q, k, v))
+
+ def test_self_attention_shapes_reach_the_ipc_path(self):
+ # Without a real IPC group this returns None either way, so patch the
+ # group lookup to prove the guard is not what rejected it.
+ q = torch.zeros(1, 128, 8, 64)
+ group = mock.MagicMock(return_value=None)
+ with mock.patch.object(
+ usp, "get_ulysses_parallel_world_size", lambda: 2
+ ), mock.patch.object(usp, "_ipc_ready_group", group):
+ self.assertIsNone(usp._ipc_input_a2a_qkv(q, q.clone(), q.clone()))
+ # Reached the group lookup, so the shape guard did not reject it.
+ self.assertEqual(group.call_count, 1)
+
+ def test_degree_other_than_two_declines(self):
+ q = torch.zeros(1, 128, 8, 64)
+ with mock.patch.object(usp, "get_ulysses_parallel_world_size", lambda: 4):
+ self.assertIsNone(usp._ipc_input_a2a_qkv(q, q.clone(), q.clone()))
+
+
+if __name__ == "__main__":
+ unittest.main()