Skip to content
Merged
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
4 changes: 4 additions & 0 deletions flashdreams/flashdreams/runtime/demo/drivers.py
Original file line number Diff line number Diff line change
Expand Up @@ -65,6 +65,8 @@ def run_one_session(
session_info = host.call(_session_info, session)
session_edges.output_sink.open(session_info)
setup_ok = True
except DriverInvariantError:
raise
except Exception as exc:
action = session_edges.error_policy.handle_setup_error(exc)
if action.drop_chunk or action.result_status == "completed":
Expand Down Expand Up @@ -206,6 +208,8 @@ async def run_one_session(
session_edges.output_sink.open(session_info)
session_edges.output_sink.begin_generation(generation)
setup_ok = True
except DriverInvariantError:
raise
except Exception as exc:
action = session_edges.error_policy.handle_setup_error(exc)
if action.drop_chunk or action.result_status == "completed":
Expand Down
2 changes: 2 additions & 0 deletions flashdreams/flashdreams/runtime/demo/outputs.py
Original file line number Diff line number Diff line change
Expand Up @@ -169,6 +169,8 @@ def open(self, session_info: SessionInfo) -> None:
def begin_generation(self, generation: int) -> None:
if generation < 0:
raise ValueError("generation must be >= 0.")
# MP4 recording is continuous across realtime resets; WebRTC is the sink
# that drops stale generations.

def write(self, result: StepResult) -> OutputDecision:
if not self._opened or self._closed or self._collector is None:
Expand Down
63 changes: 47 additions & 16 deletions flashdreams/flashdreams/runtime/demo/replay.py
Original file line number Diff line number Diff line change
Expand Up @@ -90,16 +90,12 @@ def run_replay_demo(
raise ValueError("run_replay_demo does not support WebRTC output.")

prepared = adapter.prepare_scenario(spec)
mapping = prepared.mapping or adapter.default_input_mapping()
if mapping is None:
raise ValueError(
"Demo scenario did not provide an input mapping, and the adapter "
"has no default input mapping."
)
mapping = _scenario_mapping(prepared=prepared, adapter=adapter)
if spec.config is None:
raise RuntimeError("DemoSpec.config was not initialized.")

if runner is not None:
mapping = _require_replay_mapping(mapping)
return _run_replay_demo_with_compat_runner(
spec=spec,
adapter=adapter,
Expand Down Expand Up @@ -197,18 +193,26 @@ def _run_replay_demo_with_run_mode(
spec: DemoSpec,
adapter: DemoAdapter,
prepared: "PreparedScenario",
mapping: InputMapping,
mapping: InputMapping | None,
output_sink_factory: OutputSinkFactory,
metrics: MetricsRecorder | None,
) -> RunResult:
config = _require_config(spec)
_validate_replay_mapping(
adapter=adapter,
config=config,
mapping=mapping,
source_schema=prepared.source_schema,
canonicalizer=prepared.canonicalizer,
)
if mapping is None:
if not callable(getattr(adapter, "create_model_input_provider", None)):
raise ValueError(
"Demo scenario did not provide an input mapping, and the adapter "
"has no model input provider or default input mapping."
)
adapter.validate_config(config)
else:
_validate_replay_mapping(
adapter=adapter,
config=config,
mapping=mapping,
source_schema=prepared.source_schema,
canonicalizer=prepared.canonicalizer,
)
request_state = _ReplayStepRequestState()
runtime = _ReplayRuntimeAdapter(
runtime=adapter.create_runtime(config),
Expand Down Expand Up @@ -325,7 +329,7 @@ def __init__(
self,
*,
adapter: DemoAdapter,
mapping: InputMapping,
mapping: InputMapping | None,
request_state: "_ReplayStepRequestState",
) -> None:
self._adapter = adapter
Expand Down Expand Up @@ -370,6 +374,11 @@ def create_model_input_provider(
create_provider = getattr(self._adapter, "create_model_input_provider", None)
if callable(create_provider):
return create_provider(spec, scenario)
if self._mapping is None:
raise ValueError(
"Replay adapter requires an input mapping when no model input "
"provider is available."
)
return _ReplayMappingModelInputProvider(
adapter=self._adapter,
scenario=scenario,
Expand Down Expand Up @@ -498,7 +507,7 @@ def next_window(self, request: StepRequirements) -> UserInputWindow:
return UserInputWindow(
start_s=window.start_s,
end_s=window.end_s,
inputs=self._user_inputs,
inputs=self._user_inputs.window(window),
)


Expand Down Expand Up @@ -589,6 +598,28 @@ def _require_config(spec: DemoSpec) -> InferenceConfig:
return spec.config


def _scenario_mapping(
*,
prepared: PreparedScenario,
adapter: DemoAdapter,
) -> InputMapping | None:
if prepared.mapping is not None:
return prepared.mapping
default_input_mapping = getattr(adapter, "default_input_mapping", None)
if not callable(default_input_mapping):
return None
return default_input_mapping()


def _require_replay_mapping(mapping: InputMapping | None) -> InputMapping:
if mapping is not None:
return mapping
raise ValueError(
"Compatibility replay runners require an input mapping; the prepared "
"scenario and adapter did not provide one."
)


def _all_user_inputs_window(user_inputs: UserInputs) -> TimeWindow:
if not user_inputs.events:
return TimeWindow(start_s=0.0, end_s=3600.0)
Expand Down
35 changes: 35 additions & 0 deletions flashdreams/tests/test_demo_runtime_realtime_driver.py
Original file line number Diff line number Diff line change
Expand Up @@ -223,6 +223,41 @@ async def test_realtime_driver_invariant_finalizes_edges_before_reraising() -> N
assert metrics.closed


@pytest.mark.asyncio
async def test_realtime_driver_setup_invariant_reraises_without_error_policy() -> None:
runtime = _FakeRealtimeRuntime(session=_FakeRealtimeSession(num_steps=1))
host = RuntimeHost(runtime)
provider = _FakeRealtimeProvider(
fail_initial=DriverInvariantError("setup invariant")
)
output = _RecordingOutputSink()
transport = _RecordingTransport()
metrics = InMemorySessionMetricsRecorder()
edges = _edges(
output=output,
transport=transport,
metrics=metrics,
error_policy=_SetupPolicy(result_status="failed"),
)

try:
with pytest.raises(DriverInvariantError, match="setup invariant"):
await RealtimeSessionDriver().run_one_session(
host=host,
provider=provider,
session_edges=edges,
pipeline=StepPipeline(),
)
finally:
host.close()

assert provider.close_count == 1
assert output.close_count == 1
assert transport.close_count == 1
assert metrics.closed
assert metrics.errors == []


@pytest.mark.asyncio
async def test_realtime_step_invariant_reraises_without_error_policy() -> None:
runtime = _FakeRealtimeRuntime(session=_FakeRealtimeSession(num_steps=1))
Expand Down
31 changes: 31 additions & 0 deletions flashdreams/tests/test_demo_runtime_vertical_slice.py
Original file line number Diff line number Diff line change
Expand Up @@ -572,6 +572,37 @@ def test_setup_failure_can_return_skipped_but_not_completed() -> None:
assert provider.close_count == 1


def test_batch_driver_setup_invariant_reraises_without_error_policy() -> None:
provider = _FakeVideoModelInputProvider(
fail_initial=DriverInvariantError("setup invariant")
)
output = _RecordingOutputSink()
transport = _RecordingTransport()
metrics = InMemorySessionMetricsRecorder()
edges = SessionEdges(
input_source=_FakeBatchInputSource(num_windows=1),
output_sink=output,
cleanup_tasks=set(),
metrics=metrics,
error_policy=_SetupPolicy(result_status="failed"),
transport=transport,
)

with pytest.raises(DriverInvariantError, match="setup invariant"):
BatchSessionDriver().run_one_session(
host=RuntimeHost(_FakeVideoRuntime(session=_FakeVideoSession(num_steps=1))),
provider=provider,
session_edges=edges,
pipeline=StepPipeline(),
)

assert output.close_count == 1
assert transport.close_count == 1
assert metrics.closed
assert metrics.errors == []
assert provider.close_count == 1


def test_batch_driver_invariant_finalizes_edges_when_host_closed() -> None:
runtime = _FakeVideoRuntime(session=_FakeVideoSession(num_steps=1))
host = RuntimeHost(runtime)
Expand Down
47 changes: 11 additions & 36 deletions integrations/lingbot/lingbot/demo/adapter.py
Original file line number Diff line number Diff line change
Expand Up @@ -9,7 +9,6 @@
from typing import Any

from flashdreams.runtime import (
InputCanonicalizer,
UserInputCapability,
UserInputs,
UserInputSchema,
Expand All @@ -22,10 +21,6 @@
WebRTCOutputSpec,
)
from flashdreams.runtime.interfaces import InferenceRuntime
from lingbot.input_mapping import (
KeyboardToCameraCommand,
TextEventSelection,
)
from lingbot.runtime import (
FIELD_FPS,
FIELD_PIXEL_HEIGHT,
Expand All @@ -37,7 +32,11 @@
inference_input_from_replay_inputs,
)

from .providers import LingbotInputProvider
from .providers import (
PROVIDER_INPUTS_METADATA_KEY,
LingbotInputProvider,
create_lingbot_provider_inputs,
)
from .spec import (
resolve_replay_inputs,
resolve_text_event_prompts,
Expand Down Expand Up @@ -93,27 +92,11 @@ def prepare_scenario(self, spec: DemoSpec) -> PreparedScenario:
)
text_event_prompts = resolve_text_event_prompts(scenario)
user_inputs = resolve_user_input_events(scenario)
if live_camera or _camera_source(scenario) == "events":
# Live control still needs the scenario's calibration, so the trace
# is loaded for its intrinsics and world scale and then discarded
# as a trajectory source.
trace = self.create_input_mapping(replay_inputs).camera_trace
mapping = self.create_live_input_mapping(
fps=replay_inputs.fps,
base_intrinsics=trace.intrinsics[0],
# A trace's world scale is derived from how far its poses
# travel, so a stationary example yields 0. Live control has no
# trajectory to normalize against, so it falls back to the same
# unit scale the live runtime uses.
world_scale=trace.world_scale or 1.0,
prompt=replay_inputs.prompt,
text_event_prompts=text_event_prompts,
)
else:
mapping = self.create_input_mapping(
replay_inputs,
text_event_prompts=text_event_prompts,
)
provider_inputs = create_lingbot_provider_inputs(
replay_inputs,
live_camera=live_camera or _camera_source(scenario) == "events",
text_event_prompts=text_event_prompts,
)
return PreparedScenario(
initial_inputs=inference_input_from_replay_inputs(replay_inputs),
user_inputs=user_inputs,
Expand All @@ -122,11 +105,10 @@ def prepare_scenario(self, spec: DemoSpec) -> PreparedScenario:
include_keyboard=live_camera,
include_text_events=live_camera and bool(text_event_prompts),
),
canonicalizer=_canonicalizer(text_event_prompts),
mapping=mapping,
metadata={
"model_id": self.model_id,
"preset_id": self.preset_id(spec.config),
PROVIDER_INPUTS_METADATA_KEY: provider_inputs,
},
)

Expand Down Expand Up @@ -246,13 +228,6 @@ def _source_schema(
)


def _canonicalizer(text_event_prompts: Mapping[str, str] | None) -> InputCanonicalizer:
converters: list[Any] = [KeyboardToCameraCommand()]
if text_event_prompts:
converters.append(TextEventSelection())
return InputCanonicalizer(converters)


__all__ = [
"LingbotDemoAdapter",
"ReplayRuntimeFactory",
Expand Down
Loading
Loading