chore: import upstream snapshot with attribution
PR Test (NPU) / check-changes (push) Has been cancelled
PR Test (NPU) / pr-gate (push) Has been cancelled
PR Test (NPU) / set-image-config (push) Has been cancelled
PR Test (NPU) / stage-b-test-1-npu-a2 (0) (push) Has been cancelled
PR Test (NPU) / stage-b-test-1-npu-a2 (1) (push) Has been cancelled
PR Test (NPU) / stage-b-test-2-npu-a2 (0) (push) Has been cancelled
PR Test (NPU) / stage-b-test-2-npu-a2 (1) (push) Has been cancelled
PR Test (NPU) / stage-b-test-4-npu-a3 (push) Has been cancelled
PR Test (NPU) / stage-b-test-16-npu-a3 (push) Has been cancelled
PR Test (NPU) / multimodal-gen-test-1-npu-a3 (push) Has been cancelled
PR Test (NPU) / multimodal-gen-test-2-npu-a3 (push) Has been cancelled
PR Test (Arm64) / pr-gate (push) Has been cancelled
PR Test (Arm64) / check-changes (push) Has been cancelled
PR Test (Arm64) / build-test (push) Has been cancelled
PR Test (sgl-router) / gate (push) Has been cancelled
PR Test (sgl-router) / tier-1 — lint (push) Has been cancelled
PR Test (sgl-router) / tier-2 — build + test (push) Has been cancelled
PR Test (sgl-router) / tier-3 — docker (placeholder) (push) Has been cancelled
PR Test (sgl-router) / tier-3 — k8s integration (push) Has been cancelled
PR Test (sgl-router) / tier-3 — e2e (push) Has been cancelled
PR Test (sgl-router) / finish (push) Has been cancelled
PR Test (NPU) / single-node-poc (map[name:qwen3_6_27b_w8a8_1p_in64k_out1k_50ms runner:linux-aarch64-a3-2 test_case:test/registered/ascend/performance/qwen3_6_27b/test_npu_qwen3_6_27b_w8a8_1p_in64k_out1k_50ms.py test_type:perf]) (push) Has been cancelled
PR Test (NPU) / pr-test-npu-finish (push) Has been cancelled
PR Test (Xeon) / pr-gate (push) Has been cancelled
PR Test (Xeon) / check-changes (push) Has been cancelled
PR Test (Xeon) / build-test (, xeon-gnr, base-b-test-cpu) (push) Has been cancelled
PR Test (XPU) / check-changes (push) Has been cancelled
PR Test (XPU) / pr-gate (push) Has been cancelled
PR Test (XPU) / stage-a-test-1-gpu-xpu (push) Has been cancelled
PR Test (XPU) / wait-for-stage-a (push) Has been cancelled
PR Test (XPU) / stage-b-test-1-gpu-xpu (push) Has been cancelled
PR Test (XPU) / finish (push) Has been cancelled
CI Model Inventory / build-inventory (push) Has been cancelled
Lint / lint (push) Has been cancelled
PR Benchmark (SMG Components) / Benchmark Compilation Check (push) Has been cancelled
PR Benchmark (SMG Components) / Benchmark - Manual Policy (push) Has been cancelled
PR Benchmark (SMG Components) / Benchmark - Request Processing (push) Has been cancelled
PR Benchmark (SMG Components) / Benchmark Summary (push) Has been cancelled
PR Test (SMG) / build-wheel (push) Has been cancelled
Release SGLang Model Gateway to PyPI / build on windows (x86_64 - auto) (push) Has been cancelled
Release SGLang Model Gateway to PyPI / build on macos (x86_64 - auto) (push) Has been cancelled
PR Test (SMG) / python-unit-tests (push) Has been cancelled
PR Test (SMG) / unit-tests (push) Has been cancelled
PR Test (SMG) / benchmarks (push) Has been cancelled
PR Test (SMG) / chat-completions (push) Has been cancelled
PR Test (SMG) / chat-completions-4gpu (push) Has been cancelled
PR Test (SMG) / e2e (push) Has been cancelled
PR Test (SMG) / docker-build-test (push) Has been cancelled
PR Test (SMG) / k8s-integration (push) Has been cancelled
PR Test (SMG) / finish (push) Has been cancelled
PR Test (SMG) / summarize-benchmarks (push) Has been cancelled
Release SGLang Model Gateway Docker Image / publish (push) Has been cancelled
Release SGLang Model Gateway to PyPI / build on macos (aarch64 - auto) (push) Has been cancelled
Release SGLang Model Gateway to PyPI / build on linux (aarch64 - auto) (push) Has been cancelled
Release SGLang Model Gateway to PyPI / build on linux (x86_64 - auto) (push) Has been cancelled
Release SGLang Model Gateway to PyPI / build on linux (aarch64 - musllinux_1_1) (push) Has been cancelled
Release SGLang Model Gateway to PyPI / build on linux (x86_64 - musllinux_1_1) (push) Has been cancelled
Release SGLang Model Gateway to PyPI / Build SDist (push) Has been cancelled
Release SGLang Model Gateway to PyPI / Upload to PyPI (push) Has been cancelled
Release SGLang Kernels / build-cu129-matrix (aarch64, 12.9, 3.10, arm-kernel-build-node) (push) Has been cancelled
Release SGLang Kernels / build-cu129-matrix (x86_64, 12.9, 3.10, x64-kernel-build-node) (push) Has been cancelled
Release SGLang Kernels / release-cu129 (push) Has been cancelled
Release SGLang Kernels / build-cu130-matrix (aarch64, 13.0, 3.10, arm-kernel-build-node) (push) Has been cancelled
Release SGLang Kernels / build-cu130-matrix (x86_64, 13.0, 3.10, x64-kernel-build-node) (push) Has been cancelled
Release SGLang Kernels / release-cu130 (push) Has been cancelled
Release SGLang Kernels / build-rocm-matrix (3.10, 700) (push) Has been cancelled
Release SGLang Kernels / build-rocm-matrix (3.10, 720) (push) Has been cancelled
Release SGLang Kernels / release-rocm700 (push) Has been cancelled
Release SGLang Kernels / release-rocm720 (push) Has been cancelled
Release SGLang Kernels / build-musa43 (43, 3.10) (push) Has been cancelled
Release SGLang Kernels / release-musa43 (push) Has been cancelled

