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
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:
@@ -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
|
||||
Reference in New Issue
Block a user