diff --git a/docs_new/cookbook/diffusion/LongLive/LongLive-2.0.mdx b/docs_new/cookbook/diffusion/LongLive/LongLive-2.0.mdx
new file mode 100644
index 000000000..ae03c741f
--- /dev/null
+++ b/docs_new/cookbook/diffusion/LongLive/LongLive-2.0.mdx
@@ -0,0 +1,109 @@
+---
+title: LongLive 2.0
+description: "Serve LongLive 2.0 distilled text-to-video and image-to-video models with SGLang-diffusion."
+tag: NEW
+---
+
+## 1. Model Introduction
+
+[LongLive 2.0](https://nvlabs.github.io/LongLive/LongLive2/) is a distilled few-step text-to-video and image-to-video model from NVIDIA, built on Wan2.2-TI2V-5B. SGLang serves the Diffusers-format conversion for single-prompt and multi-shot video generation.
+
+For more details, check the [LongLive 2.0 paper](https://arxiv.org/abs/2605.18739) and [LongLive 2.0 GitHub](https://github.com/NVlabs/LongLive). The model weights are released under the NVIDIA Open Model License.
+
+## 2. SGLang-diffusion Installation
+
+Please refer to the [official SGLang-diffusion installation guide](/docs/sglang-diffusion/installation) for installation instructions.
+
+## 3. Deployment
+
+```bash Command
+sglang serve --model-path Rabinovich/LongLive-2.0-5B-Diffusers
+```
+
+If the GPU runs out of memory, move the text encoder, VAE, and DiT to CPU between stages:
+
+```bash Command
+sglang serve \
+ --model-path Rabinovich/LongLive-2.0-5B-Diffusers \
+ --dit-cpu-offload \
+ --text-encoder-cpu-offload \
+ --vae-cpu-offload
+```
+
+`Rabinovich/LongLive-2.0-5B-Diffusers` is the Diffusers-format conversion of the official `Efficient-Large-Model/LongLive-2.0-5B` weights.
+
+## 4. Generation
+
+### 4.1 Single prompt
+
+Generate one clip without starting a server:
+
+```bash Command
+sglang generate \
+ --model-path Rabinovich/LongLive-2.0-5B-Diffusers \
+ --prompt "A quiet street at dusk" \
+ --num-frames 61 \
+ --save-output \
+ --output-path outputs
+```
+
+61 frames is 16 latent frames, which is two causal blocks of 8.
+
+### 4.2 Multi-shot long video
+
+Multi-shot prompts are sampling parameters, so pass them through the Python API:
+
+```python Python
+from sglang import DiffGenerator
+
+gen = DiffGenerator.from_pretrained("Rabinovich/LongLive-2.0-5B-Diffusers")
+result = gen.generate(sampling_params_kwargs={
+ "shot_prompts": [
+ "A husky walks down a sunlit hallway.",
+ "The husky turns and looks at the camera.",
+ "Two dogs play together on a carpet.",
+ ],
+ "chunks_per_shot": 4,
+ "num_frames": 381, # 3 shots x 4 chunks x 8 = 96 latent frames -> 381 frames
+ "scene_cut_prefix": "The scene transitions. ",
+ "multi_shot_sink": True,
+ "multi_shot_rope_offset": 8.0,
+ "save_output": True,
+ "output_path": "outputs",
+})
+```
+
+Each shot runs for `chunks_per_shot` causal blocks before the next prompt is used. The multi-shot defaults mirror the original LongLive prompt-block settings.
+
+### 4.3 Key parameters
+
+These are SGLang request parameters. Original LongLive configs use latent-frame `num_output_frames`; SGLang exposes output-video `num_frames`.
+
+- `num_frames`: 61 in the examples. This maps to 16 latent frames, while the original release config defaults to 128 latent frames.
+- `num_inference_steps`: 4, matching original `sampling_steps`.
+- `guidance_scale`: 1.0, matching the original inference config.
+- `height` / `width`: 704 / 1280 by default, matching original latent H/W 44 / 80 with 16x spatial compression.
+- `shot_prompts`, `chunks_per_shot`, `scene_cut_prefix`, `multi_shot_sink`, and `multi_shot_rope_offset`: SGLang request fields for the original prompt-block and multi-shot behavior.
+
+### 4.4 Image-to-video
+
+Pass a first frame with `--image-path` to condition the clip on an image:
+
+```bash Command
+sglang generate \
+ --model-path Rabinovich/LongLive-2.0-5B-Diffusers \
+ --prompt "A quiet street at dusk" \
+ --image-path first_frame.png \
+ --num-frames 61 \
+ --save-output \
+ --output-path outputs
+```
+
+The image is used as the first-frame condition.
+
+## 5. Notes
+
+- `num_frames` must map to a whole number of causal blocks. The latent frame count is `(num_frames - 1) / 4 + 1` and must be divisible by 8. For example, 61, 125, and 189 frames give 16, 32, and 48 latent frames.
+- SGLang supports T2V sizes 1280x704, 704x1280, 832x480, and 480x832.
+- I2V request images follow the Wan TI2V preprocessing path in SGLang. This is different from the original LongLive dataset resize path.
+- For multi-shot runs, set `num_frames` to match `len(shot_prompts) * chunks_per_shot * 8` latent frames, that is `num_frames = (len(shot_prompts) * chunks_per_shot * 8 - 1) * 4 + 1`.
diff --git a/docs_new/docs.json b/docs_new/docs.json
index 620879079..f286e0a86 100644
--- a/docs_new/docs.json
+++ b/docs_new/docs.json
@@ -1204,6 +1204,13 @@
"cookbook/diffusion/Wan/Wan2.2"
]
},
+ {
+ "group": "LongLive",
+ "tag": "NEW",
+ "pages": [
+ "cookbook/diffusion/LongLive/LongLive-2.0"
+ ]
+ },
{
"group": "LTX",
"pages": [
diff --git a/docs_new/docs/sglang-diffusion/compatibility_matrix.mdx b/docs_new/docs/sglang-diffusion/compatibility_matrix.mdx
index 6dea94417..b54622796 100644
--- a/docs_new/docs/sglang-diffusion/compatibility_matrix.mdx
+++ b/docs_new/docs/sglang-diffusion/compatibility_matrix.mdx
@@ -85,6 +85,12 @@ Rows are grouped when a family shares the same runtime path or optimization supp
TI2V / T2V / I2V, 480p / 720p |
SageLaserBSARain Fusion |
+
+ | LongLive 2.0 |
+ Rabinovich/LongLive-2.0-5B-Diffusers
|
+ T2V / I2V, 480p / 720p |
+ No dedicated optimization listed |
+
| HunyuanVideo |
hunyuanvideo-community/HunyuanVideoFastVideo/FastHunyuan-diffusers
|
@@ -272,6 +278,21 @@ Optimization columns are abbreviated to keep the matrix readable:
✅ |
✅ |
+
+ | LongLive 2.0 5B |
+ Rabinovich/LongLive-2.0-5B-Diffusers |
+ 480p 720p |
+ ❌ |
+ ❌ |
+ ❌ |
+ ❌ |
+ ❌ |
+ ❌ |
+ ❌ |
+ ❌ |
+ ❌ |
+ ❌ |
+
| Wan2.2 T2V A14B |
Wan-AI/Wan2.2-T2V-A14B-Diffusers
nvidia/Wan2.2-T2V-A14B-Diffusers-NVFP4 |
diff --git a/python/sglang/multimodal_gen/configs/models/dits/__init__.py b/python/sglang/multimodal_gen/configs/models/dits/__init__.py
index 39ab3fd6d..94e7b98ac 100644
--- a/python/sglang/multimodal_gen/configs/models/dits/__init__.py
+++ b/python/sglang/multimodal_gen/configs/models/dits/__init__.py
@@ -8,6 +8,7 @@ from sglang.multimodal_gen.configs.models.dits.ideogram import Ideogram4DiTConfi
from sglang.multimodal_gen.configs.models.dits.lingbot_world import (
LingBotWorldVideoConfig,
)
+from sglang.multimodal_gen.configs.models.dits.longlive2 import LongLive2VideoConfig
from sglang.multimodal_gen.configs.models.dits.mova_audio import MOVAAudioConfig
from sglang.multimodal_gen.configs.models.dits.mova_video import MOVAVideoConfig
from sglang.multimodal_gen.configs.models.dits.stablediffusion3 import (
@@ -21,6 +22,7 @@ __all__ = [
"HunyuanVideoConfig",
"Ideogram4DiTConfig",
"LingBotWorldVideoConfig",
+ "LongLive2VideoConfig",
"WanVideoConfig",
"Hunyuan3DDiTConfig",
"MOVAAudioConfig",
diff --git a/python/sglang/multimodal_gen/configs/models/dits/longlive2.py b/python/sglang/multimodal_gen/configs/models/dits/longlive2.py
new file mode 100644
index 000000000..fbc7d2e50
--- /dev/null
+++ b/python/sglang/multimodal_gen/configs/models/dits/longlive2.py
@@ -0,0 +1,85 @@
+# SPDX-License-Identifier: Apache-2.0
+# Adapted from https://github.com/NVlabs/LongLive
+
+from dataclasses import dataclass, field
+
+from sglang.multimodal_gen.configs.models.dits.base import DiTArchConfig
+from sglang.multimodal_gen.configs.models.dits.wanvideo import (
+ WanVideoArchConfig,
+ WanVideoConfig,
+)
+
+
+@dataclass
+class LongLive2ArchConfig(WanVideoArchConfig):
+ param_names_mapping: dict = field(
+ default_factory=lambda: {
+ r"^model\.patch_embedding\.(.*)$": r"patch_embedding.proj.\1",
+ r"^model\.text_embedding\.0\.(.*)$": r"condition_embedder.text_embedder.fc_in.\1",
+ r"^model\.text_embedding\.2\.(.*)$": r"condition_embedder.text_embedder.fc_out.\1",
+ r"^model\.time_embedding\.0\.(.*)$": r"condition_embedder.time_embedder.mlp.fc_in.\1",
+ r"^model\.time_embedding\.2\.(.*)$": r"condition_embedder.time_embedder.mlp.fc_out.\1",
+ r"^model\.time_projection\.1\.(.*)$": r"condition_embedder.time_modulation.linear.\1",
+ r"^model\.blocks\.(\d+)\.modulation$": r"blocks.\1.scale_shift_table",
+ r"^model\.blocks\.(\d+)\.self_attn\.q\.(.*)$": r"blocks.\1.to_q.\2",
+ r"^model\.blocks\.(\d+)\.self_attn\.k\.(.*)$": r"blocks.\1.to_k.\2",
+ r"^model\.blocks\.(\d+)\.self_attn\.v\.(.*)$": r"blocks.\1.to_v.\2",
+ r"^model\.blocks\.(\d+)\.self_attn\.o\.(.*)$": r"blocks.\1.to_out.\2",
+ r"^model\.blocks\.(\d+)\.self_attn\.norm_q\.(.*)$": r"blocks.\1.norm_q.\2",
+ r"^model\.blocks\.(\d+)\.self_attn\.norm_k\.(.*)$": r"blocks.\1.norm_k.\2",
+ r"^model\.blocks\.(\d+)\.norm3\.(.*)$": r"blocks.\1.self_attn_residual_norm.norm.\2",
+ r"^model\.blocks\.(\d+)\.cross_attn\.q\.(.*)$": r"blocks.\1.attn2.to_q.\2",
+ r"^model\.blocks\.(\d+)\.cross_attn\.k\.(.*)$": r"blocks.\1.attn2.to_k.\2",
+ r"^model\.blocks\.(\d+)\.cross_attn\.v\.(.*)$": r"blocks.\1.attn2.to_v.\2",
+ r"^model\.blocks\.(\d+)\.cross_attn\.o\.(.*)$": r"blocks.\1.attn2.to_out.\2",
+ r"^model\.blocks\.(\d+)\.cross_attn\.norm_q\.(.*)$": r"blocks.\1.attn2.norm_q.\2",
+ r"^model\.blocks\.(\d+)\.cross_attn\.norm_k\.(.*)$": r"blocks.\1.attn2.norm_k.\2",
+ r"^model\.blocks\.(\d+)\.ffn\.0\.(.*)$": r"blocks.\1.ffn.fc_in.\2",
+ r"^model\.blocks\.(\d+)\.ffn\.2\.(.*)$": r"blocks.\1.ffn.fc_out.\2",
+ r"^model\.head\.modulation$": r"scale_shift_table",
+ r"^model\.head\.head\.(.*)$": r"proj_out.\1",
+ }
+ )
+ reverse_param_names_mapping: dict = field(
+ default_factory=lambda: {
+ r"^patch_embedding\.proj\.(.*)$": r"model.patch_embedding.\1",
+ r"^condition_embedder\.text_embedder\.fc_in\.(.*)$": r"model.text_embedding.0.\1",
+ r"^condition_embedder\.text_embedder\.fc_out\.(.*)$": r"model.text_embedding.2.\1",
+ r"^condition_embedder\.time_embedder\.mlp\.fc_in\.(.*)$": r"model.time_embedding.0.\1",
+ r"^condition_embedder\.time_embedder\.mlp\.fc_out\.(.*)$": r"model.time_embedding.2.\1",
+ r"^condition_embedder\.time_modulation\.linear\.(.*)$": r"model.time_projection.1.\1",
+ r"^blocks\.(\d+)\.scale_shift_table$": r"model.blocks.\1.modulation",
+ r"^blocks\.(\d+)\.to_q\.(.*)$": r"model.blocks.\1.self_attn.q.\2",
+ r"^blocks\.(\d+)\.to_k\.(.*)$": r"model.blocks.\1.self_attn.k.\2",
+ r"^blocks\.(\d+)\.to_v\.(.*)$": r"model.blocks.\1.self_attn.v.\2",
+ r"^blocks\.(\d+)\.to_out\.(.*)$": r"model.blocks.\1.self_attn.o.\2",
+ r"^blocks\.(\d+)\.norm_q\.(.*)$": r"model.blocks.\1.self_attn.norm_q.\2",
+ r"^blocks\.(\d+)\.norm_k\.(.*)$": r"model.blocks.\1.self_attn.norm_k.\2",
+ r"^blocks\.(\d+)\.self_attn_residual_norm\.norm\.(.*)$": r"model.blocks.\1.norm3.\2",
+ r"^blocks\.(\d+)\.attn2\.to_q\.(.*)$": r"model.blocks.\1.cross_attn.q.\2",
+ r"^blocks\.(\d+)\.attn2\.to_k\.(.*)$": r"model.blocks.\1.cross_attn.k.\2",
+ r"^blocks\.(\d+)\.attn2\.to_v\.(.*)$": r"model.blocks.\1.cross_attn.v.\2",
+ r"^blocks\.(\d+)\.attn2\.to_out\.(.*)$": r"model.blocks.\1.cross_attn.o.\2",
+ r"^blocks\.(\d+)\.attn2\.norm_q\.(.*)$": r"model.blocks.\1.cross_attn.norm_q.\2",
+ r"^blocks\.(\d+)\.attn2\.norm_k\.(.*)$": r"model.blocks.\1.cross_attn.norm_k.\2",
+ r"^blocks\.(\d+)\.ffn\.fc_in\.(.*)$": r"model.blocks.\1.ffn.0.\2",
+ r"^blocks\.(\d+)\.ffn\.fc_out\.(.*)$": r"model.blocks.\1.ffn.2.\2",
+ r"^scale_shift_table$": r"model.head.modulation",
+ r"^proj_out\.(.*)$": r"model.head.head.\1",
+ }
+ )
+ num_attention_heads: int = 24
+ attention_head_dim: int = 128
+ in_channels: int = 48
+ out_channels: int = 48
+ ffn_dim: int = 14336
+ num_layers: int = 30
+ local_attn_size: int = 32
+ sink_size: int = 8
+ num_frames_per_block: int = 8
+ sliding_window_num_frames: int = 32
+
+
+@dataclass
+class LongLive2VideoConfig(WanVideoConfig):
+ arch_config: DiTArchConfig = field(default_factory=LongLive2ArchConfig)
diff --git a/python/sglang/multimodal_gen/configs/pipeline_configs/longlive2.py b/python/sglang/multimodal_gen/configs/pipeline_configs/longlive2.py
new file mode 100644
index 000000000..730f17b97
--- /dev/null
+++ b/python/sglang/multimodal_gen/configs/pipeline_configs/longlive2.py
@@ -0,0 +1,59 @@
+# SPDX-License-Identifier: Apache-2.0
+# Adapted from https://github.com/NVlabs/LongLive
+
+from dataclasses import dataclass, field
+
+from sglang.multimodal_gen.configs.models import DiTConfig
+from sglang.multimodal_gen.configs.models.dits.longlive2 import LongLive2VideoConfig
+from sglang.multimodal_gen.configs.pipeline_configs.base import ModelTaskType
+from sglang.multimodal_gen.configs.pipeline_configs.wan import Wan2_2_TI2V_5B_Config
+from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
+
+logger = init_logger(__name__)
+
+
+@dataclass
+class LongLive2T2VConfig(Wan2_2_TI2V_5B_Config):
+
+ is_causal: bool = True
+ task_type: ModelTaskType = ModelTaskType.TI2V
+ vae_precision: str = "bf16"
+
+ flow_shift: float | None = 5.0
+ dmd_denoising_steps: list[int] | None = field(
+ default_factory=lambda: [1000, 750, 500, 250]
+ )
+ expand_timesteps: bool = False
+
+ dit_config: DiTConfig = field(default_factory=LongLive2VideoConfig)
+
+ def adjust_num_frames(self, num_frames: int) -> int:
+ num_frames = super().adjust_num_frames(num_frames)
+ vae_scale_factor_temporal = self.vae_config.arch_config.scale_factor_temporal
+ latent_frames = (num_frames - 1) // vae_scale_factor_temporal + 1
+ block_size = self.dit_config.arch_config.num_frames_per_block
+ if latent_frames % block_size == 0:
+ return num_frames
+
+ adjusted_latent_frames = max(
+ block_size, latent_frames // block_size * block_size
+ )
+ adjusted_num_frames = (
+ adjusted_latent_frames - 1
+ ) * vae_scale_factor_temporal + 1
+ logger.warning(
+ "`num_frames` must map to latent frames divisible by %s for "
+ "LongLive2 causal denoising. Rounding from %s to %s.",
+ block_size,
+ num_frames,
+ adjusted_num_frames,
+ )
+ return adjusted_num_frames
+
+ def postprocess_image_latent(self, latent_condition, batch):
+ return latent_condition[:, :, :1]
+
+ def __post_init__(self) -> None:
+ super().__post_init__()
+ self.vae_config.load_encoder = True
+ self.vae_config.load_decoder = True
diff --git a/python/sglang/multimodal_gen/configs/sample/longlive2.py b/python/sglang/multimodal_gen/configs/sample/longlive2.py
new file mode 100644
index 000000000..c0dee0b40
--- /dev/null
+++ b/python/sglang/multimodal_gen/configs/sample/longlive2.py
@@ -0,0 +1,74 @@
+# SPDX-License-Identifier: Apache-2.0
+from dataclasses import dataclass, field
+
+from sglang.multimodal_gen.configs.sample.wan import Wan2_2_TI2V_5B_SamplingParam
+
+
+@dataclass
+class LongLive2SamplingParams(Wan2_2_TI2V_5B_SamplingParam):
+ height: int = 704
+ width: int = 1280
+ fps: int = 24
+ num_inference_steps: int = 4
+ guidance_scale: float = 1.0
+ num_frames: int = 61
+ shot_prompts: list[str] | None = field(
+ default=None, metadata={"batch_sig_exclude": True}
+ )
+ shot_durations: list[int] | None = field(
+ default=None, metadata={"batch_sig_exclude": True}
+ )
+ chunks_per_shot: int = field(default=0, metadata={"batch_sig_exclude": True})
+ scene_cut_prefix: str = field(
+ default="The scene transitions. ", metadata={"batch_sig_exclude": True}
+ )
+ multi_shot_sink: bool = field(default=True, metadata={"batch_sig_exclude": True})
+ multi_shot_rope_offset: float = field(
+ default=8.0, metadata={"batch_sig_exclude": True}
+ )
+
+ supported_resolutions: list[tuple[int, int]] | None = field(
+ default_factory=lambda: [
+ (1280, 704),
+ (704, 1280),
+ (832, 480),
+ (480, 832),
+ ]
+ )
+
+ def _validate(self):
+ super()._validate()
+
+ if self.shot_prompts is not None:
+ if not isinstance(self.shot_prompts, list) or not self.shot_prompts:
+ raise ValueError("shot_prompts must be a non-empty list of strings")
+ if not all(
+ isinstance(prompt, str) and prompt for prompt in self.shot_prompts
+ ):
+ raise ValueError("shot_prompts must contain non-empty strings")
+
+ if self.shot_durations is not None:
+ if not isinstance(self.shot_durations, list) or not self.shot_durations:
+ raise ValueError("shot_durations must be a non-empty list of ints")
+ if self.shot_prompts is not None and len(self.shot_durations) != len(
+ self.shot_prompts
+ ):
+ raise ValueError("shot_durations must match shot_prompts length")
+ if not all(
+ isinstance(duration, int) and duration > 0
+ for duration in self.shot_durations
+ ):
+ raise ValueError("shot_durations must contain positive ints")
+
+ if self.chunks_per_shot < 0:
+ raise ValueError("chunks_per_shot must be non-negative")
+
+ if self.scene_cut_prefix is None:
+ self.scene_cut_prefix = ""
+ if self.multi_shot_rope_offset < 0:
+ raise ValueError("multi_shot_rope_offset must be non-negative")
+
+ def _adjust(self, server_args):
+ if self.shot_prompts is not None and self.prompt is None:
+ self.prompt = self.shot_prompts[0]
+ super()._adjust(server_args)
diff --git a/python/sglang/multimodal_gen/registry.py b/python/sglang/multimodal_gen/registry.py
index c5f06ef6f..f2c8a089b 100644
--- a/python/sglang/multimodal_gen/registry.py
+++ b/python/sglang/multimodal_gen/registry.py
@@ -68,6 +68,7 @@ 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.longlive2 import LongLive2T2VConfig
from sglang.multimodal_gen.configs.pipeline_configs.ltx_2 import (
LTX2PipelineConfig,
LTX23PipelineConfig,
@@ -129,6 +130,7 @@ from sglang.multimodal_gen.configs.sample.krea2 import (
from sglang.multimodal_gen.configs.sample.lingbot_world import (
LingBotWorldSamplingParams,
)
+from sglang.multimodal_gen.configs.sample.longlive2 import LongLive2SamplingParams
from sglang.multimodal_gen.configs.sample.ltx_2 import (
LTX2SamplingParams,
LTX23HQSamplingParams,
@@ -788,6 +790,15 @@ def _register_configs():
"robbyant/lingbot-world-v2-14b-causal-fast-diffusers",
],
)
+ register_configs(
+ sampling_param_cls=LongLive2SamplingParams,
+ pipeline_config_cls=LongLive2T2VConfig,
+ hf_model_paths=[
+ # Since LongLive-2.0-5B does not have official diffusers release
+ "Rabinovich/LongLive-2.0-5B-Diffusers",
+ "Efficient-Large-Model/LongLive-2.0-5B",
+ ],
+ )
register_configs(
sampling_param_cls=FastWanT2V480PConfig,
pipeline_config_cls=FastWan2_1_T2V_480P_Config,
diff --git a/python/sglang/multimodal_gen/runtime/layers/kvcache/causal_attention_cache.py b/python/sglang/multimodal_gen/runtime/layers/kvcache/causal_attention_cache.py
index 2d6ae23d7..65e38f604 100644
--- a/python/sglang/multimodal_gen/runtime/layers/kvcache/causal_attention_cache.py
+++ b/python/sglang/multimodal_gen/runtime/layers/kvcache/causal_attention_cache.py
@@ -32,6 +32,9 @@ class CausalSelfAttentionKVCache:
sink_tokens: int = 0
attention_window_size: int = 0
allow_growth: bool = False
+ global_sink_tokens: int = 0
+ pinned_start: int = -1
+ pinned_len: int = 0
def __post_init__(self) -> None:
if self.cache_size == 0:
@@ -46,6 +49,29 @@ class CausalSelfAttentionKVCache:
self.global_end_index_int = 0
if self.local_end_index_int is not None:
self.local_end_index_int = 0
+ self.reset_pinned_sink()
+
+ def reset_pinned_sink(self) -> None:
+ self.pinned_start = -1
+ self.pinned_len = 0
+
+ def pin_current_chunk(self, current_num_tokens: int) -> None:
+ if self.sink_tokens <= 0 or current_num_tokens <= 0:
+ self.reset_pinned_sink()
+ return
+ _, local_end_index = self._read_indices()
+ self.pinned_start = local_end_index - current_num_tokens
+ self.pinned_len = min(self.sink_tokens, current_num_tokens)
+
+ def _has_pinned_sink(self) -> bool:
+ return self.pinned_start >= 0 and self.pinned_len > 0
+
+ def _effective_sink_tokens(self) -> int:
+ if self._has_pinned_sink():
+ if self.pinned_start == self.global_sink_tokens:
+ return self.global_sink_tokens + self.pinned_len
+ return self.global_sink_tokens
+ return max(self.global_sink_tokens, self.sink_tokens)
def _read_indices(self) -> tuple[int, int]:
global_end_index = self.global_end_index_int
@@ -141,7 +167,7 @@ class CausalSelfAttentionKVCache:
)
current_chunk_end = current_chunk_start + num_new_tokens
kv_cache_size = self.cache_size
- sink_tokens = self.sink_tokens
+ sink_tokens = self._effective_sink_tokens()
global_end_index, local_end_index_prev = self._read_indices()
# local_start(/end)_index: the local position of the start/end of current chunk
@@ -236,6 +262,9 @@ class CausalSelfAttentionKVCache:
:,
].clone()
+ if self._has_pinned_sink() and self.pinned_start >= sink_tokens:
+ self.pinned_start -= num_evicted_tokens
+
# if we move the minimum number of tokens, the right bound of the append token would be aligned with end of the buffer
local_end_index = kv_cache_size
else:
@@ -329,72 +358,159 @@ class CausalSelfAttentionKVCache:
heads.
"""
if recent_window_tokens is None:
- if cache_head_slice is None:
- return (
- self.k[:, attn_start_index:updated_local_end],
- self.v[:, attn_start_index:updated_local_end],
+ if self.global_sink_tokens > 0 or self._has_pinned_sink():
+ return self._pinned_attention_view(
+ attn_start_index=attn_start_index,
+ updated_local_end=updated_local_end,
+ cache_head_slice=cache_head_slice,
)
- return (
- self.k[:, attn_start_index:updated_local_end, cache_head_slice, :],
- self.v[:, attn_start_index:updated_local_end, cache_head_slice, :],
+ return self._cache_slice(
+ slice(attn_start_index, updated_local_end),
+ cache_head_slice=cache_head_slice,
)
if recent_window_tokens < 0:
raise ValueError("recent_window_tokens must be non-negative or None")
- sink_end = min(self.sink_tokens, updated_local_end)
+ sink_end = min(self._effective_sink_tokens(), updated_local_end)
recent_start = max(sink_end, local_start_index - recent_window_tokens)
if recent_start <= sink_end:
- if cache_head_slice is None:
- return self.k[:, :updated_local_end], self.v[:, :updated_local_end]
- return (
- self.k[:, :updated_local_end, cache_head_slice, :],
- self.v[:, :updated_local_end, cache_head_slice, :],
- )
- if sink_end <= 0:
- if cache_head_slice is None:
- return (
- self.k[:, recent_start:updated_local_end],
- self.v[:, recent_start:updated_local_end],
- )
- return (
- self.k[:, recent_start:updated_local_end, cache_head_slice, :],
- self.v[:, recent_start:updated_local_end, cache_head_slice, :],
+ return self._cache_slice(
+ slice(0, updated_local_end),
+ cache_head_slice=cache_head_slice,
)
+ cache_slices = []
+ if sink_end > 0:
+ cache_slices.append(slice(0, sink_end))
+ if (
+ self._has_pinned_sink()
+ and self.pinned_start >= sink_end
+ and self.pinned_start < recent_start
+ ):
+ cache_slices.append(
+ slice(self.pinned_start, self.pinned_start + self.pinned_len)
+ )
+ cache_slices.append(slice(recent_start, updated_local_end))
+ return self._cat_cache_slices(
+ cache_slices,
+ cache_head_slice=cache_head_slice,
+ )
+
+ def _cache_slice(
+ self,
+ cache_slice: slice,
+ *,
+ cache_head_slice: slice | None = None,
+ ) -> tuple[torch.Tensor, torch.Tensor]:
+ if cache_head_slice is None:
+ return self.k[:, cache_slice], self.v[:, cache_slice]
+ return (
+ self.k[:, cache_slice, cache_head_slice, :],
+ self.v[:, cache_slice, cache_head_slice, :],
+ )
+
+ def _cat_cache_slices(
+ self,
+ cache_slices: list[slice],
+ *,
+ cache_head_slice: slice | None = None,
+ ) -> tuple[torch.Tensor, torch.Tensor]:
+ if len(cache_slices) == 1:
+ return self._cache_slice(
+ cache_slices[0],
+ cache_head_slice=cache_head_slice,
+ )
if cache_head_slice is None:
return (
torch.cat(
- [
- self.k[:, :sink_end],
- self.k[:, recent_start:updated_local_end],
- ],
- dim=1,
+ [self.k[:, cache_slice] for cache_slice in cache_slices], dim=1
),
torch.cat(
- [
- self.v[:, :sink_end],
- self.v[:, recent_start:updated_local_end],
- ],
- dim=1,
+ [self.v[:, cache_slice] for cache_slice in cache_slices], dim=1
),
)
return (
torch.cat(
[
- self.k[:, :sink_end, cache_head_slice, :],
- self.k[:, recent_start:updated_local_end, cache_head_slice, :],
+ self.k[:, cache_slice, cache_head_slice, :]
+ for cache_slice in cache_slices
],
dim=1,
),
torch.cat(
[
- self.v[:, :sink_end, cache_head_slice, :],
- self.v[:, recent_start:updated_local_end, cache_head_slice, :],
+ self.v[:, cache_slice, cache_head_slice, :]
+ for cache_slice in cache_slices
],
dim=1,
),
)
+ def _pinned_attention_view(
+ self,
+ *,
+ attn_start_index: int,
+ updated_local_end: int,
+ cache_head_slice: slice | None = None,
+ ) -> tuple[torch.Tensor, torch.Tensor]:
+ effective_sink_tokens = self._effective_sink_tokens()
+ prepend_sink = effective_sink_tokens > 0 and attn_start_index > 0
+ prepend_pinned = (
+ self._has_pinned_sink()
+ and self.pinned_start >= effective_sink_tokens
+ and self.pinned_start < attn_start_index
+ )
+
+ if prepend_sink and prepend_pinned:
+ extra_tokens = effective_sink_tokens + self.pinned_len
+ local_window_size = max(0, self.attention_window_size - extra_tokens)
+ local_window_start = max(
+ effective_sink_tokens,
+ updated_local_end - local_window_size,
+ )
+ cache_slices = [
+ slice(0, effective_sink_tokens),
+ slice(self.pinned_start, self.pinned_start + self.pinned_len),
+ slice(local_window_start, updated_local_end),
+ ]
+ return self._cat_cache_slices(
+ cache_slices,
+ cache_head_slice=cache_head_slice,
+ )
+
+ if prepend_sink:
+ local_window_size = max(
+ 0,
+ self.attention_window_size - effective_sink_tokens,
+ )
+ local_window_start = max(
+ effective_sink_tokens,
+ updated_local_end - local_window_size,
+ )
+ return self._cat_cache_slices(
+ [
+ slice(0, effective_sink_tokens),
+ slice(local_window_start, updated_local_end),
+ ],
+ cache_head_slice=cache_head_slice,
+ )
+
+ if prepend_pinned:
+ local_window_size = max(0, self.attention_window_size - self.pinned_len)
+ local_window_start = max(0, updated_local_end - local_window_size)
+ return self._cat_cache_slices(
+ [
+ slice(self.pinned_start, self.pinned_start + self.pinned_len),
+ slice(local_window_start, updated_local_end),
+ ],
+ cache_head_slice=cache_head_slice,
+ )
+
+ return self._cache_slice(
+ slice(attn_start_index, updated_local_end),
+ cache_head_slice=cache_head_slice,
+ )
+
@dataclass(slots=True)
class CrossAttentionKVCache:
diff --git a/python/sglang/multimodal_gen/runtime/models/dits/causal_wanvideo.py b/python/sglang/multimodal_gen/runtime/models/dits/causal_wanvideo.py
index 0652945ef..4e8b300c2 100644
--- a/python/sglang/multimodal_gen/runtime/models/dits/causal_wanvideo.py
+++ b/python/sglang/multimodal_gen/runtime/models/dits/causal_wanvideo.py
@@ -521,7 +521,9 @@ class CausalWanTransformer3DModel(BaseDiT, LayerwiseOffloadableModuleMixin):
# Causal-specific
self.block_mask = None
self.num_frame_per_block = config.arch_config.num_frames_per_block
- assert self.num_frame_per_block <= 3
+ # Block size is bounded only by the causal block-mask construction, which
+ # supports any positive value.
+ assert self.num_frame_per_block >= 1
self.independent_first_frame = False
self.__post_init__()
diff --git a/python/sglang/multimodal_gen/runtime/models/dits/longlive2.py b/python/sglang/multimodal_gen/runtime/models/dits/longlive2.py
new file mode 100644
index 000000000..ac62a9f19
--- /dev/null
+++ b/python/sglang/multimodal_gen/runtime/models/dits/longlive2.py
@@ -0,0 +1,190 @@
+# SPDX-License-Identifier: Apache-2.0
+from typing import Any
+
+import torch
+import torch.nn as nn
+from torch.nn.attention.flex_attention import BlockMask
+
+from sglang.multimodal_gen.configs.models.dits.longlive2 import LongLive2VideoConfig
+from sglang.multimodal_gen.runtime.layers.kvcache.causal_attention_cache import (
+ CausalSelfAttentionKVCache,
+ CrossAttentionKVCache,
+)
+from sglang.multimodal_gen.runtime.layers.layernorm import (
+ LayerNormScaleShift,
+ tensor_parallel_rms_norm,
+)
+from sglang.multimodal_gen.runtime.layers.quantization.configs.base_config import (
+ QuantizationConfig,
+)
+from sglang.multimodal_gen.runtime.models.dits.causal_wanvideo import (
+ CausalWanTransformer3DModel,
+ CausalWanTransformerBlock,
+)
+
+
+class LongLive2CausalWanTransformerBlock(CausalWanTransformerBlock):
+ def __init__(self, *args, **kwargs):
+ super().__init__(*args, **kwargs)
+ self.norm1 = LayerNormScaleShift(
+ self.hidden_dim,
+ eps=self.norm1.eps,
+ elementwise_affine=False,
+ dtype=torch.float32,
+ )
+
+ def _cross_attn_with_cache(
+ self,
+ hidden_states: torch.Tensor,
+ encoder_hidden_states: torch.Tensor,
+ crossattn_cache: CrossAttentionKVCache | None,
+ ) -> torch.Tensor:
+ attn2 = self.attn2
+ q, _ = attn2.to_q(hidden_states)
+ if attn2.tp_rmsnorm:
+ q = tensor_parallel_rms_norm(q, attn2.norm_q)
+ else:
+ q = attn2.norm_q(q)
+ q = q.unflatten(2, (attn2.local_num_heads, attn2.head_dim))
+
+ if crossattn_cache is not None and crossattn_cache.is_init:
+ k = crossattn_cache.k
+ v = crossattn_cache.v
+ else:
+ k, _ = attn2.to_k(encoder_hidden_states)
+ if attn2.tp_rmsnorm:
+ k = tensor_parallel_rms_norm(k, attn2.norm_k)
+ else:
+ k = attn2.norm_k(k)
+ k = k.unflatten(2, (attn2.local_num_heads, attn2.head_dim))
+
+ v, _ = attn2.to_v(encoder_hidden_states)
+ v = v.unflatten(2, (attn2.local_num_heads, attn2.head_dim))
+
+ if crossattn_cache is not None:
+ crossattn_cache.store(k, v)
+
+ hidden_states = attn2.attn(q, k, v)
+ hidden_states = hidden_states.flatten(2)
+ hidden_states, _ = attn2.to_out(hidden_states)
+ return hidden_states
+
+ def forward(
+ self,
+ hidden_states: torch.Tensor,
+ encoder_hidden_states: torch.Tensor,
+ temb: torch.Tensor,
+ freqs_cis: tuple[torch.Tensor, torch.Tensor],
+ block_mask: BlockMask,
+ kv_cache: CausalSelfAttentionKVCache | None = None,
+ crossattn_cache: CrossAttentionKVCache | None = None,
+ current_start: int = 0,
+ cache_start: int | None = None,
+ ) -> torch.Tensor:
+ if hidden_states.dim() == 4:
+ hidden_states = hidden_states.squeeze(1)
+ num_frames = temb.shape[1]
+ bs, _, _ = hidden_states.shape
+ orig_dtype = hidden_states.dtype
+ e = self.scale_shift_table + temb.float()
+ assert e.shape == (bs, num_frames, 6, self.hidden_dim)
+ shift_msa, scale_msa, gate_msa, c_shift_msa, c_scale_msa, c_gate_msa = e.chunk(
+ 6, dim=2
+ )
+ assert shift_msa.dtype == torch.float32
+
+ norm_hidden_states = self.norm1(hidden_states, shift_msa, scale_msa)
+ query, _ = self.to_q(norm_hidden_states)
+ key, _ = self.to_k(norm_hidden_states)
+ value, _ = self.to_v(norm_hidden_states)
+
+ if self.norm_q is not None:
+ query = self.norm_q(query)
+ if self.norm_k is not None:
+ key = self.norm_k(key)
+
+ query = query.squeeze(1).unflatten(2, (self.num_attention_heads, -1))
+ key = key.squeeze(1).unflatten(2, (self.num_attention_heads, -1))
+ value = value.squeeze(1).unflatten(2, (self.num_attention_heads, -1))
+
+ attn_output = self.attn1(
+ query,
+ key,
+ value,
+ freqs_cis,
+ block_mask,
+ kv_cache,
+ current_start,
+ cache_start,
+ )
+ attn_output = attn_output.flatten(2)
+ attn_output, _ = self.to_out(attn_output)
+ attn_output = attn_output.squeeze(1)
+
+ null_shift = null_scale = torch.zeros(
+ (1,), device=hidden_states.device, dtype=hidden_states.dtype
+ )
+ norm_hidden_states, hidden_states = self.self_attn_residual_norm(
+ hidden_states, attn_output, gate_msa, null_shift, null_scale
+ )
+ norm_hidden_states, hidden_states = norm_hidden_states.to(
+ orig_dtype
+ ), hidden_states.to(orig_dtype)
+
+ attn_output = self._cross_attn_with_cache(
+ norm_hidden_states,
+ encoder_hidden_states,
+ crossattn_cache,
+ )
+ norm_hidden_states, hidden_states = self.cross_attn_residual_norm(
+ hidden_states, attn_output, 1, c_shift_msa, c_scale_msa
+ )
+ norm_hidden_states, hidden_states = norm_hidden_states.to(
+ orig_dtype
+ ), hidden_states.to(orig_dtype)
+
+ ff_output = self.ffn(norm_hidden_states)
+ hidden_states = self.mlp_residual(ff_output, c_gate_msa, hidden_states)
+ hidden_states = hidden_states.to(orig_dtype)
+
+ return hidden_states
+
+
+class LongLive2Transformer3DModel(CausalWanTransformer3DModel):
+ _fsdp_shard_conditions = LongLive2VideoConfig()._fsdp_shard_conditions
+ _compile_conditions = LongLive2VideoConfig()._compile_conditions
+ _supported_attention_backends = LongLive2VideoConfig()._supported_attention_backends
+ param_names_mapping = LongLive2VideoConfig().param_names_mapping
+ reverse_param_names_mapping = LongLive2VideoConfig().reverse_param_names_mapping
+ lora_param_names_mapping = LongLive2VideoConfig().lora_param_names_mapping
+
+ def __init__(
+ self,
+ config: LongLive2VideoConfig,
+ hf_config: dict[str, Any],
+ quant_config: QuantizationConfig | None = None,
+ ) -> None:
+ super().__init__(config=config, hf_config=hf_config, quant_config=quant_config)
+ inner_dim = config.num_attention_heads * config.attention_head_dim
+ self.blocks = nn.ModuleList(
+ [
+ LongLive2CausalWanTransformerBlock(
+ inner_dim,
+ config.ffn_dim,
+ config.num_attention_heads,
+ config.local_attn_size,
+ config.sink_size,
+ config.qk_norm,
+ config.cross_attn_norm,
+ config.eps,
+ config.added_kv_proj_dim,
+ self._supported_attention_backends,
+ prefix=f"{config.prefix}.blocks.{i}",
+ quant_config=quant_config,
+ )
+ for i in range(config.num_layers)
+ ]
+ )
+
+
+EntryClass = LongLive2Transformer3DModel
diff --git a/python/sglang/multimodal_gen/runtime/pipelines/longlive2_pipeline.py b/python/sglang/multimodal_gen/runtime/pipelines/longlive2_pipeline.py
new file mode 100644
index 000000000..ae0ee1ffa
--- /dev/null
+++ b/python/sglang/multimodal_gen/runtime/pipelines/longlive2_pipeline.py
@@ -0,0 +1,61 @@
+# SPDX-License-Identifier: Apache-2.0
+from sglang.multimodal_gen.configs.pipeline_configs.longlive2 import LongLive2T2VConfig
+from sglang.multimodal_gen.configs.sample.longlive2 import LongLive2SamplingParams
+from sglang.multimodal_gen.runtime.models.schedulers.scheduling_flow_unipc_multistep import (
+ FlowUniPCMultistepScheduler,
+)
+from sglang.multimodal_gen.runtime.pipelines.wan_causal_dmd_pipeline import (
+ WanCausalDMDPipeline,
+)
+from sglang.multimodal_gen.runtime.pipelines_core.stages import InputValidationStage
+from sglang.multimodal_gen.runtime.pipelines_core.stages.model_specific_stages.longlive2 import (
+ LongLive2CausalDenoisingStage,
+ LongLive2ImageVAEEncodingStage,
+ LongLive2LatentPreparationStage,
+ LongLive2TextEncodingStage,
+)
+from sglang.multimodal_gen.runtime.server_args import ServerArgs
+
+
+class LongLive2Pipeline(WanCausalDMDPipeline):
+ pipeline_name = "LongLive2Pipeline"
+ pipeline_config_cls = LongLive2T2VConfig
+ sampling_params_cls = LongLive2SamplingParams
+
+ def initialize_pipeline(self, server_args: ServerArgs):
+ self.modules["scheduler"] = FlowUniPCMultistepScheduler(
+ num_train_timesteps=1000,
+ shift=1,
+ use_dynamic_shifting=False,
+ )
+
+ def create_pipeline_stages(self, server_args: ServerArgs) -> None:
+ self.add_stage(InputValidationStage())
+ self.add_stage(
+ LongLive2TextEncodingStage(
+ text_encoders=[self.get_module("text_encoder")],
+ tokenizers=[self.get_module("tokenizer")],
+ )
+ )
+ self.add_stage(
+ LongLive2ImageVAEEncodingStage(
+ vae=self.get_module("vae"),
+ component_name="vae",
+ )
+ )
+ self.add_stage(
+ LongLive2LatentPreparationStage(
+ scheduler=self.get_module("scheduler"),
+ transformer=self.get_module("transformer"),
+ )
+ )
+ self.add_stage(
+ LongLive2CausalDenoisingStage(
+ transformer=self.get_module("transformer"),
+ scheduler=self.get_module("scheduler"),
+ ),
+ )
+ self.add_standard_decoding_stage()
+
+
+EntryClass = LongLive2Pipeline
diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/causal_denoising.py b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/causal_denoising.py
index af1eb5f39..b2d0cf759 100644
--- a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/causal_denoising.py
+++ b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/causal_denoising.py
@@ -41,6 +41,65 @@ from sglang.multimodal_gen.runtime.utils.precision import (
logger = init_logger(__name__)
+CAUSAL_BLOCK_PROMPTS_KEY = "causal_block_prompts"
+CAUSAL_SCENE_CUT_MASK_KEY = "causal_scene_cut_mask"
+CAUSAL_SHOT_INDICES_KEY = "causal_shot_indices"
+
+
+def expand_causal_block_prompts(
+ shot_prompts: list[str],
+ *,
+ num_blocks: int,
+ shot_durations: list[int] | None = None,
+ chunks_per_shot: int = 0,
+ scene_cut_prefix: str = "",
+) -> tuple[list[str], list[bool], list[int]]:
+ if not shot_prompts:
+ raise ValueError("shot_prompts must be non-empty")
+ if num_blocks <= 0:
+ raise ValueError("num_blocks must be positive")
+ if shot_durations is not None and len(shot_durations) != len(shot_prompts):
+ raise ValueError("shot_durations must match shot_prompts length")
+
+ if shot_durations is not None:
+ durations = shot_durations[: len(shot_prompts)]
+ elif chunks_per_shot > 0:
+ durations = [chunks_per_shot] * len(shot_prompts)
+ else:
+ base, extra = divmod(num_blocks, len(shot_prompts))
+ durations = [base + (1 if i < extra else 0) for i in range(len(shot_prompts))]
+
+ clamped: list[int] = []
+ remaining = num_blocks
+ for duration in durations:
+ if remaining <= 0:
+ break
+ take = min(int(duration), remaining)
+ clamped.append(take)
+ remaining -= take
+ if remaining > 0 and clamped:
+ clamped[-1] += remaining
+ if not clamped:
+ clamped = [num_blocks]
+
+ block_prompts: list[str] = []
+ scene_cut_mask: list[bool] = []
+ shot_indices: list[int] = []
+ for shot_idx, (caption, duration) in enumerate(zip(shot_prompts, clamped)):
+ for block_in_shot in range(duration):
+ is_scene_cut = shot_idx > 0 and block_in_shot == 0
+ if is_scene_cut and scene_cut_prefix:
+ block_prompts.append(scene_cut_prefix + caption)
+ else:
+ block_prompts.append(caption)
+ scene_cut_mask.append(is_scene_cut)
+ shot_indices.append(shot_idx)
+ return (
+ block_prompts[:num_blocks],
+ scene_cut_mask[:num_blocks],
+ shot_indices[:num_blocks],
+ )
+
@dataclass(slots=True)
class CausalDMDForwardContext:
@@ -89,6 +148,8 @@ class CausalDMDDenoisingStage(DenoisingStage):
# KV and cross-attention cache state (initialized on first forward)
self.causal_kv_cache: list | None = None
self.crossattn_cache: list | None = None
+ self.causal_kv_cache_neg: list | None = None
+ self.crossattn_cache_neg: list | None = None
# Model-dependent constants (aligned with causal_inference.py assumptions)
self.num_transformer_blocks = self.transformer.config.arch_config.num_layers
self.num_frames_per_block = (
@@ -189,6 +250,85 @@ class CausalDMDDenoisingStage(DenoisingStage):
assert torch.isnan(prompt_embeds[0]).sum() == 0
return prompt_embeds
+ @staticmethod
+ def _block_prompt_count(batch: Req) -> int | None:
+ block_prompts = batch.extra.get(CAUSAL_BLOCK_PROMPTS_KEY)
+ if block_prompts is None:
+ return None
+ return len(block_prompts)
+
+ @classmethod
+ def _select_block_conditioning(cls, value, block_index: int, block_count: int):
+ if isinstance(value, torch.Tensor) and value.shape[:1] == (block_count,):
+ return value[block_index : block_index + 1]
+ if isinstance(value, list):
+ return [
+ cls._select_block_conditioning(item, block_index, block_count)
+ for item in value
+ ]
+ if isinstance(value, tuple):
+ return tuple(
+ cls._select_block_conditioning(item, block_index, block_count)
+ for item in value
+ )
+ if isinstance(value, dict):
+ return {
+ key: cls._select_block_conditioning(item, block_index, block_count)
+ for key, item in value.items()
+ }
+ return value
+
+ @classmethod
+ def _select_block_prompt_embeds(
+ cls,
+ batch: Req,
+ prompt_embeds,
+ block_index: int,
+ ):
+ block_count = cls._block_prompt_count(batch)
+ if block_count is None:
+ return prompt_embeds
+ return cls._select_block_conditioning(prompt_embeds, block_index, block_count)
+
+ @classmethod
+ def _select_block_cond_kwargs(
+ cls,
+ batch: Req,
+ cond_kwargs: dict[str, Any],
+ block_index: int,
+ ) -> dict[str, Any]:
+ block_count = cls._block_prompt_count(batch)
+ if block_count is None:
+ return cond_kwargs
+ return {
+ key: cls._select_block_conditioning(value, block_index, block_count)
+ for key, value in cond_kwargs.items()
+ }
+
+ def _reset_crossattn_cache_for_block(self, batch: Req, *caches) -> None:
+ if self._block_prompt_count(batch) is None:
+ return
+ for cache in caches:
+ if cache is not None:
+ self._reset_crossattn_cache(cache)
+
+ def _validate_block_prompt_count(self, batch: Req, block_sizes: list[int]) -> None:
+ block_count = self._block_prompt_count(batch)
+ if block_count is None:
+ return
+ if block_count != len(block_sizes):
+ raise ValueError(
+ "causal block prompt count must match causal block count, "
+ f"got {block_count} prompts and {len(block_sizes)} blocks"
+ )
+
+ @staticmethod
+ def _shot_index(batch: Req, block_index: int) -> int:
+ shot_indices = batch.extra.get(CAUSAL_SHOT_INDICES_KEY)
+ if not isinstance(shot_indices, list) or block_index >= len(shot_indices):
+ return 0
+ return int(shot_indices[block_index])
+
def _prepare_causal_dmd_forward_context(
self,
batch: Req,
@@ -853,6 +993,93 @@ class CausalDMDDenoisingStage(DenoisingStage):
for cache_block in kv_cache:
cache_block.reset_indices()
+ def _causal_kv_cache_global_sink_tokens_for_batch(self, batch: Req) -> int:
+ return 0
+
+ def _causal_kv_cache_kwargs_for_batch(
+ self,
+ batch: Req,
+ ) -> dict[str, Any] | None:
+ global_sink_tokens = self._causal_kv_cache_global_sink_tokens_for_batch(batch)
+ if global_sink_tokens <= 0:
+ return None
+ return {"global_sink_tokens": global_sink_tokens}
+
+ def _cache_needs_reinit_for_batch(self, kv_cache, batch: Req) -> bool:
+ if kv_cache is None or len(kv_cache) != self.num_transformer_blocks:
+ return True
+ expected_global_sink_tokens = (
+ self._causal_kv_cache_global_sink_tokens_for_batch(batch)
+ )
+ return kv_cache[0].global_sink_tokens != expected_global_sink_tokens
+
+ def _pin_current_chunk(self, kv_cache, current_num_frames: int) -> None:
+ if kv_cache is None:
+ return
+ current_num_tokens = current_num_frames * self.num_token_per_frame
+ for cache_block in kv_cache:
+ cache_block.pin_current_chunk(current_num_tokens)
+
+ def _is_scene_cut(self, batch: Req, block_index: int) -> bool:
+ scene_cut_mask = batch.extra.get(CAUSAL_SCENE_CUT_MASK_KEY)
+ if not isinstance(scene_cut_mask, list) or block_index >= len(scene_cut_mask):
+ return False
+ return bool(scene_cut_mask[block_index])
+
+ def _new_causal_cache_pair(
+ self,
+ *,
+ batch_size: int,
+ max_text_len: int,
+ dtype: torch.dtype,
+ device: torch.device,
+ kv_cache_kwargs: dict[str, Any] | None = None,
+ ) -> tuple[list, list]:
+ prev_kv_cache = self.causal_kv_cache
+ prev_crossattn_cache = self.crossattn_cache
+ try:
+ return self._initialize_causal_caches(
+ batch_size=batch_size,
+ max_text_len=max_text_len,
+ dtype=dtype,
+ device=device,
+ kv_cache_kwargs=kv_cache_kwargs,
+ )
+ finally:
+ self.causal_kv_cache = prev_kv_cache
+ self.crossattn_cache = prev_crossattn_cache
+
+ def _reset_or_init_negative_caches(
+ self,
+ *,
+ batch: Req,
+ batch_size: int,
+ max_text_len: int,
+ dtype: torch.dtype,
+ device: torch.device,
+ kv_cache_kwargs: dict[str, Any] | None = None,
+ ) -> tuple[list, list]:
+ if (
+ self._cache_needs_reinit_for_batch(self.causal_kv_cache_neg, batch)
+ or self.crossattn_cache_neg is None
+ ):
+ (
+ self.causal_kv_cache_neg,
+ self.crossattn_cache_neg,
+ ) = self._new_causal_cache_pair(
+ batch_size=batch_size,
+ max_text_len=max_text_len,
+ dtype=dtype,
+ device=device,
+ kv_cache_kwargs=kv_cache_kwargs,
+ )
+ else:
+ self._reset_causal_caches(
+ kv_cache=self.causal_kv_cache_neg,
+ crossattn_cache=self.crossattn_cache_neg,
+ )
+ return self.causal_kv_cache_neg, self.crossattn_cache_neg
+
def _get_causal_kv_cache_size(
self,
*,
@@ -879,6 +1106,7 @@ class CausalDMDDenoisingStage(DenoisingStage):
device,
use_int_indices: bool = False,
sink_tokens: int = 0,
+ global_sink_tokens: int = 0,
attention_window_size: int | None = None,
allow_growth: bool = False,
) -> list[CausalSelfAttentionKVCache]:
@@ -915,6 +1143,7 @@ class CausalDMDDenoisingStage(DenoisingStage):
local_end_index_int=int_index,
cache_size=kv_cache_size,
sink_tokens=sink_tokens,
+ global_sink_tokens=global_sink_tokens,
attention_window_size=attention_window_size,
allow_growth=allow_growth,
)
@@ -1095,6 +1324,7 @@ class CausalDMDDenoisingStage(DenoisingStage):
*,
sequence_shard_enabled: bool = False,
kv_cache_size: int | None = None,
+ global_sink_tokens: int = 0,
) -> None:
"""
Initialize (but not fill) a Per-GPU KV cache aligned with the model assumptions.
@@ -1118,6 +1348,7 @@ class CausalDMDDenoisingStage(DenoisingStage):
sequence_shard_enabled=sequence_shard_enabled
),
sink_tokens=self._get_causal_sink_tokens(),
+ global_sink_tokens=global_sink_tokens,
attention_window_size=self._get_causal_attention_window_size(kv_cache_size),
)
diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/longlive2.py b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/longlive2.py
new file mode 100644
index 000000000..9f06b4b7e
--- /dev/null
+++ b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/longlive2.py
@@ -0,0 +1,902 @@
+# SPDX-License-Identifier: Apache-2.0
+from collections.abc import Callable
+from typing import Any
+
+import torch
+
+from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import Req
+from sglang.multimodal_gen.runtime.pipelines_core.stages.causal_denoising import (
+ CAUSAL_BLOCK_PROMPTS_KEY,
+ CAUSAL_SCENE_CUT_MASK_KEY,
+ CAUSAL_SHOT_INDICES_KEY,
+ CausalDMDDenoisingStage,
+ expand_causal_block_prompts,
+)
+from sglang.multimodal_gen.runtime.pipelines_core.stages.image_encoding import (
+ ImageVAEEncodingStage,
+)
+from sglang.multimodal_gen.runtime.pipelines_core.stages.latent_preparation import (
+ LatentPreparationSpec,
+ LatentPreparationStage,
+)
+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
+from sglang.multimodal_gen.runtime.utils.perf_logger import StageProfiler
+
+logger = init_logger(__name__)
+LONG_LIVE2_DEFAULT_SCENE_CUT_PREFIX = "The scene transitions. "
+
+
+def _latent_frame_count(batch: Req, server_args: ServerArgs) -> int:
+ num_frames = batch.num_frames
+ vae_config = server_args.pipeline_config.vae_config
+ if vae_config.use_temporal_scaling_frames:
+ temporal_scale_factor = vae_config.arch_config.temporal_compression_ratio
+ num_frames = (num_frames - 1) // temporal_scale_factor + 1
+ return int(num_frames)
+
+
+def _causal_block_count(batch: Req, server_args: ServerArgs) -> int:
+ latent_frames = _latent_frame_count(batch, server_args)
+ block_size = server_args.pipeline_config.dit_config.arch_config.num_frames_per_block
+ if latent_frames % block_size != 0:
+ raise ValueError(
+ "LongLive2 latent frames must be divisible by num_frames_per_block "
+ f"({block_size}), got {latent_frames}"
+ )
+ return latent_frames // block_size
+
+
+def expand_longlive2_shot_prompts(
+ shot_prompts: list[str],
+ *,
+ num_blocks: int,
+ shot_durations: list[int] | None = None,
+ chunks_per_shot: int = 0,
+ scene_cut_prefix: str = LONG_LIVE2_DEFAULT_SCENE_CUT_PREFIX,
+) -> list[str]:
+ return expand_causal_block_prompts(
+ shot_prompts,
+ num_blocks=num_blocks,
+ shot_durations=shot_durations,
+ chunks_per_shot=chunks_per_shot,
+ scene_cut_prefix=scene_cut_prefix,
+ )[0]
+
+
+class LongLive2TextEncodingStage(TextEncodingStage):
+ def build_dedup_fingerprint(self, batch: Req, server_args: ServerArgs):
+ base = super().build_dedup_fingerprint(batch, server_args)
+ return (
+ base,
+ self.freeze_for_dedup(getattr(batch, "shot_prompts", None)),
+ self.freeze_for_dedup(getattr(batch, "shot_durations", None)),
+ int(getattr(batch, "chunks_per_shot", 0) or 0),
+ getattr(batch, "scene_cut_prefix", None),
+ )
+
+ def _block_prompts(self, batch: Req, server_args: ServerArgs) -> list[str] | None:
+ shot_prompts = getattr(batch, "shot_prompts", None)
+ if shot_prompts is None:
+ return None
+ if isinstance(batch.prompt, list):
+ raise ValueError("LongLive2 shot_prompts supports one video per request")
+
+ block_prompts, scene_cut_mask, shot_indices = expand_causal_block_prompts(
+ shot_prompts,
+ num_blocks=_causal_block_count(batch, server_args),
+ shot_durations=getattr(batch, "shot_durations", None),
+ chunks_per_shot=int(getattr(batch, "chunks_per_shot", 0) or 0),
+ scene_cut_prefix=(
+ LONG_LIVE2_DEFAULT_SCENE_CUT_PREFIX
+ if getattr(batch, "scene_cut_prefix", None) is None
+ else getattr(batch, "scene_cut_prefix")
+ ),
+ )
+ batch.extra[CAUSAL_BLOCK_PROMPTS_KEY] = block_prompts
+ batch.extra[CAUSAL_SCENE_CUT_MASK_KEY] = scene_cut_mask
+ batch.extra[CAUSAL_SHOT_INDICES_KEY] = shot_indices
+ return block_prompts
+
+ @torch.no_grad()
+ def forward(self, batch: Req, server_args: ServerArgs) -> Req:
+ block_prompts = self._block_prompts(batch, server_args)
+ if block_prompts is None:
+ return super().forward(batch, server_args)
+
+ original_prompt = batch.prompt
+ batch.prompt = block_prompts
+ try:
+ return super().forward(batch, server_args)
+ finally:
+ batch.prompt = original_prompt
+
+
+class LongLive2ImageVAEEncodingStage(ImageVAEEncodingStage):
+ def preprocess(self, image):
+ image = super().preprocess(image)
+ if image.ndim == 5:
+ image = image.squeeze(2)
+ return image
+
+
+class LongLive2LatentPreparationStage(LatentPreparationStage):
+ def get_latent_preparation_spec(
+ self,
+ batch: Req,
+ server_args: ServerArgs,
+ batch_size: int,
+ num_frames: int,
+ device: torch.device | str,
+ ) -> LatentPreparationSpec:
+ b, c, t, h, w = server_args.pipeline_config.prepare_latent_shape(
+ batch, batch_size, num_frames
+ )
+ return LatentPreparationSpec(
+ shape=(b, t, c, h, w),
+ dtype=self._get_latent_dtype(batch, server_args),
+ device=device,
+ prepare_latent_ids=False,
+ pack_latents=False,
+ )
+
+ def forward(self, batch: Req, server_args: ServerArgs) -> Req:
+ batch = super().forward(batch, server_args)
+ return self._normalize_latent_layout(batch, server_args)
+
+ def _prepare_grouped_latents(
+ self,
+ batches: list[Req],
+ server_args: ServerArgs,
+ ) -> Req:
+ batch = super()._prepare_grouped_latents(batches, server_args)
+ return self._normalize_latent_layout(batch, server_args)
+
+ @staticmethod
+ def _expected_latent_channels(batch: Req, server_args: ServerArgs) -> int:
+ shape = server_args.pipeline_config.prepare_latent_shape(
+ batch,
+ batch.batch_size,
+ batch.latents.shape[1],
+ )
+ return int(shape[1])
+
+ def _normalize_latent_layout(self, batch: Req, server_args: ServerArgs) -> Req:
+ latents = batch.latents
+ if latents is None or latents.ndim != 5:
+ return batch
+ expected_channels = self._expected_latent_channels(batch, server_args)
+ if (
+ latents.shape[1] != expected_channels
+ and latents.shape[2] == expected_channels
+ ):
+ latents = latents.permute(0, 2, 1, 3, 4).contiguous()
+ batch.latents = latents
+ batch.raw_latent_shape = latents.shape
+ return batch
+
+
+class LongLive2CausalDenoisingStage(CausalDMDDenoisingStage):
+ def __init__(self, transformer, scheduler) -> None:
+ super().__init__(transformer, scheduler)
+ self._rope_temporal_offset = 0.0
+ self._i2v_image_latent: torch.Tensor | None = None
+
+ def _get_causal_dmd_latents(self, batch: Req) -> torch.Tensor:
+ latents = super()._get_causal_dmd_latents(batch)
+ if torch.is_inference(latents):
+ latents = latents.clone()
+ batch.latents = latents
+ return latents
+
+ @torch.no_grad()
+ def forward(self, batch: Req, server_args: ServerArgs) -> Req:
+ return self._forward_one_shot_common(
+ batch, server_args, use_cfg=self._use_cfg(batch)
+ )
+
+ @staticmethod
+ def _i2v_clamp_active(batch: Req) -> bool:
+ image_latent = getattr(batch, "image_latent", None)
+ return image_latent is not None and image_latent.shape[2] == 1
+
+ def _prepare_i2v_clamp(self, current_latents, start_frame):
+ clamp_latent = self._i2v_image_latent if start_frame == 0 else None
+ if clamp_latent is None:
+ return None, 0
+ clamp_latent = clamp_latent.to(
+ device=current_latents.device, dtype=current_latents.dtype
+ )
+ return clamp_latent, clamp_latent.shape[2]
+
+ @staticmethod
+ def _use_cfg(batch: Req) -> bool:
+ return bool(getattr(batch, "do_classifier_free_guidance", False))
+
+ @staticmethod
+ def _guidance_scale(batch: Req) -> float:
+ return float(getattr(batch, "guidance_scale", 1.0))
+
+ @staticmethod
+ def _denoise_step_profiler(batch: Req, start_frame: int, step_index: int):
+ return StageProfiler(
+ f"denoising_step_{start_frame}_{step_index}",
+ logger=logger,
+ metrics=batch.metrics,
+ perf_dump_path_provided=batch.perf_dump_path is not None,
+ record_as_step=True,
+ )
+
+ @staticmethod
+ def _get_negative_prompt_embeds(batch: Req):
+ negative_prompt_embeds = getattr(batch, "negative_prompt_embeds", None)
+ if negative_prompt_embeds is None or (
+ isinstance(negative_prompt_embeds, list)
+ and len(negative_prompt_embeds) == 0
+ ):
+ raise ValueError(
+ "LongLive2 classifier-free guidance requires negative_prompt_embeds"
+ )
+ return negative_prompt_embeds
+
+ def _prepare_causal_dmd_neg_cond_kwargs(
+ self,
+ batch: Req,
+ server_args: ServerArgs,
+ target_dtype: torch.dtype,
+ ) -> dict[str, Any]:
+ return self.prepare_extra_func_kwargs(
+ self.transformer.forward,
+ {
+ "encoder_attention_mask": batch.negative_attention_mask,
+ },
+ )
+
+ def _multi_shot_sink_enabled(self, batch: Req) -> bool:
+ return (
+ self._block_prompt_count(batch) is not None
+ and bool(getattr(batch, "multi_shot_sink", True))
+ and self.sink_size > 0
+ )
+
+ def _causal_kv_cache_global_sink_tokens_for_batch(self, batch: Req) -> int:
+ if not self._multi_shot_sink_enabled(batch):
+ return 0
+ return self._get_causal_sink_tokens()
+
+ def _is_scene_cut(self, batch: Req, block_index: int) -> bool:
+ if not self._multi_shot_sink_enabled(batch):
+ return False
+ return super()._is_scene_cut(batch, block_index)
+
+ def _set_rope_temporal_offset(self, batch: Req, shot_index: int) -> None:
+ offset = float(getattr(batch, "multi_shot_rope_offset", 8.0) or 0.0)
+ self._rope_temporal_offset = shot_index * offset
+
+ def _forward_one_shot_common(
+ self, batch: Req, server_args: ServerArgs, *, use_cfg: bool
+ ) -> Req:
+ ctx = self._prepare_causal_dmd_forward_context(batch, server_args)
+ target_dtype = ctx.target_dtype
+ autocast_enabled = ctx.autocast_enabled
+ scheduler = ctx.scheduler
+ device = ctx.device
+ timesteps = ctx.timesteps
+ image_kwargs = ctx.image_kwargs
+ pos_cond_kwargs = ctx.pos_cond_kwargs
+ latents = ctx.latents
+ prompt_embeds = ctx.prompt_embeds
+ t, h, w = ctx.num_frames, ctx.height, ctx.width
+
+ negative_prompt_embeds = None
+ neg_cond_kwargs = None
+ if use_cfg:
+ neg_cond_kwargs = self._prepare_causal_dmd_neg_cond_kwargs(
+ batch, server_args, target_dtype
+ )
+ negative_prompt_embeds = self._get_negative_prompt_embeds(batch)
+
+ independent_first_frame = self.transformer.independent_first_frame
+ max_text_len = self._get_max_text_len(server_args)
+ kv_cache_kwargs = self._causal_kv_cache_kwargs_for_batch(batch)
+ self._rope_temporal_offset = 0.0
+
+ if self._cache_needs_reinit_for_batch(self.causal_kv_cache, batch):
+ self._initialize_causal_caches(
+ batch_size=latents.shape[0],
+ max_text_len=max_text_len,
+ dtype=target_dtype,
+ device=latents.device,
+ kv_cache_kwargs=kv_cache_kwargs,
+ )
+ else:
+ assert self.crossattn_cache is not None
+ self._reset_causal_caches(
+ kv_cache=self.causal_kv_cache,
+ crossattn_cache=self.crossattn_cache,
+ )
+
+ kv_cache_neg = None
+ crossattn_cache_neg = None
+ if use_cfg:
+ kv_cache_neg, crossattn_cache_neg = self._reset_or_init_negative_caches(
+ batch=batch,
+ batch_size=latents.shape[0],
+ max_text_len=max_text_len,
+ dtype=target_dtype,
+ device=latents.device,
+ kv_cache_kwargs=kv_cache_kwargs,
+ )
+
+ current_start_frame = 0
+ clamp_i2v = self._i2v_clamp_active(batch)
+ self._i2v_image_latent = batch.image_latent if clamp_i2v else None
+ if getattr(batch, "image_latent", None) is not None and not clamp_i2v:
+ image_latent = batch.image_latent
+ assert image_latent is not None
+ input_frames = image_latent.shape[2]
+ warmup_prompt_embeds = self._select_block_prompt_embeds(
+ batch, prompt_embeds, 0
+ )
+ warmup_pos_cond_kwargs = self._select_block_cond_kwargs(
+ batch, pos_cond_kwargs, 0
+ )
+ warmup_neg_prompt_embeds = (
+ self._select_block_prompt_embeds(batch, negative_prompt_embeds, 0)
+ if use_cfg
+ else None
+ )
+ warmup_neg_cond_kwargs = (
+ self._select_block_cond_kwargs(batch, neg_cond_kwargs, 0)
+ if use_cfg
+ else None
+ )
+
+ def warm_up(context_input, start_frame):
+ self._warm_up_causal_context_cache(
+ batch,
+ server_args,
+ context_input=context_input,
+ prompt_embeds=warmup_prompt_embeds,
+ kv_cache=self.causal_kv_cache,
+ crossattn_cache=self.crossattn_cache,
+ current_start_frame=start_frame,
+ image_kwargs=image_kwargs,
+ pos_cond_kwargs=warmup_pos_cond_kwargs,
+ target_dtype=target_dtype,
+ autocast_enabled=autocast_enabled,
+ )
+ if use_cfg:
+ self._warm_up_causal_context_cache(
+ batch,
+ server_args,
+ context_input=context_input,
+ prompt_embeds=warmup_neg_prompt_embeds,
+ kv_cache=kv_cache_neg,
+ crossattn_cache=crossattn_cache_neg,
+ current_start_frame=start_frame,
+ image_kwargs=image_kwargs,
+ pos_cond_kwargs=warmup_neg_cond_kwargs,
+ target_dtype=target_dtype,
+ autocast_enabled=autocast_enabled,
+ )
+
+ if independent_first_frame and input_frames >= 1:
+ warm_up(image_latent[:, :, :1, :, :], current_start_frame)
+ current_start_frame += 1
+ remaining_frames = input_frames - 1
+ else:
+ remaining_frames = input_frames
+
+ while remaining_frames > 0:
+ block = min(self.num_frames_per_block, remaining_frames)
+ warm_up(
+ image_latent[
+ :, :, current_start_frame : current_start_frame + block, :, :
+ ],
+ current_start_frame,
+ )
+ current_start_frame += block
+ remaining_frames -= block
+
+ pos_start_base = current_start_frame
+
+ if not independent_first_frame or (
+ independent_first_frame and batch.image_latent is not None
+ ):
+ if t % self.num_frames_per_block != 0:
+ raise ValueError(
+ "num_frames must be divisible by num_frames_per_block for causal DMD denoising"
+ )
+ num_blocks = t // self.num_frames_per_block
+ block_sizes = [self.num_frames_per_block] * num_blocks
+ else:
+ if (t - 1) % self.num_frames_per_block != 0:
+ raise ValueError(
+ "(num_frames - 1) must be divisible by num_frame_per_block when independent_first_frame=True"
+ )
+ num_blocks = (t - 1) // self.num_frames_per_block
+ block_sizes = [1] + [self.num_frames_per_block] * num_blocks
+
+ start_index = 0
+ self._validate_block_prompt_count(batch, block_sizes)
+
+ def prepare_context_input(current_latents):
+ return current_latents
+
+ with self.progress_bar(total=len(block_sizes) * len(timesteps)) as progress_bar:
+ for block_index, current_num_frames in enumerate(block_sizes):
+ self._set_rope_temporal_offset(
+ batch, self._shot_index(batch, block_index)
+ )
+ is_scene_cut = self._is_scene_cut(batch, block_index)
+
+ current_latents = latents[
+ :, :, start_index : start_index + current_num_frames, :, :
+ ]
+ current_prompt_embeds = self._select_block_prompt_embeds(
+ batch, prompt_embeds, block_index
+ )
+ current_pos_cond_kwargs = self._select_block_cond_kwargs(
+ batch, pos_cond_kwargs, block_index
+ )
+
+ caches = [self.crossattn_cache]
+ if use_cfg:
+ caches.append(crossattn_cache_neg)
+ self._reset_crossattn_cache_for_block(batch, *caches)
+
+ def prepare_model_input(current_latents):
+ latent_model_input = current_latents
+ if (
+ batch.image_latent is not None
+ and independent_first_frame
+ and start_index == 0
+ ):
+ latent_model_input = torch.cat(
+ [latent_model_input, batch.image_latent], dim=2
+ )
+ return latent_model_input
+
+ current_start_tokens = (
+ pos_start_base + start_index
+ ) * self.num_token_per_frame
+ block_kwargs = dict(
+ chunk_latents=current_latents,
+ scheduler=scheduler,
+ timesteps=timesteps,
+ prompt_embeds=current_prompt_embeds,
+ kv_cache=self.causal_kv_cache,
+ crossattn_cache=self.crossattn_cache,
+ current_start_tokens=current_start_tokens,
+ start_frame=start_index,
+ image_kwargs=image_kwargs,
+ pos_cond_kwargs=current_pos_cond_kwargs,
+ target_dtype=target_dtype,
+ autocast_enabled=autocast_enabled,
+ device=device,
+ attn_raw_latent_shape=(current_num_frames, h, w),
+ prepare_model_input=prepare_model_input,
+ prepare_context_input=prepare_context_input,
+ progress_bar=progress_bar,
+ )
+ if use_cfg:
+ current_latents = self._denoise_and_update_causal_block_cfg(
+ batch,
+ server_args,
+ negative_prompt_embeds=self._select_block_prompt_embeds(
+ batch, negative_prompt_embeds, block_index
+ ),
+ kv_cache_neg=kv_cache_neg,
+ crossattn_cache_neg=crossattn_cache_neg,
+ neg_cond_kwargs=self._select_block_cond_kwargs(
+ batch, neg_cond_kwargs, block_index
+ ),
+ **block_kwargs,
+ )
+ else:
+ current_latents = self._denoise_and_update_causal_block(
+ batch, server_args, **block_kwargs
+ )
+
+ if is_scene_cut:
+ self._pin_current_chunk(self.causal_kv_cache, current_num_frames)
+ if use_cfg:
+ self._pin_current_chunk(kv_cache_neg, current_num_frames)
+
+ latents[:, :, start_index : start_index + current_num_frames, :, :] = (
+ current_latents
+ )
+ start_index += current_num_frames
+
+ self._rope_temporal_offset = 0.0
+ batch.latents = latents
+ return batch
+
+ def _forward_causal_transformer(
+ self,
+ batch: Req,
+ *,
+ latent_model_input: torch.Tensor,
+ prompt_embeds,
+ timestep: torch.Tensor,
+ kv_cache,
+ crossattn_cache,
+ current_start_tokens: int,
+ start_frame: int,
+ image_kwargs: dict,
+ pos_cond_kwargs: dict,
+ current_timestep: int,
+ attn_metadata,
+ target_dtype: torch.dtype,
+ autocast_enabled: bool,
+ ) -> torch.Tensor:
+ self._manage_dit_use_site(self.transformer, "transformer", batch)
+ rope_start_frame = start_frame
+ if self._rope_temporal_offset != 0.0:
+ rope_start_frame = start_frame + self._rope_temporal_offset
+ return super()._forward_causal_transformer(
+ batch,
+ latent_model_input=latent_model_input,
+ prompt_embeds=prompt_embeds,
+ timestep=timestep,
+ kv_cache=kv_cache,
+ crossattn_cache=crossattn_cache,
+ current_start_tokens=current_start_tokens,
+ start_frame=rope_start_frame,
+ image_kwargs=image_kwargs,
+ pos_cond_kwargs=pos_cond_kwargs,
+ current_timestep=current_timestep,
+ attn_metadata=attn_metadata,
+ target_dtype=target_dtype,
+ autocast_enabled=autocast_enabled,
+ )
+
+ def _prepare_causal_dmd_timesteps(
+ self,
+ batch: Req,
+ server_args: ServerArgs,
+ scheduler,
+ device: torch.device,
+ ) -> torch.Tensor:
+ scheduler.set_timesteps(
+ batch.num_inference_steps,
+ device=device,
+ shift=server_args.pipeline_config.flow_shift,
+ )
+ return scheduler.timesteps.to(device)
+
+ def _denoise_causal_dmd_chunk(
+ self,
+ batch: Req,
+ server_args: ServerArgs,
+ *,
+ chunk_latents: torch.Tensor,
+ scheduler,
+ timesteps: torch.Tensor,
+ prompt_embeds,
+ kv_cache,
+ crossattn_cache,
+ current_start_tokens: int,
+ start_frame: int,
+ image_kwargs: dict,
+ pos_cond_kwargs: dict,
+ target_dtype: torch.dtype,
+ autocast_enabled: bool,
+ device: torch.device,
+ attn_raw_latent_shape: tuple[int, int, int],
+ prepare_model_input: Callable[[torch.Tensor], torch.Tensor],
+ progress_bar=None,
+ ) -> tuple[torch.Tensor, Any | None]:
+ scheduler.set_timesteps(
+ len(timesteps),
+ device=device,
+ shift=server_args.pipeline_config.flow_shift,
+ )
+ timesteps = scheduler.timesteps.to(device)
+ current_latents = chunk_latents
+ attn_metadata = None
+ clamp_latent, context_frames = self._prepare_i2v_clamp(
+ current_latents, start_frame
+ )
+ if clamp_latent is not None:
+ current_latents = current_latents.clone()
+
+ for current_timestep, timestep in enumerate(timesteps):
+ with self._denoise_step_profiler(batch, start_frame, current_timestep):
+ if clamp_latent is not None:
+ current_latents[:, :, :context_frames] = clamp_latent
+ latent_model_input = prepare_model_input(current_latents).to(
+ target_dtype
+ )
+ attn_metadata = self._build_causal_attn_metadata(
+ batch,
+ server_args,
+ current_timestep=current_timestep,
+ raw_latent_shape=attn_raw_latent_shape,
+ device=device,
+ )
+ batch_size = latent_model_input.shape[0]
+ timestep_2d = (
+ timestep.reshape(1)
+ .to(device=latent_model_input.device, dtype=torch.float32)
+ .expand(batch_size, latent_model_input.shape[2])
+ )
+ if clamp_latent is not None:
+ timestep_2d = timestep_2d.clone()
+ timestep_2d[:, :context_frames] = 0
+ flow_pred = self._forward_causal_transformer(
+ batch,
+ latent_model_input=latent_model_input,
+ prompt_embeds=prompt_embeds,
+ timestep=timestep_2d,
+ kv_cache=kv_cache,
+ crossattn_cache=crossattn_cache,
+ current_start_tokens=current_start_tokens,
+ start_frame=start_frame,
+ image_kwargs=image_kwargs,
+ pos_cond_kwargs=pos_cond_kwargs,
+ current_timestep=current_timestep,
+ attn_metadata=attn_metadata,
+ target_dtype=target_dtype,
+ autocast_enabled=autocast_enabled,
+ )
+
+ next_latents = scheduler.step(
+ flow_pred,
+ timestep,
+ current_latents,
+ return_dict=False,
+ )[0]
+
+ current_latents = next_latents
+ if clamp_latent is not None:
+ current_latents[:, :, :context_frames] = clamp_latent
+
+ if progress_bar is not None:
+ progress_bar.update()
+
+ return current_latents, attn_metadata
+
+ def _denoise_causal_dmd_chunk_cfg(
+ self,
+ batch: Req,
+ server_args: ServerArgs,
+ *,
+ chunk_latents: torch.Tensor,
+ scheduler,
+ timesteps: torch.Tensor,
+ prompt_embeds,
+ negative_prompt_embeds,
+ kv_cache,
+ crossattn_cache,
+ kv_cache_neg,
+ crossattn_cache_neg,
+ current_start_tokens: int,
+ start_frame: int,
+ image_kwargs: dict,
+ pos_cond_kwargs: dict,
+ neg_cond_kwargs: dict,
+ target_dtype: torch.dtype,
+ autocast_enabled: bool,
+ device: torch.device,
+ attn_raw_latent_shape: tuple[int, int, int],
+ prepare_model_input: Callable[[torch.Tensor], torch.Tensor],
+ progress_bar=None,
+ ) -> tuple[torch.Tensor, Any | None]:
+ scheduler.set_timesteps(
+ len(timesteps),
+ device=device,
+ shift=server_args.pipeline_config.flow_shift,
+ )
+ timesteps = scheduler.timesteps.to(device)
+ current_latents = chunk_latents
+ attn_metadata = None
+ guidance_scale = self._guidance_scale(batch)
+ clamp_latent, context_frames = self._prepare_i2v_clamp(
+ current_latents, start_frame
+ )
+ if clamp_latent is not None:
+ current_latents = current_latents.clone()
+
+ for current_timestep, timestep in enumerate(timesteps):
+ with self._denoise_step_profiler(batch, start_frame, current_timestep):
+ if clamp_latent is not None:
+ current_latents[:, :, :context_frames] = clamp_latent
+ latent_model_input = prepare_model_input(current_latents).to(
+ target_dtype
+ )
+ attn_metadata = self._build_causal_attn_metadata(
+ batch,
+ server_args,
+ current_timestep=current_timestep,
+ raw_latent_shape=attn_raw_latent_shape,
+ device=device,
+ )
+ batch_size = latent_model_input.shape[0]
+ timestep_2d = (
+ timestep.reshape(1)
+ .to(device=latent_model_input.device, dtype=torch.float32)
+ .expand(batch_size, latent_model_input.shape[2])
+ )
+ if clamp_latent is not None:
+ timestep_2d = timestep_2d.clone()
+ timestep_2d[:, :context_frames] = 0
+ flow_pred_cond = self._forward_causal_transformer(
+ batch,
+ latent_model_input=latent_model_input,
+ prompt_embeds=prompt_embeds,
+ timestep=timestep_2d,
+ kv_cache=kv_cache,
+ crossattn_cache=crossattn_cache,
+ current_start_tokens=current_start_tokens,
+ start_frame=start_frame,
+ image_kwargs=image_kwargs,
+ pos_cond_kwargs=pos_cond_kwargs,
+ current_timestep=current_timestep,
+ attn_metadata=attn_metadata,
+ target_dtype=target_dtype,
+ autocast_enabled=autocast_enabled,
+ )
+ flow_pred_uncond = self._forward_causal_transformer(
+ batch,
+ latent_model_input=latent_model_input,
+ prompt_embeds=negative_prompt_embeds,
+ timestep=timestep_2d,
+ kv_cache=kv_cache_neg,
+ crossattn_cache=crossattn_cache_neg,
+ current_start_tokens=current_start_tokens,
+ start_frame=start_frame,
+ image_kwargs=image_kwargs,
+ pos_cond_kwargs=neg_cond_kwargs,
+ current_timestep=current_timestep,
+ attn_metadata=attn_metadata,
+ target_dtype=target_dtype,
+ autocast_enabled=autocast_enabled,
+ )
+ flow_pred = flow_pred_uncond + guidance_scale * (
+ flow_pred_cond - flow_pred_uncond
+ )
+
+ next_latents = scheduler.step(
+ flow_pred,
+ timestep,
+ current_latents,
+ return_dict=False,
+ )[0]
+ current_latents = next_latents
+ if clamp_latent is not None:
+ current_latents[:, :, :context_frames] = clamp_latent
+
+ if progress_bar is not None:
+ progress_bar.update()
+
+ return current_latents, attn_metadata
+
+ def _denoise_and_update_causal_block_cfg(
+ self,
+ batch: Req,
+ server_args: ServerArgs,
+ *,
+ chunk_latents: torch.Tensor,
+ scheduler,
+ timesteps: torch.Tensor,
+ prompt_embeds,
+ negative_prompt_embeds,
+ kv_cache,
+ crossattn_cache,
+ kv_cache_neg,
+ crossattn_cache_neg,
+ current_start_tokens: int,
+ start_frame: int,
+ image_kwargs: dict,
+ pos_cond_kwargs: dict,
+ neg_cond_kwargs: dict,
+ target_dtype: torch.dtype,
+ autocast_enabled: bool,
+ device: torch.device,
+ attn_raw_latent_shape: tuple[int, int, int],
+ prepare_model_input: Callable[[torch.Tensor], torch.Tensor],
+ prepare_context_input: Callable[[torch.Tensor], torch.Tensor],
+ progress_bar=None,
+ ) -> torch.Tensor:
+ current_latents, attn_metadata = self._denoise_causal_dmd_chunk_cfg(
+ batch,
+ server_args,
+ chunk_latents=chunk_latents,
+ scheduler=scheduler,
+ timesteps=timesteps,
+ prompt_embeds=prompt_embeds,
+ negative_prompt_embeds=negative_prompt_embeds,
+ kv_cache=kv_cache,
+ crossattn_cache=crossattn_cache,
+ kv_cache_neg=kv_cache_neg,
+ crossattn_cache_neg=crossattn_cache_neg,
+ current_start_tokens=current_start_tokens,
+ start_frame=start_frame,
+ image_kwargs=image_kwargs,
+ pos_cond_kwargs=pos_cond_kwargs,
+ neg_cond_kwargs=neg_cond_kwargs,
+ target_dtype=target_dtype,
+ autocast_enabled=autocast_enabled,
+ device=device,
+ attn_raw_latent_shape=attn_raw_latent_shape,
+ prepare_model_input=prepare_model_input,
+ progress_bar=progress_bar,
+ )
+ context_input = prepare_context_input(current_latents)
+ self._update_causal_context_cache(
+ batch,
+ server_args,
+ context_input=context_input,
+ prompt_embeds=prompt_embeds,
+ kv_cache=kv_cache,
+ crossattn_cache=crossattn_cache,
+ current_start_tokens=current_start_tokens,
+ start_frame=start_frame,
+ image_kwargs=image_kwargs,
+ pos_cond_kwargs=pos_cond_kwargs,
+ attn_metadata=attn_metadata,
+ target_dtype=target_dtype,
+ autocast_enabled=autocast_enabled,
+ )
+ self._update_causal_context_cache(
+ batch,
+ server_args,
+ context_input=context_input,
+ prompt_embeds=negative_prompt_embeds,
+ kv_cache=kv_cache_neg,
+ crossattn_cache=crossattn_cache_neg,
+ current_start_tokens=current_start_tokens,
+ start_frame=start_frame,
+ image_kwargs=image_kwargs,
+ pos_cond_kwargs=neg_cond_kwargs,
+ attn_metadata=attn_metadata,
+ target_dtype=target_dtype,
+ autocast_enabled=autocast_enabled,
+ )
+ return current_latents
+
+ def _update_causal_context_cache(
+ self,
+ batch: Req,
+ server_args: ServerArgs,
+ *,
+ context_input: torch.Tensor,
+ prompt_embeds,
+ kv_cache,
+ crossattn_cache,
+ current_start_tokens: int,
+ start_frame: int,
+ image_kwargs: dict,
+ pos_cond_kwargs: dict,
+ attn_metadata,
+ target_dtype: torch.dtype,
+ autocast_enabled: bool,
+ ) -> None:
+ context_noise = getattr(server_args.pipeline_config, "context_noise", 0)
+ timestep = torch.full(
+ (context_input.shape[0], context_input.shape[2]),
+ float(context_noise),
+ device=context_input.device,
+ dtype=torch.float32,
+ )
+ self._forward_causal_transformer(
+ batch,
+ latent_model_input=context_input.to(target_dtype),
+ prompt_embeds=prompt_embeds,
+ timestep=timestep,
+ kv_cache=kv_cache,
+ crossattn_cache=crossattn_cache,
+ current_start_tokens=current_start_tokens,
+ start_frame=start_frame,
+ image_kwargs=image_kwargs,
+ pos_cond_kwargs=pos_cond_kwargs,
+ current_timestep=0,
+ attn_metadata=attn_metadata,
+ target_dtype=target_dtype,
+ autocast_enabled=autocast_enabled,
+ )
diff --git a/python/sglang/multimodal_gen/runtime/utils/hf_diffusers_utils.py b/python/sglang/multimodal_gen/runtime/utils/hf_diffusers_utils.py
index eeae1e63f..749dbe2e7 100644
--- a/python/sglang/multimodal_gen/runtime/utils/hf_diffusers_utils.py
+++ b/python/sglang/multimodal_gen/runtime/utils/hf_diffusers_utils.py
@@ -57,6 +57,80 @@ from sglang.utils import is_in_ci
logger = init_logger(__name__)
+_NON_WEIGHT_DIFFUSERS_COMPONENT_HINTS = (
+ "tokenizer",
+ "scheduler",
+ "processor",
+ "feature_extractor",
+)
+_WEIGHT_FILE_PATTERNS = (
+ "*.safetensors",
+ "*.bin",
+ "*.pt",
+ "*.pth",
+ "*.ckpt",
+)
+
+
+def _is_diffusers_component_entry(value: Any) -> bool:
+ return (
+ isinstance(value, (list, tuple))
+ and len(value) == 2
+ and all(item is None or isinstance(item, str) for item in value)
+ )
+
+
+def _is_weight_bearing_diffusers_component(key: str, value: Any) -> bool:
+ if (
+ key.startswith("_")
+ or not _is_diffusers_component_entry(value)
+ or not any(item is not None for item in value)
+ ):
+ return False
+
+ key_lower = key.lower()
+ return not any(hint in key_lower for hint in _NON_WEIGHT_DIFFUSERS_COMPONENT_HINTS)
+
+
+def _get_declared_weight_component_dirs(model_path: str) -> list[str]:
+ model_index_path = os.path.join(model_path, "model_index.json")
+ if not os.path.exists(model_index_path):
+ return []
+
+ try:
+ with open(model_index_path) as f:
+ model_index = json.load(f)
+ except Exception as exc:
+ logger.warning(
+ "Failed to read model_index.json at %s: %s", model_index_path, exc
+ )
+ return []
+
+ return [
+ key
+ for key, value in model_index.items()
+ if _is_weight_bearing_diffusers_component(key, value)
+ ]
+
+
+def _has_local_weight_files(component_path: str) -> bool:
+ return any(
+ glob.glob(os.path.join(component_path, pattern))
+ for pattern in _WEIGHT_FILE_PATTERNS
+ )
+
+
+def _get_missing_declared_weight_components(model_path: str) -> list[str]:
+ missing_files = []
+ for component_dir in _get_declared_weight_component_dirs(model_path):
+ component_path = os.path.join(model_path, component_dir)
+ if not os.path.isdir(component_path):
+ missing_files.append(f"{component_dir}/")
+ elif not _has_local_weight_files(component_path):
+ missing_files.append(f"{component_dir}/")
+ return missing_files
+
+
def _check_index_files_for_missing_shards(
model_path: str,
) -> tuple[bool, list[str], list[str]]:
@@ -74,6 +148,15 @@ def _check_index_files_for_missing_shards(
"""
missing_files = []
checked_subdirs = []
+ checked_subdir_set = set()
+
+ def _record_checked_subdir(dir_path: str) -> None:
+ subdir = os.path.basename(dir_path)
+ if not subdir:
+ subdir = "."
+ if subdir not in checked_subdir_set:
+ checked_subdirs.append(subdir)
+ checked_subdir_set.add(subdir)
# Add common subdirectories for diffusers models
try:
@@ -85,6 +168,10 @@ def _check_index_files_for_missing_shards(
# Check the root directory and all subdirectories that might contain model weights
dirs_to_check = [model_path]
+ for component_dir in _get_declared_weight_component_dirs(model_path):
+ _record_checked_subdir(os.path.join(model_path, component_dir))
+ missing_files.extend(_get_missing_declared_weight_components(model_path))
+
for subdir in subdirs:
subdir_path = os.path.join(model_path, subdir)
if os.path.isdir(subdir_path):
@@ -95,7 +182,7 @@ def _check_index_files_for_missing_shards(
index_files = glob.glob(os.path.join(dir_path, "*.safetensors.index.json"))
for index_file in index_files:
- checked_subdirs.append(os.path.basename(dir_path))
+ _record_checked_subdir(dir_path)
try:
with open(index_file) as f:
index_data = json.load(f)
@@ -227,12 +314,13 @@ def _verify_diffusers_model_complete(path: str) -> bool:
component_keys = [
key
for key, value in model_index.items()
- if isinstance(value, (list, tuple))
- and len(value) == 2
- and all(isinstance(item, str) for item in value)
+ if _is_diffusers_component_entry(value)
+ and any(item is not None for item in value)
]
if component_keys:
- return all(os.path.exists(os.path.join(path, key)) for key in component_keys)
+ return all(
+ os.path.exists(os.path.join(path, key)) for key in component_keys
+ ) and not _get_missing_declared_weight_components(path)
return os.path.exists(os.path.join(path, "transformer")) and os.path.exists(
os.path.join(path, "vae")
diff --git a/python/sglang/multimodal_gen/test/server/gpu_cases.py b/python/sglang/multimodal_gen/test/server/gpu_cases.py
index 489223f9a..c02ad7472 100644
--- a/python/sglang/multimodal_gen/test/server/gpu_cases.py
+++ b/python/sglang/multimodal_gen/test/server/gpu_cases.py
@@ -22,7 +22,8 @@ from sglang.multimodal_gen.test.server.testcase_configs import (
DiffusionTestCase,
IDEOGRAM4_CI_sampling_params,
JOY_ECHO_T2V_CI_sampling_params,
- LINGBOT_WORLD_REALTIME_sampling_params,
+ LONGLIVE2_I2V_CI_sampling_params,
+ LONGLIVE2_T2V_CI_sampling_params,
MODELOPT_QWEN_IMAGE_2512_NVFP4_CI_sampling_params,
MODELOPT_T2I_CI_sampling_params,
MODELOPT_T2V_CI_sampling_params,
@@ -31,6 +32,7 @@ from sglang.multimodal_gen.test.server.testcase_configs import (
MULTI_IMAGE_TI2I_sampling_params,
MULTI_IMAGE_TI2I_UPLOAD_sampling_params,
PI05_ACTION_CI_sampling_params,
+ REALTIME_MODEL_sampling_params,
SANA_WM_TI2V_CI_sampling_params,
T2I_sampling_params,
T2V_sampling_params,
@@ -259,6 +261,15 @@ ONE_GPU_CASES: list[DiffusionTestCase] = [
run_consistency_check=True,
run_component_accuracy_check=False,
),
+ DiffusionTestCase(
+ "longlive2_t2v",
+ DiffusionServerArgs(
+ model_path="Rabinovich/LongLive-2.0-5B-Diffusers",
+ modality="video",
+ ),
+ LONGLIVE2_T2V_CI_sampling_params,
+ run_component_accuracy_check=False,
+ ),
# TeaCache acceleration test for Wan video model
DiffusionTestCase(
"wan2_1_t2v_1.3b_teacache_enabled",
@@ -380,6 +391,17 @@ ONE_GPU_CASES: list[DiffusionTestCase] = [
run_models_api_check=False,
run_t2v_input_reference_check=False,
),
+ DiffusionTestCase(
+ "longlive2_i2v",
+ DiffusionServerArgs(
+ model_path="Rabinovich/LongLive-2.0-5B-Diffusers",
+ modality="video",
+ ),
+ LONGLIVE2_I2V_CI_sampling_params,
+ run_component_accuracy_check=False,
+ run_models_api_check=False,
+ run_t2v_input_reference_check=False,
+ ),
# flaky
# === Helios T2V ===
# DiffusionTestCase(
@@ -439,7 +461,7 @@ ONE_GPU_CASES: list[DiffusionTestCase] = [
],
text_encoder_cpu_offload=True,
),
- LINGBOT_WORLD_REALTIME_sampling_params,
+ REALTIME_MODEL_sampling_params,
run_component_accuracy_check=False,
run_models_api_check=False,
run_t2v_input_reference_check=False,
diff --git a/python/sglang/multimodal_gen/test/server/perf_baselines/h100.json b/python/sglang/multimodal_gen/test/server/perf_baselines/h100.json
index 21508bcf9..4bb6422b4 100644
--- a/python/sglang/multimodal_gen/test/server/perf_baselines/h100.json
+++ b/python/sglang/multimodal_gen/test/server/perf_baselines/h100.json
@@ -2602,6 +2602,38 @@
"expected_median_denoise_ms": 242.79,
"estimated_full_test_time_s": 170.0
},
+ "longlive2_t2v": {
+ "stages_ms": {
+ "InputValidationStage": 0.05,
+ "LongLive2TextEncodingStage": 328.32,
+ "LongLive2ImageVAEEncodingStage": 0.0,
+ "LongLive2LatentPreparationStage": 0.17,
+ "LongLive2CausalDenoisingStage": 4879.98,
+ "DecodingStage": 1397.62,
+ "per_frame_generation": null
+ },
+ "denoise_step_ms": {},
+ "expected_e2e_ms": 6610.71,
+ "expected_avg_denoise_ms": 650.0,
+ "expected_median_denoise_ms": 650.0,
+ "estimated_full_test_time_s": 153.1
+ },
+ "longlive2_i2v": {
+ "stages_ms": {
+ "InputValidationStage": 23.02,
+ "LongLive2TextEncodingStage": 327.98,
+ "LongLive2ImageVAEEncodingStage": 1048.81,
+ "LongLive2LatentPreparationStage": 0.12,
+ "LongLive2CausalDenoisingStage": 4975.28,
+ "DecodingStage": 3051.96,
+ "per_frame_generation": null
+ },
+ "denoise_step_ms": {},
+ "expected_e2e_ms": 9431.65,
+ "expected_avg_denoise_ms": 800.0,
+ "expected_median_denoise_ms": 800.0,
+ "estimated_full_test_time_s": 149.4
+ },
"lingbot_world_realtime_plastic_beach": {
"stages_ms": {},
"denoise_step_ms": {},
diff --git a/python/sglang/multimodal_gen/test/server/testcase_configs.py b/python/sglang/multimodal_gen/test/server/testcase_configs.py
index bc5330431..ee6b39981 100644
--- a/python/sglang/multimodal_gen/test/server/testcase_configs.py
+++ b/python/sglang/multimodal_gen/test/server/testcase_configs.py
@@ -315,7 +315,13 @@ class DiffusionTestCase:
)
-LINGBOT_WORLD_REALTIME_sampling_params = DiffusionSamplingParams(
+_REALTIME_MODEL_COMMON_EXTRAS = {
+ "seed": 42,
+ "num_inference_steps": 4,
+ "guidance_scale": 1.0,
+}
+
+REALTIME_MODEL_sampling_params = DiffusionSamplingParams(
prompt=(
"A slow aerial orbit around a pastel floating island hotel in the open "
"ocean, hazy sunlight, turquoise water, toy-like architectural detail, "
@@ -336,9 +342,7 @@ LINGBOT_WORLD_REALTIME_sampling_params = DiffusionSamplingParams(
},
realtime_perf_ignore_initial_chunks=2,
extras={
- "seed": 42,
- "num_inference_steps": 4,
- "guidance_scale": 1.0,
+ **_REALTIME_MODEL_COMMON_EXTRAS,
"realtime_causal_sink_size": 9,
"realtime_causal_kv_cache_num_frames": 18,
"condition_inputs": {
@@ -607,6 +611,28 @@ SANA_WM_TI2V_CI_sampling_params = DiffusionSamplingParams(
extras={"num_inference_steps": 12, "seed": 0, "guidance_scale": 4.5},
)
+LONGLIVE2_T2V_CI_sampling_params = replace(
+ REALTIME_MODEL_sampling_params,
+ image_path=None,
+ num_frames=61,
+ realtime_num_chunks=None,
+ realtime_events=[],
+ realtime_perf_thresholds={},
+ realtime_perf_ignore_initial_chunks=0,
+ extras=dict(_REALTIME_MODEL_COMMON_EXTRAS),
+)
+
+LONGLIVE2_I2V_CI_sampling_params = replace(
+ REALTIME_MODEL_sampling_params,
+ output_size="960x928",
+ num_frames=61,
+ realtime_num_chunks=None,
+ realtime_events=[],
+ realtime_perf_thresholds={},
+ realtime_perf_ignore_initial_chunks=0,
+ extras=dict(_REALTIME_MODEL_COMMON_EXTRAS),
+)
+
TURBOWAN_I2V_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",
diff --git a/python/sglang/multimodal_gen/test/test_utils.py b/python/sglang/multimodal_gen/test/test_utils.py
index 063a96a6c..844220f1b 100644
--- a/python/sglang/multimodal_gen/test/test_utils.py
+++ b/python/sglang/multimodal_gen/test/test_utils.py
@@ -34,7 +34,7 @@ if TYPE_CHECKING:
logger = init_logger(__name__)
-SGL_TEST_FILES_CI_DATA_REVISION = "9a64abec5a7517a9f2b04ac1b4eab4173adb2d38"
+SGL_TEST_FILES_CI_DATA_REVISION = "d51ca9623e0bb27087da243a44c942fdda5aafe5"
if current_platform.is_npu():
SGL_TEST_FILES_CI_DATA_REVISION = "6b62f4b6825c76a25fd2ba28248df68f2b400e65"
diff --git a/python/sglang/multimodal_gen/test/unit/realtime/test_realtime_consistency_harness.py b/python/sglang/multimodal_gen/test/unit/realtime/test_realtime_consistency_harness.py
index 86bc5a4ee..44d6a8604 100644
--- a/python/sglang/multimodal_gen/test/unit/realtime/test_realtime_consistency_harness.py
+++ b/python/sglang/multimodal_gen/test/unit/realtime/test_realtime_consistency_harness.py
@@ -31,7 +31,9 @@ from sglang.multimodal_gen.test.server.realtime_consistency import (
from sglang.multimodal_gen.test.server.test_server_utils import get_generate_fn
from sglang.multimodal_gen.test.server.testcase_configs import (
DiffusionSamplingParams,
- LINGBOT_WORLD_REALTIME_sampling_params,
+ LONGLIVE2_I2V_CI_sampling_params,
+ LONGLIVE2_T2V_CI_sampling_params,
+ REALTIME_MODEL_sampling_params,
)
# Request construction
@@ -493,8 +495,8 @@ def test_realtime_sampling_params_route_to_realtime_video_generator():
assert generate_fn.__name__ == "generate_realtime_video"
-def test_lingbot_realtime_plastic_beach_params_are_lossless_gt_ready():
- params = LINGBOT_WORLD_REALTIME_sampling_params
+def test_realtime_model_params_are_lossless_gt_ready():
+ params = REALTIME_MODEL_sampling_params
assert "floating island hotel" in params.prompt
assert "825646291038" in str(params.image_path)
@@ -521,6 +523,33 @@ def test_lingbot_realtime_plastic_beach_params_are_lossless_gt_ready():
]
+def test_longlive2_cases_share_realtime_model_sampling_profile():
+ for params in (
+ LONGLIVE2_T2V_CI_sampling_params,
+ LONGLIVE2_I2V_CI_sampling_params,
+ ):
+ assert params.prompt == REALTIME_MODEL_sampling_params.prompt
+ assert params.fps == REALTIME_MODEL_sampling_params.fps
+ assert params.extras == {
+ "seed": 42,
+ "num_inference_steps": 4,
+ "guidance_scale": 1.0,
+ }
+ assert params.realtime_num_chunks is None
+ assert params.realtime_perf_thresholds == {}
+
+ assert LONGLIVE2_T2V_CI_sampling_params.image_path is None
+ assert (
+ LONGLIVE2_T2V_CI_sampling_params.output_size
+ == REALTIME_MODEL_sampling_params.output_size
+ )
+ assert (
+ LONGLIVE2_I2V_CI_sampling_params.image_path
+ == REALTIME_MODEL_sampling_params.image_path
+ )
+ assert LONGLIVE2_I2V_CI_sampling_params.output_size == "960x928"
+
+
def test_lingbot_realtime_case_is_registered_by_default():
from sglang.multimodal_gen.test.server.gpu_cases import ONE_GPU_CASES
diff --git a/python/sglang/multimodal_gen/test/unit/test_hf_diffusers_utils.py b/python/sglang/multimodal_gen/test/unit/test_hf_diffusers_utils.py
new file mode 100644
index 000000000..614a5134b
--- /dev/null
+++ b/python/sglang/multimodal_gen/test/unit/test_hf_diffusers_utils.py
@@ -0,0 +1,72 @@
+import json
+
+from sglang.multimodal_gen.runtime.utils.hf_diffusers_utils import (
+ _check_index_files_for_missing_shards,
+ _verify_diffusers_model_complete,
+)
+
+
+def _write_model_index(root):
+ (root / "model_index.json").write_text(
+ json.dumps(
+ {
+ "_class_name": "LongLive2Pipeline",
+ "_diffusers_version": "0.34.0",
+ "scheduler": ["diffusers", "FlowMatchEulerDiscreteScheduler"],
+ "text_encoder": ["transformers", "T5EncoderModel"],
+ "tokenizer": ["transformers", "T5TokenizerFast"],
+ "transformer": ["diffusers", "LongLive2Transformer3DModel"],
+ "transformer_2": [None, None],
+ "vae": ["diffusers", "AutoencoderKLWan"],
+ }
+ )
+ )
+
+
+def test_diffusers_cache_validation_rejects_declared_component_without_weights(
+ tmp_path,
+):
+ _write_model_index(tmp_path)
+ for subdir in ("scheduler", "text_encoder", "tokenizer", "transformer", "vae"):
+ (tmp_path / subdir).mkdir()
+ (tmp_path / "text_encoder" / "model.safetensors").write_bytes(b"weights")
+ (tmp_path / "vae" / "diffusion_pytorch_model.bin").write_bytes(b"weights")
+
+ assert not _verify_diffusers_model_complete(str(tmp_path))
+
+ is_valid, missing_files, checked_subdirs = _check_index_files_for_missing_shards(
+ str(tmp_path)
+ )
+ assert not is_valid
+ assert "transformer/" in missing_files
+ assert "transformer" in checked_subdirs
+
+
+def test_diffusers_cache_validation_checks_declared_component_shards(tmp_path):
+ _write_model_index(tmp_path)
+ for subdir in ("scheduler", "text_encoder", "tokenizer", "transformer", "vae"):
+ (tmp_path / subdir).mkdir()
+ (tmp_path / subdir / "model.safetensors").write_bytes(b"weights")
+
+ index_path = (
+ tmp_path / "transformer" / "diffusion_pytorch_model.safetensors.index.json"
+ )
+ index_path.write_text(
+ json.dumps(
+ {
+ "weight_map": {
+ "block.0.weight": "model.safetensors",
+ "block.1.weight": "missing.safetensors",
+ }
+ }
+ )
+ )
+
+ assert _verify_diffusers_model_complete(str(tmp_path))
+
+ is_valid, missing_files, checked_subdirs = _check_index_files_for_missing_shards(
+ str(tmp_path)
+ )
+ assert not is_valid
+ assert "transformer/missing.safetensors" in missing_files
+ assert "transformer" in checked_subdirs
diff --git a/python/sglang/multimodal_gen/test/unit/test_longlive2_pipeline_config.py b/python/sglang/multimodal_gen/test/unit/test_longlive2_pipeline_config.py
new file mode 100644
index 000000000..3646cbce1
--- /dev/null
+++ b/python/sglang/multimodal_gen/test/unit/test_longlive2_pipeline_config.py
@@ -0,0 +1,22 @@
+# SPDX-License-Identifier: Apache-2.0
+import unittest
+
+from sglang.multimodal_gen.configs.pipeline_configs.longlive2 import LongLive2T2VConfig
+
+
+class TestLongLive2AdjustNumFrames(unittest.TestCase):
+ def setUp(self):
+ self.config = LongLive2T2VConfig()
+
+ def test_reuses_wan_temporal_frame_adjustment(self):
+ self.assertEqual(self.config.adjust_num_frames(62), 61)
+
+ def test_keeps_frames_when_latents_match_causal_block(self):
+ self.assertEqual(self.config.adjust_num_frames(93), 93)
+
+ def test_rounds_to_causal_block_aligned_latents(self):
+ self.assertEqual(self.config.adjust_num_frames(65), 61)
+
+
+if __name__ == "__main__":
+ unittest.main()