This commit is contained in:
wehub-resource-sync
2026-07-13 12:38:16 +08:00
commit 94057c3d3e
7152 changed files with 2120455 additions and 0 deletions
@@ -0,0 +1 @@
# Copied and adapted from: https://github.com/hao-ai-lab/FastVideo
@@ -0,0 +1,689 @@
# Copied and adapted from: https://github.com/hao-ai-lab/FastVideo
# SPDX-License-Identifier: Apache-2.0
import os
import re
from typing import Any, Generator
import torch
from sglang.multimodal_gen.configs.models.dits.flux import FluxConfig
from sglang.multimodal_gen.runtime.distributed import get_local_torch_device
from sglang.multimodal_gen.runtime.models.registry import ModelRegistry
from sglang.multimodal_gen.runtime.models.schedulers.scheduling_comfyui_passthrough import (
ComfyUIPassThroughScheduler,
)
from sglang.multimodal_gen.runtime.pipelines_core import LoRAPipeline
from sglang.multimodal_gen.runtime.pipelines_core.composed_pipeline_base import (
ComposedPipelineBase,
)
from sglang.multimodal_gen.runtime.pipelines_core.stages import (
ComfyUILatentPreparationStage,
DenoisingStage,
)
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 ComfyUIFluxPipeline(LoRAPipeline, ComposedPipelineBase):
"""
Simplified pipeline for ComfyUI integration with only denoising stage.
This pipeline requires pre-processed inputs:
- prompt_embeds: Pre-encoded text embeddings (list of tensors)
- negative_prompt_embeds: Pre-encoded negative prompt embeddings (if using CFG)
- latents: Optional initial noise latents (will be generated if not provided)
Usage:
generator = DiffGenerator.from_pretrained(
model_path="path/to/model",
pipeline_class_name="ComfyUIFluxPipeline",
device="cuda",
)
"""
pipeline_name = "ComfyUIFluxPipeline"
# Configuration classes for safetensors files without model_index.json
from sglang.multimodal_gen.configs.pipeline_configs.flux import FluxPipelineConfig
from sglang.multimodal_gen.configs.sample.flux import FluxSamplingParams
pipeline_config_cls = FluxPipelineConfig
sampling_params_cls = FluxSamplingParams
_required_config_modules = [
"transformer",
"scheduler",
]
def initialize_pipeline(self, server_args: ServerArgs):
"""
Initialize the pipeline with ComfyUI pass-through scheduler.
This scheduler does not modify latents, allowing ComfyUI to handle denoising.
"""
self.modules["scheduler"] = ComfyUIPassThroughScheduler(
num_train_timesteps=1000
)
if hasattr(server_args.pipeline_config, "vae_config"):
vae_config = server_args.pipeline_config.vae_config
if hasattr(vae_config, "post_init") and not hasattr(
vae_config, "_post_init_called"
):
vae_config.post_init()
logger.info(
"Called vae_config.post_init() to set spatial_compression_ratio. "
f"spatial_compression_ratio={vae_config.arch_config.spatial_compression_ratio}"
)
def load_modules(
self,
server_args: ServerArgs,
loaded_modules: dict[str, torch.nn.Module] | None = None,
) -> dict[str, Any]:
"""
Load modules for ComfyUIFluxPipeline.
If model_path is a safetensors file, load transformer directly from it
without requiring model_index.json. Otherwise, fall back to default loading.
"""
if os.path.isfile(self.model_path) and self.model_path.endswith(".safetensors"):
logger.info(
"Detected safetensors file, loading transformer directly from: %s",
self.model_path,
)
return self._load_transformer_from_safetensors(server_args, loaded_modules)
else:
logger.info(
"Model path is a directory, using default loading method: %s",
self.model_path,
)
return super().load_modules(server_args, loaded_modules)
def _load_and_convert_weights_from_safetensors(
self,
model_cls: type,
dit_config: FluxConfig,
hf_config: dict,
safetensors_list: list[str],
updated_mapping: dict,
qkv_size: int,
mlp_hidden_dim: int,
has_guidance_embeds: bool,
default_dtype: torch.dtype,
) -> tuple[torch.nn.Module, dict]:
"""
Load and convert weights from safetensors file, then load them into the model.
"""
from sglang.multimodal_gen.runtime.loader.utils import (
get_param_names_mapping,
set_default_torch_dtype,
)
from sglang.multimodal_gen.runtime.loader.weight_utils import (
safetensors_weights_iterator,
)
logger.info(
"Converting ComfyUI Flux weights to SGLang format and loading model..."
)
# Create model on target device
device = get_local_torch_device()
with set_default_torch_dtype(default_dtype):
model = model_cls(**{"config": dit_config, "hf_config": hf_config})
model = model.to(device)
# Verify model has guidance_embedder if config says it should
has_guidance_embedder = hasattr(model.time_text_embed, "guidance_embedder")
if has_guidance_embeds and not has_guidance_embedder:
logger.warning(
"Config has guidance_embeds=True but model doesn't have guidance_embedder. "
"This may indicate a configuration mismatch."
)
elif not has_guidance_embeds and has_guidance_embedder:
logger.warning(
"Config has guidance_embeds=False but model has guidance_embedder. "
"This may indicate a configuration mismatch."
)
# Note: guidance_in mappings are already included in comfyui_flux_mappings above.
# If model doesn't support guidance embeddings, the weights will be filtered out
# in _convert_comfyui_weights() based on has_guidance_embeds flag.
param_names_mapping_fn = get_param_names_mapping(updated_mapping)
weight_iterator = safetensors_weights_iterator(safetensors_list)
converted_weights = self._convert_comfyui_weights(
weight_iterator=weight_iterator,
qkv_size=qkv_size,
mlp_hidden_dim=mlp_hidden_dim,
has_guidance_embeds=has_guidance_embeds,
)
model_state_dict = model.state_dict()
missing_keys = set(model_state_dict.keys())
unexpected_keys = []
loaded_count = 0
reverse_param_names_mapping = {}
# Handle merged parameters (collect all parts before merging)
from collections import defaultdict
to_merge_params = defaultdict(dict)
# Process weights incrementally: load immediately after conversion
for source_name, tensor in converted_weights:
target_name, merge_index, num_params_to_merge = param_names_mapping_fn(
source_name
)
reverse_param_names_mapping[target_name] = (
source_name,
merge_index,
num_params_to_merge,
)
if merge_index is not None:
# Collect parts for merging
to_merge_params[target_name][merge_index] = tensor
if len(to_merge_params[target_name]) == num_params_to_merge:
# All parts collected, merge them
sorted_tensors = [
to_merge_params[target_name][i]
for i in range(num_params_to_merge)
]
merged_tensor = torch.cat(sorted_tensors, dim=0)
# Load immediately after merging
if target_name in model_state_dict:
param = model_state_dict[target_name]
loaded_tensor = merged_tensor.to(
device=param.device, dtype=param.dtype
)
param.data.copy_(loaded_tensor)
missing_keys.discard(target_name)
loaded_count += 1
del merged_tensor, loaded_tensor
else:
unexpected_keys.append(target_name)
# Clear merged parts
del to_merge_params[target_name]
for t in sorted_tensors:
del t
else:
# Direct mapping, load immediately
if target_name in model_state_dict:
param = model_state_dict[target_name]
# Check shape compatibility
if tensor.shape != param.shape:
logger.warning(
f"Shape mismatch for {target_name}: "
f"loaded {tensor.shape} vs model {param.shape}, skipping. "
f"Source: {source_name}"
)
unexpected_keys.append(target_name)
del tensor
continue
# Debug logging for norm_out.linear to verify mapping
if (
"norm_out.linear" in target_name
or "final_layer.adaLN_modulation" in source_name
):
logger.info(
f"Loading norm_out.linear: {source_name} -> {target_name}, "
f"shape: {tensor.shape}"
)
loaded_tensor = tensor.to(device=param.device, dtype=param.dtype)
param.data.copy_(loaded_tensor)
missing_keys.discard(target_name)
loaded_count += 1
del tensor, loaded_tensor
else:
# Debug logging for unmapped parameters
if "norm_out.linear" in target_name:
logger.warning(
f"norm_out.linear parameter {target_name} not found in model state_dict. "
f"Source: {source_name}"
)
unexpected_keys.append(target_name)
optional_missing_keys = []
required_missing_keys = []
for key in missing_keys:
if key.endswith(".bias"):
# Check if corresponding weight exists (if weight exists but bias doesn't, it's optional)
weight_key = key.replace(".bias", ".weight")
if weight_key not in missing_keys:
optional_missing_keys.append(key)
else:
required_missing_keys.append(key)
else:
required_missing_keys.append(key)
if required_missing_keys:
logger.warning(
f"Required missing keys (first 10): {required_missing_keys[:10]}..."
)
if optional_missing_keys:
logger.info(
f"Optional missing keys (bias parameters, {len(optional_missing_keys)} total): "
f"These will use default values (zeros)"
)
if unexpected_keys:
logger.warning(f"Unexpected keys (first 10): {unexpected_keys[:10]}...")
logger.info(f"Successfully loaded {loaded_count} weight tensors")
return model, reverse_param_names_mapping
def _convert_comfyui_weights(
self,
weight_iterator: Generator[tuple[str, torch.Tensor], None, None],
qkv_size: int,
mlp_hidden_dim: int,
has_guidance_embeds: bool,
) -> Generator[tuple[str, torch.Tensor], None, None]:
"""
Convert ComfyUI Flux weights to SGLang format.
Splits fused qkv weights into to_q/to_k/to_v plus proj_mlp.
Filters out guidance_in weights if model doesn't support guidance embeddings.
Handles scale/shift order difference between ComfyUI and AdaLayerNormContinuous.
"""
for name, tensor in weight_iterator:
if not has_guidance_embeds and name.startswith("guidance_in."):
logger.debug(
f"Skipping {name} (model doesn't support guidance embeddings)"
)
continue
# Split fused qkv in double blocks into separate q/k/v projections
match = re.match(
r"double_blocks\.(\d+)\.(img_attn|txt_attn)\.qkv\.(weight|bias)$", name
)
if match:
block_idx, attn_type, param_type = match.groups()
hidden_size = qkv_size // 3
if tensor.shape[0] < 3 * hidden_size:
logger.warning(
f"{name} shape {tensor.shape} smaller than expected qkv size {3 * hidden_size}, skipping"
)
continue
if param_type == "bias":
q_tensor = tensor[:hidden_size]
k_tensor = tensor[hidden_size : 2 * hidden_size]
v_tensor = tensor[2 * hidden_size : 3 * hidden_size]
else:
q_tensor = tensor[:hidden_size, :]
k_tensor = tensor[hidden_size : 2 * hidden_size, :]
v_tensor = tensor[2 * hidden_size : 3 * hidden_size, :]
target_prefix = f"transformer_blocks.{block_idx}.attn"
if attn_type == "img_attn":
yield f"{target_prefix}.to_q.{param_type}", q_tensor
yield f"{target_prefix}.to_k.{param_type}", k_tensor
yield f"{target_prefix}.to_v.{param_type}", v_tensor
else:
# txt_attn corresponds to encoder projections
yield f"{target_prefix}.add_q_proj.{param_type}", q_tensor
yield f"{target_prefix}.add_k_proj.{param_type}", k_tensor
yield f"{target_prefix}.add_v_proj.{param_type}", v_tensor
continue
match = re.match(r"single_blocks\.(\d+)\.linear1\.(weight|bias)$", name)
if match:
block_idx, param_type = match.groups()
expected_size = qkv_size + mlp_hidden_dim
if tensor.shape[0] < expected_size:
logger.warning(
f"linear1.{param_type} shape {tensor.shape} doesn't match "
f"expected size {expected_size}, skipping"
)
continue
# Split tensor
qkv_tensor = (
tensor[:qkv_size] if param_type == "bias" else tensor[:qkv_size, :]
)
mlp_tensor = (
tensor[qkv_size:] if param_type == "bias" else tensor[qkv_size:, :]
)
# Split qkv into q/k/v for single blocks
hidden_size = qkv_size // 3
if param_type == "bias":
q_tensor = qkv_tensor[:hidden_size]
k_tensor = qkv_tensor[hidden_size : 2 * hidden_size]
v_tensor = qkv_tensor[2 * hidden_size : 3 * hidden_size]
else:
q_tensor = qkv_tensor[:hidden_size, :]
k_tensor = qkv_tensor[hidden_size : 2 * hidden_size, :]
v_tensor = qkv_tensor[2 * hidden_size : 3 * hidden_size, :]
yield f"single_transformer_blocks.{block_idx}.attn.to_q.{param_type}", q_tensor
yield f"single_transformer_blocks.{block_idx}.attn.to_k.{param_type}", k_tensor
yield f"single_transformer_blocks.{block_idx}.attn.to_v.{param_type}", v_tensor
yield f"single_transformer_blocks.{block_idx}.proj_mlp.{param_type}", mlp_tensor
elif name == "final_layer.adaLN_modulation.1.weight":
# ComfyUI: output order is [shift, scale]
# AdaLayerNormContinuous: expects [scale, shift]
# Need to swap the first half and second half of the weight matrix
# Weight shape: (2 * hidden_size, hidden_size)
# Split into two halves and swap them
half_size = tensor.shape[0] // 2
shift_weights = tensor[:half_size, :]
scale_weights = tensor[half_size:, :]
# Swap: put scale first, then shift
swapped_tensor = torch.cat([scale_weights, shift_weights], dim=0)
logger.info(
f"Swapped scale/shift order for {name}: "
f"shape {tensor.shape} -> {swapped_tensor.shape}"
)
yield name, swapped_tensor
elif name == "final_layer.adaLN_modulation.1.bias":
# Same swap for bias: (2 * hidden_size,)
half_size = tensor.shape[0] // 2
shift_bias = tensor[:half_size]
scale_bias = tensor[half_size:]
swapped_tensor = torch.cat([scale_bias, shift_bias], dim=0)
logger.info(
f"Swapped scale/shift order for {name}: "
f"shape {tensor.shape} -> {swapped_tensor.shape}"
)
yield name, swapped_tensor
else:
# Other weights pass through (handled by param_names_mapping)
yield name, tensor
def _load_transformer_from_safetensors(
self,
server_args: ServerArgs,
loaded_modules: dict[str, torch.nn.Module] | None = None,
) -> dict[str, Any]:
"""
Load transformer directly from safetensors file without model_index.json.
"""
if loaded_modules is not None and "transformer" in loaded_modules:
logger.info("Using provided transformer module")
components = {
"transformer": loaded_modules["transformer"],
"scheduler": self.modules.get("scheduler"),
}
return components
if hasattr(server_args.pipeline_config, "dit_config"):
dit_config = server_args.pipeline_config.dit_config
if not isinstance(dit_config, FluxConfig):
logger.warning("dit_config is not FluxConfig, creating new FluxConfig")
dit_config = FluxConfig()
server_args.pipeline_config.dit_config = dit_config
else:
logger.info("Creating default FluxConfig")
dit_config = FluxConfig()
server_args.pipeline_config.dit_config = dit_config
# Set guidance_embeds to True for ComfyUI Flux models
dit_config.arch_config.guidance_embeds = True
logger.info("Set guidance_embeds=True for ComfyUI Flux model")
if dit_config.arch_config.param_names_mapping is None:
dit_config.arch_config.param_names_mapping = {}
# ComfyUI Flux uses different parameter names than SGLang Flux
# Key differences:
# - ComfyUI: single_blocks.{i}.linear1 (fused QKV + MLP input)
# - SGLang: single_transformer_blocks.{i}.attn.to_qkv + proj_mlp (separate)
# - ComfyUI: single_blocks.{i}.linear2
# - SGLang: single_transformer_blocks.{i}.proj_out
# - ComfyUI: double_blocks.{i}.img_attn.qkv / txt_attn.qkv
# - SGLang: transformer_blocks.{i}.attn.to_qkv / attn.to_added_qkv
# Note: For fused layers like linear1, we need custom weight splitting logic
# which will be handled in the weight conversion function below
comfyui_flux_mappings = {
# Double stream blocks - attention layers
r"double_blocks\.(\d+)\.img_attn\.qkv\.(weight|bias)$": (
r"transformer_blocks.\1.attn.to_qkv.\2",
None,
None,
),
r"double_blocks\.(\d+)\.txt_attn\.qkv\.(weight|bias)$": (
r"transformer_blocks.\1.attn.to_added_qkv.\2",
None,
None,
),
r"double_blocks\.(\d+)\.img_attn\.proj\.(weight|bias)$": (
r"transformer_blocks.\1.attn.to_out.0.\2",
None,
None,
),
r"double_blocks\.(\d+)\.txt_attn\.proj\.(weight|bias)$": (
r"transformer_blocks.\1.attn.to_add_out.\2",
None,
None,
),
r"double_blocks\.(\d+)\.img_attn\.norm\.query_norm\.scale$": (
r"transformer_blocks.\1.attn.norm_q.weight",
None,
None,
),
r"double_blocks\.(\d+)\.img_attn\.norm\.key_norm\.scale$": (
r"transformer_blocks.\1.attn.norm_k.weight",
None,
None,
),
r"double_blocks\.(\d+)\.txt_attn\.norm\.query_norm\.scale$": (
r"transformer_blocks.\1.attn.norm_added_q.weight",
None,
None,
),
r"double_blocks\.(\d+)\.txt_attn\.norm\.key_norm\.scale$": (
r"transformer_blocks.\1.attn.norm_added_k.weight",
None,
None,
),
# Double stream blocks - MLP layers (map to net structure)
r"double_blocks\.(\d+)\.img_mlp\.0\.(weight|bias)$": (
r"transformer_blocks.\1.ff.net.0.proj.\2",
None,
None,
),
r"double_blocks\.(\d+)\.img_mlp\.2\.(weight|bias)$": (
r"transformer_blocks.\1.ff.net.2.\2",
None,
None,
),
r"double_blocks\.(\d+)\.txt_mlp\.0\.(weight|bias)$": (
r"transformer_blocks.\1.ff_context.net.0.proj.\2",
None,
None,
),
r"double_blocks\.(\d+)\.txt_mlp\.2\.(weight|bias)$": (
r"transformer_blocks.\1.ff_context.net.2.\2",
None,
None,
),
# Double stream blocks - modulation layers
r"double_blocks\.(\d+)\.img_mod\.lin\.(weight|bias)$": (
r"transformer_blocks.\1.norm1.linear.\2",
None,
None,
),
r"double_blocks\.(\d+)\.txt_mod\.lin\.(weight|bias)$": (
r"transformer_blocks.\1.norm1_context.linear.\2",
None,
None,
),
# Single stream blocks - linear2 maps to proj_out
r"single_blocks\.(\d+)\.linear2\.(weight|bias)$": (
r"single_transformer_blocks.\1.proj_out.\2",
None,
None,
),
# Single stream blocks - norm layers (scale -> weight)
r"single_blocks\.(\d+)\.norm\.query_norm\.scale$": (
r"single_transformer_blocks.\1.attn.norm_q.weight",
None,
None,
),
r"single_blocks\.(\d+)\.norm\.key_norm\.scale$": (
r"single_transformer_blocks.\1.attn.norm_k.weight",
None,
None,
),
# Single stream blocks - modulation (maps to norm.linear)
r"single_blocks\.(\d+)\.modulation\.lin\.(weight|bias)$": (
r"single_transformer_blocks.\1.norm.linear.\2",
None,
None,
),
# Time and guidance embeddings
r"^time_in\.in_layer\.(weight|bias)$": (
r"time_text_embed.timestep_embedder.linear_1.\1",
None,
None,
),
r"^time_in\.out_layer\.(weight|bias)$": (
r"time_text_embed.timestep_embedder.linear_2.\1",
None,
None,
),
r"^txt_in\.(weight|bias)$": (r"context_embedder.\1", None, None),
r"^vector_in\.in_layer\.(weight|bias)$": (
r"time_text_embed.text_embedder.linear_1.\1",
None,
None,
),
r"^vector_in\.out_layer\.(weight|bias)$": (
r"time_text_embed.text_embedder.linear_2.\1",
None,
None,
),
# Final layer mappings
r"^final_layer\.linear\.(weight|bias)$": (r"proj_out.\1", None, None),
r"^final_layer\.norm_final\.(weight|bias)$": (r"norm_out.\1", None, None),
r"^final_layer\.adaLN_modulation\.1\.(weight|bias)$": (
r"norm_out.linear.\1",
None,
None,
),
# Image input embedding
r"^img_in\.(weight|bias)$": (r"x_embedder.\1", None, None),
# Guidance embeddings (if model supports guidance)
r"^guidance_in\.in_layer\.(weight|bias)$": (
r"time_text_embed.guidance_embedder.linear_1.\1",
None,
None,
),
r"^guidance_in\.out_layer\.(weight|bias)$": (
r"time_text_embed.guidance_embedder.linear_2.\1",
None,
None,
),
}
# Merge ComfyUI mappings with existing mappings (ComfyUI mappings take precedence)
updated_mapping = {
**dit_config.arch_config.param_names_mapping,
**comfyui_flux_mappings,
}
dit_config.arch_config.param_names_mapping = updated_mapping
logger.info(
"Added ComfyUI weight name mappings for Flux model. "
f"Total mappings: {len(updated_mapping)}"
)
cls_name = "FluxTransformer2DModel"
model_cls, _ = ModelRegistry.resolve_model_cls(cls_name)
logger.info("Resolved transformer class: %s", cls_name)
original_mapping = None
if comfyui_flux_mappings:
original_mapping = model_cls.param_names_mapping
model_cls.param_names_mapping = updated_mapping
logger.info(
"Temporarily updated model class param_names_mapping with ComfyUI mappings. "
f"Total mappings: {len(updated_mapping)}"
)
safetensors_list = [self.model_path]
logger.info("Loading weights from: %s", safetensors_list)
default_dtype = resolve_precision(
server_args, "dit", precision_attr="dit_precision"
)
server_args.model_paths["transformer"] = os.path.dirname(self.model_path) or "."
hf_config = {}
hidden_size = (
dit_config.arch_config.num_attention_heads
* dit_config.arch_config.attention_head_dim
)
mlp_ratio = getattr(dit_config.arch_config, "mlp_ratio", 4.0)
mlp_hidden_dim = int(hidden_size * mlp_ratio)
qkv_size = 3 * hidden_size
has_guidance_embeds = True
# Load and convert weights from safetensors file
model, reverse_param_names_mapping = (
self._load_and_convert_weights_from_safetensors(
model_cls=model_cls,
dit_config=dit_config,
hf_config=hf_config,
safetensors_list=safetensors_list,
updated_mapping=updated_mapping,
qkv_size=qkv_size,
mlp_hidden_dim=mlp_hidden_dim,
has_guidance_embeds=has_guidance_embeds,
default_dtype=default_dtype,
)
)
model = model.eval()
for param in model.parameters():
param.requires_grad = False
model.reverse_param_names_mapping = reverse_param_names_mapping
if original_mapping is not None:
model_cls.param_names_mapping = original_mapping
total_params = sum(p.numel() for p in model.parameters())
logger.info("Loaded transformer with %.2fB parameters", total_params / 1e9)
components = {
"transformer": model,
"scheduler": self.modules.get("scheduler"),
}
logger.info("Successfully loaded modules: %s", list(components.keys()))
return components
def create_pipeline_stages(self, server_args: ServerArgs):
logger.info(
"ComfyUIFluxPipeline.create_pipeline_stages() called - creating latent_preparation_stage and denoising_stage"
)
self.add_stages(
[
ComfyUILatentPreparationStage(
scheduler=self.get_module("scheduler"),
transformer=self.get_module("transformer"),
),
DenoisingStage(
transformer=self.get_module("transformer"),
scheduler=self.get_module("scheduler"),
),
]
)
logger.info(
f"ComfyUIFluxPipeline stages created: {list(self._stage_name_mapping.keys())}"
)
EntryClass = ComfyUIFluxPipeline
@@ -0,0 +1,353 @@
# SPDX-License-Identifier: Apache-2.0
import os
from itertools import chain
from typing import Any
import torch
from torch.distributed import init_device_mesh
from torch.distributed.fsdp import MixedPrecisionPolicy
from sglang.multimodal_gen.configs.models.dits.qwenimage import QwenImageDitConfig
from sglang.multimodal_gen.runtime.distributed import get_local_torch_device
from sglang.multimodal_gen.runtime.loader.fsdp_load import (
load_model_from_full_model_state_dict,
shard_model,
)
from sglang.multimodal_gen.runtime.loader.utils import (
get_param_names_mapping,
set_default_torch_dtype,
)
from sglang.multimodal_gen.runtime.loader.weight_utils import (
safetensors_weights_iterator,
)
from sglang.multimodal_gen.runtime.models.registry import ModelRegistry
from sglang.multimodal_gen.runtime.models.schedulers.scheduling_comfyui_passthrough import (
ComfyUIPassThroughScheduler,
)
from sglang.multimodal_gen.runtime.pipelines_core import LoRAPipeline
from sglang.multimodal_gen.runtime.pipelines_core.composed_pipeline_base import (
ComposedPipelineBase,
)
from sglang.multimodal_gen.runtime.pipelines_core.stages import (
ComfyUILatentPreparationStage,
DenoisingStage,
)
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
from sglang.multimodal_gen.utils import set_mixed_precision_policy
logger = init_logger(__name__)
class ComfyUIQwenImagePipelineBase(LoRAPipeline, ComposedPipelineBase):
"""
Base pipeline for ComfyUI QwenImage integration with only denoising stage.
This pipeline requires pre-processed inputs:
- prompt_embeds: Pre-encoded text embeddings (list of tensors)
- latents: Pre-processed image latents in sequence format [B, S, D]
Usage:
generator = DiffGenerator.from_pretrained(
model_path="path/to/model",
pipeline_class_name="ComfyUIQwenImagePipeline",
device="cuda",
)
"""
# Subclasses should override this
zero_cond_t: bool = False
pipeline_name = "ComfyUIQwenImagePipeline"
_required_config_modules = [
"transformer",
"scheduler",
]
def initialize_pipeline(self, server_args: ServerArgs):
"""
Initialize the pipeline with ComfyUI pass-through scheduler.
This scheduler does not modify latents, allowing ComfyUI to handle denoising.
"""
self.modules["scheduler"] = ComfyUIPassThroughScheduler(
num_train_timesteps=1000
)
# Ensure VAE config is properly initialized even though we don't load the VAE model
vae_config = server_args.pipeline_config.vae_config
vae_config.post_init()
logger.info(
"Called vae_config.post_init() to set vae_scale_factor. "
f"vae_scale_factor={vae_config.arch_config.vae_scale_factor}"
)
def load_modules(
self,
server_args: ServerArgs,
loaded_modules: dict[str, torch.nn.Module] | None = None,
) -> dict[str, Any]:
"""
Load modules for ComfyUIQwenImagePipeline.
If model_path is a safetensors file, load transformer directly from it
without requiring model_index.json. Otherwise, fall back to default loading.
"""
if os.path.isfile(self.model_path) and self.model_path.endswith(".safetensors"):
logger.info(
"Detected safetensors file, loading transformer directly from: %s",
self.model_path,
)
return self._load_transformer_from_safetensors(server_args, loaded_modules)
else:
logger.info(
"Model path is a directory, using default loading method: %s",
self.model_path,
)
return super().load_modules(server_args, loaded_modules)
def _load_transformer_from_safetensors(
self,
server_args: ServerArgs,
loaded_modules: dict[str, torch.nn.Module] | None = None,
) -> dict[str, Any]:
"""Load transformer directly from safetensors without model_index.json."""
# 1) Fast path: use provided module
if loaded_modules is not None and "transformer" in loaded_modules:
logger.info("Using provided transformer module")
return {
"transformer": loaded_modules["transformer"],
"scheduler": self.modules.get("scheduler"),
}
# 2) Build config and mappings
dit_config, updated_mapping, model_cls, default_dtype = (
self._prepare_dit_config_and_mapping(server_args)
)
safetensors_list = [self.model_path]
logger.info("Loading weights from: %s", safetensors_list)
# 3) Instantiate model (meta) and optionally shard
model = self._instantiate_model(
model_cls, dit_config, default_dtype, updated_mapping, server_args
)
# 4) Load weights
self._load_weights_into_model(
model, safetensors_list, default_dtype, updated_mapping, server_args
)
components = {
"transformer": model,
"scheduler": self.modules.get("scheduler"),
}
logger.info("Successfully loaded modules: %s", list(components.keys()))
return components
def _prepare_dit_config_and_mapping(self, server_args: ServerArgs):
from sglang.multimodal_gen.configs.models.dits.qwenimage import (
QwenImageArchConfig,
)
comfyui_arch_config = QwenImageArchConfig(
patch_size=2,
in_channels=64,
out_channels=16,
num_layers=60,
attention_head_dim=128,
num_attention_heads=24,
joint_attention_dim=3584,
pooled_projection_dim=768,
guidance_embeds=False,
axes_dims_rope=(16, 56, 56),
zero_cond_t=self.zero_cond_t,
)
dit_config = QwenImageDitConfig(arch_config=comfyui_arch_config)
server_args.pipeline_config.dit_config = dit_config
if dit_config.arch_config.param_names_mapping is None:
dit_config.arch_config.param_names_mapping = {}
comfyui_qwen_mappings = {r"^model\.diffusion_model\.(.*)$": r"\1"}
updated_mapping = {
**dit_config.arch_config.param_names_mapping,
**comfyui_qwen_mappings,
}
dit_config.arch_config.param_names_mapping = updated_mapping
logger.info(
"Added ComfyUI weight name mappings to param_names_mapping. "
f"Total mappings: {len(updated_mapping)}"
)
cls_name = "QwenImageTransformer2DModel"
model_cls, _ = ModelRegistry.resolve_model_cls(cls_name)
logger.info("Resolved transformer class: %s", cls_name)
default_dtype = resolve_precision(
server_args, "dit", precision_attr="dit_precision"
)
server_args.model_paths["transformer"] = os.path.dirname(self.model_path) or "."
assert server_args.hsdp_shard_dim is not None, "hsdp_shard_dim must be set"
logger.info(
"Loading %s from safetensors file, default_dtype: %s",
cls_name,
default_dtype,
)
return dit_config, updated_mapping, model_cls, default_dtype
def _instantiate_model(
self,
model_cls,
dit_config,
default_dtype,
updated_mapping,
server_args: ServerArgs,
):
from sglang.multimodal_gen.runtime.platforms import current_platform
hf_config = {}
original_mapping = model_cls.param_names_mapping
model_cls.param_names_mapping = updated_mapping
logger.info(
"Temporarily updated model class param_names_mapping with ComfyUI mappings. "
f"Total mappings: {len(updated_mapping)}"
)
try:
# precision-constraint: FSDP mixed precision currently uses bf16
# parameters and fp32 reduction regardless of model load dtype.
mp_policy = MixedPrecisionPolicy(
torch.bfloat16, torch.float32, None, cast_forward_inputs=False
)
set_mixed_precision_policy(
param_dtype=torch.bfloat16,
reduce_dtype=torch.float32,
output_dtype=None,
mp_policy=mp_policy,
)
with set_default_torch_dtype(default_dtype), torch.device("meta"):
model = model_cls(**{"config": dit_config, "hf_config": hf_config})
use_fsdp = server_args.use_fsdp_inference
if current_platform.is_mps():
use_fsdp = False
logger.info("Disabling FSDP for MPS platform as it's not compatible")
if use_fsdp:
device_mesh = init_device_mesh(
current_platform.device_type,
mesh_shape=(
server_args.hsdp_replicate_dim,
server_args.hsdp_shard_dim,
),
mesh_dim_names=("replicate", "shard"),
)
shard_model(
model,
cpu_offload=server_args.dit_cpu_offload,
reshard_after_forward=True,
mp_policy=mp_policy,
mesh=device_mesh,
fsdp_shard_conditions=getattr(
model, "_fsdp_shard_conditions", None
),
pin_cpu_memory=server_args.pin_cpu_memory,
)
finally:
model_cls.param_names_mapping = original_mapping
return model
def _load_weights_into_model(
self,
model,
safetensors_list,
default_dtype,
updated_mapping,
server_args: ServerArgs,
):
# Create weight iterator for loading
weight_iterator = safetensors_weights_iterator(safetensors_list)
# Load weights
param_names_mapping_fn = get_param_names_mapping(updated_mapping)
load_model_from_full_model_state_dict(
model,
weight_iterator,
get_local_torch_device(),
default_dtype,
strict=True,
cpu_offload=server_args.dit_cpu_offload,
param_names_mapping=param_names_mapping_fn,
)
# Check for meta parameters
for n, p in chain(model.named_parameters(), model.named_buffers()):
if p.is_meta:
raise RuntimeError(f"Unexpected param or buffer {n} on meta device.")
if isinstance(p, torch.nn.Parameter):
p.requires_grad = False
total_params = sum(p.numel() for p in model.parameters())
logger.info("Loaded transformer with %.2fB parameters", total_params / 1e9)
def create_pipeline_stages(self, server_args: ServerArgs):
logger.info(
f"{self.__class__.__name__}.create_pipeline_stages() called - creating latent_preparation_stage and denoising_stage"
)
self.add_stages(
[
ComfyUILatentPreparationStage(
scheduler=self.get_module("scheduler"),
transformer=self.get_module("transformer"),
),
DenoisingStage(
transformer=self.get_module("transformer"),
scheduler=self.get_module("scheduler"),
),
]
)
logger.info(
f"{self.__class__.__name__} stages created: {list(self._stage_name_mapping.keys())}"
)
class ComfyUIQwenImagePipeline(ComfyUIQwenImagePipelineBase):
"""ComfyUI QwenImage pipeline for text-to-image generation."""
pipeline_name = "ComfyUIQwenImagePipeline"
zero_cond_t = False
from sglang.multimodal_gen.configs.pipeline_configs.qwen_image import (
QwenImagePipelineConfig,
)
from sglang.multimodal_gen.configs.sample.qwenimage import QwenImageSamplingParams
pipeline_config_cls = QwenImagePipelineConfig
sampling_params_cls = QwenImageSamplingParams
class ComfyUIQwenImageEditPipeline(ComfyUIQwenImagePipelineBase):
"""ComfyUI QwenImage pipeline for image-to-image editing."""
pipeline_name = "ComfyUIQwenImageEditPipeline"
zero_cond_t = True
from sglang.multimodal_gen.configs.pipeline_configs.qwen_image import (
QwenImageEditPlusPipelineConfig,
)
from sglang.multimodal_gen.configs.sample.qwenimage import (
QwenImageEditPlusSamplingParams,
)
pipeline_config_cls = QwenImageEditPlusPipelineConfig
sampling_params_cls = QwenImageEditPlusSamplingParams
EntryClass = [ComfyUIQwenImagePipeline, ComfyUIQwenImageEditPipeline]
@@ -0,0 +1,407 @@
# Copied and adapted from: https://github.com/hao-ai-lab/FastVideo
# SPDX-License-Identifier: Apache-2.0
import os
import re
from collections.abc import Generator
from itertools import chain
from typing import Any
import torch
from torch.distributed import init_device_mesh
from torch.distributed.fsdp import MixedPrecisionPolicy
from sglang.multimodal_gen.configs.models.dits.zimage import ZImageDitConfig
from sglang.multimodal_gen.runtime.distributed import get_local_torch_device
from sglang.multimodal_gen.runtime.loader.fsdp_load import (
load_model_from_full_model_state_dict,
shard_model,
)
from sglang.multimodal_gen.runtime.loader.utils import (
get_param_names_mapping,
set_default_torch_dtype,
)
from sglang.multimodal_gen.runtime.loader.weight_utils import (
safetensors_weights_iterator,
)
from sglang.multimodal_gen.runtime.models.registry import ModelRegistry
from sglang.multimodal_gen.runtime.models.schedulers.scheduling_comfyui_passthrough import (
ComfyUIPassThroughScheduler,
)
from sglang.multimodal_gen.runtime.pipelines_core import LoRAPipeline
from sglang.multimodal_gen.runtime.pipelines_core.composed_pipeline_base import (
ComposedPipelineBase,
)
from sglang.multimodal_gen.runtime.pipelines_core.stages import (
ComfyUILatentPreparationStage,
DenoisingStage,
)
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
from sglang.multimodal_gen.utils import set_mixed_precision_policy
logger = init_logger(__name__)
class ComfyUIZImagePipeline(LoRAPipeline, ComposedPipelineBase):
"""
Simplified pipeline for ComfyUI integration with only denoising stage.
This pipeline requires pre-processed inputs:
- prompt_embeds: Pre-encoded text embeddings (list of tensors)
- negative_prompt_embeds: Pre-encoded negative prompt embeddings (if using CFG)
- latents: Optional initial noise latents (will be generated if not provided)
Usage:
generator = DiffGenerator.from_pretrained(
model_path="path/to/model",
pipeline_class_name="ComfyUIZImagePipeline",
device="cuda",
)
"""
pipeline_name = "ComfyUIZImagePipeline"
from sglang.multimodal_gen.configs.pipeline_configs.zimage import (
ZImagePipelineConfig,
)
from sglang.multimodal_gen.configs.sample.zimage import ZImageSamplingParams
pipeline_config_cls = ZImagePipelineConfig
sampling_params_cls = ZImageSamplingParams
_required_config_modules = [
"transformer",
"scheduler",
]
def initialize_pipeline(self, server_args: ServerArgs):
"""
Initialize the pipeline with ComfyUI pass-through scheduler.
This scheduler does not modify latents, allowing ComfyUI to handle denoising.
"""
self.modules["scheduler"] = ComfyUIPassThroughScheduler(
num_train_timesteps=1000
)
# Ensure VAE config is properly initialized even though we don't load the VAE model
# This is necessary because get_freqs_cis uses spatial_compression_ratio
if hasattr(server_args.pipeline_config, "vae_config"):
vae_config = server_args.pipeline_config.vae_config
if hasattr(vae_config, "post_init") and not hasattr(
vae_config, "_post_init_called"
):
vae_config.post_init()
logger.info(
"Called vae_config.post_init() to set spatial_compression_ratio. "
f"spatial_compression_ratio={vae_config.arch_config.spatial_compression_ratio}"
)
def load_modules(
self,
server_args: ServerArgs,
loaded_modules: dict[str, torch.nn.Module] | None = None,
) -> dict[str, Any]:
"""
Load modules for ComfyUIZImagePipeline.
If model_path is a safetensors file, load transformer directly from it
without requiring model_index.json. Otherwise, fall back to default loading.
"""
if os.path.isfile(self.model_path) and self.model_path.endswith(".safetensors"):
logger.info(
"Detected safetensors file, loading transformer directly from: %s",
self.model_path,
)
return self._load_transformer_from_safetensors(server_args, loaded_modules)
else:
logger.info(
"Model path is a directory, using default loading method: %s",
self.model_path,
)
return super().load_modules(server_args, loaded_modules)
def _convert_comfyui_qkv_weights(
self,
weight_iterator: Generator[tuple[str, torch.Tensor], None, None],
dim: int,
num_heads: int,
num_kv_heads: int,
) -> Generator[tuple[str, torch.Tensor], None, None]:
"""
Convert ComfyUI zimage qkv weights to SGLang format.
Splits merged qkv.weight into separate to_q, to_k, to_v weights.
Args:
weight_iterator: Iterator yielding (name, tensor) pairs from safetensors
dim: Model dimension
num_heads: Number of attention heads
num_kv_heads: Number of key-value heads
Yields:
(name, tensor) pairs with qkv weights split into to_q, to_k, to_v
"""
head_dim = dim // num_heads
q_size = dim
k_size = head_dim * num_kv_heads
for name, tensor in weight_iterator:
# Match qkv weights in layers, noise_refiner, or context_refiner
# Pattern: (layers|noise_refiner|context_refiner).{i}.attention.qkv.(weight|bias)
match = re.match(
r"(layers|noise_refiner|context_refiner)\.(\d+)\.attention\.qkv\.(weight|bias)$",
name,
)
if match:
module_name, layer_idx, param_type = match.groups()
base_name = f"{module_name}.{layer_idx}.attention"
if param_type == "weight":
# Weight shape: (q_size + k_size + v_size, dim)
# Split into q, k, v
q_weight = tensor[:q_size, :]
k_weight = tensor[q_size : q_size + k_size, :]
v_weight = tensor[q_size + k_size :, :]
logger.debug(
f"Splitting {name} (shape {tensor.shape}) into "
f"to_q ({q_weight.shape}), to_k ({k_weight.shape}), to_v ({v_weight.shape})"
)
yield f"{base_name}.to_q.weight", q_weight
yield f"{base_name}.to_k.weight", k_weight
yield f"{base_name}.to_v.weight", v_weight
else: # bias
# Bias shape: (q_size + k_size + v_size,)
# Split into q, k, v
q_bias = tensor[:q_size]
k_bias = tensor[q_size : q_size + k_size]
v_bias = tensor[q_size + k_size :]
logger.debug(
f"Splitting {name} (shape {tensor.shape}) into "
f"to_q ({q_bias.shape}), to_k ({k_bias.shape}), to_v ({v_bias.shape})"
)
yield f"{base_name}.to_q.bias", q_bias
yield f"{base_name}.to_k.bias", k_bias
yield f"{base_name}.to_v.bias", v_bias
else:
# Pass through other weights unchanged
yield name, tensor
def _load_transformer_from_safetensors(
self,
server_args: ServerArgs,
loaded_modules: dict[str, torch.nn.Module] | None = None,
) -> dict[str, Any]:
"""
Load transformer directly from safetensors file without model_index.json.
This method:
1. Uses hardcoded ZImageDitConfig for zimage model
2. Loads transformer from the safetensors file
3. Uses ComfyUIPassThroughScheduler (already created in initialize_pipeline)
"""
# Check if transformer is already provided
if loaded_modules is not None and "transformer" in loaded_modules:
logger.info("Using provided transformer module")
components = {
"transformer": loaded_modules["transformer"],
"scheduler": self.modules.get("scheduler"),
}
return components
if hasattr(server_args.pipeline_config, "dit_config"):
dit_config = server_args.pipeline_config.dit_config
if not isinstance(dit_config, ZImageDitConfig):
logger.warning(
"dit_config is not ZImageDitConfig, creating new ZImageDitConfig"
)
dit_config = ZImageDitConfig()
server_args.pipeline_config.dit_config = dit_config
else:
logger.info("Creating default ZImageDitConfig")
dit_config = ZImageDitConfig()
server_args.pipeline_config.dit_config = dit_config
if dit_config.arch_config.param_names_mapping is None:
dit_config.arch_config.param_names_mapping = {}
# Add mappings for norm layers: map from ComfyUI format (k_norm/q_norm) to SGLang format (norm_k/norm_q)
# The regex matches the source name from safetensors, and the tuple specifies the target name in the model
# Note: qkv weights are handled separately by _convert_comfyui_qkv_weights function
comfyui_norm_mappings = {
r"(.*)\.attention\.k_norm\.weight$": (
r"\1.attention.norm_k.weight",
None,
None,
),
r"(.*)\.attention\.q_norm\.weight$": (
r"\1.attention.norm_q.weight",
None,
None,
),
r"(.*)\.attention\.out\.weight$": (
r"\1.attention.to_out.0.weight",
None,
None,
),
r"^final_layer\.(.*)$": (r"all_final_layer.2-1.\1", None, None),
r"^x_embedder\.(.*)$": (r"all_x_embedder.2-1.\1", None, None),
}
# Merge ComfyUI mappings with existing mappings (ComfyUI mappings take precedence)
updated_mapping = {
**dit_config.arch_config.param_names_mapping,
**comfyui_norm_mappings,
}
dit_config.arch_config.param_names_mapping = updated_mapping
logger.info(
"Added ComfyUI weight name mappings (k_norm/q_norm -> norm_k/norm_q) to param_names_mapping. "
f"Total mappings: {len(updated_mapping)}"
)
cls_name = "ZImageTransformer2DModel"
model_cls, _ = ModelRegistry.resolve_model_cls(cls_name)
logger.info("Resolved transformer class: %s", cls_name)
safetensors_list = [self.model_path]
logger.info("Loading weights from: %s", safetensors_list)
default_dtype = resolve_precision(
server_args, "dit", precision_attr="dit_precision"
)
server_args.model_paths["transformer"] = os.path.dirname(self.model_path) or "."
hf_config = {}
assert server_args.hsdp_shard_dim is not None, "hsdp_shard_dim must be set"
logger.info(
"Loading %s from safetensors file, default_dtype: %s",
cls_name,
default_dtype,
)
original_mapping = model_cls.param_names_mapping
model_cls.param_names_mapping = updated_mapping
logger.info(
"Temporarily updated model class param_names_mapping with ComfyUI mappings. "
f"Total mappings: {len(updated_mapping)}"
)
try:
# Create model first (same as maybe_load_fsdp_model)
from sglang.multimodal_gen.runtime.platforms import current_platform
# precision-constraint: FSDP mixed precision currently uses bf16
# parameters and fp32 reduction regardless of model load dtype.
mp_policy = MixedPrecisionPolicy(
torch.bfloat16, torch.float32, None, cast_forward_inputs=False
)
set_mixed_precision_policy(
param_dtype=torch.bfloat16,
reduce_dtype=torch.float32,
output_dtype=None,
mp_policy=mp_policy,
)
with set_default_torch_dtype(default_dtype), torch.device("meta"):
model = model_cls(**{"config": dit_config, "hf_config": hf_config})
# Check if we should use FSDP
use_fsdp = server_args.use_fsdp_inference
if current_platform.is_mps():
use_fsdp = False
logger.info("Disabling FSDP for MPS platform as it's not compatible")
if use_fsdp:
device_mesh = init_device_mesh(
current_platform.device_type,
mesh_shape=(
server_args.hsdp_replicate_dim,
server_args.hsdp_shard_dim,
),
mesh_dim_names=("replicate", "shard"),
)
shard_model(
model,
cpu_offload=server_args.dit_cpu_offload,
reshard_after_forward=True,
mp_policy=mp_policy,
mesh=device_mesh,
fsdp_shard_conditions=getattr(
model, "_fsdp_shard_conditions", None
),
pin_cpu_memory=server_args.pin_cpu_memory,
)
# Get model dimensions for qkv splitting
arch_config = dit_config.arch_config
dim = arch_config.dim
num_heads = arch_config.num_attention_heads
num_kv_heads = arch_config.n_kv_heads
# Create weight iterator with qkv conversion
base_weight_iterator = safetensors_weights_iterator(safetensors_list)
converted_weight_iterator = self._convert_comfyui_qkv_weights(
base_weight_iterator, dim, num_heads, num_kv_heads
)
# Load weights
param_names_mapping_fn = get_param_names_mapping(updated_mapping)
load_model_from_full_model_state_dict(
model,
converted_weight_iterator,
get_local_torch_device(),
default_dtype,
strict=True,
cpu_offload=server_args.dit_cpu_offload,
param_names_mapping=param_names_mapping_fn,
)
# Check for meta parameters
for n, p in chain(model.named_parameters(), model.named_buffers()):
if p.is_meta:
raise RuntimeError(
f"Unexpected param or buffer {n} on meta device."
)
if isinstance(p, torch.nn.Parameter):
p.requires_grad = False
finally:
model_cls.param_names_mapping = original_mapping
total_params = sum(p.numel() for p in model.parameters())
logger.info("Loaded transformer with %.2fB parameters", total_params / 1e9)
components = {
"transformer": model,
"scheduler": self.modules.get("scheduler"),
}
logger.info("Successfully loaded modules: %s", list(components.keys()))
return components
def create_pipeline_stages(self, server_args: ServerArgs):
logger.info(
"ComfyUIZImagePipeline.create_pipeline_stages() called - creating latent_preparation_stage and denoising_stage"
)
self.add_stages(
[
ComfyUILatentPreparationStage(
scheduler=self.get_module("scheduler"),
transformer=self.get_module("transformer"),
),
DenoisingStage(
transformer=self.get_module("transformer"),
scheduler=self.get_module("scheduler"),
),
]
)
logger.info(
f"ComfyUIZImagePipeline stages created: {list(self._stage_name_mapping.keys())}"
)
EntryClass = ComfyUIZImagePipeline
@@ -0,0 +1,112 @@
# SPDX-License-Identifier: Apache-2.0
"""Cosmos3 video diffusion pipeline.
Cosmos3 has no separate text encoder — the transformer embeds text directly
via its Understanding (UND) pathway, and the Generation (GEN) pathway
cross-attends to the cached UND K/V at each denoising step.
"""
import os
from sglang.multimodal_gen.runtime.pipelines_core.composed_pipeline_base import (
ComposedPipelineBase,
)
from sglang.multimodal_gen.runtime.pipelines_core.stages.model_specific_stages.cosmos3 import (
Cosmos3DecodingStage,
Cosmos3DenoisingStage,
Cosmos3ImagePreprocessStage,
Cosmos3LatentPreparationStage,
Cosmos3TimestepPreparationStage,
Cosmos3TokenizationStage,
)
from sglang.multimodal_gen.runtime.server_args import ServerArgs
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
logger = init_logger(__name__)
class Cosmos3Pipeline(ComposedPipelineBase):
"""Cosmos3 diffusion pipeline shared by T2V, I2V, and T2I.
Text is tokenized and embedded directly inside the transformer; there is
no separate text encoder. Modality is dispatched per-request inside the
stages from ``batch.data_type`` and ``batch.preprocessed_image``.
"""
pipeline_name = "Cosmos3OmniDiffusersPipeline"
is_video_pipeline = True
_required_config_modules = [
"text_tokenizer",
"vae",
"transformer",
"scheduler",
"sound_tokenizer",
]
def load_modules(self, server_args, loaded_modules=None):
# Visual-only Cosmos3 checkpoints ship no sound_tokenizer; require it
# only when the checkpoint actually provides one.
if "sound_tokenizer" not in self._load_config():
self._required_config_modules = [
m for m in self._required_config_modules if m != "sound_tokenizer"
]
return super().load_modules(server_args, loaded_modules)
def create_pipeline_stages(self, server_args: ServerArgs) -> None:
"""Create Cosmos3 pipeline stages.
Stage order:
1. Cosmos3ImagePreprocessStage - Load + aspect-resize the I2V image (no-op otherwise)
2. Cosmos3TokenizationStage - Tokenize with Qwen2 chat template
3. Cosmos3LatentPreparationStage - Noise latent (or image-conditioned for I2V)
4. Cosmos3TimestepPreparationStage - Set up scheduler timesteps
5. Cosmos3DenoisingStage - Dual-pathway denoising (UND once, GEN per step)
6. Cosmos3DecodingStage - VAE decode to video, or to a single image for T2I
"""
text_tokenizer = self.get_module("text_tokenizer")
vae = self.get_module("vae")
transformer = self.get_module("transformer")
scheduler = self.get_module("scheduler")
sound_tokenizer = self.get_module("sound_tokenizer")
guardrails_disabled = (
os.environ.get("SGLANG_DISABLE_COSMOS3_GUARDRAILS", "0") == "1"
)
guardrails_on = False
if not guardrails_disabled:
from sglang.multimodal_gen.runtime.pipelines_core.stages.model_specific_stages.cosmos3_guardrails import (
is_cosmos_guardrail_available,
)
guardrails_on = is_cosmos_guardrail_available()
if not guardrails_on:
logger.warning(
"Cosmos3 guardrails disabled because cosmos-guardrail is not "
"installed. Install it with: pip install cosmos-guardrail==0.3.1"
)
self.add_stage(Cosmos3ImagePreprocessStage())
self.add_stage(Cosmos3TokenizationStage(tokenizer=text_tokenizer))
if guardrails_on:
from sglang.multimodal_gen.runtime.pipelines_core.stages.model_specific_stages.cosmos3_guardrails import (
Cosmos3TextGuardrailStage,
)
self.add_stage(Cosmos3TextGuardrailStage())
self.add_stage(Cosmos3LatentPreparationStage(vae, transformer))
self.add_stage(Cosmos3TimestepPreparationStage(scheduler))
self.add_stage(Cosmos3DenoisingStage(transformer, scheduler, server_args))
self.add_stage(
Cosmos3DecodingStage(
vae, guardrails=guardrails_on, sound_tokenizer=sound_tokenizer
)
)
logger.info(
"Cosmos3 pipeline stages created successfully (guardrails=%s)",
guardrails_on,
)
EntryClass = Cosmos3Pipeline
@@ -0,0 +1,778 @@
# SPDX-License-Identifier: Apache-2.0
"""
Diffusers backend pipeline wrapper.
This module provides a wrapper that allows running any diffusers-supported model
through sglang's infrastructure using vanilla diffusers pipelines.
"""
import argparse
import inspect
import re
import warnings
from typing import Any
import numpy as np
import torch
import torchvision.transforms as T
from diffusers import DiffusionPipeline
from PIL import Image
from sglang.multimodal_gen.configs.pipeline_configs.base import PipelineConfig
from sglang.multimodal_gen.runtime.distributed import get_local_torch_device
from sglang.multimodal_gen.runtime.managers.memory_managers.component_manager import (
ComponentResidencyStrategy,
get_global_component_residency_manager,
)
from sglang.multimodal_gen.runtime.pipelines_core.composed_pipeline_base import (
ComposedPipelineBase,
)
from sglang.multimodal_gen.runtime.pipelines_core.executors.pipeline_executor import (
PipelineExecutor,
)
from sglang.multimodal_gen.runtime.pipelines_core.executors.sync_executor import (
SyncExecutor,
)
from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import Req
from sglang.multimodal_gen.runtime.pipelines_core.stages import PipelineStage
from sglang.multimodal_gen.runtime.platforms import (
AttentionBackendEnum,
current_platform,
)
from sglang.multimodal_gen.runtime.server_args import ServerArgs
from sglang.multimodal_gen.runtime.utils.hf_diffusers_utils import maybe_download_model
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
from sglang.multimodal_gen.runtime.utils.precision import resolve_precision
from sglang.multimodal_gen.runtime.utils.vision import load_image as load_vision_image
logger = init_logger(__name__)
class DiffusersExecutionStage(PipelineStage):
"""Pipeline stage that wraps diffusers pipeline execution."""
def __init__(self, diffusers_pipe: DiffusionPipeline):
super().__init__()
self.diffusers_pipe = diffusers_pipe
def forward(self, batch: Req, server_args: ServerArgs) -> Req:
"""Execute the diffusers pipeline."""
kwargs = self._build_pipeline_kwargs(batch)
# Filter kwargs to only those supported by the pipeline, warn about ignored args
kwargs, _ = self._filter_pipeline_kwargs(kwargs)
# Request tensor output for cleaner handling
if "output_type" not in kwargs:
kwargs["output_type"] = "pt"
with torch.no_grad(), warnings.catch_warnings(record=True):
warnings.simplefilter("always")
try:
output = self.diffusers_pipe(**kwargs)
except TypeError as e:
# Some pipelines don't support output_type="pt"
if "output_type" in str(e):
kwargs.pop("output_type", None)
output = self.diffusers_pipe(**kwargs)
else:
raise
batch.output = self._extract_output(output)
if batch.output is not None:
batch.output = self._postprocess_output(batch.output)
return batch
def _filter_pipeline_kwargs(
self, kwargs: dict[str, Any], *, strict: bool = False
) -> tuple[dict[str, Any], list[str]]:
"""Filter kwargs to those accepted by the pipeline's __call__.
Args:
kwargs: Arguments to filter
strict: If True, raise ValueError on unsupported args; otherwise warn
Returns:
Tuple of (filtered_kwargs, ignored_keys)
"""
try:
sig = inspect.signature(self.diffusers_pipe.__call__)
except (ValueError, TypeError):
return kwargs, []
params = sig.parameters
accepts_var_kwargs = any(
p.kind == inspect.Parameter.VAR_KEYWORD for p in params.values()
)
if accepts_var_kwargs:
return kwargs, []
valid = set(params.keys()) - {"self"}
filtered = {}
ignored = []
for k, v in kwargs.items():
if k in valid:
filtered[k] = v
else:
ignored.append(k)
if ignored:
pipe_name = type(self.diffusers_pipe).__name__
msg = (
f"Pipeline '{pipe_name}' does not support: {', '.join(sorted(ignored))}. "
"These arguments will be ignored."
)
if strict:
raise ValueError(msg)
logger.warning(msg)
return filtered, ignored
def _extract_output(self, output: Any) -> torch.Tensor | None:
"""Extract tensor output from pipeline result."""
for attr in ["images", "frames", "video", "sample", "pred_original_sample"]:
data = getattr(output, attr, None)
if data is None:
continue
result = self._convert_to_tensor(data)
if result is not None:
logger.debug(
"Extracted output from '%s': shape=%s, dtype=%s",
attr,
result.shape,
result.dtype,
)
return result
logger.warning("Could not extract output from pipeline result")
return None
def _convert_to_tensor(self, data: Any) -> torch.Tensor | None:
"""Convert various data formats to a tensor."""
if isinstance(data, torch.Tensor):
return data
if isinstance(data, np.ndarray):
tensor = torch.from_numpy(data).float()
if tensor.max() > 1.0:
tensor = tensor / 255.0
# (B, H, W, C) -> (B, C, H, W) or (B, T, H, W, C) -> (B, C, T, H, W)
if tensor.ndim == 4:
tensor = tensor.permute(0, 3, 1, 2)
elif tensor.ndim == 5:
tensor = tensor.permute(0, 4, 1, 2, 3)
return tensor
if isinstance(data, Image.Image):
return T.ToTensor()(data)
if isinstance(data, list) and len(data) > 0:
return self._convert_list_to_tensor(data)
return None
def _convert_list_to_tensor(self, data: list) -> torch.Tensor | None:
"""Convert a list of items to a tensor."""
first = data[0]
# Nested list (e.g., [[frame1, frame2, ...]] for video batches)
if isinstance(first, list) and len(first) > 0:
data = first
first = data[0]
if isinstance(first, Image.Image):
tensors = [T.ToTensor()(img) for img in data]
stacked = torch.stack(tensors)
if len(tensors) > 1:
return stacked.permute(1, 0, 2, 3) # (T, C, H, W) -> (C, T, H, W)
return stacked[0]
if isinstance(first, torch.Tensor):
stacked = torch.stack(data)
if len(data) > 1:
return stacked.permute(1, 0, 2, 3)
return stacked[0]
if isinstance(first, np.ndarray):
tensors = [torch.from_numpy(arr).float() for arr in data]
if tensors[0].max() > 1.0:
tensors = [t / 255.0 for t in tensors]
if tensors[0].ndim == 3:
tensors = [t.permute(2, 0, 1) for t in tensors]
stacked = torch.stack(tensors)
if len(data) > 1:
return stacked.permute(1, 0, 2, 3)
return stacked[0]
return None
def _postprocess_output(self, output: torch.Tensor) -> torch.Tensor:
"""Post-process output tensor to ensure valid values and correct shape."""
output = output.cpu().float()
# Handle NaN or Inf values
if torch.isnan(output).any() or torch.isinf(output).any():
logger.warning("Output contains invalid values, fixing...")
output = torch.nan_to_num(output, nan=0.5, posinf=1.0, neginf=0.0)
# Normalize to [0, 1] range if needed
min_val, max_val = output.min().item(), output.max().item()
if min_val < -0.5 or max_val > 1.5:
output = (output + 1) / 2
output = output.clamp(0, 1)
# Ensure correct shape for downstream processing
output = self._fix_output_shape(output)
logger.debug("Final output tensor shape: %s", output.shape)
return output
def _fix_output_shape(self, output: torch.Tensor) -> torch.Tensor:
"""Fix tensor shape for downstream processing.
Expected: (B, C, H, W) for images or (B, C, T, H, W) for videos.
"""
if output.dim() == 5:
# Video: (B, T, C, H, W) -> (B, C, T, H, W)
return output.permute(0, 2, 1, 3, 4)
if output.dim() == 4:
if output.shape[0] == 1 or output.shape[1] in [1, 3, 4]:
return output # Already (B, C, H, W)
# (T, C, H, W) -> (1, C, T, H, W)
return output.unsqueeze(0).permute(0, 2, 1, 3, 4)
if output.dim() == 3:
c, h, w = output.shape
if c > 4 and w <= 4:
output = output.permute(2, 0, 1)
if output.shape[0] == 1:
output = output.repeat(3, 1, 1)
return output.unsqueeze(0)
if output.dim() == 2:
return output.unsqueeze(0).repeat(3, 1, 1).unsqueeze(0)
return output
def _build_pipeline_kwargs(self, batch: Req) -> dict[str, Any]:
"""Build kwargs dict for diffusers pipeline call."""
kwargs = {}
if batch.prompt is not None:
kwargs["prompt"] = batch.prompt
if batch.negative_prompt:
kwargs["negative_prompt"] = batch.negative_prompt
if batch.num_inference_steps is not None:
kwargs["num_inference_steps"] = batch.num_inference_steps
if batch.guidance_scale is not None:
kwargs["guidance_scale"] = batch.guidance_scale
if batch.true_cfg_scale is not None:
kwargs["true_cfg_scale"] = batch.true_cfg_scale
if batch.height is not None:
kwargs["height"] = batch.height
if batch.width is not None:
kwargs["width"] = batch.width
if batch.num_frames is not None and batch.num_frames > 1:
kwargs["num_frames"] = batch.num_frames
# Generator for reproducibility
if batch.generator is not None:
kwargs["generator"] = batch.generator
elif batch.seed is not None:
device = self._get_generator_device(batch)
kwargs["generator"] = torch.Generator(device=device).manual_seed(batch.seed)
# Image input for img2img or inpainting
image = self._load_input_image(batch)
if image is not None:
kwargs["image"] = image
if batch.num_outputs_per_prompt > 1:
kwargs["num_images_per_prompt"] = batch.num_outputs_per_prompt
# Extra diffusers-specific kwargs
if batch.extra:
diffusers_kwargs = batch.extra.get("diffusers_kwargs", {})
if diffusers_kwargs:
kwargs.update(diffusers_kwargs)
return kwargs
def _get_generator_device(self, batch: Req) -> str:
"""Resolve RNG device consistently with the non-diffusers path.
Diffusers CPU offload can temporarily park modules on CPU, but that
should not silently switch a CUDA request to CPU RNG, otherwise the
same seed produces different outputs depending on runtime placement.
"""
if batch.generator_device == "cpu":
return "cpu"
return current_platform.device_type
def _load_input_image(self, batch: Req) -> Image.Image | None:
"""Load input image from batch."""
# Check for PIL image in condition_image or pixel_values
if batch.condition_image is not None and isinstance(
batch.condition_image, Image.Image
):
return batch.condition_image
if batch.pixel_values is not None and isinstance(
batch.pixel_values, Image.Image
):
return batch.pixel_values
if not batch.image_path:
return None
if isinstance(batch.image_path, list):
batch.image_path = batch.image_path[0]
try:
image = load_vision_image(batch.image_path)
return image.convert("RGB")
except Exception as e:
logger.error("Failed to load image from %s: %s", batch.image_path, e)
return None
class DiffusersPipeline(ComposedPipelineBase):
"""
Pipeline wrapper that uses vanilla diffusers pipelines.
This allows running any diffusers-supported model through sglang's infrastructure
without requiring native sglang implementation.
"""
pipeline_name = "DiffusersPipeline"
is_video_pipeline = False
_required_config_modules: list[str] = []
def __init__(
self,
model_path: str,
server_args: ServerArgs,
required_config_modules: list[str] | None = None,
loaded_modules: dict[str, torch.nn.Module] | None = None,
executor: PipelineExecutor | None = None,
):
self.server_args = server_args
self.model_path = model_path
self._stages: list[PipelineStage] = []
self._stage_name_mapping: dict[str, PipelineStage] = {}
self.modules: dict[str, Any] = {}
self.memory_usages: dict[str, float] = {}
self.component_residency_strategies: dict[str, ComponentResidencyStrategy] = {}
self.component_residency_manager = None
self.post_init_called = False
self.executor = executor or SyncExecutor(server_args=server_args)
self._cache_dit_enabled = False
logger.info("Loading diffusers pipeline from %s", model_path)
self.diffusers_pipe = self._load_diffusers_pipeline(model_path, server_args)
self._detect_pipeline_type()
def _load_diffusers_pipeline(
self, model_path: str, server_args: ServerArgs
) -> DiffusionPipeline:
"""Load the diffusers pipeline.
Optimizations applied:
- device_map: Loads models directly to GPU, warming up CUDA caching allocator
to avoid small tensor allocations during inference.
- Parallel shard loading: When using device_map with accelerate, model shards
are loaded in parallel for faster initialization.
"""
original_model_path = model_path # Keep original for custom_pipeline
model_path = maybe_download_model(model_path, force_diffusers_model=True)
self.model_path = model_path
dtype = self._get_dtype(server_args)
logger.info("Loading diffusers pipeline with dtype=%s", dtype)
# Build common kwargs for from_pretrained
load_kwargs = {
"torch_dtype": dtype,
"trust_remote_code": server_args.trust_remote_code,
"revision": server_args.revision,
}
# Add quantization config if provided (e.g., BitsAndBytesConfig for 4/8-bit)
quant_config = getattr(server_args.pipeline_config, "quantization_config", None)
if quant_config is not None:
load_kwargs["quantization_config"] = quant_config
logger.info("Using quantization config: %s", type(quant_config).__name__)
try:
pipe = DiffusionPipeline.from_pretrained(model_path, **load_kwargs)
except AttributeError as e:
if "has no attribute" in str(e):
# Custom pipeline class not in diffusers - try loading with custom_pipeline
logger.info(
"Pipeline class not found in diffusers, trying custom_pipeline from repo..."
)
try:
custom_kwargs = {
**load_kwargs,
"custom_pipeline": original_model_path,
}
custom_kwargs["trust_remote_code"] = True
pipe = DiffusionPipeline.from_pretrained(
model_path, **custom_kwargs
)
except Exception as e2:
match = re.search(r"has no attribute (\w+)", str(e))
class_name = match.group(1) if match else "unknown"
raise RuntimeError(
f"Pipeline class '{class_name}' not found in diffusers and no custom pipeline.py in repo. "
f"Try: pip install --upgrade diffusers (some pipelines require latest version). "
f"Original error: {e}"
) from e2
else:
raise
except Exception as e:
# Only retry with float32 for dtype-related errors
if "dtype" in str(e).lower() or "float" in str(e).lower():
logger.warning(
"Failed with dtype=%s, falling back to float32: %s", dtype, e
)
load_kwargs["torch_dtype"] = torch.float32
pipe = DiffusionPipeline.from_pretrained(model_path, **load_kwargs)
else:
raise
# Use CPU offload (all-or-nothing in diffusers) if any component offload is requested.
any_offload = (
server_args.dit_cpu_offload
or server_args.text_encoder_cpu_offload
or server_args.image_encoder_cpu_offload
or server_args.vae_cpu_offload
)
if any_offload:
device = get_local_torch_device()
gpu_id = device.index if device.index is not None else 0
pipe.enable_model_cpu_offload(gpu_id=gpu_id)
logger.info(
"Enabled model CPU offload for diffusers pipeline (gpu_id=%d)", gpu_id
)
else:
pipe = pipe.to(get_local_torch_device())
# Apply VAE memory optimizations from pipeline config
self._apply_vae_optimizations(pipe, server_args)
# Apply attention backend if specified
self._apply_attention_backend(pipe, server_args)
# Apply cache-dit acceleration if configured
pipe = self._apply_cache_dit(pipe, server_args)
# Apply torch.compile if enabled and supported
pipe = self._apply_torch_compile(pipe, server_args)
logger.info("Loaded diffusers pipeline: %s", pipe.__class__.__name__)
return pipe
def _apply_vae_optimizations(
self, pipe: DiffusionPipeline, server_args: ServerArgs
) -> None:
"""Apply VAE memory optimizations (tiling, slicing) from pipeline config."""
config = server_args.pipeline_config
# VAE slicing: decode latents slice-by-slice for lower peak memory
# https://huggingface.co/docs/diffusers/optimization/memory#vae-slicing
if config.vae_slicing:
if hasattr(pipe, "vae") and hasattr(pipe.vae, "enable_slicing"):
pipe.vae.enable_slicing()
logger.info("Enabled VAE slicing for lower memory usage")
elif hasattr(pipe, "enable_vae_slicing"):
pipe.enable_vae_slicing()
logger.info("Enabled VAE slicing for lower memory usage")
else:
logger.warning(
"VAE slicing is not available: neither "
"`pipe.vae.enable_slicing()` nor `pipe.enable_vae_slicing()` was found."
)
# VAE tiling: decode latents tile-by-tile for large images
# https://huggingface.co/docs/diffusers/optimization/memory#vae-tiling
if config.vae_tiling:
if hasattr(pipe, "vae") and hasattr(pipe.vae, "enable_tiling"):
pipe.vae.enable_tiling()
logger.info("Enabled VAE tiling for large image support")
elif hasattr(pipe, "enable_vae_tiling"):
pipe.enable_vae_tiling()
logger.info("Enabled VAE tiling for large image support")
else:
logger.warning(
"VAE tiling is not available: neither "
"`pipe.vae.enable_tiling()` nor `pipe.enable_vae_tiling()` was found."
)
def _apply_attention_backend(
self, pipe: DiffusionPipeline, server_args: ServerArgs
) -> None:
"""Apply attention backend setting from pipeline config or server_args.
See: https://huggingface.co/docs/diffusers/main/en/optimization/attention_backends
Available backends: flash, _flash_3_hub, sage, xformers, native, etc.
"""
backend = server_args.attention_backend
if backend is None:
backend = getattr(
server_args.pipeline_config, "diffusers_attention_backend", None
)
if backend is None:
return
backend = backend.lower()
sglang_backends = {e.name.lower() for e in AttentionBackendEnum} | {
"fa3",
"fa4",
}
if backend in sglang_backends:
logger.debug(
"Skipping diffusers attention backend '%s' because it matches a "
"SGLang backend name. Use diffusers backend names when running "
"the diffusers backend.",
backend,
)
return
for component_name in ["transformer", "unet"]:
component = getattr(pipe, component_name, None)
if component is not None and hasattr(component, "set_attention_backend"):
try:
component.set_attention_backend(backend)
logger.info(
"Set attention backend '%s' on %s", backend, component_name
)
except Exception as e:
logger.warning(
"Failed to set attention backend '%s' on %s: %s",
backend,
component_name,
e,
)
def _apply_cache_dit(
self, pipe: DiffusionPipeline, server_args: ServerArgs
) -> DiffusionPipeline:
"""Enable cache-dit for diffusers pipeline if configured."""
cache_dit_config = server_args.cache_dit_config
if not cache_dit_config:
return pipe
try:
import cache_dit
except ImportError as e:
raise RuntimeError(
"cache-dit is required for --cache-dit-config. "
"Install it with `pip install cache-dit`."
) from e
if not hasattr(cache_dit, "load_configs"):
raise RuntimeError(
"cache-dit>=1.2.0 is required for --cache-dit-config. "
"Please upgrade cache-dit."
)
try:
cache_options = cache_dit.load_configs(cache_dit_config)
except Exception as e:
raise ValueError(
"Failed to load cache-dit config. Provide a YAML/JSON path (or a dict "
"supported by cache-dit>=1.2.0)."
) from e
try:
pipe = cache_dit.enable_cache(pipe, **cache_options)
except Exception:
# cache-dit is an external integration and can raise a variety of errors.
logger.exception("Failed to enable cache-dit for diffusers pipeline")
raise
logger.info("Enabled cache-dit for diffusers pipeline")
self._cache_dit_enabled = True
return pipe
def _apply_torch_compile(self, pipe: Any, server_args: ServerArgs) -> Any:
"""Apply torch.compile to the pipeline if configured and supported."""
if not server_args.enable_torch_compile:
return pipe
# check if the pipeline has 'transformer' or 'unet' components which are
# typically the most expensive parts to compile. 'transformer_2' for some
# video pipelines, e.g, Wan 2.2 series, also check for that.
compilable_components = ["transformer", "transformer_2", "unet"]
if not any(hasattr(pipe, comp) for comp in compilable_components):
logger.warning(
"Pipeline does not have 'transformer' or 'unet' components. "
"torch.compile may not provide significant benefits and could increase latency."
)
return pipe
if self._cache_dit_enabled:
try:
import cache_dit
if hasattr(cache_dit, "set_compile_configs"):
cache_dit.set_compile_configs()
except Exception as e:
logger.warning(
f"Failed to set torch_compile configs for cache-dit: {e}"
)
for comp in compilable_components:
if hasattr(pipe, comp):
try:
component = getattr(pipe, comp)
repeated_blocks = getattr(component, "_repeated_blocks", None)
if (
isinstance(component, torch.nn.Module)
and repeated_blocks
and hasattr(component, "compile_repeated_blocks")
):
# Regional compilation: compile a single instance of each
# repeated transformer block and let inductor's cache reuse
# it for all repeats, instead of compiling the whole DiT as
# one graph
component.compile_repeated_blocks()
elif isinstance(component, torch.nn.Module) and hasattr(
component, "compile"
):
# Prefer in-place compilation if supported. According to PyTorch documentation:
# https://docs.pytorch.org/docs/stable/generated/torch.compile.html
component.compile()
else:
compiled_component = torch.compile(component)
setattr(pipe, comp, compiled_component)
logger.info(
f"Applied torch.compile to {comp} component of the pipeline"
)
except Exception as e:
logger.warning(f"Failed to apply torch.compile to {comp}: {e}")
return pipe
def _get_dtype(self, server_args: ServerArgs) -> torch.dtype:
"""
Determine the dtype to use for model loading.
"""
if hasattr(server_args, "pipeline_config") and server_args.pipeline_config:
return resolve_precision(server_args, "dit", precision_attr="dit_precision")
# precision-constraint: legacy fallback for callers without pipeline_config;
# prefer explicit dit_precision policy when available.
return torch.bfloat16 if torch.cuda.is_bf16_supported() else torch.float16
def _detect_pipeline_type(self) -> None:
"""Detect if this is an image or video pipeline."""
pipe_class_name = self.diffusers_pipe.__class__.__name__.lower()
video_indicators = ["video", "animat", "cogvideo", "wan", "hunyuan"]
self.is_video_pipeline = any(ind in pipe_class_name for ind in video_indicators)
logger.debug(
"Detected pipeline type: %s",
"video" if self.is_video_pipeline else "image",
)
def load_modules(
self,
server_args: ServerArgs,
loaded_modules: dict[str, torch.nn.Module] | None = None,
) -> dict[str, Any]:
"""Skip sglang's module loading - diffusers handles it."""
return {"diffusers_pipeline": self.diffusers_pipe}
def create_pipeline_stages(self, server_args: ServerArgs) -> None:
"""Create the execution stage wrapping the diffusers pipeline."""
self.add_stage(
stage_name="diffusers_execution",
stage=DiffusersExecutionStage(self.diffusers_pipe),
)
def initialize_pipeline(self, server_args: ServerArgs) -> None:
pass
def post_init(self) -> None:
"""Post initialization hook."""
if self.post_init_called:
return
self.post_init_called = True
self.initialize_pipeline(self.server_args)
self.create_pipeline_stages(self.server_args)
def add_stage(self, stage_name: str, stage: PipelineStage) -> None:
"""Add a stage to the pipeline."""
if stage_name is None:
stage_name = self._infer_stage_name(stage)
if stage_name in self._stage_name_mapping:
raise ValueError(f"Duplicate stage name detected: {stage_name}")
stage.set_registered_stage_name(stage_name)
stage.set_profile_stage_name(self._profile_stage_name(stage, stage_name))
self._stages.append(stage)
self._stage_name_mapping[stage_name] = stage
return self
@property
def stages(self) -> list[PipelineStage]:
"""List of stages in the pipeline."""
return self._stages
@torch.no_grad()
def forward(self, batch: Req, server_args: ServerArgs) -> Req:
"""Execute the pipeline on the given batch."""
if not self.post_init_called:
self.post_init()
self.component_residency_manager = get_global_component_residency_manager(
self, server_args
)
self.executor.component_residency_manager = self.component_residency_manager
return self.executor.execute_with_profiling(self.stages, batch, server_args)
@classmethod
def from_pretrained(
cls,
model_path: str,
device: str | None = None,
torch_dtype: torch.dtype | None = None,
pipeline_config: str | PipelineConfig | None = None,
args: argparse.Namespace | None = None,
required_config_modules: list[str] | None = None,
loaded_modules: dict[str, torch.nn.Module] | None = None,
**kwargs,
) -> "DiffusersPipeline":
"""Load a pipeline from a pretrained model using diffusers backend."""
kwargs["model_path"] = model_path
server_args = ServerArgs.from_kwargs(**kwargs)
pipe = cls(
model_path,
server_args,
required_config_modules=required_config_modules,
loaded_modules=loaded_modules,
)
pipe.post_init()
return pipe
def get_module(self, module_name: str, default_value: Any = None) -> Any:
"""Get a module by name."""
if module_name == "diffusers_pipeline":
return self.diffusers_pipe
return self.modules.get(module_name, default_value)
EntryClass = DiffusersPipeline
@@ -0,0 +1,232 @@
# SPDX-License-Identifier: Apache-2.0
"""ErnieImage text-to-image pipeline."""
import json
import os
from sglang.multimodal_gen.runtime.pipelines_core.composed_pipeline_base import (
ComposedPipelineBase,
)
from sglang.multimodal_gen.runtime.pipelines_core.lora_pipeline import LoRAPipeline
from sglang.multimodal_gen.runtime.pipelines_core.stages.input_validation import (
InputValidationStage,
)
from sglang.multimodal_gen.runtime.pipelines_core.stages.model_specific_stages.ernie_image_pe import (
PromptEnhancementStage,
)
from sglang.multimodal_gen.runtime.pipelines_core.stages.text_encoding import (
TextEncodingStage,
)
from sglang.multimodal_gen.runtime.utils.hf_diffusers_utils import (
maybe_download_model,
maybe_download_model_index,
)
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
logger = init_logger(__name__)
class ErnieImagePipeline(LoRAPipeline, ComposedPipelineBase):
pipeline_name = "ErnieImagePipeline"
_required_config_modules = [
"text_encoder",
"tokenizer",
"vae",
"transformer",
"scheduler",
]
def _has_pe_in_model_index(self, server_args) -> bool:
try:
model_index = maybe_download_model_index(server_args.model_path)
return "pe" in model_index and model_index["pe"] is not None
except Exception:
return False
def _read_tokenizer_model_max_length(self, model_path: str):
"""Read model_max_length from tokenizer/tokenizer_config.json.
Supports both local paths and HuggingFace Hub model IDs.
Returns None if the value cannot be determined.
"""
tokenizer_config_subpath = os.path.join("tokenizer", "tokenizer_config.json")
# Local path
if os.path.exists(model_path):
config_path = os.path.join(model_path, tokenizer_config_subpath)
if os.path.exists(config_path):
with open(config_path, encoding="utf-8") as f:
config = json.load(f)
return config.get("model_max_length")
return None
# Remote HuggingFace Hub model ID
try:
import tempfile
from huggingface_hub import hf_hub_download
with tempfile.TemporaryDirectory() as tmp_dir:
config_path = hf_hub_download(
repo_id=model_path,
filename=tokenizer_config_subpath,
local_dir=tmp_dir,
)
with open(config_path, encoding="utf-8") as f:
config = json.load(f)
return config.get("model_max_length")
except Exception as e:
logger.warning(
"Failed to read tokenizer_config.json from %s: %s", model_path, e
)
return None
def _resolve_pe_tokenizer_path(self, model_path: str, server_args) -> str:
"""Resolve the directory that contains the PE tokenizer files."""
pe_component_path = server_args.component_paths.get(
"pe", os.path.join(model_path, "pe")
)
if os.path.exists(os.path.join(pe_component_path, "tokenizer_config.json")):
return pe_component_path
pe_tokenizer_dir = os.path.join(model_path, "pe_tokenizer")
if os.path.exists(os.path.join(pe_tokenizer_dir, "tokenizer_config.json")):
return pe_tokenizer_dir
return pe_component_path
def _read_pe_model_max_length(self, model_path: str, server_args) -> int | None:
# If model_path is a Hub ID, download the full model first (or use cache)
# so that pe/tokenizer_config.json is available locally.
if not os.path.exists(model_path):
try:
model_path = maybe_download_model(
model_path, force_diffusers_model=True
)
except Exception as e:
logger.warning(
"Failed to download model to read pe/tokenizer_config.json: %s", e
)
return None
tokenizer_path = self._resolve_pe_tokenizer_path(model_path, server_args)
config_path = os.path.join(tokenizer_path, "tokenizer_config.json")
if os.path.exists(config_path):
try:
with open(config_path, encoding="utf-8") as f:
config = json.load(f)
val = config.get("model_max_length")
if val is not None:
return int(val)
except Exception as e:
logger.warning(
"Failed to read tokenizer_config.json from %s: %s",
tokenizer_path,
e,
)
return None
def load_modules(self, server_args, loaded_modules=None):
has_pe = self._has_pe_in_model_index(server_args)
if has_pe:
if "pe" not in self._required_config_modules:
self._required_config_modules.insert(0, "pe")
logger.info("PE model detected in model_index.json, will load PE module.")
pipeline_config = server_args.pipeline_config
# --- Text encoder max_length ---
text_model_max_length = self._read_tokenizer_model_max_length(
server_args.model_path
)
if text_model_max_length is not None:
# 1. Update arch_config.text_len so the model knows the true sequence length
if (
hasattr(pipeline_config, "text_encoder_configs")
and pipeline_config.text_encoder_configs
):
arch_config = pipeline_config.text_encoder_configs[0].arch_config
arch_config.text_len = text_model_max_length
arch_config.tokenizer_kwargs["max_length"] = text_model_max_length
# 2. Update text_encoder_extra_args used by TextEncodingStage tokenization
if (
hasattr(pipeline_config, "text_encoder_extra_args")
and pipeline_config.text_encoder_extra_args
):
pipeline_config.text_encoder_extra_args[0][
"max_length"
] = text_model_max_length
logger.info(
"Set text encoder model_max_length=%d from tokenizer/tokenizer_config.json",
text_model_max_length,
)
else:
logger.warning(
"Could not read model_max_length from tokenizer/tokenizer_config.json, "
"text encoder will use the default text_len from arch config."
)
# --- PE model_max_length ---
if has_pe:
pe_model_max_length = self._read_pe_model_max_length(
server_args.model_path, server_args
)
if pe_model_max_length is not None:
pipeline_config.pe_model_max_length = pe_model_max_length
logger.info(
"Set PE model_max_length=%d from pe/tokenizer_config.json",
pe_model_max_length,
)
else:
raise RuntimeError(
"PE model is present but 'model_max_length' could not be read from "
"pe/tokenizer_config.json. Please ensure the PE component directory "
"contains a valid tokenizer_config.json with a 'model_max_length' field."
)
return super().load_modules(server_args, loaded_modules)
def create_pipeline_stages(self, server_args):
self.add_stage(InputValidationStage())
pe_model = self.get_module("pe")
if pe_model is not None:
pe_tokenizer = getattr(pe_model, "pe_tokenizer", None)
if pe_tokenizer is None:
from transformers import AutoTokenizer
pe_tokenizer_path = self._resolve_pe_tokenizer_path(
self.model_path, server_args
)
logger.warning(
"pe_tokenizer not found on pe_model (%s), loading from %s",
type(pe_model).__name__,
pe_tokenizer_path,
)
pe_tokenizer = AutoTokenizer.from_pretrained(
pe_tokenizer_path,
trust_remote_code=server_args.trust_remote_code,
)
self.add_stage(
PromptEnhancementStage(
pe_model=pe_model,
pe_tokenizer=pe_tokenizer,
),
"prompt_enhancement_stage",
)
self.add_stage(
TextEncodingStage(
text_encoders=[self.get_module("text_encoder")],
tokenizers=[self.get_module("tokenizer")],
),
"prompt_encoding_stage_primary",
)
self.add_standard_timestep_preparation_stage()
self.add_standard_latent_preparation_stage()
self.add_standard_denoising_stage()
self.add_standard_decoding_stage()
EntryClass = ErnieImagePipeline
@@ -0,0 +1,78 @@
# Copied and adapted from: https://github.com/hao-ai-lab/FastVideo
# SPDX-License-Identifier: Apache-2.0
from sglang.multimodal_gen.runtime.pipelines_core import LoRAPipeline
from sglang.multimodal_gen.runtime.pipelines_core.composed_pipeline_base import (
ComposedPipelineBase,
)
from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import Req
from sglang.multimodal_gen.runtime.pipelines_core.stages.progressive_resolution.flux import (
FluxProgressiveDenoisingStage,
)
from sglang.multimodal_gen.runtime.server_args import ServerArgs
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
# TODO(will): move PRECISION_TO_TYPE to better place
logger = init_logger(__name__)
def calculate_shift(
image_seq_len,
base_seq_len: int = 256,
max_seq_len: int = 4096,
base_shift: float = 0.5,
max_shift: float = 1.15,
):
m = (max_shift - base_shift) / (max_seq_len - base_seq_len)
b = base_shift - m * base_seq_len
mu = image_seq_len * m + b
return mu
def prepare_mu(batch: Req, server_args: ServerArgs):
height = batch.height
width = batch.width
vae_scale_factor = (
server_args.pipeline_config.vae_config.arch_config.vae_scale_factor
)
image_seq_len = (int(height) // (vae_scale_factor * 2)) * (
int(width) // (vae_scale_factor * 2)
)
mu = calculate_shift(
image_seq_len,
# hard code, since scheduler_config is not in PipelineConfig now
256,
4096,
0.5,
1.15,
)
return "mu", mu
class FluxPipeline(LoRAPipeline, ComposedPipelineBase):
pipeline_name = "FluxPipeline"
_required_config_modules = [
"text_encoder",
"text_encoder_2",
"tokenizer",
"tokenizer_2",
"vae",
"transformer",
"scheduler",
]
def create_pipeline_stages(self, server_args: ServerArgs):
self.add_standard_t2i_stages(
text_encoder_key=["text_encoder", "text_encoder_2"],
tokenizer_key=["tokenizer", "tokenizer_2"],
text_encoding_stage_name="prompt_encoding_stage_primary",
prepare_extra_timestep_kwargs=[prepare_mu],
progressive_denoising_stage_cls=FluxProgressiveDenoisingStage,
)
EntryClass = FluxPipeline
@@ -0,0 +1,66 @@
# Copied and adapted from: https://github.com/hao-ai-lab/FastVideo
# SPDX-License-Identifier: Apache-2.0
from diffusers.pipelines.flux2.image_processor import Flux2ImageProcessor
from sglang.multimodal_gen.runtime.pipelines_core import LoRAPipeline, Req
from sglang.multimodal_gen.runtime.pipelines_core.composed_pipeline_base import (
ComposedPipelineBase,
)
from sglang.multimodal_gen.runtime.pipelines_core.stages.progressive_resolution.flux_2 import (
Flux2ProgressiveDenoisingStage,
)
from sglang.multimodal_gen.runtime.server_args import ServerArgs
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
logger = init_logger(__name__)
def compute_empirical_mu(batch: Req, server_args: ServerArgs):
num_steps = batch.num_inference_steps
image_seq_len = batch.raw_latent_shape[1]
a1, b1 = 8.73809524e-05, 1.89833333
a2, b2 = 0.00016927, 0.45666666
if image_seq_len > 4300:
mu = a2 * image_seq_len + b2
return "mu", float(mu)
m_200 = a2 * image_seq_len + b2
m_10 = a1 * image_seq_len + b1
a = (m_200 - m_10) / 190.0
b = m_200 - 200.0 * a
mu = a * num_steps + b
return "mu", float(mu)
class Flux2Pipeline(LoRAPipeline, ComposedPipelineBase):
pipeline_name = "Flux2Pipeline"
_required_config_modules = [
"text_encoder",
"tokenizer",
"vae",
"transformer",
"scheduler",
]
def create_pipeline_stages(self, server_args: ServerArgs):
vae_image_processor = Flux2ImageProcessor(
vae_scale_factor=server_args.pipeline_config.vae_config.arch_config.vae_scale_factor
* 2
)
self.add_standard_ti2i_stages(
include_input_validation=True,
vae_image_processor=vae_image_processor,
prompt_encoding="text",
image_vae_stage_kwargs={"vae_image_processor": vae_image_processor},
prepare_extra_timestep_kwargs=[compute_empirical_mu],
progressive_denoising_stage_cls=Flux2ProgressiveDenoisingStage,
)
EntryClass = Flux2Pipeline
@@ -0,0 +1,8 @@
from sglang.multimodal_gen.runtime.pipelines.flux_2 import Flux2Pipeline
class Flux2KleinPipeline(Flux2Pipeline):
pipeline_name = "Flux2KleinPipeline"
EntryClass = Flux2KleinPipeline
@@ -0,0 +1,130 @@
# SPDX-License-Identifier: Apache-2.0
import glob
import os
from dataclasses import dataclass
from functools import lru_cache
from typing import Any, cast
from sglang.multimodal_gen.runtime.pipelines.flux_2 import Flux2Pipeline
from sglang.multimodal_gen.runtime.server_args import ServerArgs
from sglang.multimodal_gen.runtime.utils.hf_diffusers_utils import (
maybe_download_model,
verify_model_config_and_directory,
)
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
logger = init_logger(__name__)
@dataclass(frozen=True)
class Flux2Nvfp4ModelResolution:
base_model_name: str
base_model_path: str
transformer_weights_path: str
_FLUX2_BASE_MODEL = "black-forest-labs/FLUX.2-dev"
@lru_cache(maxsize=1)
def _resolve_flux2_base_model_path() -> str:
return maybe_download_model(_FLUX2_BASE_MODEL, force_diffusers_model=True)
def _find_mixed_safetensors(local_dir: str) -> str | None:
mixed_files = sorted(glob.glob(os.path.join(local_dir, "*-mixed.safetensors")))
return mixed_files[0] if mixed_files else None
def _resolve_nvfp4_transformer_weights_path(
server_args: ServerArgs, model_path: str
) -> str:
if server_args.transformer_weights_path is not None:
return server_args.transformer_weights_path
local_nvfp4_path = maybe_download_model(model_path)
mixed_file = _find_mixed_safetensors(local_nvfp4_path)
if mixed_file is not None:
logger.info("Using mixed-precision NVFP4 weights: %s", mixed_file)
return mixed_file
logger.warning(
"No *-mixed.safetensors found in %s; falling back to full directory",
local_nvfp4_path,
)
return local_nvfp4_path
def resolve_flux2_nvfp4_model(
server_args: ServerArgs, model_path: str
) -> Flux2Nvfp4ModelResolution:
transformer_weights_path = _resolve_nvfp4_transformer_weights_path(
server_args, model_path
)
return Flux2Nvfp4ModelResolution(
base_model_name=_FLUX2_BASE_MODEL,
base_model_path=_resolve_flux2_base_model_path(),
transformer_weights_path=transformer_weights_path,
)
class Flux2NvfpPipeline(Flux2Pipeline):
pipeline_name = "Flux2NvfpPipeline"
_model_resolution: Flux2Nvfp4ModelResolution | None = None
def _get_model_resolution(
self, server_args: ServerArgs | None = None
) -> Flux2Nvfp4ModelResolution:
if self._model_resolution is None:
if server_args is None:
raise ValueError(
"server_args is required to resolve FLUX.2 NVFP4 paths"
)
self._model_resolution = resolve_flux2_nvfp4_model(
server_args, self.model_path
)
return self._model_resolution
def _load_config(self) -> dict[str, Any]:
model_resolution = self._get_model_resolution(self.server_args)
logger.info("Model path: %s", self.model_path)
logger.info(
"Using base model '%s' at %s for config and non-transformer components",
model_resolution.base_model_name,
model_resolution.base_model_path,
)
config = verify_model_config_and_directory(model_resolution.base_model_path)
return cast(dict[str, Any], config)
def _resolve_component_path(
self, server_args: ServerArgs, module_name: str, load_module_name: str
) -> str:
override_path = server_args.component_paths.get(module_name)
if override_path is not None:
return maybe_download_model(override_path)
# get non-transformer components from the base FLUX.2 repo explicitly.
# e.g.:
# transformer weights: ...FLUX.2-dev-NVFP4/.../flux2-dev-nvfp4-mixed.safetensors
# text_encoder path: ...FLUX.2-dev/.../text_encoder
component_model_path = os.path.join(
self._get_model_resolution(server_args).base_model_path, load_module_name
)
logger.debug("Resolved component path: %s", component_model_path)
return component_model_path
def load_modules(
self,
server_args: ServerArgs,
loaded_modules: dict | None = None,
) -> dict:
model_resolution = self._get_model_resolution(server_args)
server_args.transformer_weights_path = model_resolution.transformer_weights_path
logger.info(
"NVFP4 transformer weights: %s",
model_resolution.transformer_weights_path,
)
return super().load_modules(server_args, loaded_modules)
EntryClass = Flux2NvfpPipeline
@@ -0,0 +1,59 @@
from sglang.multimodal_gen.runtime.pipelines_core import LoRAPipeline
from sglang.multimodal_gen.runtime.pipelines_core.composed_pipeline_base import (
ComposedPipelineBase,
)
from sglang.multimodal_gen.runtime.pipelines_core.stages import DenoisingStage
from sglang.multimodal_gen.runtime.pipelines_core.stages.model_specific_stages.glm_image import (
GlmImageAR,
GlmImageBeforeDenoisingStage,
)
from sglang.multimodal_gen.runtime.server_args import ServerArgs
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
logger = init_logger(__name__)
class GlmImagePipeline(LoRAPipeline, ComposedPipelineBase):
pipeline_name = "GlmImagePipeline"
_required_config_modules = [
"text_encoder",
"tokenizer",
"vae",
"vision_language_encoder",
"processor",
"transformer",
"scheduler",
]
def create_pipeline_stages(self, server_args: ServerArgs):
self.add_stage(
GlmImageAR(
processor=self.get_module("processor"),
vision_language_encoder=self.get_module("vision_language_encoder"),
),
"glm_image_ar",
)
self.add_stage(
GlmImageBeforeDenoisingStage(
vae=self.get_module("vae"),
text_encoder=self.get_module("text_encoder"),
tokenizer=self.get_module("tokenizer"),
transformer=self.get_module("transformer"),
scheduler=self.get_module("scheduler"),
),
"glm_image_before_denoising_stage",
)
self.add_stage(
DenoisingStage(
transformer=self.get_module("transformer"),
scheduler=self.get_module("scheduler"),
),
)
self.add_standard_decoding_stage()
EntryClass = [GlmImagePipeline]
@@ -0,0 +1,89 @@
# SPDX-License-Identifier: Apache-2.0
"""
Helios video diffusion pipeline implementation.
This module contains an implementation of the Helios video diffusion pipeline
using the modular pipeline architecture. Phase 1: T2V only.
"""
from sglang.multimodal_gen.runtime.pipelines_core.composed_pipeline_base import (
ComposedPipelineBase,
)
from sglang.multimodal_gen.runtime.pipelines_core.lora_pipeline import LoRAPipeline
from sglang.multimodal_gen.runtime.pipelines_core.stages import (
InputValidationStage,
)
from sglang.multimodal_gen.runtime.pipelines_core.stages.model_specific_stages.helios_decoding import (
HeliosDecodingStage,
)
from sglang.multimodal_gen.runtime.pipelines_core.stages.model_specific_stages.helios_denoising import (
HeliosChunkedDenoisingStage,
)
from sglang.multimodal_gen.runtime.server_args import ServerArgs
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
logger = init_logger(__name__)
class HeliosPipeline(LoRAPipeline, ComposedPipelineBase):
"""
Helios video diffusion pipeline with LoRA support.
Implements the Helios T2V pipeline with chunked denoising,
multi-term memory history, and CFG Zero Star guidance.
"""
pipeline_name = "HeliosPipeline"
_required_config_modules = [
"text_encoder",
"tokenizer",
"vae",
"transformer",
"scheduler",
]
def initialize_pipeline(self, server_args: ServerArgs):
# Use the scheduler loaded from model's scheduler_config.json as-is.
# It contains critical config: use_dynamic_shifting=true,
# time_shift_type="exponential", etc.
scheduler = self.modules.get("scheduler")
if scheduler is not None and server_args.pipeline_config.flow_shift is not None:
scheduler.set_shift(server_args.pipeline_config.flow_shift)
# Configure scheduler for Stage 2/3 if enabled
pipeline_config = server_args.pipeline_config
if scheduler is not None and pipeline_config.is_enable_stage2:
scheduler.config.stages = pipeline_config.pyramid_num_stages
scheduler.config.scheduler_type = pipeline_config.scheduler_type
scheduler.config.gamma = pipeline_config.gamma
scheduler.init_sigmas_for_each_stage()
def create_pipeline_stages(self, server_args: ServerArgs) -> None:
self.add_stage(InputValidationStage())
self.add_standard_text_encoding_stage()
self.add_standard_latent_preparation_stage()
# Skip standard timestep preparation — the Helios denoising stage
# handles scheduler.set_timesteps internally per-chunk with mu.
self.add_stage(
HeliosChunkedDenoisingStage(
transformer=self.get_module("transformer"),
scheduler=self.modules["scheduler"],
),
"helios_chunked_denoising_stage",
)
# Helios-specific decoding: decode each chunk's latents separately
# to avoid temporal artifacts from Wan VAE causal convolutions
self.add_stage(
HeliosDecodingStage(vae=self.get_module("vae"), pipeline=self),
"helios_decoding_stage",
)
class HeliosPyramidPipeline(HeliosPipeline):
"""Helios pyramid SR pipeline (used by Helios-Mid and Helios-Distilled)."""
pipeline_name = "HeliosPyramidPipeline"
EntryClass = [HeliosPipeline, HeliosPyramidPipeline]
@@ -0,0 +1,429 @@
"""
Hunyuan3D image-to-mesh pipeline implementation.
Shape pipeline: BeforeDenoising -> Denoising -> Export -> Save
Paint pipeline (optional): Preprocess -> TexGen -> Postprocess
"""
from __future__ import annotations
import glob
import importlib
import os
from itertools import chain
from typing import Any
import torch
import torch.nn as nn
from sglang.multimodal_gen.configs.pipeline_configs.hunyuan3d import (
Hunyuan3D2PipelineConfig,
)
from sglang.multimodal_gen.runtime.disaggregation.roles import RoleType
from sglang.multimodal_gen.runtime.loader.fsdp_load import (
load_model_from_full_model_state_dict,
set_default_torch_dtype,
)
from sglang.multimodal_gen.runtime.loader.utils import get_param_names_mapping
from sglang.multimodal_gen.runtime.pipelines_core.composed_pipeline_base import (
ComposedPipelineBase,
)
from sglang.multimodal_gen.runtime.pipelines_core.stages.model_specific_stages.hunyuan3d import (
Hunyuan3DPaintPostprocessStage,
Hunyuan3DPaintPreprocessStage,
Hunyuan3DPaintTexGenStage,
Hunyuan3DShapeBeforeDenoisingStage,
Hunyuan3DShapeDenoisingStage,
Hunyuan3DShapeExportStage,
Hunyuan3DShapeSaveStage,
)
from sglang.multimodal_gen.runtime.server_args import ServerArgs
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
logger = init_logger(__name__)
class Hunyuan3D2Pipeline(ComposedPipelineBase):
"""Hunyuan3D 2.0 image-to-mesh pipeline.
Shape pipeline: BeforeDenoising -> Denoising -> Export -> Save
Paint pipeline (optional): Preprocess -> TexGen -> Postprocess
"""
pipeline_name = "Hunyuan3D2Pipeline"
_required_config_modules = [
"hy3dshape_model",
"hy3dshape_vae",
"hy3dshape_scheduler",
"hy3dshape_conditioner",
"hy3dshape_image_processor",
]
def validate_disagg_role(self, role: RoleType) -> None:
if role == RoleType.MONOLITHIC:
return
config = self.server_args.pipeline_config
if not isinstance(config, Hunyuan3D2PipelineConfig):
raise TypeError(
"Hunyuan3D2Pipeline requires Hunyuan3D2PipelineConfig, "
f"got {type(config)}"
)
if config.paint_enable:
raise ValueError(
"Hunyuan3D2Pipeline only supports shape-only disaggregation. "
"Disable paint_enable when launching encoder/denoiser/decoder roles."
)
def _load_config(self) -> dict[str, Any]:
return {
"_class_name": self.pipeline_name,
"_diffusers_version": "0.0.0",
"hy3dshape_model": ["diffusers", "Hunyuan3DShapeModel"],
"hy3dshape_vae": ["diffusers", "Hunyuan3DShapeVAE"],
"hy3dshape_scheduler": ["diffusers", "Hunyuan3DShapeScheduler"],
"hy3dshape_conditioner": ["diffusers", "Hunyuan3DShapeConditioner"],
"hy3dshape_image_processor": ["diffusers", "Hunyuan3DShapeImageProcessor"],
}
# Class resolution
@staticmethod
def _resolve_class(target: str) -> Any:
"""Resolve a YAML target string to a Python class."""
from sglang.multimodal_gen.runtime.models.registry import ModelRegistry
cls = ModelRegistry.resolve_by_alias(target)
if cls is not None:
return cls
class_name = target.rsplit(".", 1)[-1]
try:
cls, _ = ModelRegistry.resolve_model_cls(class_name)
return cls
except Exception:
pass
from sglang.multimodal_gen.runtime.utils.mesh3d_utils import (
resolve_hunyuan3d_tool,
)
for name in (target, class_name):
tool_cls = resolve_hunyuan3d_tool(name)
if tool_cls is not None:
return tool_cls
module, cls_name = target.rsplit(".", 1)
return getattr(importlib.import_module(module, package=None), cls_name)
# Path / checkpoint resolution
@staticmethod
def _resolve_shape_dir(
model_path: str,
subfolder: str,
use_safetensors: bool,
variant: str | None,
) -> tuple[str, str]:
"""Locate (or download) the shape subfolder and return (config_path, ckpt_path)."""
local_path = os.path.join(model_path, subfolder)
if not os.path.exists(local_path):
local_path = os.path.expanduser(local_path)
if not os.path.exists(local_path):
logger.info(
"Local path %s not found, downloading from HuggingFace Hub",
local_path,
)
from huggingface_hub import snapshot_download
downloaded = snapshot_download(
repo_id=model_path,
allow_patterns=[f"{subfolder}/*"],
)
local_path = os.path.join(downloaded, subfolder)
config_path = os.path.join(local_path, "config.yaml")
if not os.path.exists(config_path):
for alt in ("config.yml", "model_config.yaml"):
alt_path = os.path.join(local_path, alt)
if os.path.exists(alt_path):
config_path = alt_path
break
if use_safetensors:
ckpt_name = (
f"model.{variant}.safetensors" if variant else "model.safetensors"
)
else:
ckpt_name = f"model-{variant}.ckpt" if variant else "model.ckpt"
ckpt_path = os.path.join(local_path, ckpt_name)
if not os.path.exists(ckpt_path):
pattern = "*.safetensors" if use_safetensors else "*.ckpt"
files = glob.glob(os.path.join(local_path, pattern))
if files:
ckpt_path = files[0]
logger.info("Config path: %s", config_path)
logger.info("Checkpoint path: %s", ckpt_path)
return config_path, ckpt_path
@staticmethod
def _resolve_paint_dir(model_path: str, subfolder: str) -> str:
"""Locate (or download) the paint subfolder and return its local path."""
local_path = os.path.join(model_path, subfolder)
if not os.path.exists(local_path):
local_path = os.path.expanduser(local_path)
if not os.path.exists(local_path):
logger.info(
"Local path %s not found, downloading from HuggingFace Hub",
local_path,
)
from huggingface_hub import snapshot_download
downloaded = snapshot_download(
repo_id=model_path,
allow_patterns=[f"{subfolder}/*"],
)
local_path = os.path.join(downloaded, subfolder)
for subdir in ("vae", "unet"):
config_file = os.path.join(local_path, subdir, "config.json")
if not os.path.exists(config_file):
raise FileNotFoundError(
f"Paint model incomplete: {config_file} not found. "
"Download the model or check network connectivity."
)
logger.info("Resolved paint model directory: %s", local_path)
return local_path
@staticmethod
def _load_and_split_checkpoint(
ckpt_path: str, use_safetensors: bool
) -> dict[str, dict[str, torch.Tensor]]:
"""Load a bundled checkpoint and split by the first '.' in each key."""
if use_safetensors:
import safetensors.torch
flat = safetensors.torch.load_file(ckpt_path, device="cpu")
ckpt: dict[str, dict[str, torch.Tensor]] = {}
for key, value in flat.items():
component = key.split(".")[0]
sub_key = key[len(component) + 1 :]
ckpt.setdefault(component, {})[sub_key] = value
return ckpt
else:
return torch.load(ckpt_path, map_location="cpu", weights_only=True)
# Component loading helpers
@classmethod
def _load_dit_model(
cls,
cfg: dict[str, Any],
weights: dict[str, torch.Tensor],
device: torch.device,
dtype: torch.dtype,
) -> nn.Module:
"""Load the DiT model using meta-device instantiation + standard weight loading."""
if "target" not in cfg:
raise KeyError("Expected key 'target' in model config.")
target_cls = cls._resolve_class(cfg["target"])
params = cfg.get("params", {})
if hasattr(target_cls, "build_config_from_params"):
dit_config = target_cls.build_config_from_params(params)
init_kwargs: dict[str, Any] = {"config": dit_config, "hf_config": {}}
else:
init_kwargs = params
with set_default_torch_dtype(dtype), torch.device("meta"):
model = target_cls(**init_kwargs)
weight_iterator = ((k, v) for k, v in weights.items())
param_names_mapping_fn = get_param_names_mapping(model.param_names_mapping)
load_model_from_full_model_state_dict(
model,
weight_iterator,
device,
dtype,
strict=False,
param_names_mapping=param_names_mapping_fn,
)
for name, p in chain(model.named_parameters(), model.named_buffers()):
if p.is_meta:
raise RuntimeError(f"Unexpected param or buffer {name} on meta device.")
if isinstance(p, nn.Parameter):
p.requires_grad = False
return model.eval()
@classmethod
def _load_simple_component(
cls,
cfg: dict[str, Any],
weights: dict[str, torch.Tensor] | None,
device: torch.device,
dtype: torch.dtype,
) -> nn.Module:
"""Load a component (VAE / conditioner) with direct instantiation + state_dict."""
if "target" not in cfg:
raise KeyError("Expected key 'target' in component config.")
target_cls = cls._resolve_class(cfg["target"])
params = cfg.get("params", {})
with set_default_torch_dtype(dtype):
component = target_cls(**params)
if weights is not None:
component.load_state_dict(weights, strict=False)
component.to(device=device, dtype=dtype)
return component.eval()
@classmethod
def _instantiate_component(cls, cfg: dict[str, Any]) -> Any:
"""Instantiate a lightweight component (scheduler / image_processor) without weights."""
if "target" not in cfg:
raise KeyError("Expected key 'target' in component config.")
target_cls = cls._resolve_class(cfg["target"])
params = cfg.get("params", {})
return target_cls(**params)
# Module loading override
def load_modules(
self,
server_args: ServerArgs,
loaded_modules: dict[str, torch.nn.Module] | None = None,
) -> dict[str, Any]:
"""Load all Hunyuan3D shape components from a bundled checkpoint."""
import yaml
from sglang.multimodal_gen.runtime.distributed import get_local_torch_device
config = server_args.pipeline_config
if not isinstance(config, Hunyuan3D2PipelineConfig):
raise TypeError(f"Expected Hunyuan3D2PipelineConfig, got {type(config)}")
model_path = config.shape_model_path or server_args.model_path
logger.info("Loading Hunyuan3D shape models from %s", model_path)
config_path, ckpt_path = self._resolve_shape_dir(
model_path,
config.shape_subfolder,
config.shape_use_safetensors,
config.shape_variant,
)
with open(config_path, "r") as f:
model_config = yaml.safe_load(f)
ckpt = self._load_and_split_checkpoint(ckpt_path, config.shape_use_safetensors)
dtype = torch.float16
if config.shape_variant and "bf16" in config.shape_variant:
dtype = torch.bfloat16
device = get_local_torch_device()
components: dict[str, Any] = {}
components["hy3dshape_model"] = self._load_dit_model(
model_config["model"], ckpt["model"], device, dtype
)
components["hy3dshape_vae"] = self._load_simple_component(
model_config["vae"], ckpt.get("vae"), device, dtype
)
components["hy3dshape_conditioner"] = self._load_simple_component(
model_config["conditioner"], ckpt.get("conditioner"), device, dtype
)
components["hy3dshape_scheduler"] = self._instantiate_component(
model_config["scheduler"]
)
components["hy3dshape_image_processor"] = self._instantiate_component(
model_config["image_processor"]
)
logger.info("All Hunyuan3D shape components loaded successfully")
if config.paint_enable:
try:
paint_dir = self._resolve_paint_dir(
server_args.model_path, config.paint_subfolder
)
components["hy3dpaint_dir"] = paint_dir
except Exception as e:
logger.warning("Failed to resolve paint model path: %s", e)
return components
# Pipeline lifecycle
def initialize_pipeline(self, server_args: ServerArgs):
config = server_args.pipeline_config
if not isinstance(config, Hunyuan3D2PipelineConfig):
raise TypeError(
"Hunyuan3D2Pipeline requires Hunyuan3D2PipelineConfig, "
f"got {type(config)}"
)
def create_pipeline_stages(self, server_args: ServerArgs):
config = server_args.pipeline_config
assert isinstance(config, Hunyuan3D2PipelineConfig)
latent_shape = tuple(config.vae_config.arch_config.latent_shape)
guidance_embed = bool(config.dit_config.arch_config.guidance_embed)
# Shape: 4 stages
self.add_stage(
stage_name="shape_before_denoising",
stage=Hunyuan3DShapeBeforeDenoisingStage(
image_processor=self.get_module("hy3dshape_image_processor"),
conditioner=self.get_module("hy3dshape_conditioner"),
scheduler=self.get_module("hy3dshape_scheduler"),
config=config,
latent_shape=latent_shape,
guidance_embed=guidance_embed,
),
)
self.add_stage(
stage_name="shape_denoising",
stage=Hunyuan3DShapeDenoisingStage(
transformer=self.get_module("hy3dshape_model"),
scheduler=self.get_module("hy3dshape_scheduler"),
),
)
self.add_stage(
stage_name="shape_export",
stage=Hunyuan3DShapeExportStage(
vae=self.get_module("hy3dshape_vae"),
config=config,
),
)
self.add_stage(
stage_name="shape_save",
stage=Hunyuan3DShapeSaveStage(config=config),
)
# Paint: 3 stages (optional)
if config.paint_enable:
self.add_stage(
stage_name="paint_preprocess",
stage=Hunyuan3DPaintPreprocessStage(config=config),
)
self.add_stage(
stage_name="paint_texgen",
stage=Hunyuan3DPaintTexGenStage(
config=config,
paint_dir=self.get_module("hy3dpaint_dir"),
),
)
self.add_stage(
stage_name="paint_postprocess",
stage=Hunyuan3DPaintPostprocessStage(config=config),
)
EntryClass = Hunyuan3D2Pipeline
@@ -0,0 +1,61 @@
# Copied and adapted from: https://github.com/hao-ai-lab/FastVideo
# SPDX-License-Identifier: Apache-2.0
"""
Hunyuan video diffusion pipeline implementation.
This module contains an implementation of the Hunyuan video diffusion pipeline
using the modular pipeline architecture.
"""
from sglang.multimodal_gen.runtime.pipelines_core.composed_pipeline_base import (
ComposedPipelineBase,
)
from sglang.multimodal_gen.runtime.pipelines_core.stages import (
InputValidationStage,
TextEncodingStage,
)
from sglang.multimodal_gen.runtime.server_args import ServerArgs
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
# TODO(will): move PRECISION_TO_TYPE to better place
logger = init_logger(__name__)
class HunyuanVideoPipeline(ComposedPipelineBase):
pipeline_name = "HunyuanVideoPipeline"
_required_config_modules = [
"text_encoder",
"text_encoder_2",
"tokenizer",
"tokenizer_2",
"vae",
"transformer",
"scheduler",
]
def create_pipeline_stages(self, server_args: ServerArgs):
self.add_stage(InputValidationStage())
self.add_stage(
TextEncodingStage(
text_encoders=[
self.get_module("text_encoder"),
self.get_module("text_encoder_2"),
],
tokenizers=[
self.get_module("tokenizer"),
self.get_module("tokenizer_2"),
],
),
"prompt_encoding_stage_primary",
)
self.add_standard_timestep_preparation_stage()
self.add_standard_latent_preparation_stage()
self.add_standard_denoising_stage()
self.add_standard_decoding_stage()
EntryClass = HunyuanVideoPipeline
@@ -0,0 +1,244 @@
# SPDX-License-Identifier: Apache-2.0
import os
from dataclasses import dataclass
from functools import lru_cache
from typing import Any, cast
from sglang.multimodal_gen.runtime.disaggregation.roles import RoleType
from sglang.multimodal_gen.runtime.pipelines_core import LoRAPipeline
from sglang.multimodal_gen.runtime.pipelines_core.composed_pipeline_base import (
ComposedPipelineBase,
)
from sglang.multimodal_gen.runtime.pipelines_core.stages.input_validation import (
InputValidationStage,
)
from sglang.multimodal_gen.runtime.pipelines_core.stages.model_specific_stages.ideogram import (
Ideogram4DecodingStage,
Ideogram4DenoisingStage,
Ideogram4TextEncodingStage,
)
from sglang.multimodal_gen.runtime.pipelines_core.stages.progressive_resolution.denoising import (
ProgressiveDenoisingStageRouter,
)
from sglang.multimodal_gen.runtime.pipelines_core.stages.progressive_resolution.ideogram import (
Ideogram4ProgressiveDenoisingStage,
)
from sglang.multimodal_gen.runtime.server_args import ServerArgs
from sglang.multimodal_gen.runtime.utils.hf_diffusers_utils import (
maybe_download_model,
verify_model_config_and_directory,
)
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
logger = init_logger(__name__)
_IDEOGRAM4_BASE_MODEL = "ideogram-ai/ideogram-4-fp8"
_IDEOGRAM4_NVFP4_COND_FILE = "diffusion_models/ideogram4_nvfp4_mixed.safetensors"
_IDEOGRAM4_NVFP4_UNCOND_FILE = (
"diffusion_models/ideogram4_unconditional_nvfp4_mixed.safetensors"
)
@dataclass(frozen=True)
class Ideogram4Nvfp4ModelResolution:
base_model_name: str
base_model_path: str
transformer_weights_path: str
unconditional_transformer_weights_path: str | None
@lru_cache(maxsize=1)
def _resolve_ideogram4_base_model_path() -> str:
return maybe_download_model(_IDEOGRAM4_BASE_MODEL, force_diffusers_model=True)
def _resolve_ideogram4_unconditional_transformer_weights_path(
transformer_weights_path: str,
) -> str | None:
if os.path.basename(transformer_weights_path) != os.path.basename(
_IDEOGRAM4_NVFP4_COND_FILE
):
return None
return os.path.join(
os.path.dirname(transformer_weights_path),
os.path.basename(_IDEOGRAM4_NVFP4_UNCOND_FILE),
)
def _resolve_ideogram4_nvfp4_transformer_weights_paths(
server_args: ServerArgs, model_path: str
) -> tuple[str, str | None]:
if server_args.transformer_weights_path is not None:
transformer_weights_path = server_args.transformer_weights_path
return (
transformer_weights_path,
_resolve_ideogram4_unconditional_transformer_weights_path(
transformer_weights_path
),
)
local_nvfp4_path = maybe_download_model(
model_path,
allow_patterns=[
_IDEOGRAM4_NVFP4_COND_FILE,
_IDEOGRAM4_NVFP4_UNCOND_FILE,
],
)
return (
os.path.join(local_nvfp4_path, _IDEOGRAM4_NVFP4_COND_FILE),
os.path.join(local_nvfp4_path, _IDEOGRAM4_NVFP4_UNCOND_FILE),
)
def resolve_ideogram4_nvfp4_model(
server_args: ServerArgs, model_path: str
) -> Ideogram4Nvfp4ModelResolution:
(
transformer_weights_path,
unconditional_transformer_weights_path,
) = _resolve_ideogram4_nvfp4_transformer_weights_paths(
server_args,
model_path,
)
return Ideogram4Nvfp4ModelResolution(
base_model_name=_IDEOGRAM4_BASE_MODEL,
base_model_path=_resolve_ideogram4_base_model_path(),
transformer_weights_path=transformer_weights_path,
unconditional_transformer_weights_path=unconditional_transformer_weights_path,
)
class Ideogram4Pipeline(LoRAPipeline, ComposedPipelineBase):
pipeline_name = "Ideogram4Pipeline"
_required_config_modules = [
"text_encoder",
"tokenizer",
"vae",
"transformer",
"unconditional_transformer",
"scheduler",
]
def _create_denoising_stage(self):
transformer = self.get_module("transformer")
unconditional_transformer = self.get_module("unconditional_transformer")
return ProgressiveDenoisingStageRouter(
standard_stage=Ideogram4DenoisingStage(
transformer=transformer,
unconditional_transformer=unconditional_transformer,
pipeline=self,
),
progressive_stage_factory=lambda: Ideogram4ProgressiveDenoisingStage(
transformer=transformer,
unconditional_transformer=unconditional_transformer,
pipeline=self,
),
)
def create_pipeline_stages(self, server_args: ServerArgs):
self.add_stage(InputValidationStage())
self.add_stage_factory(
RoleType.ENCODER,
lambda: Ideogram4TextEncodingStage(
text_encoder=self.get_module("text_encoder"),
tokenizer=self.get_module("tokenizer"),
),
"ideogram4_text_encoding_stage",
)
self.add_standard_latent_preparation_stage()
self.add_stage_factory(
RoleType.DENOISER,
self._create_denoising_stage,
"ideogram4_denoising_stage",
)
self.add_stage_factory(
RoleType.DECODER,
lambda: Ideogram4DecodingStage(vae=self.get_module("vae")),
"ideogram4_decoding_stage",
)
class Ideogram4Nvfp4Pipeline(Ideogram4Pipeline):
pipeline_name = "Ideogram4Nvfp4Pipeline"
_model_resolution: Ideogram4Nvfp4ModelResolution | None = None
def _get_model_resolution(
self,
server_args: ServerArgs | None = None,
) -> Ideogram4Nvfp4ModelResolution:
if self._model_resolution is None:
if server_args is None:
raise ValueError(
"server_args is required to resolve Ideogram4 NVFP4 paths"
)
self._model_resolution = resolve_ideogram4_nvfp4_model(
server_args,
self.model_path,
)
return self._model_resolution
def _load_config(self) -> dict[str, Any]:
model_resolution = self._get_model_resolution(self.server_args)
logger.info("Model path: %s", self.model_path)
logger.info(
"Using base model '%s' at %s for config and non-transformer components",
model_resolution.base_model_name,
model_resolution.base_model_path,
)
config = verify_model_config_and_directory(model_resolution.base_model_path)
return cast(dict[str, Any], config)
def _resolve_component_path(
self,
server_args: ServerArgs,
module_name: str,
load_module_name: str,
) -> str:
override_path = server_args.component_paths.get(module_name)
if override_path is not None:
return maybe_download_model(override_path)
component_model_path = os.path.join(
self._get_model_resolution(server_args).base_model_path,
load_module_name,
)
logger.debug("Resolved component path: %s", component_model_path)
return component_model_path
def load_modules(
self,
server_args: ServerArgs,
loaded_modules: dict | None = None,
) -> dict:
model_resolution = self._get_model_resolution(server_args)
server_args.transformer_weights_path = model_resolution.transformer_weights_path
if model_resolution.unconditional_transformer_weights_path is not None:
# The loader treats transformer_weights_path as the base DiT override.
# Route the sibling unconditional DiT weights through the generic
# per-component override map instead of hard-coding Ideogram there.
component_transformer_weights_paths = dict(
getattr(server_args, "component_transformer_weights_paths", {})
)
component_transformer_weights_paths.setdefault(
"unconditional_transformer",
model_resolution.unconditional_transformer_weights_path,
)
server_args.component_transformer_weights_paths = (
component_transformer_weights_paths
)
logger.info(
"NVFP4 transformer weights: %s",
model_resolution.transformer_weights_path,
)
logger.info(
"NVFP4 unconditional transformer weights: %s",
server_args.component_transformer_weights_paths.get(
"unconditional_transformer"
),
)
return super().load_modules(server_args, loaded_modules)
EntryClass = [Ideogram4Pipeline, Ideogram4Nvfp4Pipeline]
@@ -0,0 +1,98 @@
# SPDX-License-Identifier: Apache-2.0
from sglang.multimodal_gen.configs.pipeline_configs.joy_echo import (
JoyEchoPipelineConfig,
)
from sglang.multimodal_gen.runtime.pipelines.ltx_2_pipeline import (
_add_ltx2_front_stages,
_BaseLTX2Pipeline,
prepare_ltx2_mu,
)
from sglang.multimodal_gen.runtime.pipelines_core.stages.image_encoding import (
LTX2ImageEncodingStage,
)
from sglang.multimodal_gen.runtime.pipelines_core.stages.model_specific_stages.joy_echo import (
JoyEchoAVDecodingStage,
JoyEchoDMDDenoisingStage,
JoyEchoMemoryBankFetchStage,
JoyEchoMultishotSetupStage,
JoyEchoSigmaPreparationStage,
PairedAudioVideoMemoryBank,
)
from sglang.multimodal_gen.runtime.pipelines_core.stages.model_specific_stages.ltx_2 import (
LTX2AVLatentPreparationStage,
)
from sglang.multimodal_gen.runtime.server_args import ServerArgs
class JoyEchoPipeline(_BaseLTX2Pipeline):
pipeline_name = "JoyEchoPipeline"
is_video_pipeline = True
def __init__(self, *args, **kwargs):
self._memory_bank: PairedAudioVideoMemoryBank | None = None
self.multishot_index: int = 0
self._multishot_session_id: str | None = None
super().__init__(*args, **kwargs)
def _get_or_create_memory_bank(
self, config: JoyEchoPipelineConfig
) -> PairedAudioVideoMemoryBank:
if self._memory_bank is None:
self._memory_bank = PairedAudioVideoMemoryBank(
max_size=int(config.memory_max_size),
num_fix_frames=int(config.memory_num_fix_frames),
)
return self._memory_bank
def reset_memory_bank(self) -> None:
if self._memory_bank is not None:
self._memory_bank.memory.clear()
def create_pipeline_stages(self, server_args: ServerArgs):
config = server_args.pipeline_config
if not isinstance(config, JoyEchoPipelineConfig):
raise TypeError(
f"JoyEchoPipeline requires JoyEchoPipelineConfig, got {type(config)}"
)
memory_bank = self._get_or_create_memory_bank(config)
self.add_stage(JoyEchoMultishotSetupStage(pipeline=self))
_add_ltx2_front_stages(self)
self.add_stage(JoyEchoSigmaPreparationStage())
self.add_standard_timestep_preparation_stage(
prepare_extra_kwargs=[prepare_ltx2_mu]
)
self.add_stages(
[
LTX2AVLatentPreparationStage(
scheduler=self.get_module("scheduler"),
transformer=self.get_module("transformer"),
audio_vae=self.get_module("audio_vae"),
),
LTX2ImageEncodingStage(
vae=self.get_module("vae"),
),
JoyEchoMemoryBankFetchStage(
memory_bank=memory_bank,
vae=self.get_module("vae"),
),
JoyEchoDMDDenoisingStage(
transformer=self.get_module("transformer"),
scheduler=self.get_module("scheduler"),
vae=self.get_module("vae"),
audio_vae=self.get_module("audio_vae"),
sampler_name="euler",
pipeline=self,
),
JoyEchoAVDecodingStage(
vae=self.get_module("vae"),
audio_vae=self.get_module("audio_vae"),
vocoder=self.get_module("vocoder"),
memory_bank=memory_bank,
pipeline=self,
),
]
)
EntryClass = JoyEchoPipeline
@@ -0,0 +1,30 @@
from sglang.multimodal_gen.runtime.pipelines_core import LoRAPipeline
from sglang.multimodal_gen.runtime.pipelines_core.composed_pipeline_base import (
ComposedPipelineBase,
)
from sglang.multimodal_gen.runtime.server_args import ServerArgs
class JoyImageEditPipeline(LoRAPipeline, ComposedPipelineBase):
pipeline_name = "JoyImageEditPipeline"
_required_config_modules = [
"processor",
"scheduler",
"text_encoder",
"tokenizer",
"transformer",
"vae",
]
def create_pipeline_stages(self, server_args: ServerArgs):
self.add_standard_ti2i_stages(
vae_image_processor=None,
prompt_encoding="image_encoding",
image_processor_key="processor",
prompt_text_encoder_key="text_encoder",
)
EntryClass = JoyImageEditPipeline
@@ -0,0 +1,105 @@
"""Krea-2 text-to-image pipeline (native diffusers layout).
The released repo is diffusers-style (``model_index.json`` + ``transformer/``,
``text_encoder/``, ``vae/``, ``tokenizer/``, ``scheduler/`` subfolders), so the
base ``load_modules`` loads every component from it (the MMDiT via
``Krea2Transformer2DModel``, the Qwen3-VL text encoder, the Qwen-Image VAE, and
the ``FlowMatchEulerDiscreteScheduler``). This pipeline only adds the two
K2-specific touches the base loader can't infer: dropping the unused Qwen3-VL
vision tower (K2 conditions on text only) and building the assistant-suffix
tokenizer (``processor``), which has no ``model_index.json`` entry. The stage
chain is Krea2BeforeDenoisingStage -> DenoisingStage -> DecodingStage.
"""
from typing import Any
import torch
from transformers import Qwen2TokenizerFast
from sglang.multimodal_gen.runtime.pipelines_core import LoRAPipeline
from sglang.multimodal_gen.runtime.pipelines_core.composed_pipeline_base import (
ComposedPipelineBase,
)
from sglang.multimodal_gen.runtime.pipelines_core.stages import DenoisingStage
from sglang.multimodal_gen.runtime.pipelines_core.stages.model_specific_stages.krea2 import (
Krea2BeforeDenoisingStage,
)
from sglang.multimodal_gen.runtime.server_args import ServerArgs
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
logger = init_logger(__name__)
_TEXT_MAX_LENGTH = 512
class Krea2Pipeline(LoRAPipeline, ComposedPipelineBase):
pipeline_name = "Krea2Pipeline"
from sglang.multimodal_gen.configs.pipeline_configs.krea2 import Krea2PipelineConfig
from sglang.multimodal_gen.configs.sample.krea2 import Krea2SamplingParams
pipeline_config_cls = Krea2PipelineConfig
sampling_params_cls = Krea2SamplingParams
# Every entry is a diffusers/transformers component declared in model_index.json,
# so the base loader handles them. "processor" is added in load_modules (no entry).
_required_config_modules = [
"text_encoder",
"tokenizer",
"vae",
"transformer",
"scheduler",
]
def initialize_pipeline(self, server_args: ServerArgs):
vae_config = server_args.pipeline_config.vae_config
if hasattr(vae_config, "post_init"):
vae_config.post_init()
def load_modules(
self,
server_args: ServerArgs,
loaded_modules: dict[str, torch.nn.Module] | None = None,
) -> dict[str, Any]:
modules = super().load_modules(server_args, loaded_modules)
# K2 conditions on text only: drop the unused Qwen3-VL vision tower that the
# base loader brings in with the full Qwen3VLModel (frees its weights and
# shrinks the encoder's CPU<->GPU page). It sits on the encoder or under .model.
text_encoder = modules.get("text_encoder")
if text_encoder is not None:
for owner in (text_encoder, getattr(text_encoder, "model", None)):
if owner is not None and getattr(owner, "visual", None) is not None:
del owner.visual
break
# The conditioner appends a fixed assistant suffix, tokenized separately;
# model_index.json has no "processor" entry, so build one from tokenizer/.
tok_path = self._resolve_component_path(server_args, "tokenizer", "tokenizer")
modules["processor"] = Qwen2TokenizerFast.from_pretrained(
tok_path, max_length=_TEXT_MAX_LENGTH
)
return modules
def create_pipeline_stages(self, server_args: ServerArgs):
self.add_stage(
Krea2BeforeDenoisingStage(
vae=self.get_module("vae"),
text_encoder=self.get_module("text_encoder"),
tokenizer=self.get_module("tokenizer"),
processor=self.get_module("processor"),
transformer=self.get_module("transformer"),
scheduler=self.get_module("scheduler"),
),
"k2_before_denoising_stage",
)
self.add_stage(
DenoisingStage(
transformer=self.get_module("transformer"),
scheduler=self.get_module("scheduler"),
),
)
self.add_standard_decoding_stage()
EntryClass = [Krea2Pipeline]
@@ -0,0 +1,101 @@
# Copied and adapted from: https://github.com/hao-ai-lab/FastVideo
# Adapted from: https://github.com/Robbyant/lingbot-world
# SPDX-License-Identifier: Apache-2.0
"""
LingBot-World realtime causal DMD pipeline.
"""
from sglang.multimodal_gen.runtime.models.schedulers.scheduling_self_forcing_flow_match import (
SelfForcingFlowMatchScheduler,
)
from sglang.multimodal_gen.runtime.pipelines_core.composed_pipeline_base import (
ComposedPipelineBase,
)
from sglang.multimodal_gen.runtime.pipelines_core.lora_pipeline import LoRAPipeline
from sglang.multimodal_gen.runtime.pipelines_core.stages import (
AuxiliaryConditionEncodingStage,
DMDTimestepPreparationStage,
ImageEncodingStage,
)
from sglang.multimodal_gen.runtime.pipelines_core.stages.model_specific_stages.lingbot_world.lingbot_world_causal_denoising import (
LingBotWorldCausalDMDDenoisingStage,
)
from sglang.multimodal_gen.runtime.pipelines_core.stages.realtime import (
CausalVaeDecodingStage,
RealtimeChunkLatentPreparationStage,
RealtimeImageVAEEncodingStage,
RealtimeInputValidationStage,
RealtimeTextEncodingStage,
)
from sglang.multimodal_gen.runtime.server_args import ServerArgs
class LingBotWorldCausalDMDPipeline(LoRAPipeline, ComposedPipelineBase):
pipeline_name = "LingBotWorldCausalDMDPipeline"
_required_config_modules = [
"text_encoder",
"tokenizer",
"vae",
"transformer",
"scheduler",
"image_encoder",
"image_processor",
]
def initialize_pipeline(self, server_args: ServerArgs):
self.modules["scheduler"] = SelfForcingFlowMatchScheduler(
num_inference_steps=1000,
shift=server_args.pipeline_config.flow_shift,
sigma_min=0.0,
extra_one_step=True,
)
def create_pipeline_stages(self, server_args) -> None:
self.add_stage(RealtimeInputValidationStage())
self.add_stage(
RealtimeTextEncodingStage(
text_encoders=[self.get_module("text_encoder")],
tokenizers=[self.get_module("tokenizer")],
)
)
image_encoder = self.get_module("image_encoder", None)
image_processor = self.get_module("image_processor", None)
self.add_stage_if(
image_encoder is not None and image_processor is not None,
ImageEncodingStage(
image_encoder=image_encoder,
image_processor=image_processor,
),
)
self.add_stage(AuxiliaryConditionEncodingStage())
self.add_stage(
RealtimeImageVAEEncodingStage(
vae=self.get_module("vae"),
)
)
self.add_stage(DMDTimestepPreparationStage(self.get_module("scheduler")))
self.add_stage(
RealtimeChunkLatentPreparationStage(
scheduler=self.get_module("scheduler"),
transformer=self.get_module("transformer"),
)
)
self.add_stage(
LingBotWorldCausalDMDDenoisingStage(
transformer=self.get_module("transformer"),
scheduler=self.get_module("scheduler"),
),
)
self.add_stage(
CausalVaeDecodingStage(
vae=self.get_module("vae"),
pipeline=self,
)
)
EntryClass = LingBotWorldCausalDMDPipeline
@@ -0,0 +1,928 @@
import math
import os
import numpy as np
import torch
from sglang.multimodal_gen.configs.pipeline_configs.ltx_2 import (
STAGE_2_DISTILLED_SIGMA_VALUES as _SHARED_STAGE_2_DISTILLED_SIGMA_VALUES,
)
from sglang.multimodal_gen.configs.pipeline_configs.ltx_2 import (
LTX2PipelineConfig,
is_ltx23_native_variant,
sync_ltx23_runtime_vae_markers,
)
from sglang.multimodal_gen.configs.sample.ltx_2 import LTX23HQSamplingParams
from sglang.multimodal_gen.runtime.distributed import get_local_torch_device
from sglang.multimodal_gen.runtime.loader.component_loaders.component_loader import (
PipelineComponentLoader,
)
from sglang.multimodal_gen.runtime.managers.memory_managers.component_manager import (
ComponentResidencyStrategy,
ComponentUse,
ResidencyState,
)
from sglang.multimodal_gen.runtime.models.schedulers.scheduling_flow_match_euler_discrete import (
FlowMatchEulerDiscreteScheduler,
)
from sglang.multimodal_gen.runtime.pipelines_core.composed_pipeline_base import (
ComposedPipelineBase,
)
from sglang.multimodal_gen.runtime.pipelines_core.lora_pipeline import LoRAPipeline
from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import Req
from sglang.multimodal_gen.runtime.pipelines_core.stages import (
InputValidationStage,
TextEncodingStage,
)
from sglang.multimodal_gen.runtime.pipelines_core.stages.base import PipelineStage
from sglang.multimodal_gen.runtime.pipelines_core.stages.image_encoding import (
LTX2ImageEncodingStage,
)
from sglang.multimodal_gen.runtime.pipelines_core.stages.model_specific_stages.ltx_2 import (
LTX2AVDecodingStage,
LTX2AVDenoisingStage,
LTX2AVLatentPreparationStage,
LTX2HalveResolutionStage,
LTX2LoRASwitchStage,
LTX2RefinementStage,
LTX2TextConnectorStage,
LTX2UpsampleStage,
)
from sglang.multimodal_gen.runtime.server_args import (
LTX2_TWO_STAGE_DEVICE_MODE_CHOICES,
ServerArgs,
_normalize_ltx2_two_stage_device_mode,
)
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
logger = init_logger(__name__)
BASE_SHIFT_ANCHOR = 1024
MAX_SHIFT_ANCHOR = 4096
def _resolve_ltx2_two_stage_component_paths(
model_path: str, component_paths: dict[str, str]
) -> dict[str, str]:
resolved = dict(component_paths)
auto_resolved = []
if "spatial_upsampler" not in resolved:
spatial_candidates = [
os.path.join(model_path, "ltx-2.3-spatial-upscaler-x2-1.0.safetensors"),
os.path.join(model_path, "ltx-2.3-spatial-upscaler-x2-1.1.safetensors"),
os.path.join(model_path, "latent_upsampler"),
os.path.join(model_path, "ltx-2-spatial-upscaler-x2-1.0.safetensors"),
]
for candidate in spatial_candidates:
if os.path.exists(candidate):
resolved["spatial_upsampler"] = candidate
auto_resolved.append(f"spatial_upsampler={candidate}")
break
if "distilled_lora" not in resolved:
distilled_lora_candidates = [
os.path.join(model_path, "ltx-2.3-20b-distilled-lora-384.safetensors"),
os.path.join(model_path, "ltx-2.3-22b-distilled-lora-384.safetensors"),
os.path.join(model_path, "ltx-2-19b-distilled-lora-384.safetensors"),
]
for distilled_lora in distilled_lora_candidates:
if os.path.exists(distilled_lora):
resolved["distilled_lora"] = distilled_lora
auto_resolved.append(f"distilled_lora={distilled_lora}")
break
if auto_resolved:
logger.info(
"Auto-resolved LTX2 two-stage components: %s", ", ".join(auto_resolved)
)
return resolved
def calculate_ltx2_shift(
image_seq_len: int,
base_seq_len: int = BASE_SHIFT_ANCHOR,
max_seq_len: int = MAX_SHIFT_ANCHOR,
base_shift: float = 0.95,
max_shift: float = 2.05,
) -> float:
mm = (max_shift - base_shift) / (max_seq_len - base_seq_len)
b = base_shift - mm * base_seq_len
return image_seq_len * mm + b
def prepare_ltx2_mu(batch: Req, server_args: ServerArgs):
if is_ltx23_native_variant(server_args.pipeline_config.vae_config.arch_config):
return "mu", None
latent_num_frames = (int(batch.num_frames) - 1) // int(
server_args.pipeline_config.vae_temporal_compression
) + 1
latent_height = int(batch.height) // int(
server_args.pipeline_config.vae_scale_factor
)
latent_width = int(batch.width) // int(server_args.pipeline_config.vae_scale_factor)
video_sequence_length = latent_num_frames * latent_height * latent_width
return "mu", calculate_ltx2_shift(video_sequence_length)
def build_official_ltx2_sigmas(
steps: int,
*,
max_shift: float = 2.05,
base_shift: float = 0.95,
stretch: bool = True,
terminal: float = 0.1,
default_number_of_tokens: int = MAX_SHIFT_ANCHOR,
number_of_tokens: int | None = None,
) -> list[float]:
sigmas = torch.linspace(1.0, 0.0, steps + 1, dtype=torch.float32)
mm = (max_shift - base_shift) / (MAX_SHIFT_ANCHOR - BASE_SHIFT_ANCHOR)
b = base_shift - mm * BASE_SHIFT_ANCHOR
tokens = (
int(number_of_tokens)
if number_of_tokens is not None
else int(default_number_of_tokens)
)
sigma_shift = float(tokens) * mm + b
non_zero_mask = sigmas != 0
shifted = torch.where(
non_zero_mask,
math.exp(sigma_shift) / (math.exp(sigma_shift) + (1.0 / sigmas - 1.0)),
torch.zeros_like(sigmas),
)
if stretch:
one_minus_z = 1.0 - shifted[non_zero_mask]
if bool(torch.any(one_minus_z != 0)):
scale_factor = one_minus_z[-1] / (1.0 - terminal)
shifted[non_zero_mask] = 1.0 - (one_minus_z / scale_factor)
return shifted[:-1].tolist()
class LTX2SigmaPreparationStage(PipelineStage):
"""Prepare native LTX-2 sigma schedule before timestep setup."""
def forward(self, batch: Req, server_args: ServerArgs) -> Req:
batch.extra["ltx2_phase"] = "stage1"
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
# `LTX2Scheduler.execute(latent=empty_latent, ...)` where
# `empty_latent` is built from the **half-resolution** stage-1
# shape → resolution-aware sigma shift.
# - Non-HQ two-stage (`ti2vid_two_stages.py:145`) and
# one-stage (`ti2vid_one_stage.py:138`) call
# `LTX2Scheduler.execute(steps=...)` with no `latent` →
# falls back to `default_number_of_tokens = MAX_SHIFT_ANCHOR
# = 4096` → constant-anchor sigma shift.
if server_args.pipeline_class_name == "LTX2TwoStageHQPipeline":
# batch.height/width have already been halved by
# LTX2HalveResolutionStage, so these latents are the
# half-resolution stage-1 shape (matches `empty_latent`).
latent_num_frames = (int(batch.num_frames) - 1) // int(
server_args.pipeline_config.vae_temporal_compression
) + 1
latent_height = int(batch.height) // int(
server_args.pipeline_config.vae_scale_factor
)
latent_width = int(batch.width) // int(
server_args.pipeline_config.vae_scale_factor
)
batch.sigmas = build_official_ltx2_sigmas(
int(batch.num_inference_steps),
number_of_tokens=latent_num_frames * latent_height * latent_width,
)
else:
batch.sigmas = build_official_ltx2_sigmas(
int(batch.num_inference_steps)
)
else:
batch.sigmas = np.linspace(
1.0,
1.0 / int(batch.num_inference_steps),
int(batch.num_inference_steps),
).tolist()
return batch
def _add_ltx2_front_stages(pipeline: ComposedPipelineBase):
pipeline.add_stages(
[
InputValidationStage(),
TextEncodingStage(
text_encoders=[pipeline.get_module("text_encoder")],
tokenizers=[pipeline.get_module("tokenizer")],
),
LTX2TextConnectorStage(connectors=pipeline.get_module("connectors")),
]
)
def _add_ltx2_stage1_generation_stages(
pipeline: ComposedPipelineBase,
*,
denoising_sampler_name: str = "euler",
):
pipeline.add_stage(LTX2SigmaPreparationStage())
pipeline.add_standard_timestep_preparation_stage(
prepare_extra_kwargs=[prepare_ltx2_mu]
)
pipeline.add_stages(
[
LTX2AVLatentPreparationStage(
scheduler=pipeline.get_module("scheduler"),
transformer=pipeline.get_module("transformer"),
audio_vae=pipeline.get_module("audio_vae"),
),
LTX2ImageEncodingStage(
vae=pipeline.get_module("vae"),
),
LTX2AVDenoisingStage(
transformer=pipeline.get_module("transformer"),
scheduler=pipeline.get_module("scheduler"),
vae=pipeline.get_module("vae"),
audio_vae=pipeline.get_module("audio_vae"),
sampler_name=denoising_sampler_name,
pipeline=pipeline,
),
]
)
def _add_ltx2_decoding_stage(pipeline: ComposedPipelineBase):
pipeline.add_stage(
LTX2AVDecodingStage(
vae=pipeline.get_module("vae"),
audio_vae=pipeline.get_module("audio_vae"),
vocoder=pipeline.get_module("vocoder"),
pipeline=pipeline,
)
)
class LTX2FlowMatchScheduler(FlowMatchEulerDiscreteScheduler):
"""Override ``_time_shift_exponential`` to use torch f32 instead of numpy f64."""
def set_timesteps(
self,
num_inference_steps=None,
device=None,
sigmas=None,
mu=None,
timesteps=None,
):
if sigmas is not None and timesteps is None and mu is None:
sigmas = torch.tensor(sigmas, dtype=torch.float32, device=device)
timesteps = sigmas * self.config.num_train_timesteps
sigmas = torch.cat([sigmas, torch.zeros(1, device=sigmas.device)])
self.num_inference_steps = len(timesteps)
self.timesteps = timesteps
self.sigmas = sigmas
self._step_index = None
self._begin_index = None
return
return super().set_timesteps(
num_inference_steps=num_inference_steps,
device=device,
sigmas=sigmas,
mu=mu,
timesteps=timesteps,
)
def _time_shift_exponential(self, mu, sigma, t):
if isinstance(t, np.ndarray):
t_torch = torch.from_numpy(t).to(torch.float32)
result = math.exp(mu) / (math.exp(mu) + (1 / t_torch - 1) ** sigma)
return result.numpy()
return math.exp(mu) / (math.exp(mu) + (1 / t - 1) ** sigma)
class _BaseLTX2Pipeline(LoRAPipeline):
_required_config_modules = [
"transformer",
"text_encoder",
"tokenizer",
"scheduler",
"vae",
"audio_vae",
"vocoder",
"connectors",
]
def initialize_pipeline(self, server_args: ServerArgs):
orig = self.get_module("scheduler")
self.modules["scheduler"] = LTX2FlowMatchScheduler.from_config(orig.config)
sync_ltx23_runtime_vae_markers(
server_args.pipeline_config.vae_config.arch_config,
getattr(self.get_module("vae"), "config", None),
)
class LTX2Pipeline(_BaseLTX2Pipeline):
# Must match model_index.json `_class_name`.
pipeline_name = "LTX2Pipeline"
def create_pipeline_stages(self, server_args: ServerArgs):
_add_ltx2_front_stages(self)
_add_ltx2_stage1_generation_stages(self)
_add_ltx2_decoding_stage(self)
class LTX2TwoStageResidencyStrategy(ComponentResidencyStrategy):
name = "ltx2_original"
def __init__(self, manager: "LTX2TwoStageResidencyController") -> None:
self.manager = manager
@property
def pipeline(self) -> "LTX2TwoStagePipeline":
return self.manager.pipeline
@property
def server_args(self) -> ServerArgs:
return self.manager.server_args
def _phase(self, use: ComponentUse) -> str:
if use.phase in ("stage1", "stage2"):
return use.phase
return "stage2" if use.component_name == "transformer_2" else "stage1"
def initialize(self) -> None:
pass
def prepare_for_use(
self,
module: torch.nn.Module,
use: ComponentUse,
state: ResidencyState,
) -> None:
phase = self._phase(use)
if phase != self.manager._active_phase:
self.enter_phase(phase)
def wait_for_use(
self,
module: torch.nn.Module,
use: ComponentUse,
state: ResidencyState,
) -> None:
self.ensure_phase_ready(self._phase(use))
def finish_use(
self,
module: torch.nn.Module,
use: ComponentUse,
state: ResidencyState,
) -> None:
self.exit_phase(self._phase(use))
def prepare_after_request(
self,
module: torch.nn.Module,
use: ComponentUse,
state: ResidencyState,
) -> None:
phase = self._phase(use)
if phase != self.manager._active_phase:
self.enter_phase(phase)
def enter_phase(self, phase: str) -> bool:
return False
def exit_phase(self, phase: str | None, next_phase: str | None = None) -> None:
pass
def ensure_phase_ready(self, phase: str | None) -> None:
"""wait for the preparation to be ready"""
pass
def _ensure_on_gpu(self, module_name: str) -> None:
module = self.pipeline.get_module(module_name)
if module is None:
return
param = next(module.parameters(), None)
if param is not None and param.device.type == "cpu":
module.to(get_local_torch_device(), non_blocking=True)
class LTX2OriginalResidencyStrategy(LTX2TwoStageResidencyStrategy):
pass
class LTX2ResidentResidencyStrategy(LTX2TwoStageResidencyStrategy):
"""A residency strategy for ltx two-stage pipeline with pre-merged lora, that keep both dits always resident"""
name = "ltx2_resident"
def initialize(self) -> None:
self._ensure_on_gpu("transformer")
self._ensure_on_gpu("transformer_2")
logger.info(
"Using resident LTX-2.3 two-stage transformers mode (both DiTs stay on GPU)"
)
self.manager._active_phase = "stage1"
self.manager._sync_refinement_stage_transformer("stage1")
def enter_phase(self, phase: str) -> bool:
self.manager._sync_refinement_stage_transformer(phase)
self.manager._active_phase = phase
return True
class LTX2TwoStageResidencyController:
"""
LTX-2.3 two-stage residency controller.
It builds the selected LTX2 ComponentResidencyStrategy and keeps the
thin stage adapter methods that are specific to two-stage LoRA flow.
Modes:
- resident: keep both DiTs on GPU; phase switch is pointer rebinding only.
- original: official two-stage semantics without premerged stage-2.
"""
VALID_MODES = ("original", "resident")
def __init__(self, pipeline: "LTX2TwoStagePipeline", server_args: ServerArgs):
self.pipeline = pipeline
self.server_args = server_args
self.mode = self._resolve_mode(server_args)
self._active_phase: str | None = None
self._strategy = self._build_strategy()
@classmethod
def _resolve_mode(cls, server_args: ServerArgs) -> str:
mode = server_args.ltx2_two_stage_device_mode
if mode is None:
env_mode = os.getenv("SGLANG_LTX2_TWO_STAGE_DEVICE_MODE")
mode = (
_normalize_ltx2_two_stage_device_mode(env_mode)
if env_mode
else "original"
)
else:
mode = _normalize_ltx2_two_stage_device_mode(mode)
if mode not in cls.VALID_MODES:
raise ValueError(
f"Invalid ltx2_two_stage_device_mode={mode!r}. "
f"Expected one of {LTX2_TWO_STAGE_DEVICE_MODE_CHOICES}."
)
return mode
def _build_strategy(self) -> LTX2TwoStageResidencyStrategy:
if self.mode == "resident":
return LTX2ResidentResidencyStrategy(self)
return LTX2OriginalResidencyStrategy(self)
@property
def strategy(self) -> ComponentResidencyStrategy:
return self._strategy
@property
def should_use_premerged(self) -> bool:
"""Whether to keep a pre-merged stage-2 DiT for LTX-2.3 two-stage.
We only enable this optimization for resident native LTX-2.3 two-stage
and when users did not explicitly provide a stage-1 LoRA path
"""
return (
self.mode == "resident"
and self.pipeline._should_merge_stage2_distilled_lora(self.server_args)
and self.pipeline._stage1_lora_path is None
)
def initialize(self) -> None:
if self.mode == "original":
# maybe merge the fixed stage-1 distilled LoRA into the base once so phase switches skip per-request
# merge/unmerge.
self.pipeline._maybe_merge_stage1_distilled_into_base(self.server_args)
return
if not self.should_use_premerged:
return
self.pipeline._initialize_premerged_stage2_transformer(self.server_args)
self._strategy.initialize()
def enter_phase(self, phase: str) -> bool:
"""Switch active two-stage DiT with minimal transfer/sync overhead."""
if not self.should_use_premerged:
return False
if phase == self._active_phase:
return True
return self._strategy.enter_phase(phase)
def _sync_refinement_stage_transformer(self, phase: str) -> None:
"""Keep stage-2 refinement bound to the expected DiT for current phase."""
refinement_stage = self.pipeline.get_stage("LTX2RefinementStage")
if refinement_stage is None:
return
target_name = "transformer_2" if phase == "stage2" else "transformer"
target_transformer = self.pipeline.get_module(target_name)
if target_transformer is not None:
refinement_stage.transformer = target_transformer
class LTX2TwoStagePipeline(_BaseLTX2Pipeline):
pipeline_name = "LTX2TwoStagePipeline"
STAGE_2_DISTILLED_SIGMA_VALUES = list(_SHARED_STAGE_2_DISTILLED_SIGMA_VALUES)
STAGE_1_DISTILLED_LORA_STRENGTH = 0.0
STAGE_2_DISTILLED_LORA_STRENGTH = 1.0
STAGE_1_DENOISING_SAMPLER_NAME = "euler"
STAGE_2_DENOISING_SAMPLER_NAME = "euler"
def __init__(self, *args, **kwargs):
super().__init__(*args, **kwargs)
self._ltx2_residency = LTX2TwoStageResidencyController(self, self.server_args)
self._use_premerged_stage2_transformer = (
self._ltx2_residency.should_use_premerged
)
self._ltx2_residency.initialize()
if self._use_premerged_stage2_transformer:
self.component_residency_strategies["transformer"] = (
self._ltx2_residency.strategy
)
self.component_residency_strategies["transformer_2"] = (
self._ltx2_residency.strategy
)
@staticmethod
def _should_merge_stage2_distilled_lora(server_args: ServerArgs) -> bool:
return is_ltx23_native_variant(
server_args.pipeline_config.vae_config.arch_config
)
def _should_merge_lora_for_phase(self, phase: str) -> bool:
if phase == "stage2" and self._ltx2_residency.mode == "original":
# original mode reuses one DiT for both phases; dynamic LoRA avoids
# request-time merge/unmerge without keeping another DiT resident
return False
return self._should_merge_stage2_distilled_lora(self.server_args)
def initialize_pipeline(self, server_args: ServerArgs):
super().initialize_pipeline(server_args)
server_args.component_paths = _resolve_ltx2_two_stage_component_paths(
self.model_path, server_args.component_paths
)
upsampler_path = server_args.component_paths.get("spatial_upsampler")
if not upsampler_path:
raise ValueError(
f"{self.pipeline_name} requires --spatial-upsampler-path "
"(component_paths['spatial_upsampler'])."
)
module, memory_usage = PipelineComponentLoader.load_component(
component_name="spatial_upsampler",
component_model_path=upsampler_path,
transformers_or_diffusers="diffusers",
server_args=server_args,
)
self.modules["spatial_upsampler"] = module
self.memory_usages["spatial_upsampler"] = memory_usage
distilled_lora_path = server_args.component_paths.get("distilled_lora")
if not distilled_lora_path:
raise ValueError(
f"{self.pipeline_name} requires --distilled-lora-path "
"(component_paths['distilled_lora'])."
)
self._distilled_lora_path = distilled_lora_path
self._stage1_lora_path = server_args.lora_path
self._stage1_lora_scale = float(server_args.lora_scale)
self._active_lora_phase = None
self._active_lora_signature = None
self._use_premerged_stage2_transformer = False
# set when original mode merges stage-1 distilled LoRA into the DiT base
# once at init (see _merge_stage1_distilled_into_base).
self._stage1_distilled_in_base = False
self._stage1_distilled_base_strength: float | None = None
def _initialize_premerged_stage2_transformer(self, server_args: ServerArgs) -> None:
transformer_path = self._resolve_component_path(
server_args, "transformer", "transformer"
)
module, memory_usage = PipelineComponentLoader.load_component(
component_name="transformer_2",
component_model_path=transformer_path,
transformers_or_diffusers="diffusers",
server_args=server_args,
)
self.modules["transformer_2"] = module
self.memory_usages["transformer_2"] = memory_usage
# Reuse the canonical LoRA path used by legacy switching to reduce
# precision drift against original two-stage behavior.
self.set_lora(
lora_nickname="ltx2_stage2_distilled",
lora_path=self._distilled_lora_path,
target="transformer_2",
strength=self.STAGE_2_DISTILLED_LORA_STRENGTH,
merge_weights=True,
)
def _can_merge_stage1_distilled_into_base(self, server_args: ServerArgs) -> bool:
"""Whether original mode can merge stage-1 distilled LoRA into the base once.
For a fixed non-zero stage-1 strength (HQ only), we merge it into the base once and run
stage 2 as a dynamic delta. Requires native LTX-2.3, no user stage-1
LoRA, plain (non-FSDP/DTensor, unquantized) weights.
"""
return (
self._ltx2_residency.mode == "original"
and self._should_merge_stage2_distilled_lora(server_args)
and self._stage1_lora_path is None
and float(self.STAGE_1_DISTILLED_LORA_STRENGTH) != 0.0
and not bool(getattr(server_args, "use_fsdp_inference", False))
and getattr(server_args, "quantization", None) is None
)
def _maybe_merge_stage1_distilled_into_base(self, server_args: ServerArgs) -> None:
"""Merge stage-1 distilled LoRA into the single DiT base once at init.
Stage 1 then runs on the base; stage 2 adds a dynamic delta of
``stage2 - stage1`` strength on top. No per-request merge/unmerge.
"""
self._stage1_distilled_in_base = False
self._stage1_distilled_base_strength = None
if not self._can_merge_stage1_distilled_into_base(server_args):
return
strength = float(self.STAGE_1_DISTILLED_LORA_STRENGTH)
# Canonical merge path (handles offload/TP), then commit it as the base.
self.set_lora(
lora_nickname="ltx2_stage1_distilled",
lora_path=self._distilled_lora_path,
target="transformer",
strength=strength,
merge_weights=True,
)
if self._uses_dtensor_weights(self.lora_layers):
# Unsupported layout; undo and fall back to per-request merge.
self.deactivate_lora_weights(target="transformer")
return
for layer in self.lora_layers.values():
layer.commit_merged_as_base()
# Keep the adapter loaded for the stage-2 delta; clear merged bookkeeping.
self.is_lora_merged["transformer"] = False
self.cur_adapter_strength.pop("transformer", None)
self.cur_adapter_config.pop("transformer", None)
self._stage1_distilled_in_base = True
self._stage1_distilled_base_strength = strength
self._active_lora_phase = "stage1"
self._active_lora_signature = None
logger.info(
"Merged LTX-2 stage-1 distilled LoRA (strength=%.4f) into the DiT base; "
"stage-2 uses a dynamic delta to avoid per-request merge/unmerge.",
strength,
)
def _unmerge_stage1_distilled_from_base(self) -> None:
"""Restore the base weights and revert to per-request merging.
Used when a request overrides the stage-1 strength away from the merged
value. Subtracts the merged delta, then disables the optimization.
"""
if not self._stage1_distilled_in_base:
return
self.set_lora(
lora_nickname="ltx2_stage1_distilled",
lora_path=self._distilled_lora_path,
target="transformer",
strength=-float(self._stage1_distilled_base_strength),
merge_weights=True,
)
for layer in self.lora_layers.values():
layer.commit_merged_as_base()
self.is_lora_merged["transformer"] = False
self.cur_adapter_strength.pop("transformer", None)
self.cur_adapter_config.pop("transformer", None)
self._stage1_distilled_in_base = False
self._stage1_distilled_base_strength = None
self._active_lora_signature = None
logger.info("Restored LTX-2 base; reverting to per-request stage-1 merge.")
def _switch_lora_phase_base_merged(
self, phase: str, distilled_lora_strength: float
) -> bool:
"""Phase switch when stage-1 distilled is merged into the base, unmerge or apply dynamic lora
Returns True if handled, False to fall back to the per-request path
(after restoring the base).
"""
if phase == "stage1":
if distilled_lora_strength != self._stage1_distilled_base_strength:
self._unmerge_stage1_distilled_from_base()
return False
# Base already holds stage-1 distilled; just drop the stage-2 delta.
self.deactivate_lora_weights(target="transformer")
return True
if phase == "stage2":
delta = distilled_lora_strength - float(
self._stage1_distilled_base_strength
)
if delta == 0.0:
self.deactivate_lora_weights(target="transformer")
return True
# Dynamic delta on the merged base (base + delta == stage-2 strength);
# reuse the loaded adapter, so no reload/merge/unmerge.
self.set_lora(
lora_nickname="ltx2_stage1_distilled",
lora_path=self._distilled_lora_path,
target="transformer",
strength=delta,
merge_weights=False,
)
return True
return False
def should_skip_ltx2_lora_switch_stage(self) -> bool:
return (
self._use_premerged_stage2_transformer
and self._ltx2_residency.mode == "resident"
)
def _get_stage_distilled_lora_strength(
self, phase: str, batch: Req | None
) -> float:
if phase == "stage1":
default_strength = self.STAGE_1_DISTILLED_LORA_STRENGTH
extra_key = "ltx2_distilled_lora_strength_stage_1"
elif phase == "stage2":
default_strength = self.STAGE_2_DISTILLED_LORA_STRENGTH
extra_key = "ltx2_distilled_lora_strength_stage_2"
else:
raise ValueError(f"Unknown LTX2 two-stage LoRA phase: {phase}")
if batch is None:
return float(default_strength)
request_strength = batch.extra.get(extra_key)
if request_strength is None:
return float(default_strength)
return float(request_strength)
def _can_short_circuit_lora_switch(
self, phase: str, batch: Req | None = None
) -> bool:
distilled_lora_strength = self._get_stage_distilled_lora_strength(phase, batch)
if phase == "stage1":
return (
self._use_premerged_stage2_transformer
and self._stage1_lora_path is None
and distilled_lora_strength == 0.0
)
if phase == "stage2":
return (
self._use_premerged_stage2_transformer
and self._stage1_lora_path is None
and distilled_lora_strength == self.STAGE_2_DISTILLED_LORA_STRENGTH
)
return False
def _build_lora_switch_spec(
self, phase: str, batch: Req | None = None
) -> tuple[list[str], list[str], list[float], list[str]]:
distilled_lora_strength = self._get_stage_distilled_lora_strength(phase, batch)
lora_nicknames: list[str] = []
lora_paths: list[str] = []
lora_strengths: list[float] = []
lora_targets: list[str] = []
if phase == "stage1":
if self._stage1_lora_path:
lora_nicknames.append("ltx2_stage1_base")
lora_paths.append(self._stage1_lora_path)
lora_strengths.append(self._stage1_lora_scale)
lora_targets.append("transformer")
if distilled_lora_strength != 0.0:
lora_nicknames.append("ltx2_stage1_distilled")
lora_paths.append(self._distilled_lora_path)
lora_strengths.append(distilled_lora_strength)
lora_targets.append("transformer")
elif phase == "stage2":
if self._stage1_lora_path:
lora_nicknames.append("ltx2_stage1_base")
lora_paths.append(self._stage1_lora_path)
lora_strengths.append(self._stage1_lora_scale)
lora_targets.append("transformer")
if distilled_lora_strength != 0.0:
lora_nicknames.append("ltx2_stage2_distilled")
lora_paths.append(self._distilled_lora_path)
lora_strengths.append(distilled_lora_strength)
lora_targets.append("transformer")
else:
raise ValueError(f"Unknown LTX2 two-stage LoRA phase: {phase}")
return lora_nicknames, lora_paths, lora_strengths, lora_targets
def switch_lora_phase(self, phase: str, batch: Req | None = None) -> None:
distilled_lora_strength = self._get_stage_distilled_lora_strength(phase, batch)
phase_signature = (phase, distilled_lora_strength)
if phase_signature == self._active_lora_signature:
return
if self._stage1_distilled_in_base:
if self._switch_lora_phase_base_merged(phase, distilled_lora_strength):
self._active_lora_phase = phase
self._active_lora_signature = phase_signature
return
# Base was restored (stage-1 strength override); fall through to the
# legacy per-request merge path below.
if self._ltx2_residency.enter_phase(
phase
) and self._can_short_circuit_lora_switch(phase, batch):
self._active_lora_phase = phase
self._active_lora_signature = phase_signature
return
lora_nicknames, lora_paths, lora_strengths, lora_targets = (
self._build_lora_switch_spec(phase, batch)
)
if lora_nicknames:
set_lora_kwargs = dict(
lora_nickname=lora_nicknames,
lora_path=lora_paths,
target=lora_targets,
strength=lora_strengths,
)
if phase == "stage2":
# premerged modes keep official LTX-2.3 fused stage-2 LoRA; original
# avoids single-DiT request-time merge/unmerge with dynamic LoRA
set_lora_kwargs["merge_weights"] = self._should_merge_lora_for_phase(
phase
)
elif phase == "stage1" and self.pipeline_name == "LTX2TwoStageHQPipeline":
# Official HQ also builds stage 1 with distilled LoRA fused.
set_lora_kwargs["merge_weights"] = self._should_merge_lora_for_phase(
phase
)
self.set_lora(
**set_lora_kwargs,
)
else:
# Stage 1 must run on the base transformer weights. If stage 2 left the
# distilled adapter active, stage 1 quality drifts away from the official
# two-stage pipeline immediately.
self.deactivate_lora_weights(target="transformer")
self._active_lora_phase = phase
self._active_lora_signature = phase_signature
def create_pipeline_stages(self, server_args: ServerArgs):
_add_ltx2_front_stages(self)
self.add_stage(LTX2HalveResolutionStage())
self.add_stage(
LTX2LoRASwitchStage(pipeline=self, phase="stage1"),
)
_add_ltx2_stage1_generation_stages(
self,
denoising_sampler_name=self.STAGE_1_DENOISING_SAMPLER_NAME,
)
self.add_stages(
[
LTX2UpsampleStage(
spatial_upsampler=self.get_module("spatial_upsampler"),
vae=self.get_module("vae"),
audio_vae=self.get_module("audio_vae"),
pipeline=self,
),
(
LTX2LoRASwitchStage(pipeline=self, phase="stage2"),
"ltx2_lora_switch_stage2",
),
(
LTX2ImageEncodingStage(
vae=self.get_module("vae"),
),
"ltx2_image_encoding_stage2",
),
LTX2RefinementStage(
transformer=self.get_module("transformer"),
scheduler=self.get_module("scheduler"),
distilled_sigmas=self.STAGE_2_DISTILLED_SIGMA_VALUES,
vae=self.get_module("vae"),
audio_vae=self.get_module("audio_vae"),
pipeline=self,
sampler_name=self.STAGE_2_DENOISING_SAMPLER_NAME,
),
]
)
_add_ltx2_decoding_stage(self)
class LTX2TwoStageHQPipeline(LTX2TwoStagePipeline):
pipeline_name = "LTX2TwoStageHQPipeline"
pipeline_config_cls = LTX2PipelineConfig
sampling_params_cls = LTX23HQSamplingParams
STAGE_1_DISTILLED_LORA_STRENGTH = 0.25
STAGE_2_DISTILLED_LORA_STRENGTH = 0.5
STAGE_1_DENOISING_SAMPLER_NAME = "res2s"
STAGE_2_DENOISING_SAMPLER_NAME = "res2s"
EntryClass = [LTX2Pipeline, LTX2TwoStagePipeline, LTX2TwoStageHQPipeline]
@@ -0,0 +1,110 @@
# SPDX-License-Identifier: Apache-2.0
"""
MOVA pipeline integration (native SGLang pipeline).
"""
from __future__ import annotations
from sglang.multimodal_gen.configs.pipeline_configs.mova import MOVAPipelineConfig
from sglang.multimodal_gen.configs.sample.mova import MOVASamplingParams
from sglang.multimodal_gen.runtime.pipelines_core.composed_pipeline_base import (
ComposedPipelineBase,
)
from sglang.multimodal_gen.runtime.pipelines_core.stages import (
ImageVAEEncodingStage,
InputValidationStage,
)
from sglang.multimodal_gen.runtime.pipelines_core.stages.model_specific_stages.mova import (
MOVADecodingStage,
MOVADenoisingStage,
MOVALatentPreparationStage,
MOVATimestepPreparationStage,
)
from sglang.multimodal_gen.runtime.server_args import ServerArgs
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
logger = init_logger(__name__)
class MOVAPipeline(ComposedPipelineBase):
"""MOVA pipeline with SGLang stage orchestration."""
pipeline_name = "MOVA"
is_video_pipeline = True
_required_config_modules = [
"video_vae",
"audio_vae",
"text_encoder",
"tokenizer",
"scheduler",
"video_dit",
"video_dit_2",
"audio_dit",
"dual_tower_bridge",
]
pipeline_config_cls = MOVAPipelineConfig
sampling_params_cls = MOVASamplingParams
def initialize_pipeline(self, server_args: ServerArgs) -> None:
"""
Initialize the pipeline.
MOVA supports Context Parallel (sequence parallel) through USPAttention,
which uses Ulysses-style all-to-all communication for distributed attention.
"""
if server_args.sp_degree > 1:
logger.info(
"MOVA Context Parallel enabled with sp_degree=%d. "
"Using USPAttention for distributed self-attention.",
server_args.sp_degree,
)
def create_pipeline_stages(self, server_args: ServerArgs) -> None:
self.add_stage(InputValidationStage())
self.add_standard_text_encoding_stage()
if getattr(self.get_module("video_dit"), "require_vae_embedding", True):
self.add_stage(
ImageVAEEncodingStage(
vae=self.get_module("video_vae"),
component_name="video_vae",
)
)
self.add_stage(
MOVALatentPreparationStage(
audio_vae=self.get_module("audio_vae"),
require_vae_embedding=getattr(
self.get_module("video_dit"), "require_vae_embedding", True
),
),
"mova_latent_preparation_stage",
)
self.add_stage(
MOVATimestepPreparationStage(
scheduler=self.get_module("scheduler"),
),
"mova_timestep_preparation_stage",
)
self.add_stage(
MOVADenoisingStage(
video_dit=self.get_module("video_dit"),
video_dit_2=self.get_module("video_dit_2"),
audio_dit=self.get_module("audio_dit"),
dual_tower_bridge=self.get_module("dual_tower_bridge"),
scheduler=self.get_module("scheduler"),
),
"mova_denoising_stage",
)
self.add_stage(
MOVADecodingStage(
video_vae=self.get_module("video_vae"),
audio_vae=self.get_module("audio_vae"),
),
"mova_decoding_stage",
)
class MOVAPipelineAlias(MOVAPipeline):
pipeline_name = "MOVAPipeline"
EntryClass = [MOVAPipeline, MOVAPipelineAlias]
@@ -0,0 +1,126 @@
# SPDX-License-Identifier: Apache-2.0
from __future__ import annotations
import torch
from sglang.multimodal_gen.configs.pipeline_configs.pi05 import Pi05PipelineConfig
from sglang.multimodal_gen.configs.sample.pi05 import Pi05SamplingParams
from sglang.multimodal_gen.runtime.disaggregation.roles import RoleType
from sglang.multimodal_gen.runtime.models.vlas import Pi05PolicyModel
from sglang.multimodal_gen.runtime.pipelines_core.composed_pipeline_base import (
ComposedPipelineBase,
)
from sglang.multimodal_gen.runtime.pipelines_core.stages.model_specific_stages.pi05_preprocess import (
Pi05Preprocessor,
)
from sglang.multimodal_gen.runtime.pipelines_core.stages.vla import (
VLAActionDenoisingStage,
VLAActionPostprocessStage,
VLAObservationPreprocessStage,
VLAPrefixEncodingStage,
)
from sglang.multimodal_gen.runtime.server_args import ServerArgs
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
from sglang.multimodal_gen.runtime.vla.prefix_cache import VLAPrefixCacheManager
logger = init_logger(__name__)
class Pi05Pipeline(ComposedPipelineBase):
pipeline_name = "Pi05Pipeline"
pipeline_config_cls = Pi05PipelineConfig
sampling_params_cls = Pi05SamplingParams
_required_config_modules: list[str] = []
def validate_disagg_role(self, role: RoleType) -> None:
if role != RoleType.MONOLITHIC:
raise ValueError(
"Pi05Pipeline v1 supports same-process execution only. "
"Use prefix/action logical groups inside one worker; cross-node "
"multimodal_gen disaggregation is a v2 target."
)
def load_modules(
self,
server_args: ServerArgs,
loaded_modules: dict[str, torch.nn.Module] | None = None,
) -> dict[str, torch.nn.Module]:
if loaded_modules is not None:
return loaded_modules
pipeline_config: Pi05PipelineConfig = server_args.pipeline_config
pipeline_config.offload_prefix_image_encoder = (
pipeline_config.offload_prefix_image_encoder
or bool(server_args.image_encoder_cpu_offload)
)
pipeline_config.offload_prefix_token_embedding = (
pipeline_config.offload_prefix_token_embedding
or bool(server_args.text_encoder_cpu_offload)
)
logger.info(
"Pi05 memory config: prefix_cache=%s/%s, action_cuda_graph=%s, "
"offload_image=%s, offload_image_after_embed=%s, "
"offload_tokens=%s, offload_language_layers=%s, "
"offload_language_after_prefix=%s/%s, "
"offload_action_after_denoise=%s, empty_cache_after_prefix=%s",
pipeline_config.enable_global_prefix_cache,
pipeline_config.prefix_cache_max_entries,
pipeline_config.enable_action_cuda_graph,
pipeline_config.offload_prefix_image_encoder,
pipeline_config.offload_prefix_image_encoder_after_embed,
pipeline_config.offload_prefix_token_embedding,
pipeline_config.offload_prefix_language_layers,
pipeline_config.offload_prefix_language_layers_after_prefix,
pipeline_config.offload_prefix_language_layer_count_after_prefix,
pipeline_config.offload_action_expert_after_denoise,
pipeline_config.empty_cache_after_prefix,
)
policy_model = Pi05PolicyModel.from_pretrained(
self.model_path,
pipeline_config,
)
if (
pipeline_config.prefix_parallel_strategy
== pipeline_config.action_parallel_strategy
== "tp"
):
raise ValueError(
"VLA action expert should not share the prefix TP layout. "
"Use SP, Ulysses, Ring, DP, or monolithic fallback for the "
"action path."
)
return {
"policy_model": policy_model,
}
def initialize_pipeline(self, server_args: ServerArgs) -> None:
pipeline_config: Pi05PipelineConfig = server_args.pipeline_config
self.preprocessor = Pi05Preprocessor(pipeline_config)
self.prefix_cache = VLAPrefixCacheManager(
max_entries=pipeline_config.prefix_cache_max_entries
)
def create_pipeline_stages(self, server_args: ServerArgs):
self.add_stage(
VLAObservationPreprocessStage(self.preprocessor),
"pi05_preprocess",
)
self.add_stage(
VLAPrefixEncodingStage(
self.get_module("policy_model"),
self.prefix_cache,
),
"pi05_prefix",
)
self.add_stage(
VLAActionDenoisingStage(self.get_module("policy_model")),
"pi05_action_denoise",
)
self.add_stage(
VLAActionPostprocessStage(),
"pi05_postprocess",
)
EntryClass = Pi05Pipeline
@@ -0,0 +1,158 @@
# Copied and adapted from: https://github.com/hao-ai-lab/FastVideo
# SPDX-License-Identifier: Apache-2.0
from diffusers.image_processor import VaeImageProcessor
from sglang.multimodal_gen.runtime.disaggregation.roles import RoleType
from sglang.multimodal_gen.runtime.pipelines_core import LoRAPipeline
from sglang.multimodal_gen.runtime.pipelines_core.composed_pipeline_base import (
ComposedPipelineBase,
)
from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import Req
from sglang.multimodal_gen.runtime.pipelines_core.stages.model_specific_stages.qwen_image_layered import (
QwenImageLayeredBeforeDenoisingStage,
)
from sglang.multimodal_gen.runtime.pipelines_core.stages.progressive_resolution.qwen_image import (
QwenImageProgressiveDenoisingStage,
)
from sglang.multimodal_gen.runtime.server_args import ServerArgs
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
from sglang.multimodal_gen.utils import PRECISION_TO_TYPE
# TODO(will): move PRECISION_TO_TYPE to better place
logger = init_logger(__name__)
def calculate_shift(
image_seq_len,
base_seq_len: int = 256,
max_seq_len: int = 4096,
base_shift: float = 0.5,
max_shift: float = 1.15,
):
m = (max_shift - base_shift) / (max_seq_len - base_seq_len)
b = base_shift - m * base_seq_len
mu = image_seq_len * m + b
return mu
def prepare_mu(batch: Req, server_args: ServerArgs):
height = batch.height
width = batch.width
vae_scale_factor = server_args.pipeline_config.vae_config.vae_scale_factor
image_seq_len = (int(height) // vae_scale_factor // 2) * (
int(width) // vae_scale_factor // 2
)
mu = calculate_shift(
image_seq_len,
# hard code, since scheduler_config is not in PipelineConfig now
256,
8192,
0.5,
0.9,
)
return "mu", mu
class QwenImagePipeline(LoRAPipeline, ComposedPipelineBase):
pipeline_name = "QwenImagePipeline"
_required_config_modules = [
"text_encoder",
"tokenizer",
"vae",
"transformer",
"scheduler",
]
def create_pipeline_stages(self, server_args: ServerArgs):
self.add_standard_t2i_stages(
prepare_extra_timestep_kwargs=[prepare_mu],
progressive_denoising_stage_cls=QwenImageProgressiveDenoisingStage,
)
class QwenImageEditPipeline(LoRAPipeline, ComposedPipelineBase):
pipeline_name = "QwenImageEditPipeline"
_required_config_modules = [
"processor",
"scheduler",
"text_encoder",
"tokenizer",
"transformer",
"vae",
]
def create_pipeline_stages(self, server_args: ServerArgs):
vae_image_processor = VaeImageProcessor(
vae_scale_factor=server_args.pipeline_config.vae_config.arch_config.vae_scale_factor
* 2
)
self.add_standard_ti2i_stages(
vae_image_processor=vae_image_processor,
prompt_encoding="image_encoding",
image_processor_key="processor",
prompt_text_encoder_key="text_encoder",
prepare_extra_timestep_kwargs=[prepare_mu],
)
class QwenImageEditPlusPipeline(QwenImageEditPipeline):
pipeline_name = "QwenImageEditPlusPipeline"
def prepare_mu_layered(batch: Req, server_args: ServerArgs):
base_seqlen = 256 * 256 / 16 / 16
mu = (batch.image_latent.shape[1] / base_seqlen) ** 0.5
return "mu", mu
class QwenImageLayeredPipeline(QwenImageEditPipeline):
pipeline_name = "QwenImageLayeredPipeline"
_required_config_modules = [
"vae",
"tokenizer",
"processor",
"transformer",
"scheduler",
]
def create_pipeline_stages(self, server_args: ServerArgs):
def create_before_denoising_stage():
return QwenImageLayeredBeforeDenoisingStage(
vae=self.get_module("vae"),
text_encoder=None,
tokenizer=self.get_module("tokenizer"),
processor=self.get_module("processor"),
transformer=self.get_module("transformer"),
scheduler=self.get_module("scheduler"),
model_path=self.model_path,
vae_dtype=PRECISION_TO_TYPE[server_args.pipeline_config.vae_precision],
text_encoder_dtype=PRECISION_TO_TYPE[
server_args.pipeline_config.text_encoder_precisions[0]
],
)
self.add_stage_factory(
RoleType.ENCODER,
create_before_denoising_stage,
"QwenImageLayeredBeforeDenoisingStage",
)
self.add_standard_timestep_preparation_stage(
prepare_extra_kwargs=[prepare_mu_layered]
)
self.add_standard_denoising_stage()
self.add_standard_decoding_stage()
EntryClass = [
QwenImagePipeline,
QwenImageEditPipeline,
QwenImageEditPlusPipeline,
QwenImageLayeredPipeline,
]
@@ -0,0 +1,56 @@
# SPDX-License-Identifier: Apache-2.0
#
# SANA text-to-image pipeline.
#
# Stage order matches Flux (InputValidation -> TextEncoding -> TimestepPrep ->
# LatentPrep -> Denoising -> Decoding) rather than the add_standard_t2i_stages
# helper (which puts LatentPrep before TimestepPrep). Both orderings are
# functionally equivalent since these stages are independent.
#
# SANA uses a single text encoder (Gemma2), so only one text_encoder + tokenizer
# pair is registered — unlike Flux which has text_encoder + text_encoder_2.
# The pipeline_name must match the _class_name in HF model_index.json.
from sglang.multimodal_gen.runtime.pipelines_core import LoRAPipeline
from sglang.multimodal_gen.runtime.pipelines_core.composed_pipeline_base import (
ComposedPipelineBase,
)
from sglang.multimodal_gen.runtime.pipelines_core.stages import (
InputValidationStage,
TextEncodingStage,
)
from sglang.multimodal_gen.runtime.server_args import ServerArgs
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
logger = init_logger(__name__)
class SanaPipeline(LoRAPipeline, ComposedPipelineBase):
pipeline_name = "SanaPipeline"
_required_config_modules = [
"text_encoder",
"tokenizer",
"vae",
"transformer",
"scheduler",
]
def create_pipeline_stages(self, server_args: ServerArgs):
self.add_stage(InputValidationStage())
self.add_stage(
TextEncodingStage(
text_encoders=[self.get_module("text_encoder")],
tokenizers=[self.get_module("tokenizer")],
),
"prompt_encoding_stage_primary",
)
self.add_standard_timestep_preparation_stage()
self.add_standard_latent_preparation_stage()
self.add_standard_denoising_stage()
self.add_standard_decoding_stage()
EntryClass = SanaPipeline
@@ -0,0 +1,322 @@
# SPDX-License-Identifier: Apache-2.0
import os
from sglang.multimodal_gen.configs.pipeline_configs.sana_wm import SanaWMPipelineConfig
from sglang.multimodal_gen.configs.sample.sana_wm import SanaWMSamplingParams
from sglang.multimodal_gen.runtime.loader.utils import get_memory_usage_of_component
from sglang.multimodal_gen.runtime.pipelines_core import LoRAPipeline
from sglang.multimodal_gen.runtime.pipelines_core.composed_pipeline_base import (
ComposedPipelineBase,
)
from sglang.multimodal_gen.runtime.pipelines_core.stages import (
InputValidationStage,
)
from sglang.multimodal_gen.runtime.pipelines_core.stages.model_specific_stages.sana_wm import (
SanaWMBeforeDenoisingStage,
SanaWMDecodingStage,
SanaWMDenoisingStage,
SanaWMTextEncodingStage,
)
from sglang.multimodal_gen.runtime.pipelines_core.stages.model_specific_stages.sana_wm.refiner import (
OfficialDiffusersLTX2RefinerModule,
OfficialGemma3TextEncoderModule,
SanaWMLTX2RefinerStage,
SanaWMRefinerDecodingStage,
default_sana_wm_refiner_dtype,
sana_wm_skip_refiner_enabled,
)
from sglang.multimodal_gen.runtime.pipelines_core.stages.model_specific_stages.sana_wm.streaming import (
SanaWMStreamingDecodingStage,
SanaWMStreamingDenoisingStage,
)
from sglang.multimodal_gen.runtime.pipelines_core.stages.model_specific_stages.sana_wm.streaming_refiner import (
SanaWMStreamingRefinerStage,
)
from sglang.multimodal_gen.runtime.server_args import ServerArgs
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
# Stage-2 refiner sub-modules live under `<model_path>/refiner/...`, not at the
# model root. They're loaded manually in `initialize_pipeline` rather than via
# `_required_config_modules`, because the framework verifier resolves every
# required module key as a literal top-level subdir of the materialized model.
logger = init_logger(__name__)
class SanaWMPipeline(LoRAPipeline, ComposedPipelineBase):
"""SANA-WM TI2V pipeline (single-stage)."""
pipeline_name = "SanaWMPipeline"
pipeline_config_cls = SanaWMPipelineConfig
sampling_params_cls = SanaWMSamplingParams
_required_config_modules = [
"text_encoder",
"tokenizer",
"vae",
"transformer",
"scheduler",
]
@staticmethod
def _validate_parallelism_args(server_args: ServerArgs) -> None:
tp_size = getattr(server_args, "tp_size", 1) or 1
if tp_size != 1:
raise ValueError(
"SANA-WM does not support tensor parallelism yet. "
"Use --num-gpus with FSDP/CFG parallelism instead of "
f"--tp-size {tp_size}."
)
sp_degree = getattr(server_args, "sp_degree", 1) or 1
if sp_degree != 1:
raise ValueError(
"SANA-WM does not support temporal sequence parallelism yet. "
"Stage-1 GDN/GLUMBConvTemp span frames and require halo/state "
"exchange before latents can be sharded. Use --num-gpus with "
"FSDP/CFG parallelism instead of "
f"--sp-degree {sp_degree}."
)
def create_pipeline_stages(self, server_args: ServerArgs):
self._validate_parallelism_args(server_args)
self.add_stage(InputValidationStage())
self.add_stage(
SanaWMTextEncodingStage(
text_encoders=[self.get_module("text_encoder")],
tokenizers=[self.get_module("tokenizer")],
),
"prompt_encoding_stage",
)
self.add_stage(
SanaWMBeforeDenoisingStage(
vae=self.get_module("vae"),
transformer=self.get_module("transformer"),
scheduler=self.get_module("scheduler"),
pipeline_config=server_args.pipeline_config,
),
"sana_wm_before_denoising",
)
if getattr(server_args.pipeline_config, "streaming", False):
DenoiseStage = SanaWMStreamingDenoisingStage
else:
DenoiseStage = SanaWMDenoisingStage
self.add_stage(
DenoiseStage(
transformer=self.get_module("transformer"),
scheduler=self.get_module("scheduler"),
),
)
# Subclasses (e.g. SanaWMTwoStagePipeline) insert latent-domain stages
# between denoising and VAE decoding.
self._maybe_add_refiner_stage(server_args)
self._add_decoding_stage(server_args)
def _add_decoding_stage(self, server_args: ServerArgs = None) -> None:
if server_args is not None and getattr(
server_args.pipeline_config, "streaming", False
):
DecodeStage = SanaWMStreamingDecodingStage
else:
DecodeStage = SanaWMDecodingStage
self.add_stage(
DecodeStage(
vae=self.get_module("vae"),
pipeline=self,
component_name="vae",
),
"decoding_stage",
)
def _maybe_add_refiner_stage(self, server_args: ServerArgs) -> None:
"""Hook for subclasses; single-stage pipeline is a no-op."""
return None
class SanaWMTwoStagePipeline(SanaWMPipeline):
"""SANA-WM two-stage pipeline: SANA-WM DiT + LTX-2 latent refiner.
Stage-1 produces a coarse 720p latent; the LTX-2 refiner runs 3 Euler steps
on it before VAE decode, matching the NVlabs ``inference_sana_wm.py`` default.
"""
pipeline_name = "SanaWMTwoStagePipeline"
# Stage-2 refiner sub-modules and their on-disk layout. Loaded through the
# official Diffusers/Transformers classes because NVlabs' reference refiner
# is a narrow video-only wrapper around those modules.
_REFINER_SUB_MODULES: tuple[tuple[str, str], ...] = (
("transformer_2", "refiner/transformer"),
("connectors", "refiner/connectors"),
("text_encoder_2", "refiner/text_encoder"),
# The refiner Gemma-3 ships its tokenizer files alongside the encoder.
("tokenizer_2", "refiner/text_encoder"),
)
def initialize_pipeline(self, server_args: ServerArgs) -> None:
super().initialize_pipeline(server_args)
if sana_wm_skip_refiner_enabled():
logger.info(
"SANA-WM refiner component loading skipped by "
"SGLANG_SANA_WM_SKIP_REFINER."
)
return
self._load_refiner_modules(server_args)
def _resolve_refiner_paths(self, server_args: ServerArgs) -> tuple[str, str]:
component_paths = getattr(server_args, "component_paths", {}) or {}
refiner_root = component_paths.get(
"refiner", os.path.join(self.model_path, "refiner")
)
refiner_gemma_root = component_paths.get(
"refiner_text_encoder",
component_paths.get(
"text_encoder_2", os.path.join(refiner_root, "text_encoder")
),
)
return refiner_root, refiner_gemma_root
def _resolve_refiner_component_path(
self, server_args: ServerArgs, module_name: str, subpath: str
) -> str:
component_paths = getattr(server_args, "component_paths", {}) or {}
if module_name in component_paths:
return self._resolve_component_path(server_args, module_name, subpath)
if (
"refiner" not in component_paths
and "refiner_text_encoder" not in component_paths
):
return self._resolve_component_path(server_args, module_name, subpath)
refiner_root, refiner_gemma_root = self._resolve_refiner_paths(server_args)
if module_name in ("text_encoder_2", "tokenizer_2"):
return refiner_gemma_root
rel_subpath = subpath.removeprefix("refiner/")
return os.path.join(refiner_root, rel_subpath)
def _load_refiner_modules(self, server_args: ServerArgs) -> None:
for module_name, subpath in self._REFINER_SUB_MODULES:
component_path = self._resolve_refiner_component_path(
server_args, module_name, subpath
)
logger.info(
"SANA-WM loading refiner component %s from %s",
module_name,
component_path,
)
module, memory_usage = self._load_official_refiner_component(
module_name,
component_path,
server_args,
)
self.modules[module_name] = module
self.memory_usages[module_name] = memory_usage
@staticmethod
def _load_official_refiner_component(
module_name: str,
component_path: str,
server_args: ServerArgs,
):
"""Load SANA-WM refiner modules through the same libraries as NVlabs.
The upstream wrapper (``diffusion/refiner/diffusers_ltx2_refiner.py``)
keeps the LTX-2 transformer/connectors as Diffusers modules and only
customizes the video-only forward surface; use that path for the
quality-critical stage-2 refiner instead of the experimental native port.
"""
dtype = default_sana_wm_refiner_dtype(server_args)
if module_name == "transformer_2":
from diffusers.models.transformers.transformer_ltx2 import (
LTX2VideoTransformer3DModel,
)
module = LTX2VideoTransformer3DModel.from_pretrained(
component_path,
torch_dtype=dtype,
).eval()
module = OfficialDiffusersLTX2RefinerModule(module)
elif module_name == "connectors":
from diffusers.pipelines.ltx2 import LTX2TextConnectors
module = LTX2TextConnectors.from_pretrained(
component_path,
torch_dtype=dtype,
).eval()
elif module_name == "text_encoder_2":
from transformers import Gemma3ForConditionalGeneration
module = Gemma3ForConditionalGeneration.from_pretrained(
component_path,
torch_dtype=dtype,
low_cpu_mem_usage=True,
).eval()
module = OfficialGemma3TextEncoderModule(module)
elif module_name == "tokenizer_2":
from transformers import AutoTokenizer
module = AutoTokenizer.from_pretrained(component_path)
else:
raise ValueError(f"Unsupported SANA-WM refiner component: {module_name}")
memory_usage = get_memory_usage_of_component(module)
logger.info(
"Loaded %s: %s (official native version). model size: %s GB",
module_name,
module.__class__.__name__,
memory_usage if memory_usage is not None else "NA",
)
return module, memory_usage or 0.0
def _maybe_add_refiner_stage(self, server_args: ServerArgs) -> None:
if sana_wm_skip_refiner_enabled():
return
pc = server_args.pipeline_config
common = dict(
transformer=self.get_module("transformer_2"),
connectors=self.get_module("connectors"),
text_encoder=self.get_module("text_encoder_2"),
tokenizer=self.get_module("tokenizer_2"),
dtype=default_sana_wm_refiner_dtype(server_args),
)
if getattr(pc, "streaming", False) and getattr(pc, "refiner_chunked", True):
stage = SanaWMStreamingRefinerStage(
**common,
block_size=int(getattr(pc, "refiner_block_size", 3)),
kv_max_frames=int(getattr(pc, "refiner_kv_max_frames", 11)),
sink_size=int(getattr(pc, "sink_size", 1)),
seed=int(getattr(pc, "refiner_seed", 42)),
)
else:
stage = SanaWMLTX2RefinerStage(**common)
self.add_stage(stage, "sana_wm_refiner")
def _add_decoding_stage(self, server_args: ServerArgs = None) -> None:
# Streaming and skip-refiner both route to the base decode
# (SanaWMStreamingDecodingStage / dense decode); otherwise dense refiner-decode.
streaming = server_args is not None and getattr(
server_args.pipeline_config, "streaming", False
)
if streaming or sana_wm_skip_refiner_enabled():
return super()._add_decoding_stage(server_args)
self.add_stage(
SanaWMRefinerDecodingStage(
vae=self.get_module("vae"),
pipeline=self,
component_name="vae",
),
"decoding_stage",
)
EntryClass = [SanaWMPipeline, SanaWMTwoStagePipeline]
@@ -0,0 +1,135 @@
# SPDX-License-Identifier: Apache-2.0
from __future__ import annotations
from sglang.multimodal_gen.configs.pipeline_configs.sana_wm import (
SanaWMRealtimeConfig,
)
from sglang.multimodal_gen.runtime.pipelines.sana_wm_pipeline import (
SanaWMTwoStagePipeline,
)
from sglang.multimodal_gen.runtime.pipelines_core.stages.model_specific_stages.sana_wm import (
SanaWMTextEncodingStage,
)
from sglang.multimodal_gen.runtime.pipelines_core.stages.model_specific_stages.sana_wm.realtime_chain import (
SanaWMCameraCondStage,
SanaWMCausalDecodeChainStage,
SanaWMChunkedRefinerChainStage,
SanaWMCondFrameEncodeStage,
SanaWMRealtimeLatentPrepStage,
)
from sglang.multimodal_gen.runtime.pipelines_core.stages.model_specific_stages.sana_wm.refiner import (
default_sana_wm_refiner_dtype,
sana_wm_skip_refiner_enabled,
)
from sglang.multimodal_gen.runtime.pipelines_core.stages.model_specific_stages.sana_wm.streaming import (
SanaWMStreamingDenoisingStage,
)
from sglang.multimodal_gen.runtime.pipelines_core.stages.model_specific_stages.sana_wm.streaming_refiner import (
SanaWMStreamingRefinerStage,
)
from sglang.multimodal_gen.runtime.pipelines_core.stages.realtime import (
RealtimeInputValidationStage,
RealtimeTextEncodingStage,
)
from sglang.multimodal_gen.runtime.server_args import ServerArgs
from sglang.multimodal_gen.runtime.utils.hf_diffusers_utils import maybe_download_model
DEFAULT_SANA_WM_TEXT_ENCODER = "Efficient-Large-Model/gemma-2-2b-it"
class SanaWMRealtimeTextEncodingStage(
RealtimeTextEncodingStage, SanaWMTextEncodingStage
):
"""Realtime text encoding using SANA-WM prompt processing.
MRO contract: ``RealtimeTextEncodingStage.forward`` (per-session cache) calls
``super().forward``, which must resolve to ``SanaWMTextEncodingStage.forward``.
This preserves the chi prompt prefix and official prompt window used by the
batch path.
"""
class SanaWMRealtimePipeline(SanaWMTwoStagePipeline):
"""SANA-WM realtime interactive pipeline.
Extends the two-stage pipeline to inherit refiner sub-module loading (``transformer_2`` /
``connectors`` / ``text_encoder_2`` / ``tokenizer_2``). The streaming refiner stage built
here is purely a carrier of those modules handed to ``SanaWMRealtimeStage``, not added to
the stage list (the realtime stage drives the incremental stage-1 session + chunked refiner
runner per user action).
"""
pipeline_name = "SanaWMRealtimePipeline"
is_video_pipeline = True
# Must be the realtime config so get_realtime_model_adapter() resolves the
# SANA-WM adapter (the realtime registry keys on SanaWMRealtimeConfig).
pipeline_config_cls = SanaWMRealtimeConfig
def _resolve_component_path(
self, server_args: ServerArgs, module_name: str, load_module_name: str
) -> str:
if (
module_name in {"text_encoder", "tokenizer"}
and module_name not in server_args.component_paths
):
return maybe_download_model(DEFAULT_SANA_WM_TEXT_ENCODER)
return super()._resolve_component_path(
server_args,
module_name,
load_module_name,
)
def _build_realtime_refiner_stage(self, server_args: ServerArgs):
"""Build the chunked streaming refiner carrier when refiner modules exist."""
if sana_wm_skip_refiner_enabled():
return None
if self.get_module("transformer_2") is None:
return None
pc = server_args.pipeline_config
return SanaWMStreamingRefinerStage(
transformer=self.get_module("transformer_2"),
connectors=self.get_module("connectors"),
text_encoder=self.get_module("text_encoder_2"),
tokenizer=self.get_module("tokenizer_2"),
dtype=default_sana_wm_refiner_dtype(server_args),
block_size=int(getattr(pc, "refiner_block_size", 3)),
kv_max_frames=int(getattr(pc, "refiner_kv_max_frames", 11)),
sink_size=int(getattr(pc, "sink_size", 1)),
seed=int(getattr(pc, "refiner_seed", 42)),
)
def create_pipeline_stages(self, server_args: ServerArgs):
refiner_stage = self._build_realtime_refiner_stage(server_args)
common = dict(
transformer=self.get_module("transformer"),
vae=self.get_module("vae"),
model_path=self.model_path,
)
self.add_stage(RealtimeInputValidationStage())
self.add_stage(
SanaWMRealtimeTextEncodingStage(
text_encoders=[self.get_module("text_encoder")],
tokenizers=[self.get_module("tokenizer")],
)
)
self.add_stage(SanaWMCondFrameEncodeStage(**common))
self.add_stage(
SanaWMRealtimeLatentPrepStage(
use_refiner=refiner_stage is not None, **common
)
)
self.add_stage(SanaWMCameraCondStage(**common))
self.add_stage(
SanaWMStreamingDenoisingStage(
transformer=self.get_module("transformer"), keep_resident=True
)
)
if refiner_stage is not None:
self.add_stage(
SanaWMChunkedRefinerChainStage(refiner_stage=refiner_stage, **common)
)
self.add_stage(SanaWMCausalDecodeChainStage(**common))
EntryClass = SanaWMRealtimePipeline
@@ -0,0 +1,111 @@
# Copied and adapted from: https://github.com/hao-ai-lab/FastVideo
# SPDX-License-Identifier: Apache-2.0
"""StableDiffusion3 pipeline implementation."""
import torch
from sglang.multimodal_gen.runtime.pipelines_core.composed_pipeline_base import (
ComposedPipelineBase,
)
from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import Req
from sglang.multimodal_gen.runtime.pipelines_core.stages import (
InputValidationStage,
PipelineStage,
TextEncodingStage,
)
from sglang.multimodal_gen.runtime.server_args import ServerArgs
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
logger = init_logger(__name__)
class SD3ConditioningStage(PipelineStage):
"""Merge CLIP-T, CLIP-G and T5 embeddings into unified prompt/pooled tensors."""
@torch.no_grad()
def forward(self, batch: Req, server_args: ServerArgs) -> Req:
batch.prompt_embeds, batch.pooled_embeds = self._merge(
batch.prompt_embeds, batch.pooled_embeds
)
if batch.do_classifier_free_guidance:
batch.negative_prompt_embeds, batch.neg_pooled_embeds = self._merge(
batch.negative_prompt_embeds, batch.neg_pooled_embeds
)
return batch
@staticmethod
def _merge(
embeds_list: list[torch.Tensor],
pooled_list: list[torch.Tensor],
) -> tuple[list[torch.Tensor], list[torch.Tensor]]:
"""Merge 3 encoder outputs into unified prompt/pooled tensors.
SD3-medium uses exactly 3 text encoders (CLIP-L, CLIP-G, T5).
Returns single-element lists to match the batch field format expected
by downstream stages (get_pos_prompt_embeds accesses index [0]).
"""
if len(embeds_list) != 3:
raise ValueError(
f"SD3 requires exactly 3 prompt embedding tensors, got {len(embeds_list)}."
)
if len(pooled_list) < 2:
raise ValueError(
f"SD3 requires at least 2 pooled embedding tensors, got {len(pooled_list)}."
)
clipt, clipg, t5 = embeds_list
clip_merged = torch.cat([clipt, clipg], dim=-1)
clip_merged = torch.nn.functional.pad(
clip_merged, (0, t5.shape[-1] - clip_merged.shape[-1])
)
merged_embeds = [torch.cat([clip_merged, t5], dim=-2)]
merged_pooled = [torch.cat([pooled_list[0], pooled_list[1]], dim=-1)]
return merged_embeds, merged_pooled
class StableDiffusion3Pipeline(ComposedPipelineBase):
"""StableDiffusion3 pipeline implementation."""
pipeline_name = "StableDiffusion3Pipeline"
_required_config_modules = [
"text_encoder",
"text_encoder_2",
"text_encoder_3",
"tokenizer",
"tokenizer_2",
"tokenizer_3",
"vae",
"transformer",
"scheduler",
]
def create_pipeline_stages(self, server_args: ServerArgs):
self.add_stage(InputValidationStage())
self.add_stage(
TextEncodingStage(
text_encoders=[
self.get_module("text_encoder"),
self.get_module("text_encoder_2"),
self.get_module("text_encoder_3"),
],
tokenizers=[
self.get_module("tokenizer"),
self.get_module("tokenizer_2"),
self.get_module("tokenizer_3"),
],
),
"prompt_encoding_stage_primary",
)
self.add_stage(SD3ConditioningStage())
self.add_standard_timestep_preparation_stage()
self.add_standard_latent_preparation_stage()
self.add_standard_denoising_stage()
self.add_standard_decoding_stage()
EntryClass = StableDiffusion3Pipeline
@@ -0,0 +1,54 @@
# Copied and adapted from: https://github.com/hao-ai-lab/FastVideo
# SPDX-License-Identifier: Apache-2.0
"""
Wan causal DMD pipeline implementation.
This module wires the causal DMD denoising stage into the modular pipeline.
"""
from sglang.multimodal_gen.runtime.pipelines_core.composed_pipeline_base import (
ComposedPipelineBase,
)
from sglang.multimodal_gen.runtime.pipelines_core.lora_pipeline import LoRAPipeline
# isort: off
from sglang.multimodal_gen.runtime.pipelines_core.stages import (
CausalDMDDenoisingStage,
InputValidationStage,
)
from sglang.multimodal_gen.runtime.server_args import ServerArgs
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
# isort: on
logger = init_logger(__name__)
class WanCausalDMDPipeline(LoRAPipeline, ComposedPipelineBase):
pipeline_name = "WanCausalDMDPipeline"
_required_config_modules = [
"text_encoder",
"tokenizer",
"vae",
"transformer",
"scheduler",
]
def create_pipeline_stages(self, server_args: ServerArgs) -> None:
self.add_stage(InputValidationStage())
self.add_standard_text_encoding_stage()
self.add_standard_latent_preparation_stage()
self.add_stage(
CausalDMDDenoisingStage(
transformer=self.get_module("transformer"),
scheduler=self.get_module("scheduler"),
),
)
self.add_standard_decoding_stage()
EntryClass = WanCausalDMDPipeline
@@ -0,0 +1,77 @@
# Copied and adapted from: https://github.com/hao-ai-lab/FastVideo
# SPDX-License-Identifier: Apache-2.0
"""
Wan video diffusion pipeline implementation.
This module contains an implementation of the Wan video diffusion pipeline
using the modular pipeline architecture.
"""
from sglang.multimodal_gen.runtime.models.schedulers.scheduling_flow_match_euler_discrete import (
FlowMatchEulerDiscreteScheduler,
)
from sglang.multimodal_gen.runtime.pipelines_core.composed_pipeline_base import (
ComposedPipelineBase,
)
from sglang.multimodal_gen.runtime.pipelines_core.lora_pipeline import LoRAPipeline
from sglang.multimodal_gen.runtime.server_args import ServerArgs
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
# isort: off
from sglang.multimodal_gen.runtime.pipelines_core.stages import (
DmdDenoisingStage,
InputValidationStage,
)
# isort: on
logger = init_logger(__name__)
class WanDMDPipeline(LoRAPipeline, ComposedPipelineBase):
"""
Wan video diffusion pipeline with LoRA support.
"""
pipeline_name = "WanDMDPipeline"
_required_config_modules = [
"text_encoder",
"tokenizer",
"vae",
"transformer",
"scheduler",
]
def initialize_pipeline(self, server_args: ServerArgs):
self.modules["scheduler"] = FlowMatchEulerDiscreteScheduler(
shift=server_args.pipeline_config.flow_shift
)
def create_pipeline_stages(self, server_args: ServerArgs) -> None:
self.add_stages(
[
InputValidationStage(),
]
)
self.add_standard_text_encoding_stage()
self.add_standard_timestep_preparation_stage()
self.add_standard_latent_preparation_stage()
self.add_stages(
[
DmdDenoisingStage(
transformer=self.get_module("transformer"),
scheduler=self.get_module("scheduler"),
),
]
)
self.add_standard_decoding_stage()
EntryClass = WanDMDPipeline
@@ -0,0 +1,54 @@
# Copied and adapted from: https://github.com/hao-ai-lab/FastVideo
# SPDX-License-Identifier: Apache-2.0
"""
Wan video diffusion pipeline implementation.
This module contains an implementation of the Wan video diffusion pipeline
using the modular pipeline architecture.
"""
from sglang.multimodal_gen.runtime.models.schedulers.scheduling_flow_match_euler_discrete import (
FlowMatchEulerDiscreteScheduler,
)
from sglang.multimodal_gen.runtime.pipelines_core.composed_pipeline_base import (
ComposedPipelineBase,
)
from sglang.multimodal_gen.runtime.pipelines_core.lora_pipeline import LoRAPipeline
from sglang.multimodal_gen.runtime.pipelines_core.stages import DmdDenoisingStage
from sglang.multimodal_gen.runtime.server_args import ServerArgs
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
logger = init_logger(__name__)
class WanImageToVideoDmdPipeline(LoRAPipeline, ComposedPipelineBase):
pipeline_name = "WanImageToVideoDmdPipeline"
_required_config_modules = [
"text_encoder",
"tokenizer",
"vae",
"transformer",
"scheduler",
"image_encoder",
"image_processor",
]
def initialize_pipeline(self, server_args: ServerArgs):
self.modules["scheduler"] = FlowMatchEulerDiscreteScheduler(
shift=server_args.pipeline_config.flow_shift
)
def create_pipeline_stages(self, server_args: ServerArgs):
self.add_standard_ti2v_stages(
image_vae_encoding_position="after_latent",
denoising_stage_factory=lambda: DmdDenoisingStage(
transformer=self.get_module("transformer"),
scheduler=self.get_module("scheduler"),
transformer_2=self.get_module("transformer_2"),
),
)
EntryClass = WanImageToVideoDmdPipeline
@@ -0,0 +1,46 @@
# Copied and adapted from: https://github.com/hao-ai-lab/FastVideo
# SPDX-License-Identifier: Apache-2.0
"""
Wan video diffusion pipeline implementation.
This module contains an implementation of the Wan video diffusion pipeline
using the modular pipeline architecture.
"""
from sglang.multimodal_gen.runtime.models.schedulers.scheduling_flow_unipc_multistep import (
FlowUniPCMultistepScheduler,
)
from sglang.multimodal_gen.runtime.pipelines_core.composed_pipeline_base import (
ComposedPipelineBase,
)
from sglang.multimodal_gen.runtime.pipelines_core.lora_pipeline import LoRAPipeline
from sglang.multimodal_gen.runtime.server_args import ServerArgs
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
logger = init_logger(__name__)
class WanImageToVideoPipeline(LoRAPipeline, ComposedPipelineBase):
pipeline_name = "WanImageToVideoPipeline"
_required_config_modules = [
"text_encoder",
"tokenizer",
"vae",
"transformer",
"scheduler",
"image_encoder",
"image_processor",
]
def initialize_pipeline(self, server_args: ServerArgs):
self.modules["scheduler"] = FlowUniPCMultistepScheduler(
shift=server_args.pipeline_config.flow_shift
)
def create_pipeline_stages(self, server_args: ServerArgs):
self.add_standard_ti2v_stages()
EntryClass = WanImageToVideoPipeline
@@ -0,0 +1,60 @@
# Copied and adapted from: https://github.com/hao-ai-lab/FastVideo
# SPDX-License-Identifier: Apache-2.0
"""
Wan video diffusion pipeline implementation.
This module contains an implementation of the Wan video diffusion pipeline
using the modular pipeline architecture.
"""
from sglang.multimodal_gen.runtime.models.schedulers.scheduling_flow_unipc_multistep import (
FlowUniPCMultistepScheduler,
)
from sglang.multimodal_gen.runtime.pipelines_core.composed_pipeline_base import (
ComposedPipelineBase,
)
from sglang.multimodal_gen.runtime.pipelines_core.lora_pipeline import LoRAPipeline
from sglang.multimodal_gen.runtime.pipelines_core.stages import (
InputValidationStage,
)
from sglang.multimodal_gen.runtime.pipelines_core.stages.progressive_resolution.wan import (
WanProgressiveDenoisingStage,
)
from sglang.multimodal_gen.runtime.server_args import ServerArgs
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
logger = init_logger(__name__)
class WanPipeline(LoRAPipeline, ComposedPipelineBase):
"""
Wan video diffusion pipeline with LoRA support.
"""
pipeline_name = "WanPipeline"
_required_config_modules = [
"text_encoder",
"tokenizer",
"vae",
"transformer",
"scheduler",
]
def initialize_pipeline(self, server_args: ServerArgs):
# We use UniPCMScheduler from Wan2.1 official repo, not the one in diffusers.
self.modules["scheduler"] = FlowUniPCMultistepScheduler(
shift=server_args.pipeline_config.flow_shift
)
def create_pipeline_stages(self, server_args: ServerArgs) -> None:
self.add_stage(InputValidationStage())
self.add_standard_text_encoding_stage()
self.add_standard_latent_preparation_stage()
self.add_standard_timestep_preparation_stage()
self.add_progressive_denoising_stage(WanProgressiveDenoisingStage)
self.add_standard_decoding_stage()
EntryClass = WanPipeline
@@ -0,0 +1,66 @@
# Copied and adapted from: https://github.com/hao-ai-lab/FastVideo
# SPDX-License-Identifier: Apache-2.0
from sglang.multimodal_gen.runtime.pipelines_core import LoRAPipeline, Req
from sglang.multimodal_gen.runtime.pipelines_core.composed_pipeline_base import (
ComposedPipelineBase,
)
from sglang.multimodal_gen.runtime.pipelines_core.stages.progressive_resolution.zimage import (
ZImageProgressiveDenoisingStage,
)
from sglang.multimodal_gen.runtime.server_args import ServerArgs
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
logger = init_logger(__name__)
def calculate_shift(
image_seq_len,
base_seq_len: int = 256,
max_seq_len: int = 4096,
base_shift: float = 0.5,
max_shift: float = 1.15,
):
m = (max_shift - base_shift) / (max_seq_len - base_seq_len)
b = base_shift - m * base_seq_len
mu = image_seq_len * m + b
return mu
def prepare_mu(batch: Req, server_args: ServerArgs):
height = batch.height
width = batch.width
vae_scale_factor = server_args.pipeline_config.vae_config.vae_scale_factor
image_seq_len = ((int(height) // vae_scale_factor) // 2) * (
(int(width) // vae_scale_factor) // 2
)
mu = calculate_shift(
image_seq_len,
# hard code, since scheduler_config is not in PipelineConfig now
256,
4096,
0.5,
1.15,
)
return "mu", mu
class ZImagePipeline(LoRAPipeline, ComposedPipelineBase):
pipeline_name = "ZImagePipeline"
_required_config_modules = [
"text_encoder",
"tokenizer",
"vae",
"transformer",
"scheduler",
]
def create_pipeline_stages(self, server_args: ServerArgs):
self.add_standard_t2i_stages(
prepare_extra_timestep_kwargs=[prepare_mu],
progressive_denoising_stage_cls=ZImageProgressiveDenoisingStage,
)
EntryClass = ZImagePipeline