WIP: initial multimodal-gen support (#12484)

Co-authored-by: yhyang201 <yhyang201@gmail.com>
Co-authored-by: yizhang2077 <1109276519@qq.com>
Co-authored-by: Xinyuan Tong <xinyuantong.cs@gmail.com>
Co-authored-by: ispobock <ispobaoke@gmail.com>
Co-authored-by: JiLi <leege233@gmail.com>
Co-authored-by: CHEN Xi <78632976+RubiaCx@users.noreply.github.com>
Co-authored-by: laixin <xielx@shanghaitech.edu.cn>
Co-authored-by: SolitaryThinker <wlsaidhi@gmail.com>
Co-authored-by: jzhang38 <a1286225768@gmail.com>
Co-authored-by: BrianChen1129 <yongqichcd@gmail.com>
Co-authored-by: Kevin Lin <42618777+kevin314@users.noreply.github.com>
Co-authored-by: Edenzzzz <wtan45@wisc.edu>
Co-authored-by: rlsu9 <r3su@ucsd.edu>
Co-authored-by: Jinzhe Pan <48981407+eigensystem@users.noreply.github.com>
Co-authored-by: foreverpiano <pianoqwz@qq.com>
Co-authored-by: RandNMR73 <notomatthew31@gmail.com>
Co-authored-by: PorridgeSwim <yz3883@columbia.edu>
Co-authored-by: Jiali Chen <90408393+gary-chenjl@users.noreply.github.com>
This commit is contained in:
Mick
2025-11-05 12:28:52 -08:00
committed by GitHub
co-authored by yhyang201 yizhang2077 Xinyuan Tong ispobock JiLi CHEN Xi laixin SolitaryThinker jzhang38 BrianChen1129 Kevin Lin Edenzzzz rlsu9 Jinzhe Pan foreverpiano RandNMR73 PorridgeSwim Jiali Chen
parent 4fe53e5888
commit 7bc1dae095
249 changed files with 63750 additions and 11 deletions
@@ -0,0 +1,69 @@
# 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)
param_names_mapping: dict = field(default_factory=dict)
reverse_param_names_mapping: dict = field(default_factory=dict)
lora_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.FA3,
AttentionBackendEnum.TORCH_SDPA,
AttentionBackendEnum.VIDEO_SPARSE_ATTN,
AttentionBackendEnum.VMOBA_ATTN,
AttentionBackendEnum.SAGE_ATTN_THREE,
}
)
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)
# sgl-diffusionDiT-specific parameters
prefix: str = ""
quant_config: QuantizationConfig | None = None
@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",
)
return parser