# 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.platforms import AttentionBackendEnum @dataclass class AdapterArchConfig(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) # 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.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 AdapterConfig(ModelConfig): arch_config: AdapterArchConfig = field(default_factory=AdapterArchConfig) # sglang-diffusion Adapter-specific parameters prefix: str = "" @staticmethod def add_cli_args(parser: Any, prefix: str = "dit-config") -> Any: """Add CLI arguments for AdapterConfig fields""" parser.add_argument( f"--{prefix}.prefix", type=str, dest=f"{prefix.replace('-', '_')}.prefix", default=AdapterConfig.prefix, help="Prefix for the Adapter", ) return parser