[diffusion] optimize: optimize LingBot realtime transport and camera conditioning (#27297)
This commit is contained in:
+25
-54
@@ -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,
|
||||
)
|
||||
|
||||
+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():
|
||||
|
||||
Reference in New Issue
Block a user