[diffusion] model: support lingbot-video moe 30b t2v (#32341)

Signed-off-by: Pan Li <pandalee@ustc.edu>
Co-authored-by: Mick <mickjagger19@icloud.com>
This commit is contained in:
Pan Li
2026-08-07 17:57:01 +08:00
committed by GitHub
co-authored by Mick
parent 13938fed3f
commit a42683eb62
20 changed files with 1666 additions and 0 deletions
@@ -8,6 +8,9 @@ from sglang.multimodal_gen.configs.models.dits.ideogram import (
Ideogram4DistilledDiTConfig,
Ideogram4DiTConfig,
)
from sglang.multimodal_gen.configs.models.dits.lingbot_video_moe import (
LingBotVideoMoEConfig,
)
from sglang.multimodal_gen.configs.models.dits.lingbot_world import (
LingBotWorldVideoConfig,
)
@@ -27,6 +30,7 @@ __all__ = [
"Ideogram4DiTConfig",
"Ideogram4DistilledDiTConfig",
"LingBotWorldVideoConfig",
"LingBotVideoMoEConfig",
"LongLive2VideoConfig",
"MiniMaxH3DiTConfig",
"WanVideoConfig",
@@ -0,0 +1,59 @@
# SPDX-License-Identifier: Apache-2.0
from dataclasses import dataclass, field
from sglang.multimodal_gen.configs.models.dits.base import DiTArchConfig, DiTConfig
def is_blocks(name: str, module) -> bool:
return "blocks" in name and str.isdigit(name.split(".")[-1])
@dataclass
class LingBotVideoMoEArchConfig(DiTArchConfig):
_fsdp_shard_conditions: list = field(default_factory=lambda: [is_blocks])
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)
patch_size: tuple[int, int, int] = (1, 2, 2)
in_channels: int = 16
out_channels: int = 16
hidden_size: int = 2048
num_attention_heads: int = 16
depth: int = 48
intermediate_size: int = 6144
text_dim: int = 2560
freq_dim: int = 256
norm_eps: float = 1e-6
rope_theta: float = 256.0
axes_dims: tuple[int, ...] = (32, 48, 48)
axes_lens: tuple[int, ...] = (4096, 512, 512)
qkv_bias: bool = False
out_bias: bool = True
patch_embed_bias: bool = True
timestep_mlp_bias: bool = True
num_experts: int = 128
num_experts_per_tok: int = 8
moe_intermediate_size: int = 768
decoder_sparse_step: int = 1
mlp_only_layers: tuple[int, ...] = ()
n_shared_experts: int = 1
score_func: str = "sigmoid"
norm_topk_prob: bool = True
n_group: int = 4
topk_group: int = 2
routed_scaling_factor: float = 2.5
def __post_init__(self):
super().__post_init__()
self.num_channels_latents = self.out_channels
@dataclass
class LingBotVideoMoEConfig(DiTConfig):
arch_config: DiTArchConfig = field(default_factory=LingBotVideoMoEArchConfig)
prefix: str = "LingBotVideo"
@@ -32,6 +32,9 @@ from sglang.multimodal_gen.configs.pipeline_configs.ideogram import (
Ideogram4DistilledPipelineConfig,
Ideogram4PipelineConfig,
)
from sglang.multimodal_gen.configs.pipeline_configs.lingbot_video_moe import (
LingBotVideoMoEPipelineConfig,
)
from sglang.multimodal_gen.configs.pipeline_configs.lingbot_world import (
LingBotWorldCausalDMDConfig,
LingBotWorldV2CausalDMDConfig,
@@ -89,5 +92,6 @@ __all__ = [
"LTX23PipelineConfig",
"LingBotWorldCausalDMDConfig",
"LingBotWorldV2CausalDMDConfig",
"LingBotVideoMoEPipelineConfig",
"MiniMaxH3PipelineConfig",
]
@@ -0,0 +1,86 @@
# SPDX-License-Identifier: Apache-2.0
from collections.abc import Callable
from dataclasses import dataclass, field
import torch
from sglang.multimodal_gen.configs.models import DiTConfig, EncoderConfig, VAEConfig
from sglang.multimodal_gen.configs.models.dits import LingBotVideoMoEConfig
from sglang.multimodal_gen.configs.models.encoders import (
BaseEncoderOutput,
Qwen3VLConfig,
)
from sglang.multimodal_gen.configs.models.vaes import WanVAEConfig
from sglang.multimodal_gen.configs.pipeline_configs.base import (
ModelTaskType,
PipelineConfig,
)
from sglang.multimodal_gen.configs.pipeline_configs.model_deployment_config import (
ModelDeploymentConfig,
)
def _qwen3vl_postprocess_text(
outputs: BaseEncoderOutput, _text_inputs
) -> list[torch.Tensor]:
mask: torch.Tensor = outputs.attention_mask
hidden_state: torch.Tensor = outputs.last_hidden_state
seq_lens = mask.gt(0).sum(dim=1).long()
return [u[:v] for u, v in zip(hidden_state, seq_lens, strict=True)]
@dataclass
class LingBotVideoMoEPipelineConfig(PipelineConfig):
task_type: ModelTaskType = ModelTaskType.T2V
dit_config: DiTConfig = field(default_factory=LingBotVideoMoEConfig)
vae_config: VAEConfig = field(default_factory=WanVAEConfig)
vae_tiling: bool = False
vae_sp: bool = False
flow_shift: float | None = 3.0
text_encoder_configs: tuple[EncoderConfig, ...] = field(
default_factory=lambda: (Qwen3VLConfig(),)
)
text_encoder_precisions: tuple[str, ...] = field(default_factory=lambda: ("bf16",))
preprocess_text_funcs: tuple[Callable[[str], str] | None, ...] = field(
default_factory=lambda: (None,)
)
postprocess_text_funcs: tuple[Callable[[BaseEncoderOutput], torch.Tensor], ...] = (
field(default_factory=lambda: (_qwen3vl_postprocess_text,))
)
precision: str = "bf16"
vae_precision: str = "bf16"
should_use_guidance: bool = True
embedded_cfg_scale: float = 6.0
def __post_init__(self):
self.vae_config.load_encoder = False
self.vae_config.load_decoder = True
def get_model_deployment_config(self) -> ModelDeploymentConfig:
return ModelDeploymentConfig(auto_dit_layerwise_offload=True)
def get_pos_prompt_embeds(self, batch):
return batch.prompt_embeds[0]
def get_neg_prompt_embeds(self, batch):
return batch.negative_prompt_embeds[0]
def prepare_pos_cond_kwargs(self, batch, device, rotary_emb, dtype):
return {}
def prepare_neg_cond_kwargs(self, batch, device, rotary_emb, dtype):
return {}
def get_latent_dtype(self, prompt_dtype: torch.dtype) -> torch.dtype:
return torch.float32
def get_decode_scale_and_shift(self, device, dtype, vae):
arch = self.vae_config.arch_config
mean = torch.tensor(arch.latents_mean, device=device, dtype=dtype).view(
1, -1, 1, 1, 1
)
std = torch.tensor(arch.latents_std, device=device, dtype=dtype).view(
1, -1, 1, 1, 1
)
return 1.0 / std, mean
@@ -4,6 +4,9 @@ from sglang.multimodal_gen.configs.sample.diffusers_generic import (
DiffusersGenericSamplingParams,
)
from sglang.multimodal_gen.configs.sample.ideogram import Ideogram4SamplingParams
from sglang.multimodal_gen.configs.sample.lingbot_video_moe import (
LingBotVideoMoESamplingParams,
)
from sglang.multimodal_gen.configs.sample.pi05 import Pi05SamplingParams
from sglang.multimodal_gen.configs.sample.sampling_params import SamplingParams
from sglang.multimodal_gen.configs.sample.vla import VLASamplingParams
@@ -14,4 +17,5 @@ __all__ = [
"DiffusersGenericSamplingParams",
"Ideogram4SamplingParams",
"Pi05SamplingParams",
"LingBotVideoMoESamplingParams",
]
@@ -0,0 +1,20 @@
# SPDX-License-Identifier: Apache-2.0
from dataclasses import dataclass
from sglang.multimodal_gen.configs.sample.sampling_params import SamplingParams
DEFAULT_NEGATIVE_PROMPT = '{"universal_negative": {"visual_quality": ["low quality", "worst quality", "blurry", "pixelated", "jpeg artifacts", "low resolution", "unstable color", "color flicker", "underexposed", "overexposed", "invisible subject", "subject hidden in darkness"], "artistic_style": ["painting", "illustration", "drawing", "cartoon", "3d render", "cgi", "sketch", "digital art"], "composition_and_content": ["text", "watermark", "signature", "logo", "subtitles", "pillarboxed", "side bars", "portrait image in landscape frame"], "temporal_and_motion_stability": ["flickering", "jittery", "motion blur", "temporal inconsistency", "warping", "morphing", "incoherent motion", "unnatural movement", "static object with sudden jump", "frame-to-frame inconsistency"], "material_and_structure": ["plastic-like glass", "unrealistic texture", "deformed bottle", "liquid freezing improperly", "distorted reflections"]}}'
@dataclass
class LingBotVideoMoESamplingParams(SamplingParams):
# prompt must be a structured-JSON caption; raw free-text is out-of-distribution.
num_frames: int = 81
height: int = 480
width: int = 480
fps: int = 16
num_inference_steps: int = 40
guidance_scale: float = 6.0
flow_shift: float = 3.0
negative_prompt: str | None = DEFAULT_NEGATIVE_PROMPT
seed: int = 0
+14
View File
@@ -70,6 +70,9 @@ from sglang.multimodal_gen.configs.pipeline_configs.joy_image import (
JoyImageEditPipelineConfig,
)
from sglang.multimodal_gen.configs.pipeline_configs.krea2 import Krea2PipelineConfig
from sglang.multimodal_gen.configs.pipeline_configs.lingbot_video_moe import (
LingBotVideoMoEPipelineConfig,
)
from sglang.multimodal_gen.configs.pipeline_configs.longlive2 import LongLive2T2VConfig
from sglang.multimodal_gen.configs.pipeline_configs.ltx_2 import (
LTX2PipelineConfig,
@@ -133,6 +136,9 @@ from sglang.multimodal_gen.configs.sample.joy_image import (
from sglang.multimodal_gen.configs.sample.krea2 import (
Krea2SamplingParams,
)
from sglang.multimodal_gen.configs.sample.lingbot_video_moe import (
LingBotVideoMoESamplingParams,
)
from sglang.multimodal_gen.configs.sample.lingbot_world import (
LingBotWorldSamplingParams,
)
@@ -1173,6 +1179,14 @@ def _register_configs():
],
)
register_configs(
sampling_param_cls=LingBotVideoMoESamplingParams,
pipeline_config_cls=LingBotVideoMoEPipelineConfig,
model_detectors=[
lambda hf_id: "lingbot-video-moe" in hf_id.lower(),
],
)
_register_configs()
@@ -154,6 +154,20 @@ def _clear_srt_world_group() -> None:
srt_parallel_state._WORLD = None
def _sync_srt_tp_group() -> None:
import sglang.srt.distributed.parallel_state as srt_parallel_state
if srt_parallel_state._TP is None:
srt_parallel_state._TP = _TP
def _clear_srt_tp_group() -> None:
import sglang.srt.distributed.parallel_state as srt_parallel_state
if srt_parallel_state._TP is _TP:
srt_parallel_state._TP = None
def init_parallel_group_coordinator(
group_ranks: List[List[int]],
local_rank: int,
@@ -466,6 +480,7 @@ def initialize_model_parallel(
backend=backend,
parallel_mode="tensor",
)
_sync_srt_tp_group()
global _VAE_DECODE
assert _VAE_DECODE is None, "VAE decode parallel group is already initialized"
@@ -901,6 +916,7 @@ def destroy_model_parallel() -> None:
"""Set the groups to none and destroy them."""
global _TP, _SP, _DP, _CFG, _PP, _VAE_DECODE, _DIT, _VAE
_clear_srt_tp_group()
# The IPC transport keeps CUDA mappings associated with the current
# Ulysses group. Drop them before tearing down the process groups.
from .device_communicators.ipc_a2a import IPC_A2A
@@ -0,0 +1,182 @@
# Adapted from LingBot-Video (https://github.com/Robbyant/lingbot-video).
#
# SPDX-License-Identifier: Apache-2.0
from __future__ import annotations
import torch
import torch.nn.functional as F
from torch import nn
class LingBotVideoMLP(nn.Module):
def __init__(self, hidden_size: int, intermediate_size: int) -> None:
super().__init__()
self.gate_proj = nn.Linear(hidden_size, intermediate_size, bias=False)
self.up_proj = nn.Linear(hidden_size, intermediate_size, bias=False)
self.down_proj = nn.Linear(intermediate_size, hidden_size, bias=False)
def forward(self, x: torch.Tensor) -> torch.Tensor:
return self.down_proj(F.silu(self.gate_proj(x)) * self.up_proj(x))
class LingBotVideoRouter(nn.Module):
def __init__(
self,
hidden_size: int,
num_experts: int,
top_k: int,
score_func: str,
norm_topk_prob: bool,
n_group: int | None,
topk_group: int | None,
route_scale: float,
) -> None:
super().__init__()
self.num_experts = num_experts
self.top_k = top_k
self.score_func = score_func
self.norm_topk_prob = norm_topk_prob
self.n_group = n_group
self.topk_group = topk_group
self.route_scale = route_scale
self.weight = nn.Parameter(torch.empty(num_experts, hidden_size))
self.register_buffer(
"e_score_correction_bias", torch.zeros(num_experts), persistent=True
)
def _group_limited_topk(self, scores_for_choice: torch.Tensor) -> torch.Tensor:
seq_len = scores_for_choice.shape[0]
experts_per_group = self.num_experts // self.n_group
grouped = scores_for_choice.view(seq_len, self.n_group, experts_per_group)
group_scores = grouped.topk(2, dim=-1)[0].sum(dim=-1)
group_idx = torch.topk(group_scores, k=self.topk_group, dim=-1, sorted=False)[1]
group_mask = torch.zeros_like(group_scores)
group_mask.scatter_(1, group_idx, 1)
score_mask = (
group_mask.unsqueeze(-1)
.expand(seq_len, self.n_group, experts_per_group)
.reshape(seq_len, -1)
)
masked = scores_for_choice.masked_fill(~score_mask.bool(), float("-inf"))
return torch.topk(masked, k=self.top_k, dim=-1, sorted=False)[1]
def forward(self, tokens: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
with torch.amp.autocast(tokens.device.type, enabled=False):
logits = F.linear(tokens.float(), self.weight.float())
if self.score_func == "softmax":
scores = F.softmax(logits, dim=-1)
else:
scores = logits.sigmoid()
scores_for_choice = scores + self.e_score_correction_bias.unsqueeze(0)
if self.n_group is not None and self.n_group > 1:
top_indices = self._group_limited_topk(scores_for_choice)
else:
top_indices = torch.topk(
scores_for_choice, k=self.top_k, dim=-1, sorted=False
)[1]
top_scores = scores.gather(1, top_indices)
if self.top_k > 1 and self.norm_topk_prob:
top_scores = top_scores / (top_scores.sum(dim=-1, keepdim=True) + 1e-20)
top_scores = top_scores * self.route_scale
return top_indices, top_scores.to(tokens.dtype)
class LingBotVideoGroupedExperts(nn.Module):
def __init__(
self, num_experts: int, hidden_size: int, intermediate_size: int
) -> None:
super().__init__()
self.num_experts = num_experts
self.w13_weight = nn.Parameter(
torch.empty(num_experts, 2 * intermediate_size, hidden_size)
)
self.w2 = nn.Parameter(torch.empty(num_experts, hidden_size, intermediate_size))
class LingBotVideoSparseMoeBlock(nn.Module):
def __init__(
self,
hidden_size: int,
intermediate_size: int,
num_experts: int,
top_k: int,
score_func: str,
norm_topk_prob: bool,
n_group: int | None,
topk_group: int | None,
routed_scaling_factor: float,
n_shared_experts: int | None,
) -> None:
super().__init__()
self.hidden_size = hidden_size
self.num_experts = num_experts
self.top_k = top_k
self.intermediate_size = intermediate_size
self.router = LingBotVideoRouter(
hidden_size,
num_experts,
top_k,
score_func,
norm_topk_prob,
n_group,
topk_group,
routed_scaling_factor,
)
self.experts = LingBotVideoGroupedExperts(
num_experts, hidden_size, intermediate_size
)
self.shared_experts: LingBotVideoMLP | None = None
if n_shared_experts is not None and n_shared_experts > 0:
self.shared_experts = LingBotVideoMLP(
hidden_size, intermediate_size * n_shared_experts
)
def _run_sglang_triton_experts(
self,
tokens: torch.Tensor,
top_scores: torch.Tensor,
top_indices: torch.Tensor,
) -> torch.Tensor:
from sglang.srt.layers.moe.moe_runner import MoeRunnerConfig
from sglang.srt.layers.moe.moe_runner.triton_utils.fused_moe import (
fused_experts,
)
from sglang.srt.layers.moe.topk import StandardTopKOutput
topk_output = StandardTopKOutput(
topk_weights=top_scores.float(),
topk_ids=top_indices.to(torch.int32),
router_logits=torch.empty(0, device=tokens.device),
)
# Router pre-scales the topk scores; fused_experts must not apply routed_scaling_factor.
runner_config = MoeRunnerConfig(
num_experts=self.num_experts,
num_local_experts=self.num_experts,
hidden_size=self.hidden_size,
intermediate_size_per_partition=self.intermediate_size,
top_k=self.top_k,
activation="silu",
is_gated=True,
inplace=False,
apply_router_weight_on_input=False,
routed_scaling_factor=None,
gate_up_interleaved=False,
)
return fused_experts(
tokens.contiguous().bfloat16(),
self.experts.w13_weight.bfloat16(),
self.experts.w2.bfloat16(),
topk_output,
runner_config,
).type_as(tokens)
def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
b = hidden_states.shape[0]
tokens = hidden_states.reshape(-1, self.hidden_size)
top_indices, top_scores = self.router(tokens)
out = self._run_sglang_triton_experts(tokens, top_scores, top_indices)
out = out.reshape(b, -1, self.hidden_size)
if self.shared_experts is not None:
out = out + self.shared_experts(hidden_states)
return out
@@ -422,6 +422,9 @@ class ComponentLoader(ABC):
):
transformers_or_diffusers = "diffusers"
if transformers_or_diffusers.startswith("lingbot_video"):
transformers_or_diffusers = "diffusers"
return transformers_or_diffusers
@classmethod
@@ -247,6 +247,12 @@ class GPUWorker(GPUWorkerPostTrainingMixin):
dist_timeout=self.server_args.dist_timeout,
)
from sglang.srt.runtime_context import get_context
from sglang.srt.server_args import ServerArgs as SrtServerArgs
if get_context()._server_args is None:
get_context().set_server_args(SrtServerArgs(model_path="dummy"))
# set proc title
if model_parallel_is_initialized():
suffix = ""
@@ -0,0 +1,578 @@
# SPDX-License-Identifier: Apache-2.0
from __future__ import annotations
import math
from typing import Any, Iterable, Iterator, Optional
import torch
import torch.nn.functional as F
from diffusers.models.embeddings import TimestepEmbedding, Timesteps
from torch import nn
from sglang.multimodal_gen.configs.models.dits.lingbot_video_moe import (
LingBotVideoMoEConfig,
)
from sglang.multimodal_gen.runtime.distributed import divide, get_tp_world_size
from sglang.multimodal_gen.runtime.layers.attention import (
USPAttention,
build_varlen_mask_meta,
)
from sglang.multimodal_gen.runtime.layers.linear import (
ColumnParallelLinear,
RowParallelLinear,
)
from sglang.multimodal_gen.runtime.layers.moe import (
LingBotVideoMLP,
LingBotVideoSparseMoeBlock,
)
from sglang.multimodal_gen.runtime.layers.quantization.configs.base_config import (
QuantizationConfig,
)
from sglang.multimodal_gen.runtime.layers.rotary_embedding import (
NDRotaryEmbedding,
_apply_rotary_emb,
)
from sglang.multimodal_gen.runtime.managers.memory_managers.layerwise_offload import (
LayerwiseOffloadableModuleMixin,
)
from sglang.multimodal_gen.runtime.models.dits.base import CachableDiT
from sglang.multimodal_gen.runtime.platforms import (
AttentionBackendEnum,
current_platform,
)
from sglang.srt.utils import add_prefix
LINGBOT_VIDEO_FP32_MODULES = (
"time_embedder",
"time_modulation",
"scale_shift_table",
"norm",
"norm1",
"norm2",
"norm_q",
"norm_k",
"norm_post_attn",
"norm_post_ffn",
"norm_out",
"norm_out_modulation",
"router",
)
def should_keep_in_fp32(name: str) -> bool:
return any(
module_name in name.split(".") for module_name in LINGBOT_VIDEO_FP32_MODULES
)
class LingBotVideoRMSNorm(nn.Module):
def __init__(self, dim: int, eps: float = 1e-6):
super().__init__()
self.weight = nn.Parameter(torch.ones(dim))
self.variance_epsilon = eps
def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
input_dtype = hidden_states.dtype
hidden_states = hidden_states.to(torch.float32)
variance = hidden_states.pow(2).mean(-1, keepdim=True)
hidden_states = hidden_states * torch.rsqrt(variance + self.variance_epsilon)
return (self.weight * hidden_states).to(input_dtype)
def make_joint_position_ids(
text_len: int, grid_t: int, grid_h: int, grid_w: int, device: torch.device
) -> torch.Tensor:
tt = torch.arange(grid_t, device=device, dtype=torch.int32) + (text_len + 1)
hh = torch.arange(grid_h, device=device, dtype=torch.int32)
ww = torch.arange(grid_w, device=device, dtype=torch.int32)
grid = torch.stack(torch.meshgrid(tt, hh, ww, indexing="ij"), dim=-1).flatten(0, 2)
text_t = torch.arange(text_len, device=device, dtype=torch.int32) + 1
text_pos = torch.stack(
[text_t, torch.zeros_like(text_t), torch.zeros_like(text_t)], dim=-1
)
return torch.cat([grid, text_pos], dim=0) # (Nx + L, 3)
def _joint_position_ids(
text_lens: torch.Tensor,
grid_t: int,
grid_h: int,
grid_w: int,
text_len_padded: int,
device: torch.device,
) -> torch.Tensor:
# Joint video;text positions for rotary_emb; on-device, padding masked in attention.
B = text_lens.shape[0]
n_video = grid_t * grid_h * grid_w
seq_len = n_video + text_len_padded
text_lens_i = text_lens.to(torch.int32)
tt = torch.arange(grid_t, device=device, dtype=torch.int32)
hh = torch.arange(grid_h, device=device, dtype=torch.int32)
ww = torch.arange(grid_w, device=device, dtype=torch.int32)
video_t = (text_lens_i + 1)[:, None] + tt[None, :]
t_g = video_t[:, :, None, None].expand(B, grid_t, grid_h, grid_w)
h_g = hh[None, None, :, None].expand(B, grid_t, grid_h, grid_w)
w_g = ww[None, None, None, :].expand(B, grid_t, grid_h, grid_w)
video_pos = torch.stack([t_g, h_g, w_g], dim=-1).reshape(B, n_video, 3)
text_t = (
torch.arange(text_len_padded, device=device, dtype=torch.int32)[None, :] + 1
)
real = (
torch.arange(text_len_padded, device=device, dtype=torch.int32)[None, :]
< text_lens_i[:, None]
)
text_t = torch.where(real, text_t, torch.zeros_like(text_t))
text_pos = torch.stack(
[text_t, torch.zeros_like(text_t), torch.zeros_like(text_t)], dim=-1
)
return torch.cat([video_pos, text_pos], dim=1).reshape(B * seq_len, 3)
class LingBotVideoTextEmbedder(nn.Module):
def __init__(self, text_dim: int, hidden_size: int):
super().__init__()
self.norm = LingBotVideoRMSNorm(text_dim, eps=1e-6)
self.linear_1 = nn.Linear(text_dim, hidden_size, bias=True)
self.linear_2 = nn.Linear(hidden_size, hidden_size, bias=True)
def forward(self, x: torch.Tensor) -> torch.Tensor:
x = self.norm(x)
return self.linear_2(F.silu(self.linear_1(x)))
class LingBotVideoAttention(nn.Module):
def __init__(
self,
hidden_size: int,
num_heads: int,
norm_eps: float,
qkv_bias: bool,
out_bias: bool,
prefix: str = "",
supported_attention_backends: Optional[set[AttentionBackendEnum]] = None,
quant_config: Optional[QuantizationConfig] = None,
):
super().__init__()
self.num_heads = num_heads
self.head_dim = hidden_size // num_heads
tp_size = get_tp_world_size()
self.local_num_heads = divide(num_heads, tp_size)
self.to_q = ColumnParallelLinear(
hidden_size,
hidden_size,
bias=qkv_bias,
gather_output=False,
quant_config=quant_config,
prefix=add_prefix("to_q", prefix),
)
self.to_k = ColumnParallelLinear(
hidden_size,
hidden_size,
bias=qkv_bias,
gather_output=False,
quant_config=quant_config,
prefix=add_prefix("to_k", prefix),
)
self.to_v = ColumnParallelLinear(
hidden_size,
hidden_size,
bias=qkv_bias,
gather_output=False,
quant_config=quant_config,
prefix=add_prefix("to_v", prefix),
)
self.norm_q = LingBotVideoRMSNorm(self.head_dim, norm_eps)
self.norm_k = LingBotVideoRMSNorm(self.head_dim, norm_eps)
self.to_out = RowParallelLinear(
hidden_size,
hidden_size,
bias=out_bias,
input_is_parallel=True,
quant_config=quant_config,
prefix=add_prefix("to_out", prefix),
)
self.attn = USPAttention(
num_heads=self.local_num_heads,
head_size=self.head_dim,
dropout_rate=0,
softmax_scale=None,
causal=False,
supported_attention_backends=supported_attention_backends,
skip_sequence_parallel=False,
quant_config=quant_config,
)
def forward(
self,
x: torch.Tensor,
freqs_cis: tuple[torch.Tensor, torch.Tensor],
attention_mask: Optional[torch.Tensor] = None,
attn_mask_meta: Optional[dict] = None,
) -> torch.Tensor:
cos, sin = freqs_cis
q, _ = self.to_q(x)
k, _ = self.to_k(x)
v, _ = self.to_v(x)
q = self.norm_q(q.unflatten(2, (self.local_num_heads, self.head_dim)))
k = self.norm_k(k.unflatten(2, (self.local_num_heads, self.head_dim)))
v = v.unflatten(2, (self.local_num_heads, self.head_dim))
B, S, H, D = q.shape
# RoPE over the flattened batch; one batched call, the key mask isolates samples.
q = _apply_rotary_emb(
q.reshape(1, B * S, H, D), cos, sin, is_neox_style=False
).reshape(B, S, H, D)
k = _apply_rotary_emb(
k.reshape(1, B * S, H, D), cos, sin, is_neox_style=False
).reshape(B, S, H, D)
out = self.attn(
q,
k,
v,
attn_mask=attention_mask,
attn_mask_meta=attn_mask_meta,
)
out = out.flatten(2)
out, _ = self.to_out(out)
return out
class LingBotVideoBlock(nn.Module):
def __init__(
self,
hidden_size,
num_attention_heads,
intermediate_size,
norm_eps,
qkv_bias,
out_bias,
num_experts,
num_experts_per_tok,
moe_intermediate_size,
decoder_sparse_step,
mlp_only_layers,
n_shared_experts,
score_func,
norm_topk_prob,
n_group,
topk_group,
routed_scaling_factor,
layer_idx: int,
prefix: str = "",
supported_attention_backends: Optional[set[AttentionBackendEnum]] = None,
quant_config: Optional[QuantizationConfig] = None,
):
super().__init__()
self.layer_idx = layer_idx
h = hidden_size
self.scale_shift_table = nn.Parameter(torch.zeros(1, 6 * h))
self.norm1 = LingBotVideoRMSNorm(h, norm_eps)
self.attn = LingBotVideoAttention(
h,
num_attention_heads,
norm_eps,
qkv_bias,
out_bias,
prefix=add_prefix("attn", prefix),
supported_attention_backends=supported_attention_backends,
quant_config=quant_config,
)
self.norm_post_attn = LingBotVideoRMSNorm(h, norm_eps)
self.norm2 = LingBotVideoRMSNorm(h, norm_eps)
if layer_idx not in mlp_only_layers and (
num_experts > 0 and (layer_idx + 1) % decoder_sparse_step == 0
):
self.ffn = LingBotVideoSparseMoeBlock(
hidden_size=h,
intermediate_size=moe_intermediate_size,
num_experts=num_experts,
top_k=num_experts_per_tok,
score_func=score_func,
norm_topk_prob=norm_topk_prob,
n_group=n_group,
topk_group=topk_group,
routed_scaling_factor=routed_scaling_factor,
n_shared_experts=n_shared_experts,
)
else:
self.ffn = LingBotVideoMLP(h, intermediate_size)
self.norm_post_ffn = LingBotVideoRMSNorm(h, norm_eps)
def forward(
self,
x: torch.Tensor,
temb6: torch.Tensor,
freqs_cis: tuple[torch.Tensor, torch.Tensor],
attention_mask: Optional[torch.Tensor] = None,
attn_mask_meta: Optional[dict] = None,
) -> torch.Tensor:
expected_tokens = x.shape[0] * x.shape[1]
if temb6.ndim != 2 or temb6.shape[0] != expected_tokens:
raise ValueError(
"LingBotVideoBlock expects token-level temb6 with shape "
f"(B*S, 6D); got {tuple(temb6.shape)} for hidden states {tuple(x.shape)}."
)
mod = temb6.view(x.shape[0], x.shape[1], -1) + self.scale_shift_table.unsqueeze(
0
)
shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, gate_mlp = mod.chunk(
6, dim=-1
)
gate_msa, gate_mlp = gate_msa.tanh(), gate_mlp.tanh()
scale_msa, scale_mlp = 1.0 + scale_msa, 1.0 + scale_mlp
bulk_dtype = self.attn.to_q.weight.dtype
attn_in = (self.norm1(x) * scale_msa + shift_msa).to(bulk_dtype)
attn_out = self.attn(
attn_in,
freqs_cis,
attention_mask=attention_mask,
attn_mask_meta=attn_mask_meta,
)
x = x + (gate_msa * self.norm_post_attn(attn_out)).to(x.dtype)
ffn_in = (self.norm2(x) * scale_mlp + shift_mlp).to(bulk_dtype)
ffn_out = self.ffn(ffn_in)
ffn_normed = self.norm_post_ffn(ffn_out)
x = x + (gate_mlp * ffn_normed).to(x.dtype)
return x
class LingBotVideoTransformer3DModel(CachableDiT, LayerwiseOffloadableModuleMixin):
_no_split_modules = ("LingBotVideoBlock",)
_keep_in_fp32_modules = tuple(LINGBOT_VIDEO_FP32_MODULES)
_fsdp_shard_conditions = LingBotVideoMoEConfig()._fsdp_shard_conditions
_compile_conditions = LingBotVideoMoEConfig()._compile_conditions
_supported_attention_backends = (
LingBotVideoMoEConfig()._supported_attention_backends
)
param_names_mapping = LingBotVideoMoEConfig().param_names_mapping
reverse_param_names_mapping = LingBotVideoMoEConfig().reverse_param_names_mapping
lora_param_names_mapping = LingBotVideoMoEConfig().lora_param_names_mapping
def to(self, *args, **kwargs):
# _parse_to is private but the only exact parser for .to() overloads.
device, dtype, non_blocking, _ = torch._C._nn._parse_to(*args, **kwargs)
if dtype is None or dtype == torch.float32:
return super().to(*args, **kwargs)
dtype_is_floating = torch.is_floating_point(torch.empty((), dtype=dtype))
if not dtype_is_floating:
return super().to(*args, **kwargs)
if device is not None:
super().to(device=device, non_blocking=non_blocking)
for name, param in self.named_parameters():
if not torch.is_floating_point(param):
continue
target_dtype = torch.float32 if should_keep_in_fp32(name) else dtype
param.data = param.data.to(dtype=target_dtype, non_blocking=non_blocking)
if param.grad is not None:
param.grad.data = param.grad.data.to(
dtype=target_dtype, non_blocking=non_blocking
)
for name, buffer in self.named_buffers():
if not torch.is_floating_point(buffer):
continue
target_dtype = torch.float32 if should_keep_in_fp32(name) else dtype
buffer.data = buffer.data.to(dtype=target_dtype, non_blocking=non_blocking)
return self
def preprocess_loaded_state_dict(
self, weight_iterator: Iterable[tuple[str, torch.Tensor]]
) -> Iterator[tuple[str, torch.Tensor]]:
# Pack experts.w1+w3 into experts.w13_weight, gate then up on dim 1; w2 passes through.
seen: dict[str, list[torch.Tensor | None]] = {}
for name, tensor in weight_iterator:
suffix = next(
(s for s in (".ffn.experts.w1", ".ffn.experts.w3") if name.endswith(s)),
None,
)
if suffix is None:
yield name, tensor
continue
prefix = name[: -len(suffix)]
pair = seen.setdefault(prefix, [None, None])
pair[0 if suffix.endswith(".w1") else 1] = tensor
if pair[0] is not None and pair[1] is not None:
yield f"{prefix}.ffn.experts.w13_weight", torch.cat(pair, dim=1)
del seen[prefix]
def __init__(
self,
config: LingBotVideoMoEConfig,
hf_config: dict[str, Any],
quant_config: Optional[QuantizationConfig] = None,
) -> None:
super().__init__(config=config, hf_config=hf_config)
hidden_size = config.hidden_size
num_attention_heads = config.num_attention_heads
head_dim = hidden_size // num_attention_heads
assert head_dim == sum(
config.axes_dims
), f"head_dim {head_dim} != sum(axes_dims) {sum(config.axes_dims)}"
mlp_only_layers = tuple(config.mlp_only_layers)
self.hidden_size = hidden_size
self.num_attention_heads = num_attention_heads
self.in_channels = config.in_channels
self.out_channels = config.out_channels
self.num_channels_latents = config.out_channels
self.patch_size = config.patch_size
self.patch_embedder = nn.Linear(
config.in_channels * math.prod(config.patch_size),
hidden_size,
bias=config.patch_embed_bias,
)
self.time_proj = Timesteps(
config.freq_dim, flip_sin_to_cos=True, downscale_freq_shift=0
)
self.time_embedder = TimestepEmbedding(
config.freq_dim,
hidden_size,
act_fn="silu",
sample_proj_bias=config.timestep_mlp_bias,
)
self.time_modulation = nn.Sequential(
nn.SiLU(), nn.Linear(hidden_size, 6 * hidden_size)
)
self.text_embedder = LingBotVideoTextEmbedder(config.text_dim, hidden_size)
self.rotary_emb = NDRotaryEmbedding(
rope_dim_list=list(config.axes_dims),
rope_theta=config.rope_theta,
dtype=(
torch.float64
if current_platform.is_float64_supported()
else torch.float32
),
)
self.blocks = nn.ModuleList(
[
LingBotVideoBlock(
hidden_size=hidden_size,
num_attention_heads=num_attention_heads,
intermediate_size=config.intermediate_size,
norm_eps=config.norm_eps,
qkv_bias=config.qkv_bias,
out_bias=config.out_bias,
num_experts=config.num_experts,
num_experts_per_tok=config.num_experts_per_tok,
moe_intermediate_size=config.moe_intermediate_size,
decoder_sparse_step=config.decoder_sparse_step,
mlp_only_layers=mlp_only_layers,
n_shared_experts=config.n_shared_experts,
score_func=config.score_func,
norm_topk_prob=config.norm_topk_prob,
n_group=config.n_group,
topk_group=config.topk_group,
routed_scaling_factor=config.routed_scaling_factor,
layer_idx=i,
prefix=f"blocks.{i}",
supported_attention_backends=self._supported_attention_backends,
quant_config=quant_config,
)
for i in range(config.depth)
]
)
self.norm_out = nn.LayerNorm(
hidden_size, elementwise_affine=False, eps=config.norm_eps
)
self.norm_out_modulation = nn.Sequential(
nn.SiLU(), nn.Linear(hidden_size, 2 * hidden_size)
)
self.proj_out = nn.Linear(
hidden_size, math.prod(config.patch_size) * config.out_channels
)
self.__post_init__()
self.layer_names = ["blocks"]
def forward(
self,
hidden_states: torch.Tensor, # (B, C, T, H, W)
timestep: torch.Tensor, # (B,) in [0, 1000] (= sigma*1000)
encoder_hidden_states: torch.Tensor, # (B, L, text_dim)
encoder_attention_mask: Optional[torch.Tensor] = None, # (B, L) 1=valid
**kwargs,
) -> torch.Tensor:
B, C, T, H, W = hidden_states.shape
pF, pH, pW = self.patch_size
gt, gh, gw = T // pF, H // pH, W // pW
n_video = gt * gh * gw
L = encoder_hidden_states.shape[1]
device = hidden_states.device
if encoder_attention_mask is not None:
text_lens = encoder_attention_mask.sum(dim=-1).long()
else:
text_lens = torch.full((B,), L, dtype=torch.long, device=device)
patch_tokens = hidden_states.reshape(B, C, gt, pF, gh, pH, gw, pW)
patch_tokens = patch_tokens.permute(0, 2, 4, 6, 3, 5, 7, 1).reshape(
B,
n_video,
pF * pH * pW * C,
)
x = self.patch_embedder(patch_tokens)
text = self.text_embedder(encoder_hidden_states)
joint = torch.cat([x, text], dim=1) # [video; text]
joint_seq_len = joint.shape[1]
positions = _joint_position_ids(text_lens, gt, gh, gw, L, device)
cos, sin = self.rotary_emb.forward_uncached(positions)
freqs_cis = (cos.float(), sin.float())
attention_mask = attn_mask_meta = None
# B==1 text is trimmed to true length upstream, so no mask; B>1 may pad, build a key mask.
if B > 1 and encoder_attention_mask is not None:
attention_mask = torch.cat(
[
torch.ones(B, n_video, dtype=torch.bool, device=device),
encoder_attention_mask.bool(),
],
dim=1,
)
attn_mask_meta = build_varlen_mask_meta(attention_mask)
timestep_for_embed = timestep.float()
timestep_proj = self.time_proj(timestep_for_embed)
t_emb = self.time_embedder(timestep_proj) # (B, D)
temb_input = t_emb.unsqueeze(1).expand(B, joint_seq_len, -1) # (B, S, D)
temb6 = self.time_modulation(temb_input.reshape(B * joint_seq_len, -1))
temb6 = temb6.reshape(B, joint_seq_len, -1) # (B, S, 6D)
temb6 = temb6.reshape(temb6.shape[0] * temb6.shape[1], -1)
for block in self.blocks:
joint = block(
joint,
temb6,
freqs_cis,
attention_mask,
attn_mask_meta,
)
final_mod = self.norm_out_modulation(
temb_input.reshape(joint.shape[0] * joint.shape[1], -1)
)
shift, scale = final_mod.reshape(joint.shape[0], joint.shape[1], -1).chunk(
2, dim=-1
)
final_hidden = self.norm_out(joint) * (1.0 + scale) + shift
projected = self.proj_out(final_hidden.to(self.proj_out.weight.dtype))
x = projected[:, :n_video]
Cout = self.out_channels
x = x.reshape(B, gt, gh, gw, pF, pH, pW, Cout)
x = x.permute(0, 7, 1, 4, 2, 5, 3, 6).reshape(B, Cout, T, H, W)
return x
EntryClass = [LingBotVideoTransformer3DModel]
@@ -0,0 +1,60 @@
# SPDX-License-Identifier: Apache-2.0
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 (
DenoisingStage,
InputValidationStage,
)
from sglang.multimodal_gen.runtime.pipelines_core.stages.model_specific_stages.lingbot_video_moe import (
LingBotVideoTextEncodingStage,
)
from sglang.multimodal_gen.runtime.server_args import ServerArgs
def _flow_shift_kwarg(batch, server_args: ServerArgs) -> tuple[str, float | None]:
shift = (
batch.flow_shift
if batch.flow_shift is not None
else server_args.pipeline_config.flow_shift
)
return ("shift", shift)
class LingBotVideoPipeline(LoRAPipeline, ComposedPipelineBase):
pipeline_name = "LingBotVideoPipeline"
is_video_pipeline = True
_required_config_modules = (
"text_encoder",
"processor",
"vae",
"transformer",
"scheduler",
)
def create_pipeline_stages(self, server_args: ServerArgs) -> None:
self.add_stage(InputValidationStage())
self.add_stage(
LingBotVideoTextEncodingStage(
text_encoders=[self.get_module("text_encoder")],
tokenizers=[self.get_module("processor")],
transformer=self.get_module("transformer"),
),
)
self.add_standard_latent_preparation_stage()
self.add_standard_timestep_preparation_stage(
prepare_extra_kwargs=[_flow_shift_kwarg],
)
self.add_stage(
DenoisingStage(
transformer=self.get_module("transformer"),
scheduler=self.get_module("scheduler"),
),
)
self.add_standard_decoding_stage()
EntryClass = [LingBotVideoPipeline]
@@ -0,0 +1,6 @@
# SPDX-License-Identifier: Apache-2.0
"""LingBot-Video MoE model-specific pipeline stages."""
from sglang.multimodal_gen.runtime.pipelines_core.stages.model_specific_stages.lingbot_video_moe.text_encoding import ( # noqa: F401
LingBotVideoTextEncodingStage,
)
@@ -0,0 +1,152 @@
# SPDX-License-Identifier: Apache-2.0
import torch
from sglang.multimodal_gen.runtime.distributed import get_local_torch_device
from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import Req
from sglang.multimodal_gen.runtime.pipelines_core.stages.text_encoding import (
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__)
TOKEN_LENGTH = 37698
HIDDEN_STATE_SKIP_LAYER = 0
PROMPT_TEMPLATE = (
"<|im_start|>system\nGiven a user input that may include a text prompt alone, "
"a text prompt with an image reference, or a text prompt with a video reference "
'or a video reference alone, generate an "Enhanced prompt" that provides detailed '
"visual descriptions suitable for video generation. Evaluate the level of detail "
"in the user's input: if it is simple, enrich it by adding specifics about colors, "
"shapes, sizes, textures, lighting, motion dynamics, camera movement, temporal "
"progression, and spatial relationships to create vivid, concrete, and temporally "
"coherent scenes to create vivid and concrete scenes. Please generate only the "
"enhanced description for the prompt below and avoid including any additional "
"commentary or evaluations:<|im_end|>\n<|im_start|>user\n{}<|im_end|>\n"
"<|im_start|>assistant\n"
)
IMG_PROMPT_TEMPLATE = "<|vision_start|><|image_pad|><|vision_end|>"
VIDEO_PROMPT_TEMPLATE = "<|vision_start|><|video_pad|><|vision_end|>"
class LingBotVideoTextEncodingStage(TextEncodingStage):
"""Qwen3-VL prompt/negative encoding for LingBot-Video MoE (T2V, base)."""
def __init__(self, text_encoders, tokenizers, transformer):
super().__init__(text_encoders, tokenizers)
self.transformer = transformer
self.token_length = TOKEN_LENGTH
self.hidden_state_skip_layer = HIDDEN_STATE_SKIP_LAYER
self.prompt_template = PROMPT_TEMPLATE
self._crop_start: int | None = None
@staticmethod
def check_inputs(height: int, width: int, num_frames: int) -> None:
if num_frames != 1 and (num_frames - 1) % 4 != 0:
raise ValueError(f"`num_frames` must be 1 or 4n+1, got {num_frames}.")
if height % 16 != 0 or width % 16 != 0:
raise ValueError(
f"`height` and `width` must be multiples of 16, got {height}x{width}."
)
@staticmethod
def apply_text_to_template(text: str, template: str = PROMPT_TEMPLATE) -> str:
return template.format(text)
def _compute_crop_start(self) -> int:
processor = self.tokenizers[0]
if self._crop_start is None:
marker = "<|USER_INPUT_MARKER|>"
marked = self.prompt_template.format(marker)
marker_pos = marked.find(marker)
if marker_pos < 0:
self._crop_start = 0
else:
prefix = processor(
text=marked[:marker_pos],
images=None,
videos=None,
return_tensors="pt",
)
self._crop_start = int(prefix["input_ids"].shape[1])
return self._crop_start
def _build_prompt_inputs(self, prompt: str | list[str]):
processor = self.tokenizers[0]
prompts = [prompt] if isinstance(prompt, str) else list(prompt)
texts = [
self.apply_text_to_template(text, self.prompt_template) for text in prompts
]
return processor(
text=texts,
images=None,
videos=None,
video_metadata=None,
do_resize=False,
truncation=True,
max_length=self.token_length,
padding="longest",
return_tensors="pt",
)
@torch.no_grad()
def _encode_prompt(
self,
prompt: str | list[str],
device: torch.device,
dtype: torch.dtype,
):
text_encoder = self.text_encoders[0]
if text_encoder is None or self.tokenizers[0] is None:
raise ValueError(
"`text_encoder` and `processor` are required for encode_prompt()."
)
inputs = self._build_prompt_inputs(prompt)
inputs = inputs.to(device)
outputs = text_encoder(
**inputs,
output_hidden_states=self.hidden_state_skip_layer is not None,
)
if self.hidden_state_skip_layer is not None:
prompt_embeds = outputs.hidden_states[-(self.hidden_state_skip_layer + 1)]
else:
prompt_embeds = outputs.last_hidden_state
prompt_mask = inputs["attention_mask"]
crop_start = self._compute_crop_start()
if crop_start > 0:
prompt_embeds = prompt_embeds[:, crop_start:]
prompt_mask = prompt_mask[:, crop_start:]
# B=1: drop right padding before DiT inference.
if prompt_embeds.shape[0] == 1:
true_len = int(prompt_mask[0].sum().item())
prompt_embeds = prompt_embeds[:, :true_len]
prompt_mask = prompt_mask[:, :true_len]
return prompt_embeds.to(dtype=dtype), prompt_mask
@torch.no_grad()
def forward(self, batch: Req, server_args: ServerArgs) -> Req:
device = get_local_torch_device()
dtype = next(self.transformer.parameters(), torch.tensor([])).dtype
if dtype not in (torch.bfloat16, torch.float16, torch.float32):
dtype = torch.bfloat16
self.check_inputs(int(batch.height), int(batch.width), int(batch.num_frames))
prompt_embeds, prompt_mask = self._encode_prompt(batch.prompt, device, dtype)
batch.prompt_embeds = [prompt_embeds]
batch.prompt_attention_mask = prompt_mask
if batch.do_classifier_free_guidance:
negative_embeds, negative_mask = self._encode_prompt(
batch.negative_prompt, device, dtype
)
batch.negative_prompt_embeds = [negative_embeds]
batch.negative_attention_mask = negative_mask
return batch
@@ -22,6 +22,7 @@ from sglang.multimodal_gen.test.server.testcase_configs import (
DiffusionTestCase,
IDEOGRAM4_CI_sampling_params,
JOY_ECHO_T2V_CI_sampling_params,
LINGBOT_VIDEO_T2V_CI_sampling_params,
LONGLIVE2_I2V_CI_sampling_params,
LONGLIVE2_T2V_CI_sampling_params,
MODELOPT_QWEN_IMAGE_2512_NVFP4_CI_sampling_params,
@@ -450,6 +451,21 @@ ONE_GPU_CASES: list[DiffusionTestCase] = [
),
run_component_accuracy_check=False,
),
DiffusionTestCase(
"lingbot_video_moe_t2v",
DiffusionServerArgs(
model_path="robbyant/lingbot-video-moe-30b-a3b",
modality="video",
num_gpus=1,
text_encoder_cpu_offload=True,
),
LINGBOT_VIDEO_T2V_CI_sampling_params,
run_perf_check=False,
run_consistency_check=False,
run_component_accuracy_check=False,
run_models_api_check=False,
run_t2v_input_reference_check=False,
),
DiffusionTestCase(
"lingbot_world_realtime_plastic_beach",
DiffusionServerArgs(
@@ -101,6 +101,28 @@
"expected_avg_denoise_ms": 246.97,
"expected_median_denoise_ms": 273.01,
"estimated_full_test_time_s": 329.8
},
"lingbot_video_moe_t2v": {
"stages_ms": {
"InputValidationStage": 0.1,
"LingBotVideoTextEncodingStage": 270.0,
"LatentPreparationStage": 1.2,
"TimestepPreparationStage": 0.2,
"DenoisingStage": 26414.5,
"DecodingStage": 628.3
},
"denoise_step_ms": {
"0": 2200.8,
"2": 2200.8,
"4": 2200.8,
"7": 2200.8,
"9": 2200.8,
"11": 2200.8
},
"expected_e2e_ms": 27740.0,
"expected_avg_denoise_ms": 2200.8,
"expected_median_denoise_ms": 2200.8,
"estimated_full_test_time_s": 90.0
}
}
}
@@ -2634,6 +2634,14 @@
"expected_median_denoise_ms": 800.0,
"estimated_full_test_time_s": 149.4
},
"lingbot_video_moe_t2v": {
"stages_ms": {},
"denoise_step_ms": {},
"expected_e2e_ms": 0.0,
"expected_avg_denoise_ms": 0.0,
"expected_median_denoise_ms": 0.0,
"estimated_full_test_time_s": 600.0
},
"lingbot_world_realtime_plastic_beach": {
"stages_ms": {},
"denoise_step_ms": {},
@@ -596,6 +596,65 @@ MODELOPT_T2V_CI_sampling_params = DiffusionSamplingParams(
extras={"num_inference_steps": 12, "seed": 0},
)
LINGBOT_VIDEO_T2V_CI_PROMPT = json.dumps(
{
"comprehensive_description": {
"scene_content_description": (
"A small silver robot arm on a white table slowly reaches "
"toward a red cube. The background is a plain, softly lit "
"laboratory wall."
),
"camera_movement_description": (
"The camera is static at eye level, medium shot, with the "
"robot arm centered and in sharp focus."
),
},
"camera_info": {
"color": "Neutral",
"frame_size": "Medium",
"shot_type_angle": "Eye level",
"lens_size": "Medium",
"composition": "Center",
"lighting": "Soft light",
"lighting_type": "Artificial light",
},
"world_knowledge": [],
"prominent_elements": [
{
"name": "robot arm",
"description": "A small silver robot arm with a two-finger gripper.",
"actions": [
{
"timestamp": "[0.0s - 1.0s]",
"action": "reaches toward the red cube",
}
],
"location": "center of the frame",
"relative_size": "dominant",
"shape_and_color": "articulated silver metal arm",
"texture": "brushed metal",
"appearance_details": "two-finger gripper, visible joints",
"relationship": "reaching toward the red cube on the table",
"orientation": "upright, base on the table",
"pose": "reaching",
"expression": "",
"clothing": "",
"gender": "",
"skin_tone_and_texture": "",
}
],
},
separators=(",", ":"),
)
LINGBOT_VIDEO_T2V_CI_sampling_params = DiffusionSamplingParams(
prompt=LINGBOT_VIDEO_T2V_CI_PROMPT,
output_size="384x640",
num_frames=17,
fps=16,
extras={"num_inference_steps": 12, "seed": 0},
)
TI2V_sampling_params = DiffusionSamplingParams(
prompt="The man in the picture slowly turns his head, his expression enigmatic and otherworldly. The camera performs a slow, cinematic dolly out, focusing on his face. Moody lighting, neon signs glowing in the background, shallow depth of field.",
image_path="https://is1-ssl.mzstatic.com/image/thumb/Music114/v4/5f/fa/56/5ffa56c2-ea1f-7a17-6bad-192ff9b6476d/825646124206.jpg/600x600bb.jpg",
@@ -0,0 +1,367 @@
# SPDX-License-Identifier: Apache-2.0
import json
import os
import tempfile
from types import SimpleNamespace
import torch
from sglang.multimodal_gen.configs.models.dits.lingbot_video_moe import (
LingBotVideoMoEArchConfig,
)
from sglang.multimodal_gen.configs.pipeline_configs.lingbot_video_moe import (
LingBotVideoMoEPipelineConfig,
)
from sglang.multimodal_gen.configs.sample.lingbot_video_moe import (
LingBotVideoMoESamplingParams,
)
from sglang.multimodal_gen.registry import _get_config_info, get_model_info
from sglang.multimodal_gen.runtime.layers.moe import (
LingBotVideoGroupedExperts,
LingBotVideoRouter,
)
from sglang.multimodal_gen.runtime.models.dits import (
lingbot_video_moe as dits_lingbot_video_moe,
)
from sglang.multimodal_gen.runtime.models.dits.lingbot_video_moe import (
LingBotVideoAttention,
LingBotVideoTransformer3DModel,
_joint_position_ids,
make_joint_position_ids,
)
from sglang.multimodal_gen.runtime.pipelines_core.stages.model_specific_stages.lingbot_video_moe.text_encoding import (
PROMPT_TEMPLATE,
LingBotVideoTextEncodingStage,
)
_LINGBOT_MODULE_SUBDIRS = (
"scheduler",
"text_encoder",
"processor",
"transformer",
"vae",
)
def test_moe_path_resolves_moe_configs():
get_model_info.cache_clear()
_get_config_info.cache_clear()
with tempfile.TemporaryDirectory() as tmpdir:
model_dir = os.path.join(tmpdir, "lingbot-video-moe-30b-a3b")
os.makedirs(model_dir)
with open(
os.path.join(model_dir, "model_index.json"), "w", encoding="utf-8"
) as f:
json.dump(
{"_class_name": "LingBotVideoPipeline", "_diffusers_version": "0.37.1"},
f,
)
for subdir in _LINGBOT_MODULE_SUBDIRS:
os.mkdir(os.path.join(model_dir, subdir))
info = get_model_info(model_dir, backend="sglang")
assert info.pipeline_cls.__name__ == "LingBotVideoPipeline"
assert info.pipeline_config_cls is LingBotVideoMoEPipelineConfig
assert info.sampling_param_cls is LingBotVideoMoESamplingParams
def test_arch_config_defaults_without_mlp_only_layers():
arch = LingBotVideoMoEArchConfig()
assert arch.num_experts == 128
assert arch.mlp_only_layers == ()
def test_router_bias_shifts_selection_but_not_gate_weights():
router = LingBotVideoRouter(
hidden_size=4,
num_experts=4,
top_k=2,
score_func="sigmoid",
norm_topk_prob=False,
n_group=None,
topk_group=None,
route_scale=1.0,
)
with torch.no_grad():
router.weight.copy_(
torch.tensor(
[
[4.0, 0.0, 0.0, 0.0],
[2.0, 0.0, 0.0, 0.0],
[-2.0, 0.0, 0.0, 0.0],
[-4.0, 0.0, 0.0, 0.0],
]
)
)
router.e_score_correction_bias.copy_(torch.tensor([0.0, 0.0, 0.0, 10.0]))
top_indices, top_scores = router(torch.tensor([[1.0, 0.0, 0.0, 0.0]]))
assert set(top_indices[0].tolist()) == {0, 3}
raw = torch.sigmoid(torch.tensor([4.0, -4.0]))
picked = {
int(idx): float(score.detach())
for idx, score in zip(top_indices[0], top_scores[0])
}
assert abs(picked[0] - float(raw[0])) < 1e-5
assert abs(picked[3] - float(raw[1])) < 1e-5
def _sdpa(q, k, v, attn_mask=None, attn_mask_meta=None):
q_, k_, v_ = q.transpose(1, 2), k.transpose(1, 2), v.transpose(1, 2)
if attn_mask is not None and attn_mask.dim() == 2:
attn_mask = attn_mask[:, None, None, :]
out = torch.nn.functional.scaled_dot_product_attention(
q_, k_, v_, attn_mask=attn_mask
)
return out.transpose(1, 2)
def _real_attention(num_heads, head_dim):
attn = object.__new__(LingBotVideoAttention)
attn.local_num_heads = num_heads
attn.head_dim = head_dim
attn.to_q = attn.to_k = attn.to_v = attn.to_out = lambda x: (x, None)
attn.norm_q = attn.norm_k = lambda t: t
attn.attn = _sdpa
return attn
def test_attention_isolates_samples_across_batch(monkeypatch):
monkeypatch.setattr(
dits_lingbot_video_moe, "_apply_rotary_emb", lambda t, *a, **k: t
)
num_heads, head_dim, batch, seq_len = 4, 8, 3, 8
attn = _real_attention(num_heads, head_dim)
hidden = num_heads * head_dim
torch.manual_seed(0)
x = torch.randn(batch, seq_len, hidden)
freqs = torch.zeros(batch * seq_len, head_dim // 2)
valid = [seq_len, seq_len - 2, seq_len - 5]
mask = torch.zeros(batch, seq_len, dtype=torch.bool)
for i, length in enumerate(valid):
mask[i, :length] = True
batched = attn.forward(x, (freqs, freqs), mask)
for i, length in enumerate(valid):
solo = attn.forward(
x[i : i + 1],
(freqs[i * seq_len : (i + 1) * seq_len],) * 2,
mask[i : i + 1],
)
torch.testing.assert_close(batched[i : i + 1, :length], solo[:, :length])
# Flattening the batch into one sequence lets sample 0 attend across the
# boundary; its output must differ from the isolated per-sample result.
flat = attn.forward(x.reshape(1, batch * seq_len, hidden), (freqs, freqs), None)
flat = flat.reshape(batch, seq_len, hidden)
assert (flat[0, : valid[0]] - batched[0, : valid[0]]).abs().max() > 1e-3
def test_attention_forwards_2d_mask_and_varlen_metadata(monkeypatch):
monkeypatch.setattr(
dits_lingbot_video_moe, "_apply_rotary_emb", lambda t, *a, **k: t
)
num_heads, head_dim, batch, seq_len = 4, 8, 2, 6
attn = _real_attention(num_heads, head_dim)
hidden = num_heads * head_dim
captured = {}
def capture_attention(q, k, v, attn_mask=None, attn_mask_meta=None):
captured["mask"] = attn_mask
captured["meta"] = attn_mask_meta
return _sdpa(q, k, v, attn_mask=attn_mask)
attn.attn = capture_attention
x = torch.randn(batch, seq_len, hidden)
freqs = torch.zeros(batch * seq_len, head_dim // 2)
mask = torch.tensor([[1, 1, 1, 1, 1, 1], [1, 1, 1, 1, 0, 0]], dtype=torch.bool)
metadata = {"max_seqlen": seq_len}
attn.forward(x, (freqs, freqs), mask, metadata)
assert captured["mask"] is mask
assert captured["meta"] is metadata
def test_attention_single_sample_matches_direct_attention(monkeypatch):
monkeypatch.setattr(
dits_lingbot_video_moe, "_apply_rotary_emb", lambda t, *a, **k: t
)
num_heads, head_dim, seq_len = 4, 8, 6
attn = _real_attention(num_heads, head_dim)
hidden = num_heads * head_dim
torch.manual_seed(0)
x = torch.randn(1, seq_len, hidden)
freqs = torch.zeros(seq_len, head_dim // 2)
out = attn.forward(x, (freqs, freqs), attention_mask=None)
qkv = x.unflatten(2, (num_heads, head_dim))
expected = _sdpa(qkv, qkv, qkv).flatten(2)
torch.testing.assert_close(out, expected)
class _FakeBatchEncoding(dict):
def to(self, _device):
return self
class _FakeQwenProcessor:
def __init__(self, prompt_width, prefix_width, true_len):
self.prompt_width = prompt_width
self.prefix_width = prefix_width
self.true_len = true_len
def __call__(self, **kwargs):
if "max_length" in kwargs:
width = self.prompt_width
mask = torch.zeros(1, width, dtype=torch.long)
mask[0, : self.true_len] = 1
else:
width = self.prefix_width
mask = torch.ones(1, width, dtype=torch.long)
return _FakeBatchEncoding(
input_ids=torch.zeros(1, width, dtype=torch.long),
attention_mask=mask,
)
def _text_encoding_stage(processor, encoder):
stage = object.__new__(LingBotVideoTextEncodingStage)
stage.text_encoders = [encoder]
stage.tokenizers = [processor]
stage.token_length = 128
stage.hidden_state_skip_layer = 0
stage.prompt_template = PROMPT_TEMPLATE
stage._crop_start = None
return stage
def test_text_encoding_crops_template_then_trims_padding():
prompt_width, prefix_width, true_len, channels = 10, 3, 8, 4
hidden = torch.arange(prompt_width, dtype=torch.float32)
hidden = hidden.view(1, prompt_width, 1).expand(1, prompt_width, channels)
def encoder(**kwargs):
return SimpleNamespace(hidden_states=[hidden])
stage = _text_encoding_stage(
_FakeQwenProcessor(prompt_width, prefix_width, true_len), encoder
)
embeds, mask = stage._encode_prompt(
"a structured caption", torch.device("cpu"), torch.float32
)
assert tuple(embeds.shape) == (1, true_len - prefix_width, channels)
torch.testing.assert_close(embeds, hidden[:, prefix_width:true_len])
assert int(mask.sum()) == true_len - prefix_width
assert stage._compute_crop_start() == prefix_width
def test_check_inputs_enforces_frame_and_size_contract():
check = LingBotVideoTextEncodingStage.check_inputs
check(480, 832, 1)
check(480, 832, 81)
try:
check(480, 832, 82)
raise AssertionError("expected ValueError for num_frames=82")
except ValueError:
pass
try:
check(480, 830, 81)
raise AssertionError("expected ValueError for width=830")
except ValueError:
pass
def test_decode_scale_and_shift_invert_vae_normalization():
config = LingBotVideoMoEPipelineConfig()
scale, shift = config.get_decode_scale_and_shift(
torch.device("cpu"), torch.float32, vae=None
)
arch = config.vae_config.arch_config
std = torch.tensor(arch.latents_std, dtype=torch.float32).view(1, -1, 1, 1, 1)
mean = torch.tensor(arch.latents_mean, dtype=torch.float32).view(1, -1, 1, 1, 1)
torch.testing.assert_close(scale, 1.0 / std)
torch.testing.assert_close(shift, mean)
def test_latents_stay_fp32_under_bf16_precision():
config = LingBotVideoMoEPipelineConfig()
assert config.get_latent_dtype(torch.bfloat16) == torch.float32
def test_grouped_experts_store_packed_w13_weight():
experts = LingBotVideoGroupedExperts(
num_experts=2, hidden_size=4, intermediate_size=3
)
names = {n for n, _ in experts.named_parameters()}
assert "w13_weight" in names and "w2" in names
assert "w1" not in names and "w3" not in names
assert tuple(experts.w13_weight.shape) == (2, 6, 4) # [E, 2I, H]
def test_preprocess_packs_w1_w3_into_w13_weight():
pack = LingBotVideoTransformer3DModel.preprocess_loaded_state_dict
E, I, H = 2, 3, 4
w1 = torch.arange(E * I * H, dtype=torch.float32).reshape(E, I, H)
w2 = torch.arange(E * H * I, dtype=torch.float32).reshape(E, H, I)
w3 = torch.arange(E * I * H, dtype=torch.float32).reshape(E, I, H) + 100.0
# block 0: w1 before w3; block 1: w3 before w1 (order-independence).
src = [
("blocks.0.ffn.experts.w1", w1),
("blocks.0.ffn.experts.w2", w2),
("blocks.0.ffn.experts.w3", w3),
("blocks.0.ffn.router.weight", torch.zeros(E, H)),
("blocks.1.ffn.experts.w3", w3.clone()),
("blocks.1.ffn.experts.w2", w2.clone()),
("blocks.1.ffn.experts.w1", w1.clone()),
]
out = dict(pack(None, iter(src)))
assert set(out.keys()) == {
"blocks.0.ffn.experts.w13_weight",
"blocks.0.ffn.experts.w2",
"blocks.0.ffn.router.weight",
"blocks.1.ffn.experts.w13_weight",
"blocks.1.ffn.experts.w2",
}
packed = torch.cat((w1, w3), dim=1) # gate then up, dim-1
torch.testing.assert_close(out["blocks.0.ffn.experts.w13_weight"], packed)
torch.testing.assert_close(out["blocks.1.ffn.experts.w13_weight"], packed)
torch.testing.assert_close(out["blocks.0.ffn.experts.w2"], w2)
def test_joint_position_ids_match_reference_and_cover_padding():
dev = torch.device("cpu")
gt, gh, gw = 2, 3, 4
n_video = gt * gh * gw
# B==1, no padding: byte-identical to the per-sample reference.
vec = _joint_position_ids(torch.tensor([5]), gt, gh, gw, 5, dev)
torch.testing.assert_close(vec, make_joint_position_ids(5, gt, gh, gw, dev))
# B==1 with padding: real tokens match the text_len=4 reference; the extra
# padding row is (0,0,0). vec has n_video+L rows (matches q for B*S).
vec_p = _joint_position_ids(torch.tensor([4]), gt, gh, gw, 5, dev)
torch.testing.assert_close(
vec_p[: n_video + 4], make_joint_position_ids(4, gt, gh, gw, dev)
)
torch.testing.assert_close(
vec_p[n_video + 4 :], torch.zeros((1, 3), dtype=torch.int32)
)
# B>1 with padding: covers B*S rows; each sample's real tokens match its ref.
text_lens = [5, 3, 6]
B, L = len(text_lens), 6
vec_b = _joint_position_ids(torch.tensor(text_lens), gt, gh, gw, L, dev)
assert vec_b.shape[0] == B * (n_video + L)
for i, t in enumerate(text_lens):
start = i * (n_video + L)
real = n_video + t
torch.testing.assert_close(
vec_b[start : start + real], make_joint_position_ids(t, gt, gh, gw, dev)
)