fix(perception): isolate recorded decode from realtime admission

This commit is contained in:
DCCONSTRUCTIONS
2026-08-25 21:58:12 +03:00
parent bab0bb2df2
commit fb5bf943c9
4 changed files with 416 additions and 15 deletions
@@ -47,14 +47,14 @@ from k1link.perception.m48s_reference_graph_runtime import (
from k1link.perception.motion import ClassIndependentMotionEstimator from k1link.perception.motion import ClassIndependentMotionEstimator
from k1link.perception.object_understanding import AdvisoryResponse from k1link.perception.object_understanding import AdvisoryResponse
from k1link.perception.providers import SourcePacket from k1link.perception.providers import SourcePacket
from k1link.perception.recorded_source import DecodedFrameTiming from k1link.perception.recorded_source import DecodedFrameTiming, SourcePacingTiming
from k1link.perception.reference_graph_runtime import ReferenceGraphRuntimePaths from k1link.perception.reference_graph_runtime import ReferenceGraphRuntimePaths
from k1link.perception.rolling_map import RollingLocalObstacleMapProvider from k1link.perception.rolling_map import RollingLocalObstacleMapProvider
from k1link.perception.temporal import BoundedSpatialTemporalProvider from k1link.perception.temporal import BoundedSpatialTemporalProvider
SCHEMA_VERSION: Final = "missioncore.m48s-reference-graph-shadow-load/v4" SCHEMA_VERSION: Final = "missioncore.m48s-reference-graph-shadow-load/v5"
FRAME_EVIDENCE_SCHEMA: Final = "missioncore.m48s-reference-graph-frame-evidence/v1" FRAME_EVIDENCE_SCHEMA: Final = "missioncore.m48s-reference-graph-frame-evidence/v1"
PIPELINE_TIMING_SCHEMA: Final = "missioncore.m48s-frame-pipeline-timing/v0" PIPELINE_TIMING_SCHEMA: Final = "missioncore.m48s-frame-pipeline-timing/v1"
GC_POLICY_SCHEMA: Final = "missioncore.cyclic-gc-hot-loop-policy/v0" GC_POLICY_SCHEMA: Final = "missioncore.cyclic-gc-hot-loop-policy/v0"
AUTHORITY: Final = { AUTHORITY: Final = {
"ground_truth": False, "ground_truth": False,
@@ -228,14 +228,23 @@ class FrameTimingStore:
def __init__(self) -> None: def __init__(self) -> None:
self._lock = threading.Lock() self._lock = threading.Lock()
self._decode: dict[int, int] = {} self._decode: dict[int, DecodedFrameTiming] = {}
self._pacing: dict[int, SourcePacingTiming] = {}
self._all_decode: list[DecodedFrameTiming] = []
self._all_pacing: list[SourcePacingTiming] = []
self._detector: dict[int, DetectorFrameTiming] = {} self._detector: dict[int, DetectorFrameTiming] = {}
self._providers: dict[int, dict[str, int]] = defaultdict(dict) self._providers: dict[int, dict[str, int]] = defaultdict(dict)
self._delivered: list[dict[str, object]] = [] self._delivered: list[dict[str, object]] = []
def observe_decode(self, timing: DecodedFrameTiming) -> None: def observe_decode(self, timing: DecodedFrameTiming) -> None:
with self._lock: with self._lock:
self._decode[timing.sequence] = timing.duration_ns self._decode[timing.sequence] = timing
self._all_decode.append(timing)
def observe_pacing(self, timing: SourcePacingTiming) -> None:
with self._lock:
self._pacing[timing.sequence] = timing
self._all_pacing.append(timing)
def observe_detector(self, timing: DetectorFrameTiming) -> None: def observe_detector(self, timing: DetectorFrameTiming) -> None:
with self._lock: with self._lock:
@@ -257,7 +266,8 @@ class FrameTimingStore:
) -> dict[str, object]: ) -> dict[str, object]:
with self._lock: with self._lock:
try: try:
decode_ns = self._decode.pop(sequence) decode = self._decode.pop(sequence)
pacing = self._pacing.pop(sequence)
detector = self._detector.pop(sequence) detector = self._detector.pop(sequence)
providers = self._providers.pop(sequence) providers = self._providers.pop(sequence)
except KeyError as exc: except KeyError as exc:
@@ -272,13 +282,19 @@ class FrameTimingStore:
document = { document = {
"schema_version": PIPELINE_TIMING_SCHEMA, "schema_version": PIPELINE_TIMING_SCHEMA,
"sequence": sequence, "sequence": sequence,
"decode_duration_ns": decode_ns, "decode_duration_ns": decode.duration_ns,
"decode_phase": decode.phase.value,
"source_pacing": {
"scheduled_monotonic_ns": pacing.scheduled_monotonic_ns,
"emitted_monotonic_ns": pacing.emitted_monotonic_ns,
"lateness_ns": pacing.lateness_ns,
},
"detector": detector.to_dict(), "detector": detector.to_dict(),
"providers": dict(sorted(providers.items())), "providers": dict(sorted(providers.items())),
"graph_admission_to_delivery_ns": admission_to_delivery_ns, "graph_admission_to_delivery_ns": admission_to_delivery_ns,
"graph_attributed_provider_ns": attributed_graph_ns, "graph_attributed_provider_ns": attributed_graph_ns,
"graph_unattributed_ns": unattributed_ns, "graph_unattributed_ns": unattributed_ns,
"decode_to_delivery_processing_ns": decode_ns + admission_to_delivery_ns, "decode_to_delivery_processing_ns": decode.duration_ns + admission_to_delivery_ns,
} }
with self._lock: with self._lock:
self._delivered.append(document) self._delivered.append(document)
@@ -288,6 +304,14 @@ class FrameTimingStore:
with self._lock: with self._lock:
return tuple(self._delivered) return tuple(self._delivered)
def all_decode(self) -> tuple[DecodedFrameTiming, ...]:
with self._lock:
return tuple(self._all_decode)
def all_pacing(self) -> tuple[SourcePacingTiming, ...]:
with self._lock:
return tuple(self._all_pacing)
class TimedProviderProxy: class TimedProviderProxy:
"""Record one provider's actual call duration without another inference pass.""" """Record one provider's actual call duration without another inference pass."""
@@ -420,6 +444,8 @@ def main() -> int:
map_output_ages_ms: list[float] = [] map_output_ages_ms: list[float] = []
all_deliveries: list[DeliveredFrame] = [] all_deliveries: list[DeliveredFrame] = []
all_pipeline_timings: list[dict[str, object]] = [] all_pipeline_timings: list[dict[str, object]] = []
all_decode_timings: list[DecodedFrameTiming] = []
all_pacing_timings: list[SourcePacingTiming] = []
with ( with (
progress.open("x", encoding="utf-8") as progress_stream, progress.open("x", encoding="utf-8") as progress_stream,
frame_ledger.open("x", encoding="utf-8") as frame_ledger_stream, frame_ledger.open("x", encoding="utf-8") as frame_ledger_stream,
@@ -447,6 +473,7 @@ def main() -> int:
timing_store, timing_store,
), ),
decode_timing_observer=timing_store.observe_decode, decode_timing_observer=timing_store.observe_decode,
source_pacing_observer=timing_store.observe_pacing,
detector_timing_observer=timing_store.observe_detector, detector_timing_observer=timing_store.observe_detector,
maximum_frames=arguments.maximum_frames, maximum_frames=arguments.maximum_frames,
source_rate_hz=arguments.source_rate_hz, source_rate_hz=arguments.source_rate_hz,
@@ -469,8 +496,10 @@ def main() -> int:
), ),
) )
detector_warmup = runtime.warm_up_detector() detector_warmup = runtime.warm_up_detector()
source_prefetch = runtime.prepare_source()
gc_policy = CyclicGcHotLoopPolicy() gc_policy = CyclicGcHotLoopPolicy()
with gc_policy: with gc_policy:
runtime.mark_source_admission_started()
loop_started_ns = time.monotonic_ns() loop_started_ns = time.monotonic_ns()
result = runtime.graph.run() result = runtime.graph.run()
loop_completed_ns = time.monotonic_ns() loop_completed_ns = time.monotonic_ns()
@@ -523,6 +552,7 @@ def main() -> int:
setup_seconds=setup_seconds, setup_seconds=setup_seconds,
gc_policy=gc_policy.to_dict(), gc_policy=gc_policy.to_dict(),
detector_warmup=detector_warmup, detector_warmup=detector_warmup,
source_prefetch=asdict(source_prefetch),
) )
loop_documents.append(loop_document) loop_documents.append(loop_document)
completion_ages_ms.extend(value / 1_000_000.0 for value in loop_completion_ages_ns) completion_ages_ms.extend(value / 1_000_000.0 for value in loop_completion_ages_ns)
@@ -531,6 +561,8 @@ def main() -> int:
) )
all_deliveries.extend(result.deliveries) all_deliveries.extend(result.deliveries)
all_pipeline_timings.extend(loop_pipeline_timings) all_pipeline_timings.extend(loop_pipeline_timings)
all_decode_timings.extend(timing_store.all_decode())
all_pacing_timings.extend(timing_store.all_pacing())
frame_ledger_stream.flush() frame_ledger_stream.flush()
progress_row = { progress_row = {
"loop": loop_index + 1, "loop": loop_index + 1,
@@ -608,6 +640,12 @@ def main() -> int:
and cast(dict[str, object], loop["detector_warmup"])["inference_passes"] == 1 and cast(dict[str, object], loop["detector_warmup"])["inference_passes"] == 1
for loop in loop_documents for loop in loop_documents
), ),
"source_prefetch_completed_before_source_admission": all(
cast(dict[str, object], loop["source_prefetch"])["buffered_frames"]
== cast(dict[str, object], loop["source_prefetch"])["ready_frames"]
for loop in loop_documents
),
"source_pacing_attribution_complete": len(all_pacing_timings) == admitted,
"authority_remains_false": all(value is False for value in AUTHORITY.values()), "authority_remains_false": all(value is False for value in AUTHORITY.values()),
} }
operating_target_checks = { operating_target_checks = {
@@ -672,6 +710,8 @@ def main() -> int:
"identity_continuity": identity, "identity_continuity": identity,
"semantic_advisory": semantic, "semantic_advisory": semantic,
"pipeline_timing": _pipeline_timing_metrics(all_pipeline_timings), "pipeline_timing": _pipeline_timing_metrics(all_pipeline_timings),
"source_decode": _source_decode_metrics(all_decode_timings),
"source_pacing": _source_pacing_metrics(all_pacing_timings),
"python_gc": _gc_telemetry_summary(gc_telemetry.events), "python_gc": _gc_telemetry_summary(gc_telemetry.events),
"gpu": _telemetry_summary(gpu.samples), "gpu": _telemetry_summary(gpu.samples),
"process_peak_rss_before_mib": round(rss_before_kib / 1024.0, 6), "process_peak_rss_before_mib": round(rss_before_kib / 1024.0, 6),
@@ -762,6 +802,7 @@ def _loop_document(
setup_seconds: float, setup_seconds: float,
gc_policy: dict[str, object], gc_policy: dict[str, object],
detector_warmup: DetectorWarmupSnapshot, detector_warmup: DetectorWarmupSnapshot,
source_prefetch: dict[str, object],
) -> dict[str, object]: ) -> dict[str, object]:
outcomes = Counter(item.outcome.value for item in result.terminal_outcomes) outcomes = Counter(item.outcome.value for item in result.terminal_outcomes)
outcome_stages = Counter( outcome_stages = Counter(
@@ -774,6 +815,7 @@ def _loop_document(
"setup_seconds": round(setup_seconds, 6), "setup_seconds": round(setup_seconds, 6),
"cyclic_gc_hot_loop": gc_policy, "cyclic_gc_hot_loop": gc_policy,
"detector_warmup": asdict(detector_warmup), "detector_warmup": asdict(detector_warmup),
"source_prefetch": source_prefetch,
"admitted_count": result.admitted_count, "admitted_count": result.admitted_count,
"delivered_count": len(result.deliveries), "delivered_count": len(result.deliveries),
"effective_world_state_fps": round(len(result.deliveries) / wall_seconds, 6), "effective_world_state_fps": round(len(result.deliveries) / wall_seconds, 6),
@@ -899,6 +941,8 @@ def _pipeline_timing_metrics(
"decode_to_delivery_processing_ns", "decode_to_delivery_processing_ns",
) )
top_level_values: dict[str, list[float]] = {key: [] for key in top_level_fields} top_level_values: dict[str, list[float]] = {key: [] for key in top_level_fields}
decode_by_phase: dict[str, list[float]] = defaultdict(list)
delivered_pacing_lateness_ms: list[float] = []
for document in documents: for document in documents:
detector = cast(Mapping[str, int], document["detector"]) detector = cast(Mapping[str, int], document["detector"])
providers = cast(Mapping[str, int], document["providers"]) providers = cast(Mapping[str, int], document["providers"])
@@ -908,6 +952,11 @@ def _pipeline_timing_metrics(
provider_values[key].append(providers[key] / 1_000_000.0) provider_values[key].append(providers[key] / 1_000_000.0)
for key in top_level_fields: for key in top_level_fields:
top_level_values[key].append(cast(int, document[key]) / 1_000_000.0) top_level_values[key].append(cast(int, document[key]) / 1_000_000.0)
decode_by_phase[cast(str, document["decode_phase"])].append(
cast(int, document["decode_duration_ns"]) / 1_000_000.0
)
source_pacing = cast(Mapping[str, int], document["source_pacing"])
delivered_pacing_lateness_ms.append(source_pacing["lateness_ns"] / 1_000_000.0)
maximum = max( maximum = max(
documents, documents,
key=lambda document: cast(int, document["graph_admission_to_delivery_ns"]), key=lambda document: cast(int, document["graph_admission_to_delivery_ns"]),
@@ -916,6 +965,12 @@ def _pipeline_timing_metrics(
return { return {
"sample_count": len(documents), "sample_count": len(documents),
"decode_duration_ms": _distribution(top_level_values["decode_duration_ns"]), "decode_duration_ms": _distribution(top_level_values["decode_duration_ns"]),
"decode_duration_by_phase_ms": {
phase: _distribution(values) for phase, values in sorted(decode_by_phase.items())
},
"delivered_source_pacing_lateness_ms": _distribution(
delivered_pacing_lateness_ms
),
"detector_ms": { "detector_ms": {
key.removesuffix("_duration_ns"): _distribution(values) key.removesuffix("_duration_ns"): _distribution(values)
for key, values in detector_values.items() for key, values in detector_values.items()
@@ -936,6 +991,50 @@ def _pipeline_timing_metrics(
} }
def _source_decode_metrics(samples: list[DecodedFrameTiming]) -> dict[str, object]:
by_phase: dict[str, list[float]] = defaultdict(list)
for sample in samples:
by_phase[sample.phase.value].append(sample.duration_ns / 1_000_000.0)
return {
"sample_count": len(samples),
"phase_counts": {
phase: len(values) for phase, values in sorted(by_phase.items())
},
"duration_by_phase_ms": {
phase: _distribution(values) for phase, values in sorted(by_phase.items())
},
}
def _source_pacing_metrics(samples: list[SourcePacingTiming]) -> dict[str, object]:
ordered = sorted(samples, key=lambda sample: sample.sequence)
lateness_ms = [sample.lateness_ns / 1_000_000.0 for sample in ordered]
scheduled_intervals_ms = [
(current.scheduled_monotonic_ns - previous.scheduled_monotonic_ns) / 1_000_000.0
for previous, current in zip(ordered, ordered[1:], strict=False)
]
emitted_intervals_ms = [
(current.emitted_monotonic_ns - previous.emitted_monotonic_ns) / 1_000_000.0
for previous, current in zip(ordered, ordered[1:], strict=False)
]
catch_up_emissions = sum(
emitted < scheduled * 0.5
for scheduled, emitted in zip(
scheduled_intervals_ms,
emitted_intervals_ms,
strict=True,
)
if scheduled > 0
)
return {
"sample_count": len(ordered),
"lateness_ms": _distribution(lateness_ms),
"scheduled_interval_ms": _distribution(scheduled_intervals_ms),
"emitted_interval_ms": _distribution(emitted_intervals_ms),
"catch_up_emission_count": catch_up_emissions,
}
def _telemetry_summary(samples: list[dict[str, float]]) -> dict[str, Any]: def _telemetry_summary(samples: list[dict[str, float]]) -> dict[str, Any]:
result: dict[str, Any] = {"sample_count": len(samples)} result: dict[str, Any] = {"sample_count": len(samples)}
for key in ( for key in (
@@ -5,7 +5,7 @@ from __future__ import annotations
import hashlib import hashlib
import json import json
from collections.abc import Callable, Iterator from collections.abc import Callable, Iterator
from dataclasses import dataclass from dataclasses import dataclass, field
from pathlib import Path from pathlib import Path
from threading import Event from threading import Event
@@ -32,10 +32,13 @@ from .providers import (
) )
from .recorded_source import ( from .recorded_source import (
DecodedRecordedSource, DecodedRecordedSource,
DecodePrefetchSnapshot,
DecodeTimingObserver, DecodeTimingObserver,
PrefetchedRecordedImageDecoder,
PyAvRecordedImageDecoder, PyAvRecordedImageDecoder,
RecordedRavnoves00Source, RecordedRavnoves00Source,
ReplayPacing, ReplayPacing,
SourcePacingObserver,
) )
from .reference_graph_runtime import ReferenceGraphRuntimePaths from .reference_graph_runtime import ReferenceGraphRuntimePaths
from .rf_detr_object_detector import ( from .rf_detr_object_detector import (
@@ -64,6 +67,8 @@ class M48sReferenceGraphRuntime:
graph: ReferencePerceptionGraphV2 graph: ReferencePerceptionGraphV2
inference_backend: TritonRfDetrHttpInferenceBackend inference_backend: TritonRfDetrHttpInferenceBackend
source_prefetch: PrefetchedRecordedImageDecoder
_preparation_stop_event: Event = field(default_factory=Event)
def warm_up_detector(self) -> DetectorWarmupSnapshot: def warm_up_detector(self) -> DetectorWarmupSnapshot:
detector = self.graph.detector detector = self.graph.detector
@@ -71,8 +76,18 @@ class M48sReferenceGraphRuntime:
raise M48sReferenceGraphRuntimeError("RF-DETR runtime detector changed before warmup") raise M48sReferenceGraphRuntimeError("RF-DETR runtime detector changed before warmup")
return detector.warm_up() return detector.warm_up()
def prepare_source(self) -> DecodePrefetchSnapshot:
return self.source_prefetch.prepare(self._preparation_stop_event)
def mark_source_admission_started(self) -> None:
self.source_prefetch.mark_admission_started()
def close(self) -> None: def close(self) -> None:
self.inference_backend.close() self._preparation_stop_event.set()
try:
self.source_prefetch.close()
finally:
self.inference_backend.close()
def __enter__(self) -> M48sReferenceGraphRuntime: def __enter__(self) -> M48sReferenceGraphRuntime:
return self return self
@@ -90,9 +105,12 @@ def build_m48s_reference_graph_runtime(
delivery_observer: Callable[[DeliveredFrame, int], None] | None = None, delivery_observer: Callable[[DeliveredFrame, int], None] | None = None,
delivery_evidence_observer: DeliveryEvidenceObserver | None = None, delivery_evidence_observer: DeliveryEvidenceObserver | None = None,
decode_timing_observer: DecodeTimingObserver | None = None, decode_timing_observer: DecodeTimingObserver | None = None,
source_pacing_observer: SourcePacingObserver | None = None,
detector_timing_observer: DetectorTimingObserver | None = None, detector_timing_observer: DetectorTimingObserver | None = None,
maximum_frames: int | None = None, maximum_frames: int | None = None,
source_rate_hz: float | None = None, source_rate_hz: float | None = None,
source_prefetch_capacity_frames: int = 64,
source_prefetch_ready_frames: int = 64,
) -> M48sReferenceGraphRuntime: ) -> M48sReferenceGraphRuntime:
"""Instantiate the complete graph with only its detector pin replaced.""" """Instantiate the complete graph with only its detector pin replaced."""
@@ -117,6 +135,12 @@ def build_m48s_reference_graph_runtime(
if maximum_frames is not None and maximum_frames < 1: if maximum_frames is not None and maximum_frames < 1:
raise M48sReferenceGraphRuntimeError("maximum frame count must be positive") raise M48sReferenceGraphRuntimeError("maximum frame count must be positive")
prefetch = PrefetchedRecordedImageDecoder(
PyAvRecordedImageDecoder(paths.video),
capacity_frames=source_prefetch_capacity_frames,
ready_frames=source_prefetch_ready_frames,
timing_observer=decode_timing_observer,
)
source: SourceProvider = DecodedRecordedSource( source: SourceProvider = DecodedRecordedSource(
source=RecordedRavnoves00Source( source=RecordedRavnoves00Source(
camera_index_path=paths.camera_index, camera_index_path=paths.camera_index,
@@ -131,9 +155,9 @@ def build_m48s_reference_graph_runtime(
if run_mode is GraphRunMode.SOURCE_PACED_LATEST_WINS if run_mode is GraphRunMode.SOURCE_PACED_LATEST_WINS
else None else None
), ),
pacing_observer=source_pacing_observer,
), ),
decoder=PyAvRecordedImageDecoder(paths.video), decoder=prefetch,
timing_observer=decode_timing_observer,
) )
if maximum_frames is not None: if maximum_frames is not None:
source = _LimitedSource(source, maximum_frames) source = _LimitedSource(source, maximum_frames)
@@ -175,9 +199,16 @@ def build_m48s_reference_graph_runtime(
delivery_evidence_observer=delivery_evidence_observer, delivery_evidence_observer=delivery_evidence_observer,
) )
except Exception: except Exception:
backend.close() try:
prefetch.close()
finally:
backend.close()
raise raise
return M48sReferenceGraphRuntime(graph=graph, inference_backend=backend) return M48sReferenceGraphRuntime(
graph=graph,
inference_backend=backend,
source_prefetch=prefetch,
)
class _LimitedSource: class _LimitedSource:
+227 -1
View File
@@ -10,7 +10,8 @@ from collections.abc import Callable, Iterator
from dataclasses import dataclass, replace from dataclasses import dataclass, replace
from enum import StrEnum from enum import StrEnum
from pathlib import Path from pathlib import Path
from threading import Event from queue import Empty, Full, Queue
from threading import Event, Lock, Thread
from typing import Any, Final, Protocol, cast from typing import Any, Final, Protocol, cast
import numpy as np import numpy as np
@@ -53,6 +54,11 @@ class ReplayPacing(StrEnum):
UNCAPPED = "uncapped" UNCAPPED = "uncapped"
class DecodePhase(StrEnum):
PREADMISSION = "preadmission"
HOT_LOOP = "hot-loop"
@dataclass(frozen=True, slots=True) @dataclass(frozen=True, slots=True)
class RecordedFrameReference: class RecordedFrameReference:
"""Opaque reference passed to a provider without decoding sensor data.""" """Opaque reference passed to a provider without decoding sensor data."""
@@ -89,6 +95,7 @@ WaitFunction = Callable[[Event, float], bool]
class DecodedFrameTiming: class DecodedFrameTiming:
sequence: int sequence: int
duration_ns: int duration_ns: int
phase: DecodePhase = DecodePhase.HOT_LOOP
def __post_init__(self) -> None: def __post_init__(self) -> None:
if self.sequence < 0 or self.duration_ns < 0: if self.sequence < 0 or self.duration_ns < 0:
@@ -98,6 +105,43 @@ class DecodedFrameTiming:
DecodeTimingObserver = Callable[[DecodedFrameTiming], None] DecodeTimingObserver = Callable[[DecodedFrameTiming], None]
@dataclass(frozen=True, slots=True)
class SourcePacingTiming:
sequence: int
scheduled_monotonic_ns: int
emitted_monotonic_ns: int
lateness_ns: int
def __post_init__(self) -> None:
if (
self.sequence < 0
or self.scheduled_monotonic_ns < 0
or self.emitted_monotonic_ns < 0
or self.lateness_ns < 0
):
raise RecordedSourceError("source pacing timing must be nonnegative")
SourcePacingObserver = Callable[[SourcePacingTiming], None]
@dataclass(frozen=True, slots=True)
class DecodePrefetchSnapshot:
capacity_frames: int
ready_frames: int
buffered_frames: int
preparation_duration_ns: int
producer_alive: bool
@dataclass(frozen=True, slots=True)
class _DecodeFailure:
error: BaseException
_DECODE_END: Final = object()
class RecordedRavnoves00Source: class RecordedRavnoves00Source:
"""Emit the admitted timeline at its recorded or an explicit replay rate. """Emit the admitted timeline at its recorded or an explicit replay rate.
@@ -116,6 +160,7 @@ class RecordedRavnoves00Source:
expected_frame_count: int = DEFAULT_FRAME_COUNT, expected_frame_count: int = DEFAULT_FRAME_COUNT,
expected_source_pack_sha256: str | None = RECORDED_SOURCE_PACK_SHA256, expected_source_pack_sha256: str | None = RECORDED_SOURCE_PACK_SHA256,
target_rate_hz: float | None = None, target_rate_hz: float | None = None,
pacing_observer: SourcePacingObserver | None = None,
clock_ns: Callable[[], int] = time.monotonic_ns, clock_ns: Callable[[], int] = time.monotonic_ns,
wait: WaitFunction | None = None, wait: WaitFunction | None = None,
) -> None: ) -> None:
@@ -133,6 +178,7 @@ class RecordedRavnoves00Source:
self.expected_frame_count = expected_frame_count self.expected_frame_count = expected_frame_count
self.expected_source_pack_sha256 = expected_source_pack_sha256 self.expected_source_pack_sha256 = expected_source_pack_sha256
self.target_rate_hz = target_rate_hz self.target_rate_hz = target_rate_hz
self.pacing_observer = pacing_observer
self._clock_ns = clock_ns self._clock_ns = clock_ns
self._wait = wait or _event_wait self._wait = wait or _event_wait
@@ -189,6 +235,16 @@ class RecordedRavnoves00Source:
target_ns = started_ns + round(source_elapsed_ns * pacing_scale) target_ns = started_ns + round(source_elapsed_ns * pacing_scale)
if not self._pace_until(stop_event, target_ns): if not self._pace_until(stop_event, target_ns):
return return
emitted_ns = int(self._clock_ns())
if self.pacing_observer is not None:
self.pacing_observer(
SourcePacingTiming(
sequence=packet.envelope.sequence,
scheduled_monotonic_ns=target_ns,
emitted_monotonic_ns=emitted_ns,
lateness_ns=max(0, emitted_ns - target_ns),
)
)
yield packet yield packet
def _pace_until(self, stop_event: Event, target_ns: int) -> bool: def _pace_until(self, stop_event: Event, target_ns: int) -> bool:
@@ -295,6 +351,176 @@ class PyAvRecordedImageDecoder:
container.close() container.close()
class PrefetchedRecordedImageDecoder:
"""Decode into a bounded queue before source admission and during replay.
The queue absorbs ordinary storage/codec jitter without retaining the whole
recording in RAM. Decode work remains single-pass and source order remains
exact. The explicit phase boundary keeps cold codec initialization out of
hot-loop latency attribution.
"""
def __init__(
self,
decoder: RecordedImageDecoder,
*,
capacity_frames: int = 64,
ready_frames: int = 64,
timing_observer: DecodeTimingObserver | None = None,
clock_ns: Callable[[], int] = time.perf_counter_ns,
) -> None:
if capacity_frames < 1:
raise RecordedSourceError("decode prefetch capacity must be positive")
if ready_frames < 1 or ready_frames > capacity_frames:
raise RecordedSourceError("decode prefetch readiness must fit its capacity")
self.decoder = decoder
self.capacity_frames = capacity_frames
self.ready_frames = ready_frames
self.timing_observer = timing_observer
self._clock_ns = clock_ns
self._queue: Queue[NDArray[np.uint8] | _DecodeFailure | object] = Queue(
capacity_frames
)
self._stop = Event()
self._ready = Event()
self._guard = Lock()
self._thread: Thread | None = None
self._phase = DecodePhase.PREADMISSION
self._produced = 0
self._started_ns: int | None = None
self._ready_ns: int | None = None
self._closed = False
def prepare(self, stop_event: Event) -> DecodePrefetchSnapshot:
with self._guard:
if self._closed:
raise RecordedSourceError("decode prefetch is closed")
if self._thread is None:
self._started_ns = int(self._clock_ns())
self._thread = Thread(
target=self._produce,
name="m48s-recorded-decode-prefetch",
daemon=True,
)
self._thread.start()
while not self._ready.wait(0.05):
if stop_event.is_set():
raise RecordedSourceError("decode prefetch preparation was cancelled")
with self._guard:
started_ns = self._started_ns
ready_ns = self._ready_ns
thread = self._thread
produced = self._produced
if started_ns is None or ready_ns is None or thread is None:
raise RecordedSourceError("decode prefetch readiness is incomplete")
first = self._peek()
if isinstance(first, _DecodeFailure):
raise RecordedSourceError("recorded camera prefetch failed") from first.error
return DecodePrefetchSnapshot(
capacity_frames=self.capacity_frames,
ready_frames=self.ready_frames,
buffered_frames=min(produced, self.capacity_frames),
preparation_duration_ns=max(0, ready_ns - started_ns),
producer_alive=thread.is_alive(),
)
def mark_admission_started(self) -> None:
with self._guard:
if not self._ready.is_set():
raise RecordedSourceError("source admission started before decode prefetch")
self._phase = DecodePhase.HOT_LOOP
def frames(self, stop_event: Event) -> Iterator[NDArray[np.uint8]]:
self.prepare(stop_event)
while not stop_event.is_set() and not self._stop.is_set():
try:
item = self._queue.get(timeout=0.05)
except Empty:
continue
try:
if item is _DECODE_END:
return
if isinstance(item, _DecodeFailure):
raise RecordedSourceError("recorded camera decode failed") from item.error
yield cast(NDArray[np.uint8], item)
finally:
self._queue.task_done()
def close(self, *, timeout_seconds: float = 5.0) -> None:
with self._guard:
if self._closed:
return
self._closed = True
self._stop.set()
thread = self._thread
if thread is not None:
thread.join(timeout=max(0.0, timeout_seconds))
if thread.is_alive():
raise RecordedSourceError("decode prefetch worker did not stop")
def _produce(self) -> None:
sequence = 0
images = self.decoder.frames(self._stop)
try:
while not self._stop.is_set():
while self._queue.full() and not self._stop.wait(0.01):
pass
if self._stop.is_set():
return
started_ns = int(self._clock_ns())
try:
image = next(images)
except StopIteration:
self._signal_ready()
self._put(_DECODE_END)
return
completed_ns = int(self._clock_ns())
with self._guard:
phase = self._phase
if self.timing_observer is not None:
self.timing_observer(
DecodedFrameTiming(
sequence=sequence,
duration_ns=max(0, completed_ns - started_ns),
phase=phase,
)
)
if not self._put(np.asarray(image, dtype=np.uint8)):
return
sequence += 1
with self._guard:
self._produced = sequence
ready = sequence >= self.ready_frames
if ready:
self._signal_ready()
except BaseException as exc:
self._put(_DecodeFailure(exc))
self._signal_ready()
finally:
close = getattr(images, "close", None)
if callable(close):
close()
def _put(self, item: NDArray[np.uint8] | _DecodeFailure | object) -> bool:
while not self._stop.is_set():
try:
self._queue.put(item, timeout=0.05)
return True
except Full:
continue
return False
def _signal_ready(self) -> None:
with self._guard:
if self._ready_ns is None:
self._ready_ns = int(self._clock_ns())
self._ready.set()
def _peek(self) -> NDArray[np.uint8] | _DecodeFailure | object | None:
with self._queue.mutex:
return self._queue.queue[0] if self._queue.queue else None
def _packet( def _packet(
frame_index: int, frame_index: int,
camera: dict[str, object], camera: dict[str, object],
+45
View File
@@ -62,9 +62,12 @@ from k1link.perception.providers import (
from k1link.perception.recorded_source import ( from k1link.perception.recorded_source import (
DecodedFrameTiming, DecodedFrameTiming,
DecodedRecordedSource, DecodedRecordedSource,
DecodePhase,
PrefetchedRecordedImageDecoder,
RecordedRavnoves00Source, RecordedRavnoves00Source,
RecordedSourceError, RecordedSourceError,
ReplayPacing, ReplayPacing,
SourcePacingTiming,
) )
@@ -839,6 +842,7 @@ def test_recorded_source_target_rate_changes_only_wall_clock_pacing(tmp_path: Pa
camera_path, timeline_path = _write_recorded_fixture(tmp_path) camera_path, timeline_path = _write_recorded_fixture(tmp_path)
now = [1_000_000_000] now = [1_000_000_000]
waits: list[float] = [] waits: list[float] = []
pacing: list[SourcePacingTiming] = []
def wait(stop_event: Event, seconds: float) -> bool: def wait(stop_event: Event, seconds: float) -> bool:
waits.append(seconds) waits.append(seconds)
@@ -852,6 +856,7 @@ def test_recorded_source_target_rate_changes_only_wall_clock_pacing(tmp_path: Pa
target_rate_hz=20.0, target_rate_hz=20.0,
expected_frame_count=2, expected_frame_count=2,
expected_source_pack_sha256=None, expected_source_pack_sha256=None,
pacing_observer=pacing.append,
clock_ns=lambda: now[0], clock_ns=lambda: now[0],
wait=wait, wait=wait,
) )
@@ -859,6 +864,8 @@ def test_recorded_source_target_rate_changes_only_wall_clock_pacing(tmp_path: Pa
packets = list(source.packets(Event())) packets = list(source.packets(Event()))
assert waits == pytest.approx([0.05]) assert waits == pytest.approx([0.05])
assert [sample.sequence for sample in pacing] == [0, 1]
assert [sample.lateness_ns for sample in pacing] == [0, 0]
assert ( assert (
packets[1].envelope.timestamps.source_ns packets[1].envelope.timestamps.source_ns
- packets[0].envelope.timestamps.source_ns - packets[0].envelope.timestamps.source_ns
@@ -879,6 +886,44 @@ def test_recorded_source_rejects_target_rate_for_uncapped_replay(tmp_path: Path)
) )
def test_prefetched_decoder_moves_cold_decode_before_source_admission() -> None:
observed: list[DecodedFrameTiming] = []
class Decoder:
def frames(self, stop_event: Event) -> Iterator[np.ndarray]:
for index in range(4):
if stop_event.is_set():
return
yield np.full((2, 3, 3), index, dtype=np.uint8)
decoder = PrefetchedRecordedImageDecoder(
Decoder(),
capacity_frames=2,
ready_frames=2,
timing_observer=observed.append,
)
stop_event = Event()
snapshot = decoder.prepare(stop_event)
assert snapshot.buffered_frames == 2
assert snapshot.capacity_frames == 2
assert [sample.phase for sample in observed] == [
DecodePhase.PREADMISSION,
DecodePhase.PREADMISSION,
]
decoder.mark_admission_started()
frames = list(decoder.frames(stop_event))
decoder.close()
assert [int(frame[0, 0, 0]) for frame in frames] == [0, 1, 2, 3]
assert [sample.sequence for sample in observed] == [0, 1, 2, 3]
assert [sample.phase for sample in observed[2:]] == [
DecodePhase.HOT_LOOP,
DecodePhase.HOT_LOOP,
]
def test_recorded_source_rejects_timeline_mismatch(tmp_path: Path) -> None: def test_recorded_source_rejects_timeline_mismatch(tmp_path: Path) -> None:
camera_path, timeline_path = _write_recorded_fixture(tmp_path, mismatched=True) camera_path, timeline_path = _write_recorded_fixture(tmp_path, mismatched=True)
source = RecordedRavnoves00Source( source = RecordedRavnoves00Source(