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()