diff --git a/scripts/benchmark/rl/benchmark_offpolicy_collector_active.py b/scripts/benchmark/rl/benchmark_offpolicy_collector_active.py index 57f0cc66a..91d9d38d1 100644 --- a/scripts/benchmark/rl/benchmark_offpolicy_collector_active.py +++ b/scripts/benchmark/rl/benchmark_offpolicy_collector_active.py @@ -246,6 +246,12 @@ class CollectorResult: # optional benchmark-side monkeypatch used for RNG/noise-buffer profiling. numpy_random_ms_per_vector_step: TimingStats | None = None numpy_random_calls_per_vector_step: TimingStats | None = None + # Public backend runtime diagnostics, captured after construction and before + # the measured window. They make CUDA graph fallbacks visible in cross-host + # A/B reports instead of requiring private backend inspection. + backend_runtime_diagnostics: dict[str, dict[str, bool | str | None]] = field( + default_factory=dict + ) def _stats(samples_ms: list[float]) -> TimingStats: @@ -454,6 +460,21 @@ def _configure_collector_cpu_threads(cap: int | None = None) -> int: return n_threads +def _backend_runtime_diagnostics(env: Any) -> dict[str, dict[str, bool | str | None]]: + backend = getattr(env, "_backend", None) + diagnostics = getattr(backend, "get_tensor_runtime_diagnostics", None) + if diagnostics is None: + return {} + return { + name: { + "requested": bool(diagnostic.requested), + "enabled": bool(diagnostic.enabled), + "disable_reason": diagnostic.disable_reason, + } + for name, diagnostic in diagnostics().items() + } + + def _runtime_sim_backend(sim: str) -> str: return BACKEND_ALIASES.get(sim, sim) @@ -646,6 +667,7 @@ def _run_active_window_case( ) env_device = env.device + backend_runtime_diagnostics = _backend_runtime_diagnostics(env) actions = torch.zeros((case.num_envs, case.action_dim), dtype=torch.float32, device=env_device) state = env.step(actions) obs, critic = split_obs_dict(state.obs) @@ -847,6 +869,7 @@ def _run_active_window_case( numpy_random_calls_per_vector_step=( _stats(numpy_random_call_samples) if numpy_random_call_samples else None ), + backend_runtime_diagnostics=backend_runtime_diagnostics, ) diff --git a/src/unilab/tasks/motion_tracking/common/tensor_state_store.py b/src/unilab/tasks/motion_tracking/common/tensor_state_store.py index 2b1cb640e..d1853050b 100644 --- a/src/unilab/tasks/motion_tracking/common/tensor_state_store.py +++ b/src/unilab/tasks/motion_tracking/common/tensor_state_store.py @@ -180,8 +180,8 @@ def _read_device_resident(self, rows: torch.Tensor | None) -> None: self._device_sensor_views = sensor_views self._device_linvel_view = linvel_view self._device_gyro_view = gyro_view - self.linvel = linvel_view.clone() - self.gyro = gyro_view.clone() + self.linvel = linvel_view + self.gyro = gyro_view else: # Named sensor views are not guaranteed to be stable snapshots on # every DEVICE_RESIDENT adapter. Crossing the public sensor-read @@ -191,17 +191,14 @@ def _read_device_resident(self, rows: torch.Tensor | None) -> None: gyro_view = self.backend.get_sensor_view("torso_gyro", device=self.device) self._device_linvel_view = linvel_view self._device_gyro_view = gyro_view - self.linvel.copy_(linvel_view) - self.gyro.copy_(gyro_view) + self.linvel = linvel_view + self.gyro = gyro_view # One tracked-sensor read refreshes all injected frame sensors. self.backend.get_sensor_view(f"track_pos_w_{self.body_names[0]}", device=self.device) - if rows is not None: - linvel_view = self._device_linvel_view - gyro_view = self._device_gyro_view - if linvel_view is None or gyro_view is None: - raise RuntimeError("scalar sensor views were not negotiated before row read") - self.linvel[rows] = linvel_view.index_select(0, rows) - self.gyro[rows] = gyro_view.index_select(0, rows) + if rows is not None and ( + self._device_linvel_view is None or self._device_gyro_view is None + ): + raise RuntimeError("scalar sensor views were not negotiated before row read") self.qpos = state_views["qpos"] self.qvel = state_views["qvel"] @@ -333,6 +330,22 @@ def step_tensor(self, ctrl: torch.Tensor, nsteps: int) -> dict | None: return cast(dict | None, self._host_bridge_plan.step(nsteps)) return cast(dict | None, self.backend.step_tensor(ctrl, nsteps=nsteps)) + def refresh_named_sensors(self) -> None: + """Cross the backend's named-sensor projection boundary. + + DEVICE_RESIDENT adapters may recompute derived projections when a named + sensor view is requested. Unlike :meth:`read`, this method performs no + row validation, state-view negotiation, or task-buffer copy. + """ + if self._execution is not TensorExecution.DEVICE_RESIDENT: + raise RuntimeError( + "named-sensor refresh requires DEVICE_RESIDENT tensor execution; " + f"received {self._execution}" + ) + self.backend.get_sensor_view("pelvis_local_linvel", device=self.device) + self.backend.get_sensor_view("torso_gyro", device=self.device) + self.backend.get_sensor_view(f"track_pos_w_{self.body_names[0]}", device=self.device) + def apply_reset( self, rows: torch.Tensor, qpos: torch.Tensor, qvel: torch.Tensor ) -> dict | None: @@ -367,11 +380,24 @@ def refresh_after_selected_reset(self, ctrl: torch.Tensor, nsteps: int) -> dict ) if not self._views_require_readiness_barrier: return None + if self._uses_named_sensor_reset_refresh(): + self.refresh_named_sensors() + self._views_require_readiness_barrier = False + return None result = self.step_tensor(ctrl, nsteps=nsteps) self._views_require_readiness_barrier = False self.last_backend_result = result return result + def _uses_named_sensor_reset_refresh(self) -> bool: + """Read the public optional optimization declaration, if present.""" + try: + diagnostics = self.backend.get_tensor_runtime_diagnostics() + except (AttributeError, KeyError, TypeError, ValueError): + return False + diagnostic = diagnostics.get("selected_reset_sensor_refresh") + return bool(getattr(diagnostic, "enabled", False)) + def validate_finite(self) -> None: qpos, qvel = self._require_qviews() values = ( diff --git a/tests/tasks/test_tensor_runtime_components.py b/tests/tasks/test_tensor_runtime_components.py index 2431bdedb..ae1bb36e3 100644 --- a/tests/tasks/test_tensor_runtime_components.py +++ b/tests/tasks/test_tensor_runtime_components.py @@ -1,5 +1,7 @@ from __future__ import annotations +from types import SimpleNamespace + import numpy as np import pytest import torch @@ -292,6 +294,25 @@ def get_sensor_view(self, name, device=None): return torch.full((1, 3), value, device=device) +class _NamedSensorResetRefreshBackend(_SelectedResetReadinessBackend): + """A DEVICE_RESIDENT adapter whose named views publish lazy reset state.""" + + backend_type = "fake-device-named-sensor-reset" + + def get_tensor_runtime_diagnostics(self): + return { + "selected_reset_sensor_refresh": SimpleNamespace( + requested=True, + enabled=True, + disable_reason=None, + ) + } + + def set_state_tensor(self, env_indices, qpos, qvel, randomization=None): + self.stale = False + return {"timing": {"selected_reset_ms": 1.0}} + + def test_tensor_state_store_full_read_validates_layout_and_finite_state() -> None: store = TensorDeviceStateStore( backend=_HostBridgeBackend(), # pyright: ignore[reportArgumentType] @@ -435,6 +456,30 @@ def test_host_bridge_selected_reset_has_no_device_resident_readiness_step() -> N assert backend.selected_read_calls == 1 +def test_tensor_state_store_uses_declared_named_sensor_reset_refresh() -> None: + backend = _NamedSensorResetRefreshBackend() + store = TensorDeviceStateStore( + backend=backend, # pyright: ignore[reportArgumentType] + device=torch.device("cpu"), + num_envs=1, + joint_qpos_ids=np.array([7], dtype=np.int64), + joint_qvel_ids=np.array([6], dtype=np.int64), + body_names=("pelvis",), + body_ids=np.array([0], dtype=np.intp), + ) + ctrl = torch.tensor([[0.25]], dtype=torch.float32) + qpos = torch.tensor([[0.0, 0.0, 0.0, 0.0, 1.0, 0.0, 0.0, 0.25]]) + qvel = torch.zeros((1, 7)) + + store.apply_reset(torch.tensor([0], dtype=torch.int64), qpos, qvel) + result = store.refresh_after_selected_reset(ctrl, nsteps=3) + store.read() + + assert result is None + assert backend.stale is False + assert backend.step_calls == [] + + def test_tensor_state_store_empty_rows_validate_and_read_without_sync_or_backend_read( monkeypatch: pytest.MonkeyPatch, ) -> None: