Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
562 changes: 562 additions & 0 deletions flashdreams/flashdreams/serving/webrtc/encoders.py

Large diffs are not rendered by default.

248 changes: 238 additions & 10 deletions flashdreams/flashdreams/serving/webrtc/manager.py
Original file line number Diff line number Diff line change
Expand Up @@ -9,18 +9,29 @@
import contextlib
import inspect
import json
import math
from collections import deque
from collections.abc import Set as AbstractSet
from dataclasses import dataclass, field
from enum import IntEnum
from typing import Any

import torch
from aiortc import RTCConfiguration, RTCPeerConnection, RTCSessionDescription
from aiortc import (
RTCConfiguration,
RTCPeerConnection,
RTCRtpSender,
RTCSessionDescription,
)
from aiortc.sdp import SessionDescription
from loguru import logger

from flashdreams.serving.realtime.input import KeyboardResampler
from flashdreams.serving.webrtc.media import BufferedVideoTrack
from flashdreams.serving.webrtc.encoders import (
DefaultRTCEncoder,
VideoEncoder,
)
from flashdreams.serving.webrtc.media import BufferedVideoTrack, NVENCVideoTrack
from flashdreams.serving.webrtc.server import SessionBusyError
from flashdreams.serving.webrtc.warmup import (
run_loopback_warmup_session,
Expand All @@ -35,6 +46,45 @@
_CLIENT_LIVENESS_CHECK_INTERVAL_S = 1.0


def _performance_stats_payload(stats: dict[str, float] | None) -> dict[str, float]:
if not stats:
return {}
payload: dict[str, float] = {}
for key, value in stats.items():
if not isinstance(key, str):
continue
try:
number = float(value)
except (TypeError, ValueError):
continue
if math.isfinite(number):
payload[key] = round(number, 1)
return payload


def _format_performance_stats(stats: dict[str, float]) -> str:
return " ".join(f"{key}={value:.1f}" for key, value in sorted(stats.items()))


def _sdp_video_codecs(sdp: str) -> tuple[str, ...] | None:
"""Return offered video codec MIME types, or ``None`` if parsing fails."""
try:
description = SessionDescription.parse(sdp)
except Exception:
logger.debug("Could not parse remote SDP while checking video codecs.")
return None

codecs: list[str] = []
for media in description.media:
if getattr(media, "kind", None) != "video":
continue
for codec in media.rtp.codecs:
mime_type = str(getattr(codec, "mimeType", "")).strip()
if mime_type:
codecs.append(mime_type)
return tuple(dict.fromkeys(codecs))


class WebRTCControlSignal(IntEnum):
"""Rank-orchestration signals shared by the single-session runtimes."""

Expand All @@ -61,7 +111,8 @@ class ManagedWebRTCSession:
"""Per-session state for the single active WebRTC peer connection."""

runtime: Any
video_track: BufferedVideoTrack
video_track: BufferedVideoTrack | NVENCVideoTrack
video_encoder: VideoEncoder
peer_connection: Any
resampler: KeyboardResampler
control_channel: Any | None = None
Expand Down Expand Up @@ -158,6 +209,108 @@ def _make_resampler(self, *, start_v: float) -> KeyboardResampler:
def _register_extra_peer_handlers(self, peer_connection: Any) -> None:
"""Register optional extra peer-connection event handlers."""

def _prefer_h264_video_codec(self, *, transceiver: Any) -> None:
"""Constrain the transceiver's codec preferences to H.264 variants.

Required when the selected encoder emits pre-encoded H.264 packets
(``av.Packet`` route through ``H264Encoder.pack()``): if the SDP
negotiates VP8/VP9 instead, aiortc will pack the H.264 bitstream
under the wrong codec header and the receiver will fail to decode.

If the local aiortc build does not advertise H.264, no preference
is set; the SDP-time fallback in ``_enforce_h264_or_fallback``
will then swap the encoder to :class:`DefaultRTCEncoder`.
"""
caps = RTCRtpSender.getCapabilities("video")
h264_codecs = [c for c in caps.codecs if c.mimeType.lower() == "video/h264"]
if not h264_codecs:
return
transceiver.setCodecPreferences(h264_codecs)

def _prepare_video_encoder_for_offer(
self,
*,
video_encoder: VideoEncoder,
offer_sdp: str,
) -> VideoEncoder:
"""Select a session encoder compatible with the browser offer."""
if video_encoder.prefers_codec != "h264":
return video_encoder

offered_video_codecs = _sdp_video_codecs(offer_sdp)
if offered_video_codecs is None:
return video_encoder
if any(codec.lower() == "video/h264" for codec in offered_video_codecs):
return video_encoder

offered = ", ".join(offered_video_codecs) or "<none>"
if getattr(self.runtime_config, "encoder_backend", None) == "nvenc":
video_encoder.close()
raise RuntimeError(
"encoder_backend='nvenc' requested but the browser offer does "
f"not advertise H.264; offered video codecs: {offered}."
)

logger.warning(
"Hardware encoder emits H.264, but the browser offer does not "
"advertise H.264 (offered: {}). Using aiortc software encoder "
"for this session.",
offered,
)
video_encoder.close()
return DefaultRTCEncoder(fps=self.fps)

def _enforce_h264_or_fallback(
self,
*,
transceiver: Any,
managed_session: ManagedWebRTCSession,
num_frames: int,
) -> None:
"""Verify H.264 was negotiated; swap to the software encoder if not.

aiortc exposes the negotiated codec set on
``RTCRtpTransceiver._codecs`` after ``setLocalDescription``. We
read it via that attribute (aiortc-internal, but stable in the
pinned version) and, if H.264 did not land, close the hardware
encoder and install a :class:`DefaultRTCEncoder` with a
:class:`BufferedVideoTrack` on the same sender before the first
RTP packet flies. ``replaceTrack`` does not renegotiate; aiortc's
RTP loop will encode raw ``av.VideoFrame`` output with whatever
codec (VP8/VP9/H.264) actually landed in the SDP.
"""
negotiated = getattr(transceiver, "_codecs", None) or []
if negotiated and negotiated[0].mimeType.lower() == "video/h264":
logger.info(
"Video codec negotiated: {} (hardware encoder path active).",
negotiated[0].mimeType,
)
return

chosen = negotiated[0].mimeType if negotiated else "<none>"
if getattr(self.runtime_config, "encoder_backend", None) == "nvenc":
managed_session.video_encoder.close()
raise RuntimeError(
"encoder_backend='nvenc' requested but SDP negotiation landed on "
f"{chosen!r}; cannot stream pre-encoded H.264 packets."
)
logger.warning(
"H.264 preferred by hardware encoder but SDP negotiation "
"landed on {!r}; swapping to the software encoder before "
"streaming begins.",
chosen,
)
# Close the hardware encoder so its NVENC session is released
# promptly; the software adapter has no hardware resources to
# release itself.
managed_session.video_encoder.close()

fallback_encoder = DefaultRTCEncoder(fps=self.fps)
fallback_track = fallback_encoder.create_track(maxsize=num_frames)
transceiver.sender.replaceTrack(fallback_track)
managed_session.video_encoder = fallback_encoder
managed_session.video_track = fallback_track

def _on_offer_received(self, offer_sdp: str) -> None:
"""Hook invoked with the remote offer SDP before negotiation."""

Expand Down Expand Up @@ -283,8 +436,20 @@ async def _create_answer_with_runtime_ready_locked(
# frames than steady state; sizing to it would force a per-chunk
# stall, so we size to the steady-state count.
num_frames = self._runtime.peek_steady_chunk_num_frames()
video_track = BufferedVideoTrack(fps=self.fps, maxsize=num_frames)
peer_connection.addTrack(video_track)
video_encoder: VideoEncoder = self._prepare_video_encoder_for_offer(
video_encoder=self._runtime.video_encoder,
offer_sdp=offer_sdp,
)
video_track = video_encoder.create_track(maxsize=num_frames)
# Use ``addTransceiver`` (not ``addTrack``) so we can constrain the
# SDP m-line's codec list via ``setCodecPreferences`` when the
# encoder emits pre-encoded H.264 packets.
video_transceiver = peer_connection.addTransceiver(
video_track,
direction="sendonly",
)
if video_encoder.prefers_codec == "h264":
self._prefer_h264_video_codec(transceiver=video_transceiver)
# Start the resampler's virtual clock at 0; the real anchor is set
# in the ``on_datachannel`` handler so chunk 0's window starts when
# input can actually arrive.
Expand All @@ -293,6 +458,7 @@ async def _create_answer_with_runtime_ready_locked(
managed_session = ManagedWebRTCSession(
runtime=self._runtime,
video_track=video_track,
video_encoder=video_encoder,
peer_connection=peer_connection,
resampler=resampler,
last_client_message_at=loop.time(),
Expand Down Expand Up @@ -350,6 +516,12 @@ async def on_connectionstatechange() -> None:
answer = await peer_connection.createAnswer()
await peer_connection.setLocalDescription(answer)
await wait_for_ice_gathering_complete(peer_connection)
if video_encoder.prefers_codec == "h264":
self._enforce_h264_or_fallback(
transceiver=video_transceiver,
managed_session=managed_session,
num_frames=num_frames,
)
local_description = peer_connection.localDescription
if local_description is None:
raise RuntimeError("Peer connection did not produce local description.")
Expand Down Expand Up @@ -548,6 +720,7 @@ async def _generation_worker(
runtime = managed_session.runtime
resampler = managed_session.resampler
video_track = managed_session.video_track
video_encoder = managed_session.video_encoder

# Stay idle until the user interacts. Generating eagerly would burn
# GPU cycles on a still scene the viewer never sees. Once an event
Expand Down Expand Up @@ -603,6 +776,7 @@ async def _generation_worker(
consumed_action_arrivals.append(
managed_session.pending_action_arrivals.popleft()
)
t_after_sample = loop.time()
try:
result = await runtime.generate_chunk(
segments=segments, frame_times=frame_times
Expand All @@ -617,31 +791,74 @@ async def _generation_worker(
return
continue
t_after_gen = loop.time()
enqueued = await video_track.enqueue_chunk(result.video_chunk)
delivery = await video_encoder.deliver_chunk(
result.video_chunk,
video_track,
force_keyframe=False,
)
enqueued = delivery.num_frames
t_after_enqueue = loop.time()

gen_ms = (t_after_gen - t_before_gen) * 1e3
sample_ms = (t_after_sample - t_before_gen) * 1e3
runtime_call_ms = (t_after_gen - t_after_sample) * 1e3
enqueue_ms = (t_after_enqueue - t_after_gen) * 1e3
chunk_total_ms = (t_after_enqueue - t_before_gen) * 1e3
chunk_fps = (
result.num_frames * 1000.0 / chunk_total_ms
if chunk_total_ms > 0
else 0.0
)
play_ms = result.num_frames * 1000.0 / video_track.fps
lag_ms = (t_after_enqueue - resampler.next_chunk_start_v) * 1e3
control_latency_ms = (
(t_after_enqueue - consumed_action_arrivals[0]) * 1e3
if consumed_action_arrivals
else None
)
logger.debug(
"Chunk done chunk={} num_frames={} segments={} enqueued={} "
"gen_ms={:.1f} enqueue_ms={:.1f} play_ms={:.1f} queue_depth={} "
"lag_ms={:.1f}",
track_dropped_packets = getattr(
video_track,
"dropped_packets",
None,
)
runtime_stats = _performance_stats_payload(result.stats)
log_stats = {
**runtime_stats,
"delivery_encode_ms": round(delivery.encode_ms, 1),
}
if isinstance(track_dropped_packets, int):
log_stats["track_dropped_packets"] = float(track_dropped_packets)
extra_stats = _format_performance_stats(log_stats)
log_level = (
logger.info
if (
result.chunk_index <= 2
or result.chunk_index % 10 == 0
or chunk_total_ms > play_ms
or lag_ms > play_ms
)
else logger.debug
)
log_level(
"WebRTC chunk done chunk={} num_frames={} segments={} "
"enqueued={} encoder={} gen_ms={:.1f} sample_ms={:.1f} "
"runtime_call_ms={:.1f} delivery_ms={:.1f} "
"play_ms={:.1f} chunk_fps={:.1f} queue_depth={} "
"lag_ms={:.1f} {}",
result.chunk_index,
result.num_frames,
len(segments),
enqueued,
delivery.backend,
gen_ms,
sample_ms,
runtime_call_ms,
enqueue_ms,
play_ms,
chunk_fps,
video_track.qsize(),
lag_ms,
extra_stats,
)

channel = managed_session.control_channel
Expand All @@ -658,11 +875,22 @@ async def _generation_worker(
},
"model": self._model_name(),
"gen_ms": round(gen_ms, 1),
"sample_ms": round(sample_ms, 1),
"runtime_call_ms": round(runtime_call_ms, 1),
"enqueue_ms": round(enqueue_ms, 1),
"delivery_ms": round(enqueue_ms, 1),
"delivery_encode_ms": round(delivery.encode_ms, 1),
"encoder_backend": delivery.backend,
"keyframes": delivery.num_keyframes,
"chunk_total_ms": round(chunk_total_ms, 1),
"chunk_fps": round(chunk_fps, 1),
"play_ms": round(play_ms, 1),
"queue_depth": video_track.qsize(),
"lag_ms": round(lag_ms, 1),
}
if isinstance(track_dropped_packets, int):
payload["track_dropped_packets"] = track_dropped_packets
payload.update(runtime_stats)
payload.update(self._chunk_done_extra())
if control_latency_ms is not None:
payload["latency_ms"] = round(control_latency_ms, 1)
Expand Down
Loading
Loading