# Copied and adapted from: https://github.com/hao-ai-lab/FastVideo # SPDX-License-Identifier: Apache-2.0 from dataclasses import dataclass, field from typing import Any from sglang.multimodal_gen.configs.models.base import ArchConfig, ModelConfig from sglang.multimodal_gen.runtime.layers.quantization import QuantizationConfig from sglang.multimodal_gen.runtime.platforms import AttentionBackendEnum @dataclass class DiTArchConfig(ArchConfig): _fsdp_shard_conditions: list = field(default_factory=list) _compile_conditions: list = field(default_factory=list) # convert weights name from HF-format to SGLang-dit-format param_names_mapping: dict = field(default_factory=dict) # convert weights name from misc-format to HF-format # usually applicable if the LoRA is trained with official repo implementation lora_param_names_mapping: dict = field(default_factory=dict) # Reverse mapping for saving checkpoints: custom -> hf reverse_param_names_mapping: dict = field(default_factory=dict) _supported_attention_backends: set[AttentionBackendEnum] = field( default_factory=lambda: { AttentionBackendEnum.SLIDING_TILE_ATTN, AttentionBackendEnum.SAGE_ATTN, AttentionBackendEnum.FA, AttentionBackendEnum.AITER, AttentionBackendEnum.AITER_SAGE, AttentionBackendEnum.TORCH_SDPA, AttentionBackendEnum.VIDEO_SPARSE_ATTN, AttentionBackendEnum.SPARSE_VIDEO_GEN_2_ATTN, AttentionBackendEnum.VMOBA_ATTN, AttentionBackendEnum.SAGE_ATTN_3, AttentionBackendEnum.LASER_ATTN, AttentionBackendEnum.BLOCK_SPARSE_ATTN, AttentionBackendEnum.RAIN_FUSION_ATTN, } ) hidden_size: int = 0 num_attention_heads: int = 0 num_channels_latents: int = 0 exclude_lora_layers: list[str] = field(default_factory=list) boundary_ratio: float | None = None def __post_init__(self) -> None: if not self._compile_conditions: self._compile_conditions = self._fsdp_shard_conditions.copy() @dataclass class DiTConfig(ModelConfig): arch_config: DiTArchConfig = field(default_factory=DiTArchConfig) # sglang-diffusion DiT-specific parameters prefix: str = "" quant_config: QuantizationConfig | None = None torch_compile_mode: str = "max-autotune-no-cudagraphs" @staticmethod def add_cli_args(parser: Any, prefix: str = "dit-config") -> Any: """Add CLI arguments for DiTConfig fields""" parser.add_argument( f"--{prefix}.prefix", type=str, dest=f"{prefix.replace('-', '_')}.prefix", default=DiTConfig.prefix, help="Prefix for the DiT model", ) parser.add_argument( f"--{prefix}.quant-config", type=str, dest=f"{prefix.replace('-', '_')}.quant_config", default=None, help="Quantization configuration for the DiT model", ) parser.add_argument( f"--{prefix}.torch-compile-mode", type=str, dest=f"{prefix.replace('-', '_')}.torch_compile_mode", default=DiTConfig.torch_compile_mode, help="torch.compile mode for the DiT model", ) return parser