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 ( + + ); + })} +
+
+ ))} + +
+
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()