From 2fc548f2502b00897b19d6949732f1097c387449 Mon Sep 17 00:00:00 2001 From: Mick Date: Tue, 2 Jun 2026 13:52:49 +0800 Subject: [PATCH] [diffusion] model: support lingot-world (#26954) --- python/pyproject.toml | 2 +- .../configs/models/dits/__init__.py | 4 + .../configs/models/dits/lingbot_world.py | 91 ++ .../configs/pipeline_configs/__init__.py | 4 + .../configs/pipeline_configs/base.py | 18 + .../configs/pipeline_configs/lingbot_world.py | 368 +++++ .../configs/sample/lingbot_world.py | 39 + .../configs/sample/sampling_params.py | 7 + python/sglang/multimodal_gen/registry.py | 16 +- .../runtime/entrypoints/http_server.py | 9 +- .../runtime/entrypoints/openai/protocol.py | 26 +- .../entrypoints/openai/realtime/__init__.py | 1 + .../lingbot_world_realtime_adapter.py | 456 ++++++ .../openai/realtime/generate_session.py | 78 + .../openai/realtime/realtime_adapter.py | 74 + .../realtime/realtime_output_adapter.py | 496 ++++++ .../openai/realtime/realtime_video_api.py | 427 ++++++ .../entrypoints/openai/realtime/registry.py | 59 + .../entrypoints/openai/realtime/timing.py | 21 + .../runtime/entrypoints/openai/utils.py | 34 +- .../runtime/entrypoints/utils.py | 343 +++-- .../multimodal_gen/runtime/launch_server.py | 1 + .../runtime/layers/kvcache/__init__.py | 13 + .../layers/kvcache/causal_attention_cache.py | 238 +++ .../multimodal_gen/runtime/layers/usp.py | 178 +++ .../runtime/layers/visual_embedding.py | 55 + .../runtime/managers/gpu_worker.py | 111 +- .../runtime/managers/scheduler.py | 19 + .../runtime/models/dits/causal_wanvideo.py | 245 +-- .../runtime/models/dits/lingbot_world.py | 1360 +++++++++++++++++ .../runtime/models/vaes/wanvae.py | 42 +- .../lingbot_world_causal_dmd_pipeline.py | 99 ++ .../runtime/pipelines_core/schedule_batch.py | 26 +- .../runtime/pipelines_core/stages/__init__.py | 24 + .../pipelines_core/stages/causal_denoising.py | 1221 +++++++++++---- .../stages/condition_encoding.py | 35 + .../stages/latent_preparation.py | 195 ++- .../lingbot_world/__init__.py | 10 + .../lingbot_world_causal_denoising.py | 315 ++++ .../stages/realtime_input_validation.py | 139 ++ .../stages/realtime_text_encoding.py | 152 ++ .../pipelines_core/stages/realtime_vae.py | 204 +++ .../pipelines_core/stages/text_encoding.py | 6 +- .../stages/timestep_preparation.py | 60 + .../runtime/postprocess/__init__.py | 2 + .../postprocess/realesrgan_upscaler.py | 204 ++- .../runtime/realtime/__init__.py | 31 + .../runtime/realtime/causal_state.py | 20 + .../runtime/realtime/condition_events.py | 289 ++++ .../runtime/realtime/session.py | 111 ++ .../runtime/utils/camera_geometry.py | 118 ++ .../runtime/utils/realtime_video.py | 211 +++ .../multimodal_gen/test/server/gpu_cases.py | 21 + .../test/server/perf_baselines.json | 8 + .../test/server/realtime_consistency.py | 403 +++++ .../test/server/test_server_common.py | 43 +- .../test/server/test_server_utils.py | 63 +- .../test/server/testcase_configs.py | 53 +- .../sglang/multimodal_gen/test/test_utils.py | 2 +- .../unit/realtime/test_causal_denoising.py | 503 ++++++ .../realtime/test_lingbot_causal_denoising.py | 297 ++++ .../realtime/test_output_materialization.py | 116 ++ .../test_realtime_condition_events.py | 250 +++ .../test_realtime_consistency_harness.py | 522 +++++++ .../test_realtime_output_transport.py | 631 ++++++++ .../unit/realtime/test_realtime_runtime.py | 883 +++++++++++ .../test/unit/realtime/test_realtime_vae.py | 171 +++ .../test/unit/test_openai_utils.py | 23 + 68 files changed, 11557 insertions(+), 739 deletions(-) create mode 100644 python/sglang/multimodal_gen/configs/models/dits/lingbot_world.py create mode 100644 python/sglang/multimodal_gen/configs/pipeline_configs/lingbot_world.py create mode 100644 python/sglang/multimodal_gen/configs/sample/lingbot_world.py create mode 100644 python/sglang/multimodal_gen/runtime/entrypoints/openai/realtime/__init__.py create mode 100644 python/sglang/multimodal_gen/runtime/entrypoints/openai/realtime/adapters/lingbot_world_realtime_adapter.py create mode 100644 python/sglang/multimodal_gen/runtime/entrypoints/openai/realtime/generate_session.py create mode 100644 python/sglang/multimodal_gen/runtime/entrypoints/openai/realtime/realtime_adapter.py create mode 100644 python/sglang/multimodal_gen/runtime/entrypoints/openai/realtime/realtime_output_adapter.py create mode 100644 python/sglang/multimodal_gen/runtime/entrypoints/openai/realtime/realtime_video_api.py create mode 100644 python/sglang/multimodal_gen/runtime/entrypoints/openai/realtime/registry.py create mode 100644 python/sglang/multimodal_gen/runtime/entrypoints/openai/realtime/timing.py create mode 100644 python/sglang/multimodal_gen/runtime/layers/kvcache/__init__.py create mode 100644 python/sglang/multimodal_gen/runtime/layers/kvcache/causal_attention_cache.py create mode 100644 python/sglang/multimodal_gen/runtime/models/dits/lingbot_world.py create mode 100644 python/sglang/multimodal_gen/runtime/pipelines/lingbot_world_causal_dmd_pipeline.py create mode 100644 python/sglang/multimodal_gen/runtime/pipelines_core/stages/condition_encoding.py create mode 100644 python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/lingbot_world/__init__.py create mode 100644 python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/lingbot_world/lingbot_world_causal_denoising.py create mode 100644 python/sglang/multimodal_gen/runtime/pipelines_core/stages/realtime_input_validation.py create mode 100644 python/sglang/multimodal_gen/runtime/pipelines_core/stages/realtime_text_encoding.py create mode 100644 python/sglang/multimodal_gen/runtime/pipelines_core/stages/realtime_vae.py create mode 100644 python/sglang/multimodal_gen/runtime/realtime/__init__.py create mode 100644 python/sglang/multimodal_gen/runtime/realtime/causal_state.py create mode 100644 python/sglang/multimodal_gen/runtime/realtime/condition_events.py create mode 100644 python/sglang/multimodal_gen/runtime/realtime/session.py create mode 100644 python/sglang/multimodal_gen/runtime/utils/camera_geometry.py create mode 100644 python/sglang/multimodal_gen/runtime/utils/realtime_video.py create mode 100644 python/sglang/multimodal_gen/test/server/realtime_consistency.py create mode 100644 python/sglang/multimodal_gen/test/unit/realtime/test_causal_denoising.py create mode 100644 python/sglang/multimodal_gen/test/unit/realtime/test_lingbot_causal_denoising.py create mode 100644 python/sglang/multimodal_gen/test/unit/realtime/test_output_materialization.py create mode 100644 python/sglang/multimodal_gen/test/unit/realtime/test_realtime_condition_events.py create mode 100644 python/sglang/multimodal_gen/test/unit/realtime/test_realtime_consistency_harness.py create mode 100644 python/sglang/multimodal_gen/test/unit/realtime/test_realtime_output_transport.py create mode 100644 python/sglang/multimodal_gen/test/unit/realtime/test_realtime_runtime.py create mode 100644 python/sglang/multimodal_gen/test/unit/realtime/test_realtime_vae.py create mode 100644 python/sglang/multimodal_gen/test/unit/test_openai_utils.py diff --git a/python/pyproject.toml b/python/pyproject.toml index be2f05b97..9662dd289 100755 --- a/python/pyproject.toml +++ b/python/pyproject.toml @@ -82,7 +82,7 @@ dependencies = [ "watchfiles", "xgrammar==0.2.1", "smg-grpc-servicer>=0.5.0", - "kernels<0.15", + "kernels>=0.14.1,<0.15", ] [[tool.uv.index]] diff --git a/python/sglang/multimodal_gen/configs/models/dits/__init__.py b/python/sglang/multimodal_gen/configs/models/dits/__init__.py index c3ee938b6..bc540c18d 100644 --- a/python/sglang/multimodal_gen/configs/models/dits/__init__.py +++ b/python/sglang/multimodal_gen/configs/models/dits/__init__.py @@ -4,6 +4,9 @@ from sglang.multimodal_gen.configs.models.dits.cosmos3video import Cosmos3VideoC from sglang.multimodal_gen.configs.models.dits.helios import HeliosConfig from sglang.multimodal_gen.configs.models.dits.hunyuan3d import Hunyuan3DDiTConfig from sglang.multimodal_gen.configs.models.dits.hunyuanvideo import HunyuanVideoConfig +from sglang.multimodal_gen.configs.models.dits.lingbot_world import ( + LingBotWorldVideoConfig, +) 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 ( @@ -15,6 +18,7 @@ __all__ = [ "Cosmos3VideoConfig", "HeliosConfig", "HunyuanVideoConfig", + "LingBotWorldVideoConfig", "WanVideoConfig", "Hunyuan3DDiTConfig", "MOVAAudioConfig", diff --git a/python/sglang/multimodal_gen/configs/models/dits/lingbot_world.py b/python/sglang/multimodal_gen/configs/models/dits/lingbot_world.py new file mode 100644 index 000000000..6897a8449 --- /dev/null +++ b/python/sglang/multimodal_gen/configs/models/dits/lingbot_world.py @@ -0,0 +1,91 @@ +# Copied and adapted from: https://github.com/hao-ai-lab/FastVideo +# Adapted from: https://github.com/Robbyant/lingbot-world + +# SPDX-License-Identifier: Apache-2.0 +from dataclasses import dataclass, field + +from sglang.multimodal_gen.configs.models.dits.base import DiTArchConfig, DiTConfig + + +def is_blocks(n: str, m) -> bool: + return "blocks" in n and str.isdigit(n.split(".")[-1]) + + +@dataclass +class LingBotWorldArchConfig(DiTArchConfig): + _fsdp_shard_conditions: list = field(default_factory=lambda: [is_blocks]) + + param_names_mapping: dict = field( + default_factory=lambda: { + r"^patch_embedding\.(.*)$": r"patch_embedding.proj.\1", + r"^patch_embedding_wancamctrl\.(.*)$": r"patch_embedding_wancamctrl.proj.\1", + r"^c2ws_hidden_states_layer1\.(.*)$": r"c2ws_mlp.fc_in.\1", + r"^c2ws_hidden_states_layer2\.(.*)$": r"c2ws_mlp.fc_out.\1", + r"^text_embedding\.0\.(.*)$": r"condition_embedder.text_embedder.fc_in.\1", + r"^text_embedding\.2\.(.*)$": r"condition_embedder.text_embedder.fc_out.\1", + r"^time_embedding\.0\.(.*)$": r"condition_embedder.time_embedder.mlp.fc_in.\1", + r"^time_embedding\.2\.(.*)$": r"condition_embedder.time_embedder.mlp.fc_out.\1", + r"^time_projection\.1\.(.*)$": r"condition_embedder.time_modulation.linear.\1", + r"^blocks\.(\d+)\.modulation$": r"blocks.\1.scale_shift_table", + r"^blocks\.(\d+)\.self_attn\.q\.(.*)$": r"blocks.\1.to_q.\2", + r"^blocks\.(\d+)\.self_attn\.k\.(.*)$": r"blocks.\1.to_k.\2", + r"^blocks\.(\d+)\.self_attn\.v\.(.*)$": r"blocks.\1.to_v.\2", + r"^blocks\.(\d+)\.self_attn\.o\.(.*)$": r"blocks.\1.to_out.\2", + r"^blocks\.(\d+)\.self_attn\.norm_q\.(.*)$": r"blocks.\1.norm_q.\2", + r"^blocks\.(\d+)\.self_attn\.norm_k\.(.*)$": r"blocks.\1.norm_k.\2", + r"^blocks\.(\d+)\.norm3\.(.*)$": r"blocks.\1.self_attn_residual_norm.norm.\2", + r"^blocks\.(\d+)\.cross_attn\.q\.(.*)$": r"blocks.\1.attn2.to_q.\2", + r"^blocks\.(\d+)\.cross_attn\.k\.(.*)$": r"blocks.\1.attn2.to_k.\2", + r"^blocks\.(\d+)\.cross_attn\.v\.(.*)$": r"blocks.\1.attn2.to_v.\2", + r"^blocks\.(\d+)\.cross_attn\.o\.(.*)$": r"blocks.\1.attn2.to_out.\2", + r"^blocks\.(\d+)\.cross_attn\.norm_q\.(.*)$": r"blocks.\1.attn2.norm_q.\2", + r"^blocks\.(\d+)\.cross_attn\.norm_k\.(.*)$": r"blocks.\1.attn2.norm_k.\2", + r"^blocks\.(\d+)\.ffn\.0\.(.*)$": r"blocks.\1.ffn.fc_in.\2", + r"^blocks\.(\d+)\.ffn\.2\.(.*)$": r"blocks.\1.ffn.fc_out.\2", + r"^blocks\.(\d+)\.cam_injector_layer1\.(.*)$": r"blocks.\1.cam_conditioner.cam_injector.fc_in.\2", + r"^blocks\.(\d+)\.cam_injector_layer2\.(.*)$": r"blocks.\1.cam_conditioner.cam_injector.fc_out.\2", + r"^blocks\.(\d+)\.cam_scale_layer\.(.*)$": r"blocks.\1.cam_conditioner.cam_scale_layer.\2", + r"^blocks\.(\d+)\.cam_shift_layer\.(.*)$": r"blocks.\1.cam_conditioner.cam_shift_layer.\2", + r"^head\.modulation$": r"scale_shift_table", + r"^head\.head\.(.*)$": r"proj_out.\1", + } + ) + reverse_param_names_mapping: dict = field(default_factory=lambda: {}) + lora_param_names_mapping: dict = field(default_factory=lambda: {}) + + patch_size: tuple[int, int, int] = (1, 2, 2) + text_len: int = 512 + num_attention_heads: int = 40 + attention_head_dim: int = 128 + in_channels: int = 36 + out_channels: int = 16 + text_dim: int = 4096 + freq_dim: int = 256 + ffn_dim: int = 13824 + num_layers: int = 40 + cross_attn_norm: bool = True + qk_norm: str = "rms_norm_across_heads" + eps: float = 1e-6 + image_dim: int | None = None + added_kv_proj_dim: int | None = None + rope_max_seq_len: int = 1024 + pos_embed_seq_len: int | None = None + exclude_lora_layers: list[str] = field(default_factory=lambda: ["embedder"]) + boundary_ratio: float | None = None + local_attn_size: int = -1 + sink_size: int = 3 + num_frames_per_block: int = 3 + sliding_window_num_frames: int = 45 + + def __post_init__(self): + super().__post_init__() + self.out_channels = self.out_channels or self.in_channels + self.hidden_size = self.num_attention_heads * self.attention_head_dim + self.num_channels_latents = self.out_channels + + +@dataclass +class LingBotWorldVideoConfig(DiTConfig): + arch_config: DiTArchConfig = field(default_factory=LingBotWorldArchConfig) + + prefix: str = "Wan" diff --git a/python/sglang/multimodal_gen/configs/pipeline_configs/__init__.py b/python/sglang/multimodal_gen/configs/pipeline_configs/__init__.py index 831f6bfcf..9dde6536c 100644 --- a/python/sglang/multimodal_gen/configs/pipeline_configs/__init__.py +++ b/python/sglang/multimodal_gen/configs/pipeline_configs/__init__.py @@ -28,6 +28,9 @@ from sglang.multimodal_gen.configs.pipeline_configs.hunyuan import ( from sglang.multimodal_gen.configs.pipeline_configs.hunyuan3d import ( Hunyuan3D2PipelineConfig, ) +from sglang.multimodal_gen.configs.pipeline_configs.lingbot_world import ( + LingBotWorldCausalDMDConfig, +) from sglang.multimodal_gen.configs.pipeline_configs.ltx_2 import LTX2PipelineConfig from sglang.multimodal_gen.configs.pipeline_configs.mova import MOVAPipelineConfig from sglang.multimodal_gen.configs.pipeline_configs.sana import SanaPipelineConfig @@ -68,4 +71,5 @@ __all__ = [ "SelfForcingWanT2V480PConfig", "ZImagePipelineConfig", "LTX2PipelineConfig", + "LingBotWorldCausalDMDConfig", ] diff --git a/python/sglang/multimodal_gen/configs/pipeline_configs/base.py b/python/sglang/multimodal_gen/configs/pipeline_configs/base.py index 5590f6182..70a177157 100644 --- a/python/sglang/multimodal_gen/configs/pipeline_configs/base.py +++ b/python/sglang/multimodal_gen/configs/pipeline_configs/base.py @@ -311,6 +311,9 @@ class PipelineConfig: (target_width, target_height), PIL.Image.Resampling.LANCZOS ), (target_width, target_height) + def preprocess_realtime_condition_image(self, batch, _vae_image_processor) -> bool: + return False + def prepare_calculated_size(self, image): return self.calculate_condition_image_size(image, image.width, image.height) @@ -686,6 +689,9 @@ class PipelineConfig: def prepare_neg_cond_kwargs(self, batch, device, rotary_emb, dtype): return {} + def prepare_world_condition(self, batch, device, dtype): + return None + def _unpad_and_unpack_latents(self, latents, audio_latents, batch, vae, audio_vae): raise NotImplementedError("not yet implemented") @@ -806,6 +812,18 @@ class PipelineConfig: default=PipelineConfig.dmd_denoising_steps, help="Comma-separated list of denoising steps (e.g., '1000,757,522')", ) + parser.add_argument( + f"--{prefix_with_dot}realtime-causal-sink-size", + type=int, + default=None, + help="Override the number of sink frames kept by realtime causal DiT pipelines that support it.", + ) + parser.add_argument( + f"--{prefix_with_dot}realtime-causal-kv-cache-num-frames", + type=int, + default=None, + help="Override the total frame capacity of realtime causal DiT KV cache for pipelines that support it.", + ) # Add VAE configuration arguments from sglang.multimodal_gen.configs.models.vaes.base import VAEConfig diff --git a/python/sglang/multimodal_gen/configs/pipeline_configs/lingbot_world.py b/python/sglang/multimodal_gen/configs/pipeline_configs/lingbot_world.py new file mode 100644 index 000000000..49b81b416 --- /dev/null +++ b/python/sglang/multimodal_gen/configs/pipeline_configs/lingbot_world.py @@ -0,0 +1,368 @@ +# Copied and adapted from: https://github.com/hao-ai-lab/FastVideo +# Adapted from: https://github.com/Robbyant/lingbot-world + +# SPDX-License-Identifier: Apache-2.0 +import html +import re +from collections.abc import Callable +from dataclasses import dataclass, field +from typing import Any + +import numpy as np +import torch + +from sglang.multimodal_gen.configs.models import DiTConfig +from sglang.multimodal_gen.configs.models.dits import LingBotWorldVideoConfig +from sglang.multimodal_gen.configs.pipeline_configs.wan import Wan2_2_I2V_A14B_Config +from sglang.multimodal_gen.runtime.realtime.session import ( + BaseRealtimeState, +) +from sglang.multimodal_gen.runtime.utils.camera_geometry import ( + camera_poses_to_plucker, + compute_relative_poses, +) +from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger + +logger = init_logger(__name__) + + +def lingbot_prompt_clean(text: str) -> str: + try: + import ftfy + + text = ftfy.fix_text(text) + except ImportError: + pass + text = html.unescape(html.unescape(text)) + text = re.sub(r"\s+", " ", text) + return text.strip() + + +class _LingBotWorldCameraState(BaseRealtimeState): + def __init__(self): + super().__init__() + self.action_history: list[list[str]] = [] + self.last_actions: list[str] = [] + + def reset_camera_actions(self): + self.action_history.clear() + self.last_actions = [] + + def append_camera_actions(self, camera_actions: list[list[str]]) -> None: + for actions in camera_actions: + normalized = list(actions) + self.action_history.append(normalized) + self.last_actions = normalized + + def dispose(self): + super().dispose() + self.reset_camera_actions() + + +def _validate_actions(actions: Any) -> list[list[str]]: + if not isinstance(actions, list): + raise TypeError("actions must be a list[list[str]]") + result: list[list[str]] = [] + for frame_actions in actions: + if not isinstance(frame_actions, list): + raise TypeError("actions must be a list[list[str]]") + result.append(list(frame_actions)) + return result + + +def _pad_actions_to_chunk( + action_history: list[list[str]], chunk_size: int +) -> list[list[str]]: + if len(action_history) >= chunk_size: + return action_history + fill_item = action_history[-1] if action_history else [] + return action_history + [ + list(fill_item) for _ in range(chunk_size - len(action_history)) + ] + + +def _get_rotation_matrix(axis: str, angle_rad: float) -> np.ndarray: + def calculate_c_s(): + return np.cos(angle_rad), np.sin(angle_rad) + + if axis == "x": + c, s = calculate_c_s() + return np.array([[1, 0, 0], [0, c, -s], [0, s, c]]) + if axis == "y": + c, s = calculate_c_s() + return np.array([[c, 0, s], [0, 1, 0], [-s, 0, c]]) + if axis == "z": + c, s = calculate_c_s() + return np.array([[c, -s, 0], [s, c, 0], [0, 0, 1]]) + return np.eye(3) + + +def _actions_to_c2ws(action_history: list[list[str]]) -> list[np.ndarray]: + move_speed = 0.05 + rotate_speed_rad_ik = np.deg2rad(4.0) + rotate_speed_rad_jl = np.deg2rad(6.0) + + current_c2w = np.eye(4) + current_pitch = 0.0 + pitch_limit = np.deg2rad(85) + all_matrices = [current_c2w] + + for frame_keys in action_history: + R = current_c2w[:3, :3] + T = current_c2w[:3, 3] + + pitch_delta = 0.0 + if "i" in frame_keys: + pitch_delta += rotate_speed_rad_ik + if "k" in frame_keys: + pitch_delta -= rotate_speed_rad_ik + + new_pitch = current_pitch + pitch_delta + if -pitch_limit <= new_pitch <= pitch_limit: + current_pitch = new_pitch + else: + pitch_delta = 0.0 + + yaw_delta = 0.0 + if "j" in frame_keys: + yaw_delta -= rotate_speed_rad_jl + if "l" in frame_keys: + yaw_delta += rotate_speed_rad_jl + + R_pitch = _get_rotation_matrix("x", pitch_delta) + R_yaw = _get_rotation_matrix("y", yaw_delta) + R_new = R_yaw @ R @ R_pitch + + vec_right = R_new[:, 0] + vec_forward = R_new[:, 2] + forward_flat = np.array([vec_forward[0], 0, vec_forward[2]]) + right_flat = np.array([vec_right[0], 0, vec_right[2]]) + + f_norm = np.linalg.norm(forward_flat) + r_norm = np.linalg.norm(right_flat) + if f_norm > 0: + forward_flat = forward_flat / (f_norm + 1e-6) + if r_norm > 0: + right_flat = right_flat / (r_norm + 1e-6) + + move_vec = np.zeros(3) + if "w" in frame_keys: + move_vec += forward_flat * move_speed + if "s" in frame_keys: + move_vec -= forward_flat * move_speed + if "d" in frame_keys: + move_vec += right_flat * move_speed + if "a" in frame_keys: + move_vec -= right_flat * move_speed + + T_new = T + move_vec + current_c2w = np.eye(4) + current_c2w[:3, :3] = R_new + current_c2w[:3, 3] = T_new + all_matrices.append(current_c2w) + + return all_matrices + + +def _get_camera_control( + action_history: list[list[str]], + *, + chunk_size: int, + width: int, + height: int, + device: torch.device | str, + dtype: torch.dtype, +) -> tuple[torch.Tensor, torch.Tensor]: + c2ws_list = _actions_to_c2ws(action_history) + c2ws_np = np.stack(c2ws_list[1:]) + c2ws = torch.from_numpy(c2ws_np).to(device=device, dtype=dtype) + Ks = torch.tensor( + [[500.0, 500.0, width / 2, height / 2]], + device=device, + dtype=dtype, + ).repeat(chunk_size, 1) + logger.debug("prefix c2ws shape: %s, Ks shape: %s", c2ws.shape, Ks.shape) + return c2ws, Ks + + +def _build_camera_condition( + *, + action_history: list[list[str]], + width: int, + height: int, + spatial_scale: int, + device: torch.device | str, + dtype: torch.dtype, + tail_chunk_size: int, +) -> torch.Tensor: + action_history = _pad_actions_to_chunk(action_history, tail_chunk_size) + c2ws_prefix, Ks = _get_camera_control( + action_history, + chunk_size=tail_chunk_size, + width=width, + height=height, + device=device, + dtype=dtype, + ) + c2ws_prefix = compute_relative_poses(c2ws_prefix, framewise=True) + c2ws_prefix = c2ws_prefix[-tail_chunk_size:] + + return camera_poses_to_plucker( + c2ws=c2ws_prefix, + Ks=Ks, + height=height, + width=width, + spatial_scale=spatial_scale, + device=device, + dtype=dtype, + ) + + +def _prepare_lingbot_world_condition( + *, + batch, + pipeline_config, + device: torch.device | str, + dtype: torch.dtype, +) -> torch.Tensor | None: + if batch.c2ws_plucker_emb is not None: + return batch.c2ws_plucker_emb.to(device=device, dtype=dtype) + + actions = batch.condition_inputs.get("camera_actions") + if actions is None: + return None + + spatial_scale = pipeline_config.vae_config.arch_config.spatial_compression_ratio + chunk_size = batch.realtime_chunk_size or max( + 1, + int(pipeline_config.dit_config.arch_config.num_frames_per_block), + ) + + normalized_actions = _validate_actions(actions) + if len(normalized_actions) == 0: + return None + + if batch.session is None: + action_history = normalized_actions + else: + state = batch.session.get_or_create_state(_LingBotWorldCameraState) + if batch.block_idx == 0: + state.reset_camera_actions() + state.append_camera_actions(normalized_actions) + action_history = state.action_history + + if len(action_history) == 0: + return None + + c2ws_plucker_emb = _build_camera_condition( + action_history=action_history, + width=int(batch.width), + height=int(batch.height), + spatial_scale=spatial_scale, + device=device, + dtype=dtype, + tail_chunk_size=chunk_size, + ) + logger.debug( + "LingBot action condition prepared: session_id=%s, block_idx=%s, new_action_count=%s, total_history=%s", + batch.realtime_session_id, + batch.block_idx, + len(normalized_actions), + len(action_history), + ) + return c2ws_plucker_emb + + +@dataclass +class LingBotWorldI2VConfig(Wan2_2_I2V_A14B_Config): + dit_config: DiTConfig = field(default_factory=LingBotWorldVideoConfig) + flow_shift: float | None = 10.0 + boundary_ratio: float | None = 0.947 + text_encoder_precisions: tuple[str, ...] = field(default_factory=lambda: ("bf16",)) + preprocess_text_funcs: tuple[Callable[[str], str] | None, ...] = field( + default_factory=lambda: (lingbot_prompt_clean,) + ) + + def prepare_pos_cond_kwargs(self, batch, device, rotary_emb, dtype): + kwargs = super().prepare_pos_cond_kwargs(batch, device, rotary_emb, dtype) + if batch.c2ws_plucker_emb is not None: + kwargs["c2ws_plucker_emb"] = batch.c2ws_plucker_emb.to( + device=device, dtype=dtype + ) + return kwargs + + def preprocess_realtime_condition_image(self, batch, _vae_image_processor) -> bool: + if batch.condition_image is None: + return False + if isinstance(batch.condition_image, list): + batch.condition_image = batch.condition_image[0] + + width = int(batch.width or 832) + height = int(batch.height or 480) + batch.condition_image = batch.condition_image.resize((width, height)) + batch.width = width + batch.height = height + return True + + def prepare_world_condition(self, batch, device, dtype): + c2ws_plucker_emb = _prepare_lingbot_world_condition( + batch=batch, + pipeline_config=self, + device=device, + dtype=dtype, + ) + if c2ws_plucker_emb is None: + return None + return {"c2ws_plucker_emb": c2ws_plucker_emb} + + +@dataclass +class LingBotWorldCausalDMDConfig(LingBotWorldI2VConfig): + is_causal: bool = True + dmd_denoising_steps: list[int] | None = field( + default_factory=lambda: [1000, 821, 642, 321] + ) + warp_denoising_step: bool = True + realtime_causal_sink_size: int | None = None + realtime_causal_kv_cache_num_frames: int | None = None + + def postprocess_image_latent(self, latent_condition, batch): + """Build condition tensor aligned to chunk_size (num_frames_per_block). + + Matches lingbot_fast_server's _prepare_latents_causal: + condition = [mask(temporal_ratio ch), latent(z_dim ch)] -> 20ch total, + with temporal dim aligned to chunk_size. + """ + vae_arch = self.vae_config.arch_config + temporal_ratio = vae_arch.temporal_compression_ratio + spatial_ratio = vae_arch.spatial_compression_ratio + chunk_size = self.dit_config.arch_config.num_frames_per_block + + latent_height = batch.height // spatial_ratio + latent_width = batch.width // spatial_ratio + + # Align num_latent_frames to chunk_size + num_latent_frames = latent_condition.shape[2] + num_latent_frames = num_latent_frames - (num_latent_frames % chunk_size) + latent_condition = latent_condition[:, :, :num_latent_frames, :, :] + + # Number of initial frames that have actual image content + # (latent_condition from VAE encode of [image, zeros...]) + # First frame is real, rest are zero-padded + initial_latent_frames = 1 # single image -> 1 latent frame + + # Build mask: [B, temporal_ratio, num_latent_frames, H, W] + mask = torch.ones( + 1, + temporal_ratio, + num_latent_frames, + latent_height, + latent_width, + dtype=latent_condition.dtype, + device=latent_condition.device, + ) + # Zero out mask for frames beyond the initial image + if initial_latent_frames < num_latent_frames: + mask[:, :, initial_latent_frames:] = 0 + + return torch.cat([mask, latent_condition], dim=1) diff --git a/python/sglang/multimodal_gen/configs/sample/lingbot_world.py b/python/sglang/multimodal_gen/configs/sample/lingbot_world.py new file mode 100644 index 000000000..1165ca62f --- /dev/null +++ b/python/sglang/multimodal_gen/configs/sample/lingbot_world.py @@ -0,0 +1,39 @@ +# Copied and adapted from: https://github.com/hao-ai-lab/FastVideo +# Adapted from: https://github.com/Robbyant/lingbot-world + +# SPDX-License-Identifier: Apache-2.0 +from dataclasses import dataclass + +from sglang.multimodal_gen.configs.sample.wan import Wan2_2_I2V_A14B_SamplingParam + + +@dataclass +class LingBotWorldSamplingParams(Wan2_2_I2V_A14B_SamplingParam): + negative_prompt: str | None = None + actions: list[list[str]] | None = None + chunk_size: int | None = None + guidance_scale: float = 5.0 + guidance_scale_2: float = 5.0 + num_inference_steps: int = 70 + num_frames: int = 117 + fps: int = 16 + + def _adjust(self, server_args): + enable_sequence_shard = self.enable_sequence_shard + if enable_sequence_shard is None or enable_sequence_shard: + self.adjust_frames = False + super()._adjust(server_args) + if enable_sequence_shard is None or enable_sequence_shard: + self.enable_sequence_shard = True + self.adjust_frames = False + if self.chunk_size is None: + self.chunk_size = max( + 1, + int( + server_args.pipeline_config.dit_config.arch_config.num_frames_per_block + ), + ) + if self.actions is not None: + self.condition_inputs["camera_actions"] = self.actions + if self.chunk_size is not None: + self.realtime_chunk_size = self.chunk_size diff --git a/python/sglang/multimodal_gen/configs/sample/sampling_params.py b/python/sglang/multimodal_gen/configs/sample/sampling_params.py index 50f446480..aab3c62c1 100644 --- a/python/sglang/multimodal_gen/configs/sample/sampling_params.py +++ b/python/sglang/multimodal_gen/configs/sample/sampling_params.py @@ -216,6 +216,7 @@ class SamplingParams: # if True, suppress verbose logging for this request suppress_logs: bool = field(default=False, metadata={"batch_sig_exclude": True}) + # return output file paths directly to client return_file_paths_only: bool = True enable_sequence_shard: bool | None = None diffusers_kwargs: dict | None = None @@ -227,6 +228,8 @@ class SamplingParams: use_resolution_template: bool | None = None use_system_prompt: bool | None = None use_guardrails: bool | None = None + condition_inputs: dict[str, Any] = field(default_factory=dict) + realtime_chunk_size: int | None = None # Prompt enhancement (ErnieImage) use_pe: bool | None = None @@ -309,6 +312,10 @@ class SamplingParams: def apply_request_extra(self, req: Any) -> None: """Merge request extras (model specific, e.g., LTX2.3) into an already-created pipeline request.""" req.extra.update(self.build_request_extra()) + if self.condition_inputs: + req.condition_inputs.update(self.condition_inputs) + if self.realtime_chunk_size is not None: + req.realtime_chunk_size = self.realtime_chunk_size def _adjust_output_quality(self, output_quality: str, data_type: DataType) -> int: """Convert output_quality string to compression level.""" diff --git a/python/sglang/multimodal_gen/registry.py b/python/sglang/multimodal_gen/registry.py index d557516fb..384cb9059 100644 --- a/python/sglang/multimodal_gen/registry.py +++ b/python/sglang/multimodal_gen/registry.py @@ -36,6 +36,7 @@ from sglang.multimodal_gen.configs.pipeline_configs import ( HeliosMidConfig, HeliosT2VConfig, HunyuanConfig, + LingBotWorldCausalDMDConfig, WanI2V480PConfig, WanI2V720PConfig, WanT2V480PConfig, @@ -104,7 +105,12 @@ from sglang.multimodal_gen.configs.sample.hunyuan import ( HunyuanSamplingParams, ) from sglang.multimodal_gen.configs.sample.hunyuan3d import Hunyuan3DSamplingParams -from sglang.multimodal_gen.configs.sample.joy_image import JoyImageEditSamplingParams +from sglang.multimodal_gen.configs.sample.joy_image import ( + JoyImageEditSamplingParams, +) +from sglang.multimodal_gen.configs.sample.lingbot_world import ( + LingBotWorldSamplingParams, +) from sglang.multimodal_gen.configs.sample.ltx_2 import ( LTX2SamplingParams, LTX23HQSamplingParams, @@ -751,6 +757,14 @@ def _register_configs(): pipeline_config_cls=Wan2_2_I2V_A14B_Config, hf_model_paths=["Wan-AI/Wan2.2-I2V-A14B-Diffusers"], ) + register_configs( + sampling_param_cls=LingBotWorldSamplingParams, + pipeline_config_cls=LingBotWorldCausalDMDConfig, + hf_model_paths=[ + "IPostYellow/lingbot-world-fast-diffusers", + "robbyant/lingbot-world-fast-diffusers", + ], + ) register_configs( sampling_param_cls=FastWanT2V480PConfig, pipeline_config_cls=FastWan2_1_T2V_480P_Config, diff --git a/python/sglang/multimodal_gen/runtime/entrypoints/http_server.py b/python/sglang/multimodal_gen/runtime/entrypoints/http_server.py index 6f56cb91e..dfc55a6b1 100644 --- a/python/sglang/multimodal_gen/runtime/entrypoints/http_server.py +++ b/python/sglang/multimodal_gen/runtime/entrypoints/http_server.py @@ -13,10 +13,16 @@ import torch from fastapi import APIRouter, FastAPI, Request from sglang.multimodal_gen.configs.sample.sampling_params import SamplingParams -from sglang.multimodal_gen.runtime.entrypoints.openai import image_api, video_api +from sglang.multimodal_gen.runtime.entrypoints.openai import ( + image_api, + video_api, +) from sglang.multimodal_gen.runtime.entrypoints.openai.protocol import ( VertexGenerateReqInput, ) +from sglang.multimodal_gen.runtime.entrypoints.openai.realtime import ( + realtime_video_api, +) from sglang.multimodal_gen.runtime.entrypoints.openai.utils import build_sampling_params from sglang.multimodal_gen.runtime.entrypoints.post_training import ( rollout_api, @@ -401,6 +407,7 @@ def create_app(server_args: ServerArgs): app.include_router(common_api.router) app.include_router(image_api.router) app.include_router(video_api.router) + app.include_router(realtime_video_api.router) app.include_router(mesh_api.router) app.include_router(weights_api.router) app.include_router(rollout_api.router) diff --git a/python/sglang/multimodal_gen/runtime/entrypoints/openai/protocol.py b/python/sglang/multimodal_gen/runtime/entrypoints/openai/protocol.py index 3d6540546..b029a648c 100644 --- a/python/sglang/multimodal_gen/runtime/entrypoints/openai/protocol.py +++ b/python/sglang/multimodal_gen/runtime/entrypoints/openai/protocol.py @@ -2,7 +2,7 @@ import time import uuid from abc import ABC from dataclasses import dataclass, field -from typing import Any, Dict, List, Optional, Union +from typing import Any, Dict, List, Literal, Optional, Union from pydantic import BaseModel, ConfigDict, Field @@ -138,6 +138,30 @@ class VideoRemixRequest(BaseModel): prompt: str +class RealtimeVideoGenerationsRequest(VideoGenerationsRequest): + type: Literal["init"] + # WebSocket does not support multipart/form-data image uploads + first_frame: Optional[bytes | str] = None + condition_inputs: Optional[Dict[str, Any]] = None + max_chunks: Optional[int] = Field(default=None, ge=1) + seed: Optional[int] = 42 + guidance_scale: Optional[float] = 1.0 + size: Optional[str] = "832x480" + profile: Optional[bool] = False + num_profiled_timesteps: Optional[int] = None + profile_all_stages: Optional[bool] = False + realtime_output_format: Optional[Literal["raw", "webp", "jpeg"]] = None + realtime_causal_sink_size: Optional[int] = None + realtime_causal_kv_cache_num_frames: Optional[int] = None + + +class RealtimeEvent(BaseModel): + type: Literal["event"] + kind: str + payload: Any = None + event_id: Optional[int] = None + + # Mesh API protocol models class MeshResponse(BaseModel): id: str diff --git a/python/sglang/multimodal_gen/runtime/entrypoints/openai/realtime/__init__.py b/python/sglang/multimodal_gen/runtime/entrypoints/openai/realtime/__init__.py new file mode 100644 index 000000000..988131360 --- /dev/null +++ b/python/sglang/multimodal_gen/runtime/entrypoints/openai/realtime/__init__.py @@ -0,0 +1 @@ +# SPDX-License-Identifier: Apache-2.0 diff --git a/python/sglang/multimodal_gen/runtime/entrypoints/openai/realtime/adapters/lingbot_world_realtime_adapter.py b/python/sglang/multimodal_gen/runtime/entrypoints/openai/realtime/adapters/lingbot_world_realtime_adapter.py new file mode 100644 index 000000000..b78df7f12 --- /dev/null +++ b/python/sglang/multimodal_gen/runtime/entrypoints/openai/realtime/adapters/lingbot_world_realtime_adapter.py @@ -0,0 +1,456 @@ +# SPDX-License-Identifier: Apache-2.0 + +from __future__ import annotations + +import os +import tempfile +from collections import deque +from typing import TYPE_CHECKING, Any + +from fastapi import WebSocket + +from sglang.multimodal_gen.runtime.entrypoints.openai.protocol import ( + RealtimeEvent, + RealtimeVideoGenerationsRequest, +) +from sglang.multimodal_gen.runtime.entrypoints.openai.realtime.realtime_adapter import ( + RealtimeChunkInputs, + RealtimeModelAdapter, +) +from sglang.multimodal_gen.runtime.entrypoints.openai.realtime.realtime_output_adapter import ( + RawRGBRealtimeOutputAdapter, + RealtimeFrameSendStats, +) +from sglang.multimodal_gen.runtime.entrypoints.openai.utils import ( + build_sampling_params, + save_image_to_path, +) +from sglang.multimodal_gen.runtime.entrypoints.utils import ( + prepare_request, +) +from sglang.multimodal_gen.runtime.realtime.condition_events import ( + ConditionEvent, + ConditionEventQueue, + ControlSignal, + ControlStateSamplingQueue, + ControlStateTransition, +) +from sglang.multimodal_gen.runtime.server_args import get_global_server_args + +if TYPE_CHECKING: + from sglang.multimodal_gen.runtime.entrypoints.openai.realtime.generate_session import ( + GenerateSession, + RealtimeChunkContext, + ) + from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import ( + OutputBatch, + Req, + ) + from sglang.multimodal_gen.runtime.server_args import ServerArgs + + +LINGBOT_REALTIME_DEFAULT_NUM_INFERENCE_STEPS = 4 +LINGBOT_REALTIME_MIN_CONDITION_CHUNKS = 2 + + +class LingBotWorldRealtimeState: + def __init__(self): + self.events = ConditionEventQueue(max_events={"prompt": 1}) + self.camera_state = ControlStateSamplingQueue( + default_item=[], + min_pulse_items=1, + max_transitions=512, + ) + self.camera_script_queue: deque[ControlSignal] = deque(maxlen=512) + self.latest_sampled_event_id: int | None = None + + def clear(self) -> None: + self.events.clear() + self.camera_state.clear() + self.camera_script_queue.clear() + self.latest_sampled_event_id = None + + def receive_prompt(self, prompt: str, *, event_id: int | None = None) -> None: + self.events.push( + ConditionEvent( + kind="prompt", + payload=ControlSignal( + kind="prompt", + payload=prompt, + seq_id=event_id, + ), + ) + ) + + def receive_camera_script( + self, + camera_actions: list[list[str]], + *, + event_id: int | None = None, + ) -> None: + self.camera_script_queue.clear() + self.camera_state.clear() + for actions in camera_actions: + self.camera_script_queue.append( + ControlSignal( + kind="camera_actions", + payload=list(actions), + seq_id=event_id, + ) + ) + + def receive_camera_state_transitions( + self, + transitions: list[ControlStateTransition], + ) -> None: + self.camera_script_queue.clear() + self.camera_state.push_many(transitions) + + def receive_camera_actions( + self, + camera_actions: list[list[str]], + *, + event_id: int | None = None, + ) -> None: + self.receive_camera_script(camera_actions, event_id=event_id) + + def receive_camera_state( + self, + actions: list[str], + *, + event_id: int | None = None, + timestamp_ms: int | None = None, + ) -> None: + self.receive_camera_state_transitions( + [ + ControlStateTransition( + payload=list(actions), + seq_id=event_id, + timestamp_ms=timestamp_ms, + ) + ] + ) + + def _sample_camera_script(self, chunk_size: int) -> list[list[str]]: + chunk: list[list[str]] = [] + latest_event_id = self.latest_sampled_event_id + while self.camera_script_queue and len(chunk) < chunk_size: + signal = self.camera_script_queue.popleft() + chunk.append(list(signal.payload)) + latest_event_id = signal.seq_id + while len(chunk) < chunk_size: + chunk.append([]) + self.latest_sampled_event_id = latest_event_id + return chunk + + def _camera_state_transition( + self, + actions: list[str], + *, + event_id: int | None, + timestamp_ms: int | None, + ) -> ControlStateTransition: + return ControlStateTransition( + payload=list(actions), + seq_id=event_id, + timestamp_ms=timestamp_ms, + ) + + def _camera_transitions_from_event_payload( + self, + payload: dict[str, Any], + *, + event_id: int | None, + ) -> list[ControlStateTransition]: + transitions = payload.get("transitions") + if not isinstance(transitions, list): + raise ValueError("camera_actions state payload requires transitions") + result = [] + for transition in transitions: + if not isinstance(transition, dict): + raise ValueError("camera_actions transition must be a map") + actions = transition.get("actions") + if not isinstance(actions, list): + raise ValueError("camera_actions transition actions must be a list") + timestamp_ms = transition.get("client_ts_ms") + if timestamp_ms is not None: + timestamp_ms = int(timestamp_ms) + result.append( + self._camera_state_transition( + list(actions), + event_id=event_id, + timestamp_ms=timestamp_ms, + ) + ) + return result + + def receive_camera_event_payload( + self, + payload: Any, + *, + event_id: int | None, + ) -> str: + if isinstance(payload, dict) and payload.get("mode") == "state": + transitions = self._camera_transitions_from_event_payload( + payload, + event_id=event_id, + ) + self.receive_camera_state_transitions(transitions) + return f"kind=camera_actions, mode=state, transitions={len(transitions)}" + + camera_actions = LingBotWorldRealtimeAdapter._validate_camera_actions(payload) + self.receive_camera_script(camera_actions, event_id=event_id) + return f"kind=camera_actions, mode=script, frames={len(camera_actions)}" + + def sample_prompt(self) -> str: + prompt = self.events.pop_latest("prompt") + if not isinstance(prompt, str): + raise ValueError("prompt event payload must be a string") + self.latest_sampled_event_id = self.events.last_sampled_seq_id("prompt") + return prompt + + def sample_camera_actions(self, chunk_size: int) -> list[list[str]] | None: + """samples a sequence of camera actions for the chunk with chunk_size frames + + Args: + chunk_size: number of frames + """ + if self.camera_script_queue: + return self._sample_camera_script(chunk_size) + action_list = self.camera_state.sample_chunk(chunk_size) + if action_list is None: + return None + self.latest_sampled_event_id = self.camera_state.latest_sampled_seq_id() + return [list(actions) for actions in action_list] + + def has_prompt(self) -> bool: + return self.events.has_events("prompt") + + +class LingBotWorldRealtimeAdapter(RealtimeModelAdapter): + name = "lingbot_world" + + def __init__(self): + self.output_adapter = RawRGBRealtimeOutputAdapter() + + def create_state(self) -> LingBotWorldRealtimeState: + return LingBotWorldRealtimeState() + + def _state(self, session: GenerateSession) -> LingBotWorldRealtimeState: + state = session.adapter_state + if not isinstance(state, LingBotWorldRealtimeState): + raise TypeError("LingBot realtime adapter state is not initialized") + return state + + async def on_init( + self, + session: GenerateSession, + request: RealtimeVideoGenerationsRequest, + ) -> None: + condition_inputs = request.condition_inputs or {} + camera_actions = condition_inputs.get("camera_actions") + if camera_actions is not None: + state = self._state(session) + state.receive_camera_script(self._validate_camera_actions(camera_actions)) + + if request.first_frame is None: + return + + server_args = get_global_server_args() + if server_args.input_save_path is not None: + uploads_dir = server_args.input_save_path + os.makedirs(uploads_dir, exist_ok=True) + else: + if session.input_temp_dir is None: + session.input_temp_dir = tempfile.mkdtemp(prefix="sglang_input_") + uploads_dir = session.input_temp_dir + + target_path = os.path.join(uploads_dir, f"{session.id}_first_frame") + image_path = await save_image_to_path(request.first_frame, target_path) + request.first_frame = image_path + + async def wait_for_next_chunk(self, session: GenerateSession) -> None: + del session + + @staticmethod + def _validate_camera_actions(payload: Any) -> list[list[str]]: + if not isinstance(payload, list): + raise ValueError("camera_actions event payload must be list[list[str]]") + normalized = [] + for frame_actions in payload: + if not isinstance(frame_actions, list): + raise ValueError("camera_actions event payload must be list[list[str]]") + normalized.append(list(frame_actions)) + return normalized + + def ingest_event( + self, + session: GenerateSession, + event: RealtimeEvent, + ) -> str: + state = self._state(session) + if event.kind == "camera_actions": + return state.receive_camera_event_payload( + event.payload, + event_id=event.event_id, + ) + elif event.kind == "prompt": + if not isinstance(event.payload, str) or not event.payload: + raise ValueError("prompt event payload must be a non-empty string") + state.receive_prompt(event.payload, event_id=event.event_id) + return f"kind=prompt, prompt_len={len(event.payload)}" + raise ValueError(f"unsupported event kind: {event.kind}") + + def _sample_chunk_inputs( + self, + session: GenerateSession, + chunk: RealtimeChunkContext, + chunk_size: int, + ) -> RealtimeChunkInputs: + """Samples user inputs (conditions) for the current RealtimeChunk from RealtimeStates""" + state = self._state(session) + request = session.request + if request is None: + raise ValueError("realtime request is not initialized") + + if chunk.index == 0: + prompt = request.prompt + elif state.has_prompt(): + prompt = state.sample_prompt() + request.prompt = prompt + else: + prompt = request.prompt + + condition_inputs = {} + camera_actions = state.sample_camera_actions(chunk_size) + if camera_actions is not None: + condition_inputs["camera_actions"] = camera_actions + return RealtimeChunkInputs(prompt=prompt, condition_inputs=condition_inputs) + + def _build_sampling_params( + self, + session: GenerateSession, + chunk: RealtimeChunkContext, + chunk_inputs: RealtimeChunkInputs, + chunk_size: int, + server_args: ServerArgs, + ): + request = session.request + if request is None: + raise ValueError("realtime request is not initialized") + + num_frames = self._condition_num_frames( + request=request, + server_args=server_args, + chunk_size=chunk_size, + ) + + return build_sampling_params( + chunk.request_id, + prompt=chunk_inputs.prompt, + size=request.size, + num_frames=num_frames, + fps=request.fps, + image_path=request.first_frame, + output_file_name=chunk.request_id, + save_output=False, + seed=request.seed, + generator_device=request.generator_device, + num_inference_steps=( + request.num_inference_steps + or LINGBOT_REALTIME_DEFAULT_NUM_INFERENCE_STEPS + ), + guidance_scale=request.guidance_scale, + guidance_scale_2=request.guidance_scale_2, + negative_prompt=request.negative_prompt, + enable_teacache=request.enable_teacache, + enable_frame_interpolation=request.enable_frame_interpolation, + frame_interpolation_exp=request.frame_interpolation_exp, + frame_interpolation_scale=request.frame_interpolation_scale, + frame_interpolation_model_path=request.frame_interpolation_model_path, + enable_upscaling=request.enable_upscaling, + upscaling_model_path=request.upscaling_model_path, + upscaling_scale=request.upscaling_scale, + diffusers_kwargs=request.diffusers_kwargs, + profile=request.profile, + num_profiled_timesteps=request.num_profiled_timesteps, + profile_all_stages=request.profile_all_stages, + perf_dump_path=request.perf_dump_path, + output_path=request.output_path, + output_compression=request.output_compression, + output_quality=request.output_quality, + condition_inputs=chunk_inputs.condition_inputs, + realtime_chunk_size=chunk_size, + ) + + @staticmethod + def _condition_num_frames( + *, + request: RealtimeVideoGenerationsRequest, + server_args: ServerArgs | None, + chunk_size: int, + ) -> int: + if server_args is None: + return int(request.num_frames or 0) + + # encode one extra blank condition chunk so repeat-last never reuses + # the first-frame image mask on later realtime chunks + temporal_ratio = int( + server_args.pipeline_config.vae_config.arch_config.temporal_compression_ratio + ) + required_latent_frames = chunk_size * LINGBOT_REALTIME_MIN_CONDITION_CHUNKS + required_num_frames = (required_latent_frames - 1) * temporal_ratio + 1 + return max(int(request.num_frames or 0), required_num_frames) + + def prepare_next_request( + self, + session: GenerateSession, + server_args: ServerArgs, + chunk: RealtimeChunkContext, + ) -> Req: + """build a new request for the next chunk""" + pipeline_config = server_args.pipeline_config + chunk_size = int(pipeline_config.dit_config.arch_config.num_frames_per_block) + chunk_inputs = self._sample_chunk_inputs(session, chunk, chunk_size) + sampling_params = self._build_sampling_params( + session, + chunk, + chunk_inputs, + chunk_size, + server_args, + ) + batch = prepare_request( + server_args=server_args, + sampling_params=sampling_params, + ) + batch.session = session.realtime_session + batch.realtime_session_id = session.id + batch.return_raw_frames = True + batch.block_idx = chunk.index + batch.realtime_event_id = self._state(session).latest_sampled_event_id + if session.request is not None: + batch.realtime_output_format = session.request.realtime_output_format + batch.realtime_causal_sink_size = session.request.realtime_causal_sink_size + batch.realtime_causal_kv_cache_num_frames = ( + session.request.realtime_causal_kv_cache_num_frames + ) + return batch + + async def send_output( + self, + ws: WebSocket, + session: GenerateSession, + result: OutputBatch, + batch: Req, + ) -> RealtimeFrameSendStats: + return await self.output_adapter.send(ws, session, result, batch) + + def on_chunk_complete(self, session: GenerateSession, result: OutputBatch) -> None: + del result + session.generate_chunk_completed() + + def dispose(self, session: GenerateSession) -> None: + state = session.adapter_state + if isinstance(state, LingBotWorldRealtimeState): + state.clear() + self.output_adapter.reset() diff --git a/python/sglang/multimodal_gen/runtime/entrypoints/openai/realtime/generate_session.py b/python/sglang/multimodal_gen/runtime/entrypoints/openai/realtime/generate_session.py new file mode 100644 index 000000000..af362fd4d --- /dev/null +++ b/python/sglang/multimodal_gen/runtime/entrypoints/openai/realtime/generate_session.py @@ -0,0 +1,78 @@ +# SPDX-License-Identifier: Apache-2.0 + +from __future__ import annotations + +from dataclasses import dataclass +from typing import TYPE_CHECKING, Any +from uuid import uuid4 + +from sglang.multimodal_gen.runtime.entrypoints.openai.protocol import ( + RealtimeVideoGenerationsRequest, +) +from sglang.multimodal_gen.runtime.realtime.session import ( + RealtimeSession, +) + +if TYPE_CHECKING: + from sglang.multimodal_gen.runtime.entrypoints.openai.realtime.realtime_adapter import ( + RealtimeModelAdapter, + ) + + +@dataclass(frozen=True, slots=True) +class RealtimeChunkContext: + session_id: str + index: int + request_id: str + + +class GenerateSession: + def __init__(self): + self.id = uuid4().hex + self.request: RealtimeVideoGenerationsRequest | None = None + self.input_temp_dir: str | None = None + self.generate_chunk_cnt = 0 + self.current_chunk: RealtimeChunkContext | None = None + self.realtime_session = RealtimeSession() + self.adapter: RealtimeModelAdapter | None = None + self.adapter_state: Any = None + + def set_adapter(self, adapter: RealtimeModelAdapter): + self.adapter = adapter + self.adapter_state = adapter.create_state() + + def set_request(self, request: RealtimeVideoGenerationsRequest): + self.request = request + + def dispose(self): + if self.adapter is not None: + self.adapter.dispose(self) + self.request = None + self.input_temp_dir = None + self.generate_chunk_cnt = 0 + self.current_chunk = None + self.adapter = None + self.adapter_state = None + self.realtime_session.dispose() + + def new_chunk(self) -> RealtimeChunkContext: + if self.current_chunk is not None: + raise RuntimeError("previous realtime chunk is still active") + chunk = RealtimeChunkContext( + session_id=self.id, + index=self.generate_chunk_cnt, + request_id=f"{self.id}_{uuid4().hex}", + ) + self.current_chunk = chunk + return chunk + + def generate_chunk_completed(self): + self.generate_chunk_cnt += 1 + self.current_chunk = None + + def reached_max_chunks(self) -> bool: + return ( + self.request is not None + and self.request.max_chunks is not None + and self.generate_chunk_cnt >= self.request.max_chunks + ) diff --git a/python/sglang/multimodal_gen/runtime/entrypoints/openai/realtime/realtime_adapter.py b/python/sglang/multimodal_gen/runtime/entrypoints/openai/realtime/realtime_adapter.py new file mode 100644 index 000000000..6a95a9116 --- /dev/null +++ b/python/sglang/multimodal_gen/runtime/entrypoints/openai/realtime/realtime_adapter.py @@ -0,0 +1,74 @@ +# SPDX-License-Identifier: Apache-2.0 + +from __future__ import annotations + +from dataclasses import dataclass, field +from typing import TYPE_CHECKING, Any, Protocol + +from fastapi import WebSocket + +from sglang.multimodal_gen.runtime.entrypoints.openai.protocol import ( + RealtimeEvent, + RealtimeVideoGenerationsRequest, +) +from sglang.multimodal_gen.runtime.entrypoints.openai.realtime.realtime_output_adapter import ( + RealtimeFrameSendStats, +) + +if TYPE_CHECKING: + from sglang.multimodal_gen.runtime.entrypoints.openai.realtime.generate_session import ( + GenerateSession, + RealtimeChunkContext, + ) + from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import ( + OutputBatch, + Req, + ) + from sglang.multimodal_gen.runtime.server_args import ServerArgs + + +@dataclass(slots=True) +class RealtimeChunkInputs: + prompt: str + condition_inputs: dict[str, Any] = field(default_factory=dict) + + +class RealtimeModelAdapter(Protocol): + name: str + + def create_state(self) -> Any: ... + + async def on_init( + self, + session: GenerateSession, + request: RealtimeVideoGenerationsRequest, + ) -> None: ... + + async def wait_for_next_chunk(self, session: GenerateSession) -> None: ... + + def ingest_event( + self, + session: GenerateSession, + event: RealtimeEvent, + ) -> str: ... + + def prepare_next_request( + self, + session: GenerateSession, + server_args: ServerArgs, + chunk: RealtimeChunkContext, + ) -> Req: ... + + async def send_output( + self, + ws: WebSocket, + session: GenerateSession, + result: OutputBatch, + batch: Req, + ) -> RealtimeFrameSendStats: ... + + def on_chunk_complete( + self, session: GenerateSession, result: OutputBatch + ) -> None: ... + + def dispose(self, session: GenerateSession) -> None: ... diff --git a/python/sglang/multimodal_gen/runtime/entrypoints/openai/realtime/realtime_output_adapter.py b/python/sglang/multimodal_gen/runtime/entrypoints/openai/realtime/realtime_output_adapter.py new file mode 100644 index 000000000..35eb10c34 --- /dev/null +++ b/python/sglang/multimodal_gen/runtime/entrypoints/openai/realtime/realtime_output_adapter.py @@ -0,0 +1,496 @@ +# SPDX-License-Identifier: Apache-2.0 + +from __future__ import annotations + +import asyncio +import io +from dataclasses import dataclass +from typing import TYPE_CHECKING, TypedDict + +import msgspec.msgpack +from fastapi import WebSocket +from PIL import Image + +from sglang.multimodal_gen.runtime.entrypoints.openai.realtime.timing import ( + RealtimeStageTimer, +) +from sglang.multimodal_gen.runtime.utils.realtime_video import ( + JPEG_FRAME_CONTENT_TYPE, + RAW_RGB_CHANNELS, + RAW_RGB_CONTENT_TYPE, + RAW_RGB_DELTA_GZIP_CONTENT_TYPE, + WEBP_FRAME_CONTENT_TYPE, + build_delta_gzip_raw_rgb_payload, +) + +if TYPE_CHECKING: + from sglang.multimodal_gen.runtime.entrypoints.openai.realtime.generate_session import ( + GenerateSession, + ) + from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import ( + OutputBatch, + Req, + ) + + +class RealtimeFrameBatchHeader(TypedDict, total=False): + type: str + request_id: str + chunk_index: int + content_type: str + num_frames: int + total_size: int + format: str + width: int + height: int + channels: int + bytes_per_frame: int + raw_size: int + encoding: str + delta_reference: str + event_id: int + frame_batch_index: int + num_frame_batches: int + is_final_frame_batch: bool + + +class RealtimeFrameBatchMessage(RealtimeFrameBatchHeader, total=False): + payload: bytes + + +class RealtimeFrameSendStats(TypedDict): + header_pack_ms: float + header_write_ms: float + raw_payload_build_ms: float + raw_write_ms: float + ws_write_ms: float + raw_bytes: int + ws_payload_bytes: int + num_frames: int + num_batches: int + frame_shape: tuple[int, int, int] | None + content_type: str + + +def empty_frame_send_stats(content_type: str = "") -> RealtimeFrameSendStats: + return { + "header_pack_ms": 0.0, + "header_write_ms": 0.0, + "raw_payload_build_ms": 0.0, + "raw_write_ms": 0.0, + "ws_write_ms": 0.0, + "raw_bytes": 0, + "ws_payload_bytes": 0, + "num_frames": 0, + "num_batches": 0, + "frame_shape": None, + "content_type": content_type, + } + + +def _raw_rgb_frame_metadata(batch: Req) -> dict[str, int | str]: + frame_width = batch.width + frame_height = batch.height + if frame_width is None or frame_height is None: + return {} + + frame_width = int(frame_width) + frame_height = int(frame_height) + if batch.enable_upscaling: + upscaling_scale = int(batch.upscaling_scale or 1) + frame_width *= upscaling_scale + frame_height *= upscaling_scale + + return { + "format": "rgb24", + "width": frame_width, + "height": frame_height, + "channels": RAW_RGB_CHANNELS, + "bytes_per_frame": frame_width * frame_height * RAW_RGB_CHANNELS, + } + + +def _frame_shape_from_metadata( + metadata: dict[str, int | str] | None, +) -> tuple[int, int, int] | None: + if not metadata: + return None + return ( + int(metadata["height"]), + int(metadata["width"]), + int(metadata["channels"]), + ) + + +RAW_RGB_FRAMES_PER_WS_MESSAGE = 16 +FRAME_BATCH_PACK_OFFLOAD_BYTES = 64 * 1024 +WEBP_DEFAULT_QUALITY = 90 +JPEG_DEFAULT_QUALITY = 95 +JPEG_SUBSAMPLING = 0 +RAW_LOSSLESS_OUTPUT_FORMAT = "raw" +ENCODED_PREVIEW_FORMATS = {"webp", "jpeg"} + + +@dataclass(frozen=True) +class _TransportPayload: + content_type: str + payload: bytes + metadata: dict[str, int | str | bool] + last_raw_rgb_frame: bytes | None = None + last_event_id: int | None = None + + +def _split_frame_batch(frames: list[bytes]) -> list[list[bytes]]: + if not frames: + return [frames] + return [ + frames[i : i + RAW_RGB_FRAMES_PER_WS_MESSAGE] + for i in range(0, len(frames), RAW_RGB_FRAMES_PER_WS_MESSAGE) + ] + + +def _encode_rgb_frame_to_webp( + frame: bytes, + *, + width: int, + height: int, + quality: int, +) -> bytes: + buffer = io.BytesIO() + Image.frombytes("RGB", (width, height), frame).save( + buffer, + format="WEBP", + quality=quality, + method=0, + ) + return buffer.getvalue() + + +def _encode_rgb_frame_to_jpeg( + frame: bytes, + *, + width: int, + height: int, + quality: int, +) -> bytes: + buffer = io.BytesIO() + Image.frombytes("RGB", (width, height), frame).save( + buffer, + format="JPEG", + quality=quality, + subsampling=JPEG_SUBSAMPLING, + ) + return buffer.getvalue() + + +def _pack_frame_batch_message( + header: RealtimeFrameBatchHeader, + payload: bytes, +) -> bytes: + message: RealtimeFrameBatchMessage = { + **header, + "type": "frame_batch", + "payload": payload, + } + return msgspec.msgpack.encode(message) + + +def _build_transport_payload( + transport_frames: list[bytes], + *, + content_type: str, + metadata: dict[str, int | str], + output_format: str | None, + transport_quality: int | None, + reference_frame: bytes | None, + event_id: int | None, +) -> _TransportPayload: + payload_content_type = content_type + payload_metadata: dict[str, int | str | bool] = {} + raw_payload = b"" + + if ( + output_format in ENCODED_PREVIEW_FORMATS + and content_type == RAW_RGB_CONTENT_TYPE + and transport_frames + ): + if output_format == "webp": + raw_payload = _encode_rgb_frame_to_webp( + transport_frames[0], + width=int(metadata["width"]), + height=int(metadata["height"]), + quality=int(transport_quality or WEBP_DEFAULT_QUALITY), + ) + payload_content_type = WEBP_FRAME_CONTENT_TYPE + else: + raw_payload = _encode_rgb_frame_to_jpeg( + transport_frames[0], + width=int(metadata["width"]), + height=int(metadata["height"]), + quality=int(transport_quality or JPEG_DEFAULT_QUALITY), + ) + payload_content_type = JPEG_FRAME_CONTENT_TYPE + payload_metadata = { + "format": output_format, + "encoding": output_format, + } + elif ( + output_format == RAW_LOSSLESS_OUTPUT_FORMAT + and content_type == RAW_RGB_CONTENT_TYPE + and transport_frames + ): + raw_payload = b"".join(transport_frames) + payload_metadata = { + "raw_size": len(raw_payload), + "encoding": RAW_LOSSLESS_OUTPUT_FORMAT, + } + elif content_type == RAW_RGB_CONTENT_TYPE and transport_frames: + raw_payload = build_delta_gzip_raw_rgb_payload( + transport_frames, + reference_frame=reference_frame, + ) + payload_content_type = RAW_RGB_DELTA_GZIP_CONTENT_TYPE + payload_metadata = { + "raw_size": sum(len(frame) for frame in transport_frames), + "encoding": "delta-gzip", + } + if reference_frame is not None: + payload_metadata["delta_reference"] = "previous-frame" + return _TransportPayload( + content_type=payload_content_type, + payload=raw_payload, + metadata=payload_metadata, + last_raw_rgb_frame=transport_frames[-1], + last_event_id=event_id, + ) + else: + raw_payload = b"".join(transport_frames) + + return _TransportPayload( + content_type=payload_content_type, + payload=raw_payload, + metadata=payload_metadata, + ) + + +def _should_build_payload_off_loop( + *, + content_type: str, + output_format: str | None, + transport_frames: list[bytes], +) -> bool: + if content_type != RAW_RGB_CONTENT_TYPE or not transport_frames: + return False + return output_format in ENCODED_PREVIEW_FORMATS or output_format is None + + +def _is_encoded_preview_transport( + *, + content_type: str, + output_format: str | None, +) -> bool: + return ( + output_format in ENCODED_PREVIEW_FORMATS + and content_type == RAW_RGB_CONTENT_TYPE + ) + + +async def _build_encoded_preview_payloads( + split_batches: list[list[bytes]], + *, + content_type: str, + metadata: dict[str, int | str], + output_format: str, + transport_quality: int | None, + event_id: int | None, +) -> list[_TransportPayload]: + return list( + await asyncio.gather( + *( + asyncio.to_thread( + _build_transport_payload, + transport_frames, + content_type=content_type, + metadata=metadata, + output_format=output_format, + transport_quality=transport_quality, + reference_frame=None, + event_id=event_id, + ) + for transport_frames in split_batches + ) + ) + ) + + +class RawRGBRealtimeOutputAdapter: + """send raw RGB over WebSocket using lossless transport compression""" + + def __init__(self) -> None: + self._last_raw_rgb_frame: bytes | None = None + self._last_event_id: int | None = None + + def reset(self) -> None: + self._last_raw_rgb_frame = None + self._last_event_id = None + + async def send( + self, + ws: WebSocket, + session: GenerateSession, + result: OutputBatch, + batch: Req, + ) -> RealtimeFrameSendStats: + """send frames through ws""" + content_type = result.raw_frame_content_type + if result.raw_frame_batches is None: + return empty_frame_send_stats(content_type) + if batch.block_idx == 0: + self.reset() + + frame_metadata = ( + result.raw_frame_metadata or _raw_rgb_frame_metadata(batch) + if content_type == RAW_RGB_CONTENT_TYPE + else {} + ) + output_format = getattr(batch, "realtime_output_format", None) + stats = await self._send_frame_batches( + ws, + result.raw_frame_batches, + content_type=content_type, + chunk_index_start=batch.block_idx, + request_id=batch.request_id, + event_id=getattr(batch, "realtime_event_id", None), + frame_metadata=frame_metadata, + output_format=output_format, + transport_quality=getattr(batch, "output_compression", None), + ) + stats["frame_shape"] = _frame_shape_from_metadata(frame_metadata) + return stats + + async def _send_frame_batches( + self, + ws: WebSocket, + frame_batches: list[list[bytes]], + *, + content_type: str, + chunk_index_start: int, + request_id: str, + event_id: int | None = None, + frame_metadata: dict[str, int | str] | None = None, + output_format: str | None = None, + transport_quality: int | None = None, + ) -> RealtimeFrameSendStats: + chunk_index = chunk_index_start + metadata = frame_metadata or {} + stats = empty_frame_send_stats(content_type) + for frames in frame_batches: + split_batches = ( + [[frame] for frame in frames] + if _is_encoded_preview_transport( + content_type=content_type, + output_format=output_format, + ) + else ( + _split_frame_batch(frames) + if content_type == RAW_RGB_CONTENT_TYPE + else [frames] + ) + ) + num_frame_batches = len(split_batches) + encoded_preview_payloads: list[_TransportPayload] | None = None + if _is_encoded_preview_transport( + content_type=content_type, + output_format=output_format, + ): + timer = RealtimeStageTimer() + encoded_preview_payloads = await _build_encoded_preview_payloads( + split_batches, + content_type=content_type, + metadata=metadata, + output_format=output_format, + transport_quality=transport_quality, + event_id=event_id, + ) + stats["raw_payload_build_ms"] += timer.mark_ms() + for frame_batch_index, transport_frames in enumerate(split_batches): + timer = RealtimeStageTimer() + transport_metadata = metadata + if encoded_preview_payloads is not None: + transport_payload = encoded_preview_payloads[frame_batch_index] + else: + reference_frame = self._last_raw_rgb_frame + if event_id != self._last_event_id: + reference_frame = None + if _should_build_payload_off_loop( + content_type=content_type, + output_format=output_format, + transport_frames=transport_frames, + ): + transport_payload = await asyncio.to_thread( + _build_transport_payload, + transport_frames, + content_type=content_type, + metadata=metadata, + output_format=output_format, + transport_quality=transport_quality, + reference_frame=reference_frame, + event_id=event_id, + ) + else: + transport_payload = _build_transport_payload( + transport_frames, + content_type=content_type, + metadata=metadata, + output_format=output_format, + transport_quality=transport_quality, + reference_frame=reference_frame, + event_id=event_id, + ) + if transport_payload.last_raw_rgb_frame is not None: + self._last_raw_rgb_frame = transport_payload.last_raw_rgb_frame + self._last_event_id = transport_payload.last_event_id + stats["raw_payload_build_ms"] += timer.mark_ms() + + header: RealtimeFrameBatchHeader = { + "type": "frame_batch_header", + "request_id": request_id, + "chunk_index": chunk_index, + "content_type": transport_payload.content_type, + "num_frames": len(transport_frames), + "total_size": len(transport_payload.payload), + "frame_batch_index": frame_batch_index, + "num_frame_batches": num_frame_batches, + "is_final_frame_batch": frame_batch_index == num_frame_batches - 1, + } + if event_id is not None: + header["event_id"] = event_id + header.update(transport_metadata) + header.update(transport_payload.metadata) + + if len(transport_payload.payload) >= FRAME_BATCH_PACK_OFFLOAD_BYTES: + message_payload = await asyncio.to_thread( + _pack_frame_batch_message, + header, + transport_payload.payload, + ) + else: + message_payload = _pack_frame_batch_message( + header, + transport_payload.payload, + ) + stats["header_pack_ms"] += timer.mark_ms() + + stats["header_write_ms"] += timer.mark_ms() + await ws.send_bytes(message_payload) + stats["raw_write_ms"] += timer.mark_ms() + + stats["raw_bytes"] += sum(len(frame) for frame in transport_frames) + stats["ws_payload_bytes"] += len(message_payload) + stats["num_frames"] += len(transport_frames) + stats["num_batches"] += 1 + stats["content_type"] = transport_payload.content_type + chunk_index += 1 + + stats["ws_write_ms"] = stats["header_write_ms"] + stats["raw_write_ms"] + return stats diff --git a/python/sglang/multimodal_gen/runtime/entrypoints/openai/realtime/realtime_video_api.py b/python/sglang/multimodal_gen/runtime/entrypoints/openai/realtime/realtime_video_api.py new file mode 100644 index 000000000..39152a36c --- /dev/null +++ b/python/sglang/multimodal_gen/runtime/entrypoints/openai/realtime/realtime_video_api.py @@ -0,0 +1,427 @@ +# SPDX-License-Identifier: Apache-2.0 + +import asyncio +import shutil +import time +from typing import TYPE_CHECKING + +import msgspec.msgpack +from fastapi import APIRouter, WebSocket, WebSocketDisconnect + +from sglang.multimodal_gen.runtime.entrypoints.openai.protocol import ( + RealtimeEvent, + RealtimeVideoGenerationsRequest, +) +from sglang.multimodal_gen.runtime.entrypoints.openai.realtime.generate_session import ( + GenerateSession, + RealtimeChunkContext, +) +from sglang.multimodal_gen.runtime.entrypoints.openai.realtime.realtime_output_adapter import ( + RealtimeFrameSendStats, +) +from sglang.multimodal_gen.runtime.entrypoints.openai.realtime.registry import ( + get_realtime_model_adapter, +) +from sglang.multimodal_gen.runtime.entrypoints.openai.realtime.timing import ( + RealtimeStageTimer, +) +from sglang.multimodal_gen.runtime.entrypoints.openai.utils import ( + process_generation_batch, +) +from sglang.multimodal_gen.runtime.entrypoints.utils import ( + ReleaseRealtimeSessionReq, +) +from sglang.multimodal_gen.runtime.scheduler_client import async_scheduler_client +from sglang.multimodal_gen.runtime.server_args import get_global_server_args +from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger + +if TYPE_CHECKING: + from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import Req + +logger = init_logger(__name__) +router = APIRouter(prefix="/v1/realtime_video", tags=["realtime"]) +_ACTIVE_SESSION_IDS: set[str] = set() +_ACTIVE_SESSION_WAIT_SECONDS = 15.0 +_ACTIVE_SESSION_WAIT_INTERVAL_SECONDS = 0.1 + + +def _transport_ms(value: float) -> int: + return max(0, int(value + 0.5)) + + +async def _wait_for_active_session_slot( + *, + timeout_s: float = _ACTIVE_SESSION_WAIT_SECONDS, + interval_s: float = _ACTIVE_SESSION_WAIT_INTERVAL_SECONDS, +) -> bool: + deadline = time.monotonic() + timeout_s + while _ACTIVE_SESSION_IDS and time.monotonic() < deadline: + await asyncio.sleep(interval_s) + return not _ACTIVE_SESSION_IDS + + +def _log_realtime_chunk_timing( + session: GenerateSession, + chunk: RealtimeChunkContext, + batch: "Req", + request_prepare_ms: float, + scheduler_forward_ms: float, + chunk_total_ms: float, + send_stats: RealtimeFrameSendStats, +) -> None: + logger.info( + "realtime chunk timing: session_id=%s request_id=%s " + "chunk_idx=%s event_id=%s condition_kinds=%s " + "request_prepare=%.2fms scheduler_forward=%.2fms " + "header_pack=%.2fms " + "header_write=%.2fms raw_payload_build=%.2fms raw_write=%.2fms " + "ws_write=%.2fms chunk_total=%.2fms batches=%d frames=%d " + "frame_shape=%s raw_bytes=%d ws_payload_bytes=%d content_type=%s", + session.id, + chunk.request_id, + batch.block_idx, + getattr(batch, "realtime_event_id", None), + sorted(batch.condition_inputs) if batch.condition_inputs else [], + request_prepare_ms, + scheduler_forward_ms, + send_stats["header_pack_ms"], + send_stats["header_write_ms"], + send_stats["raw_payload_build_ms"], + send_stats["raw_write_ms"], + send_stats["ws_write_ms"], + chunk_total_ms, + send_stats["num_batches"], + send_stats["num_frames"], + send_stats["frame_shape"], + send_stats["raw_bytes"], + send_stats["ws_payload_bytes"], + send_stats["content_type"], + ) + + +async def _send_realtime_chunk_stats( + ws: WebSocket, + session: GenerateSession, + chunk: RealtimeChunkContext, + batch: "Req", + request_prepare_ms: float, + scheduler_forward_ms: float, + chunk_total_ms: float, + send_stats: RealtimeFrameSendStats, +) -> None: + await ws.send_bytes( + msgspec.msgpack.encode( + { + "type": "chunk_stats", + "session_id": session.id, + "request_id": chunk.request_id, + "chunk_index": batch.block_idx, + "event_id": getattr(batch, "realtime_event_id", None), + "request_prepare_ms": _transport_ms(request_prepare_ms), + "scheduler_forward_ms": _transport_ms(scheduler_forward_ms), + "header_write_ms": _transport_ms(send_stats["header_write_ms"]), + "raw_payload_build_ms": _transport_ms( + send_stats["raw_payload_build_ms"] + ), + "raw_write_ms": _transport_ms(send_stats["raw_write_ms"]), + "ws_write_ms": _transport_ms(send_stats["ws_write_ms"]), + "chunk_total_ms": _transport_ms(chunk_total_ms), + "num_batches": send_stats["num_batches"], + "num_frames": send_stats["num_frames"], + "raw_bytes": send_stats["raw_bytes"], + "ws_payload_bytes": send_stats["ws_payload_bytes"], + "content_type": send_stats["content_type"], + } + ) + ) + + +async def _generate_loop(ws: WebSocket, session: GenerateSession): + adapter = session.adapter + if adapter is None: + raise ValueError("realtime adapter is not initialized") + + pending_send_task = None + while not session.reached_max_chunks(): + try: + if pending_send_task is not None and pending_send_task.done(): + await pending_send_task + pending_send_task = None + + # send to scheduler and generate video chunk + server_args = get_global_server_args() + + await adapter.wait_for_next_chunk(session) + + timer = RealtimeStageTimer() + chunk_started = time.perf_counter() + + chunk = session.new_chunk() + batch = adapter.prepare_next_request( + session, + server_args, + chunk, + ) + if batch.condition_inputs: + logger.debug( + "consume realtime conditions, session_id=%s, block_idx=%s, kinds=%s", + session.id, + batch.block_idx, + sorted(batch.condition_inputs), + ) + request_prepare_ms = timer.mark_ms() + + _, result = await process_generation_batch(async_scheduler_client, batch) + scheduler_forward_ms = timer.mark_ms() + + # finish + adapter.on_chunk_complete(session, result) + if pending_send_task is not None: + await pending_send_task + pending_send_task = asyncio.create_task( + _send_output_and_log( + ws, + session, + chunk, + batch, + result, + request_prepare_ms, + scheduler_forward_ms, + chunk_started, + ) + ) + + except asyncio.CancelledError: + if pending_send_task is not None: + pending_send_task.cancel() + await _await_realtime_task(pending_send_task) + logger.info("generation completed, session_id=%s", session.id) + break + except WebSocketDisconnect: + if pending_send_task is not None: + pending_send_task.cancel() + await _await_realtime_task(pending_send_task) + logger.info( + "client disconnected during generation, session_id=%s", session.id + ) + break + except Exception as e: + if pending_send_task is not None: + pending_send_task.cancel() + await _await_realtime_task(pending_send_task) + err_msg = str(e).splitlines()[0] + logger.error("error during generate loop: %s", err_msg) + try: + await write_error_msg(f"error during generate loop: {err_msg}", ws) + except Exception as send_error: + logger.error( + "error during sending complete msg: %s", + send_error, + ) + break + else: + if pending_send_task is not None: + await pending_send_task + logger.info( + "generation reached max chunks, session_id=%s, max_chunks=%s", + session.id, + session.request.max_chunks if session.request is not None else None, + ) + + +async def _send_output_and_log( + ws: WebSocket, + session: GenerateSession, + chunk: RealtimeChunkContext, + batch: "Req", + result, + request_prepare_ms: float, + scheduler_forward_ms: float, + chunk_started: float, +) -> RealtimeFrameSendStats: + if session.adapter is None: + raise ValueError("realtime adapter is not initialized") + send_stats = await session.adapter.send_output( + ws, + session, + result, + batch, + ) + chunk_total_ms = (time.perf_counter() - chunk_started) * 1000 + _log_realtime_chunk_timing( + session, + chunk, + batch, + request_prepare_ms, + scheduler_forward_ms, + chunk_total_ms, + send_stats, + ) + await _send_realtime_chunk_stats( + ws, + session, + chunk, + batch, + request_prepare_ms, + scheduler_forward_ms, + chunk_total_ms, + send_stats, + ) + return send_stats + + +async def _await_realtime_task(task: asyncio.Task | None) -> None: + if task is None: + return + try: + await task + except (asyncio.CancelledError, WebSocketDisconnect): + pass + except Exception as e: + logger.debug("realtime task exited with error: %s", e) + + +async def _listen_events(ws: WebSocket, session: GenerateSession): + """listen for user events: usually condition inputs""" + async for message in ws.iter_bytes(): + data = None + try: + data = msgspec.msgpack.decode(message) + if not isinstance(data, dict): + raise ValueError("realtime event must be a map") + realtime_event = RealtimeEvent.model_validate(data) + if session.adapter is None: + raise ValueError("realtime adapter is not initialized") + event_log = session.adapter.ingest_event(session, realtime_event) + logger.info( + "receive realtime event, session_id=%s, event_id=%s, %s", + session.id, + realtime_event.event_id, + event_log, + ) + except Exception as e: + event_kind = data.get("kind") if isinstance(data, dict) else None + logger.warning("invalid event, kind=%s, error=%s", event_kind, e) + await write_error_msg("invalid event", ws) + continue + + +async def _listen_generate_request(ws: WebSocket, session: GenerateSession): + while True: + try: + data = msgspec.msgpack.decode(await ws.receive_bytes()) + if not isinstance(data, dict): + raise ValueError("generate request must be a map") + + realtime_req = RealtimeVideoGenerationsRequest.model_validate(data) + adapter = get_realtime_model_adapter(get_global_server_args()) + session.set_adapter(adapter) + await adapter.on_init(session, realtime_req) + + # Keep session state update atomic with validated request. + session.set_request(realtime_req) + break + except WebSocketDisconnect: + raise + except Exception as e: + logger.warning( + "invalid generate request, session_id=%s, error=%s", + session.id, + e, + ) + await write_error_msg("invalid generate request", ws) + continue + + +async def _cleanup_realtime_session( + session: GenerateSession, + generate_task: asyncio.Task | None, + listen_task: asyncio.Task | None, +) -> None: + logger.info("terminating session, session_id=%s", session.id) + for task in (generate_task, listen_task): + if task and not task.done(): + task.cancel() + for task in (generate_task, listen_task): + if task is None: + continue + await _await_realtime_task(task) + try: + await async_scheduler_client.forward( + ReleaseRealtimeSessionReq(session_id=session.id) + ) + except Exception as e: + logger.warning( + "failed to release realtime session on scheduler, session_id=%s, error=%s", + session.id, + e, + ) + if session.input_temp_dir is not None: + shutil.rmtree(session.input_temp_dir, ignore_errors=True) + session.dispose() + + +async def _close_realtime_websocket( + websocket: WebSocket, + *, + code: int, + reason: str, +) -> None: + try: + await websocket.close(code=code, reason=reason) + except (RuntimeError, WebSocketDisconnect): + pass + + +@router.websocket("/generate") +async def generate(websocket: WebSocket): + """endpoint for creating a new realtime session""" + await websocket.accept() + if _ACTIVE_SESSION_IDS and not await _wait_for_active_session_slot(): + logger.warning( + "reject realtime session because another session is active: %s", + sorted(_ACTIVE_SESSION_IDS), + ) + try: + await write_error_msg( + "another realtime session is already active", websocket + ) + finally: + await websocket.close(code=1008) + return + + session = GenerateSession() + _ACTIVE_SESSION_IDS.add(session.id) + generate_task = None + listen_task = None + try: + # receive new generate request + await _listen_generate_request(websocket, session) + + # continuously generate video chunk + generate_task = asyncio.create_task(_generate_loop(websocket, session)) + # continuously listen for user events + listen_task = asyncio.create_task(_listen_events(websocket, session)) + + wait_tasks = [generate_task, listen_task] + await asyncio.wait(wait_tasks, return_when=asyncio.FIRST_COMPLETED) + if generate_task.done() and session.reached_max_chunks(): + await _close_realtime_websocket( + websocket, + code=1000, + reason="generation complete", + ) + + except WebSocketDisconnect: + logger.info("client disconnected, session_id=%s", session.id) + finally: + try: + await _cleanup_realtime_session(session, generate_task, listen_task) + finally: + _ACTIVE_SESSION_IDS.discard(session.id) + + +async def write_error_msg(error_msg: str, websocket: WebSocket): + await websocket.send_bytes( + msgspec.msgpack.encode({"type": "error", "content": error_msg}) + ) diff --git a/python/sglang/multimodal_gen/runtime/entrypoints/openai/realtime/registry.py b/python/sglang/multimodal_gen/runtime/entrypoints/openai/realtime/registry.py new file mode 100644 index 000000000..44181d52e --- /dev/null +++ b/python/sglang/multimodal_gen/runtime/entrypoints/openai/realtime/registry.py @@ -0,0 +1,59 @@ +# SPDX-License-Identifier: Apache-2.0 + +from __future__ import annotations + +from typing import TYPE_CHECKING + +from sglang.multimodal_gen.runtime.entrypoints.openai.realtime.realtime_adapter import ( + RealtimeModelAdapter, +) + +if TYPE_CHECKING: + from sglang.multimodal_gen.runtime.server_args import ServerArgs + + +_REALTIME_ADAPTER_REGISTRY: dict[type, type[RealtimeModelAdapter]] = {} +_BUILTIN_ADAPTERS_REGISTERED = False + + +def register_realtime_model_adapter( + pipeline_config_cls: type, + adapter_cls: type[RealtimeModelAdapter], +) -> None: + _REALTIME_ADAPTER_REGISTRY[pipeline_config_cls] = adapter_cls + + +def _register_builtin_realtime_model_adapters() -> None: + global _BUILTIN_ADAPTERS_REGISTERED + if _BUILTIN_ADAPTERS_REGISTERED: + return + + from sglang.multimodal_gen.configs.pipeline_configs.lingbot_world import ( + LingBotWorldCausalDMDConfig, + ) + from sglang.multimodal_gen.runtime.entrypoints.openai.realtime.adapters.lingbot_world_realtime_adapter import ( + LingBotWorldRealtimeAdapter, + ) + + register_realtime_model_adapter( + LingBotWorldCausalDMDConfig, + LingBotWorldRealtimeAdapter, + ) + _BUILTIN_ADAPTERS_REGISTERED = True + + +def get_realtime_model_adapter( + server_args: ServerArgs, +) -> RealtimeModelAdapter: + _register_builtin_realtime_model_adapters() + + pipeline_config = server_args.pipeline_config + for config_cls in type(pipeline_config).__mro__: + adapter_cls = _REALTIME_ADAPTER_REGISTRY.get(config_cls) + if adapter_cls is not None: + return adapter_cls() + + raise ValueError( + "Realtime video is not supported for pipeline config " + f"{type(pipeline_config).__name__}; no realtime adapter is registered." + ) diff --git a/python/sglang/multimodal_gen/runtime/entrypoints/openai/realtime/timing.py b/python/sglang/multimodal_gen/runtime/entrypoints/openai/realtime/timing.py new file mode 100644 index 000000000..f2e975c4d --- /dev/null +++ b/python/sglang/multimodal_gen/runtime/entrypoints/openai/realtime/timing.py @@ -0,0 +1,21 @@ +# SPDX-License-Identifier: Apache-2.0 + +import time + + +class RealtimeStageTimer: + __slots__ = ("_last", "_start") + + def __init__(self): + now = time.perf_counter() + self._start = now + self._last = now + + def mark_ms(self) -> float: + now = time.perf_counter() + elapsed_ms = (now - self._last) * 1000.0 + self._last = now + return elapsed_ms + + def total_ms(self) -> float: + return (time.perf_counter() - self._start) * 1000.0 diff --git a/python/sglang/multimodal_gen/runtime/entrypoints/openai/utils.py b/python/sglang/multimodal_gen/runtime/entrypoints/openai/utils.py index ccca4d3e4..f304f7ed5 100644 --- a/python/sglang/multimodal_gen/runtime/entrypoints/openai/utils.py +++ b/python/sglang/multimodal_gen/runtime/entrypoints/openai/utils.py @@ -1,6 +1,7 @@ # Copied and adapted from: https://github.com/hao-ai-lab/FastVideo import asyncio import base64 +import inspect import json import os import re @@ -165,7 +166,7 @@ def build_sampling_params(request_id: str, **kwargs) -> SamplingParams: async def save_image_to_path( - image: Union[UploadFile, str], + image: Union[UploadFile, bytes, str], target_path: str, *, prefer_remote_source: bool = False, @@ -179,9 +180,27 @@ async def save_image_to_path( # Helpers -async def _save_upload_to_path(upload: UploadFile, target_path: str) -> str: +async def _save_upload_to_path( + upload: Union[UploadFile, bytes], target_path: str +) -> str: os.makedirs(os.path.dirname(target_path), exist_ok=True) - content = await upload.read() + if isinstance(upload, bytes): + content = upload + elif isinstance(upload, (bytearray, memoryview)): + content = bytes(upload) + else: + read = getattr(upload, "read", None) + if not callable(read): + raise TypeError(f"Unsupported image upload type: {type(upload).__name__}") + content = read() + if inspect.isawaitable(content): + content = await content + if isinstance(content, (bytearray, memoryview)): + content = bytes(content) + if not isinstance(content, bytes): + raise TypeError( + f"Image upload read() returned {type(content).__name__}, expected bytes" + ) with open(target_path, "wb") as f: f.write(content) return target_path @@ -353,15 +372,20 @@ async def process_generation_batch( with trace_req(batch.trace_ctx), log_generation_timer(logger, batch.prompt): result = await scheduler_client.forward([batch]) - if result.output is None and result.output_file_paths is None: + if ( + result.output is None + and result.output_file_paths is None + and result.raw_frame_batches is None + ): error_msg = result.error or "Unknown error" raise RuntimeError( f"Model generation returned no output. Error from scheduler: {error_msg}" ) + save_file_path_list = [] if result.output_file_paths: save_file_path_list = result.output_file_paths - else: + elif result.output is not None: num_outputs = len(result.output) save_file_path_list = save_outputs( result.output, diff --git a/python/sglang/multimodal_gen/runtime/entrypoints/utils.py b/python/sglang/multimodal_gen/runtime/entrypoints/utils.py index 509030551..535f44cba 100644 --- a/python/sglang/multimodal_gen/runtime/entrypoints/utils.py +++ b/python/sglang/multimodal_gen/runtime/entrypoints/utils.py @@ -73,6 +73,11 @@ class ShutdownReq: pass +@dataclass +class ReleaseRealtimeSessionReq: + session_id: str + + @dataclass class GetDisaggStatsReq: """Request to get disagg pipeline metrics from the scheduler.""" @@ -125,6 +130,14 @@ class GenerationResult: output_file_path: str | None = None +@dataclass +class MaterializedOutput: + sample: Any + frames: list[Any] + audio: Any = None + fps: int = 0 + + def normalize_output_seeds( seed: int | list[int], *, @@ -187,6 +200,7 @@ def _copy_req_for_output( output_req = copy(req) output_req.sampling_params = copy(req.sampling_params) output_req.extra = dict(req.extra) + output_req.condition_inputs = dict(req.condition_inputs) output_req.trace_ctx = _copy_trace_ctx_for_output(req, request_id, output_index) return output_req @@ -455,20 +469,173 @@ def attach_audio_to_video_sample( output_idx: int, ) -> Any: """Attach per-sample audio for video outputs when available.""" + audio = select_output_audio(audio, output_idx) if audio is None: return sample - if isinstance(audio, torch.Tensor) and audio.ndim >= 2: - audio = audio[output_idx] if audio.shape[0] > output_idx else None - elif isinstance(audio, np.ndarray) and audio.ndim >= 2: - audio = audio[output_idx] if audio.shape[0] > output_idx else None - - if audio is not None and not ( - isinstance(sample, (tuple, list)) and len(sample) == 2 - ): + if not (isinstance(sample, (tuple, list)) and len(sample) == 2): return (sample, audio) return sample +def select_output_audio(audio: Any, output_idx: int) -> Any: + if isinstance(audio, torch.Tensor) and audio.ndim >= 2: + return audio[output_idx] if audio.shape[0] > output_idx else None + if isinstance(audio, np.ndarray) and audio.ndim >= 2: + return audio[output_idx] if audio.shape[0] > output_idx else None + return audio + + +def _split_sample_audio(sample: Any) -> tuple[Any, Any]: + if isinstance(sample, (tuple, list)) and len(sample) == 2: + return sample[0], sample[1] + return sample, None + + +def _sample_to_uint8_frames(sample: Any) -> list[Any]: + """return numpy frames in THCW format""" + if isinstance(sample, torch.Tensor): + # sample is raw tensor + if sample.dim() == 3: + sample = sample.unsqueeze(1) + sample = (sample * 255).clamp(0, 255).to(torch.uint8) + videos = sample.permute(1, 2, 3, 0).contiguous().cpu().numpy() + return list(videos) + + if not isinstance(sample, np.ndarray): + raise TypeError(f"Unsupported sample type: {type(sample)}") + + # sample is numpy frames + arr = sample + if arr.ndim == 3: + if arr.shape[-1] in (1, 3, 4): + arr = arr[None, ...] + else: + arr = arr[..., None] + if arr.ndim != 4: + raise ValueError(f"Unexpected numpy sample shape: {tuple(arr.shape)}") + + if arr.shape[-1] not in (1, 3, 4) and arr.shape[0] in (1, 3, 4): + t = torch.from_numpy(arr) + if t.dim() == 3: + t = t.unsqueeze(1) + t = (t * 255).clamp(0, 255).to(torch.uint8) + videos = t.permute(1, 2, 3, 0).contiguous().cpu().numpy() + return list(videos) + + if arr.dtype != np.uint8: + arr = (np.clip(arr, 0.0, 1.0) * 255.0).astype(np.uint8) + return list(arr) + + +def materialize_output_sample( + sample: Any, + data_type: DataType, + fps: int, + *, + enable_frame_interpolation: bool = False, + frame_interpolation_exp: int = 1, + frame_interpolation_scale: float = 1.0, + frame_interpolation_model_path: Optional[str] = None, + enable_upscaling: bool = False, + upscaling_model_path: Optional[str] = None, + upscaling_scale: int = 4, +) -> MaterializedOutput: + """materialize samples, apply postprocessing if applicable""" + sample_without_audio, audio = _split_sample_audio(sample) + frames = _sample_to_uint8_frames(sample_without_audio) + + # frames are uint8 numpy arrays in THWC format at this point + if enable_frame_interpolation and data_type == DataType.VIDEO and len(frames) > 1: + from sglang.multimodal_gen.runtime.postprocess import ( + interpolate_video_frames, + ) + + frames, multiplier = interpolate_video_frames( + frames, + exp=frame_interpolation_exp, + scale=frame_interpolation_scale, + model_path=frame_interpolation_model_path, + ) + fps = fps * multiplier + + if enable_upscaling and frames: + from sglang.multimodal_gen.runtime.postprocess import upscale_frames + + frames = upscale_frames( + frames, + model_path=upscaling_model_path, + scale=upscaling_scale, + ) + + return MaterializedOutput(sample=sample, frames=frames, audio=audio, fps=fps) + + +def save_materialized_output( + materialized: MaterializedOutput, + data_type: DataType, + save_file_path: Optional[str], + *, + save_output: bool = True, + audio_sample_rate: Optional[int] = None, + output_compression: Optional[int] = None, +) -> None: + if not save_output: + return + if not save_file_path: + logger.info(f"No output path provided, output not saved") + return + + os.makedirs(os.path.dirname(save_file_path), exist_ok=True) + if data_type == DataType.VIDEO: + quality = output_compression / 10 if output_compression is not None else 5 + imageio.mimsave( + save_file_path, + materialized.frames, + fps=materialized.fps, + format=data_type.get_default_extension(), + codec="libx264", + quality=quality, + ) + + _maybe_mux_audio_into_mp4( + save_file_path=save_file_path, + audio=materialized.audio, + frames=materialized.frames, + fps=materialized.fps, + audio_sample_rate=audio_sample_rate, + ) + else: + quality = output_compression if output_compression is not None else 75 + if len(materialized.frames) > 1: + for i, image in enumerate(materialized.frames): + parts = save_file_path.rsplit(".", 1) + if len(parts) == 2: + indexed_path = f"{parts[0]}_{i}.{parts[1]}" + else: + indexed_path = f"{save_file_path}_{i}" + _save_image_frame(indexed_path, image, quality, output_compression) + else: + _save_image_frame( + save_file_path, materialized.frames[0], quality, output_compression + ) + logger.info(f"Output saved to {CYAN}{save_file_path}{RESET}") + + +def _save_image_frame( + path: str, frame: np.ndarray, quality: int | None, output_compression: int | None +) -> None: + ext = os.path.splitext(path)[1].lower() + if ext == ".png": + compress_level = 1 + if output_compression is not None and output_compression != 75: + compress_level = max(0, min(9, round(output_compression / 100 * 9))) + if frame.ndim == 3 and frame.shape[-1] == 1: + frame = frame[..., 0] + Image.fromarray(frame).save(path, format="PNG", compress_level=compress_level) + else: + imageio.imwrite(path, frame, quality=quality) + + def save_outputs( outputs: Sequence[Any], data_type: DataType, @@ -490,11 +657,9 @@ def save_outputs( upscaling_model_path: Optional[str] = None, upscaling_scale: int = 4, ) -> list[str]: - """Save outputs to files and return the list of file paths.""" output_paths: list[str] = [] - for idx, output in enumerate(outputs): + for idx, sample in enumerate(outputs): save_file_path = build_output_path(idx) - sample = output if data_type == DataType.VIDEO: sample = attach_audio_to_video_sample(sample, audio, idx) @@ -519,12 +684,7 @@ def save_outputs( samples_out.append(sample) if audios_out is not None: if data_type == DataType.VIDEO: - audio_item = audio - if isinstance(audio, torch.Tensor) and audio.ndim >= 2: - audio_item = audio[idx] if audio.shape[0] > idx else None - elif isinstance(audio, np.ndarray) and audio.ndim >= 2: - audio_item = audio[idx] if audio.shape[0] > idx else None - audios_out.append(audio_item) + audios_out.append(select_output_audio(audio, idx)) else: audios_out.append(audio) if frames_out is not None: @@ -548,129 +708,26 @@ def post_process_sample( enable_upscaling: bool = False, upscaling_model_path: Optional[str] = None, upscaling_scale: int = 4, -): - """ - Process sample output, optionally interpolate video frames, and save. - """ - audio = None - if isinstance(sample, (tuple, list)) and len(sample) == 2: - sample, audio = sample - - # 1. Convert tensor / array to list of uint8 HWC frames - frames = None - if isinstance(sample, torch.Tensor): - if sample.dim() == 3: - sample = sample.unsqueeze(1) - sample = (sample * 255).clamp(0, 255).to(torch.uint8) - videos = sample.permute(1, 2, 3, 0).cpu().numpy() - frames = list(videos) - else: - if not isinstance(sample, np.ndarray): - raise TypeError(f"Unsupported sample type: {type(sample)}") - - arr = sample - if arr.ndim == 3: - if arr.shape[-1] in (1, 3, 4): - arr = arr[None, ...] - else: - arr = arr[..., None] - if arr.ndim != 4: - raise ValueError(f"Unexpected numpy sample shape: {tuple(arr.shape)}") - - if arr.shape[-1] not in (1, 3, 4) and arr.shape[0] in (1, 3, 4): - t = torch.from_numpy(arr) - if t.dim() == 3: - t = t.unsqueeze(1) - t = (t * 255).clamp(0, 255).to(torch.uint8) - videos = t.permute(1, 2, 3, 0).cpu().numpy() - frames = list(videos) - else: - if arr.dtype != np.uint8: - arr = (np.clip(arr, 0.0, 1.0) * 255.0).astype(np.uint8) - frames = list(arr) - - # 2. Frame interpolation (video only) - if enable_frame_interpolation and data_type == DataType.VIDEO and len(frames) > 1: - from sglang.multimodal_gen.runtime.postprocess import ( - interpolate_video_frames, - ) - - frames, multiplier = interpolate_video_frames( - frames, - exp=frame_interpolation_exp, - scale=frame_interpolation_scale, - model_path=frame_interpolation_model_path, - ) - fps = fps * multiplier - - # 3. Upscaling (images and videos) - if enable_upscaling and frames: - from sglang.multimodal_gen.runtime.postprocess import upscale_frames - - frames = upscale_frames( - frames, - model_path=upscaling_model_path, - scale=upscaling_scale, - ) - - # 4. Save outputs if requested - if save_output: - if save_file_path: - os.makedirs(os.path.dirname(save_file_path), exist_ok=True) - if data_type == DataType.VIDEO: - quality = ( - output_compression / 10 if output_compression is not None else 5 - ) - imageio.mimsave( - save_file_path, - frames, - fps=fps, - format=data_type.get_default_extension(), - codec="libx264", - quality=quality, - ) - - _maybe_mux_audio_into_mp4( - save_file_path=save_file_path, - audio=audio, - frames=frames, - fps=fps, - audio_sample_rate=audio_sample_rate, - ) - - else: - quality = output_compression if output_compression is not None else 75 - if len(frames) > 1: - for i, image in enumerate(frames): - parts = save_file_path.rsplit(".", 1) - if len(parts) == 2: - indexed_path = f"{parts[0]}_{i}.{parts[1]}" - else: - indexed_path = f"{save_file_path}_{i}" - _save_image_frame( - indexed_path, image, quality, output_compression - ) - else: - _save_image_frame( - save_file_path, frames[0], quality, output_compression - ) - logger.info(f"Output saved to {CYAN}{save_file_path}{RESET}") - else: - logger.info(f"No output path provided, output not saved") - - return frames - - -def _save_image_frame( - path: str, frame: np.ndarray, quality: int | None, output_compression: int | None -) -> None: - ext = os.path.splitext(path)[1].lower() - if ext == ".png": - compress_level = 1 - if output_compression is not None and output_compression != 75: - compress_level = max(0, min(9, round(output_compression / 100 * 9))) - if frame.ndim == 3 and frame.shape[-1] == 1: - frame = frame[..., 0] - Image.fromarray(frame).save(path, format="PNG", compress_level=compress_level) - else: - imageio.imwrite(path, frame, quality=quality) +) -> list[Any]: + """materialize frames and save outputs (optional)""" + materialized = materialize_output_sample( + sample, + data_type, + fps, + enable_frame_interpolation=enable_frame_interpolation, + frame_interpolation_exp=frame_interpolation_exp, + frame_interpolation_scale=frame_interpolation_scale, + frame_interpolation_model_path=frame_interpolation_model_path, + enable_upscaling=enable_upscaling, + upscaling_model_path=upscaling_model_path, + upscaling_scale=upscaling_scale, + ) + save_materialized_output( + materialized, + data_type, + save_file_path, + save_output=save_output, + audio_sample_rate=audio_sample_rate, + output_compression=output_compression, + ) + return materialized.frames diff --git a/python/sglang/multimodal_gen/runtime/launch_server.py b/python/sglang/multimodal_gen/runtime/launch_server.py index 44c5791a0..9a17a17c9 100644 --- a/python/sglang/multimodal_gen/runtime/launch_server.py +++ b/python/sglang/multimodal_gen/runtime/launch_server.py @@ -457,6 +457,7 @@ def launch_http_server_only(server_args): host=server_args.host, port=server_args.port, reload=False, + ws_per_message_deflate=False, ) diff --git a/python/sglang/multimodal_gen/runtime/layers/kvcache/__init__.py b/python/sglang/multimodal_gen/runtime/layers/kvcache/__init__.py new file mode 100644 index 000000000..1e0b2b3b7 --- /dev/null +++ b/python/sglang/multimodal_gen/runtime/layers/kvcache/__init__.py @@ -0,0 +1,13 @@ +# SPDX-License-Identifier: Apache-2.0 + +from sglang.multimodal_gen.runtime.layers.kvcache.causal_attention_cache import ( + CausalAttentionKVView, + CausalSelfAttentionKVCache, + CrossAttentionKVCache, +) + +__all__ = [ + "CausalAttentionKVView", + "CausalSelfAttentionKVCache", + "CrossAttentionKVCache", +] 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 new file mode 100644 index 000000000..823c28722 --- /dev/null +++ b/python/sglang/multimodal_gen/runtime/layers/kvcache/causal_attention_cache.py @@ -0,0 +1,238 @@ +# SPDX-License-Identifier: Apache-2.0 + +from dataclasses import dataclass + +import torch + + +@dataclass(slots=True) +class CausalAttentionKVView: + k: torch.Tensor + v: torch.Tensor + local_start_index: int + local_end_index: int + visible_local_end: int + visible_global_end: int + + +@dataclass(slots=True) +class CausalSelfAttentionKVCache: + """one transformer block's causal self-attn K/V cache and write cursors""" + + k: torch.Tensor + v: torch.Tensor + # the right bound of the valid global token range + # e.g., 12000 means [0, 12000) has been generated and cached + global_end_index: torch.Tensor + # the right bound of the valid local token range within the buffer (when cache is unfilled) + local_end_index: torch.Tensor + global_end_index_int: int | None = None + local_end_index_int: int | None = None + cache_size: int = 0 + sink_tokens: int = 0 + attention_window_size: int = 0 + allow_growth: bool = False + + def __post_init__(self) -> None: + if self.cache_size == 0: + self.cache_size = self.k.shape[1] + if self.attention_window_size == 0: + self.attention_window_size = self.cache_size + + def reset_indices(self) -> None: + self.global_end_index.zero_() + self.local_end_index.zero_() + if self.global_end_index_int is not None: + self.global_end_index_int = 0 + if self.local_end_index_int is not None: + self.local_end_index_int = 0 + + def _read_indices(self) -> tuple[int, int]: + global_end_index = self.global_end_index_int + local_end_index = self.local_end_index_int + if global_end_index is None or local_end_index is None: + global_end_index = int(self.global_end_index.item()) + local_end_index = int(self.local_end_index.item()) + self.global_end_index_int = global_end_index + self.local_end_index_int = local_end_index + return global_end_index, local_end_index + + def _write_indices(self, *, global_end_index: int, local_end_index: int) -> None: + if ( + self.global_end_index_int == global_end_index + and self.local_end_index_int == local_end_index + ): + return + if self.global_end_index_int is not None: + self.global_end_index_int = global_end_index + if self.local_end_index_int is not None: + self.local_end_index_int = local_end_index + self.global_end_index.fill_(global_end_index) + self.local_end_index.fill_(local_end_index) + + def _grow_to_fit(self, required_tokens: int) -> None: + if required_tokens <= self.cache_size: + return + old_cache_size = self.cache_size + new_cache_size = max(required_tokens, old_cache_size * 2) + + new_k = self.k.new_zeros( + self.k.shape[0], + new_cache_size, + self.k.shape[2], + self.k.shape[3], + ) + new_v = self.v.new_zeros( + self.v.shape[0], + new_cache_size, + self.v.shape[2], + self.v.shape[3], + ) + new_k[:, :old_cache_size] = self.k + new_v[:, :old_cache_size] = self.v + self.k = new_k + self.v = new_v + self.cache_size = new_cache_size + if self.attention_window_size == old_cache_size: + self.attention_window_size = new_cache_size + + def update_and_get_attention_kv( + self, + *, + key: torch.Tensor, + value: torch.Tensor, + current_chunk_start: int, + debug_name: str = "causal KV cache", + ) -> CausalAttentionKVView: + """write kv into the cache, returns the part visible to the current chunk + + Args: + current_chunk_start: the global position of the start of the chunk + + """ + num_new_tokens = key.shape[1] + current_chunk_end = current_chunk_start + num_new_tokens + kv_cache_size = self.cache_size + sink_tokens = self.sink_tokens + global_end_index, local_end_index_prev = self._read_indices() + + # local_end_index: the local position of the end of current chunk + # updated_local_end: the updated local end + # updated_global_end: the updated global end + + # the global position of the start of the buffer + window_start = global_end_index - local_end_index_prev + + if current_chunk_end <= global_end_index: + # the window stays as previous + # cache layout: + # [sink tokens, recent window tokens, current chunk tokens, uninitialized tokens (optional)] + local_start_index = current_chunk_start - window_start + local_end_index = local_start_index + num_new_tokens + + # the local end and global end remains unchanged (since the chunk hasn't proceed) + updated_local_end = local_end_index_prev + updated_global_end = global_end_index + else: + # the chunk window has proceed, append new tokens, and evict earliest (if have to) + appended_tokens = current_chunk_end - global_end_index + if self.allow_growth: + self._grow_to_fit(local_end_index_prev + appended_tokens) + kv_cache_size = self.cache_size + if local_end_index_prev + appended_tokens > kv_cache_size: + # the new tokens can't fit in the remaining space (after local_end_index_prev), start evicting: + # before: + # [sink tokens, evicted tokens, rolled tokens, remaining space] + # ^ end of previous chunk + # after: + # [sink tokens, rolled tokens, remaining space ] + + # 1. keep sink tokens ([0: sink_tokens]) untouched + # 2. evict obsolete tokens in: [sink_tokens:sink_tokens + num_evicted_tokens] + num_evicted_tokens = ( + local_end_index_prev + appended_tokens - kv_cache_size + ) + + # number of tokens to move + num_rolled_tokens = max( + 0, + local_end_index_prev - num_evicted_tokens - sink_tokens, + ) + if num_rolled_tokens > 0: + self.k[:, sink_tokens : sink_tokens + num_rolled_tokens] = self.k[ + :, + sink_tokens + + num_evicted_tokens : sink_tokens + + num_evicted_tokens + + num_rolled_tokens, + ].clone() + self.v[:, sink_tokens : sink_tokens + num_rolled_tokens] = self.v[ + :, + sink_tokens + + num_evicted_tokens : sink_tokens + + num_evicted_tokens + + num_rolled_tokens, + ].clone() + + # 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: + # enough space, directly append new tokens after end of previous chunk + local_end_index = local_end_index_prev + appended_tokens + local_start_index = local_end_index - num_new_tokens + updated_local_end = local_end_index + # after filling in the proceeded new chunk, the global end aligns with the global end of the current chunk + updated_global_end = current_chunk_end + + if ( + local_start_index < 0 + or local_end_index > kv_cache_size + or local_end_index - local_start_index != num_new_tokens + ): + raise RuntimeError( + f"Invalid {debug_name} write range: " + f"local=[{local_start_index}, {local_end_index}), " + f"global_end={global_end_index}, " + f"prev_local_end={local_end_index_prev}, " + f"kv_cache_size={kv_cache_size}, " + f"num_new_tokens={num_new_tokens}, " + f"current_start={current_chunk_start}, current_end={current_chunk_end}" + ) + + if self.k.requires_grad: + self.k = self.k.detach() + if self.v.requires_grad: + self.v = self.v.detach() + self.k[:, local_start_index:local_end_index] = key + self.v[:, local_start_index:local_end_index] = value + + attn_start_index = max(0, updated_local_end - self.attention_window_size) + self._write_indices( + global_end_index=updated_global_end, + local_end_index=updated_local_end, + ) + return CausalAttentionKVView( + k=self.k[:, attn_start_index:updated_local_end], + v=self.v[:, attn_start_index:updated_local_end], + local_start_index=local_start_index, + local_end_index=local_end_index, + visible_local_end=updated_local_end, + visible_global_end=updated_global_end, + ) + + +@dataclass(slots=True) +class CrossAttentionKVCache: + """one transformer block's cross-attn condition K/V cache""" + + k: torch.Tensor + v: torch.Tensor + is_init: bool = False + + def store(self, k: torch.Tensor, v: torch.Tensor) -> None: + self.k = k.detach() + self.v = v.detach() + self.is_init = True + + def reset(self) -> None: + self.is_init = False diff --git a/python/sglang/multimodal_gen/runtime/layers/usp.py b/python/sglang/multimodal_gen/runtime/layers/usp.py index 7053a7de3..aaabb3303 100644 --- a/python/sglang/multimodal_gen/runtime/layers/usp.py +++ b/python/sglang/multimodal_gen/runtime/layers/usp.py @@ -4,11 +4,13 @@ import logging from typing import TYPE_CHECKING import torch +import torch.distributed as dist import torch.distributed._functional_collectives as ft_c from torch.distributed.tensor.experimental._attention import _cp_options from sglang.multimodal_gen.runtime.distributed.parallel_state import ( get_sp_group, + get_ulysses_parallel_rank, get_ulysses_parallel_world_size, ) from sglang.srt.utils.common import torch_release @@ -45,6 +47,25 @@ def _usp_all_to_all_single(x: torch.Tensor) -> torch.Tensor: return output.reshape(x_shape) +def _usp_all_to_all_single_varlen( + x: torch.Tensor, + output_split_sizes: list[int], + input_split_sizes: list[int], +) -> torch.Tensor: + ulysses_pg = get_sp_group().ulysses_group + assert ulysses_pg is not None, "Ulysses process group is not initialized." + x = x.flatten().contiguous() + output = torch.empty(sum(output_split_sizes), dtype=x.dtype, device=x.device) + dist.all_to_all_single( + output, + x, + output_split_sizes=output_split_sizes, + input_split_sizes=input_split_sizes, + group=ulysses_pg, + ) + return output + + def _usp_input_all_to_all(x: torch.Tensor, head_dim: int = 1) -> torch.Tensor: """ Perform Ulysses-style input all-to-all over the head dimension. @@ -101,6 +122,83 @@ def _usp_input_all_to_all(x: torch.Tensor, head_dim: int = 1) -> torch.Tensor: return x +def _usp_input_all_to_all_varlen( + x: torch.Tensor, seq_lens: list[int], head_dim: int = 1 +) -> torch.Tensor: + """ + Perform Ulysses-style input all-to-all over the head dimension with variable + local sequence lengths. + + Default layout expects heads at dim=1 and sequence at dim=2: + [b, h, s_local, d] -> [b, h_local, s_global, d] + + If heads are at dim=2 (input is [b, s_local, h, d]), set head_dim=2, and the + function returns [b, s_global, h_local, d], preserving the original + head/sequence dim ordering. + + Args: + x: A 4D tensor with layout [b, *, *, d] where '*' are sequence and heads + seq_lens: Local sequence lengths for each rank in the Ulysses group + head_dim: Which dimension index corresponds to heads (1 or 2) + + Returns: + Tensor with the same dim order as input, with heads sharded and sequence gathered. + """ + world_size = get_ulysses_parallel_world_size() + if world_size <= 1: + return x + + assert x.ndim == 4, f"x must have 4 dimensions, got {x.ndim}" + assert head_dim in (1, 2), f"head_dim must be 1 or 2, got {head_dim}" + assert ( + len(seq_lens) == world_size + ), f"seq_lens must have length {world_size}, got {len(seq_lens)}" + + rank = get_ulysses_parallel_rank() + + # Move the dimension to be split (h_global) to dim 0 for all_to_all_single + if head_dim == 1: + b, h_global, s_local, d = x.shape + # Shape transition: [b, h_global, s_local, d] -> [h_global, b, s_local, d] + permute_order = (1, 0, 2, 3) + else: # head_dim == 2 + b, s_local, h_global, d = x.shape + # Shape transition: [b, s_local, h_global, d] -> [h_global, b, s_local, d] + permute_order = (2, 0, 1, 3) + + assert ( + s_local == seq_lens[rank] + ), f"s_local ({s_local}) must equal seq_lens[{rank}] ({seq_lens[rank]})" + assert ( + h_global % world_size == 0 + ), f"h_global ({h_global}) must be divisible by world_size ({world_size})" + + h_local = h_global // world_size + + x = x.permute(permute_order).contiguous() + x = x.reshape(world_size, h_local, b, s_local, d) + input_split_sizes = [h_local * b * s_local * d] * world_size + output_split_sizes = [h_local * b * seq_len * d for seq_len in seq_lens] + x = _usp_all_to_all_single_varlen(x, output_split_sizes, input_split_sizes) + + chunks = [] + offset = 0 + for seq_len, split_size in zip(seq_lens, output_split_sizes): + chunk = x[offset : offset + split_size].reshape(h_local, b, seq_len, d) + chunks.append(chunk) + offset += split_size + x = torch.cat(chunks, dim=2) + + if head_dim == 1: + # Shape transition: [h_local, b, s_global, d] -> [b, h_local, s_global, d] + x = x.permute(1, 0, 2, 3).contiguous() + else: # head_dim == 2 + # Shape transition: [h_local, b, s_global, d] -> [b, s_global, h_local, d] + x = x.permute(1, 2, 0, 3).contiguous() + + return x + + def _usp_output_all_to_all(x: torch.Tensor, head_dim: int = 1) -> torch.Tensor: """ Perform Ulysses-style output all-to-all over the head dimension (inverse of input). @@ -157,6 +255,86 @@ def _usp_output_all_to_all(x: torch.Tensor, head_dim: int = 1) -> torch.Tensor: return x +def _usp_output_all_to_all_varlen( + x: torch.Tensor, seq_lens: list[int], head_dim: int = 1 +) -> torch.Tensor: + """ + Perform Ulysses-style output all-to-all over the head dimension (inverse of input) + with variable local sequence lengths. + + Default layout expects heads at dim=1 and sequence at dim=2: + [b, h_local, s, d] -> [b, h, s_local, d] + + If heads are at dim=2 (input is [b, s_global, h // world_size, d]), set head_dim=2, + and the function returns [b, s_local, h, d], preserving the original head/sequence + dim ordering. + + Args: + x: A 4D tensor with layout [b, *, *, d] where '*' are sequence and heads + seq_lens: Local sequence lengths for each rank in the Ulysses group + head_dim: Which dimension index corresponds to heads (1 or 2) + + Returns: + Tensor with the same dim order as input, with heads gathered and sequence sharded. + """ + world_size = get_ulysses_parallel_world_size() + if world_size <= 1: + return x + + assert x.ndim == 4, f"x must have 4 dimensions, got {x.ndim}" + assert head_dim in (1, 2), f"head_dim must be 1 or 2, got {head_dim}" + assert ( + len(seq_lens) == world_size + ), f"seq_lens must have length {world_size}, got {len(seq_lens)}" + + rank = get_ulysses_parallel_rank() + + # Move the sequence dimension to dim 2 for splitting across seq_lens + if head_dim == 1: + b, h_local, s_global, d = x.shape + # Shape transition: [b, h_local, s_global, d] -> [h_local, b, s_global, d] + permute_order = (1, 0, 2, 3) + else: # head_dim == 2 + b, s_global, h_local, d = x.shape + # Shape transition: [b, s_global, h_local, d] -> [h_local, b, s_global, d] + permute_order = (2, 0, 1, 3) + + assert s_global == sum( + seq_lens + ), f"s_global ({s_global}) must equal sum(seq_lens) ({sum(seq_lens)})" + + s_local = seq_lens[rank] + + x = x.permute(permute_order).contiguous() + input_chunks = [] + start = 0 + for seq_len in seq_lens: + end = start + seq_len + input_chunks.append(x[:, :, start:end, :].contiguous().reshape(-1)) + start = end + x = torch.cat(input_chunks, dim=0) + input_split_sizes = [h_local * b * seq_len * d for seq_len in seq_lens] + output_split_sizes = [h_local * b * s_local * d] * world_size + x = _usp_all_to_all_single_varlen(x, output_split_sizes, input_split_sizes) + + chunks = [] + offset = 0 + for split_size in output_split_sizes: + chunk = x[offset : offset + split_size].reshape(h_local, b, s_local, d) + chunks.append(chunk) + offset += split_size + x = torch.cat(chunks, dim=0) + + if head_dim == 1: + # Shape transition: [h_global, b, s_local, d] -> [b, h_global, s_local, d] + x = x.permute(1, 0, 2, 3).contiguous() + else: # head_dim == 2 + # Shape transition: [h_global, b, s_local, d] -> [b, s_local, h_global, d] + x = x.permute(1, 2, 0, 3).contiguous() + + return x + + def ring_attn( query: torch.Tensor, key: torch.Tensor, diff --git a/python/sglang/multimodal_gen/runtime/layers/visual_embedding.py b/python/sglang/multimodal_gen/runtime/layers/visual_embedding.py index fe9669e23..96de21f7b 100644 --- a/python/sglang/multimodal_gen/runtime/layers/visual_embedding.py +++ b/python/sglang/multimodal_gen/runtime/layers/visual_embedding.py @@ -111,6 +111,61 @@ class PatchEmbed(nn.Module): return x +class WanCamControlPatchEmbedding(nn.Module): + """Patch embedding used by LingBotWorld camera/plucker controls.""" + + def __init__( + self, + patch_size=(1, 2, 2), + in_chans=384, + embed_dim=2048, + bias=True, + dtype=None, + prefix: str = "", + ): + super().__init__() + del prefix + if isinstance(patch_size, list | tuple): + if len(patch_size) != 3: + raise ValueError( + f"patch_size must have length 3, got {len(patch_size)}" + ) + patch_size = tuple(patch_size) + else: + raise ValueError(f"Unsupported patch_size type: {type(patch_size)}") + + self.patch_size = patch_size + pt, ph, pw = self.patch_size + self.in_features = in_chans * pt * ph * pw + self.proj = nn.Linear(self.in_features, embed_dim, bias=bias, dtype=dtype) + + def forward(self, x: torch.Tensor) -> torch.Tensor: + if x.dim() != 5: + raise ValueError( + f"Expected camera embedding shape [B, C, F, H, W], got {tuple(x.shape)}" + ) + + bsz, channels, frames, height, width = x.shape + pt, ph, pw = self.patch_size + if (frames % pt) != 0 or (height % ph) != 0 or (width % pw) != 0: + raise ValueError( + f"Input shape {tuple(x.shape)} must be divisible by patch_size {self.patch_size}" + ) + + x = x.view( + bsz, + channels, + frames // pt, + pt, + height // ph, + ph, + width // pw, + pw, + ) + x = x.permute(0, 2, 4, 6, 1, 3, 5, 7).reshape(bsz, -1, self.in_features) + return self.proj(x) + + class Timesteps(_Timesteps): def forward(self, timesteps: torch.Tensor) -> torch.Tensor: if _is_cuda: diff --git a/python/sglang/multimodal_gen/runtime/managers/gpu_worker.py b/python/sglang/multimodal_gen/runtime/managers/gpu_worker.py index ee4f1cd07..e9a6af9e4 100644 --- a/python/sglang/multimodal_gen/runtime/managers/gpu_worker.py +++ b/python/sglang/multimodal_gen/runtime/managers/gpu_worker.py @@ -33,6 +33,7 @@ from sglang.multimodal_gen.runtime.distributed.parallel_state import ( get_ulysses_parallel_world_size, ) from sglang.multimodal_gen.runtime.entrypoints.utils import ( + materialize_output_sample, post_process_sample, save_outputs, ) @@ -53,6 +54,9 @@ from sglang.multimodal_gen.runtime.pipelines_core import ( ) from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import OutputBatch from sglang.multimodal_gen.runtime.platforms import current_platform +from sglang.multimodal_gen.runtime.realtime.session import ( + RealtimeSessionCache, +) from sglang.multimodal_gen.runtime.server_args import PortArgs, ServerArgs from sglang.multimodal_gen.runtime.utils.common import set_cuda_arch, set_musa_arch from sglang.multimodal_gen.runtime.utils.logging_utils import ( @@ -64,6 +68,10 @@ from sglang.multimodal_gen.runtime.utils.perf_logger import ( PerformanceLogger, capture_memory_snapshot, ) +from sglang.multimodal_gen.runtime.utils.realtime_video import ( + RAW_RGB_CONTENT_TYPE, + build_raw_rgb_frame_batches, +) from sglang.multimodal_gen.runtime.utils.trace_wrapper import DiffStage, trace_slice from sglang.srt.observability.trace import process_tracing_init, trace_set_thread_info from sglang.srt.utils.network import NetworkAddress @@ -118,6 +126,24 @@ class GPUWorker: self.cfg_group = get_cfg_group() self.cfg_cpu_group = self.cfg_group.cpu_group + self._realtime_sessions = RealtimeSessionCache(max_sessions=1) + + def release_realtime_session(self, session_id: str) -> OutputBatch: + """release the session of a realtime connection""" + if not session_id: + return OutputBatch( + output={ + "released": False, + "session_id": session_id, + "reason": "empty_session_id", + } + ) + + released = self._realtime_sessions.release(session_id) + if released: + if torch.cuda.is_initialized(): + torch.cuda.empty_cache() + return OutputBatch(output={"released": released, "session_id": session_id}) def init_device_and_model(self) -> None: """Initialize the device and load the model.""" @@ -273,7 +299,7 @@ class GPUWorker: error_context=f"request {req.request_id}", ) - def _execute_forward_batch(self, batch: list[Req]) -> OutputBatch: + def _execute_forward_batch(self, batch: list[Req]) -> OutputBatch | Req: """Execute expanded multi-output requests as one grouped forward.""" # TODO: support early return or mix-stage execution for reqs in a group assert self.pipeline is not None @@ -309,6 +335,7 @@ class GPUWorker: torch.get_device_module().reset_peak_memory_stats() start_time = time.monotonic() + self._realtime_sessions.attach(req) # capture memory baseline for each req in grouped forward on rank-0 request_metrics = [ @@ -354,21 +381,13 @@ class GPUWorker: for metrics in output_metrics: metrics.total_duration_ms = duration_ms - # file-path-only responses avoid serializing generated tensors between - # scheduler_client and gpu_worker. - if req.save_output and req.return_file_paths_only: - save_output_paths(output_batch) - output_batch.output = None - output_batch.audio = None - output_batch.audio_sample_rate = None + self._materialize_output_transport(output_batch, req, save_output_paths) - if torch.cuda.is_initialized(): - torch.cuda.empty_cache() - - # Keep return_frames payloads off the scheduler's tensor ZMQ path. - self._materialize_frame_outputs_for_return(output_batch, req) - - if torch.cuda.is_initialized() and output_batch.output is None: + if ( + torch.cuda.is_initialized() + and output_batch.output is None + and not req.return_raw_frames + ): torch.cuda.empty_cache() if req.perf_dump_path is not None or envs.SGLANG_DIFFUSION_STAGE_LOGGING: @@ -403,9 +422,54 @@ class GPUWorker: torch.cuda.empty_cache() return output_batch + def _materialize_output_transport( + self, + output_batch: OutputBatch, + req: Req, + save_output_paths: Callable[[OutputBatch], None], + ) -> None: + if req.return_raw_frames: + self._materialize_raw_frame_transport(output_batch, req) + elif req.save_output and req.return_file_paths_only: + self._materialize_file_path_transport(output_batch, save_output_paths) + elif req.return_frames: + self._materialize_frame_outputs_for_return(output_batch, req) + + def _materialize_raw_frame_transport( + self, output_batch: OutputBatch, req: Req + ) -> None: + if self.rank != 0: + return + if output_batch.output is not None: + output_batch.raw_frame_content_type = RAW_RGB_CONTENT_TYPE + ( + output_batch.raw_frame_batches, + output_batch.raw_frame_metadata, + ) = build_raw_rgb_frame_batches( + output_batch.output, + req, + output_batch, + post_process_sample, + ) + output_batch.output = None + output_batch.audio = None + output_batch.audio_sample_rate = None + + def _materialize_file_path_transport( + self, + output_batch: OutputBatch, + save_output_paths: Callable[[OutputBatch], None], + ) -> None: + if self.rank == 0: + save_output_paths(output_batch) + output_batch.output = None + output_batch.audio = None + output_batch.audio_sample_rate = None + def _materialize_frame_outputs_for_return( self, output_batch: OutputBatch, req: Req ) -> None: + """materialize the output from tensor to numpy frames for faster serialization""" if self.rank != 0 or output_batch.output is None or not req.return_frames: return @@ -452,13 +516,10 @@ class GPUWorker: ): return output - frames = post_process_sample( + materialized = materialize_output_sample( output, req.data_type, req.fps, - save_output=False, - audio_sample_rate=output_batch.audio_sample_rate, - output_compression=req.output_compression, enable_frame_interpolation=req.enable_frame_interpolation, frame_interpolation_exp=req.frame_interpolation_exp, frame_interpolation_scale=req.frame_interpolation_scale, @@ -467,7 +528,7 @@ class GPUWorker: upscaling_model_path=req.upscaling_model_path, upscaling_scale=req.upscaling_scale, ) - return np.asarray(frames) + return np.asarray(materialized.frames) def _record_output_peak_memory(self, output_batch: OutputBatch) -> None: if self.rank != 0 or current_platform.is_cpu(): @@ -482,6 +543,7 @@ class GPUWorker: return self._merge_expanded_output_batches(output_batches) def _save_output_paths(self, req: Req, output_batch: OutputBatch) -> None: + """save outputs to files""" if self.rank != 0 or output_batch.output is None: return @@ -500,10 +562,15 @@ class GPUWorker: dynamic_output_paths = None if dynamic_output_paths is not None: - build_output_path = lambda idx: dynamic_output_paths[idx] + + def build_output_path(idx: int) -> str: + return dynamic_output_paths[idx] + else: num_outputs = len(output_batch.output) - build_output_path = lambda idx: req.output_file_path(num_outputs, idx) + + def build_output_path(idx: int) -> str: + return req.output_file_path(num_outputs, idx) output_batch.output_file_paths = save_outputs( output_batch.output, diff --git a/python/sglang/multimodal_gen/runtime/managers/scheduler.py b/python/sglang/multimodal_gen/runtime/managers/scheduler.py index a763e5bd4..d46e76f2f 100644 --- a/python/sglang/multimodal_gen/runtime/managers/scheduler.py +++ b/python/sglang/multimodal_gen/runtime/managers/scheduler.py @@ -25,6 +25,7 @@ from sglang.multimodal_gen.runtime.entrypoints.utils import ( GetDisaggStatsReq, ListLorasReq, MergeLoraWeightsReq, + ReleaseRealtimeSessionReq, SetLoraReq, ShutdownReq, UnmergeLoraWeightsReq, @@ -128,6 +129,7 @@ class Scheduler(SchedulerDisaggMixin): Req: self._handle_generation, ListLorasReq: self._handle_list_loras, ShutdownReq: self._handle_shutdown, + ReleaseRealtimeSessionReq: self._handle_release_realtime_session, GetDisaggStatsReq: self._handle_get_disagg_stats, UpdateWeightFromDiskReqInput: self._handle_update_weights_from_disk, GetWeightsChecksumReqInput: self._handle_get_weights_checksum, @@ -205,6 +207,10 @@ class Scheduler(SchedulerDisaggMixin): self._running = False return OutputBatch() + def _handle_release_realtime_session(self, reqs: List[Any]) -> OutputBatch: + req = reqs[0] + return self.worker.release_realtime_session(req.session_id) + def _handle_update_weights_from_disk(self, reqs: List[Any]) -> OutputBatch: """Handle update_weights_from_disk request for RL workflows.""" req = reqs[0] @@ -487,6 +493,10 @@ class Scheduler(SchedulerDisaggMixin): if base_req.is_warmup or candidate_req.is_warmup: return "warmup" + if self._has_realtime_session(base_req) or self._has_realtime_session( + candidate_req + ): + return "realtime_session" if not isinstance(base_req.prompt, str) or not isinstance( candidate_req.prompt, str ): @@ -506,11 +516,20 @@ class Scheduler(SchedulerDisaggMixin): or "signature_mismatch" ) + @staticmethod + def _has_realtime_session(req: Req) -> bool: + return bool(req.realtime_session_id) or req.session is not None + def _can_dynamic_batch(self, base_req: Req, candidate_req: Req) -> bool: """Return whether `candidate_req` can be merged into a batch with `base_req`.""" if base_req.is_warmup or candidate_req.is_warmup: return False + if self._has_realtime_session(base_req) or self._has_realtime_session( + candidate_req + ): + return False + if not isinstance(base_req.prompt, str) or not isinstance( candidate_req.prompt, str ): 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 dea47f10f..370968609 100644 --- a/python/sglang/multimodal_gen/runtime/models/dits/causal_wanvideo.py +++ b/python/sglang/multimodal_gen/runtime/models/dits/causal_wanvideo.py @@ -29,6 +29,10 @@ from sglang.multimodal_gen.configs.models.dits import WanVideoConfig from sglang.multimodal_gen.runtime.distributed.parallel_state import get_sp_world_size from sglang.multimodal_gen.runtime.layers.attention import LocalAttention from sglang.multimodal_gen.runtime.layers.elementwise import MulAdd +from sglang.multimodal_gen.runtime.layers.kvcache.causal_attention_cache import ( + CausalSelfAttentionKVCache, + CrossAttentionKVCache, +) from sglang.multimodal_gen.runtime.layers.layernorm import ( FP32LayerNorm, LayerNormScaleShift, @@ -80,9 +84,6 @@ class CausalWanSelfAttention(nn.Module): self.qk_norm = qk_norm self.eps = eps self.parallel_attention = parallel_attention - self.max_attention_size = ( - 32760 if local_attn_size == -1 else local_attn_size * 1560 - ) # Scaled dot product attention self.attn = LocalAttention( @@ -105,7 +106,7 @@ class CausalWanSelfAttention(nn.Module): v: torch.Tensor, freqs_cis: tuple[torch.Tensor, torch.Tensor], block_mask: BlockMask, - kv_cache: dict | None = None, + kv_cache: CausalSelfAttentionKVCache | None = None, current_start: int = 0, cache_start: int | None = None, ): @@ -116,9 +117,6 @@ class CausalWanSelfAttention(nn.Module): grid_sizes(Tensor): Shape [B, 3], the second dimension contains (F, H, W) freqs(Tensor): Rope freqs, shape [1024, C / num_heads / 2] """ - if cache_start is None: - cache_start = current_start - cos, sin = freqs_cis roped_query = _apply_rotary_emb(q, cos, sin, is_neox_style=False).type_as(v) roped_key = _apply_rotary_emb(k, cos, sin, is_neox_style=False).type_as(v) @@ -169,84 +167,17 @@ class CausalWanSelfAttention(nn.Module): block_mask=block_mask, )[:, :, :-padded_length].transpose(2, 1) else: - frame_seqlen = q.shape[1] - current_end = current_start + roped_query.shape[1] - sink_tokens = self.sink_size * frame_seqlen - # If we are using local attention and the current KV cache size is larger than the local attention size, we need to truncate the KV cache - kv_cache_size = kv_cache["k"].shape[1] - num_new_tokens = roped_query.shape[1] - if ( - self.local_attn_size != -1 - and (current_end > kv_cache["global_end_index"].item()) - and ( - num_new_tokens + kv_cache["local_end_index"].item() > kv_cache_size - ) - ): - # Calculate the number of new tokens added in this step - # Shift existing cache content left to discard oldest tokens - # Clone the source slice to avoid overlapping memory error - num_evicted_tokens = ( - num_new_tokens + kv_cache["local_end_index"].item() - kv_cache_size - ) - num_rolled_tokens = ( - kv_cache["local_end_index"].item() - - num_evicted_tokens - - sink_tokens - ) - kv_cache["k"][ - :, sink_tokens : sink_tokens + num_rolled_tokens - ] = kv_cache["k"][ - :, - sink_tokens - + num_evicted_tokens : sink_tokens - + num_evicted_tokens - + num_rolled_tokens, - ].clone() - kv_cache["v"][ - :, sink_tokens : sink_tokens + num_rolled_tokens - ] = kv_cache["v"][ - :, - sink_tokens - + num_evicted_tokens : sink_tokens - + num_evicted_tokens - + num_rolled_tokens, - ].clone() - # Insert the new keys/values at the end - local_end_index = ( - kv_cache["local_end_index"].item() - + current_end - - kv_cache["global_end_index"].item() - - num_evicted_tokens - ) - local_start_index = local_end_index - num_new_tokens - kv_cache["k"][:, local_start_index:local_end_index] = roped_key - kv_cache["v"][:, local_start_index:local_end_index] = v - else: - # Assign new keys/values directly up to current_end - local_end_index = ( - kv_cache["local_end_index"].item() - + current_end - - kv_cache["global_end_index"].item() - ) - local_start_index = local_end_index - num_new_tokens - kv_cache["k"] = kv_cache["k"].detach() - kv_cache["v"] = kv_cache["v"].detach() - # logger.info("kv_cache['k'] is in comp graph: %s", kv_cache["k"].requires_grad or kv_cache["k"].grad_fn is not None) - kv_cache["k"][:, local_start_index:local_end_index] = roped_key - kv_cache["v"][:, local_start_index:local_end_index] = v + cache_view = kv_cache.update_and_get_attention_kv( + key=roped_key, + value=v, + current_chunk_start=current_start, + debug_name="CausalWan KV cache", + ) x = self.attn( roped_query, - kv_cache["k"][ - :, - max(0, local_end_index - self.max_attention_size) : local_end_index, - ], - kv_cache["v"][ - :, - max(0, local_end_index - self.max_attention_size) : local_end_index, - ], + cache_view.k, + cache_view.v, ) - kv_cache["global_end_index"].fill_(current_end) - kv_cache["local_end_index"].fill_(local_end_index) return x @@ -335,8 +266,8 @@ class CausalWanTransformerBlock(nn.Module): temb: torch.Tensor, freqs_cis: tuple[torch.Tensor, torch.Tensor], block_mask: BlockMask, - kv_cache: dict | None = None, - crossattn_cache: dict | None = None, + kv_cache: CausalSelfAttentionKVCache | None = None, + crossattn_cache: CrossAttentionKVCache | None = None, current_start: int = 0, cache_start: int | None = None, ) -> torch.Tensor: @@ -594,18 +525,17 @@ class CausalWanTransformer3DModel(BaseDiT, LayerwiseOffloadableModuleMixin): return block_mask - def _forward_inference( + def forward( self, hidden_states: torch.Tensor, encoder_hidden_states: torch.Tensor | list[torch.Tensor], timestep: torch.LongTensor, encoder_hidden_states_image: torch.Tensor | list[torch.Tensor] | None = None, - kv_cache: dict = None, - crossattn_cache: dict = None, + kv_cache: list[CausalSelfAttentionKVCache] | None = None, + crossattn_cache: list[CrossAttentionKVCache] | None = None, current_start: int = 0, cache_start: int = 0, start_frame: int = 0, - **kwargs, ) -> torch.Tensor: r""" Run the diffusion model with kv caching. @@ -739,144 +669,5 @@ class CausalWanTransformer3DModel(BaseDiT, LayerwiseOffloadableModuleMixin): return output - def _forward_train( - self, - hidden_states: torch.Tensor, - encoder_hidden_states: torch.Tensor | list[torch.Tensor], - timestep: torch.LongTensor, - encoder_hidden_states_image: torch.Tensor | list[torch.Tensor] | None = None, - start_frame: int = 0, - **kwargs, - ) -> torch.Tensor: - - orig_dtype = hidden_states.dtype - if not isinstance(encoder_hidden_states, torch.Tensor): - encoder_hidden_states = encoder_hidden_states[0] - if ( - isinstance(encoder_hidden_states_image, list) - and len(encoder_hidden_states_image) > 0 - ): - encoder_hidden_states_image = encoder_hidden_states_image[0] - else: - encoder_hidden_states_image = None - - batch_size, num_channels, num_frames, height, width = hidden_states.shape - p_t, p_h, p_w = self.patch_size - post_patch_num_frames = num_frames // p_t - post_patch_height = height // p_h - post_patch_width = width // p_w - - # Get rotary embeddings - d = self.hidden_size // self.num_attention_heads - rope_dim_list = [d - 4 * (d // 6), 2 * (d // 6), 2 * (d // 6)] - freqs_cos, freqs_sin = get_rotary_pos_embed( - ( - post_patch_num_frames * get_sp_world_size(), - post_patch_height, - post_patch_width, - ), - self.hidden_size, - self.num_attention_heads, - rope_dim_list, - dtype=( - torch.float64 - if current_platform.is_float64_supported() - else torch.float32 - ), - rope_theta=10000, - start_frame=start_frame, - ) - freqs_cos = freqs_cos.to(hidden_states.device) - freqs_sin = freqs_sin.to(hidden_states.device) - freqs_cis = ( - (freqs_cos.float(), freqs_sin.float()) if freqs_cos is not None else None - ) - - # Construct blockwise causal attn mask - if self.block_mask is None: - self.block_mask = self._prepare_blockwise_causal_attn_mask( - device=hidden_states.device, - num_frames=num_frames, - frame_seqlen=post_patch_height * post_patch_width, - num_frame_per_block=self.num_frame_per_block, - local_attn_size=self.local_attn_size, - ) - - hidden_states = self.patch_embedding(hidden_states) - hidden_states = hidden_states.flatten(2).transpose(1, 2) - - ( - temb, - timestep_proj, - encoder_hidden_states, - encoder_hidden_states_image, - ) = self.condition_embedder( - timestep.flatten(), encoder_hidden_states, encoder_hidden_states_image - ) - timestep_proj = timestep_proj.unflatten(1, (6, self.hidden_size)).unflatten( - dim=0, sizes=timestep.shape - ) - - if encoder_hidden_states_image is not None: - encoder_hidden_states = torch.concat( - [encoder_hidden_states_image, encoder_hidden_states], dim=1 - ) - - encoder_hidden_states = ( - encoder_hidden_states.to(orig_dtype) - if current_platform.is_mps() - else encoder_hidden_states - ) # cast to orig_dtype for MPS - - assert encoder_hidden_states.dtype == orig_dtype - - # 4. Transformer blocks - if torch.is_grad_enabled() and self.gradient_checkpointing: - for block in self.blocks: - hidden_states = self._gradient_checkpointing_func( - block, - hidden_states, - encoder_hidden_states, - timestep_proj, - freqs_cis, - block_mask=self.block_mask, - ) - else: - for block in self.blocks: - hidden_states = block( - hidden_states, - encoder_hidden_states, - timestep_proj, - freqs_cis, - block_mask=self.block_mask, - ) - - # 5. Output norm, projection & unpatchify - temb = temb.unflatten(dim=0, sizes=timestep.shape).unsqueeze(2) - shift, scale = (self.scale_shift_table.unsqueeze(1) + temb).chunk(2, dim=2) - hidden_states = self.norm_out(hidden_states, shift, scale) - hidden_states = self.proj_out(hidden_states) - - hidden_states = hidden_states.reshape( - batch_size, - post_patch_num_frames, - post_patch_height, - post_patch_width, - p_t, - p_h, - p_w, - -1, - ) - hidden_states = hidden_states.permute(0, 7, 1, 4, 2, 5, 3, 6) - output = hidden_states.flatten(6, 7).flatten(4, 5).flatten(2, 3) - - return output - - def forward(self, *args, **kwargs): - if kwargs.get("kv_cache") is not None: - return self._forward_inference(*args, **kwargs) - else: - return self._forward_train(*args, **kwargs) - EntryClass = CausalWanTransformer3DModel diff --git a/python/sglang/multimodal_gen/runtime/models/dits/lingbot_world.py b/python/sglang/multimodal_gen/runtime/models/dits/lingbot_world.py new file mode 100644 index 000000000..5df820e65 --- /dev/null +++ b/python/sglang/multimodal_gen/runtime/models/dits/lingbot_world.py @@ -0,0 +1,1360 @@ +# Copied and adapted from: https://github.com/hao-ai-lab/FastVideo +# Adapted from: https://github.com/Robbyant/lingbot-world + +# SPDX-License-Identifier: Apache-2.0 + +import math +from functools import lru_cache +from typing import Any + +import torch +import torch.distributed as dist +import torch.nn as nn +import torch.nn.functional as F + +from sglang.multimodal_gen.configs.models.dits import LingBotWorldVideoConfig +from sglang.multimodal_gen.runtime.distributed import ( + divide, + get_sp_group, + get_sp_parallel_rank, + get_sp_world_size, + get_tp_world_size, + sequence_model_parallel_all_gather, +) +from sglang.multimodal_gen.runtime.distributed.parallel_state import ( + get_ring_parallel_world_size, + get_ulysses_parallel_world_size, +) +from sglang.multimodal_gen.runtime.layers.attention import LocalAttention, USPAttention +from sglang.multimodal_gen.runtime.layers.elementwise import MulAdd +from sglang.multimodal_gen.runtime.layers.kvcache.causal_attention_cache import ( + CausalSelfAttentionKVCache, + CrossAttentionKVCache, +) +from sglang.multimodal_gen.runtime.layers.layernorm import ( + LayerNormScaleShift, + RMSNorm, + ScaleResidualLayerNormScaleShift, + tensor_parallel_rms_norm, +) +from sglang.multimodal_gen.runtime.layers.linear import ( + ColumnParallelLinear, + RowParallelLinear, +) +from sglang.multimodal_gen.runtime.layers.mlp import MLP +from sglang.multimodal_gen.runtime.layers.quantization.configs.base_config import ( + QuantizationConfig, +) +from sglang.multimodal_gen.runtime.layers.rotary_embedding import ( + NDRotaryEmbedding, + _apply_rotary_emb, + apply_flashinfer_rope_qk_inplace, + get_rotary_pos_embed, +) +from sglang.multimodal_gen.runtime.layers.usp import ( + _usp_input_all_to_all_varlen, + _usp_output_all_to_all_varlen, +) +from sglang.multimodal_gen.runtime.layers.visual_embedding import ( + PatchEmbed, + WanCamControlPatchEmbedding, +) +from sglang.multimodal_gen.runtime.managers.forward_context import get_forward_context +from sglang.multimodal_gen.runtime.managers.memory_managers.layerwise_offload import ( + LayerwiseOffloadableModuleMixin, +) +from sglang.multimodal_gen.runtime.models.dits.base import CachableDiT +from sglang.multimodal_gen.runtime.models.dits.causal_wanvideo import ( + CausalWanSelfAttention, + CausalWanTransformer3DModel, + CausalWanTransformerBlock, +) +from sglang.multimodal_gen.runtime.models.dits.wanvideo import ( + WanI2VCrossAttention, + WanT2VCrossAttention, + WanTimeTextImageEmbedding, + WanTransformer3DModel, +) +from sglang.multimodal_gen.runtime.models.utils import _use_aiter +from sglang.multimodal_gen.runtime.platforms import ( + AttentionBackendEnum, + current_platform, +) +from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger +from sglang.srt.utils import add_prefix + +logger = init_logger(__name__) +_is_cuda = current_platform.is_cuda() + +if _use_aiter: + from aiter.ops.rope import rope_cached_2c_fwd_inplace + + +def _compute_sequence_splits(total_len: int, world_size: int) -> list[int]: + base = total_len // world_size + remainder = total_len % world_size + return [base + (1 if rank < remainder else 0) for rank in range(world_size)] + + +def _sequence_shard_tensor( + x: torch.Tensor, seq_splits: list[int], rank: int +) -> torch.Tensor: + start = sum(seq_splits[:rank]) + end = start + seq_splits[rank] + return x[:, start:end, ...].contiguous() + + +def _sequence_all_gather_varlen( + x: torch.Tensor, + seq_splits: list[int], + group: dist.ProcessGroup, +) -> torch.Tensor: + rank = get_sp_parallel_rank() + max_seq = max(seq_splits) + local_seq = seq_splits[rank] + if local_seq < max_seq: + pad_shape = list(x.shape) + pad_shape[1] = max_seq - local_seq + pad = torch.zeros(pad_shape, dtype=x.dtype, device=x.device) + x = torch.cat([x, pad], dim=1) + gathered = [torch.empty_like(x) for _ in seq_splits] + dist.all_gather(gathered, x.contiguous(), group=group) + return torch.cat( + [chunk[:, :seq_len, ...] for chunk, seq_len in zip(gathered, seq_splits)], dim=1 + ) + + +class LingBotWorldCamConditioner(nn.Module): + def __init__(self, dim: int) -> None: + super().__init__() + self.cam_injector = MLP(dim, dim, dim, bias=True, act_type="silu") + self.cam_scale_layer = nn.Linear(dim, dim) + self.cam_shift_layer = nn.Linear(dim, dim) + + def forward( + self, hidden_states: torch.Tensor, c2ws_plucker_emb: torch.Tensor | None + ) -> torch.Tensor: + if c2ws_plucker_emb is None: + return hidden_states + if c2ws_plucker_emb.shape != hidden_states.shape: + raise ValueError( + "c2ws_plucker_emb shape must match hidden_states shape, " + f"got {tuple(c2ws_plucker_emb.shape)} vs {tuple(hidden_states.shape)}" + ) + c2ws_hidden_states = self.cam_injector(c2ws_plucker_emb) + c2ws_hidden_states = c2ws_hidden_states + c2ws_plucker_emb + cam_scale = self.cam_scale_layer(c2ws_hidden_states) + cam_shift = self.cam_shift_layer(c2ws_hidden_states) + return (1.0 + cam_scale) * hidden_states + cam_shift + + +class LingBotWorldCausalSelfAttention(CausalWanSelfAttention): + def __init__(self, *args, **kwargs): + super().__init__(*args, **kwargs) + ulysses_world_size = max(get_ulysses_parallel_world_size(), 1) + if self.num_heads % ulysses_world_size != 0: + raise ValueError( + f"num_heads ({self.num_heads}) must be divisible by ulysses_degree ({ulysses_world_size})." + ) + self.ulysses_num_heads = self.num_heads // ulysses_world_size + self.ulysses_attn = LocalAttention( + num_heads=self.ulysses_num_heads, + head_size=self.head_dim, + dropout_rate=0, + softmax_scale=None, + causal=False, + supported_attention_backends=( + AttentionBackendEnum.FA, + AttentionBackendEnum.AITER, + AttentionBackendEnum.TORCH_SDPA, + ), + ) + + def forward( + self, + q: torch.Tensor, + k: torch.Tensor, + v: torch.Tensor, + freqs_cis: tuple[torch.Tensor, ...], + block_mask, + kv_cache: CausalSelfAttentionKVCache | None = None, + current_start: int = 0, + cache_start: int | None = None, + ): + cos, sin = freqs_cis[:2] + cos_sin_cache = freqs_cis[2] if len(freqs_cis) > 2 else None + if _is_cuda and q.dim() == 4 and q.shape == k.shape: + if cos_sin_cache is None: + cos_sin_cache = torch.cat( + [ + cos.to(dtype=torch.float32).contiguous(), + sin.to(dtype=torch.float32).contiguous(), + ], + dim=-1, + ) + roped_query, roped_key = apply_flashinfer_rope_qk_inplace( + q, k, cos_sin_cache, is_neox=False + ) + roped_query = roped_query.type_as(v) + roped_key = roped_key.type_as(v) + else: + roped_query = _apply_rotary_emb(q, cos, sin, is_neox_style=False).type_as(v) + roped_key = _apply_rotary_emb(k, cos, sin, is_neox_style=False).type_as(v) + forward_batch = get_forward_context().forward_batch + seq_splits = None + sequence_shard_enabled = ( + kv_cache is not None + and forward_batch is not None + and getattr(forward_batch, "enable_sequence_shard", False) + and get_ulysses_parallel_world_size() > 1 + ) + + if kv_cache is None: + if sequence_shard_enabled: + raise NotImplementedError( + "LingBot causal sequence sharding currently requires kv_cache-backed inference." + ) + return super().forward( + q, + k, + v, + (cos, sin), + block_mask, + kv_cache, + current_start, + cache_start, + ) + + if sequence_shard_enabled: + seq_splits = getattr(forward_batch, "sequence_shard_splits", None) + if seq_splits is None: + raise ValueError( + "LingBot causal sequence sharding requires forward_batch.sequence_shard_splits." + ) + seq_splits = list(seq_splits) + # Pack Q/K/V to avoid launching three Ulysses all-to-all collectives. + qkv = torch.cat([roped_query, roped_key, v], dim=-1) + qkv = _usp_input_all_to_all_varlen(qkv, seq_splits, head_dim=2) + roped_query, roped_key, v = qkv.chunk(3, dim=-1) + + cache_view = kv_cache.update_and_get_attention_kv( + key=roped_key, + value=v, + current_chunk_start=current_start, + debug_name="LingBot KV cache", + ) + attn_impl = self.ulysses_attn if sequence_shard_enabled else self.attn + x = attn_impl( + roped_query, + cache_view.k, + cache_view.v, + ) + if sequence_shard_enabled: + assert seq_splits is not None + x = _usp_output_all_to_all_varlen(x, seq_splits, head_dim=2) + return x + + +class LingBotWorldTransformerBlock(nn.Module): + def __init__( + self, + dim: int, + ffn_dim: int, + num_heads: int, + qk_norm: str = "rms_norm_across_heads", + cross_attn_norm: bool = False, + eps: float = 1e-6, + added_kv_proj_dim: int | None = None, + supported_attention_backends=None, + prefix: str = "", + attention_type: str = "original", + sla_topk: float = 0.1, + quant_config: QuantizationConfig | None = None, + ): + super().__init__() + + self.norm1 = LayerNormScaleShift( + dim, + eps=eps, + elementwise_affine=False, + dtype=torch.float32, + ) + self.to_q = ColumnParallelLinear( + dim, + dim, + bias=True, + gather_output=False, + quant_config=quant_config, + prefix=add_prefix("attn1.to_q", prefix), + ) + self.to_k = ColumnParallelLinear( + dim, + dim, + bias=True, + gather_output=False, + quant_config=quant_config, + prefix=add_prefix("attn1.to_k", prefix), + ) + self.to_v = ColumnParallelLinear( + dim, + dim, + bias=True, + gather_output=False, + quant_config=quant_config, + prefix=add_prefix("attn1.to_v", prefix), + ) + self.to_out = RowParallelLinear( + dim, + dim, + bias=True, + reduce_results=True, + quant_config=quant_config, + prefix=add_prefix("attn1.to_out.0", prefix), + ) + + tp_size = get_tp_world_size() + self.local_num_heads = divide(num_heads, tp_size) + self.hidden_dim = dim + self.num_attention_heads = num_heads + self.dim_head = dim // num_heads + self.tp_rmsnorm = qk_norm == "rms_norm_across_heads" and tp_size > 1 + + self.attn1 = USPAttention( + num_heads=self.local_num_heads, + head_size=self.dim_head, + causal=False, + supported_attention_backends=supported_attention_backends, + prefix=add_prefix("attn1", prefix), + quant_config=quant_config, + is_cross_attention=False, + ) + if qk_norm == "rms_norm": + self.norm_q = RMSNorm(self.dim_head, eps=eps) + self.norm_k = RMSNorm(self.dim_head, eps=eps) + elif qk_norm == "rms_norm_across_heads": + self.norm_q = RMSNorm(dim, eps=eps) + self.norm_k = RMSNorm(dim, eps=eps) + else: + raise ValueError(f"Unsupported qk_norm: {qk_norm}") + if not cross_attn_norm: + raise ValueError("LingBotWorld requires cross_attn_norm=True") + self.self_attn_residual_norm = ScaleResidualLayerNormScaleShift( + dim, + eps=eps, + elementwise_affine=True, + dtype=torch.float32, + ) + + cross_attn_backends = { + b for b in supported_attention_backends if not b.is_sparse + } + if added_kv_proj_dim is not None: + self.attn2 = WanI2VCrossAttention( + dim, + num_heads, + qk_norm=qk_norm, + eps=eps, + prefix=add_prefix("attn2", prefix), + supported_attention_backends=cross_attn_backends, + quant_config=quant_config, + ) + else: + self.attn2 = WanT2VCrossAttention( + dim, + num_heads, + qk_norm=qk_norm, + eps=eps, + prefix=add_prefix("attn2", prefix), + supported_attention_backends=cross_attn_backends, + quant_config=quant_config, + ) + self.cross_attn_residual_norm = ScaleResidualLayerNormScaleShift( + dim, + eps=eps, + elementwise_affine=False, + dtype=torch.float32, + ) + self.ffn = MLP( + dim, + ffn_dim, + act_type="gelu_pytorch_tanh", + prefix=add_prefix("ffn.net", prefix), + quant_config=quant_config, + ) + self.mlp_residual = MulAdd() + self.scale_shift_table = nn.Parameter(torch.randn(1, 6, dim) / dim**0.5) + self.cam_conditioner = LingBotWorldCamConditioner(dim) + + def forward( + self, + hidden_states: torch.Tensor, + encoder_hidden_states: torch.Tensor, + temb: torch.Tensor, + freqs_cis: tuple[torch.Tensor, ...], + c2ws_plucker_emb: torch.Tensor | None = None, + ) -> torch.Tensor: + if hidden_states.dim() == 4: + hidden_states = hidden_states.squeeze(1) + orig_dtype = hidden_states.dtype + if temb.dim() == 4: + shift_msa, scale_msa, gate_msa, c_shift_msa, c_scale_msa, c_gate_msa = ( + self.scale_shift_table.unsqueeze(0) + temb.float() + ).chunk(6, dim=2) + shift_msa = shift_msa.squeeze(2) + scale_msa = scale_msa.squeeze(2) + gate_msa = gate_msa.squeeze(2) + c_shift_msa = c_shift_msa.squeeze(2) + c_scale_msa = c_scale_msa.squeeze(2) + c_gate_msa = c_gate_msa.squeeze(2) + else: + shift_msa, scale_msa, gate_msa, c_shift_msa, c_scale_msa, c_gate_msa = ( + self.scale_shift_table + temb.float() + ).chunk(6, dim=1) + + 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.tp_rmsnorm: + query = tensor_parallel_rms_norm(query, self.norm_q) + key = tensor_parallel_rms_norm(key, self.norm_k) + else: + query = self.norm_q(query) + key = self.norm_k(key) + + query = query.squeeze(1).unflatten(2, (self.local_num_heads, self.dim_head)) + key = key.squeeze(1).unflatten(2, (self.local_num_heads, self.dim_head)) + value = value.squeeze(1).unflatten(2, (self.local_num_heads, self.dim_head)) + + cos, sin = freqs_cis + if _is_cuda and query.shape == key.shape: + cos_sin_cache = torch.cat( + [ + cos.to(dtype=torch.float32).contiguous(), + sin.to(dtype=torch.float32).contiguous(), + ], + dim=-1, + ) + query, key = apply_flashinfer_rope_qk_inplace( + query, key, cos_sin_cache, is_neox=False + ) + elif _use_aiter: + query_shape = query.shape + key_shape = key.shape + num_tokens = query.shape[:-2].numel() + q_sbhd = query.view(num_tokens, 1, query_shape[-2], query_shape[-1]) + k_sbhd = key.view(num_tokens, 1, key_shape[-2], key_shape[-1]) + cos_sbhd = cos.contiguous().view(num_tokens, 1, 1, -1) + sin_sbhd = sin.contiguous().view(num_tokens, 1, 1, -1) + rope_cached_2c_fwd_inplace( + q_sbhd, + k_sbhd, + cos_sbhd, + sin_sbhd, + 1, + True, + False, + ) + query = q_sbhd.view(query_shape) + key = k_sbhd.view(key_shape) + else: + query = _apply_rotary_emb(query, cos, sin, is_neox_style=False) + key = _apply_rotary_emb(key, cos, sin, is_neox_style=False) + + attn_output = self.attn1(query, key, value) + attn_output = attn_output.flatten(2) + attn_output, _ = self.to_out(attn_output) + attn_output = attn_output.squeeze(1) + + null_shift = torch.zeros( + (1,), device=hidden_states.device, dtype=hidden_states.dtype + ) + 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 + ) + hidden_states = self.cam_conditioner( + hidden_states.to(orig_dtype), c2ws_plucker_emb + ) + norm_hidden_states = self.self_attn_residual_norm.norm(hidden_states).to( + orig_dtype + ) + + attn_output = self.attn2( + norm_hidden_states, context=encoder_hidden_states, context_lens=None + ) + norm_hidden_states, hidden_states = self.cross_attn_residual_norm( + hidden_states, attn_output, 1, c_shift_msa, c_scale_msa + ) + norm_hidden_states = norm_hidden_states.to(orig_dtype) + hidden_states = hidden_states.to(orig_dtype) + + ff_output = self.ffn(norm_hidden_states) + hidden_states = self.mlp_residual(ff_output, c_gate_msa, hidden_states) + return hidden_states.to(orig_dtype) + + +class LingBotWorldTransformer3DModel(CachableDiT, LayerwiseOffloadableModuleMixin): + _fsdp_shard_conditions = LingBotWorldVideoConfig()._fsdp_shard_conditions + _compile_conditions = LingBotWorldVideoConfig()._compile_conditions + _supported_attention_backends = ( + LingBotWorldVideoConfig()._supported_attention_backends + ) + param_names_mapping = LingBotWorldVideoConfig().param_names_mapping + reverse_param_names_mapping = LingBotWorldVideoConfig().reverse_param_names_mapping + lora_param_names_mapping = LingBotWorldVideoConfig().lora_param_names_mapping + + def __init__( + self, + config: LingBotWorldVideoConfig, + hf_config: dict[str, Any], + quant_config: QuantizationConfig | None = None, + ) -> None: + super().__init__(config=config, hf_config=hf_config) + + inner_dim = config.num_attention_heads * config.attention_head_dim + self.hidden_size = config.hidden_size + self.num_attention_heads = config.num_attention_heads + self.in_channels = config.in_channels + self.out_channels = config.out_channels + self.num_channels_latents = config.num_channels_latents + self.patch_size = config.patch_size + self.text_len = config.text_len + self.patch_embedding = PatchEmbed( + in_chans=config.in_channels, + embed_dim=inner_dim, + patch_size=config.patch_size, + flatten=False, + ) + self.patch_embedding_wancamctrl = WanCamControlPatchEmbedding( + in_chans=6 * 64, + embed_dim=inner_dim, + patch_size=config.patch_size, + ) + self.c2ws_mlp = MLP(inner_dim, inner_dim, inner_dim, bias=True, act_type="silu") + self.condition_embedder = WanTimeTextImageEmbedding( + dim=inner_dim, + time_freq_dim=config.freq_dim, + text_embed_dim=config.text_dim, + image_embed_dim=config.image_dim, + ) + self.blocks = nn.ModuleList( + [ + LingBotWorldTransformerBlock( + inner_dim, + config.ffn_dim, + config.num_attention_heads, + config.qk_norm, + config.cross_attn_norm, + config.eps, + config.added_kv_proj_dim, + self._supported_attention_backends, + prefix=f"blocks.{i}", + quant_config=quant_config, + ) + for i in range(config.num_layers) + ] + ) + self.norm_out = LayerNormScaleShift( + inner_dim, + eps=config.eps, + elementwise_affine=False, + dtype=torch.float32, + ) + self.proj_out = ColumnParallelLinear( + inner_dim, + config.out_channels * math.prod(config.patch_size), + bias=True, + gather_output=True, + prefix="proj_out", + quant_config=quant_config, + ) + self.scale_shift_table = nn.Parameter( + torch.randn(1, 2, inner_dim) / inner_dim**0.5 + ) + self.cnt = 0 + self.__post_init__() + self.sp_size = get_sp_world_size() + d = self.hidden_size // self.num_attention_heads + self.rope_dim_list = [d - 4 * (d // 6), 2 * (d // 6), 2 * (d // 6)] + self.rotary_emb = NDRotaryEmbedding( + rope_dim_list=self.rope_dim_list, + rope_theta=10000, + dtype=( + torch.float64 + if current_platform.is_float64_supported() + else torch.float32 + ), + ) + self.layer_names = ["blocks"] + + @lru_cache(maxsize=1) + def _compute_rope_for_sequence_shard( + self, + local_len: int, + rank: int, + frame_stride_local: int, + width_local: int, + device: torch.device, + ) -> tuple[torch.Tensor, torch.Tensor]: + token_start = rank * local_len + token_indices = torch.arange( + token_start, token_start + local_len, device=device, dtype=torch.long + ) + t_idx = token_indices // frame_stride_local + rem = token_indices % frame_stride_local + h_idx = rem // width_local + w_idx = rem % width_local + positions = torch.stack((t_idx, h_idx, w_idx), dim=1) + return self.rotary_emb.forward_uncached(positions) + + def forward( + self, + hidden_states: torch.Tensor, + encoder_hidden_states: torch.Tensor | list[torch.Tensor], + timestep: torch.LongTensor, + encoder_hidden_states_image: torch.Tensor | list[torch.Tensor] | None = None, + guidance=None, + c2ws_plucker_emb: torch.Tensor | None = None, + **kwargs, + ) -> torch.Tensor: + forward_batch = get_forward_context().forward_batch + sequence_shard_enabled = ( + forward_batch is not None + and forward_batch.enable_sequence_shard + and self.sp_size > 1 + ) + self.enable_teacache = ( + forward_batch is not None and forward_batch.enable_teacache + ) + + orig_dtype = hidden_states.dtype + if not isinstance(encoder_hidden_states, torch.Tensor): + encoder_hidden_states = encoder_hidden_states[0] + if ( + isinstance(encoder_hidden_states_image, list) + and len(encoder_hidden_states_image) > 0 + ): + encoder_hidden_states_image = encoder_hidden_states_image[0] + else: + encoder_hidden_states_image = None + + batch_size, _, num_frames, height, width = hidden_states.shape + p_t, p_h, p_w = self.patch_size + post_patch_num_frames = num_frames // p_t + post_patch_height = height // p_h + post_patch_width = width // p_w + + if not sequence_shard_enabled: + freqs_cos, freqs_sin = self.rotary_emb.forward_from_grid( + ( + post_patch_num_frames * self.sp_size, + post_patch_height, + post_patch_width, + ), + shard_dim=0, + start_frame=0, + device=hidden_states.device, + ) + freqs_cis = (freqs_cos.float(), freqs_sin.float()) + + hidden_states = self.patch_embedding(hidden_states) + hidden_states = hidden_states.flatten(2).transpose(1, 2) + if c2ws_plucker_emb is not None: + c2ws_plucker_emb = self.patch_embedding_wancamctrl( + c2ws_plucker_emb.to( + device=hidden_states.device, dtype=hidden_states.dtype + ) + ) + c2ws_plucker_emb = c2ws_plucker_emb + self.c2ws_mlp(c2ws_plucker_emb) + + seq_len_orig = hidden_states.shape[1] + seq_shard_pad = 0 + if sequence_shard_enabled: + if seq_len_orig % self.sp_size != 0: + seq_shard_pad = self.sp_size - (seq_len_orig % self.sp_size) + pad = torch.zeros( + (batch_size, seq_shard_pad, hidden_states.shape[2]), + dtype=hidden_states.dtype, + device=hidden_states.device, + ) + hidden_states = torch.cat([hidden_states, pad], dim=1) + if c2ws_plucker_emb is not None: + c2ws_plucker_emb = torch.cat([c2ws_plucker_emb, pad], dim=1) + sp_rank = get_sp_group().rank_in_group + local_seq_len = hidden_states.shape[1] // self.sp_size + hidden_states = hidden_states.view( + batch_size, self.sp_size, local_seq_len, hidden_states.shape[2] + )[:, sp_rank, :, :] + if c2ws_plucker_emb is not None: + c2ws_plucker_emb = c2ws_plucker_emb.view( + batch_size, self.sp_size, local_seq_len, c2ws_plucker_emb.shape[2] + )[:, sp_rank, :, :] + frame_stride = post_patch_height * post_patch_width + freqs_cos, freqs_sin = self._compute_rope_for_sequence_shard( + local_seq_len, + sp_rank, + frame_stride, + post_patch_width, + hidden_states.device, + ) + freqs_cis = (freqs_cos.float(), freqs_sin.float()) + + if timestep.dim() == 2: + ts_seq_len = timestep.shape[1] + timestep = timestep.flatten() + else: + ts_seq_len = None + + temb, timestep_proj, encoder_hidden_states, encoder_hidden_states_image = ( + self.condition_embedder( + timestep, + encoder_hidden_states, + encoder_hidden_states_image, + timestep_seq_len=ts_seq_len, + ) + ) + timestep_proj = ( + timestep_proj.unflatten(2, (6, -1)) + if ts_seq_len is not None + else timestep_proj.unflatten(1, (6, -1)) + ) + if sequence_shard_enabled and ts_seq_len is not None: + if seq_shard_pad > 0: + pad = torch.zeros( + ( + batch_size, + seq_shard_pad, + timestep_proj.shape[2], + timestep_proj.shape[3], + ), + dtype=timestep_proj.dtype, + device=timestep_proj.device, + ) + timestep_proj = torch.cat([timestep_proj, pad], dim=1) + timestep_proj = timestep_proj.view( + batch_size, + self.sp_size, + hidden_states.shape[1], + timestep_proj.shape[2], + timestep_proj.shape[3], + )[:, get_sp_parallel_rank(), :, :, :] + + if encoder_hidden_states_image is not None: + encoder_hidden_states = torch.concat( + [encoder_hidden_states_image, encoder_hidden_states], dim=1 + ) + encoder_hidden_states = ( + encoder_hidden_states.to(orig_dtype) + if not current_platform.is_amp_supported() + else encoder_hidden_states + ) + + should_skip_forward = self.should_skip_forward_for_cached_states( + timestep_proj=timestep_proj, temb=temb + ) + if should_skip_forward: + hidden_states = self.retrieve_cached_states(hidden_states) + else: + if self.enable_teacache: + original_hidden_states = hidden_states.clone() + for block in self.blocks: + hidden_states = block( + hidden_states, + encoder_hidden_states, + timestep_proj, + freqs_cis, + c2ws_plucker_emb, + ) + if self.enable_teacache: + self.maybe_cache_states(hidden_states, original_hidden_states) + self.cnt += 1 + + if sequence_shard_enabled: + hidden_states = sequence_model_parallel_all_gather( + hidden_states.contiguous(), dim=1 + ) + if seq_shard_pad > 0: + hidden_states = hidden_states[:, :seq_len_orig, :] + + if temb.dim() == 3: + shift, scale = ( + self.scale_shift_table.unsqueeze(0) + temb.unsqueeze(2) + ).chunk(2, dim=2) + shift = shift.squeeze(2) + scale = scale.squeeze(2) + else: + shift, scale = (self.scale_shift_table + temb.unsqueeze(1)).chunk(2, dim=1) + + hidden_states = self.norm_out(hidden_states, shift, scale) + hidden_states, _ = self.proj_out(hidden_states) + hidden_states = hidden_states.reshape( + batch_size, + post_patch_num_frames, + post_patch_height, + post_patch_width, + p_t, + p_h, + p_w, + -1, + ) + hidden_states = hidden_states.permute(0, 7, 1, 4, 2, 5, 3, 6) + return hidden_states.flatten(6, 7).flatten(4, 5).flatten(2, 3) + + maybe_cache_states = WanTransformer3DModel.maybe_cache_states + should_skip_forward_for_cached_states = ( + WanTransformer3DModel.should_skip_forward_for_cached_states + ) + retrieve_cached_states = WanTransformer3DModel.retrieve_cached_states + + +class CausalLingBotWorldTransformerBlock(CausalWanTransformerBlock): + def __init__(self, *args, **kwargs): + super().__init__(*args, **kwargs) + self.attn1 = LingBotWorldCausalSelfAttention( + dim=self.hidden_dim, + num_heads=self.num_attention_heads, + local_attn_size=self.local_attn_size, + sink_size=self.attn1.sink_size, + qk_norm=self.attn1.qk_norm, + eps=self.attn1.eps, + ) + self.cam_conditioner = LingBotWorldCamConditioner(self.hidden_dim) + self._fused_qkv_weight = None + self._fused_qkv_bias = None + + def _can_fuse_qkv_projection(self) -> bool: + if self._fused_qkv_weight is not None: + return True + + layers = (self.to_q, self.to_k, self.to_v) + biases = [layer.bias for layer in layers] + return ( + all(bias is None for bias in biases) + or all(bias is not None for bias in biases) + ) and all( + getattr(layer, "quant_config", None) is None + and hasattr(layer, "weight") + and layer.weight is not None + for layer in layers + ) + + def fuse_qkv_projection(self) -> bool: + if not self._can_fuse_qkv_projection(): + return False + if self._fused_qkv_weight is not None: + return True + + layers = (self.to_q, self.to_k, self.to_v) + with torch.no_grad(): + self._fused_qkv_weight = nn.Parameter( + torch.cat([layer.weight.detach() for layer in layers], dim=0) + .contiguous() + .to(self.to_q.weight.device), + requires_grad=False, + ) + if all(layer.bias is not None for layer in layers): + self._fused_qkv_bias = nn.Parameter( + torch.cat([layer.bias.detach() for layer in layers], dim=0) + .contiguous() + .to(self.to_q.weight.device), + requires_grad=False, + ) + + return True + + def _project_qkv( + self, hidden_states: torch.Tensor + ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]: + if self.fuse_qkv_projection(): + qkv = F.linear(hidden_states, self._fused_qkv_weight, self._fused_qkv_bias) + return qkv.chunk(3, dim=-1) + + query, _ = self.to_q(hidden_states) + key, _ = self.to_k(hidden_states) + value, _ = self.to_v(hidden_states) + return query, key, value + + 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, ...], + block_mask, + kv_cache: CausalSelfAttentionKVCache | None = None, + crossattn_cache: CrossAttentionKVCache | None = None, + current_start: int = 0, + cache_start: int | None = None, + c2ws_plucker_emb: torch.Tensor | None = None, + ) -> torch.Tensor: + if hidden_states.dim() == 4: + hidden_states = hidden_states.squeeze(1) + num_frames = temb.shape[1] + seqlen_per_frame = hidden_states.shape[1] // num_frames + orig_dtype = hidden_states.dtype + e = self.scale_shift_table + temb.float() + shift_msa, scale_msa, gate_msa, c_shift_msa, c_scale_msa, c_gate_msa = e.chunk( + 6, dim=2 + ) + norm_hidden_states = ( + ( + self.norm1(hidden_states.float()).unflatten( + dim=1, sizes=(num_frames, seqlen_per_frame) + ) + * (1 + scale_msa) + + shift_msa + ) + .flatten(1, 2) + .to(orig_dtype) + ) + query, key, value = self._project_qkv(norm_hidden_states) + query = self.norm_q(query) + 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) + + residual_zero = 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, residual_zero, residual_zero + ) + hidden_states = self.cam_conditioner( + hidden_states.to(orig_dtype), c2ws_plucker_emb + ) + norm_hidden_states = self.self_attn_residual_norm.norm(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 + ) + ff_output = self.ffn(norm_hidden_states.to(orig_dtype)) + hidden_states = self.mlp_residual( + ff_output, c_gate_msa, hidden_states.to(orig_dtype) + ) + return hidden_states.to(orig_dtype) + + +class CausalLingBotWorldTransformer3DModel(CausalWanTransformer3DModel): + _fsdp_shard_conditions = LingBotWorldVideoConfig()._fsdp_shard_conditions + _compile_conditions = LingBotWorldVideoConfig()._compile_conditions + _supported_attention_backends = ( + LingBotWorldVideoConfig()._supported_attention_backends + ) + param_names_mapping = LingBotWorldVideoConfig().param_names_mapping + reverse_param_names_mapping = LingBotWorldVideoConfig().reverse_param_names_mapping + lora_param_names_mapping = LingBotWorldVideoConfig().lora_param_names_mapping + + def __init__( + self, + config: LingBotWorldVideoConfig, + 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.patch_embedding_wancamctrl = WanCamControlPatchEmbedding( + in_chans=6 * 64, + embed_dim=inner_dim, + patch_size=config.patch_size, + ) + self.c2ws_mlp = MLP(inner_dim, inner_dim, inner_dim, bias=True, act_type="silu") + self.sp_size = get_sp_world_size() + d = self.hidden_size // self.num_attention_heads + self.rope_dim_list = [d - 4 * (d // 6), 2 * (d // 6), 2 * (d // 6)] + self.rotary_emb = NDRotaryEmbedding( + rope_dim_list=self.rope_dim_list, + rope_theta=10000, + dtype=( + torch.float64 + if current_platform.is_float64_supported() + else torch.float32 + ), + ) + self.blocks = nn.ModuleList( + [ + CausalLingBotWorldTransformerBlock( + 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) + ] + ) + + def post_load_weights(self) -> None: + super().post_load_weights() + for block in self.blocks: + block.fuse_qkv_projection() + + @lru_cache(maxsize=8) + def _compute_rope_for_sequence_shard_with_offset( + self, + local_len: int, + token_start: int, + frame_stride_local: int, + width_local: int, + device: torch.device, + ) -> tuple[torch.Tensor, torch.Tensor]: + token_indices = torch.arange( + token_start, + token_start + local_len, + device=device, + dtype=torch.long, + ) + t_idx = token_indices // frame_stride_local + rem = token_indices % frame_stride_local + h_idx = rem // width_local + w_idx = rem % width_local + positions = torch.stack((t_idx, h_idx, w_idx), dim=1) + return self.rotary_emb.forward_uncached(positions) + + def _prepare_c2ws_plucker_emb( + self, + hidden_states: torch.Tensor, + c2ws_plucker_emb: torch.Tensor | None, + forward_batch=None, + ) -> torch.Tensor | None: + if c2ws_plucker_emb is None: + return None + cache = self._get_request_cache(forward_batch, "lingbot_c2ws_plucker_emb") + cache_key = ( + c2ws_plucker_emb.data_ptr(), + tuple(c2ws_plucker_emb.shape), + tuple(c2ws_plucker_emb.stride()), + c2ws_plucker_emb.dtype, + c2ws_plucker_emb.device.type, + c2ws_plucker_emb.device.index, + hidden_states.dtype, + hidden_states.device.type, + hidden_states.device.index, + ) + if cache is not None and cache_key in cache: + return cache[cache_key] + + c2ws_plucker_emb = self.patch_embedding_wancamctrl( + c2ws_plucker_emb.to(device=hidden_states.device, dtype=hidden_states.dtype) + ) + c2ws_plucker_emb = c2ws_plucker_emb + self.c2ws_mlp(c2ws_plucker_emb) + if cache is not None: + cache.clear() + cache[cache_key] = c2ws_plucker_emb + return c2ws_plucker_emb + + @staticmethod + def _get_request_cache(forward_batch, name: str) -> dict | None: + if forward_batch is None: + return None + extra = getattr(forward_batch, "extra", None) + if extra is None: + return None + return extra.setdefault(name, {}) + + @staticmethod + def _all_crossattn_caches_initialized( + crossattn_cache: list[CrossAttentionKVCache] | None, + ) -> bool: + return crossattn_cache is not None and all( + cache.is_init for cache in crossattn_cache + ) + + def _prepare_cached_rope( + self, + *, + forward_batch, + post_patch_num_frames: int, + post_patch_height: int, + post_patch_width: int, + start_frame: int, + device: torch.device, + ) -> tuple[torch.Tensor, ...]: + cache = self._get_request_cache(forward_batch, "lingbot_rope") + cache_key = ( + post_patch_num_frames, + post_patch_height, + post_patch_width, + start_frame, + device.type, + device.index, + ) + if cache is not None and cache_key in cache: + return cache[cache_key] + + freqs_cos, freqs_sin = get_rotary_pos_embed( + ( + post_patch_num_frames * get_sp_world_size(), + post_patch_height, + post_patch_width, + ), + self.hidden_size, + self.num_attention_heads, + self.rope_dim_list, + dtype=( + torch.float64 + if current_platform.is_float64_supported() + else torch.float32 + ), + rope_theta=10000, + start_frame=start_frame, + ) + freqs_cos = freqs_cos.to(device).float() + freqs_sin = freqs_sin.to(device).float() + freqs_cis: tuple[torch.Tensor, ...] = (freqs_cos, freqs_sin) + if _is_cuda: + freqs_cis = ( + freqs_cos, + freqs_sin, + torch.cat([freqs_cos.contiguous(), freqs_sin.contiguous()], dim=-1), + ) + if cache is not None: + cache.clear() + cache[cache_key] = freqs_cis + return freqs_cis + + def _prepare_condition_embeddings( + self, + *, + timestep: torch.LongTensor, + encoder_hidden_states: torch.Tensor, + encoder_hidden_states_image: torch.Tensor | None, + crossattn_cache: list[CrossAttentionKVCache] | None, + ): + if self._all_crossattn_caches_initialized(crossattn_cache): + temb = self.condition_embedder.time_embedder(timestep.flatten()) + timestep_proj = self.condition_embedder.time_modulation(temb) + return temb, timestep_proj, encoder_hidden_states, None + + return self.condition_embedder( + timestep.flatten(), encoder_hidden_states, encoder_hidden_states_image + ) + + def forward( + self, + hidden_states: torch.Tensor, + encoder_hidden_states: torch.Tensor | list[torch.Tensor], + timestep: torch.LongTensor, + encoder_hidden_states_image: torch.Tensor | list[torch.Tensor] | None = None, + kv_cache: list[CausalSelfAttentionKVCache] | None = None, + crossattn_cache: list[CrossAttentionKVCache] | None = None, + current_start: int = 0, + cache_start: int = 0, + start_frame: int = 0, + c2ws_plucker_emb: torch.Tensor | None = None, + skip_final_projection: bool = False, + ) -> torch.Tensor: + forward_batch = get_forward_context().forward_batch + sequence_shard_enabled = ( + forward_batch is not None + and getattr(forward_batch, "enable_sequence_shard", False) + and self.sp_size > 1 + ) + if sequence_shard_enabled: + if get_ring_parallel_world_size() > 1: + raise NotImplementedError( + "LingBot causal sequence sharding currently supports ulysses_degree > 1 with ring_degree = 1 only." + ) + if get_ulysses_parallel_world_size() <= 1: + raise NotImplementedError( + "LingBot causal sequence sharding currently requires ulysses_degree > 1." + ) + orig_dtype = hidden_states.dtype + if not isinstance(encoder_hidden_states, torch.Tensor): + encoder_hidden_states = encoder_hidden_states[0] + if ( + isinstance(encoder_hidden_states_image, list) + and len(encoder_hidden_states_image) > 0 + ): + encoder_hidden_states_image = encoder_hidden_states_image[0] + else: + encoder_hidden_states_image = None + + batch_size, _, num_frames, height, width = hidden_states.shape + p_t, p_h, p_w = self.patch_size + post_patch_num_frames = num_frames // p_t + post_patch_height = height // p_h + post_patch_width = width // p_w + if sequence_shard_enabled: + seq_shard_splits = _compute_sequence_splits( + post_patch_num_frames * post_patch_height * post_patch_width, + self.sp_size, + ) + forward_batch.sequence_shard_splits = tuple(seq_shard_splits) + if not sequence_shard_enabled: + freqs_cis = self._prepare_cached_rope( + forward_batch=forward_batch, + post_patch_num_frames=post_patch_num_frames, + post_patch_height=post_patch_height, + post_patch_width=post_patch_width, + start_frame=start_frame, + device=hidden_states.device, + ) + + hidden_states = self.patch_embedding(hidden_states).flatten(2).transpose(1, 2) + c2ws_plucker_emb = self._prepare_c2ws_plucker_emb( + hidden_states, c2ws_plucker_emb, forward_batch + ) + if sequence_shard_enabled: + sp_rank = get_sp_parallel_rank() + seq_shard_splits = list(forward_batch.sequence_shard_splits) + local_seq_len = seq_shard_splits[sp_rank] + hidden_states = _sequence_shard_tensor( + hidden_states, seq_shard_splits, sp_rank + ) + if c2ws_plucker_emb is not None: + c2ws_plucker_emb = _sequence_shard_tensor( + c2ws_plucker_emb, seq_shard_splits, sp_rank + ) + frame_stride = post_patch_height * post_patch_width + token_start = start_frame * frame_stride + sum(seq_shard_splits[:sp_rank]) + freqs_cos, freqs_sin = self._compute_rope_for_sequence_shard_with_offset( + local_seq_len, + token_start, + frame_stride, + post_patch_width, + hidden_states.device, + ) + freqs_cos = freqs_cos.float() + freqs_sin = freqs_sin.float() + freqs_cis = (freqs_cos, freqs_sin) + if _is_cuda: + freqs_cis = ( + freqs_cos, + freqs_sin, + torch.cat([freqs_cos.contiguous(), freqs_sin.contiguous()], dim=-1), + ) + + temb, timestep_proj, encoder_hidden_states, encoder_hidden_states_image = ( + self._prepare_condition_embeddings( + timestep=timestep, + encoder_hidden_states=encoder_hidden_states, + encoder_hidden_states_image=encoder_hidden_states_image, + crossattn_cache=crossattn_cache, + ) + ) + timestep_proj = timestep_proj.unflatten(1, (6, self.hidden_size)).unflatten( + dim=0, sizes=timestep.shape + ) + if encoder_hidden_states_image is not None: + encoder_hidden_states = torch.concat( + [encoder_hidden_states_image, encoder_hidden_states], dim=1 + ) + encoder_hidden_states = ( + encoder_hidden_states.to(orig_dtype) + if current_platform.is_mps() + else encoder_hidden_states + ) + + for block_index, block in enumerate(self.blocks): + hidden_states = block( + hidden_states, + encoder_hidden_states, + timestep_proj, + freqs_cis, + block_mask=self.block_mask, + kv_cache=kv_cache[block_index], + crossattn_cache=crossattn_cache[block_index], + current_start=current_start, + cache_start=cache_start, + c2ws_plucker_emb=c2ws_plucker_emb, + ) + + if skip_final_projection: + return hidden_states + + if sequence_shard_enabled: + hidden_states = _sequence_all_gather_varlen( + hidden_states.contiguous(), + list(forward_batch.sequence_shard_splits), + get_sp_group().device_group, + ) + + temb = temb.unflatten(dim=0, sizes=timestep.shape).unsqueeze(2) + shift, scale = (self.scale_shift_table.unsqueeze(1) + temb).chunk(2, dim=2) + hidden_states = self.norm_out(hidden_states, shift, scale) + hidden_states = self.proj_out(hidden_states) + hidden_states = hidden_states.reshape( + batch_size, + post_patch_num_frames, + post_patch_height, + post_patch_width, + p_t, + p_h, + p_w, + -1, + ) + hidden_states = hidden_states.permute(0, 7, 1, 4, 2, 5, 3, 6) + return hidden_states.flatten(6, 7).flatten(4, 5).flatten(2, 3) + + +EntryClass = [ + LingBotWorldTransformer3DModel, + CausalLingBotWorldTransformer3DModel, +] diff --git a/python/sglang/multimodal_gen/runtime/models/vaes/wanvae.py b/python/sglang/multimodal_gen/runtime/models/vaes/wanvae.py index 7279c2ed8..23cbeede7 100644 --- a/python/sglang/multimodal_gen/runtime/models/vaes/wanvae.py +++ b/python/sglang/multimodal_gen/runtime/models/vaes/wanvae.py @@ -863,6 +863,7 @@ class AutoencoderKLWan(ParallelTiledVAE): ) self.use_feature_cache = config.use_feature_cache + self._causal_decode_initialized = False def clear_cache(self) -> None: @@ -883,6 +884,41 @@ class AutoencoderKLWan(ParallelTiledVAE): self._enc_conv_idx = 0 self._enc_feat_map = [None] * self._enc_conv_num + def reset_causal_decode_state(self) -> None: + """Reset decoder feature cache before a new causal video session.""" + self._causal_decode_initialized = False + if self.use_feature_cache: + self.clear_cache() + + def causal_decode(self, z: torch.Tensor) -> torch.Tensor: + """Decode latents while preserving decoder feature cache across chunks.""" + if not self.use_feature_cache: + return self.decode(z) + + is_first_chunk = not self._causal_decode_initialized + if is_first_chunk: + self.clear_cache() + + iter_ = z.shape[2] + x = self.post_quant_conv(z) + outs = [] + with forward_context( + feat_cache_arg=self._feat_map, feat_idx_arg=self._conv_idx + ): + for i in range(iter_): + feat_idx.set(0) + first_chunk.set(is_first_chunk and i == 0) + outs.append(self.decoder(x[:, :, i : i + 1, :, :])) + out = torch.cat(outs, 2) + + if self.config.patch_size is not None: + out = unpatchify(out, patch_size=self.config.patch_size) + + out = out.float() + out = torch.clamp(out, min=-1.0, max=1.0) + self._causal_decode_initialized = True + return out + def encode(self, x: torch.Tensor) -> torch.Tensor: if self.use_feature_cache: self.clear_cache() @@ -946,6 +982,7 @@ class AutoencoderKLWan(ParallelTiledVAE): self.clear_cache() iter_ = z.shape[2] x = self.post_quant_conv(z) + outs = [] with forward_context( feat_cache_arg=self._feat_map, feat_idx_arg=self._conv_idx ): @@ -953,11 +990,10 @@ class AutoencoderKLWan(ParallelTiledVAE): feat_idx.set(0) if i == 0: first_chunk.set(True) - out = self.decoder(x[:, :, i : i + 1, :, :]) else: first_chunk.set(False) - out_ = self.decoder(x[:, :, i : i + 1, :, :]) - out = torch.cat([out, out_], 2) + outs.append(self.decoder(x[:, :, i : i + 1, :, :])) + out = torch.cat(outs, 2) if self.config.patch_size is not None: out = unpatchify(out, patch_size=self.config.patch_size) diff --git a/python/sglang/multimodal_gen/runtime/pipelines/lingbot_world_causal_dmd_pipeline.py b/python/sglang/multimodal_gen/runtime/pipelines/lingbot_world_causal_dmd_pipeline.py new file mode 100644 index 000000000..e3f1300f5 --- /dev/null +++ b/python/sglang/multimodal_gen/runtime/pipelines/lingbot_world_causal_dmd_pipeline.py @@ -0,0 +1,99 @@ +# Copied and adapted from: https://github.com/hao-ai-lab/FastVideo +# Adapted from: https://github.com/Robbyant/lingbot-world + +# SPDX-License-Identifier: Apache-2.0 +""" +LingBot-World realtime causal DMD pipeline. +""" + +from sglang.multimodal_gen.runtime.models.schedulers.scheduling_self_forcing_flow_match import ( + SelfForcingFlowMatchScheduler, +) +from sglang.multimodal_gen.runtime.pipelines_core.composed_pipeline_base import ( + ComposedPipelineBase, +) +from sglang.multimodal_gen.runtime.pipelines_core.lora_pipeline import LoRAPipeline +from sglang.multimodal_gen.runtime.pipelines_core.stages import ( + AuxiliaryConditionEncodingStage, + CausalVaeDecodingStage, + DMDTimestepPreparationStage, + ImageEncodingStage, + RealtimeChunkLatentPreparationStage, + RealtimeImageVAEEncodingStage, + RealtimeInputValidationStage, + RealtimeTextEncodingStage, +) +from sglang.multimodal_gen.runtime.pipelines_core.stages.model_specific_stages.lingbot_world import ( + LingBotWorldCausalDMDDenoisingStage, +) +from sglang.multimodal_gen.runtime.server_args import ServerArgs + + +class LingBotWorldCausalDMDPipeline(LoRAPipeline, ComposedPipelineBase): + pipeline_name = "LingBotWorldCausalDMDPipeline" + + _required_config_modules = [ + "text_encoder", + "tokenizer", + "vae", + "transformer", + "scheduler", + "image_encoder", + "image_processor", + ] + + def initialize_pipeline(self, server_args: ServerArgs): + self.modules["scheduler"] = SelfForcingFlowMatchScheduler( + num_inference_steps=1000, + shift=server_args.pipeline_config.flow_shift, + sigma_min=0.0, + extra_one_step=True, + ) + + def create_pipeline_stages(self, server_args) -> None: + self.add_stage(RealtimeInputValidationStage()) + self.add_stage( + RealtimeTextEncodingStage( + text_encoders=[self.get_module("text_encoder")], + tokenizers=[self.get_module("tokenizer")], + ) + ) + + image_encoder = self.get_module("image_encoder", None) + image_processor = self.get_module("image_processor", None) + self.add_stage_if( + image_encoder is not None and image_processor is not None, + ImageEncodingStage( + image_encoder=image_encoder, + image_processor=image_processor, + ), + ) + + self.add_stage(AuxiliaryConditionEncodingStage()) + self.add_stage( + RealtimeImageVAEEncodingStage( + vae=self.get_module("vae"), + ) + ) + self.add_stage(DMDTimestepPreparationStage(self.get_module("scheduler"))) + self.add_stage( + RealtimeChunkLatentPreparationStage( + scheduler=self.get_module("scheduler"), + transformer=self.get_module("transformer"), + ) + ) + self.add_stage( + LingBotWorldCausalDMDDenoisingStage( + transformer=self.get_module("transformer"), + scheduler=self.get_module("scheduler"), + ), + ) + self.add_stage( + CausalVaeDecodingStage( + vae=self.get_module("vae"), + pipeline=self, + ) + ) + + +EntryClass = LingBotWorldCausalDMDPipeline diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/schedule_batch.py b/python/sglang/multimodal_gen/runtime/pipelines_core/schedule_batch.py index 0491dc24f..de4b14a9a 100644 --- a/python/sglang/multimodal_gen/runtime/pipelines_core/schedule_batch.py +++ b/python/sglang/multimodal_gen/runtime/pipelines_core/schedule_batch.py @@ -17,7 +17,7 @@ import pprint from collections import Counter from copy import deepcopy from dataclasses import MISSING, asdict, dataclass, field, fields -from typing import Any, Optional, Union +from typing import Any, Optional, Sequence, Union import PIL.Image import torch @@ -26,6 +26,9 @@ from sglang.multimodal_gen.configs.sample.sampling_params import SamplingParams from sglang.multimodal_gen.runtime.post_training.rl_dataclasses import ( RolloutTrajectoryData, ) +from sglang.multimodal_gen.runtime.realtime.session import ( + RealtimeSession, +) from sglang.multimodal_gen.runtime.server_args import ServerArgs from sglang.multimodal_gen.runtime.utils.logging_utils import ( _sanitize_for_logging, @@ -141,6 +144,7 @@ class Req: image_latent: torch.Tensor | list[torch.Tensor] | None = None condition_image_latent_ids: torch.Tensor | list[torch.Tensor] | None = None vae_image_sizes: list[tuple[int, int]] | None = None + c2ws_plucker_emb: torch.Tensor | None = None # Latent dimensions height_latents: list[int] | int | None = None @@ -179,6 +183,8 @@ class Req: # Extra parameters that might be needed by specific pipeline implementations (e.g., LTX2.3 DenoisingAVStage) extra: dict[str, Any] = field(default_factory=dict) + condition_inputs: dict[str, Any] = field(default_factory=dict) + is_warmup: bool = False # STA parameters @@ -198,6 +204,18 @@ class Req: default_factory=TraceNullContext ) + # realtime + realtime_session_id: str | None = None + session: RealtimeSession | None = None + block_idx: int = 0 + realtime_chunk_size: int | None = None + realtime_event_id: int | None = None + realtime_output_format: str | None = None + realtime_causal_sink_size: int | None = None + realtime_causal_kv_cache_num_frames: int | None = None + # return websocket-friendly raw RGB frame bytes instead of rwa tensors + return_raw_frames: bool = False + # results output: torch.Tensor | None = None audio: torch.Tensor | None = None @@ -386,7 +404,11 @@ class OutputBatch: Final output (after pipeline completion) """ - output: Any | None = None + # tensors or numpy frames + output: Sequence[Any] | None = None + raw_frame_batches: list[list[bytes]] | None = None + raw_frame_content_type: str = "application/x-raw-rgb" + raw_frame_metadata: dict[str, Any] | None = None audio: torch.Tensor | None = None audio_sample_rate: int | None = None trajectory_timesteps: torch.Tensor | None = None diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/__init__.py b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/__init__.py index 09adde89a..e1157c507 100644 --- a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/__init__.py +++ b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/__init__.py @@ -15,6 +15,10 @@ from sglang.multimodal_gen.runtime.pipelines_core.stages.causal_denoising import from sglang.multimodal_gen.runtime.pipelines_core.stages.comfyui_latent_preparation import ( ComfyUILatentPreparationStage, ) +from sglang.multimodal_gen.runtime.pipelines_core.stages.condition_encoding import ( + AuxiliaryConditionEncodingStage, + ConditionEncodingStage, +) from sglang.multimodal_gen.runtime.pipelines_core.stages.decoding import DecodingStage from sglang.multimodal_gen.runtime.pipelines_core.stages.decoding_av import ( LTX2AVDecodingStage, @@ -53,6 +57,7 @@ from sglang.multimodal_gen.runtime.pipelines_core.stages.input_validation import ) from sglang.multimodal_gen.runtime.pipelines_core.stages.latent_preparation import ( LatentPreparationStage, + RealtimeChunkLatentPreparationStage, ) from sglang.multimodal_gen.runtime.pipelines_core.stages.latent_preparation_av import ( LTX2AVLatentPreparationStage, @@ -60,6 +65,16 @@ from sglang.multimodal_gen.runtime.pipelines_core.stages.latent_preparation_av i from sglang.multimodal_gen.runtime.pipelines_core.stages.ltx_2_denoising import ( LTX2DenoisingStage, ) +from sglang.multimodal_gen.runtime.pipelines_core.stages.realtime_input_validation import ( + RealtimeInputValidationStage, +) +from sglang.multimodal_gen.runtime.pipelines_core.stages.realtime_text_encoding import ( + RealtimeTextEncodingStage, +) +from sglang.multimodal_gen.runtime.pipelines_core.stages.realtime_vae import ( + CausalVaeDecodingStage, + RealtimeImageVAEEncodingStage, +) from sglang.multimodal_gen.runtime.pipelines_core.stages.text_connector import ( LTX2TextConnectorStage, ) @@ -67,6 +82,7 @@ from sglang.multimodal_gen.runtime.pipelines_core.stages.text_encoding import ( TextEncodingStage, ) from sglang.multimodal_gen.runtime.pipelines_core.stages.timestep_preparation import ( + DMDTimestepPreparationStage, TimestepPreparationStage, ) from sglang.multimodal_gen.runtime.pipelines_core.stages.upsampling import ( @@ -78,8 +94,11 @@ from sglang.multimodal_gen.runtime.pipelines_core.stages.upsampling import ( __all__ = [ "PipelineStage", "InputValidationStage", + "RealtimeInputValidationStage", "TimestepPreparationStage", + "DMDTimestepPreparationStage", "LatentPreparationStage", + "RealtimeChunkLatentPreparationStage", "ComfyUILatentPreparationStage", "LTX2AVLatentPreparationStage", "DenoisingStage", @@ -88,12 +107,17 @@ __all__ = [ "LTX2AVDenoisingStage", "CausalDMDDenoisingStage", "EncodingStage", + "ConditionEncodingStage", + "AuxiliaryConditionEncodingStage", "DecodingStage", + "CausalVaeDecodingStage", "LTX2AVDecodingStage", "ImageEncodingStage", "ImageVAEEncodingStage", + "RealtimeImageVAEEncodingStage", "LTX2ImageEncodingStage", "TextEncodingStage", + "RealtimeTextEncodingStage", "LTX2TextConnectorStage", # Hunyuan3D shape stages "Hunyuan3DShapeBeforeDenoisingStage", 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 414968e66..eaae177dc 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 @@ -1,8 +1,17 @@ # Copied and adapted from: https://github.com/hao-ai-lab/FastVideo +from collections.abc import Callable +from contextlib import nullcontext +from dataclasses import dataclass +from typing import Any + import torch # type: ignore from sglang.multimodal_gen.runtime.distributed import get_local_torch_device +from sglang.multimodal_gen.runtime.layers.kvcache.causal_attention_cache import ( + CausalSelfAttentionKVCache, + CrossAttentionKVCache, +) from sglang.multimodal_gen.runtime.managers.forward_context import set_forward_context from sglang.multimodal_gen.runtime.models.utils import pred_noise_to_pred_video from sglang.multimodal_gen.runtime.pipelines_core.diffusion_scheduler_utils import ( @@ -20,12 +29,50 @@ from sglang.multimodal_gen.runtime.platforms import ( AttentionBackendEnum, current_platform, ) +from sglang.multimodal_gen.runtime.realtime.causal_state import RealtimeCausalDiTState from sglang.multimodal_gen.runtime.server_args import ServerArgs from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger logger = init_logger(__name__) +@dataclass(slots=True) +class CausalDMDForwardContext: + target_dtype: torch.dtype + autocast_enabled: bool + device: torch.device + scheduler: Any + timesteps: torch.Tensor + latents: torch.Tensor + prompt_embeds: Any + image_kwargs: dict[str, Any] + pos_cond_kwargs: dict[str, Any] + batch_size: int + channels: int + num_frames: int + height: int + width: int + + +@dataclass(slots=True) +class CausalDMDCachePolicy: + sequence_shard_enabled: bool + num_attention_heads: int + expected_cache_tokens: int + expected_sink_tokens: int + kv_cache_kwargs: dict[str, Any] + + +@dataclass(slots=True) +class CausalDMDRealtimeCacheContext: + cache_state: RealtimeCausalDiTState + persist_state: bool + kv_cache: list[CausalSelfAttentionKVCache] + crossattn_cache: list[CrossAttentionKVCache] + current_start_frame: int + chunk_idx: int + + class CausalDMDDenoisingStage(DenoisingStage): """ Denoising stage for causal diffusion. @@ -34,7 +81,7 @@ class CausalDMDDenoisingStage(DenoisingStage): def __init__(self, transformer, scheduler) -> None: super().__init__(transformer, scheduler) # KV and cross-attention cache state (initialized on first forward) - self.kv_cache1: list | None = None + self.causal_kv_cache: list | None = None self.crossattn_cache: list | None = None # Model-dependent constants (aligned with causal_inference.py assumptions) self.num_transformer_blocks = self.transformer.config.arch_config.num_layers @@ -51,28 +98,43 @@ class CausalDMDDenoisingStage(DenoisingStage): ) # type: ignore except Exception: self.local_attn_size = -1 + self.sink_size = self.transformer.config.arch_config.sink_size - def forward( + self._causal_attn_metadata_builder_cls = None + self._causal_attn_metadata_builder = None + + def _target_dtype(self) -> torch.dtype: + return torch.bfloat16 + + def _autocast_enabled( self, - batch: Req, + target_dtype: torch.dtype, server_args: ServerArgs, - ) -> Req: - target_dtype = torch.bfloat16 - autocast_enabled = ( - target_dtype != torch.float32 - ) and not server_args.disable_autocast - scheduler = get_or_create_request_scheduler(batch, self.scheduler) + ) -> bool: + return (target_dtype != torch.float32) and not server_args.disable_autocast - latent_seq_length = batch.latents.shape[-1] * batch.latents.shape[-2] + def _prepare_frame_seq_length(self, h: int, w: int) -> int: patch_ratio = ( self.transformer.config.arch_config.patch_size[-1] * self.transformer.config.arch_config.patch_size[-2] ) - self.frame_seq_length = latent_seq_length // patch_ratio - # TODO(will): make this a parameter once we add i2v support - independent_first_frame = self.transformer.independent_first_frame + self.num_token_per_frame = (h * w) // patch_ratio + return self.num_token_per_frame - # Timesteps for DMD + def _get_causal_dmd_latents(self, batch: Req) -> torch.Tensor: + assert batch.latents is not None, "latents must be provided" + return batch.latents + + def _get_causal_dmd_scheduler(self, batch: Req, server_args: ServerArgs): + return get_or_create_request_scheduler(batch, self.scheduler) + + def _prepare_causal_dmd_timesteps( + self, + batch: Req, + server_args: ServerArgs, + scheduler, + device: torch.device, + ) -> torch.Tensor: timesteps = torch.tensor( server_args.pipeline_config.dmd_denoising_steps, dtype=torch.long ).cpu() @@ -83,13 +145,25 @@ class CausalDMDDenoisingStage(DenoisingStage): (scheduler.timesteps.cpu(), torch.tensor([0], dtype=torch.float32)) ) timesteps = scheduler_timesteps[1000 - timesteps] - timesteps = timesteps.to(get_local_torch_device()) + timesteps = timesteps.to(device) logger.info("Using timesteps: %s", timesteps) + return timesteps - # Image kwargs (kept empty unless caller provides compatible args) - image_kwargs: dict = {} + def _prepare_causal_dmd_image_kwargs( + self, + batch: Req, + server_args: ServerArgs, + target_dtype: torch.dtype, + ) -> dict[str, Any]: + return {} - pos_cond_kwargs = self.prepare_extra_func_kwargs( + def _prepare_causal_dmd_pos_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_hidden_states_2": batch.clip_embedding_pos, @@ -97,47 +171,782 @@ class CausalDMDDenoisingStage(DenoisingStage): }, ) - # STA + def _prepare_causal_dmd_prompt_embeds( + self, + batch: Req, + server_args: ServerArgs, + target_dtype: torch.dtype, + ): + prompt_embeds = batch.prompt_embeds + assert torch.isnan(prompt_embeds[0]).sum() == 0 + return prompt_embeds + + def _prepare_causal_dmd_forward_context( + self, + batch: Req, + server_args: ServerArgs, + ) -> CausalDMDForwardContext: + target_dtype = self._target_dtype() + autocast_enabled = self._autocast_enabled(target_dtype, server_args) + device = get_local_torch_device() + scheduler = self._get_causal_dmd_scheduler(batch, server_args) + latents = self._get_causal_dmd_latents(batch) + b, c, t, h, w = latents.shape + self._prepare_frame_seq_length(h, w) + timesteps = self._prepare_causal_dmd_timesteps( + batch, + server_args, + scheduler, + device, + ) + image_kwargs = self._prepare_causal_dmd_image_kwargs( + batch, + server_args, + target_dtype, + ) + pos_cond_kwargs = self._prepare_causal_dmd_pos_cond_kwargs( + batch, + server_args, + target_dtype, + ) + if self.attn_backend.get_enum() == AttentionBackendEnum.SLIDING_TILE_ATTN: self.prepare_sta_param(batch, server_args) - # Latents and prompts - assert batch.latents is not None, "latents must be provided" - latents = batch.latents # [B, C, T, H, W] - b, c, t, h, w = latents.shape - prompt_embeds = batch.prompt_embeds - assert torch.isnan(prompt_embeds[0]).sum() == 0 + prompt_embeds = self._prepare_causal_dmd_prompt_embeds( + batch, + server_args, + target_dtype, + ) + return CausalDMDForwardContext( + target_dtype=target_dtype, + autocast_enabled=autocast_enabled, + device=device, + scheduler=scheduler, + timesteps=timesteps, + latents=latents, + prompt_embeds=prompt_embeds, + image_kwargs=image_kwargs, + pos_cond_kwargs=pos_cond_kwargs, + batch_size=b, + channels=c, + num_frames=t, + height=h, + width=w, + ) + + def _apply_causal_cache_overrides( + self, + batch: Req, + server_args: ServerArgs, + ) -> None: + pipeline_config = server_args.pipeline_config + sink_size = getattr(batch, "realtime_causal_sink_size", None) + if sink_size is None: + sink_size = getattr(pipeline_config, "realtime_causal_sink_size", None) + if sink_size is not None: + if sink_size < 0: + raise ValueError("realtime_causal_sink_size must be non-negative") + self.sink_size = int(sink_size) + + kv_cache_num_frames = getattr( + batch, + "realtime_causal_kv_cache_num_frames", + None, + ) + if kv_cache_num_frames is None: + kv_cache_num_frames = getattr( + pipeline_config, + "realtime_causal_kv_cache_num_frames", + None, + ) + if kv_cache_num_frames is not None: + if kv_cache_num_frames <= 0: + raise ValueError("realtime_causal_kv_cache_num_frames must be positive") + self.sliding_window_num_frames = int(kv_cache_num_frames) + + def _causal_sequence_shard_enabled(self, batch: Req) -> bool: + return False + + def _num_causal_cache_attention_heads( + self, + *, + sequence_shard_enabled: bool, + ) -> int: + return self.transformer.num_attention_heads + + def _causal_kv_cache_kwargs( + self, + policy: CausalDMDCachePolicy, + ) -> dict[str, Any]: + return {} + + def _use_causal_cache_int_indices( + self, + *, + sequence_shard_enabled: bool, + ) -> bool: + return False + + def _build_realtime_causal_cache_policy( + self, + batch: Req, + server_args: ServerArgs, + ) -> CausalDMDCachePolicy: + self._apply_causal_cache_overrides(batch, server_args) + sequence_shard_enabled = self._causal_sequence_shard_enabled(batch) + policy = CausalDMDCachePolicy( + sequence_shard_enabled=sequence_shard_enabled, + num_attention_heads=self._num_causal_cache_attention_heads( + sequence_shard_enabled=sequence_shard_enabled + ), + expected_cache_tokens=self._get_causal_kv_cache_size( + sequence_shard_enabled=sequence_shard_enabled + ), + expected_sink_tokens=self._get_causal_sink_tokens(), + kv_cache_kwargs={}, + ) + policy.kv_cache_kwargs = self._causal_kv_cache_kwargs(policy) + return policy + + def _get_realtime_causal_cache_state( + self, + batch: Req, + ) -> tuple[RealtimeCausalDiTState, bool]: + if batch.session is not None: + state = batch.session.get_or_create_state(RealtimeCausalDiTState) + return state, True + return RealtimeCausalDiTState(), False + + def _clear_stage_causal_cache_refs(self) -> None: + self.causal_kv_cache = None + self.crossattn_cache = None + + def _should_reset_realtime_causal_caches( + self, + batch: Req, + *, + cache_state: RealtimeCausalDiTState, + policy: CausalDMDCachePolicy, + ) -> bool: + causal_kv_cache = cache_state.kv_cache + crossattn_cache = cache_state.crossattn_cache + return ( + batch.block_idx == 0 + or causal_kv_cache is None + or crossattn_cache is None + or len(causal_kv_cache) != self.num_transformer_blocks + or len(crossattn_cache) != self.num_transformer_blocks + or causal_kv_cache[0].k.shape[1] != policy.expected_cache_tokens + or causal_kv_cache[0].k.shape[2] != policy.num_attention_heads + or causal_kv_cache[0].sink_tokens != policy.expected_sink_tokens + ) + + def _prepare_realtime_causal_caches( + self, + batch: Req, + server_args: ServerArgs, + ctx: CausalDMDForwardContext, + ) -> CausalDMDRealtimeCacheContext: + policy = self._build_realtime_causal_cache_policy(batch, server_args) + cache_state, persist_state = self._get_realtime_causal_cache_state(batch) + + if self._should_reset_realtime_causal_caches( + batch, + cache_state=cache_state, + policy=policy, + ): + causal_kv_cache, crossattn_cache = self._initialize_causal_caches( + batch_size=ctx.batch_size, + max_text_len=self._get_max_text_len(server_args), + dtype=ctx.target_dtype, + device=ctx.device, + kv_cache_kwargs=policy.kv_cache_kwargs, + ) + cache_state.kv_cache = causal_kv_cache + cache_state.crossattn_cache = crossattn_cache + self._clear_stage_causal_cache_refs() + cache_state.current_chunk_start_frame = 0 + cache_state.chunk_idx = 0 + else: + causal_kv_cache = cache_state.kv_cache + crossattn_cache = cache_state.crossattn_cache + + assert causal_kv_cache is not None + assert crossattn_cache is not None + return CausalDMDRealtimeCacheContext( + cache_state=cache_state, + persist_state=persist_state, + kv_cache=causal_kv_cache, + crossattn_cache=crossattn_cache, + current_start_frame=cache_state.current_chunk_start_frame, + chunk_idx=cache_state.chunk_idx, + ) + + def _advance_realtime_causal_cache( + self, + cache_ctx: CausalDMDRealtimeCacheContext, + *, + num_frames: int, + ) -> None: + cache_ctx.cache_state.current_chunk_start_frame += num_frames + cache_ctx.cache_state.chunk_idx += 1 + + def _realtime_causal_progress_bar(self, batch: Req, timesteps: torch.Tensor): + if batch.session is not None: + return nullcontext(None) + return self.progress_bar(total=len(timesteps)) + + def _denoise_realtime_causal_chunk( + self, + batch: Req, + server_args: ServerArgs, + *, + ctx: CausalDMDForwardContext, + cache_ctx: CausalDMDRealtimeCacheContext, + chunk_latents: torch.Tensor, + prepare_model_input: Callable[[torch.Tensor], torch.Tensor], + prepare_context_input: Callable[[torch.Tensor], torch.Tensor], + ) -> torch.Tensor: + current_start_frame = cache_ctx.current_start_frame + with self._realtime_causal_progress_bar(batch, ctx.timesteps) as progress_bar: + return self._denoise_and_update_causal_block( + batch, + server_args, + chunk_latents=chunk_latents, + scheduler=ctx.scheduler, + timesteps=ctx.timesteps, + prompt_embeds=ctx.prompt_embeds, + kv_cache=cache_ctx.kv_cache, + crossattn_cache=cache_ctx.crossattn_cache, + current_start_tokens=current_start_frame * self.num_token_per_frame, + start_frame=current_start_frame, + image_kwargs=ctx.image_kwargs, + pos_cond_kwargs=ctx.pos_cond_kwargs, + target_dtype=ctx.target_dtype, + autocast_enabled=ctx.autocast_enabled, + device=ctx.device, + attn_raw_latent_shape=(ctx.num_frames, ctx.height, ctx.width), + prepare_model_input=prepare_model_input, + prepare_context_input=prepare_context_input, + progress_bar=progress_bar, + ) + + def _build_causal_attn_metadata( + self, + batch: Req, + server_args: ServerArgs, + *, + current_timestep: int, + raw_latent_shape: tuple[int, int, int], + device: torch.device, + ): + if self.attn_backend.get_enum() != AttentionBackendEnum.VIDEO_SPARSE_ATTN: + return None + + builder_cls = self.attn_backend.get_builder_cls() + if builder_cls is None: + return None + if self._causal_attn_metadata_builder_cls is not builder_cls: + self._causal_attn_metadata_builder_cls = builder_cls + self._causal_attn_metadata_builder = builder_cls() + attn_metadata = self._causal_attn_metadata_builder.build( + current_timestep=current_timestep, + raw_latent_shape=raw_latent_shape, + patch_size=server_args.pipeline_config.dit_config.patch_size, + STA_param=batch.STA_param, + VSA_sparsity=server_args.attention_backend_config.VSA_sparsity, + device=device, + ) + assert attn_metadata is not None, "attn_metadata cannot be None" + return attn_metadata + + @staticmethod + def _single_generator(batch: Req): + if isinstance(batch.generator, list): + return batch.generator[0] + return batch.generator + + @staticmethod + def _expand_timestep( + timestep: torch.Tensor, batch_size: int, device + ) -> torch.Tensor: + return timestep.reshape(1).to(device=device).expand(batch_size) + + 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: + with ( + torch.autocast( + device_type=current_platform.device_type, + dtype=target_dtype, + enabled=autocast_enabled, + ), + set_forward_context( + current_timestep=current_timestep, + attn_metadata=attn_metadata, + forward_batch=batch, + ), + ): + return self.transformer( + latent_model_input, + prompt_embeds, + timestep, + kv_cache=kv_cache, + crossattn_cache=crossattn_cache, + current_start=current_start_tokens, + start_frame=start_frame, + **image_kwargs, + **pos_cond_kwargs, + ) + + def _predict_x0_btchw( + self, + batch: Req, + server_args: ServerArgs, + *, + latent_model_input: torch.Tensor, + noise_latents_btchw: torch.Tensor, + timestep: torch.Tensor, + scheduler, + prompt_embeds, + kv_cache, + crossattn_cache, + current_start_tokens: int, + start_frame: int, + image_kwargs: dict, + pos_cond_kwargs: dict, + attn_raw_latent_shape: tuple[int, int, int], + current_timestep: int, + target_dtype: torch.dtype, + autocast_enabled: bool, + device: torch.device, + ) -> tuple[torch.Tensor, object | None]: + 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 = self._expand_timestep( + timestep, batch_size, latent_model_input.device + ) + pred_noise = self._forward_causal_transformer( + batch, + latent_model_input=latent_model_input, + prompt_embeds=prompt_embeds, + timestep=timestep_2d.unsqueeze(1), + 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, + ) + pred_noise_btchw = pred_noise.permute(0, 2, 1, 3, 4) + x0_btchw = pred_noise_to_pred_video( + pred_noise=pred_noise_btchw.flatten(0, 1), + noise_input_latent=noise_latents_btchw.flatten(0, 1), + timestep=timestep_2d, + scheduler=scheduler, + ).unflatten(0, pred_noise_btchw.shape[:2]) + return x0_btchw, attn_metadata + + def _add_noise_for_next_timestep( + self, + batch: Req, + *, + x0_btchw: torch.Tensor, + raw_latent_shape: torch.Size, + next_timestep: torch.Tensor, + scheduler, + device, + ) -> torch.Tensor: + noise = torch.randn( + raw_latent_shape, + dtype=x0_btchw.dtype, + generator=self._single_generator(batch), + device=device, + ) + return scheduler.add_noise( + x0_btchw.flatten(0, 1), + noise.flatten(0, 1), + next_timestep, + ).unflatten(0, x0_btchw.shape[:2]) + + 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, object | None]: + current_latents = chunk_latents + noise_latents_btchw = current_latents.permute(0, 2, 1, 3, 4) + raw_latent_shape = noise_latents_btchw.shape + attn_metadata = None + + for i, timestep in enumerate(timesteps): + noise_latents = noise_latents_btchw + latent_model_input = prepare_model_input(current_latents).to(target_dtype) + x0_btchw, attn_metadata = self._predict_x0_btchw( + batch, + server_args, + latent_model_input=latent_model_input, + noise_latents_btchw=noise_latents, + timestep=timestep, + scheduler=scheduler, + 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_raw_latent_shape=attn_raw_latent_shape, + current_timestep=i, + target_dtype=target_dtype, + autocast_enabled=autocast_enabled, + device=device, + ) + + if i < len(timesteps) - 1: + next_timestep = timesteps[i + 1 : i + 2] + noise_latents_btchw = self._add_noise_for_next_timestep( + batch, + x0_btchw=x0_btchw, + raw_latent_shape=raw_latent_shape, + next_timestep=next_timestep, + scheduler=scheduler, + device=device, + ) + current_latents = noise_latents_btchw.permute(0, 2, 1, 3, 4) + else: + current_latents = x0_btchw.permute(0, 2, 1, 3, 4) + + if progress_bar is not None: + progress_bar.update() + + return current_latents, attn_metadata + + 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: + """fill the self-attn KV cache by performing a one-time forward on DiT""" + context_noise = getattr(server_args.pipeline_config, "context_noise", 0) + timestep = torch.full( + (context_input.shape[0], 1), + int(context_noise), + device=context_input.device, + dtype=torch.long, + ) + 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, + ) + + def _warm_up_causal_context_cache( + self, + batch: Req, + server_args: ServerArgs, + *, + context_input: torch.Tensor, + prompt_embeds, + kv_cache, + crossattn_cache, + current_start_frame: int, + image_kwargs: dict, + pos_cond_kwargs: dict, + target_dtype: torch.dtype, + autocast_enabled: bool, + ) -> None: + 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_frame * self.num_token_per_frame, + start_frame=current_start_frame, + image_kwargs=image_kwargs, + pos_cond_kwargs=pos_cond_kwargs, + attn_metadata=None, + target_dtype=target_dtype, + autocast_enabled=autocast_enabled, + ) + + def _denoise_and_update_causal_block( + 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], + prepare_context_input: Callable[[torch.Tensor], torch.Tensor], + progress_bar=None, + ) -> torch.Tensor: + current_latents, attn_metadata = self._denoise_causal_dmd_chunk( + batch, + server_args, + chunk_latents=chunk_latents, + scheduler=scheduler, + timesteps=timesteps, + 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, + 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, + ) + self._update_causal_context_cache( + batch, + server_args, + context_input=prepare_context_input(current_latents), + 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, + ) + return current_latents + + def _get_max_text_len(self, server_args: ServerArgs) -> int: + return server_args.pipeline_config.text_encoder_configs[0].arch_config.text_len + + def _initialize_causal_caches( + self, + *, + batch_size: int, + max_text_len: int, + dtype: torch.dtype, + device: torch.device, + kv_cache_kwargs: dict | None = None, + ) -> tuple[list, list]: + self._initialize_kv_cache( + batch_size=batch_size, + dtype=dtype, + device=device, + **(kv_cache_kwargs or {}), + ) + self._initialize_crossattn_cache( + batch_size=batch_size, + max_text_len=max_text_len, + dtype=dtype, + device=device, + ) + assert self.causal_kv_cache is not None + assert self.crossattn_cache is not None + return self.causal_kv_cache, self.crossattn_cache + + def _reset_causal_caches( + self, + *, + kv_cache, + crossattn_cache, + ) -> None: + self._reset_crossattn_cache(crossattn_cache) + self._reset_kv_cache(kv_cache) + + def _reset_crossattn_cache(self, crossattn_cache) -> None: + for cache_block in crossattn_cache: + cache_block.reset() + + @staticmethod + def _reset_kv_cache(kv_cache) -> None: + for cache_block in kv_cache: + cache_block.reset_indices() + + def _get_causal_kv_cache_size( + self, + *, + sequence_shard_enabled: bool = False, + ) -> int: + if self.local_attn_size != -1: + return self.local_attn_size * self.num_token_per_frame + return self.num_token_per_frame * self.sliding_window_num_frames + + def _get_causal_sink_tokens(self) -> int: + return self.sink_size * self.num_token_per_frame + + def _get_causal_attention_window_size(self, kv_cache_size: int) -> int: + return kv_cache_size + + def _allocate_causal_kv_cache( + self, + *, + batch_size: int, + kv_cache_size: int, + num_attention_heads: int, + attention_head_dim: int, + dtype: torch.dtype, + device, + use_int_indices: bool = False, + sink_tokens: int = 0, + attention_window_size: int | None = None, + allow_growth: bool = False, + ) -> list[CausalSelfAttentionKVCache]: + causal_kv_cache = [] + int_index = 0 if use_int_indices else None + if attention_window_size is None: + attention_window_size = kv_cache_size + for _ in range(self.num_transformer_blocks): + causal_kv_cache.append( + CausalSelfAttentionKVCache( + k=torch.zeros( + [ + batch_size, + kv_cache_size, + num_attention_heads, + attention_head_dim, + ], + dtype=dtype, + device=device, + ), + v=torch.zeros( + [ + batch_size, + kv_cache_size, + num_attention_heads, + attention_head_dim, + ], + dtype=dtype, + device=device, + ), + global_end_index=torch.zeros(1, dtype=torch.long, device=device), + local_end_index=torch.zeros(1, dtype=torch.long, device=device), + global_end_index_int=int_index, + local_end_index_int=int_index, + cache_size=kv_cache_size, + sink_tokens=sink_tokens, + attention_window_size=attention_window_size, + allow_growth=allow_growth, + ) + ) + return causal_kv_cache + + @torch.no_grad() + def forward( + self, + batch: Req, + server_args: ServerArgs, + ) -> 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 + + # TODO(will): make this a parameter once we add i2v support + independent_first_frame = self.transformer.independent_first_frame # Initialize or reset caches - if self.kv_cache1 is None: - self._initialize_kv_cache( - batch_size=latents.shape[0], dtype=target_dtype, device=latents.device - ) - self._initialize_crossattn_cache( + if self.causal_kv_cache is None: + self._initialize_causal_caches( batch_size=latents.shape[0], - max_text_len=server_args.pipeline_config.text_encoder_configs[ - 0 - ].arch_config.text_len, + max_text_len=self._get_max_text_len(server_args), dtype=target_dtype, device=latents.device, ) else: assert self.crossattn_cache is not None - # reset cross-attention cache - for block_index in range(self.num_transformer_blocks): - self.crossattn_cache[block_index]["is_init"] = False # type: ignore - # reset kv cache pointers - for block_index in range(len(self.kv_cache1)): - self.kv_cache1[block_index]["global_end_index"] = ( - torch.tensor( # type: ignore - [0], dtype=torch.long, device=latents.device - ) - ) - self.kv_cache1[block_index]["local_end_index"] = ( - torch.tensor( # type: ignore - [0], dtype=torch.long, device=latents.device - ) - ) + self._reset_causal_caches( + kv_cache=self.causal_kv_cache, + crossattn_cache=self.crossattn_cache, + ) # Optional: cache context features from provided image latents prior to generation current_start_frame = 0 @@ -145,30 +954,21 @@ class CausalDMDDenoisingStage(DenoisingStage): image_latent = batch.image_latent assert image_latent is not None input_frames = image_latent.shape[2] - # timestep zero (or configured context noise) for cache warm-up - t_zero = torch.zeros( - [latents.shape[0]], device=latents.device, dtype=torch.long - ) if independent_first_frame and input_frames >= 1: # warm-up with the very first frame independently - image_first_btchw = ( - image_latent[:, :, :1, :, :].to(target_dtype).permute(0, 2, 1, 3, 4) + self._warm_up_causal_context_cache( + batch, + server_args, + context_input=image_latent[:, :, :1, :, :], + prompt_embeds=prompt_embeds, + kv_cache=self.causal_kv_cache, + crossattn_cache=self.crossattn_cache, + current_start_frame=current_start_frame, + image_kwargs=image_kwargs, + pos_cond_kwargs=pos_cond_kwargs, + target_dtype=target_dtype, + autocast_enabled=autocast_enabled, ) - with torch.autocast( - device_type=current_platform.device_type, - dtype=target_dtype, - enabled=autocast_enabled, - ): - _ = self.transformer( - image_first_btchw, - prompt_embeds, - t_zero, - kv_cache=self.kv_cache1, - crossattn_cache=self.crossattn_cache, - current_start=current_start_frame * self.frame_seq_length, - **image_kwargs, - **pos_cond_kwargs, - ) current_start_frame += 1 remaining_frames = input_frames - 1 else: @@ -177,28 +977,21 @@ class CausalDMDDenoisingStage(DenoisingStage): # process remaining input frames in blocks of num_frame_per_block while remaining_frames > 0: block = min(self.num_frames_per_block, remaining_frames) - ref_btchw = ( - image_latent[ + self._warm_up_causal_context_cache( + batch, + server_args, + context_input=image_latent[ :, :, current_start_frame : current_start_frame + block, :, : - ] - .to(target_dtype) - .permute(0, 2, 1, 3, 4) + ], + prompt_embeds=prompt_embeds, + kv_cache=self.causal_kv_cache, + crossattn_cache=self.crossattn_cache, + current_start_frame=current_start_frame, + image_kwargs=image_kwargs, + pos_cond_kwargs=pos_cond_kwargs, + target_dtype=target_dtype, + autocast_enabled=autocast_enabled, ) - with torch.autocast( - device_type=current_platform.device_type, - dtype=target_dtype, - enabled=autocast_enabled, - ): - _ = self.transformer( - ref_btchw, - prompt_embeds, - t_zero, - kv_cache=self.kv_cache1, - crossattn_cache=self.crossattn_cache, - current_start=current_start_frame * self.frame_seq_length, - **image_kwargs, - **pos_cond_kwargs, - ) current_start_frame += block remaining_frames -= block @@ -225,218 +1018,95 @@ class CausalDMDDenoisingStage(DenoisingStage): block_sizes = [1] + [self.num_frames_per_block] * num_blocks start_index = 0 + def prepare_context_input(current_latents): + return current_latents + # DMD loop in causal blocks with self.progress_bar(total=len(block_sizes) * len(timesteps)) as progress_bar: for current_num_frames in block_sizes: current_latents = latents[ :, :, start_index : start_index + current_num_frames, :, : ] - # use BTCHW for DMD conversion routines - noise_latents_btchw = current_latents.permute(0, 2, 1, 3, 4) - video_raw_latent_shape = noise_latents_btchw.shape - - for i, t_cur in enumerate(timesteps): - # Copy for pred conversion - noise_latents = noise_latents_btchw.clone() - latent_model_input = current_latents.to(target_dtype) + 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.to(target_dtype)], + [latent_model_input, batch.image_latent], dim=2, ) + return latent_model_input - # Prepare inputs - t_expand = t_cur.repeat(latent_model_input.shape[0]) - - # Attention metadata if needed - if ( - self.attn_backend.get_enum() - == AttentionBackendEnum.VIDEO_SPARSE_ATTN - ): - self.attn_metadata_builder_cls = ( - self.attn_backend.get_builder_cls() - ) - if self.attn_metadata_builder_cls is not None: - self.attn_metadata_builder = ( - self.attn_metadata_builder_cls() - ) - attn_metadata = self.attn_metadata_builder.build( # type: ignore - current_timestep=i, # type: ignore - raw_latent_shape=( - current_num_frames, - h, - w, - ), # type: ignore - patch_size=server_args.pipeline_config.dit_config.patch_size, # type: ignore - STA_param=batch.STA_param, # type: ignore - VSA_sparsity=server_args.attention_backend_config.VSA_sparsity, # type: ignore - device=get_local_torch_device(), # type: ignore - ) # type: ignore - assert ( - attn_metadata is not None - ), "attn_metadata cannot be None" - else: - attn_metadata = None - else: - attn_metadata = None - - with ( - torch.autocast( - device_type=current_platform.device_type, - dtype=target_dtype, - enabled=autocast_enabled, - ), - set_forward_context( - current_timestep=i, - attn_metadata=attn_metadata, - forward_batch=batch, - ), - ): - # Run transformer; follow DMD stage pattern - t_expanded_noise = t_cur * torch.ones( - (latent_model_input.shape[0], 1), - device=latent_model_input.device, - dtype=torch.long, - ) - pred_noise_btchw = self.transformer( - latent_model_input, - prompt_embeds, - t_expanded_noise, - kv_cache=self.kv_cache1, - crossattn_cache=self.crossattn_cache, - current_start=(pos_start_base + start_index) - * self.frame_seq_length, - start_frame=start_index, - **image_kwargs, - **pos_cond_kwargs, - ).permute(0, 2, 1, 3, 4) - - # Convert pred noise to pred video with FM Euler scheduler utilities - pred_video_btchw = pred_noise_to_pred_video( - pred_noise=pred_noise_btchw.flatten(0, 1), - noise_input_latent=noise_latents.flatten(0, 1), - timestep=t_expand, - scheduler=scheduler, - ).unflatten(0, pred_noise_btchw.shape[:2]) - - if i < len(timesteps) - 1: - next_timestep = timesteps[i + 1] * torch.ones( - [1], dtype=torch.long, device=pred_video_btchw.device - ) - noise = torch.randn( - video_raw_latent_shape, - dtype=pred_video_btchw.dtype, - generator=( - batch.generator[0] - if isinstance(batch.generator, list) - else batch.generator - ), - device=self.device, - ) - noise_btchw = noise - noise_latents_btchw = scheduler.add_noise( - pred_video_btchw.flatten(0, 1), - noise_btchw.flatten(0, 1), - next_timestep, - ).unflatten(0, pred_video_btchw.shape[:2]) - current_latents = noise_latents_btchw.permute(0, 2, 1, 3, 4) - else: - current_latents = pred_video_btchw.permute(0, 2, 1, 3, 4) - - if progress_bar is not None: - progress_bar.update() + current_start_tokens = ( + pos_start_base + start_index + ) * self.num_token_per_frame + current_latents = self._denoise_and_update_causal_block( + batch, + server_args, + chunk_latents=current_latents, + scheduler=scheduler, + timesteps=timesteps, + prompt_embeds=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=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, + ) # Write back and advance latents[:, :, start_index : start_index + current_num_frames, :, :] = ( current_latents ) - # Re-run with context timestep to update KV cache using clean context - context_noise = getattr(server_args.pipeline_config, "context_noise", 0) - t_context = torch.ones( - [latents.shape[0]], device=latents.device, dtype=torch.long - ) * int(context_noise) - context_bcthw = current_latents.to(target_dtype) - with ( - torch.autocast( - device_type=current_platform.device_type, - dtype=target_dtype, - enabled=autocast_enabled, - ), - set_forward_context( - current_timestep=0, - attn_metadata=attn_metadata, - forward_batch=batch, - ), - ): - t_expanded_context = t_context.unsqueeze(1) - _ = self.transformer( - context_bcthw, - prompt_embeds, - t_expanded_context, - kv_cache=self.kv_cache1, - crossattn_cache=self.crossattn_cache, - current_start=(pos_start_base + start_index) - * self.frame_seq_length, - start_frame=start_index, - **image_kwargs, - **pos_cond_kwargs, - ) start_index += current_num_frames batch.latents = latents return batch - def _initialize_kv_cache(self, batch_size, dtype, device) -> None: + def _initialize_kv_cache( + self, + batch_size, + dtype, + device, + *, + sequence_shard_enabled: bool = False, + ) -> None: """ - Initialize a Per-GPU KV cache aligned with the Wan model assumptions. + Initialize (but not fill) a Per-GPU KV cache aligned with the model assumptions. """ - kv_cache1 = [] - num_attention_heads = self.transformer.num_attention_heads + num_attention_heads = self._num_causal_cache_attention_heads( + sequence_shard_enabled=sequence_shard_enabled + ) attention_head_dim = self.transformer.attention_head_dim - if self.local_attn_size != -1: - kv_cache_size = self.local_attn_size * self.frame_seq_length - else: - kv_cache_size = self.frame_seq_length * self.sliding_window_num_frames - - for _ in range(self.num_transformer_blocks): - kv_cache1.append( - { - "k": torch.zeros( - [ - batch_size, - kv_cache_size, - num_attention_heads, - attention_head_dim, - ], - dtype=dtype, - device=device, - ), - "v": torch.zeros( - [ - batch_size, - kv_cache_size, - num_attention_heads, - attention_head_dim, - ], - dtype=dtype, - device=device, - ), - "global_end_index": torch.tensor( - [0], dtype=torch.long, device=device - ), - "local_end_index": torch.tensor( - [0], dtype=torch.long, device=device - ), - } - ) - - self.kv_cache1 = kv_cache1 + kv_cache_size = self._get_causal_kv_cache_size( + sequence_shard_enabled=sequence_shard_enabled + ) + self.causal_kv_cache = self._allocate_causal_kv_cache( + batch_size=batch_size, + kv_cache_size=kv_cache_size, + num_attention_heads=num_attention_heads, + attention_head_dim=attention_head_dim, + dtype=dtype, + device=device, + use_int_indices=self._use_causal_cache_int_indices( + sequence_shard_enabled=sequence_shard_enabled + ), + sink_tokens=self._get_causal_sink_tokens(), + attention_window_size=self._get_causal_attention_window_size(kv_cache_size), + ) def _initialize_crossattn_cache( self, batch_size, max_text_len, dtype, device @@ -449,8 +1119,8 @@ class CausalDMDDenoisingStage(DenoisingStage): attention_head_dim = self.transformer.attention_head_dim for _ in range(self.num_transformer_blocks): crossattn_cache.append( - { - "k": torch.zeros( + CrossAttentionKVCache( + k=torch.zeros( [ batch_size, max_text_len, @@ -460,7 +1130,7 @@ class CausalDMDDenoisingStage(DenoisingStage): dtype=dtype, device=device, ), - "v": torch.zeros( + v=torch.zeros( [ batch_size, max_text_len, @@ -470,8 +1140,7 @@ class CausalDMDDenoisingStage(DenoisingStage): dtype=dtype, device=device, ), - "is_init": False, - } + ) ) self.crossattn_cache = crossattn_cache diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/condition_encoding.py b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/condition_encoding.py new file mode 100644 index 000000000..fd4e6ed39 --- /dev/null +++ b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/condition_encoding.py @@ -0,0 +1,35 @@ +# SPDX-License-Identifier: Apache-2.0 + +from collections.abc import Mapping + +import torch + +from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import Req +from sglang.multimodal_gen.runtime.pipelines_core.stages.base import PipelineStage +from sglang.multimodal_gen.runtime.server_args import ServerArgs + + +class ConditionEncodingStage(PipelineStage): + """Base class for stages that materialize denoiser conditions on the batch.""" + + +class AuxiliaryConditionEncodingStage(ConditionEncodingStage): + """Apply pipeline-config prepared auxiliary conditioning tensors to the request.""" + + def __init__(self, dtype: torch.dtype = torch.bfloat16): + super().__init__() + self.dtype = dtype + + def forward(self, batch: Req, server_args: ServerArgs) -> Req: + condition = server_args.pipeline_config.prepare_world_condition( + batch=batch, + device=self.device, + dtype=self.dtype, + ) + if condition is None: + return batch + if not isinstance(condition, Mapping): + raise TypeError("prepare_world_condition must return a mapping or None") + for field_name, value in condition.items(): + setattr(batch, field_name, value) + return batch diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/latent_preparation.py b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/latent_preparation.py index 2fd915231..3d8ef095e 100644 --- a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/latent_preparation.py +++ b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/latent_preparation.py @@ -38,6 +38,17 @@ class LatentPreparationFingerprint: generator_device: str | None +@dataclass(frozen=True) +class LatentPreparationSpec: + """ "dataclass for controlling the LatentPreparationStage runtime semantics""" + + shape: tuple[int, ...] + dtype: torch.dtype + device: torch.device | str + prepare_latent_ids: bool = True + pack_latents: bool = True + + class LatentPreparationStage(PipelineStage): """ Stage for preparing initial latent variables for the diffusion process. @@ -60,6 +71,37 @@ class LatentPreparationStage(PipelineStage): batch.prompt_embeds[0].dtype ) + def get_forward_latent_num_frames( + self, + batch: Req, + server_args: ServerArgs, + ) -> int: + """get the number of frames to generate for the current batch""" + return self.adjust_video_length(batch, server_args) + + def get_latent_preparation_spec( + self, + batch: Req, + server_args: ServerArgs, + batch_size: int, + num_frames: int, + device: torch.device | str, + ) -> LatentPreparationSpec: + shape = server_args.pipeline_config.prepare_latent_shape( + batch, batch_size, num_frames + ) + return LatentPreparationSpec( + shape=shape, + dtype=self._get_latent_dtype(batch, server_args), + device=device, + ) + + def should_scale_initial_noise(self, batch: Req, server_args: ServerArgs) -> bool: + return True + + def requires_batch_height_width(self, batch: Req, server_args: ServerArgs) -> bool: + return True + def forward( self, batch: Req, @@ -75,23 +117,21 @@ class LatentPreparationStage(PipelineStage): """ # Adjust video length based on VAE version if needed - latent_num_frames = self.adjust_video_length(batch, server_args) + latent_num_frames = self.get_forward_latent_num_frames(batch, server_args) batch_size = batch.batch_size # Get required parameters - dtype = self._get_latent_dtype(batch, server_args) device = get_local_torch_device() generator = batch.generator latents = batch.latents - num_frames = ( - latent_num_frames if latent_num_frames is not None else batch.num_frames - ) height = batch.height width = batch.width # TODO(will): remove this once we add input/output validation for stages - if height is None or width is None: + if self.requires_batch_height_width(batch, server_args) and ( + height is None or width is None + ): raise ValueError("Height and width must be provided") # Validate generator if it's a list @@ -103,26 +143,36 @@ class LatentPreparationStage(PipelineStage): # Generate or use provided latents if latents is None: - shape = server_args.pipeline_config.prepare_latent_shape( - batch, batch_size, num_frames + spec = self.get_latent_preparation_spec( + batch, server_args, batch_size, latent_num_frames, device ) latents = randn_tensor( - shape, generator=generator, device=device, dtype=dtype + spec.shape, + generator=generator, + device=spec.device, + dtype=spec.dtype, ) - latent_ids = server_args.pipeline_config.maybe_prepare_latent_ids(latents) + latent_ids = ( + server_args.pipeline_config.maybe_prepare_latent_ids(latents) + if spec.prepare_latent_ids + else None + ) if latent_ids is not None: batch.latent_ids = latent_ids.to(device=device) - latents = server_args.pipeline_config.maybe_pack_latents( - latents, batch_size, batch - ) + if spec.pack_latents: + latents = server_args.pipeline_config.maybe_pack_latents( + latents, batch_size, batch + ) else: latents = latents.to(device) # Scale the initial noise if needed - if hasattr(self.scheduler, "init_noise_sigma"): + if self.should_scale_initial_noise(batch, server_args) and hasattr( + self.scheduler, "init_noise_sigma" + ): latents = latents * self.scheduler.init_noise_sigma # Update batch with prepared latents batch.latents = latents @@ -175,7 +225,7 @@ class LatentPreparationStage(PipelineStage): if isinstance(batch.prompt_embeds, list) and batch.prompt_embeds else None ) - latent_num_frames = self.adjust_video_length(batch, server_args) + latent_num_frames = self.get_forward_latent_num_frames(batch, server_args) return LatentPreparationFingerprint( height=batch.height, width=batch.width, @@ -206,51 +256,58 @@ class LatentPreparationStage(PipelineStage): deterministic packing/scaling work. """ first_batch = batches[0] - latent_num_frames = self.adjust_video_length(first_batch, server_args) + latent_num_frames = self.get_forward_latent_num_frames(first_batch, server_args) batch_size = len(batches) - dtype = self._get_latent_dtype(first_batch, server_args) device = get_local_torch_device() - num_frames = ( - latent_num_frames - if latent_num_frames is not None - else first_batch.num_frames + first_spec = self.get_latent_preparation_spec( + first_batch, + server_args, + batch_size, + latent_num_frames, + device, ) - height = first_batch.height - width = first_batch.width - - if height is None or width is None: - raise ValueError("Height and width must be provided") raw_latents = [] for batch in batches: - shape = server_args.pipeline_config.prepare_latent_shape( - batch, 1, num_frames + spec = self.get_latent_preparation_spec( + batch, + server_args, + 1, + latent_num_frames, + device, ) raw_latents.append( randn_tensor( - shape, + spec.shape, generator=self._single_generator(batch), - device=device, - dtype=dtype, + device=spec.device, + dtype=spec.dtype, ) ) latents = torch.cat(raw_latents, dim=0) - latent_ids = server_args.pipeline_config.maybe_prepare_latent_ids(latents) + latent_ids = ( + server_args.pipeline_config.maybe_prepare_latent_ids(latents) + if first_spec.prepare_latent_ids + else None + ) if latent_ids is not None: first_batch.latent_ids = latent_ids.to(device=device) - original_num_outputs = first_batch.num_outputs_per_prompt - try: - first_batch.num_outputs_per_prompt = batch_size - latents = server_args.pipeline_config.maybe_pack_latents( - latents, batch_size, first_batch - ) - finally: - first_batch.num_outputs_per_prompt = original_num_outputs + if first_spec.pack_latents: + original_num_outputs = first_batch.num_outputs_per_prompt + try: + first_batch.num_outputs_per_prompt = batch_size + latents = server_args.pipeline_config.maybe_pack_latents( + latents, batch_size, first_batch + ) + finally: + first_batch.num_outputs_per_prompt = original_num_outputs - if hasattr(self.scheduler, "init_noise_sigma"): + if self.should_scale_initial_noise(first_batch, server_args) and hasattr( + self.scheduler, "init_noise_sigma" + ): latents = latents * self.scheduler.init_noise_sigma first_batch.latents = latents @@ -320,3 +377,59 @@ class LatentPreparationStage(PipelineStage): # result.add_check("latents", batch.latents, [V.is_tensor, V.with_dims(5)]) result.add_check("raw_latent_shape", batch.raw_latent_shape, V.is_tuple) return result + + +class RealtimeChunkLatentPreparationStage(LatentPreparationStage): + """Prepare one realtime causal DiT chunk from the encoded condition shape.""" + + def get_forward_latent_num_frames( + self, + batch: Req, + server_args: ServerArgs, + ) -> int: + return int( + batch.realtime_chunk_size + or self.transformer.config.arch_config.num_frames_per_block + ) + + def get_latent_preparation_spec( + self, + batch: Req, + server_args: ServerArgs, + batch_size: int, + num_frames: int, + device: torch.device | str, + ) -> LatentPreparationSpec: + condition_latent = batch.image_latent + assert condition_latent is not None, ( + "Realtime chunk latent preparation requires image_latent. " + "Ensure the condition VAE encoding stage runs before this stage." + ) + return LatentPreparationSpec( + shape=( + condition_latent.shape[0], + self.transformer.config.arch_config.out_channels, + num_frames, + condition_latent.shape[3], + condition_latent.shape[4], + ), + dtype=condition_latent.dtype, + device=device, + prepare_latent_ids=False, + pack_latents=False, + ) + + def should_scale_initial_noise(self, batch: Req, server_args: ServerArgs) -> bool: + return False + + def requires_batch_height_width(self, batch: Req, server_args: ServerArgs) -> bool: + return False + + def verify_input(self, batch: Req, server_args: ServerArgs) -> VerificationResult: + result = VerificationResult() + result.add_check( + "image_latent", batch.image_latent, [V.is_tensor, V.with_dims(5)] + ) + result.add_check("generator", batch.generator, V.generator_or_list_generators) + result.add_check("latents", batch.latents, V.none_or_tensor) + return result diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/lingbot_world/__init__.py b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/lingbot_world/__init__.py new file mode 100644 index 000000000..5e7b442f3 --- /dev/null +++ b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/lingbot_world/__init__.py @@ -0,0 +1,10 @@ +# SPDX-License-Identifier: Apache-2.0 +"""LingBot-World-specific pipeline stages.""" + +from sglang.multimodal_gen.runtime.pipelines_core.stages.model_specific_stages.lingbot_world.lingbot_world_causal_denoising import ( + LingBotWorldCausalDMDDenoisingStage, +) + +__all__ = [ + "LingBotWorldCausalDMDDenoisingStage", +] diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/lingbot_world/lingbot_world_causal_denoising.py b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/lingbot_world/lingbot_world_causal_denoising.py new file mode 100644 index 000000000..3fa583d22 --- /dev/null +++ b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/lingbot_world/lingbot_world_causal_denoising.py @@ -0,0 +1,315 @@ +# SPDX-License-Identifier: Apache-2.0 +# Adapted from: https://github.com/Robbyant/lingbot-world + +"""LingBot-World causal DMD denoising stage.""" + +import torch + +from sglang.multimodal_gen.runtime.distributed.parallel_state import ( + get_ring_parallel_world_size, + get_ulysses_parallel_world_size, +) +from sglang.multimodal_gen.runtime.managers.forward_context import set_forward_context +from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import Req +from sglang.multimodal_gen.runtime.pipelines_core.stages.causal_denoising import ( + CausalDMDCachePolicy, + CausalDMDDenoisingStage, +) +from sglang.multimodal_gen.runtime.pipelines_core.stages.validators import ( + StageValidators as V, +) +from sglang.multimodal_gen.runtime.pipelines_core.stages.validators import ( + VerificationResult, +) +from sglang.multimodal_gen.runtime.platforms import current_platform +from sglang.multimodal_gen.runtime.server_args import ServerArgs + + +class LingBotWorldCausalDMDDenoisingStage(CausalDMDDenoisingStage): + """Causal DMD denoising with I2V condition concatenation for LingBot-World. + + The LingBot-World transformer has ``in_channels = 36`` and expects + ``[noise(16ch), condition(20ch)]`` concatenated along channel dim. + Each call processes one chunk (num_frames_per_block frames). + """ + + def _get_causal_kv_cache_size( + self, + *, + sequence_shard_enabled: bool = False, + ) -> int: + if self.local_attn_size != -1: + return self.local_attn_size * self.num_token_per_frame + + return self.sliding_window_num_frames * self.num_token_per_frame + + def _causal_sequence_shard_enabled(self, batch: Req) -> bool: + return bool( + getattr(batch, "enable_sequence_shard", False) + and get_ulysses_parallel_world_size() > 1 + ) + + def _num_causal_cache_attention_heads( + self, + *, + sequence_shard_enabled: bool, + ) -> int: + num_attention_heads = self.transformer.num_attention_heads + if not sequence_shard_enabled: + return num_attention_heads + + ulysses_world_size = get_ulysses_parallel_world_size() + if get_ring_parallel_world_size() > 1: + raise NotImplementedError( + "LingBot causal sequence sharding currently supports ulysses_degree > 1 with ring_degree = 1 only." + ) + if ulysses_world_size <= 1: + raise ValueError( + "LingBot causal sequence sharding requires ulysses_degree > 1." + ) + if num_attention_heads % ulysses_world_size != 0: + raise ValueError( + f"num_attention_heads ({num_attention_heads}) must be divisible by ulysses_degree ({ulysses_world_size})." + ) + return num_attention_heads // ulysses_world_size + + def _causal_kv_cache_kwargs( + self, + policy: CausalDMDCachePolicy, + ) -> dict[str, bool]: + return {"sequence_shard_enabled": policy.sequence_shard_enabled} + + def _use_causal_cache_int_indices( + self, + *, + sequence_shard_enabled: bool, + ) -> bool: + return True + + def verify_input(self, batch: Req, server_args: ServerArgs) -> VerificationResult: + result = VerificationResult() + result.add_check( + "image_latent", batch.image_latent, [V.is_tensor, V.with_dims(5)] + ) + result.add_check("latents", batch.latents, [V.is_tensor, V.with_dims(5)]) + result.add_check("timesteps", batch.timesteps, [V.is_tensor, V.with_dims(1)]) + result.add_check("scheduler", batch.scheduler, V.not_none) + result.add_check("prompt_embeds", batch.prompt_embeds, V.list_not_empty) + return result + + def _get_causal_dmd_latents(self, batch: Req) -> torch.Tensor: + latents = batch.latents + assert latents is not None, ( + "LingBot-World causal DMD requires prepared chunk latents. " + "Ensure RealtimeChunkLatentPreparationStage runs before this stage." + ) + return latents + + def _get_causal_dmd_scheduler(self, batch: Req, server_args: ServerArgs): + scheduler = batch.scheduler + assert scheduler is not None, ( + "LingBot-World causal DMD requires prepared DMD timesteps. " + "Ensure DMDTimestepPreparationStage runs before this stage." + ) + return scheduler + + def _prepare_causal_dmd_timesteps( + self, + batch: Req, + server_args: ServerArgs, + scheduler, + device: torch.device, + ) -> torch.Tensor: + timesteps = batch.timesteps + assert timesteps is not None + return timesteps.to(device) + + def _prepare_causal_dmd_image_kwargs( + self, + batch: Req, + server_args: ServerArgs, + target_dtype: torch.dtype, + ) -> dict: + image_embeds = getattr(batch, "image_embeds", []) + if len(image_embeds) > 0: + image_embeds = [ie.to(target_dtype) for ie in image_embeds] + return { + "encoder_hidden_states_image": image_embeds, + } + + def _prepare_causal_dmd_pos_cond_kwargs( + self, + batch: Req, + server_args: ServerArgs, + target_dtype: torch.dtype, + ) -> dict: + # lingbot transformer forward uses varargs, so inspect filtering drops valid kwargs + return server_args.pipeline_config.prepare_pos_cond_kwargs( + batch, + self.device, + getattr(self.transformer, "rotary_emb", None), + dtype=target_dtype, + ) + + def _prepare_causal_dmd_prompt_embeds( + self, + batch: Req, + server_args: ServerArgs, + target_dtype: torch.dtype, + ): + return server_args.pipeline_config.get_pos_prompt_embeds(batch) + + 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], 1), + int(context_noise), + device=context_input.device, + dtype=torch.long, + ) + with ( + torch.autocast( + device_type=current_platform.device_type, + dtype=target_dtype, + enabled=autocast_enabled, + ), + set_forward_context( + current_timestep=0, + attn_metadata=attn_metadata, + forward_batch=batch, + ), + ): + self.transformer( + context_input.to(target_dtype), + prompt_embeds, + timestep, + kv_cache=kv_cache, + crossattn_cache=crossattn_cache, + current_start=current_start_tokens, + start_frame=start_frame, + skip_final_projection=True, + **image_kwargs, + **pos_cond_kwargs, + ) + + @staticmethod + def _select_i2v_condition_chunk( + condition_full: torch.Tensor, + chunk_idx: int, + chunk_size: int, + ) -> torch.Tensor: + condition_chunks = condition_full.split(chunk_size, dim=2) + condition = condition_chunks[min(chunk_idx, len(condition_chunks) - 1)] + + if condition.shape[2] == chunk_size: + return condition + pad_frames = chunk_size - condition.shape[2] + return torch.cat( + [ + condition, + condition.new_zeros( + condition.shape[0], + condition.shape[1], + pad_frames, + condition.shape[3], + condition.shape[4], + ), + ], + dim=2, + ) + + @staticmethod + def _build_i2v_model_input_writer( + *, + latents: torch.Tensor, + condition: torch.Tensor, + target_dtype: torch.dtype, + device: torch.device, + ): + b, latent_channels, t, h, w = latents.shape + condition = condition.to(device=device, dtype=target_dtype) + model_input = torch.empty( + ( + b, + latent_channels + condition.shape[1], + t, + h, + w, + ), + dtype=target_dtype, + device=device, + ) + model_input[:, latent_channels:].copy_(condition) + + def write(current_latents: torch.Tensor) -> torch.Tensor: + model_input[:, :latent_channels].copy_(current_latents) + return model_input + + return write + + @torch.no_grad() + def forward(self, batch: Req, server_args: ServerArgs) -> Req: + # --- Condition: take current chunk's slice --- + condition_full = batch.image_latent + assert condition_full is not None, ( + "LingBot-World causal DMD requires image_latent as condition. " + "Ensure ImageVAEEncodingStage runs before this stage." + ) + ctx = self._prepare_causal_dmd_forward_context(batch, server_args) + latents = ctx.latents + cache_ctx = self._prepare_realtime_causal_caches(batch, server_args, ctx) + + # Keep cross-attention K/V cache across realtime chunks; LingBot text/image + # conditions are session-static and are invalidated by cache reset. + + # Slice condition to current chunk + condition = self._select_i2v_condition_chunk( + condition_full, + cache_ctx.chunk_idx, + ctx.num_frames, + ) + + # --- Denoising loop (single chunk) --- + current_latents = latents + prepare_model_input = self._build_i2v_model_input_writer( + latents=current_latents, + condition=condition, + target_dtype=ctx.target_dtype, + device=ctx.device, + ) + + current_latents = self._denoise_realtime_causal_chunk( + batch, + server_args, + ctx=ctx, + cache_ctx=cache_ctx, + chunk_latents=current_latents, + prepare_model_input=prepare_model_input, + prepare_context_input=prepare_model_input, + ) + + # Advance cumulative frame position + self._advance_realtime_causal_cache(cache_ctx, num_frames=ctx.num_frames) + + # Output denoised latents for decoder + batch.latents = current_latents + batch.raw_latent_shape = current_latents.shape + if not cache_ctx.persist_state: + cache_ctx.cache_state.dispose() + return batch diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/realtime_input_validation.py b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/realtime_input_validation.py new file mode 100644 index 000000000..5b9727458 --- /dev/null +++ b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/realtime_input_validation.py @@ -0,0 +1,139 @@ +# SPDX-License-Identifier: Apache-2.0 + +from __future__ import annotations + +from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import Req +from sglang.multimodal_gen.runtime.pipelines_core.stages.input_validation import ( + InputValidationStage, +) +from sglang.multimodal_gen.runtime.realtime.session import ( + BaseRealtimeState, +) +from sglang.multimodal_gen.runtime.server_args import ServerArgs + + +class RealtimeInputValidationState(BaseRealtimeState): + def __init__(self): + super().__init__() + self.image_path = None + self.condition_image = None + self.original_condition_image_size = None + self.height = None + self.width = None + self.generator = None + self.seeds = None + self.generator_seed = None + self.generator_device = None + self.num_outputs_per_prompt = None + + def dispose(self): + super().dispose() + self.image_path = None + self.condition_image = None + self.original_condition_image_size = None + self.height = None + self.width = None + self.generator = None + self.seeds = None + self.generator_seed = None + self.generator_device = None + self.num_outputs_per_prompt = None + + +class RealtimeInputValidationStage(InputValidationStage): + """Reuse validated image and generator inputs across chunks in a realtime session.""" + + def preprocess_condition_image( + self, + batch: Req, + server_args: ServerArgs, + condition_image_width, + condition_image_height, + ): + if server_args.pipeline_config.preprocess_realtime_condition_image( + batch, + self.vae_image_processor, + ): + return + return super().preprocess_condition_image( + batch, + server_args, + condition_image_width, + condition_image_height, + ) + + def _cache_batch(self, batch: Req, state: RealtimeInputValidationState) -> None: + state.image_path = batch.image_path + state.condition_image = batch.condition_image + state.original_condition_image_size = batch.original_condition_image_size + state.height = batch.height + state.width = batch.width + + def _can_reuse_cached_image( + self, batch: Req, state: RealtimeInputValidationState + ) -> bool: + if batch.block_idx == 0 or state.condition_image is None: + return False + return ( + batch.image_path in (None, state.image_path) + and batch.height in (None, state.height) + and batch.width in (None, state.width) + ) + + def _cache_generator(self, batch: Req, state: RealtimeInputValidationState) -> None: + state.generator = batch.generator + state.seeds = batch.seeds + state.generator_seed = batch.seed + state.generator_device = batch.generator_device + state.num_outputs_per_prompt = batch.num_outputs_per_prompt + + def _can_reuse_generator( + self, batch: Req, state: RealtimeInputValidationState + ) -> bool: + if batch.block_idx == 0 or state.generator is None: + return False + return ( + state.generator_seed == batch.seed + and state.generator_device == batch.generator_device + and state.num_outputs_per_prompt == batch.num_outputs_per_prompt + ) + + def _reuse_or_cache_generator( + self, batch: Req, state: RealtimeInputValidationState + ) -> None: + if self._can_reuse_generator(batch, state): + batch.generator = state.generator + batch.seeds = state.seeds + return + + self._cache_generator(batch, state) + + def forward( + self, + batch: Req, + server_args: ServerArgs, + ) -> Req: + if batch.session is None: + return super().forward(batch, server_args) + + state = batch.session.get_or_create_state(RealtimeInputValidationState) + if self._can_reuse_cached_image(batch, state): + original_image_path = batch.image_path + batch.image_path = None + batch.condition_image = state.condition_image + batch.original_condition_image_size = state.original_condition_image_size + if batch.height is None: + batch.height = state.height + if batch.width is None: + batch.width = state.width + try: + batch = super().forward(batch, server_args) + self._reuse_or_cache_generator(batch, state) + return batch + finally: + batch.image_path = original_image_path + + batch = super().forward(batch, server_args) + self._cache_batch(batch, state) + self._reuse_or_cache_generator(batch, state) + return batch diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/realtime_text_encoding.py b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/realtime_text_encoding.py new file mode 100644 index 000000000..ddd1636c3 --- /dev/null +++ b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/realtime_text_encoding.py @@ -0,0 +1,152 @@ +# SPDX-License-Identifier: Apache-2.0 +# Adapted from: https://github.com/Robbyant/lingbot-world + +""" +LingBot-World realtime text stages. + +The reference lingbot_fast_server initializes prompt embeddings once per session. +Cache text encoder outputs across realtime chunks so condition sampling stays +closer to the actual denoising step. +""" + +from __future__ import annotations + +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.text_encoding import ( + TextEncodingStage, +) +from sglang.multimodal_gen.runtime.realtime.session import ( + BaseRealtimeState, +) +from sglang.multimodal_gen.runtime.server_args import ServerArgs + + +def _normalize_prompt_value( + value: str | list[str] | None, +) -> str | tuple[str, ...] | None: + if isinstance(value, list): + return tuple(value) + return value + + +def _copy_tensor_list( + value: list[torch.Tensor] | None, +) -> list[torch.Tensor] | None: + if value is None: + return None + return list(value) + + +def _copy_seq_lens( + value: list[list[int]] | None, +) -> list[list[int]] | None: + if value is None: + return None + return [list(seq_lens) for seq_lens in value] + + +class RealtimeTextState(BaseRealtimeState): + def __init__(self): + super().__init__() + self.cache_key: tuple[Any, ...] | None = None + self.prompt_embeds: list[torch.Tensor] | None = None + self.pooled_embeds: list[torch.Tensor] | None = None + self.prompt_attention_mask: list[torch.Tensor] | None = None + self.prompt_embeds_mask: list[torch.Tensor] | None = None + self.prompt_seq_lens: list[list[int]] | None = None + self.negative_prompt_embeds: list[torch.Tensor] | None = None + self.neg_pooled_embeds: list[torch.Tensor] | None = None + self.negative_attention_mask: list[torch.Tensor] | None = None + self.negative_prompt_embeds_mask: list[torch.Tensor] | None = None + self.negative_prompt_seq_lens: list[list[int]] | None = None + + def clear_text_cache(self): + self.cache_key = None + self.prompt_embeds = None + self.pooled_embeds = None + self.prompt_attention_mask = None + self.prompt_embeds_mask = None + self.prompt_seq_lens = None + self.negative_prompt_embeds = None + self.neg_pooled_embeds = None + self.negative_attention_mask = None + self.negative_prompt_embeds_mask = None + self.negative_prompt_seq_lens = None + + def dispose(self): + super().dispose() + self.clear_text_cache() + + +class RealtimeTextEncodingStage(TextEncodingStage): + """Cache text encoder outputs across realtime chunks by prompt identity.""" + + def _make_cache_key(self, batch: Req) -> tuple[Any, ...]: + return ( + _normalize_prompt_value(batch.prompt), + bool(batch.do_classifier_free_guidance), + ( + _normalize_prompt_value(batch.negative_prompt) + if batch.do_classifier_free_guidance + else None + ), + ) + + def _restore_cached_outputs(self, batch: Req, state: RealtimeTextState) -> Req: + batch.prompt_embeds = _copy_tensor_list(state.prompt_embeds) or [] + batch.pooled_embeds = _copy_tensor_list(state.pooled_embeds) or [] + batch.prompt_attention_mask = _copy_tensor_list(state.prompt_attention_mask) + batch.prompt_embeds_mask = _copy_tensor_list(state.prompt_embeds_mask) + batch.prompt_seq_lens = _copy_seq_lens(state.prompt_seq_lens) + batch.negative_prompt_embeds = _copy_tensor_list(state.negative_prompt_embeds) + batch.neg_pooled_embeds = _copy_tensor_list(state.neg_pooled_embeds) or [] + batch.negative_attention_mask = _copy_tensor_list(state.negative_attention_mask) + batch.negative_prompt_embeds_mask = _copy_tensor_list( + state.negative_prompt_embeds_mask + ) + batch.negative_prompt_seq_lens = _copy_seq_lens(state.negative_prompt_seq_lens) + return batch + + def _store_outputs(self, batch: Req, state: RealtimeTextState) -> None: + state.prompt_embeds = _copy_tensor_list(batch.prompt_embeds) + state.pooled_embeds = _copy_tensor_list(batch.pooled_embeds) + state.prompt_attention_mask = _copy_tensor_list(batch.prompt_attention_mask) + state.prompt_embeds_mask = _copy_tensor_list(batch.prompt_embeds_mask) + state.prompt_seq_lens = _copy_seq_lens(batch.prompt_seq_lens) + state.negative_prompt_embeds = _copy_tensor_list(batch.negative_prompt_embeds) + state.neg_pooled_embeds = _copy_tensor_list(batch.neg_pooled_embeds) + state.negative_attention_mask = _copy_tensor_list(batch.negative_attention_mask) + state.negative_prompt_embeds_mask = _copy_tensor_list( + batch.negative_prompt_embeds_mask + ) + state.negative_prompt_seq_lens = _copy_seq_lens(batch.negative_prompt_seq_lens) + + @torch.no_grad() + def forward( + self, + batch: Req, + server_args: ServerArgs, + ) -> Req: + if batch.session is None: + return super().forward(batch, server_args) + + state = batch.session.get_or_create_state(RealtimeTextState) + assert isinstance(state, RealtimeTextState) + + # cache the encoder results into BaseRealtimeState, restore when encoder inputs hits the cache + cache_key = self._make_cache_key(batch) + if state.cache_key == cache_key and state.prompt_embeds is not None: + return self._restore_cached_outputs(batch, state) + + state.clear_text_cache() + + # perform regular text encoding + batch = super().forward(batch, server_args) + + state.cache_key = cache_key + self._store_outputs(batch, state) + return batch diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/realtime_vae.py b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/realtime_vae.py new file mode 100644 index 000000000..888334dc7 --- /dev/null +++ b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/realtime_vae.py @@ -0,0 +1,204 @@ +# SPDX-License-Identifier: Apache-2.0 + +import torch + +from sglang.multimodal_gen.runtime.distributed import get_local_torch_device +from sglang.multimodal_gen.runtime.models.vaes.wanvae import ( + unpatchify as wan_unpatchify, +) +from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import OutputBatch, Req +from sglang.multimodal_gen.runtime.pipelines_core.stages.decoding import ( + DecodingStage, + _ensure_tensor_decode_output, +) +from sglang.multimodal_gen.runtime.pipelines_core.stages.image_encoding import ( + ImageVAEEncodingStage, +) +from sglang.multimodal_gen.runtime.platforms import current_platform +from sglang.multimodal_gen.runtime.realtime.session import ( + BaseRealtimeState, +) +from sglang.multimodal_gen.runtime.server_args import ServerArgs +from sglang.multimodal_gen.utils import PRECISION_TO_TYPE + + +class RealtimeVAEState(BaseRealtimeState): + def __init__(self): + super().__init__() + self.image_latent: torch.Tensor | None = None + + def dispose(self): + super().dispose() + self.image_latent = None + + +class RealtimeVAEDecodeState(BaseRealtimeState): + def __init__(self): + super().__init__() + self.reset_causal_decode_state = None + + def dispose(self): + reset_causal_decode_state = self.reset_causal_decode_state + self.reset_causal_decode_state = None + if callable(reset_causal_decode_state): + reset_causal_decode_state() + + +class RealtimeImageVAEEncodingStage(ImageVAEEncodingStage): + """Reuse the first chunk's conditioning image latent across a realtime session.""" + + def forward( + self, + batch: Req, + server_args: ServerArgs, + ) -> Req: + state = None + if batch.session is not None: + state = batch.session.get_or_create_state(RealtimeVAEState) + if batch.block_idx == 0: + state.image_latent = None + elif state.image_latent is not None: + batch.image_latent = state.image_latent + return batch + + if batch.condition_image is None: + if state is not None and state.image_latent is not None: + batch.image_latent = state.image_latent + return batch + + batch = super().forward(batch, server_args) + + if state is not None and batch.image_latent is not None: + state.image_latent = batch.image_latent + return batch + + +class CausalVaeDecodingStage(DecodingStage): + """Decode realtime chunks with a persistent causal VAE cache when available.""" + + @staticmethod + def _supports_wan_decoder_cache(vae) -> bool: + return all( + hasattr(vae, attr) + for attr in ( + "clear_cache", + "post_quant_conv", + "decoder", + "_feat_map", + "_conv_idx", + ) + ) + + def _get_causal_decode_reset_fn(self): + reset_causal_state = getattr(self.vae, "reset_causal_decode_state", None) + if callable(reset_causal_state): + return reset_causal_state + if self._supports_wan_decoder_cache(self.vae): + return self.vae.clear_cache + return None + + def _decode_wan_with_persistent_cache( + self, + latents: torch.Tensor, + *, + first_chunk: bool, + ) -> torch.Tensor: + x = self.vae.post_quant_conv(latents) + decoded_frames = [] + for frame_idx in range(x.shape[2]): + self.vae._conv_idx = [0] + decoded = self.vae.decoder( + x[:, :, frame_idx : frame_idx + 1], + feat_cache=self.vae._feat_map, + feat_idx=self.vae._conv_idx, + first_chunk=first_chunk and frame_idx == 0, + ) + decoded_frames.append(decoded) + + image = torch.cat(decoded_frames, dim=2) + if getattr(self.vae.config, "patch_size", None) is not None: + image = wan_unpatchify(image, patch_size=self.vae.config.patch_size) + return image.clamp(-1.0, 1.0) + + @torch.no_grad() + def decode_causal( + self, + latents: torch.Tensor, + server_args: ServerArgs, + *, + first_chunk: bool, + ) -> torch.Tensor: + vae_dtype = PRECISION_TO_TYPE[server_args.pipeline_config.vae_precision] + self.vae = self.vae.to(device=get_local_torch_device(), dtype=vae_dtype) + latents = latents.to(get_local_torch_device()) + vae_autocast_enabled = ( + vae_dtype != torch.float32 + ) and not server_args.disable_autocast + + latents = self.scale_and_shift(latents, server_args) + latents = server_args.pipeline_config.preprocess_decoding( + latents, server_args, vae=self.vae + ) + + with torch.autocast( + device_type=current_platform.device_type, + dtype=vae_dtype, + enabled=vae_autocast_enabled, + ): + try: + if server_args.pipeline_config.vae_tiling: + self.vae.enable_tiling() + except Exception: + pass + + if not vae_autocast_enabled: + latents = latents.to(vae_dtype) + + decode_fn = getattr(self.vae, "causal_decode", None) + if callable(decode_fn): + decode_output = decode_fn(latents) + image = _ensure_tensor_decode_output(decode_output) + elif self._supports_wan_decoder_cache(self.vae): + image = self._decode_wan_with_persistent_cache( + latents, + first_chunk=first_chunk, + ) + else: + decode_output = self.vae.decode(latents) + image = _ensure_tensor_decode_output(decode_output) + + return (image / 2 + 0.5).clamp(0, 1) + + @torch.no_grad() + def forward( + self, + batch: Req, + server_args: ServerArgs, + ) -> OutputBatch: + if batch.session is None: + return super().forward(batch, server_args) + + self.load_model() + + reset_causal_state = self._get_causal_decode_reset_fn() + decode_state = batch.session.get_or_create_state(RealtimeVAEDecodeState) + decode_state.reset_causal_decode_state = reset_causal_state + if batch.block_idx == 0 and callable(reset_causal_state): + reset_causal_state() + + frames = self.decode_causal( + batch.latents, + server_args, + first_chunk=batch.block_idx == 0, + ) + frames = server_args.pipeline_config.post_decoding(frames, server_args) + + return OutputBatch( + output=frames, + trajectory_timesteps=batch.trajectory_timesteps, + trajectory_latents=batch.trajectory_latents, + rollout_trajectory_data=batch.rollout_trajectory_data, + trajectory_decoded=None, + metrics=batch.metrics, + noise_pred=None, + ) diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/text_encoding.py b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/text_encoding.py index c332c1652..e98b1ce93 100644 --- a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/text_encoding.py +++ b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/text_encoding.py @@ -22,7 +22,9 @@ from sglang.multimodal_gen.runtime.managers.memory_managers.component_manager im ComponentUse, ) from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import Req -from sglang.multimodal_gen.runtime.pipelines_core.stages.base import PipelineStage +from sglang.multimodal_gen.runtime.pipelines_core.stages.condition_encoding import ( + ConditionEncodingStage, +) from sglang.multimodal_gen.runtime.pipelines_core.stages.validators import ( StageValidators as V, ) @@ -66,7 +68,7 @@ def stack_tensors(name: str, tensors: list[torch.Tensor]) -> torch.Tensor: return torch.stack(tensors, dim=0) -class TextEncodingStage(PipelineStage): +class TextEncodingStage(ConditionEncodingStage): """ Stage for encoding text prompts into embeddings for diffusion models. diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/timestep_preparation.py b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/timestep_preparation.py index da08f28ea..ea581f9b8 100644 --- a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/timestep_preparation.py +++ b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/timestep_preparation.py @@ -199,3 +199,63 @@ class TimestepPreparationStage(PipelineStage): result = VerificationResult() result.add_check("timesteps", batch.timesteps, [V.is_tensor, V.with_dims(1)]) return result + + +class DMDTimestepPreparationStage(PipelineStage): + """Prepare distilled DMD timesteps from pipeline config.""" + + deduplicated_tensor_tree_output_fields = ("timesteps",) + deduplicated_deepcopy_output_fields = ("scheduler",) + + def __init__(self, scheduler) -> None: + super().__init__() + self.scheduler = scheduler + + @property + def parallelism_type(self) -> StageParallelismType: + return StageParallelismType.REPLICATED + + def forward( + self, + batch: Req, + server_args: ServerArgs, + ) -> Req: + if batch.scheduler is not None and batch.timesteps is not None: + return batch + + scheduler = get_or_create_request_scheduler(batch, self.scheduler) + num_train_timesteps = getattr(scheduler, "num_train_timesteps", None) + if num_train_timesteps is None: + num_train_timesteps = scheduler.config.num_train_timesteps + num_train_timesteps = int(num_train_timesteps) + scheduler.set_timesteps(num_train_timesteps) + + timesteps = torch.tensor( + server_args.pipeline_config.dmd_denoising_steps, dtype=torch.long + ).cpu() + if server_args.pipeline_config.warp_denoising_step: + scheduler_timesteps = torch.cat( + (scheduler.timesteps.cpu(), torch.tensor([0], dtype=torch.float32)) + ) + timesteps = scheduler_timesteps[num_train_timesteps - timesteps] + + batch.timesteps = timesteps.to(get_local_torch_device()) + batch.scheduler = scheduler + if not batch.is_warmup: + self.log_debug("DMD timesteps: %s", batch.timesteps) + return batch + + def verify_input(self, batch: Req, server_args: ServerArgs) -> VerificationResult: + result = VerificationResult() + result.add_check( + "dmd_denoising_steps", + server_args.pipeline_config.dmd_denoising_steps, + V.list_not_empty, + ) + return result + + def verify_output(self, batch: Req, server_args: ServerArgs) -> VerificationResult: + result = VerificationResult() + result.add_check("timesteps", batch.timesteps, [V.is_tensor, V.with_dims(1)]) + result.add_check("scheduler", batch.scheduler, V.not_none) + return result diff --git a/python/sglang/multimodal_gen/runtime/postprocess/__init__.py b/python/sglang/multimodal_gen/runtime/postprocess/__init__.py index 139f98766..99c416072 100644 --- a/python/sglang/multimodal_gen/runtime/postprocess/__init__.py +++ b/python/sglang/multimodal_gen/runtime/postprocess/__init__.py @@ -3,6 +3,7 @@ from sglang.multimodal_gen.runtime.postprocess.realesrgan_upscaler import ( ImageUpscaler, + batch_upscale_frames, upscale_frames, ) from sglang.multimodal_gen.runtime.postprocess.rife_interpolator import ( @@ -14,5 +15,6 @@ __all__ = [ "FrameInterpolator", "interpolate_video_frames", "ImageUpscaler", + "batch_upscale_frames", "upscale_frames", ] diff --git a/python/sglang/multimodal_gen/runtime/postprocess/realesrgan_upscaler.py b/python/sglang/multimodal_gen/runtime/postprocess/realesrgan_upscaler.py index 9b6a6a056..873e9d9c1 100644 --- a/python/sglang/multimodal_gen/runtime/postprocess/realesrgan_upscaler.py +++ b/python/sglang/multimodal_gen/runtime/postprocess/realesrgan_upscaler.py @@ -11,6 +11,7 @@ The ImageUpscaler wrapper and integration code are original work. import math import os +import time from typing import Optional import numpy as np @@ -270,6 +271,44 @@ class UpscalerModel: def dtype(self) -> torch.dtype: return next(self.net.parameters()).dtype + def _copy_input_to_device(self, frames: np.ndarray) -> torch.Tensor: + return torch.from_numpy(frames).to(self.device) + + def _preprocess_input_tensor(self, imgs_t: torch.Tensor) -> torch.Tensor: + imgs_t = imgs_t.permute(0, 3, 1, 2).to(dtype=self.dtype).mul_(1.0 / 255.0) + if self.device.type == "cuda": + imgs_t = imgs_t.contiguous(memory_format=torch.channels_last) + return imgs_t + + @staticmethod + def _postprocess_output_tensor(out: torch.Tensor) -> torch.Tensor: + out = out.permute(0, 2, 3, 1).clamp(0.0, 1.0).mul_(255.0) + return out.to(torch.uint8).contiguous() + + @staticmethod + def _copy_output_to_host(out: torch.Tensor) -> np.ndarray: + return out.cpu().numpy() + + def _start_cuda_timer(self): + if self.device.type != "cuda": + return None + start = torch.cuda.Event(enable_timing=True) + end = torch.cuda.Event(enable_timing=True) + start.record() + return start, end + + @staticmethod + def _stop_cuda_timer(timer) -> None: + if timer is not None: + timer[1].record() + + @staticmethod + def _cuda_elapsed_s(timer, fallback_s: float) -> float: + if timer is None: + return fallback_s + timer[1].synchronize() + return timer[0].elapsed_time(timer[1]) / 1000.0 + def _should_use_tiled_upscale(self, h: int, w: int) -> bool: if self.device.type != "cuda": return False @@ -374,6 +413,111 @@ class UpscalerModel: out_np = out.squeeze(0).permute(1, 2, 0).clamp(0.0, 1.0).cpu().numpy() return (out_np * 255.0).astype(np.uint8) + def upscale_batch( + self, frames: list[np.ndarray], outscale: float | None = None + ) -> list[np.ndarray]: + """Upscale same-resolution HWC uint8 frames in one batched forward pass.""" + if not frames: + return [] + + h, w = frames[0].shape[:2] + if any(frame.shape[:2] != (h, w) for frame in frames): + raise ValueError("All frames in a batch must have the same resolution") + + total_start_time = time.perf_counter() + + start_time = time.perf_counter() + imgs = np.stack(frames, axis=0) + stack_duration_s = time.perf_counter() - start_time + + start_time = time.perf_counter() + h2d_timer = self._start_cuda_timer() + imgs_t = self._copy_input_to_device(imgs) + self._stop_cuda_timer(h2d_timer) + h2d_wall_duration_s = time.perf_counter() - start_time + + start_time = time.perf_counter() + input_preprocess_timer = self._start_cuda_timer() + imgs_t = self._preprocess_input_tensor(imgs_t) + self._stop_cuda_timer(input_preprocess_timer) + input_preprocess_wall_duration_s = time.perf_counter() - start_time + + start_time = time.perf_counter() + forward_timer = self._start_cuda_timer() + with torch.inference_mode(): + out = self.net(imgs_t) + self._stop_cuda_timer(forward_timer) + forward_wall_duration_s = time.perf_counter() - start_time + + resize_timer = None + resize_wall_duration_s = 0.0 + if outscale is not None and outscale != self.scale: + start_time = time.perf_counter() + resize_timer = self._start_cuda_timer() + target_h = int(h * outscale) + target_w = int(w * outscale) + out = F.interpolate( + out, size=(target_h, target_w), mode="bicubic", align_corners=False + ) + self._stop_cuda_timer(resize_timer) + resize_wall_duration_s = time.perf_counter() - start_time + + start_time = time.perf_counter() + output_postprocess_timer = self._start_cuda_timer() + out = self._postprocess_output_tensor(out) + self._stop_cuda_timer(output_postprocess_timer) + output_postprocess_wall_duration_s = time.perf_counter() - start_time + + start_time = time.perf_counter() + output_d2h_timer = self._start_cuda_timer() + out_np = self._copy_output_to_host(out) + self._stop_cuda_timer(output_d2h_timer) + output_d2h_wall_duration_s = time.perf_counter() - start_time + + start_time = time.perf_counter() + outputs = [frame for frame in out_np] + post_duration_s = time.perf_counter() - start_time + + h2d_duration_s = self._cuda_elapsed_s(h2d_timer, h2d_wall_duration_s) + input_preprocess_duration_s = self._cuda_elapsed_s( + input_preprocess_timer, input_preprocess_wall_duration_s + ) + forward_duration_s = self._cuda_elapsed_s( + forward_timer, forward_wall_duration_s + ) + resize_duration_s = self._cuda_elapsed_s(resize_timer, resize_wall_duration_s) + output_postprocess_duration_s = self._cuda_elapsed_s( + output_postprocess_timer, output_postprocess_wall_duration_s + ) + output_d2h_duration_s = self._cuda_elapsed_s( + output_d2h_timer, output_d2h_wall_duration_s + ) + total_duration_s = time.perf_counter() - total_start_time + timing_source = "cuda_event" if self.device.type == "cuda" else "wall" + logger.info( + "RealESRGAN batch upscale: batch=%d input=%dx%d native_scale=%dx outscale=%s " + "dtype=%s timing=%s total=%.3fs stack=%.3fs input_h2d=%.3fs " + "input_pre=%.3fs forward=%.3fs resize=%.3fs output_post=%.3fs " + "output_d2h=%.3fs python_post=%.3fs", + len(frames), + w, + h, + self.scale, + outscale if outscale is not None else self.scale, + self.dtype, + timing_source, + total_duration_s, + stack_duration_s, + h2d_duration_s, + input_preprocess_duration_s, + forward_duration_s, + resize_duration_s, + output_postprocess_duration_s, + output_d2h_duration_s, + post_duration_s, + ) + return outputs + # --------------------------------------------------------------------------- # ImageUpscaler public class @@ -472,6 +616,42 @@ class ImageUpscaler: outscale = self._scale if self._scale != model.scale else None return [model.upscale(frame, outscale=outscale) for frame in frames] + def upscale_batched(self, frames: list[np.ndarray]) -> list[np.ndarray]: + """Upscale HWC uint8 frames with batched forwards grouped by resolution.""" + if not frames: + return frames + total_start_time = time.perf_counter() + model = self._ensure_model_loaded() + outscale = self._scale if self._scale != model.scale else None + output_frames: list[np.ndarray | None] = [None] * len(frames) + groups: dict[tuple[int, ...], list[int]] = {} + for idx, frame in enumerate(frames): + groups.setdefault(tuple(frame.shape), []).append(idx) + + for shape, indices in groups.items(): + logger.info( + "RealESRGAN upscale group: frames=%d shape=%s indices=%s", + len(indices), + shape, + indices, + ) + group_frames = [frames[idx] for idx in indices] + group_outputs = model.upscale_batch(group_frames, outscale=outscale) + for idx, output in zip(indices, group_outputs): + output_frames[idx] = output + + if any(frame is None for frame in output_frames): + raise RuntimeError("RealESRGAN batch upscale did not produce all frames") + + total_duration_s = time.perf_counter() - total_start_time + logger.info( + "RealESRGAN batch_upscale_frames completed in %.3f seconds for %d frames across %d groups", + total_duration_s, + len(frames), + len(groups), + ) + return [frame for frame in output_frames if frame is not None] + # --------------------------------------------------------------------------- # HF download helper @@ -561,6 +741,28 @@ def upscale_frames( List of upscaled uint8 HWC numpy frames. """ upscaler = ImageUpscaler( - model_path=model_path, scale=scale, half_precision=half_precision + model_path=model_path, + scale=scale, + half_precision=half_precision, ) return upscaler.upscale(frames) + + +def batch_upscale_frames( + frames: list[np.ndarray], + model_path: Optional[str] = None, + scale: int = 4, +) -> list[np.ndarray]: + """ + Batched Real-ESRGAN upscaling for realtime video paths. + + The default ``upscale_frames`` API intentionally keeps its original + per-frame behavior. Call this helper only when the caller can tolerate + batched execution and same-shape grouping semantics. + """ + upscaler = ImageUpscaler( + model_path=model_path, + scale=scale, + half_precision=current_platform.is_cuda(), + ) + return upscaler.upscale_batched(frames) diff --git a/python/sglang/multimodal_gen/runtime/realtime/__init__.py b/python/sglang/multimodal_gen/runtime/realtime/__init__.py new file mode 100644 index 000000000..2dd5abad1 --- /dev/null +++ b/python/sglang/multimodal_gen/runtime/realtime/__init__.py @@ -0,0 +1,31 @@ +# SPDX-License-Identifier: Apache-2.0 + +"""session-scoped realtime state, control events, and runtime-only helpers""" + +from sglang.multimodal_gen.runtime.realtime.causal_state import RealtimeCausalDiTState +from sglang.multimodal_gen.runtime.realtime.condition_events import ( + ConditionEvent, + ConditionEventQueue, + ConditionSamplingParams, + ControlSignal, + ControlStateSamplingQueue, + ControlStateTransition, +) +from sglang.multimodal_gen.runtime.realtime.session import ( + BaseRealtimeState, + RealtimeSession, + RealtimeSessionCache, +) + +__all__ = [ + "BaseRealtimeState", + "ConditionEvent", + "ConditionEventQueue", + "ConditionSamplingParams", + "ControlSignal", + "ControlStateSamplingQueue", + "ControlStateTransition", + "RealtimeCausalDiTState", + "RealtimeSession", + "RealtimeSessionCache", +] diff --git a/python/sglang/multimodal_gen/runtime/realtime/causal_state.py b/python/sglang/multimodal_gen/runtime/realtime/causal_state.py new file mode 100644 index 000000000..89ebb98fd --- /dev/null +++ b/python/sglang/multimodal_gen/runtime/realtime/causal_state.py @@ -0,0 +1,20 @@ +# SPDX-License-Identifier: Apache-2.0 + +from sglang.multimodal_gen.runtime.realtime.session import BaseRealtimeState + + +class RealtimeCausalDiTState(BaseRealtimeState): + """persist causal DiT cache and frame position across realtime chunks""" + + def __init__(self): + super().__init__() + self.kv_cache = None + self.crossattn_cache = None + self.current_chunk_start_frame: int = 0 + self.chunk_idx: int = 0 + + def dispose(self) -> None: + self.kv_cache = None + self.crossattn_cache = None + self.current_chunk_start_frame = 0 + self.chunk_idx = 0 diff --git a/python/sglang/multimodal_gen/runtime/realtime/condition_events.py b/python/sglang/multimodal_gen/runtime/realtime/condition_events.py new file mode 100644 index 000000000..d18c589a6 --- /dev/null +++ b/python/sglang/multimodal_gen/runtime/realtime/condition_events.py @@ -0,0 +1,289 @@ +# SPDX-License-Identifier: Apache-2.0 + +from __future__ import annotations + +import copy +from collections import deque +from collections.abc import Mapping, Sequence +from dataclasses import dataclass +from typing import Any + +_MISSING = object() + + +@dataclass(frozen=True) +class ControlSignal: + kind: str + payload: Any + timestamp_ms: int | None = None + seq_id: int | None = None + + +@dataclass(frozen=True) +class ControlStateTransition: + payload: Any + timestamp_ms: int | None = None + seq_id: int | None = None + + +@dataclass(frozen=True) +class ConditionEvent: + """transport envelope for one or more same-kind control signals""" + + kind: str + payload: Any + + def iter_signals(self, expand_payload: bool = True): + items = ( + self.payload + if self._should_expand_payload(self.payload, expand_payload) + else (self.payload,) + ) + for item in items: + if isinstance(item, ControlSignal): + if item.kind != self.kind: + raise ValueError( + "control signal kind " + f"{item.kind!r} does not match event kind {self.kind!r}" + ) + yield item + else: + yield ControlSignal(kind=self.kind, payload=item) + + @staticmethod + def _should_expand_payload(payload: Any, expand_payload: bool) -> bool: + return ( + expand_payload + and isinstance(payload, Sequence) + and not isinstance(payload, (str, bytes, bytearray)) + ) + + +@dataclass(frozen=True) +class ConditionSamplingParams: + chunk_size: int + default_item: Any = _MISSING + repeat_last: bool = True + repeat_last_across_empty_chunks: bool = False + expand_payload: bool = True + + +class ConditionEventQueue: + """per-session queue for prompt, camera, audio, and future events + + all events are stored here for per-chunk sampling + """ + + def __init__( + self, + max_events: int | Mapping[str, int] = 512, + ) -> None: + self._max_events = max_events + self._events: dict[str, deque[ConditionEvent]] = {} + self._pending_signals: dict[str, deque[ControlSignal]] = {} + self._last_payloads: dict[str, Any] = {} + self._last_sampled_seq_ids: dict[str, int | None] = {} + self._seen_kinds: set[str] = set() + + def push(self, event: ConditionEvent) -> None: + queue = self._queue_for(event.kind) + queue.append(event) + self._seen_kinds.add(event.kind) + + def replace(self, event: ConditionEvent) -> None: + self.clear_kind(event.kind) + self.push(event) + + def pop_latest(self, kind: str) -> Any | None: + queue = self._events.get(kind) + if not queue: + return None + latest_payload = None + has_signal = False + for signal in queue.pop().iter_signals(): + latest_payload = signal.payload + has_signal = True + self._last_sampled_seq_ids[kind] = signal.seq_id + queue.clear() + self._seen_kinds.add(kind) + if not has_signal: + return None + return latest_payload + + def has_events(self, kind: str) -> bool: + queue = self._events.get(kind) + pending = self._pending_signals.get(kind) + return bool(queue) or bool(pending) + + def sample_chunk( + self, + kind: str, + params: ConditionSamplingParams, + ) -> list[Any] | None: + """samples a list of actions for a chunk + + Args: + params: the sampling strategy + + """ + if params.chunk_size <= 0: + return None + + chunk: list[Any] = [] + pending = self._pending_signals.get(kind) + self._drain_signals(kind, pending, chunk, params.chunk_size) + + queue = self._events.get(kind) + while len(chunk) < params.chunk_size and queue: + event = queue.popleft() + signals = deque(event.iter_signals(params.expand_payload)) + self._drain_signals(kind, signals, chunk, params.chunk_size) + if signals: + self._pending_signals[kind] = signals + + if len(chunk) == 0 and kind not in self._seen_kinds: + if params.default_item is _MISSING: + return None + return [params.default_item for _ in range(params.chunk_size)] + + if len(chunk) == 0: + if params.repeat_last_across_empty_chunks and kind in self._last_payloads: + return [self._last_payloads[kind] for _ in range(params.chunk_size)] + if params.default_item is _MISSING: + return None + return [params.default_item for _ in range(params.chunk_size)] + + if not params.repeat_last: + return chunk + + pad_item = self._last_payloads.get(kind, params.default_item) + if pad_item is _MISSING: + return chunk + while len(chunk) < params.chunk_size: + chunk.append(pad_item) + return chunk + + def clear(self) -> None: + self._events.clear() + self._pending_signals.clear() + self._last_payloads.clear() + self._last_sampled_seq_ids.clear() + self._seen_kinds.clear() + + def clear_kind(self, kind: str) -> None: + self._events.pop(kind, None) + self._pending_signals.pop(kind, None) + self._last_payloads.pop(kind, None) + self._last_sampled_seq_ids.pop(kind, None) + self._seen_kinds.discard(kind) + + def last_sampled_seq_id(self, kind: str) -> int | None: + return self._last_sampled_seq_ids.get(kind) + + def _queue_for(self, kind: str) -> deque[ConditionEvent]: + queue = self._events.get(kind) + if queue is None: + if isinstance(self._max_events, Mapping): + maxlen = self._max_events.get(kind, 512) + else: + maxlen = self._max_events + queue = deque(maxlen=maxlen) + self._events[kind] = queue + return queue + + def _drain_signals( + self, + kind: str, + signals: deque[ControlSignal] | None, + chunk: list[Any], + chunk_size: int, + ) -> None: + while signals and len(chunk) < chunk_size: + signal = signals.popleft() + chunk.append(signal.payload) + self._last_payloads[kind] = signal.payload + self._last_sampled_seq_ids[kind] = signal.seq_id + + +class ControlStateSamplingQueue: + """state-based control sampler for realtime inputs""" + + def __init__( + self, + *, + default_item: Any, + min_pulse_items: int = 1, + max_transitions: int = 512, + ) -> None: + self.default_item = default_item + self.min_pulse_items = min_pulse_items + self._pending: deque[ControlStateTransition] = deque(maxlen=max_transitions) + self._current_item = default_item + self._current_seq_id: int | None = None + self._latest_sampled_seq_id: int | None = None + + def clear(self) -> None: + self._pending.clear() + self._current_item = self.default_item + self._current_seq_id = None + self._latest_sampled_seq_id = None + + def push(self, transition: ControlStateTransition) -> None: + self._pending.append(transition) + + def push_many(self, transitions: Sequence[ControlStateTransition]) -> None: + for transition in transitions: + self.push(transition) + + def sample_chunk(self, chunk_size: int) -> list[Any] | None: + if chunk_size <= 0: + return None + + transitions = self._drain_pending() + if not transitions: + self._latest_sampled_seq_id = self._current_seq_id + return [self._copy_item(self._current_item) for _ in range(chunk_size)] + + pulse = self._latest_non_default_transition(transitions) + final = transitions[-1] + self._current_item = final.payload + self._current_seq_id = final.seq_id + + if pulse is not None and pulse.payload != final.payload: + pulse_items = min(self.min_pulse_items, chunk_size) + chunk = [self._copy_item(pulse.payload) for _ in range(pulse_items)] + chunk.extend( + self._copy_item(final.payload) for _ in range(chunk_size - pulse_items) + ) + self._latest_sampled_seq_id = ( + final.seq_id if len(chunk) > pulse_items else pulse.seq_id + ) + return chunk + + self._latest_sampled_seq_id = final.seq_id + return [self._copy_item(final.payload) for _ in range(chunk_size)] + + def latest_sampled_seq_id(self) -> int | None: + return self._latest_sampled_seq_id + + def _drain_pending(self) -> list[ControlStateTransition]: + transitions = list(self._pending) + self._pending.clear() + return transitions + + def _latest_non_default_transition( + self, + transitions: Sequence[ControlStateTransition], + ) -> ControlStateTransition | None: + for transition in reversed(transitions): + if transition.payload != self.default_item: + return transition + return None + + @staticmethod + def _copy_item(item: Any) -> Any: + if isinstance(item, list): + return copy.deepcopy(item) + if isinstance(item, dict): + return copy.deepcopy(item) + return item diff --git a/python/sglang/multimodal_gen/runtime/realtime/session.py b/python/sglang/multimodal_gen/runtime/realtime/session.py new file mode 100644 index 000000000..e1f2a9e19 --- /dev/null +++ b/python/sglang/multimodal_gen/runtime/realtime/session.py @@ -0,0 +1,111 @@ +# SPDX-License-Identifier: Apache-2.0 + +from __future__ import annotations + +from collections import OrderedDict +from typing import Any + +from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger + +logger = init_logger(__name__) + + +class BaseRealtimeState: + """per-session state owned by pipeline stages""" + + def dispose(self) -> None: + pass + + +class RealtimeSession: + """reusable state container across realtime request chunks""" + + def __init__(self) -> None: + self._states: dict[type[BaseRealtimeState], BaseRealtimeState] = {} + + @staticmethod + def resolve_session_id(req: Any) -> str | None: + session_id = req.realtime_session_id + if isinstance(session_id, str) and session_id: + return session_id + return None + + def get_or_create_state( + self, state_cls: type[BaseRealtimeState] + ) -> BaseRealtimeState: + state = self._states.get(state_cls) + if state is None: + state = state_cls() + self._states[state_cls] = state + return state + + def get_state(self, state_cls: type[BaseRealtimeState]) -> BaseRealtimeState | None: + return self._states.get(state_cls) + + def dispose(self) -> None: + for state in list(self._states.values()): + state.dispose() + self._states.clear() + + +class RealtimeSessionCache: + """lru cache that binds incoming chunks to persistent realtime sessions""" + + def __init__(self, max_sessions: int = 64) -> None: + self.max_sessions = max_sessions + self._sessions: OrderedDict[str, RealtimeSession] = OrderedDict() + + def _dispose_session( + self, session_id: str, session: RealtimeSession | None + ) -> None: + if session is None: + return + try: + session.dispose() + except Exception as e: + logger.warning( + "Failed to dispose realtime session cache entry %s: %s", + session_id, + e, + ) + + def release(self, session_id: str) -> bool: + session = self._sessions.pop(session_id, None) + released = session is not None + self._dispose_session(session_id, session) + logger.info( + "Realtime session release: session_id=%s released=%s", + session_id, + released, + ) + return released + + def attach(self, req: Any) -> None: + session_id = RealtimeSession.resolve_session_id(req) + if session_id is None: + return + + if session_id not in self._sessions: + if req.block_idx > 0: + raise ValueError( + "Missing realtime session state for " + f"session_id={session_id} block_idx={req.block_idx}." + ) + self._sessions[session_id] = req.session or RealtimeSession() + elif req.block_idx == 0: + old_session = self._sessions[session_id] + new_session = req.session or RealtimeSession() + if old_session is not new_session: + self._dispose_session(session_id, old_session) + self._sessions[session_id] = new_session + logger.info("Realtime session reset: session_id=%s", session_id) + + req.session = self._sessions[session_id] + self._sessions.move_to_end(session_id) + self._evict_stale_sessions() + + def _evict_stale_sessions(self) -> None: + while len(self._sessions) > self.max_sessions: + stale_session_id, stale_session = self._sessions.popitem(last=False) + self._dispose_session(stale_session_id, stale_session) + logger.debug("Evicted stale realtime session cache: %s", stale_session_id) diff --git a/python/sglang/multimodal_gen/runtime/utils/camera_geometry.py b/python/sglang/multimodal_gen/runtime/utils/camera_geometry.py new file mode 100644 index 000000000..3f78af22b --- /dev/null +++ b/python/sglang/multimodal_gen/runtime/utils/camera_geometry.py @@ -0,0 +1,118 @@ +# SPDX-License-Identifier: Apache-2.0 + +"""Camera pose and Plucker ray utilities.""" + +from __future__ import annotations + +import torch + + +def se3_inverse(T: torch.Tensor) -> torch.Tensor: + rot = T[:, :3, :3] + trans = T[:, :3, 3:] + r_inv = rot.transpose(-1, -2) + t_inv = -torch.bmm(r_inv, trans) + T_inv = torch.eye(4, device=T.device, dtype=T.dtype)[None, :, :].repeat( + T.shape[0], 1, 1 + ) + T_inv[:, :3, :3] = r_inv + T_inv[:, :3, 3:] = t_inv + return T_inv + + +def compute_relative_poses( + c2ws_mat: torch.Tensor, + framewise: bool = False, + normalize_trans: bool = True, +) -> torch.Tensor: + ref_w2cs = se3_inverse(c2ws_mat[0:1]) + relative_poses = torch.matmul(ref_w2cs, c2ws_mat) + relative_poses[0] = torch.eye(4, device=c2ws_mat.device, dtype=c2ws_mat.dtype) + if framewise and len(relative_poses) > 1: + relative_poses_framewise = torch.bmm( + se3_inverse(relative_poses[:-1]), relative_poses[1:] + ) + relative_poses[1:] = relative_poses_framewise + if normalize_trans: + translations = relative_poses[:, :3, 3] + max_norm = torch.norm(translations, dim=-1).max() + if max_norm > 0: + relative_poses[:, :3, 3] = translations / max_norm + return relative_poses + + +@torch.no_grad() +def create_meshgrid( + n_frames: int, + height: int, + width: int, + *, + bias: float = 0.5, + device: torch.device | str, + dtype: torch.dtype, +) -> torch.Tensor: + x_range = torch.arange(width, device=device, dtype=dtype) + y_range = torch.arange(height, device=device, dtype=dtype) + grid_y, grid_x = torch.meshgrid(y_range, x_range, indexing="ij") + grid_xy = torch.stack([grid_x, grid_y], dim=-1).view([-1, 2]) + bias + return grid_xy[None, ...].repeat(n_frames, 1, 1) + + +def get_plucker_embeddings( + c2ws_mat: torch.Tensor, + Ks: torch.Tensor, + height: int, + width: int, +) -> torch.Tensor: + n_frames = c2ws_mat.shape[0] + grid_xy = create_meshgrid( + n_frames, height, width, device=c2ws_mat.device, dtype=c2ws_mat.dtype + ) + fx, fy, cx, cy = Ks.chunk(4, dim=-1) + i = grid_xy[..., 0] + j = grid_xy[..., 1] + zs = torch.ones_like(i) + xs = (i - cx) / fx * zs + ys = (j - cy) / fy * zs + directions = torch.stack([xs, ys, zs], dim=-1) + directions = directions / directions.norm(dim=-1, keepdim=True) + rays_d = directions @ c2ws_mat[:, :3, :3].transpose(-1, -2) + rays_o = c2ws_mat[:, :3, 3][:, None, :].expand_as(rays_d) + plucker_embeddings = torch.cat([rays_o, rays_d], dim=-1) + return plucker_embeddings.view([n_frames, height, width, 6]) + + +def camera_poses_to_plucker( + *, + c2ws: torch.Tensor, + Ks: torch.Tensor, + height: int, + width: int, + spatial_scale: int = 8, + device: torch.device | str, + dtype: torch.dtype, +) -> torch.Tensor: + plucker = get_plucker_embeddings(c2ws, Ks, height, width) + latent_height = height // spatial_scale + latent_width = width // spatial_scale + plucker = plucker.view( + c2ws.shape[0], + latent_height, + spatial_scale, + latent_width, + spatial_scale, + 6, + ) + plucker = plucker.permute(0, 1, 3, 5, 2, 4).contiguous() + plucker = plucker.view( + c2ws.shape[0], + latent_height, + latent_width, + 6 * spatial_scale * spatial_scale, + ) + return ( + plucker.permute(3, 0, 1, 2) + .contiguous() + .unsqueeze(0) + .to(device=device, dtype=dtype) + ) diff --git a/python/sglang/multimodal_gen/runtime/utils/realtime_video.py b/python/sglang/multimodal_gen/runtime/utils/realtime_video.py new file mode 100644 index 000000000..c651209ee --- /dev/null +++ b/python/sglang/multimodal_gen/runtime/utils/realtime_video.py @@ -0,0 +1,211 @@ +# SPDX-License-Identifier: Apache-2.0 + +from __future__ import annotations + +import time +import zlib +from collections.abc import Callable, Sequence +from typing import TYPE_CHECKING, Any + +import numpy as np +import torch + +from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger + +if TYPE_CHECKING: + from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import ( + OutputBatch, + Req, + ) + +logger = init_logger(__name__) + +RAW_RGB_CONTENT_TYPE = "application/x-raw-rgb" +RAW_RGB_DELTA_GZIP_CONTENT_TYPE = "application/x-raw-rgb-delta-gzip" +RAW_RGBA_DELTA_GZIP_CONTENT_TYPE = "application/x-raw-rgba-delta-gzip" +WEBP_FRAME_CONTENT_TYPE = "image/webp" +JPEG_FRAME_CONTENT_TYPE = "image/jpeg" +RAW_RGB_CHANNELS = 3 +RAW_RGBA_CHANNELS = 4 + + +def build_delta_gzip_raw_rgb_payload( + frames: list[bytes], + *, + reference_frame: bytes | None = None, +) -> bytes: + if not frames: + return b"" + + frame_size = len(frames[0]) + if reference_frame is not None and len(reference_frame) != frame_size: + raise ValueError("raw RGB delta gzip reference frame size mismatch") + + previous = ( + np.frombuffer(reference_frame, dtype=np.uint8) + if reference_frame is not None + else None + ) + compressor = zlib.compressobj(level=1, method=zlib.DEFLATED, wbits=31) + compressed_chunks = [] + for frame in frames: + if len(frame) != frame_size: + raise ValueError("raw RGB delta gzip requires fixed-size frames") + current = np.frombuffer(frame, dtype=np.uint8) + if previous is None: + delta_frame = frame + else: + delta_frame = np.bitwise_xor(current, previous).tobytes() + compressed_chunks.append(compressor.compress(delta_frame)) + previous = current + + compressed_chunks.append(compressor.flush()) + return b"".join(compressed_chunks) + + +def restore_delta_gzip_raw_rgb_payload( + payload: bytes, + *, + bytes_per_frame: int, + num_frames: int, + reference_frame: bytes | None = None, +) -> bytes: + if reference_frame is not None and len(reference_frame) != bytes_per_frame: + raise ValueError("delta gzip reference frame size mismatch") + + delta_payload = zlib.decompress(payload, wbits=31) + expected_size = bytes_per_frame * num_frames + if len(delta_payload) != expected_size: + raise ValueError( + "delta gzip payload size mismatch: " + f"expected {expected_size}, got {len(delta_payload)}" + ) + + restored = bytearray(delta_payload) + previous = ( + np.frombuffer(reference_frame, dtype=np.uint8) + if reference_frame is not None + else None + ) + for frame_idx in range(num_frames): + offset = frame_idx * bytes_per_frame + current = np.frombuffer( + restored, dtype=np.uint8, count=bytes_per_frame, offset=offset + ) + if previous is not None: + current ^= previous + previous = current + return bytes(restored) + + +def build_raw_rgb_frame_batches( + output: Any, + req: Req, + output_batch: OutputBatch, + post_process_sample_fn: Callable[..., Any], +) -> tuple[list[list[bytes]], dict[str, Any]]: + """post-process for realtime responses, returns only the batched frames and metadata""" + start = time.monotonic() + sample_to_frames_ms = 0.0 + frames_to_bytes_ms = 0.0 + raw_bytes = 0 + num_frames = 0 + frame_shape = None + frame_batches = [] + if isinstance(output, torch.Tensor): + outputs = list(output) + else: + outputs = output if isinstance(output, Sequence) else [output] + + for sample in outputs: + stage_start = time.monotonic() + if ( + isinstance(sample, torch.Tensor) + and not req.enable_frame_interpolation + and not req.enable_upscaling + ): + frames = _tensor_sample_to_rgb24_array(sample) + else: + frames = post_process_sample_fn( + sample, + req.data_type, + req.fps, + False, + None, + audio_sample_rate=output_batch.audio_sample_rate, + output_compression=req.output_compression, + enable_frame_interpolation=req.enable_frame_interpolation, + frame_interpolation_exp=req.frame_interpolation_exp, + frame_interpolation_scale=req.frame_interpolation_scale, + frame_interpolation_model_path=req.frame_interpolation_model_path, + enable_upscaling=False, + upscaling_model_path=req.upscaling_model_path, + upscaling_scale=req.upscaling_scale, + ) + if req.enable_upscaling and frames: + from sglang.multimodal_gen.runtime.postprocess import ( + batch_upscale_frames, + ) + + frames = batch_upscale_frames( + frames, + model_path=req.upscaling_model_path, + scale=req.upscaling_scale, + ) + sample_to_frames_ms += (time.monotonic() - stage_start) * 1000.0 + + stage_start = time.monotonic() + + # numpy frames to RGB24 bytes + raw_frames = [] + for frame in frames: + if frame.ndim == 2: + frame = frame[:, :, None] + if frame.shape[-1] == 1: + frame = np.repeat(frame, 3, axis=-1) + elif frame.shape[-1] > RAW_RGB_CHANNELS: + frame = frame[:, :, :RAW_RGB_CHANNELS] + frame = np.ascontiguousarray(frame) + frame_shape = tuple(int(dim) for dim in frame.shape) + frame_bytes = frame.tobytes() + raw_bytes += len(frame_bytes) + num_frames += 1 + raw_frames.append(frame_bytes) + frames_to_bytes_ms += (time.monotonic() - stage_start) * 1000.0 + frame_batches.append(raw_frames) + + total_ms = (time.monotonic() - start) * 1000.0 + logger.info( + "realtime raw RGB frame batch timing: request_id=%s " + "chunk_idx=%s sample_to_frames=%.2fms frames_to_bytes=%.2fms " + "total=%.2fms batches=%d frames=%d frame_shape=%s " + "raw_bytes=%d content_type=%s", + req.request_id, + req.block_idx, + sample_to_frames_ms, + frames_to_bytes_ms, + total_ms, + len(frame_batches), + num_frames, + frame_shape, + raw_bytes, + RAW_RGB_CONTENT_TYPE, + ) + frame_metadata: dict[str, Any] = {} + if frame_shape is not None and len(frame_shape) == 3: + frame_height, frame_width, channels = frame_shape + frame_metadata = { + "format": "rgb24", + "width": frame_width, + "height": frame_height, + "channels": channels, + "bytes_per_frame": frame_width * frame_height * channels, + } + return frame_batches, frame_metadata + + +def _tensor_sample_to_rgb24_array(sample: torch.Tensor) -> np.ndarray: + if sample.dim() == 3: + sample = sample.unsqueeze(1) + sample = (sample * 255).clamp(0, 255).to(torch.uint8) + return sample.permute(1, 2, 3, 0).contiguous().cpu().numpy() diff --git a/python/sglang/multimodal_gen/test/server/gpu_cases.py b/python/sglang/multimodal_gen/test/server/gpu_cases.py index f141f4efd..5f7a101f4 100644 --- a/python/sglang/multimodal_gen/test/server/gpu_cases.py +++ b/python/sglang/multimodal_gen/test/server/gpu_cases.py @@ -18,6 +18,7 @@ from sglang.multimodal_gen.test.server.testcase_configs import ( DiffusionSamplingParams, DiffusionServerArgs, DiffusionTestCase, + LINGBOT_WORLD_REALTIME_sampling_params, MODELOPT_T2I_CI_sampling_params, MODELOPT_T2V_CI_sampling_params, MODELOPT_TI2I_CI_sampling_params, @@ -54,6 +55,24 @@ from sglang.multimodal_gen.test.test_utils import ( _CACHE_DIT_CONFIG_DIR = Path(__file__).parent / "configs" + +def _make_lingbot_realtime_plastic_beach_case() -> DiffusionTestCase: + return DiffusionTestCase( + "lingbot_world_realtime_plastic_beach", + DiffusionServerArgs( + model_path="robbyant/lingbot-world-fast-diffusers", + modality="video", + num_gpus=1, + extras=["--pipeline-class-name LingBotWorldCausalDMDPipeline"], + text_encoder_cpu_offload=True, + ), + LINGBOT_WORLD_REALTIME_sampling_params, + run_component_accuracy_check=False, + run_models_api_check=False, + run_t2v_input_reference_check=False, + ) + + # All test cases with clean default values # To test different models, simply add more DiffusionCase entries ONE_GPU_CASES: list[DiffusionTestCase] = [ @@ -743,5 +762,7 @@ if not current_platform.is_hip(): ) ) +ONE_GPU_CASES.append(_make_lingbot_realtime_plastic_beach_case()) + ONE_GPU_CASES += ONE_GPU_MODELOPT_FP8_CASES TWO_GPU_CASES = _with_default_num_gpus(TWO_GPU_CASES, 2) diff --git a/python/sglang/multimodal_gen/test/server/perf_baselines.json b/python/sglang/multimodal_gen/test/server/perf_baselines.json index 5baea2765..e136acbd8 100644 --- a/python/sglang/multimodal_gen/test/server/perf_baselines.json +++ b/python/sglang/multimodal_gen/test/server/perf_baselines.json @@ -2605,6 +2605,14 @@ "expected_median_denoise_ms": 246.85, "estimated_full_test_time_s": 170.0 }, + "lingbot_world_realtime_plastic_beach": { + "stages_ms": {}, + "denoise_step_ms": {}, + "expected_e2e_ms": 0.0, + "expected_avg_denoise_ms": 0.0, + "expected_median_denoise_ms": 0.0, + "estimated_full_test_time_s": 126.0 + }, "ltx_2_3_hq_pipeline": { "stages_ms": { "InputValidationStage": 0.08, diff --git a/python/sglang/multimodal_gen/test/server/realtime_consistency.py b/python/sglang/multimodal_gen/test/server/realtime_consistency.py new file mode 100644 index 000000000..bd64fe27c --- /dev/null +++ b/python/sglang/multimodal_gen/test/server/realtime_consistency.py @@ -0,0 +1,403 @@ +# SPDX-License-Identifier: Apache-2.0 + +from __future__ import annotations + +import asyncio +import os +import statistics +import tempfile +from dataclasses import dataclass +from pathlib import Path +from typing import Any + +import imageio +import msgspec.msgpack +import numpy as np +import pytest +from openai import Client + +from sglang.multimodal_gen.runtime.utils.realtime_video import ( + RAW_RGB_CONTENT_TYPE, + RAW_RGB_DELTA_GZIP_CONTENT_TYPE, + RAW_RGBA_DELTA_GZIP_CONTENT_TYPE, + restore_delta_gzip_raw_rgb_payload, +) +from sglang.multimodal_gen.test.server.testcase_configs import DiffusionSamplingParams +from sglang.multimodal_gen.test.test_utils import is_image_url + +_REALTIME_WS_TIMEOUT_SECS = float( + os.environ.get("SGLANG_TEST_REALTIME_WS_TIMEOUT_SECS", "1200") +) + + +@dataclass(frozen=True) +class RealtimeChunkStats: + chunk_index: int + request_id: str | None + content_type: str + num_frames: int + raw_bytes: int + ws_payload_bytes: int + request_prepare_ms: float + scheduler_forward_ms: float + raw_payload_build_ms: float + raw_write_ms: float + ws_write_ms: float + chunk_total_ms: float + + +@dataclass(frozen=True) +class RealtimeCollectionResult: + frames: list[np.ndarray] + chunk_stats: list[RealtimeChunkStats] + + +_REALTIME_CHUNK_STATS_BY_CASE: dict[str, list[RealtimeChunkStats]] = {} +_REALTIME_KEY_FRAMES_BY_CASE: dict[str, list[np.ndarray]] = {} + + +def realtime_ws_url(client: Client) -> str: + base_url = str(client.base_url).rstrip("/") + if base_url.startswith("https://"): + return "wss://" + base_url[len("https://") :] + "/realtime_video/generate" + if base_url.startswith("http://"): + return "ws://" + base_url[len("http://") :] + "/realtime_video/generate" + raise ValueError(f"Unsupported realtime client base_url: {base_url}") + + +def prepare_realtime_first_frame( + image_path: Path | str | list[Path | str] | None, +) -> bytes | str | None: + if image_path is None: + return None + if isinstance(image_path, list): + image_path = image_path[0] + if isinstance(image_path, str) and is_image_url(image_path): + return image_path + path = Path(image_path) + if not path.exists(): + raise FileNotFoundError(f"Realtime first frame file missing: {path}") + return path.read_bytes() + + +def build_realtime_init_payload( + *, + model_path: str, + sampling_params: DiffusionSamplingParams, + output_size: str, + first_frame: bytes | str | None, +) -> dict[str, Any]: + payload: dict[str, Any] = { + "type": "init", + "model": model_path, + "prompt": sampling_params.prompt, + "size": output_size, + "seconds": sampling_params.seconds, + "first_frame": first_frame, + } + optional_fields = { + "fps": sampling_params.fps, + "num_frames": sampling_params.num_frames, + "realtime_output_format": sampling_params.realtime_output_format, + } + payload.update({k: v for k, v in optional_fields.items() if v is not None}) + payload.update(dict(sampling_params.extras)) + return {k: v for k, v in payload.items() if v is not None} + + +def build_realtime_event_payload(event: dict[str, Any]) -> dict[str, Any]: + payload = dict(event) + payload.pop("after_chunk", None) + payload.setdefault("type", "event") + if "kind" not in payload: + raise ValueError("realtime event config must include kind") + return payload + + +def parse_realtime_chunk_stats(header: dict[str, Any]) -> RealtimeChunkStats: + if header.get("type") != "chunk_stats": + raise ValueError(f"Unexpected realtime chunk stats message: {header}") + return RealtimeChunkStats( + chunk_index=int(header["chunk_index"]), + request_id=header.get("request_id"), + content_type=str(header.get("content_type", "")), + num_frames=int(header.get("num_frames", 0)), + raw_bytes=int(header.get("raw_bytes", 0)), + ws_payload_bytes=int(header.get("ws_payload_bytes", 0)), + request_prepare_ms=float(header.get("request_prepare_ms", 0.0)), + scheduler_forward_ms=float(header.get("scheduler_forward_ms", 0.0)), + raw_payload_build_ms=float(header.get("raw_payload_build_ms", 0.0)), + raw_write_ms=float(header.get("raw_write_ms", 0.0)), + ws_write_ms=float(header.get("ws_write_ms", 0.0)), + chunk_total_ms=float(header.get("chunk_total_ms", 0.0)), + ) + + +def summarize_realtime_perf_stats( + chunk_stats: list[RealtimeChunkStats], + *, + ignore_initial_chunks: int = 0, +) -> dict[str, float]: + if not chunk_stats: + return {} + if ignore_initial_chunks < 0: + raise ValueError("ignore_initial_chunks must be non-negative") + if ignore_initial_chunks >= len(chunk_stats): + raise ValueError( + "ignore_initial_chunks must leave at least one realtime chunk to guard" + ) + + ignored_stats = chunk_stats[:ignore_initial_chunks] + guarded_stats = chunk_stats[ignore_initial_chunks:] + + metrics = { + "request_prepare_ms": [s.request_prepare_ms for s in guarded_stats], + "scheduler_forward_ms": [s.scheduler_forward_ms for s in guarded_stats], + "raw_payload_build_ms": [s.raw_payload_build_ms for s in guarded_stats], + "raw_write_ms": [s.raw_write_ms for s in guarded_stats], + "ws_write_ms": [s.ws_write_ms for s in guarded_stats], + "chunk_total_ms": [s.chunk_total_ms for s in guarded_stats], + "ws_payload_mb": [s.ws_payload_bytes / (1024 * 1024) for s in guarded_stats], + } + summary: dict[str, float] = { + "num_chunks": float(len(chunk_stats)), + "total_frames": float(sum(s.num_frames for s in chunk_stats)), + "ignored_initial_chunks": float(ignore_initial_chunks), + "guarded_chunks": float(len(guarded_stats)), + } + if ignored_stats: + summary["ignored_max_chunk_total_ms"] = max( + s.chunk_total_ms for s in ignored_stats + ) + summary["ignored_max_scheduler_forward_ms"] = max( + s.scheduler_forward_ms for s in ignored_stats + ) + for name, values in metrics.items(): + sorted_values = sorted(values) + p95_idx = min(len(sorted_values) - 1, int(len(sorted_values) * 0.95)) + summary[f"avg_{name}"] = statistics.fmean(values) + summary[f"p95_{name}"] = sorted_values[p95_idx] + summary[f"max_{name}"] = max(values) + return summary + + +def validate_realtime_perf_stats( + case_id: str, + chunk_stats: list[RealtimeChunkStats], + thresholds: dict[str, float], + *, + ignore_initial_chunks: int = 0, +) -> None: + if not thresholds: + return + summary = summarize_realtime_perf_stats( + chunk_stats, ignore_initial_chunks=ignore_initial_chunks + ) + if not summary: + pytest.fail(f"{case_id}: no realtime chunk stats were received") + + failures = [] + for metric_name, threshold in thresholds.items(): + if metric_name not in summary: + raise ValueError( + f"{case_id}: unknown realtime perf metric {metric_name!r}; " + f"available metrics: {sorted(summary)}" + ) + actual = summary[metric_name] + if actual > threshold: + failures.append( + f"{metric_name}: actual={actual:.2f}, limit={threshold:.2f}" + ) + + if failures: + pytest.fail( + f"Realtime performance guard failed for {case_id}:\n" + + "\n".join(f" - {failure}" for failure in failures) + ) + + +def record_realtime_perf_stats( + case_id: str, chunk_stats: list[RealtimeChunkStats] +) -> None: + _REALTIME_CHUNK_STATS_BY_CASE[case_id] = list(chunk_stats) + + +def pop_realtime_perf_stats(case_id: str) -> list[RealtimeChunkStats]: + return _REALTIME_CHUNK_STATS_BY_CASE.pop(case_id, []) + + +def select_realtime_key_frames(frames: list[np.ndarray]) -> list[np.ndarray]: + if not frames: + return [] + key_indices = [0, len(frames) // 2, len(frames) - 1] + return [frames[idx].copy() for idx in key_indices] + + +def record_realtime_key_frames(case_id: str, frames: list[np.ndarray]) -> None: + _REALTIME_KEY_FRAMES_BY_CASE[case_id] = select_realtime_key_frames(frames) + + +def pop_realtime_key_frames(case_id: str) -> list[np.ndarray] | None: + return _REALTIME_KEY_FRAMES_BY_CASE.pop(case_id, None) + + +def decode_realtime_raw_rgb_frames( + header: dict[str, Any], + payload: bytes, + previous_frame: bytes | None = None, +) -> list[np.ndarray]: + content_type = header.get("content_type") + if content_type not in ( + RAW_RGB_CONTENT_TYPE, + RAW_RGB_DELTA_GZIP_CONTENT_TYPE, + RAW_RGBA_DELTA_GZIP_CONTENT_TYPE, + ): + raise ValueError(f"Unsupported realtime frame content type: {content_type}") + + width = int(header["width"]) + height = int(header["height"]) + channels = int(header["channels"]) + num_frames = int(header["num_frames"]) + bytes_per_frame = int(header["bytes_per_frame"]) + expected_size = num_frames * bytes_per_frame + if content_type in ( + RAW_RGB_DELTA_GZIP_CONTENT_TYPE, + RAW_RGBA_DELTA_GZIP_CONTENT_TYPE, + ): + if header.get("delta_reference") != "previous-frame": + previous_frame = None + payload = restore_delta_gzip_raw_rgb_payload( + payload, + bytes_per_frame=bytes_per_frame, + num_frames=num_frames, + reference_frame=previous_frame, + ) + if len(payload) != expected_size: + raise ValueError( + f"Realtime payload size mismatch: expected {expected_size}, got {len(payload)}" + ) + + frames = [] + for frame_idx in range(num_frames): + offset = frame_idx * bytes_per_frame + frame = np.frombuffer( + payload[offset : offset + bytes_per_frame], dtype=np.uint8 + ) + frame = frame.reshape(height, width, channels) + if channels > 3: + frame = frame[:, :, :3] + frames.append(frame.copy()) + return frames + + +def encode_realtime_frames_to_mp4(frames: list[np.ndarray], fps: int) -> bytes: + if not frames: + raise ValueError("Cannot encode empty realtime frame list") + + with tempfile.NamedTemporaryFile(suffix=".mp4", delete=False) as tmp: + output_path = tmp.name + try: + imageio.mimsave( + output_path, + frames, + fps=fps, + format="mp4", + codec="libx264", + quality=5, + ) + return Path(output_path).read_bytes() + finally: + try: + os.remove(output_path) + except OSError: + pass + + +async def collect_realtime_frames( + *, + ws_url: str, + init_payload: dict[str, Any], + events: list[dict[str, Any]], + num_chunks: int, +) -> list[np.ndarray]: + return ( + await collect_realtime_output( + ws_url=ws_url, + init_payload=init_payload, + events=events, + num_chunks=num_chunks, + ) + ).frames + + +async def collect_realtime_output( + *, + ws_url: str, + init_payload: dict[str, Any], + events: list[dict[str, Any]], + num_chunks: int, + require_chunk_stats: bool = False, +) -> RealtimeCollectionResult: + try: + import websockets + except ImportError: + pytest.skip("websockets is required for realtime consistency checks") + + frames: list[np.ndarray] = [] + chunk_stats: list[RealtimeChunkStats] = [] + sent_event_indices: set[int] = set() + + async def send_events_for_boundary(ws, completed_chunk: int) -> None: + for event_idx, event in enumerate(events): + if event_idx in sent_event_indices: + continue + if int(event.get("after_chunk", 0)) != completed_chunk: + continue + await ws.send(msgspec.msgpack.encode(build_realtime_event_payload(event))) + sent_event_indices.add(event_idx) + + async with websockets.connect(ws_url, max_size=None, ping_interval=None) as ws: + await ws.send(msgspec.msgpack.encode(init_payload)) + await send_events_for_boundary(ws, -1) + + received_chunks: set[int] = set() + previous_frame: bytes | None = None + while len(received_chunks) < num_chunks or ( + require_chunk_stats and len(chunk_stats) < len(received_chunks) + ): + header_payload = await asyncio.wait_for( + ws.recv(), timeout=_REALTIME_WS_TIMEOUT_SECS + ) + header = msgspec.msgpack.decode(header_payload) + message_type = header.get("type") + if message_type == "error": + pytest.fail(f"Realtime generation failed: {header.get('content')}") + if message_type == "chunk_stats": + chunk_stats.append(parse_realtime_chunk_stats(header)) + continue + if message_type == "frame_batch": + raw_payload = header.pop("payload", None) + header["type"] = "frame_batch_header" + elif message_type == "frame_batch_header": + raw_payload = await asyncio.wait_for( + ws.recv(), timeout=_REALTIME_WS_TIMEOUT_SECS + ) + else: + raise ValueError(f"Unexpected realtime message: {header}") + if not isinstance(raw_payload, bytes): + raise ValueError("Realtime frame payload must be bytes") + + chunk_frames = decode_realtime_raw_rgb_frames( + header, + raw_payload, + previous_frame, + ) + frames.extend(chunk_frames) + if chunk_frames: + previous_frame = chunk_frames[-1].tobytes() + chunk_index = int(header["chunk_index"]) + if header.get("is_final_frame_batch", True): + received_chunks.add(chunk_index) + await send_events_for_boundary(ws, chunk_index) + + return RealtimeCollectionResult(frames=frames, chunk_stats=chunk_stats) diff --git a/python/sglang/multimodal_gen/test/server/test_server_common.py b/python/sglang/multimodal_gen/test/server/test_server_common.py index 3390dc5c3..1f8da49e3 100644 --- a/python/sglang/multimodal_gen/test/server/test_server_common.py +++ b/python/sglang/multimodal_gen/test/server/test_server_common.py @@ -23,6 +23,11 @@ from sglang.multimodal_gen.runtime.platforms import current_platform from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger from sglang.multimodal_gen.runtime.utils.perf_logger import RequestPerfRecord from sglang.multimodal_gen.test.server import conftest +from sglang.multimodal_gen.test.server.realtime_consistency import ( + pop_realtime_key_frames, + pop_realtime_perf_stats, + validate_realtime_perf_stats, +) from sglang.multimodal_gen.test.server.test_server_utils import ( VALIDATOR_REGISTRY, PerformanceValidator, @@ -624,7 +629,9 @@ Pinned revision used by this check: {SGL_TEST_FILES_CI_DATA_REVISION} thresholds = get_consistency_thresholds(case.id, is_video=is_video) if is_video: - output_frames = extract_key_frames_from_video(content) + output_frames = pop_realtime_key_frames(case.id) + if output_frames is None: + output_frames = extract_key_frames_from_video(content) else: output_frames = [image_bytes_to_numpy(content)] @@ -727,10 +734,12 @@ Pinned revision used by this check: {SGL_TEST_FILES_CI_DATA_REVISION} is_video = case.server_args.modality == "video" if is_video: - # Extract key frames from video - frames = extract_key_frames_from_video( - content, num_frames=case.sampling_params.num_frames - ) + # realtime consistency uses websocket raw frames to avoid lossy mp4 drift + frames = pop_realtime_key_frames(case.id) + if frames is None: + frames = extract_key_frames_from_video( + content, num_frames=case.sampling_params.num_frames + ) if len(frames) != 3: logger.warning( @@ -1201,11 +1210,12 @@ Pinned revision used by this check: {SGL_TEST_FILES_CI_DATA_REVISION} ) # Single generation - output is reused for both validations + is_realtime_case = case.sampling_params.realtime_num_chunks is not None perf_record, content = self.run_and_collect( diffusion_server, case.id, generate_fn, - collect_perf=not is_gt_gen_mode, + collect_perf=not is_gt_gen_mode and not is_realtime_case, ) if is_gt_gen_mode: @@ -1223,10 +1233,23 @@ Pinned revision used by this check: {SGL_TEST_FILES_CI_DATA_REVISION} raise failures.append((name, str(exc))) - run_case_check( - "performance", - lambda: self._validate_and_record(case, perf_record), - ) + if is_realtime_case: + run_case_check( + "performance", + lambda: validate_realtime_perf_stats( + case.id, + pop_realtime_perf_stats(case.id), + case.sampling_params.realtime_perf_thresholds, + ignore_initial_chunks=( + case.sampling_params.realtime_perf_ignore_initial_chunks + ), + ), + ) + else: + run_case_check( + "performance", + lambda: self._validate_and_record(case, perf_record), + ) if case.server_args.custom_validator == "mesh": from sglang.multimodal_gen.test.server.test_server_utils import ( diff --git a/python/sglang/multimodal_gen/test/server/test_server_utils.py b/python/sglang/multimodal_gen/test/server/test_server_utils.py index 59be5da5a..7b47ec808 100644 --- a/python/sglang/multimodal_gen/test/server/test_server_utils.py +++ b/python/sglang/multimodal_gen/test/server/test_server_utils.py @@ -4,6 +4,7 @@ Server management and performance validation for diffusion tests. from __future__ import annotations +import asyncio import base64 import os import shlex @@ -28,6 +29,15 @@ from sglang.multimodal_gen.runtime.utils.logging_utils import ( init_logger, ) from sglang.multimodal_gen.runtime.utils.perf_logger import RequestPerfRecord +from sglang.multimodal_gen.test.server.realtime_consistency import ( + build_realtime_init_payload, + collect_realtime_output, + encode_realtime_frames_to_mp4, + prepare_realtime_first_frame, + realtime_ws_url, + record_realtime_key_frames, + record_realtime_perf_stats, +) from sglang.multimodal_gen.test.server.testcase_configs import ( DiffusionSamplingParams, PerformanceSummary, @@ -1305,6 +1315,55 @@ def get_generate_fn( }, ) + def generate_realtime_video(case_id, client) -> tuple[str, bytes]: + """Realtime video generation folded back into mp4 for consistency checks.""" + if not sampling_params.prompt: + pytest.skip(f"{case_id}: no realtime prompt configured") + if sampling_params.realtime_num_chunks is None: + pytest.skip(f"{case_id}: realtime_num_chunks is not configured") + if sampling_params.realtime_num_chunks <= 0: + pytest.fail(f"{case_id}: realtime_num_chunks must be positive") + + first_frame = prepare_realtime_first_frame(sampling_params.image_path) + init_payload = build_realtime_init_payload( + model_path=model_path, + sampling_params=sampling_params, + output_size=output_size, + first_frame=first_frame, + ) + realtime_output = asyncio.run( + collect_realtime_output( + ws_url=realtime_ws_url(client), + init_payload=init_payload, + events=list(sampling_params.realtime_events), + num_chunks=sampling_params.realtime_num_chunks, + require_chunk_stats=bool(sampling_params.realtime_perf_thresholds), + ) + ) + record_realtime_perf_stats(case_id, realtime_output.chunk_stats) + record_realtime_key_frames(case_id, realtime_output.frames) + fps = int(sampling_params.fps or 24) + video_bytes = encode_realtime_frames_to_mp4(realtime_output.frames, fps=fps) + validate_openai_video(video_bytes) + + rid = f"{case_id}-realtime" + expected_filename = f"{rid}.mp4" + tmp_path = expected_filename + Path(tmp_path).write_bytes(video_bytes) + expected_width, expected_height = parse_dimensions(output_size) + validate_video_file( + tmp_path, expected_filename, expected_width, expected_height + ) + upload_file_to_slack( + case_id=case_id, + model=model_path, + prompt=sampling_params.prompt, + file_path=tmp_path, + origin_file_path=sampling_params.image_path, + ) + os.remove(tmp_path) + return (rid, video_bytes) + def generate_mesh(case_id, client) -> tuple[str, bytes]: """I2M: Image to Mesh generation using async /v1/meshes API.""" import requests as http_requests @@ -1402,7 +1461,9 @@ def get_generate_fn( if modality == "3d": fn = generate_mesh elif modality == "video": - if sampling_params.image_path and sampling_params.prompt: + if sampling_params.realtime_num_chunks is not None: + fn = generate_realtime_video + elif sampling_params.image_path and sampling_params.prompt: if getattr(sampling_params, "direct_url_test", False): fn = generate_text_url_image_to_video else: diff --git a/python/sglang/multimodal_gen/test/server/testcase_configs.py b/python/sglang/multimodal_gen/test/server/testcase_configs.py index 3a8633014..34d4600d8 100644 --- a/python/sglang/multimodal_gen/test/server/testcase_configs.py +++ b/python/sglang/multimodal_gen/test/server/testcase_configs.py @@ -26,7 +26,7 @@ import statistics from dataclasses import dataclass, field, replace from functools import lru_cache from pathlib import Path -from typing import Sequence +from typing import Any, Sequence from sglang.multimodal_gen.configs.pipeline_configs.base import ModelTaskType from sglang.multimodal_gen.registry import ( @@ -246,6 +246,15 @@ class DiffusionSamplingParams: num_outputs_per_prompt: int = 1 + # Realtime video consistency harness. When set, server tests use + # /v1/realtime_video/generate and fold streamed chunks back into mp4 bytes. + realtime_num_chunks: int | None = None + realtime_events: list[dict[str, Any]] = field(default_factory=list) + realtime_perf_thresholds: dict[str, float] = field(default_factory=dict) + realtime_perf_ignore_initial_chunks: int = 0 + # None keeps the lossless/raw transport used by GT-backed consistency checks. + realtime_output_format: str | None = None + # Additional request-level parameters (e.g. enable_teacache, enable_upscaling, …) # merged directly into the OpenAI extra_body dict. extras: dict = field(default_factory=dict) @@ -301,6 +310,48 @@ class DiffusionTestCase: ) +LINGBOT_WORLD_REALTIME_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, " + "clean horizon, cinematic but playful." + ), + image_path=( + "https://is1-ssl.mzstatic.com/image/thumb/Music/v4/b8/f9/b9/" + "b8f9b9f8-a609-bde2-0302-349436ffc508/825646291038.jpg/600x600bb.jpg" + ), + output_size="832x480", + num_frames=9, + fps=16, + realtime_num_chunks=4, + realtime_events=[ + { + "after_chunk": 0, + "kind": "camera_actions", + "payload": {"mode": "state", "transitions": [{"actions": ["w"]}]}, + }, + { + "after_chunk": 2, + "kind": "camera_actions", + "payload": {"mode": "state", "transitions": [{"actions": []}]}, + }, + ], + realtime_perf_thresholds={ + "p95_chunk_total_ms": 5000.0, + "p95_scheduler_forward_ms": 4500.0, + "p95_ws_payload_mb": 16.0, + }, + realtime_perf_ignore_initial_chunks=2, + extras={ + "seed": 42, + "num_inference_steps": 4, + "guidance_scale": 1.0, + "realtime_causal_sink_size": 9, + "realtime_causal_kv_cache_num_frames": 18, + }, +) + + def sample_step_indices( step_map: dict[int, float], fractions: Sequence[float] ) -> list[int]: diff --git a/python/sglang/multimodal_gen/test/test_utils.py b/python/sglang/multimodal_gen/test/test_utils.py index 51ead5afd..1e9636158 100644 --- a/python/sglang/multimodal_gen/test/test_utils.py +++ b/python/sglang/multimodal_gen/test/test_utils.py @@ -33,7 +33,7 @@ if TYPE_CHECKING: logger = init_logger(__name__) -SGL_TEST_FILES_CI_DATA_REVISION = "20874fb9018d082c613a18ba92ab4f32479d3a32" +SGL_TEST_FILES_CI_DATA_REVISION = "4a62a21f0f8bcc54c3bc6de5dbf25515bcb92b7e" SGL_TEST_FILES_CONSISTENCY_GT_ROOT = ( "https://raw.githubusercontent.com/" f"sgl-project/ci-data/{SGL_TEST_FILES_CI_DATA_REVISION}/" diff --git a/python/sglang/multimodal_gen/test/unit/realtime/test_causal_denoising.py b/python/sglang/multimodal_gen/test/unit/realtime/test_causal_denoising.py new file mode 100644 index 000000000..ca0e0d678 --- /dev/null +++ b/python/sglang/multimodal_gen/test/unit/realtime/test_causal_denoising.py @@ -0,0 +1,503 @@ +# SPDX-License-Identifier: Apache-2.0 + +from types import MethodType, SimpleNamespace + +import torch + +from sglang.multimodal_gen.runtime.layers.kvcache.causal_attention_cache import ( + CausalSelfAttentionKVCache, + CrossAttentionKVCache, +) +from sglang.multimodal_gen.runtime.pipelines_core.stages.causal_denoising import ( + CausalDMDDenoisingStage, +) + + +class _Progress: + def __init__(self): + self.count = 0 + + def update(self): + self.count += 1 + + +def test_causal_dmd_chunk_loop_uses_model_input_builder(): + stage = CausalDMDDenoisingStage.__new__(CausalDMDDenoisingStage) + predict_calls = [] + add_noise_calls = [] + + def fake_predict(self, *args, **kwargs): + del self, args + latent_model_input = kwargs["latent_model_input"] + predict_calls.append( + ( + latent_model_input.shape, + latent_model_input.dtype, + int(kwargs["current_timestep"]), + ) + ) + x0_btchw = latent_model_input[:, :2].permute(0, 2, 1, 3, 4).float() + return x0_btchw + int(kwargs["current_timestep"]), kwargs["current_timestep"] + + def fake_add_noise(self, *args, **kwargs): + del self, args + add_noise_calls.append(int(kwargs["next_timestep"].item())) + return kwargs["x0_btchw"] + 10 + + stage._predict_x0_btchw = MethodType(fake_predict, stage) + stage._add_noise_for_next_timestep = MethodType(fake_add_noise, stage) + + chunk_latents = torch.zeros(1, 2, 2, 1, 1) + condition = torch.ones(1, 1, 2, 1, 1) + prepare_call_count = 0 + + def prepare_model_input(current_latents): + nonlocal prepare_call_count + prepare_call_count += 1 + return torch.cat([current_latents, condition], dim=1) + + progress = _Progress() + result, attn_metadata = stage._denoise_causal_dmd_chunk( + SimpleNamespace(generator=None), + SimpleNamespace(), + chunk_latents=chunk_latents, + scheduler=SimpleNamespace(), + timesteps=torch.tensor([7, 3]), + prompt_embeds=None, + kv_cache=[], + crossattn_cache=[], + current_start_tokens=0, + start_frame=0, + image_kwargs={}, + pos_cond_kwargs={}, + target_dtype=torch.float16, + autocast_enabled=False, + device=torch.device("cpu"), + attn_raw_latent_shape=(2, 1, 1), + prepare_model_input=prepare_model_input, + progress_bar=progress, + ) + + assert prepare_call_count == 2 + assert predict_calls == [ + (torch.Size([1, 3, 2, 1, 1]), torch.float16, 0), + (torch.Size([1, 3, 2, 1, 1]), torch.float16, 1), + ] + assert add_noise_calls == [3] + assert progress.count == 2 + assert attn_metadata == 1 + assert torch.equal(result, torch.full_like(result, 11)) + + +def test_causal_dmd_forward_context_uses_prepare_hooks(monkeypatch): + from sglang.multimodal_gen.runtime.pipelines_core.stages import causal_denoising + + stage = CausalDMDDenoisingStage.__new__(CausalDMDDenoisingStage) + stage.attn_backend = SimpleNamespace(get_enum=lambda: None) + monkeypatch.setattr( + causal_denoising, "get_local_torch_device", lambda: torch.device("cpu") + ) + latents = torch.zeros(2, 3, 4, 5, 6) + seen = {} + + def fake_prepare_frame_seq_length(self, h, w): + self.num_token_per_frame = h * w + seen["frame_shape"] = (h, w) + return self.num_token_per_frame + + stage._target_dtype = MethodType(lambda self: torch.float16, stage) + stage._autocast_enabled = MethodType(lambda self, dtype, server_args: False, stage) + stage._get_causal_dmd_scheduler = MethodType( + lambda self, batch, server_args: "scheduler", stage + ) + stage._get_causal_dmd_latents = MethodType(lambda self, batch: latents, stage) + stage._prepare_frame_seq_length = MethodType(fake_prepare_frame_seq_length, stage) + stage._prepare_causal_dmd_timesteps = MethodType( + lambda self, batch, server_args, scheduler, device: torch.tensor( + [9], device=device + ), + stage, + ) + stage._prepare_causal_dmd_image_kwargs = MethodType( + lambda self, batch, server_args, target_dtype: {"image": target_dtype}, + stage, + ) + stage._prepare_causal_dmd_pos_cond_kwargs = MethodType( + lambda self, batch, server_args, target_dtype: {"pos": target_dtype}, + stage, + ) + stage._prepare_causal_dmd_prompt_embeds = MethodType( + lambda self, batch, server_args, target_dtype: ["prompt"], + stage, + ) + + ctx = stage._prepare_causal_dmd_forward_context( + SimpleNamespace(), + SimpleNamespace(), + ) + + assert ctx.target_dtype == torch.float16 + assert ctx.autocast_enabled is False + assert ctx.device == torch.device("cpu") + assert ctx.scheduler == "scheduler" + assert ctx.timesteps.tolist() == [9] + assert ctx.latents is latents + assert ctx.prompt_embeds == ["prompt"] + assert ctx.image_kwargs == {"image": torch.float16} + assert ctx.pos_cond_kwargs == {"pos": torch.float16} + assert (ctx.batch_size, ctx.channels, ctx.num_frames, ctx.height, ctx.width) == ( + 2, + 3, + 4, + 5, + 6, + ) + assert seen["frame_shape"] == (5, 6) + + +def test_causal_dmd_block_updates_context_after_denoising(): + stage = CausalDMDDenoisingStage.__new__(CausalDMDDenoisingStage) + seen = {} + + def fake_denoise(self, *args, **kwargs): + del self, args + seen["model_input"] = kwargs["prepare_model_input"]( + kwargs["chunk_latents"] + ).clone() + seen["denoise_pos"] = ( + kwargs["current_start_tokens"], + kwargs["start_frame"], + ) + return kwargs["chunk_latents"] + 2, "metadata" + + def fake_update(self, *args, **kwargs): + del self, args + seen["context_input"] = kwargs["context_input"].clone() + seen["update_pos"] = ( + kwargs["current_start_tokens"], + kwargs["start_frame"], + kwargs["attn_metadata"], + ) + + stage._denoise_causal_dmd_chunk = MethodType(fake_denoise, stage) + stage._update_causal_context_cache = MethodType(fake_update, stage) + + chunk_latents = torch.ones(1, 2, 1, 1, 1) + condition = torch.full((1, 1, 1, 1, 1), 4.0) + result = stage._denoise_and_update_causal_block( + SimpleNamespace(), + SimpleNamespace(), + chunk_latents=chunk_latents, + scheduler=SimpleNamespace(), + timesteps=torch.tensor([1]), + prompt_embeds=None, + kv_cache=[], + crossattn_cache=[], + current_start_tokens=8, + start_frame=2, + image_kwargs={}, + pos_cond_kwargs={}, + target_dtype=torch.float16, + autocast_enabled=False, + device=torch.device("cpu"), + attn_raw_latent_shape=(1, 1, 1), + prepare_model_input=lambda x: torch.cat([x, condition], dim=1), + prepare_context_input=lambda x: torch.cat([x, condition], dim=1), + ) + + assert torch.equal(result, chunk_latents + 2) + assert seen["denoise_pos"] == (8, 2) + assert seen["update_pos"] == (8, 2, "metadata") + assert torch.equal( + seen["model_input"], torch.cat([chunk_latents, condition], dim=1) + ) + assert torch.equal( + seen["context_input"], torch.cat([chunk_latents + 2, condition], dim=1) + ) + + +def test_causal_context_warmup_uses_context_cache_update_path(): + stage = CausalDMDDenoisingStage.__new__(CausalDMDDenoisingStage) + stage.num_token_per_frame = 4 + seen = {} + + def fake_update(self, *args, **kwargs): + del self, args + seen.update(kwargs) + + stage._update_causal_context_cache = MethodType(fake_update, stage) + context_input = torch.randn(1, 2, 3, 4, 4) + stage._warm_up_causal_context_cache( + SimpleNamespace(), + SimpleNamespace(), + context_input=context_input, + prompt_embeds="prompt", + kv_cache=["kv"], + crossattn_cache=["cross"], + current_start_frame=3, + image_kwargs={}, + pos_cond_kwargs={}, + target_dtype=torch.float16, + autocast_enabled=False, + ) + + assert seen["context_input"] is context_input + assert seen["current_start_tokens"] == 12 + assert seen["start_frame"] == 3 + assert seen["attn_metadata"] is None + + +def test_causal_dmd_timestep_expansion_preserves_fractional_dtype(): + timestep = torch.tensor(996.3375, dtype=torch.float32) + + expanded = CausalDMDDenoisingStage._expand_timestep( + timestep, + batch_size=3, + device=torch.device("cpu"), + ) + + assert expanded.dtype == torch.float32 + assert torch.equal(expanded, torch.tensor([996.3375, 996.3375, 996.3375])) + + +def test_causal_cache_helpers_reset_and_forward_model_specific_kwargs(): + stage = CausalDMDDenoisingStage.__new__(CausalDMDDenoisingStage) + stage.num_transformer_blocks = 2 + kv_cache = [ + CausalSelfAttentionKVCache( + k=torch.empty(1, 1, 1, 1), + v=torch.empty(1, 1, 1, 1), + global_end_index=torch.tensor([7]), + local_end_index=torch.tensor([3]), + global_end_index_int=7, + local_end_index_int=3, + ), + CausalSelfAttentionKVCache( + k=torch.empty(1, 1, 1, 1), + v=torch.empty(1, 1, 1, 1), + global_end_index=torch.tensor([5]), + local_end_index=torch.tensor([2]), + ), + ] + crossattn_cache = [ + CrossAttentionKVCache( + k=torch.empty(1, 1, 1, 1), + v=torch.empty(1, 1, 1, 1), + is_init=True, + ), + CrossAttentionKVCache( + k=torch.empty(1, 1, 1, 1), + v=torch.empty(1, 1, 1, 1), + is_init=True, + ), + ] + stage._reset_causal_caches( + kv_cache=kv_cache, + crossattn_cache=crossattn_cache, + ) + + assert [block.is_init for block in crossattn_cache] == [False, False] + assert [int(block.global_end_index.item()) for block in kv_cache] == [0, 0] + assert [int(block.local_end_index.item()) for block in kv_cache] == [0, 0] + assert kv_cache[0].global_end_index_int == 0 + assert kv_cache[0].local_end_index_int == 0 + + calls = [] + + def fake_initialize_kv_cache(self, batch_size, dtype, device, **kwargs): + calls.append(("kv", batch_size, dtype, device, kwargs)) + self.causal_kv_cache = ["kv-cache"] + + def fake_initialize_crossattn_cache(self, batch_size, max_text_len, dtype, device): + calls.append(("cross", batch_size, max_text_len, dtype, device)) + self.crossattn_cache = ["cross-cache"] + + stage._initialize_kv_cache = MethodType(fake_initialize_kv_cache, stage) + stage._initialize_crossattn_cache = MethodType( + fake_initialize_crossattn_cache, stage + ) + + kv_cache, crossattn_cache = stage._initialize_causal_caches( + batch_size=2, + max_text_len=128, + dtype=torch.bfloat16, + device=torch.device("cpu"), + kv_cache_kwargs={"sequence_shard_enabled": True}, + ) + + assert kv_cache == ["kv-cache"] + assert crossattn_cache == ["cross-cache"] + assert calls == [ + ( + "kv", + 2, + torch.bfloat16, + torch.device("cpu"), + {"sequence_shard_enabled": True}, + ), + ("cross", 2, 128, torch.bfloat16, torch.device("cpu")), + ] + + +def test_causal_kv_cache_block_resets_indices_in_place(): + stage = CausalDMDDenoisingStage.__new__(CausalDMDDenoisingStage) + k = torch.ones(1, 4, 2, 3) + v = torch.ones(1, 4, 2, 3) + global_end_index = torch.tensor([7]) + local_end_index = torch.tensor([4]) + cache = CausalSelfAttentionKVCache( + k=k, + v=v, + global_end_index=global_end_index, + local_end_index=local_end_index, + global_end_index_int=7, + local_end_index_int=4, + ) + + assert cache.k is k + assert cache.global_end_index_int == 7 + detached_k = k.detach() + cache.k = detached_k + assert cache.k is detached_k + + stage._reset_kv_cache([cache]) + + assert cache.global_end_index is global_end_index + assert cache.local_end_index is local_end_index + assert int(cache.global_end_index.item()) == 0 + assert int(cache.local_end_index.item()) == 0 + assert cache.global_end_index_int == 0 + assert cache.local_end_index_int == 0 + + +def test_causal_kv_cache_allocation_sets_shapes_and_optional_int_indices(): + stage = CausalDMDDenoisingStage.__new__(CausalDMDDenoisingStage) + stage.num_transformer_blocks = 2 + + cache = stage._allocate_causal_kv_cache( + batch_size=3, + kv_cache_size=5, + num_attention_heads=7, + attention_head_dim=11, + dtype=torch.float16, + device=torch.device("cpu"), + use_int_indices=True, + sink_tokens=13, + attention_window_size=3, + ) + + assert len(cache) == 2 + assert cache[0].k.shape == (3, 5, 7, 11) + assert cache[0].v.shape == (3, 5, 7, 11) + assert cache[0].global_end_index.shape == (1,) + assert cache[0].local_end_index.shape == (1,) + assert cache[0].global_end_index_int == 0 + assert cache[0].local_end_index_int == 0 + assert cache[0].cache_size == 5 + assert cache[0].sink_tokens == 13 + assert cache[0].attention_window_size == 3 + assert cache[0].allow_growth is False + + +def test_causal_kv_cache_update_handles_append_roll_and_recompute(): + cache = CausalSelfAttentionKVCache( + k=torch.zeros(1, 4, 1, 1), + v=torch.zeros(1, 4, 1, 1), + global_end_index=torch.zeros(1, dtype=torch.long), + local_end_index=torch.zeros(1, dtype=torch.long), + global_end_index_int=0, + local_end_index_int=0, + ) + + first_view = cache.update_and_get_attention_kv( + key=torch.tensor([[[[1.0]], [[2.0]], [[3.0]]]]), + value=torch.tensor([[[[10.0]], [[20.0]], [[30.0]]]]), + current_chunk_start=0, + ) + + assert first_view.visible_global_end == 3 + assert first_view.visible_local_end == 3 + assert cache.global_end_index_int == 3 + assert cache.local_end_index_int == 3 + assert first_view.k.flatten().tolist() == [1.0, 2.0, 3.0] + + rolled_view = cache.update_and_get_attention_kv( + key=torch.tensor([[[[4.0]], [[5.0]], [[6.0]]]]), + value=torch.tensor([[[[40.0]], [[50.0]], [[60.0]]]]), + current_chunk_start=3, + ) + + assert rolled_view.visible_global_end == 6 + assert rolled_view.visible_local_end == 4 + assert cache.k.flatten().tolist() == [3.0, 4.0, 5.0, 6.0] + + cache.attention_window_size = 2 + recompute_view = cache.update_and_get_attention_kv( + key=torch.tensor([[[[50.0]]]]), + value=torch.tensor([[[[500.0]]]]), + current_chunk_start=4, + ) + + assert recompute_view.local_start_index == 2 + assert recompute_view.local_end_index == 3 + assert recompute_view.visible_global_end == 6 + assert recompute_view.visible_local_end == 4 + assert recompute_view.k.flatten().tolist() == [50.0, 6.0] + + +def test_causal_kv_cache_update_grows_without_rolling_when_enabled(): + cache = CausalSelfAttentionKVCache( + k=torch.zeros(1, 2, 1, 1), + v=torch.zeros(1, 2, 1, 1), + global_end_index=torch.zeros(1, dtype=torch.long), + local_end_index=torch.zeros(1, dtype=torch.long), + global_end_index_int=0, + local_end_index_int=0, + allow_growth=True, + ) + + cache.update_and_get_attention_kv( + key=torch.tensor([[[[1.0]], [[2.0]]]]), + value=torch.tensor([[[[10.0]], [[20.0]]]]), + current_chunk_start=0, + ) + view = cache.update_and_get_attention_kv( + key=torch.tensor([[[[3.0]], [[4.0]]]]), + value=torch.tensor([[[[30.0]], [[40.0]]]]), + current_chunk_start=2, + ) + + assert cache.cache_size == 4 + assert cache.attention_window_size == 4 + assert cache.global_end_index_int == 4 + assert cache.local_end_index_int == 4 + assert view.local_start_index == 2 + assert view.local_end_index == 4 + assert view.k.flatten().tolist() == [1.0, 2.0, 3.0, 4.0] + + +def test_crossattn_cache_block_stores_detached_tensors_and_resets(): + stage = CausalDMDDenoisingStage.__new__(CausalDMDDenoisingStage) + stage.num_transformer_blocks = 1 + k = torch.ones(1, 8, 2, 3) + v = torch.ones(1, 8, 2, 3) + cache = CrossAttentionKVCache(k=k, v=v, is_init=True) + + assert cache.k is k + assert cache.is_init is True + detached_v = v.detach() + cache.v = detached_v + assert cache.v is detached_v + new_k = torch.full_like(k, 2.0, requires_grad=True) + new_v = torch.full_like(v, 3.0, requires_grad=True) + cache.store(new_k, new_v) + assert cache.k is not new_k + assert cache.v is not new_v + assert cache.k.requires_grad is False + assert cache.v.requires_grad is False + + stage._reset_crossattn_cache([cache]) + + assert cache.k.flatten().tolist() == [2.0] * cache.k.numel() + assert cache.v.flatten().tolist() == [3.0] * cache.v.numel() + assert cache.is_init is False diff --git a/python/sglang/multimodal_gen/test/unit/realtime/test_lingbot_causal_denoising.py b/python/sglang/multimodal_gen/test/unit/realtime/test_lingbot_causal_denoising.py new file mode 100644 index 000000000..d02485192 --- /dev/null +++ b/python/sglang/multimodal_gen/test/unit/realtime/test_lingbot_causal_denoising.py @@ -0,0 +1,297 @@ +# SPDX-License-Identifier: Apache-2.0 + +from types import SimpleNamespace + +import torch + +from sglang.multimodal_gen.runtime.layers.kvcache.causal_attention_cache import ( + CrossAttentionKVCache, +) +from sglang.multimodal_gen.runtime.models.dits.lingbot_world import ( + CausalLingBotWorldTransformer3DModel, +) +from sglang.multimodal_gen.runtime.pipelines_core.stages.model_specific_stages.lingbot_world import ( + LingBotWorldCausalDMDDenoisingStage, +) + + +def test_lingbot_denoising_stage_does_not_own_realtime_cache_refs(): + stage = LingBotWorldCausalDMDDenoisingStage.__new__( + LingBotWorldCausalDMDDenoisingStage + ) + stage.causal_kv_cache = [object()] + stage.crossattn_cache = [object()] + + stage._clear_stage_causal_cache_refs() + + assert stage.causal_kv_cache is None + assert stage.crossattn_cache is None + + +def test_lingbot_realtime_attention_cache_uses_bounded_sink_window(): + stage = LingBotWorldCausalDMDDenoisingStage.__new__( + LingBotWorldCausalDMDDenoisingStage + ) + stage.local_attn_size = -1 + stage.sink_size = 9 + stage.sliding_window_num_frames = 18 + stage.num_frames_per_block = 3 + stage.num_token_per_frame = 10 + + assert stage._get_causal_sink_tokens() == 9 * 10 + assert stage._get_causal_kv_cache_size(sequence_shard_enabled=False) == 18 * 10 + assert stage._get_causal_kv_cache_size(sequence_shard_enabled=True) == 18 * 10 + + +def test_lingbot_realtime_cache_config_overrides_checkpoint_defaults(): + stage = LingBotWorldCausalDMDDenoisingStage.__new__( + LingBotWorldCausalDMDDenoisingStage + ) + stage.local_attn_size = -1 + stage.sink_size = 9 + stage.sliding_window_num_frames = 18 + stage.num_token_per_frame = 10 + + server_args = SimpleNamespace( + pipeline_config=SimpleNamespace( + realtime_causal_sink_size=3, + realtime_causal_kv_cache_num_frames=45, + ) + ) + stage._apply_causal_cache_overrides(SimpleNamespace(), server_args) + + assert stage._get_causal_sink_tokens() == 3 * 10 + assert stage._get_causal_kv_cache_size(sequence_shard_enabled=False) == 45 * 10 + + +def test_lingbot_realtime_cache_config_uses_request_overrides(): + stage = LingBotWorldCausalDMDDenoisingStage.__new__( + LingBotWorldCausalDMDDenoisingStage + ) + stage.local_attn_size = -1 + stage.sink_size = 9 + stage.sliding_window_num_frames = 18 + stage.num_token_per_frame = 10 + + batch = SimpleNamespace( + realtime_causal_sink_size=4, + realtime_causal_kv_cache_num_frames=12, + ) + server_args = SimpleNamespace( + pipeline_config=SimpleNamespace( + realtime_causal_sink_size=3, + realtime_causal_kv_cache_num_frames=45, + ) + ) + stage._apply_causal_cache_overrides(batch, server_args) + + assert stage._get_causal_sink_tokens() == 4 * 10 + assert stage._get_causal_kv_cache_size(sequence_shard_enabled=False) == 12 * 10 + + +def test_lingbot_realtime_attention_cache_rolls_with_sink_window(): + stage = LingBotWorldCausalDMDDenoisingStage.__new__( + LingBotWorldCausalDMDDenoisingStage + ) + stage.num_transformer_blocks = 1 + stage.local_attn_size = -1 + stage.sink_size = 2 + stage.num_token_per_frame = 1 + stage.num_frames_per_block = 3 + stage.sliding_window_num_frames = 6 + stage.transformer = SimpleNamespace( + num_attention_heads=1, + attention_head_dim=1, + config=SimpleNamespace( + arch_config=SimpleNamespace( + sink_size=2, + ) + ), + ) + + stage._initialize_kv_cache( + batch_size=1, + dtype=torch.float32, + device=torch.device("cpu"), + ) + assert stage.causal_kv_cache is not None + cache = stage.causal_kv_cache[0] + assert cache.allow_growth is False + assert cache.cache_size == 6 + assert cache.sink_tokens == 2 + + cache.update_and_get_attention_kv( + key=torch.ones(1, 3, 1, 1), + value=torch.ones(1, 3, 1, 1), + current_chunk_start=0, + ) + + assert cache.local_end_index_int == 3 + assert cache.global_end_index_int == 3 + + cache.update_and_get_attention_kv( + key=torch.full((1, 3, 1, 1), 2.0), + value=torch.full((1, 3, 1, 1), 2.0), + current_chunk_start=3, + ) + third_view = cache.update_and_get_attention_kv( + key=torch.full((1, 3, 1, 1), 3.0), + value=torch.full((1, 3, 1, 1), 3.0), + current_chunk_start=6, + ) + + assert cache.cache_size == 6 + assert cache.local_end_index_int == 6 + assert cache.global_end_index_int == 9 + assert torch.equal(cache.k[:, :2], torch.ones(1, 2, 1, 1)) + assert torch.equal(cache.k[:, 2:3], torch.full((1, 1, 1, 1), 2.0)) + assert torch.equal(cache.k[:, 3:6], torch.full((1, 3, 1, 1), 3.0)) + assert third_view.k.flatten().tolist() == [ + 1.0, + 1.0, + 2.0, + 3.0, + 3.0, + 3.0, + ] + + fourth_view = cache.update_and_get_attention_kv( + key=torch.full((1, 3, 1, 1), 4.0), + value=torch.full((1, 3, 1, 1), 4.0), + current_chunk_start=9, + ) + + assert cache.local_end_index_int == 6 + assert cache.global_end_index_int == 12 + assert torch.equal(cache.k[:, :2], torch.ones(1, 2, 1, 1)) + assert torch.equal(cache.k[:, 2:3], torch.full((1, 1, 1, 1), 3.0)) + assert torch.equal(cache.k[:, 3:6], torch.full((1, 3, 1, 1), 4.0)) + assert fourth_view.k.flatten().tolist() == [ + 1.0, + 1.0, + 3.0, + 4.0, + 4.0, + 4.0, + ] + + +def test_lingbot_i2v_model_input_writer_reuses_buffer(): + latents = torch.ones(1, 16, 3, 2, 2) + condition = torch.full((1, 20, 3, 2, 2), 2.0) + + write = LingBotWorldCausalDMDDenoisingStage._build_i2v_model_input_writer( + latents=latents, + condition=condition, + target_dtype=torch.float32, + device=latents.device, + ) + first = write(latents) + first_ptr = first.data_ptr() + second = write(latents + 3.0) + + assert first_ptr == second.data_ptr() + assert second.shape == (1, 36, 3, 2, 2) + assert torch.equal(second[:, :16], latents + 3.0) + assert torch.equal(second[:, 16:], condition) + + +def test_lingbot_condition_embedding_skips_text_when_crossattn_cache_ready(): + class _ConditionEmbedder: + def __init__(self): + self.full_calls = 0 + self.time_calls = 0 + + def time_embedder(self, timestep): + self.time_calls += 1 + return timestep.float().unsqueeze(-1) + + def time_modulation(self, temb): + return temb + 1.0 + + def __call__( + self, timestep, encoder_hidden_states, encoder_hidden_states_image + ): + self.full_calls += 1 + return timestep.float(), timestep.float(), encoder_hidden_states, None + + model = CausalLingBotWorldTransformer3DModel.__new__( + CausalLingBotWorldTransformer3DModel + ) + model.condition_embedder = _ConditionEmbedder() + crossattn_cache = [ + CrossAttentionKVCache( + k=torch.empty(1, 1, 1, 1), + v=torch.empty(1, 1, 1, 1), + is_init=True, + ) + ] + + temb, timestep_proj, _, image_states = model._prepare_condition_embeddings( + timestep=torch.tensor([[7]]), + encoder_hidden_states=torch.ones(1, 2, 3), + encoder_hidden_states_image=torch.ones(1, 2, 3), + crossattn_cache=crossattn_cache, + ) + + assert model.condition_embedder.time_calls == 1 + assert model.condition_embedder.full_calls == 0 + assert torch.equal(temb, torch.tensor([[7.0]])) + assert torch.equal(timestep_proj, torch.tensor([[8.0]])) + assert image_states is None + + +def test_lingbot_context_cache_update_skips_unused_projection(monkeypatch): + from sglang.multimodal_gen.runtime.pipelines_core.stages.model_specific_stages.lingbot_world import ( + lingbot_world_causal_denoising as lingbot_denoising, + ) + + stage = LingBotWorldCausalDMDDenoisingStage.__new__( + LingBotWorldCausalDMDDenoisingStage + ) + calls = [] + + class _Transformer: + def __call__(self, latent_input, prompt_embeds, timestep, **kwargs): + calls.append((latent_input, prompt_embeds, timestep, kwargs)) + return latent_input + + stage.transformer = _Transformer() + monkeypatch.setattr( + lingbot_denoising, + "current_platform", + SimpleNamespace(device_type="cpu"), + ) + + context_input = torch.ones(2, 3, 4, 5, 6) + batch = SimpleNamespace() + stage._update_causal_context_cache( + batch, + SimpleNamespace(pipeline_config=SimpleNamespace(context_noise=7)), + context_input=context_input, + prompt_embeds="prompt", + kv_cache="kv", + crossattn_cache="cross", + current_start_tokens=12, + start_frame=4, + image_kwargs={"encoder_hidden_states_image": "image"}, + pos_cond_kwargs={"c2ws_plucker_emb": "pose"}, + attn_metadata="metadata", + target_dtype=torch.float32, + autocast_enabled=False, + ) + + assert len(calls) == 1 + latent_input, prompt_embeds, timestep, kwargs = calls[0] + assert latent_input is context_input + assert prompt_embeds == "prompt" + assert timestep.shape == (2, 1) + assert timestep.dtype == torch.long + assert timestep.tolist() == [[7], [7]] + assert kwargs["kv_cache"] == "kv" + assert kwargs["crossattn_cache"] == "cross" + assert kwargs["current_start"] == 12 + assert kwargs["start_frame"] == 4 + assert kwargs["encoder_hidden_states_image"] == "image" + assert kwargs["c2ws_plucker_emb"] == "pose" + assert kwargs["skip_final_projection"] is True diff --git a/python/sglang/multimodal_gen/test/unit/realtime/test_output_materialization.py b/python/sglang/multimodal_gen/test/unit/realtime/test_output_materialization.py new file mode 100644 index 000000000..37f9fc7e4 --- /dev/null +++ b/python/sglang/multimodal_gen/test/unit/realtime/test_output_materialization.py @@ -0,0 +1,116 @@ +# SPDX-License-Identifier: Apache-2.0 + +import numpy as np +import torch + +from sglang.multimodal_gen.configs.sample.sampling_params import DataType +from sglang.multimodal_gen.runtime.entrypoints.utils import ( + materialize_output_sample, + save_outputs, +) +from sglang.multimodal_gen.runtime.managers.gpu_worker import GPUWorker +from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import OutputBatch +from sglang.multimodal_gen.runtime.utils.realtime_video import ( + build_raw_rgb_frame_batches, +) + + +def test_materialize_output_sample_converts_tensor_to_uint8_frames(): + sample = torch.zeros(3, 1, 2, 2) + sample[0] = 1.0 + sample[1] = 0.5 + + materialized = materialize_output_sample(sample, DataType.VIDEO, fps=24) + + assert materialized.fps == 24 + assert materialized.audio is None + assert len(materialized.frames) == 1 + frame = materialized.frames[0] + assert frame.shape == (2, 2, 3) + assert frame.dtype == np.uint8 + assert np.all(frame[..., 0] == 255) + assert np.all(frame[..., 1] == 127) + assert np.all(frame[..., 2] == 0) + + +def test_save_outputs_can_materialize_without_saving(tmp_path): + sample = np.full((2, 2, 3), 0.25, dtype=np.float32) + output_path = tmp_path / "image.png" + samples_out = [] + frames_out = [] + + paths = save_outputs( + [sample], + DataType.IMAGE, + fps=1, + save_output=False, + build_output_path=lambda _idx: str(output_path), + samples_out=samples_out, + frames_out=frames_out, + ) + + assert paths == [str(output_path)] + assert not output_path.exists() + assert samples_out[0] is sample + assert len(frames_out) == 1 + assert len(frames_out[0]) == 1 + assert frames_out[0][0].dtype == np.uint8 + assert np.all(frames_out[0][0] == 63) + + +def test_file_path_transport_clears_in_memory_outputs(): + worker = GPUWorker.__new__(GPUWorker) + worker.rank = 0 + output_batch = OutputBatch( + output=[object()], + audio=torch.zeros(1), + audio_sample_rate=16000, + ) + + def save_output_paths(batch): + batch.output_file_paths = ["/tmp/output.png"] + + worker._materialize_file_path_transport(output_batch, save_output_paths) + + assert output_batch.output_file_paths == ["/tmp/output.png"] + assert output_batch.output is None + assert output_batch.audio is None + assert output_batch.audio_sample_rate is None + + +def test_raw_rgb_frame_batches_convert_batched_video_tensor_to_thwc_bytes(): + output = torch.zeros(1, 3, 2, 2, 2) + output[0, 0] = 1.0 + output[0, 1] = 0.5 + req = type( + "Req", + (), + { + "enable_frame_interpolation": False, + "enable_upscaling": False, + "request_id": "req", + "block_idx": 0, + }, + )() + output_batch = OutputBatch(audio_sample_rate=None) + + frame_batches, metadata = build_raw_rgb_frame_batches( + output, + req, + output_batch, + post_process_sample_fn=lambda *args, **kwargs: None, + ) + + assert metadata == { + "format": "rgb24", + "width": 2, + "height": 2, + "channels": 3, + "bytes_per_frame": 12, + } + assert len(frame_batches) == 1 + assert len(frame_batches[0]) == 2 + first = np.frombuffer(frame_batches[0][0], dtype=np.uint8).reshape(2, 2, 3) + assert np.all(first[..., 0] == 255) + assert np.all(first[..., 1] == 127) + assert np.all(first[..., 2] == 0) diff --git a/python/sglang/multimodal_gen/test/unit/realtime/test_realtime_condition_events.py b/python/sglang/multimodal_gen/test/unit/realtime/test_realtime_condition_events.py new file mode 100644 index 000000000..2c4091020 --- /dev/null +++ b/python/sglang/multimodal_gen/test/unit/realtime/test_realtime_condition_events.py @@ -0,0 +1,250 @@ +# SPDX-License-Identifier: Apache-2.0 + +from sglang.multimodal_gen.runtime.realtime.condition_events import ( + ConditionEvent, + ConditionEventQueue, + ConditionSamplingParams, + ControlSignal, + ControlStateSamplingQueue, + ControlStateTransition, +) + + +def test_condition_event_queue_samples_chunk_and_repeats_last_item(): + queue = ConditionEventQueue() + queue.push(ConditionEvent(kind="camera_actions", payload=[["w"], ["d"]])) + + chunk = queue.sample_chunk( + "camera_actions", + ConditionSamplingParams(chunk_size=4, default_item=[]), + ) + + assert chunk == [["w"], ["d"], ["d"], ["d"]] + + +def test_condition_event_contains_multiple_same_kind_control_signals(): + event = ConditionEvent( + kind="camera_actions", + payload=[ + ControlSignal(kind="camera_actions", payload=["w"]), + ControlSignal(kind="camera_actions", payload=["d"]), + ], + ) + signals = list(event.iter_signals()) + + assert [signal.kind for signal in signals] == [ + "camera_actions", + "camera_actions", + ] + assert [signal.payload for signal in signals] == [["w"], ["d"]] + + +def test_condition_event_queue_samples_control_signal_payloads(): + queue = ConditionEventQueue() + queue.push( + ConditionEvent( + kind="camera_actions", + payload=[ + ControlSignal(kind="camera_actions", payload=["w"]), + ControlSignal(kind="camera_actions", payload=["a"]), + ControlSignal(kind="camera_actions", payload=["s"]), + ], + ) + ) + + first = queue.sample_chunk( + "camera_actions", + ConditionSamplingParams(chunk_size=2, default_item=[]), + ) + second = queue.sample_chunk( + "camera_actions", + ConditionSamplingParams(chunk_size=2, default_item=[]), + ) + + assert first == [["w"], ["a"]] + assert second == [["s"], ["s"]] + + +def test_condition_event_queue_preserves_event_remainder_across_chunks(): + queue = ConditionEventQueue() + queue.push(ConditionEvent(kind="camera_actions", payload=[["w"], ["a"], ["s"]])) + + first = queue.sample_chunk( + "camera_actions", + ConditionSamplingParams(chunk_size=2, default_item=[]), + ) + second = queue.sample_chunk( + "camera_actions", + ConditionSamplingParams(chunk_size=2, default_item=[]), + ) + + assert first == [["w"], ["a"]] + assert second == [["s"], ["s"]] + + +def test_condition_event_queue_does_not_persist_last_signal_across_empty_chunks(): + queue = ConditionEventQueue() + queue.push(ConditionEvent(kind="camera_actions", payload=[["w"]])) + + first = queue.sample_chunk( + "camera_actions", + ConditionSamplingParams(chunk_size=2, default_item=[]), + ) + second = queue.sample_chunk( + "camera_actions", + ConditionSamplingParams(chunk_size=2, default_item=[]), + ) + + assert first == [["w"], ["w"]] + assert second == [[], []] + + +def test_condition_event_queue_can_repeat_last_signal_across_empty_chunks(): + queue = ConditionEventQueue() + queue.push(ConditionEvent(kind="camera_actions", payload=[["w"]])) + + first = queue.sample_chunk( + "camera_actions", + ConditionSamplingParams( + chunk_size=2, + default_item=[], + repeat_last_across_empty_chunks=True, + ), + ) + second = queue.sample_chunk( + "camera_actions", + ConditionSamplingParams( + chunk_size=2, + default_item=[], + repeat_last_across_empty_chunks=True, + ), + ) + queue.push(ConditionEvent(kind="camera_actions", payload=[[]])) + third = queue.sample_chunk( + "camera_actions", + ConditionSamplingParams( + chunk_size=2, + default_item=[], + repeat_last_across_empty_chunks=True, + ), + ) + + assert first == [["w"], ["w"]] + assert second == [["w"], ["w"]] + assert third == [[], []] + + +def test_condition_event_queue_tracks_sampled_signal_seq_id(): + queue = ConditionEventQueue() + queue.push( + ConditionEvent( + kind="camera_actions", + payload=[ + ControlSignal(kind="camera_actions", payload=["w"], seq_id=7), + ControlSignal(kind="camera_actions", payload=[], seq_id=8), + ], + ) + ) + + first = queue.sample_chunk( + "camera_actions", + ConditionSamplingParams( + chunk_size=1, + default_item=[], + repeat_last_across_empty_chunks=True, + ), + ) + first_seq_id = queue.last_sampled_seq_id("camera_actions") + second = queue.sample_chunk( + "camera_actions", + ConditionSamplingParams( + chunk_size=1, + default_item=[], + repeat_last_across_empty_chunks=True, + ), + ) + second_seq_id = queue.last_sampled_seq_id("camera_actions") + + assert first == [["w"]] + assert first_seq_id == 7 + assert second == [[]] + assert second_seq_id == 8 + + +def test_condition_event_queue_replace_clears_pending_signals(): + queue = ConditionEventQueue() + queue.push( + ConditionEvent(kind="camera_actions", payload=[["w"], ["w"], ["w"], ["w"]]) + ) + + first = queue.sample_chunk( + "camera_actions", + ConditionSamplingParams(chunk_size=2, default_item=[]), + ) + queue.replace(ConditionEvent(kind="camera_actions", payload=[["d"]])) + second = queue.sample_chunk( + "camera_actions", + ConditionSamplingParams(chunk_size=3, default_item=[]), + ) + third = queue.sample_chunk( + "camera_actions", + ConditionSamplingParams(chunk_size=3, default_item=[]), + ) + + assert first == [["w"], ["w"]] + assert second == [["d"], ["d"], ["d"]] + assert third == [[], [], []] + + +def test_condition_event_queue_returns_none_without_default_item(): + queue = ConditionEventQueue() + + chunk = queue.sample_chunk("audio", ConditionSamplingParams(chunk_size=2)) + + assert chunk is None + + +def test_condition_event_queue_empty_event_switches_to_default_item(): + queue = ConditionEventQueue() + queue.push(ConditionEvent(kind="camera_actions", payload=[])) + + chunk = queue.sample_chunk( + "camera_actions", + ConditionSamplingParams(chunk_size=3, default_item=[]), + ) + + assert chunk == [[], [], []] + + +def test_control_state_sampling_queue_preserves_short_pulse(): + queue = ControlStateSamplingQueue(default_item=[], min_pulse_items=1) + queue.push(ControlStateTransition(payload=["w"], seq_id=7)) + queue.push(ControlStateTransition(payload=[], seq_id=8)) + + chunk = queue.sample_chunk(3) + + assert chunk == [["w"], [], []] + assert queue.latest_sampled_seq_id() == 8 + assert queue.sample_chunk(3) == [[], [], []] + + +def test_control_state_sampling_queue_holds_current_state_without_backlog(): + queue = ControlStateSamplingQueue(default_item=[], min_pulse_items=1) + queue.push(ControlStateTransition(payload=["w"], seq_id=7)) + + assert queue.sample_chunk(3) == [["w"], ["w"], ["w"]] + assert queue.latest_sampled_seq_id() == 7 + assert queue.sample_chunk(3) == [["w"], ["w"], ["w"]] + assert queue.latest_sampled_seq_id() == 7 + + +def test_control_state_sampling_queue_compacts_many_transitions(): + queue = ControlStateSamplingQueue(default_item=[], min_pulse_items=1) + queue.push(ControlStateTransition(payload=["w"], seq_id=7)) + queue.push(ControlStateTransition(payload=["w", "d"], seq_id=8)) + queue.push(ControlStateTransition(payload=["d"], seq_id=9)) + + chunk = queue.sample_chunk(3) + + assert chunk == [["d"], ["d"], ["d"]] + assert queue.latest_sampled_seq_id() == 9 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 new file mode 100644 index 000000000..67ecdc79e --- /dev/null +++ b/python/sglang/multimodal_gen/test/unit/realtime/test_realtime_consistency_harness.py @@ -0,0 +1,522 @@ +# SPDX-License-Identifier: Apache-2.0 + +import asyncio +import sys +from types import SimpleNamespace + +import msgspec.msgpack +import numpy as np +import pytest + +from sglang.multimodal_gen.runtime.utils.realtime_video import ( + RAW_RGB_CONTENT_TYPE, + RAW_RGB_DELTA_GZIP_CONTENT_TYPE, + RAW_RGBA_DELTA_GZIP_CONTENT_TYPE, + build_delta_gzip_raw_rgb_payload, +) +from sglang.multimodal_gen.test.server.realtime_consistency import ( + build_realtime_event_payload, + build_realtime_init_payload, + collect_realtime_output, + decode_realtime_raw_rgb_frames, + parse_realtime_chunk_stats, + pop_realtime_key_frames, + prepare_realtime_first_frame, + realtime_ws_url, + record_realtime_key_frames, + select_realtime_key_frames, + summarize_realtime_perf_stats, + validate_realtime_perf_stats, +) +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, +) + +# Request construction + + +def test_realtime_ws_url_uses_existing_openai_base_url(): + client = SimpleNamespace(base_url="http://127.0.0.1:30000/v1") + + assert realtime_ws_url(client) == "ws://127.0.0.1:30000/v1/realtime_video/generate" + + +def test_realtime_init_payload_uses_sampling_params_and_extras(): + params = DiffusionSamplingParams( + prompt="turn camera left", + seconds=1, + fps=8, + num_frames=6, + extras={"seed": 7, "num_inference_steps": 4}, + realtime_num_chunks=2, + ) + + payload = build_realtime_init_payload( + model_path="robbyant/lingbot-world-fast-diffusers", + sampling_params=params, + output_size="832x480", + first_frame="https://example.com/first.png", + ) + + assert payload == { + "type": "init", + "model": "robbyant/lingbot-world-fast-diffusers", + "prompt": "turn camera left", + "size": "832x480", + "seconds": 1, + "first_frame": "https://example.com/first.png", + "fps": 8, + "num_frames": 6, + "seed": 7, + "num_inference_steps": 4, + } + + +def test_realtime_init_payload_can_request_preview_transport(): + params = DiffusionSamplingParams( + prompt="preview transport", + realtime_num_chunks=1, + realtime_output_format="webp", + ) + + payload = build_realtime_init_payload( + model_path="robbyant/lingbot-world-fast-diffusers", + sampling_params=params, + output_size="832x480", + first_frame=None, + ) + + assert payload["realtime_output_format"] == "webp" + + +def test_realtime_first_frame_accepts_url_or_file(tmp_path): + frame_path = tmp_path / "first.png" + frame_path.write_bytes(b"png-bytes") + + assert prepare_realtime_first_frame("https://example.com/first.png") == ( + "https://example.com/first.png" + ) + assert prepare_realtime_first_frame(frame_path) == b"png-bytes" + + +def test_realtime_key_frames_are_selected_from_raw_websocket_frames(): + frames = [np.full((2, 2, 3), idx, dtype=np.uint8) for idx in range(5)] + + selected = select_realtime_key_frames(frames) + assert [int(frame[0, 0, 0]) for frame in selected] == [0, 2, 4] + + record_realtime_key_frames("unit-raw-frames", frames) + frames[2][:] = 99 + popped = pop_realtime_key_frames("unit-raw-frames") + + assert popped is not None + assert [int(frame[0, 0, 0]) for frame in popped] == [0, 2, 4] + assert pop_realtime_key_frames("unit-raw-frames") is None + + +def test_realtime_event_payload_strips_test_schedule_metadata(): + payload = build_realtime_event_payload( + { + "after_chunk": 0, + "kind": "camera_actions", + "payload": [["w"], ["d"]], + } + ) + + assert payload == { + "type": "event", + "kind": "camera_actions", + "payload": [["w"], ["d"]], + } + + +# Raw RGB frame decoding + + +def test_decode_realtime_raw_rgb_frames_splits_payload_by_header_metadata(): + first = np.arange(12, dtype=np.uint8).reshape(2, 2, 3) + second = (np.arange(12, dtype=np.uint8) + 20).reshape(2, 2, 3) + header = { + "type": "frame_batch_header", + "content_type": RAW_RGB_CONTENT_TYPE, + "num_frames": 2, + "width": 2, + "height": 2, + "channels": 3, + "bytes_per_frame": 12, + } + + frames = decode_realtime_raw_rgb_frames( + header, + first.tobytes() + second.tobytes(), + ) + + assert len(frames) == 2 + np.testing.assert_array_equal(frames[0], first) + np.testing.assert_array_equal(frames[1], second) + + +def test_decode_realtime_raw_rgb_frames_rejects_truncated_payload(): + header = { + "type": "frame_batch_header", + "content_type": RAW_RGB_CONTENT_TYPE, + "num_frames": 1, + "width": 2, + "height": 2, + "channels": 3, + "bytes_per_frame": 12, + } + + with pytest.raises(ValueError, match="payload size mismatch"): + decode_realtime_raw_rgb_frames(header, b"too-short") + + +def test_decode_realtime_delta_gzip_raw_rgb_frames_roundtrips(): + first = np.arange(12, dtype=np.uint8).reshape(2, 2, 3) + second = (np.arange(12, dtype=np.uint8) + 1).reshape(2, 2, 3) + header = { + "type": "frame_batch_header", + "content_type": RAW_RGB_DELTA_GZIP_CONTENT_TYPE, + "num_frames": 2, + "width": 2, + "height": 2, + "channels": 3, + "bytes_per_frame": 12, + } + + frames = decode_realtime_raw_rgb_frames( + header, + build_delta_gzip_raw_rgb_payload([first.tobytes(), second.tobytes()]), + ) + + assert len(frames) == 2 + np.testing.assert_array_equal(frames[0], first) + np.testing.assert_array_equal(frames[1], second) + + +def test_decode_realtime_rgba_delta_gzip_strips_alpha(): + first = np.array( + [[[1, 2, 3, 255], [4, 5, 6, 255]]], + dtype=np.uint8, + ) + second = np.array( + [[[1, 2, 4, 255], [4, 6, 6, 255]]], + dtype=np.uint8, + ) + header = { + "type": "frame_batch_header", + "content_type": RAW_RGBA_DELTA_GZIP_CONTENT_TYPE, + "num_frames": 2, + "width": 2, + "height": 1, + "channels": 4, + "bytes_per_frame": 8, + } + + frames = decode_realtime_raw_rgb_frames( + header, + build_delta_gzip_raw_rgb_payload([first.tobytes(), second.tobytes()]), + ) + + assert len(frames) == 2 + np.testing.assert_array_equal(frames[0], first[:, :, :3]) + np.testing.assert_array_equal(frames[1], second[:, :, :3]) + + +# Stream collection and realtime performance stats + + +class _FakeRealtimeWebSocket: + def __init__(self, messages): + self.messages = list(messages) + self.sent = [] + + async def __aenter__(self): + return self + + async def __aexit__(self, exc_type, exc, tb): + return False + + async def send(self, payload): + self.sent.append(msgspec.msgpack.decode(payload)) + + async def recv(self): + if not self.messages: + raise AssertionError("fake websocket received too many recv calls") + return self.messages.pop(0) + + +def _packed_realtime_frame_message(chunk_index: int, frame: np.ndarray): + header = { + "type": "frame_batch_header", + "content_type": RAW_RGB_CONTENT_TYPE, + "chunk_index": chunk_index, + "is_final_frame_batch": True, + "num_frames": 1, + "width": frame.shape[1], + "height": frame.shape[0], + "channels": frame.shape[2], + "bytes_per_frame": frame.nbytes, + } + return msgspec.msgpack.encode(header), frame.tobytes() + + +def _packed_realtime_combined_frame_message(chunk_index: int, frame: np.ndarray): + header, payload = _packed_realtime_frame_message(chunk_index, frame) + message = msgspec.msgpack.decode(header) + message["type"] = "frame_batch" + message["payload"] = payload + return msgspec.msgpack.encode(message) + + +def _packed_realtime_chunk_stats(chunk_index: int, **overrides): + payload = { + "type": "chunk_stats", + "request_id": f"req-{chunk_index}", + "chunk_index": chunk_index, + "content_type": RAW_RGB_CONTENT_TYPE, + "num_frames": 1, + "raw_bytes": 12, + "ws_payload_bytes": 128, + "request_prepare_ms": 1, + "scheduler_forward_ms": 20, + "raw_payload_build_ms": 2, + "raw_write_ms": 3, + "ws_write_ms": 4, + "chunk_total_ms": 30, + } + payload.update(overrides) + for key, value in list(payload.items()): + if key.endswith("_ms"): + payload[key] = max(0, int(value + 0.5)) + packed = msgspec.msgpack.encode(payload) + assert bytes([0xCB]) not in packed + return packed + + +def test_collect_realtime_output_skips_and_records_chunk_stats(monkeypatch): + first = np.arange(12, dtype=np.uint8).reshape(2, 2, 3) + second = first + 1 + chunk0_header, chunk0_payload = _packed_realtime_frame_message(0, first) + chunk1_header, chunk1_payload = _packed_realtime_frame_message(1, second) + websocket = _FakeRealtimeWebSocket( + [ + chunk0_header, + chunk0_payload, + _packed_realtime_chunk_stats(0, chunk_total_ms=31), + chunk1_header, + chunk1_payload, + _packed_realtime_chunk_stats(1, chunk_total_ms=32), + ] + ) + monkeypatch.setitem( + sys.modules, + "websockets", + SimpleNamespace(connect=lambda *args, **kwargs: websocket), + ) + + result = asyncio.run( + collect_realtime_output( + ws_url="ws://example.test/v1/realtime_video/generate", + init_payload={"type": "init", "prompt": "test"}, + events=[ + { + "after_chunk": 0, + "kind": "camera_actions", + "payload": [["w"]], + } + ], + num_chunks=2, + require_chunk_stats=True, + ) + ) + + assert len(result.frames) == 2 + np.testing.assert_array_equal(result.frames[0], first) + np.testing.assert_array_equal(result.frames[1], second) + assert [stat.chunk_index for stat in result.chunk_stats] == [0, 1] + assert [stat.chunk_total_ms for stat in result.chunk_stats] == [31.0, 32.0] + assert websocket.sent == [ + {"type": "init", "prompt": "test"}, + { + "type": "event", + "kind": "camera_actions", + "payload": [["w"]], + }, + ] + + +def test_collect_realtime_output_accepts_combined_frame_batch(monkeypatch): + frame = np.arange(12, dtype=np.uint8).reshape(2, 2, 3) + websocket = _FakeRealtimeWebSocket( + [ + _packed_realtime_combined_frame_message(0, frame), + ] + ) + monkeypatch.setitem( + sys.modules, + "websockets", + SimpleNamespace(connect=lambda *args, **kwargs: websocket), + ) + + result = asyncio.run( + collect_realtime_output( + ws_url="ws://example.test/v1/realtime_video/generate", + init_payload={"type": "init", "prompt": "test"}, + events=[], + num_chunks=1, + ) + ) + + assert len(result.frames) == 1 + np.testing.assert_array_equal(result.frames[0], frame) + + +def test_realtime_perf_stats_summary_and_thresholds(): + stats = [ + parse_realtime_chunk_stats( + msgspec.msgpack.decode( + _packed_realtime_chunk_stats( + 0, + scheduler_forward_ms=10.0, + raw_write_ms=2.0, + ws_write_ms=3.0, + ws_payload_bytes=1024 * 1024, + chunk_total_ms=20.0, + ) + ) + ), + parse_realtime_chunk_stats( + msgspec.msgpack.decode( + _packed_realtime_chunk_stats( + 1, + scheduler_forward_ms=30.0, + raw_write_ms=4.0, + ws_write_ms=5.0, + ws_payload_bytes=2 * 1024 * 1024, + chunk_total_ms=40.0, + ) + ) + ), + ] + + summary = summarize_realtime_perf_stats(stats) + + assert summary["num_chunks"] == 2 + assert summary["total_frames"] == 2 + assert summary["guarded_chunks"] == 2 + assert summary["avg_scheduler_forward_ms"] == 20.0 + assert summary["p95_chunk_total_ms"] == 40.0 + assert summary["avg_ws_payload_mb"] == 1.5 + validate_realtime_perf_stats( + "case", + stats, + { + "avg_scheduler_forward_ms": 25.0, + "p95_chunk_total_ms": 45.0, + "avg_ws_payload_mb": 2.0, + }, + ) + with pytest.raises(pytest.fail.Exception, match="p95_chunk_total_ms"): + validate_realtime_perf_stats( + "case", + stats, + {"p95_chunk_total_ms": 35.0}, + ) + + +def test_realtime_perf_stats_can_ignore_startup_chunks(): + stats = [ + parse_realtime_chunk_stats( + msgspec.msgpack.decode( + _packed_realtime_chunk_stats( + 0, + scheduler_forward_ms=20000.0, + chunk_total_ms=21000.0, + ) + ) + ), + parse_realtime_chunk_stats( + msgspec.msgpack.decode( + _packed_realtime_chunk_stats( + 1, + scheduler_forward_ms=7000.0, + chunk_total_ms=7200.0, + ) + ) + ), + parse_realtime_chunk_stats( + msgspec.msgpack.decode( + _packed_realtime_chunk_stats( + 2, + scheduler_forward_ms=2300.0, + chunk_total_ms=2800.0, + ) + ) + ), + ] + + summary = summarize_realtime_perf_stats(stats, ignore_initial_chunks=2) + + assert summary["num_chunks"] == 3 + assert summary["ignored_initial_chunks"] == 2 + assert summary["guarded_chunks"] == 1 + assert summary["ignored_max_chunk_total_ms"] == 21000.0 + assert summary["p95_chunk_total_ms"] == 2800.0 + validate_realtime_perf_stats( + "case", + stats, + {"p95_chunk_total_ms": 5000.0, "p95_scheduler_forward_ms": 4500.0}, + ignore_initial_chunks=2, + ) + with pytest.raises(ValueError, match="leave at least one"): + summarize_realtime_perf_stats(stats, ignore_initial_chunks=3) + + +# Generate function routing + + +def test_realtime_sampling_params_route_to_realtime_video_generator(): + params = DiffusionSamplingParams( + prompt="turn camera left", + realtime_num_chunks=2, + ) + + generate_fn = get_generate_fn( + "robbyant/lingbot-world-fast-diffusers", + "video", + params, + ) + + assert generate_fn.__name__ == "generate_realtime_video" + + +def test_lingbot_realtime_plastic_beach_params_are_lossless_gt_ready(): + params = LINGBOT_WORLD_REALTIME_sampling_params + + assert "floating island hotel" in params.prompt + assert "825646291038" in str(params.image_path) + assert params.output_size == "832x480" + assert params.realtime_num_chunks == 4 + assert params.realtime_output_format is None + assert params.realtime_perf_ignore_initial_chunks == 2 + assert params.realtime_perf_thresholds["p95_chunk_total_ms"] == 5000.0 + assert params.realtime_perf_thresholds["p95_scheduler_forward_ms"] == 4500.0 + assert params.realtime_events[0]["kind"] == "camera_actions" + assert params.realtime_events[0]["payload"]["mode"] == "state" + + +def test_lingbot_realtime_case_is_registered_by_default(): + from sglang.multimodal_gen.test.server.gpu_cases import ( + ONE_GPU_CASES, + _make_lingbot_realtime_plastic_beach_case, + ) + + case = _make_lingbot_realtime_plastic_beach_case() + assert case.id == "lingbot_world_realtime_plastic_beach" + assert case.run_consistency_check is True + assert case.run_perf_check is True + assert case.sampling_params.realtime_output_format is None + assert any(item.id == case.id for item in ONE_GPU_CASES) diff --git a/python/sglang/multimodal_gen/test/unit/realtime/test_realtime_output_transport.py b/python/sglang/multimodal_gen/test/unit/realtime/test_realtime_output_transport.py new file mode 100644 index 000000000..1c9dd982f --- /dev/null +++ b/python/sglang/multimodal_gen/test/unit/realtime/test_realtime_output_transport.py @@ -0,0 +1,631 @@ +# SPDX-License-Identifier: Apache-2.0 + +import asyncio +from types import SimpleNamespace + +import msgspec.msgpack +import numpy as np +import torch + +from sglang.multimodal_gen.runtime.entrypoints.openai.realtime import ( + realtime_output_adapter, +) +from sglang.multimodal_gen.runtime.entrypoints.openai.realtime.realtime_output_adapter import ( + RawRGBRealtimeOutputAdapter, +) +from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import OutputBatch +from sglang.multimodal_gen.runtime.utils.realtime_video import ( + JPEG_FRAME_CONTENT_TYPE, + RAW_RGB_CONTENT_TYPE, + RAW_RGB_DELTA_GZIP_CONTENT_TYPE, + WEBP_FRAME_CONTENT_TYPE, + build_delta_gzip_raw_rgb_payload, + build_raw_rgb_frame_batches, + restore_delta_gzip_raw_rgb_payload, +) + + +def _unpack_frame_batch_messages(payloads): + messages = [] + for payload in payloads: + message = msgspec.msgpack.decode(payload) + assert message.pop("type") == "frame_batch" + frame_payload = message.pop("payload") + messages.append((message, frame_payload)) + return messages + + +def test_raw_rgb_frame_batches_preserve_frame_bytes_and_metadata(): + req = SimpleNamespace( + request_id="req-1", + block_idx=2, + data_type="video", + fps=24, + output_compression=None, + enable_frame_interpolation=False, + frame_interpolation_exp=1, + frame_interpolation_scale=1.0, + frame_interpolation_model_path=None, + enable_upscaling=False, + upscaling_model_path=None, + upscaling_scale=1, + ) + output_batch = OutputBatch(audio_sample_rate=None) + grayscale = np.array([[1, 2], [3, 4]], dtype=np.uint8) + rgba = np.array( + [ + [[5, 6, 7, 8], [9, 10, 11, 12]], + [[13, 14, 15, 16], [17, 18, 19, 20]], + ], + dtype=np.uint8, + ) + + def post_process_sample(*_args, **_kwargs): + return [grayscale, rgba] + + frame_batches, metadata = build_raw_rgb_frame_batches( + object(), + req, + output_batch, + post_process_sample, + ) + + assert metadata == { + "format": "rgb24", + "width": 2, + "height": 2, + "channels": 3, + "bytes_per_frame": 12, + } + assert len(frame_batches) == 1 + assert frame_batches[0][0] == bytes([1, 1, 1, 2, 2, 2, 3, 3, 3, 4, 4, 4]) + assert frame_batches[0][1] == bytes([5, 6, 7, 9, 10, 11, 13, 14, 15, 17, 18, 19]) + assert RAW_RGB_CONTENT_TYPE == "application/x-raw-rgb" + + +def test_raw_rgb_frame_batches_use_tensor_fast_path_without_postprocess(): + req = SimpleNamespace( + request_id="req-1", + block_idx=2, + data_type="video", + fps=24, + output_compression=None, + enable_frame_interpolation=False, + frame_interpolation_exp=1, + frame_interpolation_scale=1.0, + frame_interpolation_model_path=None, + enable_upscaling=False, + upscaling_model_path=None, + upscaling_scale=1, + ) + output_batch = OutputBatch(audio_sample_rate=None) + output = torch.tensor( + [[[[[0.0]], [[0.25]]], [[[0.5]], [[0.75]]], [[[1.0]], [[1.0]]]]] + ) + + def post_process_sample(*_args, **_kwargs): + raise AssertionError("tensor realtime output should not use postprocess") + + frame_batches, metadata = build_raw_rgb_frame_batches( + output, + req, + output_batch, + post_process_sample, + ) + + assert metadata == { + "format": "rgb24", + "width": 1, + "height": 1, + "channels": 3, + "bytes_per_frame": 3, + } + assert frame_batches == [[bytes([0, 127, 255]), bytes([63, 191, 255])]] + + +def test_output_batch_uses_raw_frame_transport_names(): + output_batch = OutputBatch( + raw_frame_batches=[[b"rgb"]], + raw_frame_content_type=RAW_RGB_CONTENT_TYPE, + raw_frame_metadata={"format": "rgb24"}, + ) + + assert output_batch.raw_frame_batches == [[b"rgb"]] + assert output_batch.raw_frame_content_type == RAW_RGB_CONTENT_TYPE + assert output_batch.raw_frame_metadata == {"format": "rgb24"} + + +def test_delta_gzip_raw_rgb_payload_roundtrips_exactly(): + frames = [ + bytes([1, 2, 3, 4, 5, 6]), + bytes([1, 2, 4, 4, 6, 6]), + bytes([2, 2, 4, 5, 6, 7]), + ] + + payload = build_delta_gzip_raw_rgb_payload(frames) + restored = restore_delta_gzip_raw_rgb_payload( + payload, + bytes_per_frame=6, + num_frames=3, + ) + + assert restored == b"".join(frames) + + +def test_raw_rgb_realtime_output_adapter_uses_lossless_compressed_payload(): + class _WebSocket: + def __init__(self): + self.payloads = [] + + async def send_bytes(self, payload): + self.payloads.append(payload) + + async def run(): + ws = _WebSocket() + adapter = RawRGBRealtimeOutputAdapter() + frame0 = bytes([1, 2, 3]) * 1000 + frame1 = bytes([1, 2, 4]) * 1000 + batch = SimpleNamespace( + block_idx=0, + request_id="req-1", + width=1000, + height=1, + enable_upscaling=False, + realtime_event_id=3, + ) + result = OutputBatch( + raw_frame_batches=[[frame0, frame1]], + raw_frame_content_type=RAW_RGB_CONTENT_TYPE, + raw_frame_metadata={ + "format": "rgb24", + "width": 1000, + "height": 1, + "channels": 3, + "bytes_per_frame": 3000, + }, + ) + + stats = await adapter.send(ws, SimpleNamespace(), result, batch) + return ws.payloads, stats, frame0 + frame1 + + payloads, stats, expected_frames = asyncio.run(run()) + + [(first_header, first_payload)] = _unpack_frame_batch_messages(payloads) + assert first_header["content_type"] == RAW_RGB_DELTA_GZIP_CONTENT_TYPE + assert first_header["encoding"] == "delta-gzip" + assert first_header["event_id"] == 3 + assert first_header["format"] == "rgb24" + assert first_header["channels"] == 3 + assert first_header["bytes_per_frame"] == 3000 + assert first_header["raw_size"] == 6000 + assert first_header["total_size"] == len(first_payload) + assert first_header["num_frames"] == 2 + assert first_header["num_frame_batches"] == 1 + assert first_header["frame_batch_index"] == 0 + assert "delta_reference" not in first_header + assert stats["raw_bytes"] == 6000 + assert stats["num_batches"] == 1 + assert stats["num_frames"] == 2 + restored_frames = restore_delta_gzip_raw_rgb_payload( + first_payload, + bytes_per_frame=3000, + num_frames=2, + ) + assert restored_frames == expected_frames + + +def test_raw_rgb_realtime_output_adapter_offloads_delta_payload_build( + monkeypatch, +): + calls = [] + + async def fake_to_thread(fn, *args, **kwargs): + calls.append((fn, args, kwargs)) + return fn(*args, **kwargs) + + monkeypatch.setattr( + realtime_output_adapter.asyncio, + "to_thread", + fake_to_thread, + ) + + class _WebSocket: + def __init__(self): + self.payloads = [] + + async def send_bytes(self, payload): + self.payloads.append(payload) + + async def run(): + ws = _WebSocket() + adapter = RawRGBRealtimeOutputAdapter() + frame0 = bytes([1, 2, 3]) * 1000 + frame1 = bytes([1, 2, 4]) * 1000 + batch = SimpleNamespace( + block_idx=0, + request_id="req-offload-delta", + width=1000, + height=1, + enable_upscaling=False, + realtime_event_id=3, + ) + result = OutputBatch( + raw_frame_batches=[[frame0, frame1]], + raw_frame_content_type=RAW_RGB_CONTENT_TYPE, + raw_frame_metadata={ + "format": "rgb24", + "width": 1000, + "height": 1, + "channels": 3, + "bytes_per_frame": 3000, + }, + ) + + await adapter.send(ws, SimpleNamespace(), result, batch) + return ws.payloads, frame0 + frame1 + + payloads, expected_frames = asyncio.run(run()) + + assert [call[0] for call in calls] == [ + realtime_output_adapter._build_transport_payload, + ] + [(first_header, first_payload)] = _unpack_frame_batch_messages(payloads) + assert first_header["encoding"] == "delta-gzip" + assert "delta_reference" not in first_header + restored_frames = restore_delta_gzip_raw_rgb_payload( + first_payload, + bytes_per_frame=3000, + num_frames=2, + ) + assert restored_frames == expected_frames + + +def test_raw_rgb_realtime_output_adapter_can_send_uncompressed_raw_frames(): + class _WebSocket: + def __init__(self): + self.payloads = [] + + async def send_bytes(self, payload): + self.payloads.append(payload) + + async def run(): + ws = _WebSocket() + adapter = RawRGBRealtimeOutputAdapter() + frame0 = bytes([1, 2, 3]) * 1000 + frame1 = bytes([1, 2, 4]) * 1000 + batch = SimpleNamespace( + block_idx=0, + request_id="req-raw", + width=1000, + height=1, + enable_upscaling=False, + realtime_event_id=3, + realtime_output_format="raw", + ) + result = OutputBatch( + raw_frame_batches=[[frame0, frame1]], + raw_frame_content_type=RAW_RGB_CONTENT_TYPE, + raw_frame_metadata={ + "format": "rgb24", + "width": 1000, + "height": 1, + "channels": 3, + "bytes_per_frame": 3000, + }, + ) + + stats = await adapter.send(ws, SimpleNamespace(), result, batch) + return ws.payloads, stats, frame0 + frame1 + + payloads, stats, expected_frames = asyncio.run(run()) + + [(first_header, first_payload)] = _unpack_frame_batch_messages(payloads) + assert first_header["content_type"] == RAW_RGB_CONTENT_TYPE + assert first_header["encoding"] == "raw" + assert first_header["raw_size"] == 6000 + assert first_header["total_size"] == 6000 + assert first_header["num_frames"] == 2 + assert first_header["num_frame_batches"] == 1 + assert first_header["frame_batch_index"] == 0 + assert first_payload == expected_frames + assert stats["raw_bytes"] == 6000 + assert stats["num_batches"] == 1 + assert stats["num_frames"] == 2 + assert stats["ws_payload_bytes"] == sum(len(payload) for payload in payloads) + + +def test_raw_rgb_realtime_output_adapter_uses_previous_frame_reference(): + class _WebSocket: + def __init__(self): + self.payloads = [] + + async def send_bytes(self, payload): + self.payloads.append(payload) + + async def run(): + ws = _WebSocket() + adapter = RawRGBRealtimeOutputAdapter() + base_batch = SimpleNamespace( + block_idx=0, + request_id="req-1", + width=2, + height=1, + enable_upscaling=False, + realtime_event_id=9, + ) + next_batch = SimpleNamespace( + block_idx=1, + request_id="req-2", + width=2, + height=1, + enable_upscaling=False, + realtime_event_id=9, + ) + metadata = { + "format": "rgb24", + "width": 2, + "height": 1, + "channels": 3, + "bytes_per_frame": 6, + } + first = OutputBatch( + raw_frame_batches=[[bytes([1, 2, 3, 4, 5, 6])]], + raw_frame_content_type=RAW_RGB_CONTENT_TYPE, + raw_frame_metadata=metadata, + ) + second = OutputBatch( + raw_frame_batches=[[bytes([1, 2, 4, 4, 6, 6])]], + raw_frame_content_type=RAW_RGB_CONTENT_TYPE, + raw_frame_metadata=metadata, + ) + + await adapter.send(ws, SimpleNamespace(), first, base_batch) + await adapter.send(ws, SimpleNamespace(), second, next_batch) + return ws.payloads + + payloads = asyncio.run(run()) + + (first_header, first_payload), (second_header, second_payload) = ( + _unpack_frame_batch_messages(payloads) + ) + assert "delta_reference" not in first_header + assert second_header["delta_reference"] == "previous-frame" + first_frame = restore_delta_gzip_raw_rgb_payload( + first_payload, + bytes_per_frame=6, + num_frames=1, + ) + second_frame = restore_delta_gzip_raw_rgb_payload( + second_payload, + bytes_per_frame=6, + num_frames=1, + reference_frame=first_frame, + ) + assert first_frame == bytes([1, 2, 3, 4, 5, 6]) + assert second_frame == bytes([1, 2, 4, 4, 6, 6]) + + +def test_raw_rgb_realtime_output_adapter_splits_large_frame_batches(): + class _WebSocket: + def __init__(self): + self.payloads = [] + + async def send_bytes(self, payload): + self.payloads.append(payload) + + async def run(): + ws = _WebSocket() + adapter = RawRGBRealtimeOutputAdapter() + frames = [bytes([idx, idx + 1, idx + 2]) for idx in range(17)] + batch = SimpleNamespace( + block_idx=4, + request_id="req-split", + width=1, + height=1, + enable_upscaling=False, + realtime_event_id=12, + ) + result = OutputBatch( + raw_frame_batches=[frames], + raw_frame_content_type=RAW_RGB_CONTENT_TYPE, + raw_frame_metadata={ + "format": "rgb24", + "width": 1, + "height": 1, + "channels": 3, + "bytes_per_frame": 3, + }, + ) + + stats = await adapter.send(ws, SimpleNamespace(), result, batch) + return ws.payloads, stats + + payloads, stats = asyncio.run(run()) + + headers = [header for header, _ in _unpack_frame_batch_messages(payloads)] + assert len(headers) == 2 + assert [header["chunk_index"] for header in headers] == [4, 4] + assert [header["frame_batch_index"] for header in headers] == [0, 1] + assert [header["num_frame_batches"] for header in headers] == [2, 2] + assert [header["num_frames"] for header in headers] == [16, 1] + assert [header["is_final_frame_batch"] for header in headers] == [False, True] + assert "delta_reference" not in headers[0] + assert headers[1]["delta_reference"] == "previous-frame" + assert stats["num_batches"] == 2 + assert stats["num_frames"] == 17 + + +def test_raw_rgb_realtime_output_adapter_can_send_webp_preview_frames(): + class _WebSocket: + def __init__(self): + self.payloads = [] + + async def send_bytes(self, payload): + self.payloads.append(payload) + + async def run(): + ws = _WebSocket() + adapter = RawRGBRealtimeOutputAdapter() + batch = SimpleNamespace( + block_idx=0, + request_id="req-webp", + width=2, + height=1, + enable_upscaling=False, + realtime_event_id=5, + realtime_output_format="webp", + output_compression=90, + ) + result = OutputBatch( + raw_frame_batches=[[bytes([255, 0, 0, 0, 255, 0])]], + raw_frame_content_type=RAW_RGB_CONTENT_TYPE, + raw_frame_metadata={ + "format": "rgb24", + "width": 2, + "height": 1, + "channels": 3, + "bytes_per_frame": 6, + }, + ) + + stats = await adapter.send(ws, SimpleNamespace(), result, batch) + return ws.payloads, stats + + payloads, stats = asyncio.run(run()) + + [(header, frame_payload)] = _unpack_frame_batch_messages(payloads) + assert header["content_type"] == WEBP_FRAME_CONTENT_TYPE + assert header["format"] == "webp" + assert header["encoding"] == "webp" + assert header["num_frames"] == 1 + assert header["is_final_frame_batch"] is True + assert frame_payload.startswith(b"RIFF") + assert stats["num_batches"] == 1 + assert stats["num_frames"] == 1 + + +def test_raw_rgb_realtime_output_adapter_offloads_preview_encoding(monkeypatch): + calls = [] + + async def fake_to_thread(fn, *args, **kwargs): + calls.append((fn, args, kwargs)) + return fn(*args, **kwargs) + + monkeypatch.setattr( + realtime_output_adapter.asyncio, + "to_thread", + fake_to_thread, + ) + + class _WebSocket: + def __init__(self): + self.payloads = [] + + async def send_bytes(self, payload): + self.payloads.append(payload) + + async def run(): + ws = _WebSocket() + adapter = RawRGBRealtimeOutputAdapter() + batch = SimpleNamespace( + block_idx=0, + request_id="req-webp-offload", + width=2, + height=1, + enable_upscaling=False, + realtime_event_id=5, + realtime_output_format="webp", + output_compression=90, + ) + result = OutputBatch( + raw_frame_batches=[ + [ + bytes([255, 0, 0, 0, 255, 0]), + bytes([0, 0, 255, 255, 255, 0]), + ] + ], + raw_frame_content_type=RAW_RGB_CONTENT_TYPE, + raw_frame_metadata={ + "format": "rgb24", + "width": 2, + "height": 1, + "channels": 3, + "bytes_per_frame": 6, + }, + ) + + await adapter.send(ws, SimpleNamespace(), result, batch) + return ws.payloads + + payloads = asyncio.run(run()) + + assert [call[0] for call in calls] == [ + realtime_output_adapter._build_transport_payload, + realtime_output_adapter._build_transport_payload, + ] + (first_header, first_payload), (second_header, second_payload) = ( + _unpack_frame_batch_messages(payloads) + ) + assert first_header["content_type"] == WEBP_FRAME_CONTENT_TYPE + assert first_header["encoding"] == "webp" + assert first_header["num_frames"] == 1 + assert first_header["frame_batch_index"] == 0 + assert first_header["num_frame_batches"] == 2 + assert first_header["is_final_frame_batch"] is False + assert second_header["content_type"] == WEBP_FRAME_CONTENT_TYPE + assert second_header["encoding"] == "webp" + assert second_header["num_frames"] == 1 + assert second_header["frame_batch_index"] == 1 + assert second_header["num_frame_batches"] == 2 + assert second_header["is_final_frame_batch"] is True + assert first_payload.startswith(b"RIFF") + assert second_payload.startswith(b"RIFF") + + +def test_raw_rgb_realtime_output_adapter_can_send_jpeg_preview_frames(): + class _WebSocket: + def __init__(self): + self.payloads = [] + + async def send_bytes(self, payload): + self.payloads.append(payload) + + async def run(): + ws = _WebSocket() + adapter = RawRGBRealtimeOutputAdapter() + batch = SimpleNamespace( + block_idx=0, + request_id="req-jpeg", + width=2, + height=1, + enable_upscaling=False, + realtime_event_id=5, + realtime_output_format="jpeg", + output_compression=85, + ) + result = OutputBatch( + raw_frame_batches=[[bytes([255, 0, 0, 0, 255, 0])]], + raw_frame_content_type=RAW_RGB_CONTENT_TYPE, + raw_frame_metadata={ + "format": "rgb24", + "width": 2, + "height": 1, + "channels": 3, + "bytes_per_frame": 6, + }, + ) + + stats = await adapter.send(ws, SimpleNamespace(), result, batch) + return ws.payloads, stats + + payloads, stats = asyncio.run(run()) + + [(header, frame_payload)] = _unpack_frame_batch_messages(payloads) + assert header["content_type"] == JPEG_FRAME_CONTENT_TYPE + assert header["format"] == "jpeg" + assert header["encoding"] == "jpeg" + assert header["num_frames"] == 1 + assert header["is_final_frame_batch"] is True + assert frame_payload.startswith(b"\xff\xd8") + assert stats["num_batches"] == 1 + assert stats["num_frames"] == 1 diff --git a/python/sglang/multimodal_gen/test/unit/realtime/test_realtime_runtime.py b/python/sglang/multimodal_gen/test/unit/realtime/test_realtime_runtime.py new file mode 100644 index 000000000..db75c3d46 --- /dev/null +++ b/python/sglang/multimodal_gen/test/unit/realtime/test_realtime_runtime.py @@ -0,0 +1,883 @@ +# SPDX-License-Identifier: Apache-2.0 + +import asyncio +import time +from types import SimpleNamespace + +import msgspec.msgpack +import numpy as np +import torch + +from sglang.multimodal_gen.configs.pipeline_configs.lingbot_world import ( + LingBotWorldCausalDMDConfig, + _actions_to_c2ws, +) +from sglang.multimodal_gen.runtime.entrypoints.openai.protocol import ( + RealtimeEvent, + RealtimeVideoGenerationsRequest, +) +from sglang.multimodal_gen.runtime.entrypoints.openai.realtime import ( + realtime_video_api, +) +from sglang.multimodal_gen.runtime.entrypoints.openai.realtime.adapters import ( + lingbot_world_realtime_adapter as lingbot_realtime, +) +from sglang.multimodal_gen.runtime.entrypoints.openai.realtime.generate_session import ( + GenerateSession, +) +from sglang.multimodal_gen.runtime.entrypoints.openai.realtime.realtime_output_adapter import ( + empty_frame_send_stats, +) +from sglang.multimodal_gen.runtime.entrypoints.openai.realtime.registry import ( + get_realtime_model_adapter, +) +from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import OutputBatch +from sglang.multimodal_gen.runtime.pipelines_core.stages.model_specific_stages.lingbot_world import ( + LingBotWorldCausalDMDDenoisingStage, +) +from sglang.multimodal_gen.runtime.pipelines_core.stages.realtime_input_validation import ( + RealtimeInputValidationStage, + RealtimeInputValidationState, +) +from sglang.multimodal_gen.runtime.realtime.condition_events import ( + ControlStateTransition, +) +from sglang.multimodal_gen.runtime.realtime.session import ( + BaseRealtimeState, + RealtimeSessionCache, +) +from sglang.multimodal_gen.runtime.utils.realtime_video import ( + RAW_RGB_CONTENT_TYPE, +) + + +class _Req(SimpleNamespace): + realtime_session_id: str | None = None + block_idx: int = 0 + session = None + + +class _State(BaseRealtimeState): + def __init__(self): + self.disposed = False + + def dispose(self) -> None: + self.disposed = True + + +def test_realtime_session_cache_reuses_and_releases_state(): + cache = RealtimeSessionCache(max_sessions=1) + first = _Req(realtime_session_id="session-a", block_idx=0, session=None) + cache.attach(first) + state = first.session.get_or_create_state(_State) + + second = _Req(realtime_session_id="session-a", block_idx=1, session=None) + cache.attach(second) + + assert second.session is first.session + assert second.session.get_state(_State) is state + assert cache.release("session-a") + assert state.disposed + assert not cache.release("session-a") + + +def test_realtime_session_cache_rejects_missing_nonzero_chunk(): + cache = RealtimeSessionCache(max_sessions=1) + req = _Req(realtime_session_id="missing", block_idx=1, session=None) + + try: + cache.attach(req) + except ValueError as exc: + assert "Missing realtime session state" in str(exc) + else: + raise AssertionError("expected missing realtime session to fail") + + +def test_lingbot_realtime_state_uses_generic_condition_queue(): + state = lingbot_realtime.LingBotWorldRealtimeState() + + assert state.sample_camera_actions(3) == [[], [], []] + state.receive_camera_actions([["w"], ["a"], ["s"], ["d"]]) + assert state.sample_camera_actions(3) == [["w"], ["a"], ["s"]] + assert state.sample_camera_actions(3) == [["d"], [], []] + assert state.sample_camera_actions(3) == [[], [], []] + + state.receive_prompt("turn left") + assert state.has_prompt() + assert state.sample_prompt() == "turn left" + assert not state.has_prompt() + + +def test_lingbot_realtime_camera_events_preserve_short_presses(): + state = lingbot_realtime.LingBotWorldRealtimeState() + + state.receive_camera_state(["w"], event_id=7) + state.receive_camera_state([], event_id=8) + + assert state.sample_camera_actions(3) == [["w"], [], []] + assert state.latest_sampled_event_id == 8 + assert state.sample_camera_actions(3) == [[], [], []] + + +def test_lingbot_realtime_camera_state_holds_until_release(): + state = lingbot_realtime.LingBotWorldRealtimeState() + + state.receive_camera_state(["w"], event_id=7) + + assert state.sample_camera_actions(3) == [["w"], ["w"], ["w"]] + assert state.latest_sampled_event_id == 7 + assert state.sample_camera_actions(3) == [["w"], ["w"], ["w"]] + + state.receive_camera_state([], event_id=8) + + assert state.sample_camera_actions(3) == [[], [], []] + assert state.latest_sampled_event_id == 8 + + +def test_lingbot_realtime_camera_state_compacts_multiple_pending_updates(): + state = lingbot_realtime.LingBotWorldRealtimeState() + state.receive_camera_state_transitions( + [ + ControlStateTransition(payload=["w"], seq_id=7), + ControlStateTransition(payload=["w", "d"], seq_id=8), + ControlStateTransition(payload=["d"], seq_id=9), + ] + ) + + assert state.sample_camera_actions(3) == [["d"], ["d"], ["d"]] + assert state.latest_sampled_event_id == 9 + + +def test_lingbot_realtime_adapter_ingests_generic_events(): + adapter = lingbot_realtime.LingBotWorldRealtimeAdapter() + session = GenerateSession() + session.set_adapter(adapter) + + camera_event = RealtimeEvent( + type="event", + kind="camera_actions", + payload=[["w"], ["d"]], + event_id=7, + ) + prompt_event = RealtimeEvent( + type="event", + kind="prompt", + payload="turn left", + event_id=8, + ) + + assert ( + adapter.ingest_event(session, camera_event) + == "kind=camera_actions, mode=script, frames=2" + ) + assert adapter.ingest_event(session, prompt_event) == "kind=prompt, prompt_len=9" + state = adapter._state(session) + assert state.sample_camera_actions(3) == [["w"], ["d"], []] + assert state.sample_prompt() == "turn left" + assert state.latest_sampled_event_id == 8 + + +def test_lingbot_realtime_adapter_ingests_state_camera_events(): + adapter = lingbot_realtime.LingBotWorldRealtimeAdapter() + session = GenerateSession() + session.set_adapter(adapter) + + camera_event = RealtimeEvent( + type="event", + kind="camera_actions", + payload={ + "mode": "state", + "transitions": [ + {"actions": ["w"], "client_ts_ms": 100}, + {"actions": [], "client_ts_ms": 120}, + ], + }, + event_id=11, + ) + + assert ( + adapter.ingest_event(session, camera_event) + == "kind=camera_actions, mode=state, transitions=2" + ) + state = adapter._state(session) + assert state.sample_camera_actions(3) == [["w"], [], []] + assert state.latest_sampled_event_id == 11 + + +def test_generate_session_tracks_active_chunk_context(): + session = GenerateSession() + + chunk = session.new_chunk() + + assert chunk.session_id == session.id + assert chunk.index == 0 + assert chunk.request_id.startswith(f"{session.id}_") + try: + session.new_chunk() + except RuntimeError as exc: + assert "previous realtime chunk" in str(exc) + else: + raise AssertionError("expected active chunk to block new chunk") + + session.generate_chunk_completed() + next_chunk = session.new_chunk() + + assert next_chunk.index == 1 + assert next_chunk.request_id.startswith(f"{session.id}_") + assert next_chunk.request_id != chunk.request_id + + +def test_generate_session_respects_max_chunks(): + session = GenerateSession() + session.set_request( + RealtimeVideoGenerationsRequest( + type="init", + prompt="walk forward", + max_chunks=1, + ) + ) + + assert not session.reached_max_chunks() + session.new_chunk() + session.generate_chunk_completed() + + assert session.reached_max_chunks() + + +def test_generate_loop_overlaps_send_with_next_generation(monkeypatch): + events = [] + + class _Adapter: + async def wait_for_next_chunk(self, session): + del session + + def prepare_next_request(self, session, server_args, chunk): + del session, server_args + return SimpleNamespace( + block_idx=chunk.index, + request_id=chunk.request_id, + condition_inputs={}, + ) + + async def send_output(self, ws, session, result, batch): + del ws, session, result + events.append(f"send_start_{batch.block_idx}") + await asyncio.sleep(0.05) + events.append(f"send_end_{batch.block_idx}") + return empty_frame_send_stats("test") + + def on_chunk_complete(self, session, result): + del result + session.generate_chunk_completed() + + async def fake_process_generation_batch(client, batch): + del client + events.append(f"generate_start_{batch.block_idx}") + await asyncio.sleep(0.01) + events.append(f"generate_end_{batch.block_idx}") + return None, SimpleNamespace() + + monkeypatch.setattr( + realtime_video_api, + "get_global_server_args", + lambda: SimpleNamespace(), + ) + monkeypatch.setattr( + realtime_video_api, + "process_generation_batch", + fake_process_generation_batch, + ) + + session = GenerateSession() + session.adapter = _Adapter() + session.set_request( + RealtimeVideoGenerationsRequest( + type="init", + prompt="walk forward", + max_chunks=2, + ) + ) + + class _Ws: + async def send_bytes(self, _message): + pass + + asyncio.run(realtime_video_api._generate_loop(_Ws(), session)) + + assert events.index("generate_start_1") < events.index("send_end_0") + assert events.index("send_end_0") < events.index("send_start_1") + assert events[-1] == "send_end_1" + + +def test_send_output_emits_chunk_stats_message(): + sent_messages = [] + + class _Ws: + async def send_bytes(self, message): + sent_messages.append(message) + + class _Adapter: + async def send_output(self, ws, session, result, batch): + del ws, session, result, batch + return { + "header_pack_ms": 0.1, + "header_write_ms": 0.2, + "raw_payload_build_ms": 3.0, + "raw_write_ms": 42.0, + "ws_write_ms": 42.2, + "raw_bytes": 1200, + "ws_payload_bytes": 450, + "num_frames": 3, + "num_batches": 3, + "frame_shape": (1, 2, 3), + "content_type": "image/webp", + } + + session = GenerateSession() + session.adapter = _Adapter() + chunk = session.new_chunk() + batch = SimpleNamespace( + block_idx=7, + realtime_event_id=11, + condition_inputs={"camera_actions": [["w"]]}, + ) + + stats = asyncio.run( + realtime_video_api._send_output_and_log( + _Ws(), + session, + chunk, + batch, + SimpleNamespace(), + request_prepare_ms=1.0, + scheduler_forward_ms=2.0, + chunk_started=time.perf_counter(), + ) + ) + + message = msgspec.msgpack.decode(sent_messages[-1]) + assert stats["raw_write_ms"] == 42.0 + assert message["type"] == "chunk_stats" + assert message["chunk_index"] == 7 + assert message["event_id"] == 11 + assert message["raw_write_ms"] == 42 + assert message["ws_write_ms"] == 42 + assert message["ws_payload_bytes"] == 450 + assert message["content_type"] == "image/webp" + assert bytes([0xCB]) not in sent_messages[-1] + + +def test_listen_generate_request_propagates_disconnect_without_error_write(): + sent_messages = [] + + class _Ws: + async def receive_bytes(self): + raise realtime_video_api.WebSocketDisconnect(1000, "client close") + + async def send_bytes(self, message): + sent_messages.append(message) + + try: + asyncio.run( + realtime_video_api._listen_generate_request(_Ws(), GenerateSession()) + ) + except realtime_video_api.WebSocketDisconnect: + pass + else: + raise AssertionError("expected websocket disconnect to propagate") + + assert sent_messages == [] + + +def test_wait_for_active_session_slot_observes_release(monkeypatch): + async def run(): + realtime_video_api._ACTIVE_SESSION_IDS.clear() + realtime_video_api._ACTIVE_SESSION_IDS.add("old-session") + + async def fake_sleep(_seconds): + realtime_video_api._ACTIVE_SESSION_IDS.clear() + + monkeypatch.setattr(realtime_video_api.asyncio, "sleep", fake_sleep) + try: + return await realtime_video_api._wait_for_active_session_slot( + timeout_s=1.0, + interval_s=0.1, + ) + finally: + realtime_video_api._ACTIVE_SESSION_IDS.clear() + + assert asyncio.run(run()) + + +def test_cleanup_realtime_session_keeps_active_slot_during_scheduler_release( + monkeypatch, +): + async def run(): + session = GenerateSession() + realtime_video_api._ACTIVE_SESSION_IDS.clear() + realtime_video_api._ACTIVE_SESSION_IDS.add(session.id) + release_seen = [] + + async def fake_forward(req): + release_seen.append( + (req.session_id, session.id in realtime_video_api._ACTIVE_SESSION_IDS) + ) + + monkeypatch.setattr( + realtime_video_api.async_scheduler_client, + "forward", + fake_forward, + ) + try: + await realtime_video_api._cleanup_realtime_session(session, None, None) + still_active_after_cleanup = ( + session.id in realtime_video_api._ACTIVE_SESSION_IDS + ) + finally: + realtime_video_api._ACTIVE_SESSION_IDS.clear() + + return session.id, release_seen, still_active_after_cleanup + + session_id, release_seen, still_active_after_cleanup = asyncio.run(run()) + + assert release_seen == [(session_id, True)] + assert still_active_after_cleanup + + +def test_lingbot_realtime_adapter_prepares_chunk_request(monkeypatch): + adapter = lingbot_realtime.LingBotWorldRealtimeAdapter() + session = GenerateSession() + session.set_adapter(adapter) + session.set_request( + RealtimeVideoGenerationsRequest( + type="init", + prompt="walk forward", + num_frames=9, + fps=24, + realtime_causal_sink_size=3, + realtime_causal_kv_cache_num_frames=45, + ) + ) + state = adapter._state(session) + state.receive_camera_state(["w"], event_id=11) + server_args = SimpleNamespace( + pipeline_config=SimpleNamespace( + dit_config=SimpleNamespace( + arch_config=SimpleNamespace(num_frames_per_block=3) + ), + vae_config=SimpleNamespace( + arch_config=SimpleNamespace(temporal_compression_ratio=4) + ), + ) + ) + seen = {} + + def fake_build_sampling_params(request_id, **kwargs): + seen["request_id"] = request_id + seen["kwargs"] = kwargs + return SimpleNamespace( + request_id=request_id, + prompt=kwargs["prompt"], + condition_inputs=kwargs["condition_inputs"], + realtime_chunk_size=kwargs["realtime_chunk_size"], + ) + + def fake_prepare_backend_request(server_args, sampling_params): + seen["server_args"] = server_args + return SimpleNamespace( + request_id=sampling_params.request_id, + prompt=sampling_params.prompt, + condition_inputs=dict(sampling_params.condition_inputs), + realtime_chunk_size=sampling_params.realtime_chunk_size, + ) + + monkeypatch.setattr( + lingbot_realtime, + "build_sampling_params", + fake_build_sampling_params, + ) + monkeypatch.setattr( + lingbot_realtime, + "prepare_request", + fake_prepare_backend_request, + ) + chunk = session.new_chunk() + + batch = adapter.prepare_next_request(session, server_args, chunk) + + assert seen["request_id"] == chunk.request_id + assert seen["kwargs"]["prompt"] == "walk forward" + assert seen["kwargs"]["num_frames"] == 21 + assert seen["kwargs"]["num_inference_steps"] == 4 + assert seen["kwargs"]["guidance_scale"] == 1.0 + assert seen["kwargs"]["condition_inputs"] == { + "camera_actions": [["w"], ["w"], ["w"]] + } + assert batch.request_id == chunk.request_id + assert batch.condition_inputs == {"camera_actions": [["w"], ["w"], ["w"]]} + assert batch.realtime_chunk_size == 3 + assert batch.session is session.realtime_session + assert batch.realtime_session_id == session.id + assert batch.block_idx == 0 + assert batch.return_raw_frames is True + assert batch.realtime_event_id == 11 + assert batch.realtime_causal_sink_size == 3 + assert batch.realtime_causal_kv_cache_num_frames == 45 + + +def test_lingbot_realtime_condition_horizon_repeats_blank_tail_chunk(): + config = LingBotWorldCausalDMDConfig() + chunk_size = config.dit_config.arch_config.num_frames_per_block + latent_channels = config.vae_config.arch_config.z_dim + temporal_ratio = config.vae_config.arch_config.temporal_compression_ratio + spatial_ratio = config.vae_config.arch_config.spatial_compression_ratio + request = RealtimeVideoGenerationsRequest( + type="init", + prompt="walk forward", + num_frames=9, + ) + server_args = SimpleNamespace(pipeline_config=config) + + num_frames = lingbot_realtime.LingBotWorldRealtimeAdapter._condition_num_frames( + request=request, + server_args=server_args, + chunk_size=chunk_size, + ) + latent_frames = (num_frames - 1) // temporal_ratio + 1 + raw_request_latent_frames = (request.num_frames - 1) // temporal_ratio + 1 + + assert raw_request_latent_frames == chunk_size + assert latent_frames == chunk_size * 2 + + latent_condition = torch.ones(1, latent_channels, latent_frames, 2, 2) + batch = SimpleNamespace( + height=2 * spatial_ratio, + width=2 * spatial_ratio, + ) + condition_full = config.postprocess_image_latent(latent_condition, batch) + first_chunk = LingBotWorldCausalDMDDenoisingStage._select_i2v_condition_chunk( + condition_full, + chunk_idx=0, + chunk_size=chunk_size, + ) + tail_chunk = LingBotWorldCausalDMDDenoisingStage._select_i2v_condition_chunk( + condition_full, + chunk_idx=1, + chunk_size=chunk_size, + ) + repeated_tail = LingBotWorldCausalDMDDenoisingStage._select_i2v_condition_chunk( + condition_full, + chunk_idx=2, + chunk_size=chunk_size, + ) + + assert torch.count_nonzero(first_chunk[:, :temporal_ratio]) > 0 + assert torch.count_nonzero(tail_chunk[:, :temporal_ratio]) == 0 + assert not torch.equal(tail_chunk, first_chunk) + assert torch.equal(repeated_tail, tail_chunk) + + long_request = RealtimeVideoGenerationsRequest( + type="init", + prompt="walk forward", + num_frames=45, + ) + assert ( + lingbot_realtime.LingBotWorldRealtimeAdapter._condition_num_frames( + request=long_request, + server_args=server_args, + chunk_size=chunk_size, + ) + == 45 + ) + + +def test_lingbot_realtime_adapter_ingests_initial_condition_inputs(): + adapter = lingbot_realtime.LingBotWorldRealtimeAdapter() + session = GenerateSession() + session.set_adapter(adapter) + request = RealtimeVideoGenerationsRequest( + type="init", + prompt="walk forward", + condition_inputs={"camera_actions": [["w"], ["d"], ["a"], ["s"]]}, + ) + + state = adapter._state(session) + + asyncio.run(adapter.on_init(session, request)) + + assert state.sample_camera_actions(3) == [["w"], ["d"], ["a"]] + assert state.sample_camera_actions(3) == [["s"], [], []] + assert state.sample_camera_actions(3) == [[], [], []] + + +def test_realtime_video_request_accepts_raw_lossless_output_format(): + request = RealtimeVideoGenerationsRequest( + type="init", + prompt="walk forward", + realtime_output_format="raw", + realtime_causal_sink_size=3, + realtime_causal_kv_cache_num_frames=45, + ) + + assert request.realtime_output_format == "raw" + assert request.realtime_causal_sink_size == 3 + assert request.realtime_causal_kv_cache_num_frames == 45 + + +def test_lingbot_realtime_adapter_does_not_wait_for_idle_chunks(): + async def run(): + adapter = lingbot_realtime.LingBotWorldRealtimeAdapter() + session = GenerateSession() + session.set_adapter(adapter) + await adapter.wait_for_next_chunk(session) + + session.generate_chunk_cnt = 1 + await asyncio.wait_for(adapter.wait_for_next_chunk(session), timeout=1) + + asyncio.run(run()) + + +def test_lingbot_realtime_adapter_sends_stale_output_for_client_cutover(): + adapter = lingbot_realtime.LingBotWorldRealtimeAdapter() + session = GenerateSession() + session.set_adapter(adapter) + state = adapter._state(session) + state.receive_camera_actions([["d"]], event_id=7) + calls = [] + + async def fake_send(ws, session_arg, result_arg, batch_arg): + calls.append((ws, session_arg, result_arg, batch_arg)) + return empty_frame_send_stats("sent") + + adapter.output_adapter = SimpleNamespace(send=fake_send) + batch = SimpleNamespace(block_idx=3, realtime_event_id=6) + result = OutputBatch( + raw_frame_batches=[[b"stale"]], + raw_frame_content_type=RAW_RGB_CONTENT_TYPE, + ) + + stats = asyncio.run(adapter.send_output(SimpleNamespace(), session, result, batch)) + + assert stats == empty_frame_send_stats("sent") + assert calls[0][1] is session + assert calls[0][2] is result + assert calls[0][3] is batch + + +def test_lingbot_i2v_condition_repeats_last_chunk(): + condition_full = torch.ones(1, 20, 3, 2, 2) + + first = LingBotWorldCausalDMDDenoisingStage._select_i2v_condition_chunk( + condition_full, + chunk_idx=0, + chunk_size=3, + ) + second = LingBotWorldCausalDMDDenoisingStage._select_i2v_condition_chunk( + condition_full, + chunk_idx=1, + chunk_size=3, + ) + + assert torch.equal(first, condition_full) + assert second.shape == condition_full.shape + assert torch.equal(second, condition_full) + + +def test_lingbot_i2v_condition_pads_tail_then_repeats_it(): + condition_full = torch.ones(1, 20, 1, 2, 2) + + first = LingBotWorldCausalDMDDenoisingStage._select_i2v_condition_chunk( + condition_full, + chunk_idx=0, + chunk_size=3, + ) + second = LingBotWorldCausalDMDDenoisingStage._select_i2v_condition_chunk( + condition_full, + chunk_idx=1, + chunk_size=3, + ) + + assert first.shape == (1, 20, 3, 2, 2) + assert torch.equal(first[:, :, :1], condition_full) + assert torch.count_nonzero(first[:, :, 1:]) == 0 + assert second.shape == first.shape + assert torch.equal(second, first) + + +def test_lingbot_i2v_condition_uses_available_non_initial_chunks(): + first_chunk = torch.ones(1, 20, 3, 2, 2) + second_chunk = first_chunk * 2 + condition_full = torch.cat([first_chunk, second_chunk], dim=2) + + second = LingBotWorldCausalDMDDenoisingStage._select_i2v_condition_chunk( + condition_full, + chunk_idx=1, + chunk_size=3, + ) + third = LingBotWorldCausalDMDDenoisingStage._select_i2v_condition_chunk( + condition_full, + chunk_idx=2, + chunk_size=3, + ) + + assert torch.equal(second, second_chunk) + assert third.shape == second_chunk.shape + assert torch.equal(third, second_chunk) + + +def test_realtime_input_validation_reuses_generator_across_chunks(): + stage = RealtimeInputValidationStage.__new__(RealtimeInputValidationStage) + state = RealtimeInputValidationState() + generator = torch.Generator(device="cpu").manual_seed(123) + + first = SimpleNamespace( + block_idx=0, + generator=generator, + seeds=None, + seed=123, + generator_device="cpu", + num_outputs_per_prompt=1, + ) + second = SimpleNamespace( + block_idx=1, + generator=torch.Generator(device="cpu").manual_seed(123), + seeds=None, + seed=123, + generator_device="cpu", + num_outputs_per_prompt=1, + ) + + stage._cache_generator(first, state) + stage._reuse_or_cache_generator(second, state) + + assert second.generator is generator + + +def test_realtime_registry_resolves_lingbot_adapter(): + server_args = SimpleNamespace(pipeline_config=LingBotWorldCausalDMDConfig()) + + adapter = get_realtime_model_adapter(server_args) + + assert isinstance(adapter, lingbot_realtime.LingBotWorldRealtimeAdapter) + + +def test_sampling_params_apply_condition_inputs_to_req(): + from sglang.multimodal_gen.configs.sample.sampling_params import SamplingParams + + sampling_params = SamplingParams( + prompt="test", + num_inference_steps=1, + condition_inputs={"camera_actions": [["w"]]}, + realtime_chunk_size=3, + ) + req = SimpleNamespace(extra={}, condition_inputs={}, realtime_chunk_size=None) + + sampling_params.apply_request_extra(req) + + assert req.condition_inputs == {"camera_actions": [["w"]]} + assert req.realtime_chunk_size == 3 + + +def test_realtime_chunk_latent_preparation_uses_chunk_spec(): + from unittest.mock import patch + + import torch + + from sglang.multimodal_gen.runtime.pipelines_core.stages import ( + RealtimeChunkLatentPreparationStage, + ) + + transformer = SimpleNamespace( + config=SimpleNamespace( + arch_config=SimpleNamespace(out_channels=16, num_frames_per_block=3) + ) + ) + stage = RealtimeChunkLatentPreparationStage.__new__( + RealtimeChunkLatentPreparationStage + ) + stage.scheduler = SimpleNamespace(init_noise_sigma=10.0) + stage.transformer = transformer + batch = SimpleNamespace( + batch_size=1, + generator=None, + height=None, + width=None, + image_latent=torch.zeros(2, 20, 6, 4, 5, dtype=torch.float32), + latents=None, + realtime_chunk_size=2, + ) + + def fake_randn_tensor(shape, generator, device, dtype): + del generator + return torch.ones(shape, device=device, dtype=dtype) + + with ( + patch( + "sglang.multimodal_gen.runtime.pipelines_core.stages.latent_preparation.get_local_torch_device", + return_value=torch.device("cpu"), + ), + patch( + "sglang.multimodal_gen.runtime.pipelines_core.stages.latent_preparation.randn_tensor", + side_effect=fake_randn_tensor, + ), + ): + result = stage.forward(batch, SimpleNamespace()) + + assert tuple(result.latents.shape) == (2, 16, 2, 4, 5) + assert result.latents.dtype == torch.float32 + assert torch.all(result.latents == 1) + assert result.raw_latent_shape == result.latents.shape + + +def test_lingbot_camera_actions_have_deterministic_pose_precision(): + poses = _actions_to_c2ws([["w"], ["d"]]) + + np.testing.assert_allclose(poses[1][:3, 3], [0.0, 0.0, 0.05], atol=1e-6) + np.testing.assert_allclose(poses[2][:3, 3], [0.05, 0.0, 0.05], atol=1e-6) + + yaw_pose = _actions_to_c2ws([["l"]])[1] + expected_yaw = np.array( + [ + [np.cos(np.deg2rad(6.0)), 0.0, np.sin(np.deg2rad(6.0))], + [0.0, 1.0, 0.0], + [-np.sin(np.deg2rad(6.0)), 0.0, np.cos(np.deg2rad(6.0))], + ] + ) + np.testing.assert_allclose(yaw_pose[:3, :3], expected_yaw, atol=1e-6) + + +def test_lingbot_camera_condition_uses_condition_inputs_without_session(): + import torch + + from sglang.multimodal_gen.configs.pipeline_configs.lingbot_world import ( + LingBotWorldCausalDMDConfig, + ) + + batch = SimpleNamespace( + c2ws_plucker_emb=None, + condition_inputs={"camera_actions": [["w"]]}, + session=None, + realtime_chunk_size=3, + width=16, + height=16, + block_idx=0, + realtime_session_id=None, + ) + pipeline_config = LingBotWorldCausalDMDConfig() + + condition = pipeline_config.prepare_world_condition( + batch=batch, + device="cpu", + dtype=torch.float32, + ) + + spatial_scale = pipeline_config.vae_config.arch_config.spatial_compression_ratio + expected_shape = ( + 1, + 6 * spatial_scale * spatial_scale, + batch.realtime_chunk_size, + batch.height // spatial_scale, + batch.width // spatial_scale, + ) + assert tuple(condition["c2ws_plucker_emb"].shape) == expected_shape + assert condition["c2ws_plucker_emb"].dtype == torch.float32 diff --git a/python/sglang/multimodal_gen/test/unit/realtime/test_realtime_vae.py b/python/sglang/multimodal_gen/test/unit/realtime/test_realtime_vae.py new file mode 100644 index 000000000..d33aaff25 --- /dev/null +++ b/python/sglang/multimodal_gen/test/unit/realtime/test_realtime_vae.py @@ -0,0 +1,171 @@ +# SPDX-License-Identifier: Apache-2.0 + +from types import SimpleNamespace + +import torch + +from sglang.multimodal_gen.runtime.pipelines_core.stages.realtime_vae import ( + CausalVaeDecodingStage, + RealtimeVAEDecodeState, +) + + +def test_realtime_vae_decode_state_clears_model_cache_on_dispose(): + calls = [] + state = RealtimeVAEDecodeState() + state.reset_causal_decode_state = lambda: calls.append("reset") + + state.dispose() + + assert calls == ["reset"] + assert state.reset_causal_decode_state is None + + +def test_causal_vae_decoding_stage_keeps_wan_decoder_cache(monkeypatch): + from sglang.multimodal_gen.runtime.pipelines_core.stages import realtime_vae + + class _WanVAE: + def __init__(self): + self.config = SimpleNamespace(patch_size=None) + self.clear_calls = 0 + self.decoder_first_chunk_flags = [] + self._feat_map = [] + self._conv_idx = [0] + + def to(self, device=None, dtype=None): + del device, dtype + return self + + def clear_cache(self): + self.clear_calls += 1 + self._feat_map = [None] + self._conv_idx = [0] + + def post_quant_conv(self, latents): + return latents + + def decoder(self, x, *, feat_cache, feat_idx, first_chunk=False): + self.decoder_first_chunk_flags.append(first_chunk) + if feat_cache[0] is None: + feat_cache[0] = x.detach().clone() + else: + feat_cache[0] = torch.cat([feat_cache[0], x.detach().clone()], dim=2) + feat_idx[0] += 1 + return x + + class _PipelineConfig: + vae_precision = "fp32" + vae_tiling = False + + def get_decode_scale_and_shift(self, device, dtype, vae): + del device, dtype, vae + return 1.0, None + + def preprocess_decoding(self, latents, server_args, vae=None): + del server_args, vae + return latents + + def post_decoding(self, frames, server_args): + del server_args + return frames + + monkeypatch.setattr( + realtime_vae, + "get_local_torch_device", + lambda: torch.device("cpu"), + ) + + vae = _WanVAE() + vae.clear_cache() + vae.clear_calls = 0 + stage = CausalVaeDecodingStage.__new__(CausalVaeDecodingStage) + stage.vae = vae + server_args = SimpleNamespace( + pipeline_config=_PipelineConfig(), + disable_autocast=True, + ) + + first = stage.decode_causal( + torch.zeros(1, 1, 2, 1, 1), + server_args, + first_chunk=True, + ) + second = stage.decode_causal( + torch.ones(1, 1, 1, 1, 1), + server_args, + first_chunk=False, + ) + + assert tuple(first.shape) == (1, 1, 2, 1, 1) + assert tuple(second.shape) == (1, 1, 1, 1, 1) + assert vae.clear_calls == 0 + assert vae.decoder_first_chunk_flags == [True, False, False] + assert tuple(vae._feat_map[0].shape) == (1, 1, 3, 1, 1) + + +def test_causal_vae_decoding_stage_prefers_native_causal_decode(monkeypatch): + from sglang.multimodal_gen.runtime.pipelines_core.stages import realtime_vae + + class _NativeCausalVAE: + def __init__(self): + self.config = SimpleNamespace(patch_size=None) + self.calls = [] + self._feat_map = [None] + self._conv_idx = [0] + + def to(self, device=None, dtype=None): + del device, dtype + return self + + def clear_cache(self): + self.calls.append("clear_cache") + + def reset_causal_decode_state(self): + self.calls.append("reset") + + def post_quant_conv(self, latents): + self.calls.append("post_quant_conv") + return latents + + def decoder(self, x, *, feat_cache, feat_idx, first_chunk=False): + del x, feat_cache, feat_idx, first_chunk + self.calls.append("decoder") + + def causal_decode(self, latents): + self.calls.append("causal_decode") + return latents + + class _PipelineConfig: + vae_precision = "fp32" + vae_tiling = False + + def get_decode_scale_and_shift(self, device, dtype, vae): + del device, dtype, vae + return 1.0, None + + def preprocess_decoding(self, latents, server_args, vae=None): + del server_args, vae + return latents + + monkeypatch.setattr( + realtime_vae, + "get_local_torch_device", + lambda: torch.device("cpu"), + ) + + vae = _NativeCausalVAE() + stage = CausalVaeDecodingStage.__new__(CausalVaeDecodingStage) + stage.vae = vae + server_args = SimpleNamespace( + pipeline_config=_PipelineConfig(), + disable_autocast=True, + ) + + frames = stage.decode_causal( + torch.zeros(1, 1, 1, 1, 1), + server_args, + first_chunk=True, + ) + + assert tuple(frames.shape) == (1, 1, 1, 1, 1) + assert vae.calls == ["causal_decode"] diff --git a/python/sglang/multimodal_gen/test/unit/test_openai_utils.py b/python/sglang/multimodal_gen/test/unit/test_openai_utils.py new file mode 100644 index 000000000..9a00e06cd --- /dev/null +++ b/python/sglang/multimodal_gen/test/unit/test_openai_utils.py @@ -0,0 +1,23 @@ +# SPDX-License-Identifier: Apache-2.0 + +import asyncio +import io + +from starlette.datastructures import UploadFile as StarletteUploadFile + +from sglang.multimodal_gen.runtime.entrypoints.openai.utils import ( + _save_upload_to_path, +) + + +def test_save_upload_to_path_accepts_starlette_upload_file(tmp_path): + upload = StarletteUploadFile( + io.BytesIO(b"image-bytes"), + filename="input.png", + ) + target_path = tmp_path / "input.png" + + saved_path = asyncio.run(_save_upload_to_path(upload, str(target_path))) + + assert saved_path == str(target_path) + assert target_path.read_bytes() == b"image-bytes"