[diffusion] optimize: optimize LingBot realtime transport and camera conditioning (#27297)

This commit is contained in:
Mick
2026-06-05 16:00:48 +08:00
committed by GitHub
parent 00fefef16b
commit 4ef081b903
7 changed files with 462 additions and 119 deletions
@@ -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
@@ -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",
@@ -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,
)
@@ -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)
@@ -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,
@@ -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),
@@ -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():