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 index e036a1d86..decc09363 100644 --- 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 @@ -18,9 +18,7 @@ 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: @@ -139,8 +137,6 @@ class _TransportPayload: content_type: str payload: bytes metadata: dict[str, int | str | bool | list[int]] - last_raw_rgb_frame: bytes | None = None - last_event_id: int | None = None def _split_frame_batch(frames: list[bytes]) -> list[list[bytes]]: @@ -241,6 +237,10 @@ def _pack_frame_batch_message( return msgspec.msgpack.encode(message) +def _pack_frame_batch_header(header: RealtimeFrameBatchHeader) -> bytes: + return msgspec.msgpack.encode(header) + + def _build_transport_payload( transport_frames: list[bytes], *, @@ -249,8 +249,6 @@ def _build_transport_payload( output_format: str | None, transport_quality: int | None, preview_max_width: int | None, - reference_frame: bytes | None, - event_id: int | None, ) -> _TransportPayload: payload_content_type = content_type payload_metadata: dict[str, int | str | bool | list[int]] = {} @@ -302,35 +300,12 @@ def _build_transport_payload( "height": preview_height, "payload_lengths": [len(frame) for frame in encoded_frames], } - elif ( - output_format == RAW_LOSSLESS_OUTPUT_FORMAT - and content_type == RAW_RGB_CONTENT_TYPE - and transport_frames - ): + elif 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) @@ -457,15 +432,13 @@ async def _build_encoded_preview_payload( class RawRGBRealtimeOutputAdapter: - """send raw RGB over WebSocket using lossless transport compression""" + """send raw RGB over WebSocket using lossless transport""" def __init__(self) -> None: - self._last_raw_rgb_frame: bytes | None = None - self._last_event_id: int | None = None + pass def reset(self) -> None: - self._last_raw_rgb_frame = None - self._last_event_id = None + pass async def send( self, @@ -556,9 +529,6 @@ class RawRGBRealtimeOutputAdapter: 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, @@ -572,8 +542,6 @@ class RawRGBRealtimeOutputAdapter: output_format=output_format, transport_quality=transport_quality, preview_max_width=preview_max_width, - reference_frame=reference_frame, - event_id=event_id, ) else: transport_payload = _build_transport_payload( @@ -583,12 +551,7 @@ class RawRGBRealtimeOutputAdapter: output_format=output_format, transport_quality=transport_quality, preview_max_width=preview_max_width, - 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 = { @@ -608,24 +571,32 @@ class RawRGBRealtimeOutputAdapter: 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, + header_payload = _pack_frame_batch_header(header) + stats["header_pack_ms"] += timer.mark_ms() + + await ws.send_bytes(header_payload) + stats["header_write_ms"] += timer.mark_ms() + + await ws.send_bytes(transport_payload.payload) + stats["raw_write_ms"] += timer.mark_ms() + + stats["ws_payload_bytes"] += len(header_payload) + len( + transport_payload.payload ) else: message_payload = _pack_frame_batch_message( header, transport_payload.payload, ) - stats["header_pack_ms"] += timer.mark_ms() + 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["header_write_ms"] += timer.mark_ms() + await ws.send_bytes(message_payload) + stats["raw_write_ms"] += timer.mark_ms() + + stats["ws_payload_bytes"] += len(message_payload) 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 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 index e7598bfc5..0aade381a 100644 --- 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 @@ -439,10 +439,17 @@ async def _close_realtime_websocket( pass +async def _wait_for_server_warmup(websocket: WebSocket) -> None: + warmup_done = getattr(websocket.app.state, "server_warmup_done", None) + if warmup_done is not None and not warmup_done.is_set(): + await warmup_done.wait() + + @router.websocket("/generate") async def generate(websocket: WebSocket): """endpoint for creating a new realtime session""" await websocket.accept() + await _wait_for_server_warmup(websocket) if _ACTIVE_SESSION_IDS and not await _wait_for_active_session_slot(): logger.warning( "reject realtime session because another session is active: %s", diff --git a/python/sglang/multimodal_gen/runtime/models/dits/lingbot_world.py b/python/sglang/multimodal_gen/runtime/models/dits/lingbot_world.py index 11c17d9e0..6a2bc5b38 100644 --- a/python/sglang/multimodal_gen/runtime/models/dits/lingbot_world.py +++ b/python/sglang/multimodal_gen/runtime/models/dits/lingbot_world.py @@ -132,8 +132,20 @@ class LingBotWorldCamConditioner(nn.Module): self.cam_scale_layer = nn.Linear(dim, dim) self.cam_shift_layer = nn.Linear(dim, dim) + def compute_scale_shift( + self, c2ws_plucker_emb: torch.Tensor + ) -> tuple[torch.Tensor, torch.Tensor]: + 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 cam_scale, cam_shift + def forward( - self, hidden_states: torch.Tensor, c2ws_plucker_emb: torch.Tensor | None + self, + hidden_states: torch.Tensor, + c2ws_plucker_emb: torch.Tensor | None, + scale_shift: tuple[torch.Tensor, torch.Tensor] | None = None, ) -> torch.Tensor: if c2ws_plucker_emb is None: return hidden_states @@ -142,10 +154,9 @@ class LingBotWorldCamConditioner(nn.Module): "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) + if scale_shift is None: + scale_shift = self.compute_scale_shift(c2ws_plucker_emb) + cam_scale, cam_shift = scale_shift return (1.0 + cam_scale) * hidden_states + cam_shift @@ -918,6 +929,51 @@ class CausalLingBotWorldTransformerBlock(CausalWanTransformerBlock): hidden_states, _ = attn2.to_out(hidden_states) return hidden_states + def _cam_conditioner_scale_shift( + self, + c2ws_plucker_emb: torch.Tensor | None, + ) -> tuple[torch.Tensor, torch.Tensor] | None: + if c2ws_plucker_emb is None: + return None + + forward_context = get_forward_context() + if forward_context.current_timestep < 0: + return self.cam_conditioner.compute_scale_shift(c2ws_plucker_emb) + + forward_batch = forward_context.forward_batch + if not CausalLingBotWorldTransformer3DModel._should_cache_cam_conditioner( + forward_batch + ): + return self.cam_conditioner.compute_scale_shift(c2ws_plucker_emb) + + cache = CausalLingBotWorldTransformer3DModel._get_request_cache( + forward_batch, "lingbot_cam_conditioner" + ) + if cache is None: + return self.cam_conditioner.compute_scale_shift(c2ws_plucker_emb) + + source_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, + c2ws_plucker_emb._version, + ) + if cache.get("source_key") != source_key: + cache.clear() + cache["source_key"] = source_key + cache["entries"] = {} + + entries = cache["entries"] + entry_key = id(self.cam_conditioner) + scale_shift = entries.get(entry_key) + if scale_shift is None: + scale_shift = self.cam_conditioner.compute_scale_shift(c2ws_plucker_emb) + entries[entry_key] = scale_shift + return scale_shift + def forward( self, hidden_states: torch.Tensor, @@ -930,6 +986,7 @@ class CausalLingBotWorldTransformerBlock(CausalWanTransformerBlock): current_start: int = 0, cache_start: int | None = None, c2ws_plucker_emb: torch.Tensor | None = None, + cam_conditioner_scale_shift: tuple[torch.Tensor, torch.Tensor] | None = None, update_cache_only: bool = False, ) -> torch.Tensor: if hidden_states.dim() == 4: @@ -983,7 +1040,13 @@ class CausalLingBotWorldTransformerBlock(CausalWanTransformerBlock): hidden_states, attn_output, gate_msa, residual_zero, residual_zero ) hidden_states = self.cam_conditioner( - hidden_states.to(orig_dtype), c2ws_plucker_emb + hidden_states.to(orig_dtype), + c2ws_plucker_emb, + ( + cam_conditioner_scale_shift + if cam_conditioner_scale_shift is not None + else self._cam_conditioner_scale_shift(c2ws_plucker_emb) + ), ) norm_hidden_states = self.self_attn_residual_norm.norm(hidden_states).to( orig_dtype @@ -1130,6 +1193,14 @@ class CausalLingBotWorldTransformer3DModel(CausalWanTransformer3DModel): return None return extra.setdefault(name, {}) + @staticmethod + def _should_cache_cam_conditioner(forward_batch) -> bool: + return ( + forward_batch is not None + and getattr(forward_batch, "enable_sequence_shard", False) + and get_ulysses_parallel_world_size() > 1 + ) + @staticmethod def _all_crossattn_caches_initialized( crossattn_cache: list[CrossAttentionKVCache] | None, @@ -1285,6 +1356,52 @@ class CausalLingBotWorldTransformer3DModel(CausalWanTransformer3DModel): cache[cache_key] = (temb, timestep_proj) return temb, timestep_proj + def _prepare_cam_conditioner_scale_shifts( + self, + c2ws_plucker_emb: torch.Tensor | None, + forward_batch, + ) -> list[tuple[torch.Tensor, torch.Tensor]] | None: + if c2ws_plucker_emb is None: + return None + + forward_context = get_forward_context() + if forward_context.current_timestep < 0: + return None + + if not self._should_cache_cam_conditioner(forward_batch): + return None + + cache = self._get_request_cache(forward_batch, "lingbot_cam_conditioner") + if cache is None: + return None + + source_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, + c2ws_plucker_emb._version, + ) + if cache.get("source_key") != source_key: + cache.clear() + cache["source_key"] = source_key + cache["entries"] = {} + + entries = cache["entries"] + scale_shifts = [] + for block in self.blocks: + entry_key = id(block.cam_conditioner) + scale_shift = entries.get(entry_key) + if scale_shift is None: + scale_shift = block.cam_conditioner.compute_scale_shift( + c2ws_plucker_emb + ) + entries[entry_key] = scale_shift + scale_shifts.append(scale_shift) + return scale_shifts + def forward( self, hidden_states: torch.Tensor, @@ -1392,6 +1509,9 @@ class CausalLingBotWorldTransformer3DModel(CausalWanTransformer3DModel): if current_platform.is_mps() else encoder_hidden_states ) + cam_conditioner_scale_shifts = self._prepare_cam_conditioner_scale_shifts( + c2ws_plucker_emb, forward_batch + ) for block_index, block in enumerate(self.blocks): hidden_states = block( @@ -1405,6 +1525,11 @@ class CausalLingBotWorldTransformer3DModel(CausalWanTransformer3DModel): current_start=current_start, cache_start=cache_start, c2ws_plucker_emb=c2ws_plucker_emb, + cam_conditioner_scale_shift=( + None + if cam_conditioner_scale_shifts is None + else cam_conditioner_scale_shifts[block_index] + ), update_cache_only=skip_final_projection and block_index == len(self.blocks) - 1, ) 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 index 64b5fab41..263166af3 100644 --- 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 @@ -303,6 +303,7 @@ class LingBotWorldCausalDMDDenoisingStage(CausalDMDDenoisingStage): prepare_model_input=prepare_model_input, prepare_context_input=prepare_model_input, ) + cache_ctx.cache_state.runtime_cache.pop("lingbot_cam_conditioner", None) # Advance cumulative frame position self._advance_realtime_causal_cache(cache_ctx, num_frames=ctx.num_frames) diff --git a/python/sglang/multimodal_gen/test/server/gpu_cases.py b/python/sglang/multimodal_gen/test/server/gpu_cases.py index 9f622374f..879b7ccc7 100644 --- a/python/sglang/multimodal_gen/test/server/gpu_cases.py +++ b/python/sglang/multimodal_gen/test/server/gpu_cases.py @@ -430,7 +430,9 @@ ONE_GPU_CASES: list[DiffusionTestCase] = [ model_path="robbyant/lingbot-world-fast-diffusers", modality="video", num_gpus=1, - extras=["--pipeline-class-name LingBotWorldCausalDMDPipeline"], + extras=[ + "--pipeline-class-name LingBotWorldCausalDMDPipeline --warmup false" + ], text_encoder_cpu_offload=True, ), LINGBOT_WORLD_REALTIME_sampling_params, 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 index d02485192..71441094b 100644 --- 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 @@ -7,8 +7,13 @@ import torch from sglang.multimodal_gen.runtime.layers.kvcache.causal_attention_cache import ( CrossAttentionKVCache, ) +from sglang.multimodal_gen.runtime.models.dits import ( + lingbot_world as lingbot_world_module, +) from sglang.multimodal_gen.runtime.models.dits.lingbot_world import ( CausalLingBotWorldTransformer3DModel, + CausalLingBotWorldTransformerBlock, + LingBotWorldCamConditioner, ) from sglang.multimodal_gen.runtime.pipelines_core.stages.model_specific_stages.lingbot_world import ( LingBotWorldCausalDMDDenoisingStage, @@ -196,7 +201,204 @@ def test_lingbot_i2v_model_input_writer_reuses_buffer(): assert torch.equal(second[:, 16:], condition) -def test_lingbot_condition_embedding_skips_text_when_crossattn_cache_ready(): +def test_lingbot_cam_conditioner_scale_shift_matches_forward(): + hidden_states = torch.randn(2, 3, 4) + c2ws_plucker_emb = torch.randn(2, 3, 4) + conditioner = SimpleNamespace( + compute_scale_shift=lambda c2ws: (c2ws * 0.25, c2ws - 0.5) + ) + + expected = LingBotWorldCamConditioner.forward( + conditioner, + hidden_states, + c2ws_plucker_emb, + ) + actual = LingBotWorldCamConditioner.forward( + conditioner, + hidden_states, + c2ws_plucker_emb, + conditioner.compute_scale_shift(c2ws_plucker_emb), + ) + + assert torch.allclose(actual, expected) + + +def test_lingbot_cam_conditioner_cache_reuses_source_tensor(monkeypatch): + class _CamConditioner: + def __init__(self): + self.calls = 0 + + def compute_scale_shift(self, c2ws_plucker_emb): + self.calls += 1 + return c2ws_plucker_emb + 1, c2ws_plucker_emb + 2 + + block = CausalLingBotWorldTransformerBlock.__new__( + CausalLingBotWorldTransformerBlock + ) + block.cam_conditioner = _CamConditioner() + forward_batch = SimpleNamespace(extra={}, enable_sequence_shard=True) + monkeypatch.setattr( + lingbot_world_module, "get_ulysses_parallel_world_size", lambda: 2 + ) + monkeypatch.setattr( + lingbot_world_module, + "get_forward_context", + lambda: SimpleNamespace(forward_batch=forward_batch, current_timestep=7), + ) + + c2ws_plucker_emb = torch.ones(1, 2, 3) + first = block._cam_conditioner_scale_shift(c2ws_plucker_emb) + second = block._cam_conditioner_scale_shift(c2ws_plucker_emb) + next_source = c2ws_plucker_emb.clone() + third = block._cam_conditioner_scale_shift(next_source) + + assert first is second + assert third is not first + assert block.cam_conditioner.calls == 2 + cache = forward_batch.extra["lingbot_cam_conditioner"] + assert cache["source_key"][0] == next_source.data_ptr() + assert len(cache["entries"]) == 1 + + +def test_lingbot_cam_conditioner_cache_skips_non_sequence_shard(monkeypatch): + class _CamConditioner: + def __init__(self): + self.calls = 0 + + def compute_scale_shift(self, c2ws_plucker_emb): + self.calls += 1 + return c2ws_plucker_emb + 1, c2ws_plucker_emb + 2 + + block = CausalLingBotWorldTransformerBlock.__new__( + CausalLingBotWorldTransformerBlock + ) + block.cam_conditioner = _CamConditioner() + forward_batch = SimpleNamespace(extra={}, enable_sequence_shard=False) + monkeypatch.setattr( + lingbot_world_module, + "get_forward_context", + lambda: SimpleNamespace(forward_batch=forward_batch, current_timestep=7), + ) + + c2ws_plucker_emb = torch.ones(1, 2, 3) + first = block._cam_conditioner_scale_shift(c2ws_plucker_emb) + second = block._cam_conditioner_scale_shift(c2ws_plucker_emb) + + assert first is not second + assert first[0] is not second[0] + assert block.cam_conditioner.calls == 2 + assert "lingbot_cam_conditioner" not in forward_batch.extra + + +def test_lingbot_cam_conditioner_cache_skips_single_ulysses_world(monkeypatch): + class _CamConditioner: + def __init__(self): + self.calls = 0 + + def compute_scale_shift(self, c2ws_plucker_emb): + self.calls += 1 + return c2ws_plucker_emb + 1, c2ws_plucker_emb + 2 + + block = CausalLingBotWorldTransformerBlock.__new__( + CausalLingBotWorldTransformerBlock + ) + block.cam_conditioner = _CamConditioner() + forward_batch = SimpleNamespace(extra={}, enable_sequence_shard=True) + monkeypatch.setattr( + lingbot_world_module, "get_ulysses_parallel_world_size", lambda: 1 + ) + monkeypatch.setattr( + lingbot_world_module, + "get_forward_context", + lambda: SimpleNamespace(forward_batch=forward_batch, current_timestep=7), + ) + + c2ws_plucker_emb = torch.ones(1, 2, 3) + first = block._cam_conditioner_scale_shift(c2ws_plucker_emb) + second = block._cam_conditioner_scale_shift(c2ws_plucker_emb) + + assert first is not second + assert block.cam_conditioner.calls == 2 + assert "lingbot_cam_conditioner" not in forward_batch.extra + + +def test_lingbot_cam_conditioner_cache_skips_context_update(monkeypatch): + class _CamConditioner: + def __init__(self): + self.calls = 0 + + def compute_scale_shift(self, c2ws_plucker_emb): + self.calls += 1 + return c2ws_plucker_emb + 1, c2ws_plucker_emb + 2 + + block = CausalLingBotWorldTransformerBlock.__new__( + CausalLingBotWorldTransformerBlock + ) + block.cam_conditioner = _CamConditioner() + forward_batch = SimpleNamespace(extra={}, enable_sequence_shard=True) + monkeypatch.setattr( + lingbot_world_module, + "get_forward_context", + lambda: SimpleNamespace(forward_batch=forward_batch, current_timestep=-1), + ) + + c2ws_plucker_emb = torch.ones(1, 2, 3) + first = block._cam_conditioner_scale_shift(c2ws_plucker_emb) + second = block._cam_conditioner_scale_shift(c2ws_plucker_emb) + + assert first is not second + assert block.cam_conditioner.calls == 2 + assert "lingbot_cam_conditioner" not in forward_batch.extra + + +def test_lingbot_model_prepares_cam_conditioner_scale_shifts(monkeypatch): + class _CamConditioner: + def __init__(self, offset): + self.offset = offset + self.calls = 0 + + def compute_scale_shift(self, c2ws_plucker_emb): + self.calls += 1 + return ( + c2ws_plucker_emb + self.offset, + c2ws_plucker_emb + self.offset + 10, + ) + + model = CausalLingBotWorldTransformer3DModel.__new__( + CausalLingBotWorldTransformer3DModel + ) + model.blocks = [ + SimpleNamespace(cam_conditioner=_CamConditioner(1)), + SimpleNamespace(cam_conditioner=_CamConditioner(2)), + ] + forward_batch = SimpleNamespace(extra={}, enable_sequence_shard=True) + monkeypatch.setattr( + lingbot_world_module, "get_ulysses_parallel_world_size", lambda: 2 + ) + monkeypatch.setattr( + lingbot_world_module, + "get_forward_context", + lambda: SimpleNamespace(forward_batch=forward_batch, current_timestep=7), + ) + + c2ws_plucker_emb = torch.ones(1, 2, 3) + first = model._prepare_cam_conditioner_scale_shifts(c2ws_plucker_emb, forward_batch) + second = model._prepare_cam_conditioner_scale_shifts( + c2ws_plucker_emb, forward_batch + ) + next_source = c2ws_plucker_emb.clone() + third = model._prepare_cam_conditioner_scale_shifts(next_source, forward_batch) + + assert first is not None + assert second is not None + assert third is not None + assert first[0] is second[0] + assert first[1] is second[1] + assert third[0] is not first[0] + assert [block.cam_conditioner.calls for block in model.blocks] == [2, 2] + + +def test_lingbot_condition_embedding_skips_text_when_crossattn_cache_ready(monkeypatch): class _ConditionEmbedder: def __init__(self): self.full_calls = 0 @@ -219,6 +421,14 @@ def test_lingbot_condition_embedding_skips_text_when_crossattn_cache_ready(): CausalLingBotWorldTransformer3DModel ) model.condition_embedder = _ConditionEmbedder() + monkeypatch.setattr( + lingbot_world_module, + "get_forward_context", + lambda: SimpleNamespace( + forward_batch=SimpleNamespace(extra={}), + current_timestep=7, + ), + ) crossattn_cache = [ CrossAttentionKVCache( k=torch.empty(1, 1, 1, 1), 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 index 1c9dd982f..0a36c2009 100644 --- 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 @@ -17,7 +17,6 @@ from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import OutputBa 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, @@ -27,10 +26,15 @@ from sglang.multimodal_gen.runtime.utils.realtime_video import ( def _unpack_frame_batch_messages(payloads): messages = [] - for payload in payloads: + payload_iter = iter(payloads) + for payload in payload_iter: message = msgspec.msgpack.decode(payload) - assert message.pop("type") == "frame_batch" - frame_payload = message.pop("payload") + message_type = message.pop("type") + if message_type == "frame_batch": + frame_payload = message.pop("payload") + else: + assert message_type == "frame_batch_header" + frame_payload = next(payload_iter) messages.append((message, frame_payload)) return messages @@ -152,7 +156,7 @@ def test_delta_gzip_raw_rgb_payload_roundtrips_exactly(): assert restored == b"".join(frames) -def test_raw_rgb_realtime_output_adapter_uses_lossless_compressed_payload(): +def test_raw_rgb_realtime_output_adapter_uses_lossless_raw_payload_by_default(): class _WebSocket: def __init__(self): self.payloads = [] @@ -191,8 +195,8 @@ def test_raw_rgb_realtime_output_adapter_uses_lossless_compressed_payload(): 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["content_type"] == RAW_RGB_CONTENT_TYPE + assert first_header["encoding"] == "raw" assert first_header["event_id"] == 3 assert first_header["format"] == "rgb24" assert first_header["channels"] == 3 @@ -206,15 +210,10 @@ def test_raw_rgb_realtime_output_adapter_uses_lossless_compressed_payload(): 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 + assert first_payload == expected_frames -def test_raw_rgb_realtime_output_adapter_offloads_delta_payload_build( +def test_raw_rgb_realtime_output_adapter_offloads_default_lossless_payload_build( monkeypatch, ): calls = [] @@ -243,7 +242,7 @@ def test_raw_rgb_realtime_output_adapter_offloads_delta_payload_build( frame1 = bytes([1, 2, 4]) * 1000 batch = SimpleNamespace( block_idx=0, - request_id="req-offload-delta", + request_id="req-offload-raw", width=1000, height=1, enable_upscaling=False, @@ -270,14 +269,9 @@ def test_raw_rgb_realtime_output_adapter_offloads_delta_payload_build( realtime_output_adapter._build_transport_payload, ] [(first_header, first_payload)] = _unpack_frame_batch_messages(payloads) - assert first_header["encoding"] == "delta-gzip" + assert first_header["encoding"] == "raw" 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 + assert first_payload == expected_frames def test_raw_rgb_realtime_output_adapter_can_send_uncompressed_raw_frames(): @@ -334,7 +328,7 @@ def test_raw_rgb_realtime_output_adapter_can_send_uncompressed_raw_frames(): assert stats["ws_payload_bytes"] == sum(len(payload) for payload in payloads) -def test_raw_rgb_realtime_output_adapter_uses_previous_frame_reference(): +def test_raw_rgb_realtime_output_adapter_does_not_require_previous_frame_reference(): class _WebSocket: def __init__(self): self.payloads = [] @@ -388,21 +382,12 @@ def test_raw_rgb_realtime_output_adapter_uses_previous_frame_reference(): (first_header, first_payload), (second_header, second_payload) = ( _unpack_frame_batch_messages(payloads) ) + assert first_header["content_type"] == RAW_RGB_CONTENT_TYPE + assert second_header["content_type"] == RAW_RGB_CONTENT_TYPE 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]) + assert "delta_reference" not in second_header + assert first_payload == bytes([1, 2, 3, 4, 5, 6]) + assert second_payload == bytes([1, 2, 4, 4, 6, 6]) def test_raw_rgb_realtime_output_adapter_splits_large_frame_batches(): @@ -450,11 +435,61 @@ def test_raw_rgb_realtime_output_adapter_splits_large_frame_batches(): 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 "delta_reference" not in headers[1] assert stats["num_batches"] == 2 assert stats["num_frames"] == 17 +def test_raw_rgb_realtime_output_adapter_sends_large_payload_separately(): + class _WebSocket: + def __init__(self): + self.payloads = [] + + async def send_bytes(self, payload): + self.payloads.append(payload) + + async def run(): + ws = _WebSocket() + adapter = RawRGBRealtimeOutputAdapter() + frame = bytes([7]) * (72 * 1024) + batch = SimpleNamespace( + block_idx=0, + request_id="req-large", + width=len(frame) // 3, + height=1, + enable_upscaling=False, + realtime_event_id=3, + ) + result = OutputBatch( + raw_frame_batches=[[frame]], + raw_frame_content_type=RAW_RGB_CONTENT_TYPE, + raw_frame_metadata={ + "format": "rgb24", + "width": len(frame) // 3, + "height": 1, + "channels": 3, + "bytes_per_frame": len(frame), + }, + ) + + stats = await adapter.send(ws, SimpleNamespace(), result, batch) + return ws.payloads, stats + + payloads, stats = asyncio.run(run()) + + assert len(payloads) == 2 + header = msgspec.msgpack.decode(payloads[0]) + assert header["type"] == "frame_batch_header" + assert "payload" not in header + assert header["content_type"] == RAW_RGB_CONTENT_TYPE + assert header["total_size"] == len(payloads[1]) + assert payloads[1] == bytes([7]) * (72 * 1024) + assert stats["raw_bytes"] == len(payloads[1]) + assert stats["ws_payload_bytes"] == len(payloads[0]) + len(payloads[1]) + assert stats["num_batches"] == 1 + assert stats["num_frames"] == 1 + + def test_raw_rgb_realtime_output_adapter_can_send_webp_preview_frames(): class _WebSocket: def __init__(self): @@ -560,26 +595,18 @@ def test_raw_rgb_realtime_output_adapter_offloads_preview_encoding(monkeypatch): payloads = asyncio.run(run()) assert [call[0] for call in calls] == [ - realtime_output_adapter._build_transport_payload, - realtime_output_adapter._build_transport_payload, + realtime_output_adapter._encode_rgb_frame_to_webp, + realtime_output_adapter._encode_rgb_frame_to_webp, ] - (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") + [(header, payload)] = _unpack_frame_batch_messages(payloads) + assert header["content_type"] == WEBP_FRAME_CONTENT_TYPE + assert header["encoding"] == "webp" + assert header["num_frames"] == 2 + assert header["frame_batch_index"] == 0 + assert header["num_frame_batches"] == 1 + assert header["is_final_frame_batch"] is True + assert len(header["payload_lengths"]) == 2 + assert payload.startswith(b"RIFF") def test_raw_rgb_realtime_output_adapter_can_send_jpeg_preview_frames():