"""Cross-backend sim2sim contract snapshot and resolution."""
from __future__ import annotations
import json
from collections.abc import Iterator
from contextlib import contextmanager
from pathlib import Path
from typing import Any
from omegaconf import DictConfig, OmegaConf
[docs]
class CrossBackendIncompatibleError(RuntimeError):
"""Raised when a target play config diverges from the source training contract."""
ALLOWLIST: list[str] = [
"training.sim_backend",
"env.scene",
"training.play_steps",
"env.domain_rand",
"env.noise_config",
"env.commands.vel_limit",
]
WARNING_LIST: list[str] = [
"reward.scales",
"reward.base_height_target",
"reward.max_tilt_deg",
"reward.min_base_height",
"env.control_config.simulate_action_latency",
"env.ctrl_dt",
]
DENYLIST: list[str] = [
"algo.obs_groups",
"env.control_config.action_scale",
"env.observations",
"env.actions",
"env.policy_observation_group",
"env.critic_observation_group",
"algo.policy.actor_hidden_dims",
"algo.policy.critic_hidden_dims",
"algo.empirical_normalization",
"algo.obs_normalization",
"env.commands.motion.params.sampling_mode",
]
SNAPSHOT_FIELDS: list[str] = DENYLIST + WARNING_LIST
ENV_STRUCTURAL_DENYLIST: list[str] = [path for path in DENYLIST if path.startswith("env.")]
LEGACY_PATH_ALIASES: dict[str, str] = {
"env.sampling_mode": "env.commands.motion.params.sampling_mode",
}
_CANONICAL_PATH_FALLBACKS: dict[str, str] = {
canonical: legacy for legacy, canonical in LEGACY_PATH_ALIASES.items()
}
def _select(cfg: Any, path: str) -> Any:
"""Return the effective value at a dotted path (or ``None`` if absent)."""
return OmegaConf.select(cfg, path)
def _select_contract(cfg: Any, path: str) -> Any:
value = _select(cfg, path)
if value is not None:
return value
legacy_path = _CANONICAL_PATH_FALLBACKS.get(path)
return None if legacy_path is None else _select(cfg, legacy_path)
def _to_plain(value: Any) -> Any:
if OmegaConf.is_config(value):
return OmegaConf.to_container(value, resolve=True)
return value
def _normalize(value: Any) -> Any:
"""Canonicalize a value for order-insensitive, type-tolerant comparison."""
if OmegaConf.is_config(value):
value = OmegaConf.to_container(value, resolve=True)
if isinstance(value, bool): # must precede int: bool is a subclass of int
return value
if isinstance(value, dict):
return {str(k): _normalize(v) for k, v in sorted(value.items())}
if isinstance(value, (list, tuple)):
return [_normalize(v) for v in value]
if isinstance(value, (int, float)):
return float(value) # 0 == 0.0; YAML-int vs JSON-float parity
return value
def _values_equal(a: Any, b: Any) -> bool:
return bool(_normalize(a) == _normalize(b))
def _format_value(value: Any) -> str:
return json.dumps(_normalize(value), ensure_ascii=False, sort_keys=True)
def _diff_line(path: str, source_value: Any, target_value: Any) -> str:
return f"{path}: source={_format_value(source_value)} target={_format_value(target_value)}"
def _asymmetric_line(path: str, present_value: Any, *, source_present: bool) -> str:
"""Format a denial for an env-structural field set on exactly one side."""
value = _format_value(present_value)
if source_present:
return (
f"{path}: source={value} target=<absent> (target omits this field and "
"falls back to the env default, which may differ; set it explicitly in the "
"target task YAML to make the contract verifiable)"
)
return (
f"{path}: source=<absent> target={value} (the trained run omitted this field "
"and used the env default; set it explicitly so the contract can be verified)"
)
def _read_snapshot(run_dir: Path) -> dict[str, Any] | None:
"""Read ``contract_snapshot`` from ``run_dir/run_config.json`` (``None`` if absent)."""
path = run_dir / "run_config.json"
if not path.is_file():
return None
try:
parsed = json.loads(path.read_text(encoding="utf-8"))
except (OSError, ValueError):
return None
if not isinstance(parsed, dict):
return None
snapshot = parsed.get("contract_snapshot")
if not isinstance(snapshot, dict):
return None
return snapshot
[docs]
def resolve_sim2sim_config(
source_run_dir: str | Path | None,
target_cfg: DictConfig,
*,
algo_name: str | None = None,
strict: bool = True,
) -> DictConfig | None:
"""Validate a target play config against the source training contract.
Returns ``None`` if ``source_run_dir`` is ``None``; otherwise returns ``target_cfg``
unchanged (never mutated). Raises :class:`CrossBackendIncompatibleError` under
``strict`` when any DENYLIST field differs, including asymmetric presence for
:data:`ENV_STRUCTURAL_DENYLIST` paths.
"""
if source_run_dir is None:
print("[sim2sim] no source run dir; skipping cross-backend contract check")
return None
run_dir = Path(source_run_dir)
snapshot = _read_snapshot(run_dir)
if snapshot is None:
print(
f"[sim2sim] {run_dir}/run_config.json has no contract_snapshot "
"(old run); skipping cross-backend enforcement"
)
return target_cfg
denials: list[str] = []
canonical_snapshot: dict[str, Any] = {}
for raw_path, source_value in snapshot.items():
path = LEGACY_PATH_ALIASES.get(raw_path, raw_path)
canonical_snapshot[path] = source_value
for path, source_value in canonical_snapshot.items():
target_value = _select_contract(target_cfg, path)
if target_value is None:
if path in ENV_STRUCTURAL_DENYLIST:
denials.append(_asymmetric_line(path, source_value, source_present=True))
continue
if _values_equal(source_value, target_value):
continue
line = _diff_line(path, source_value, target_value)
if path in DENYLIST:
denials.append(line)
else:
print(f"[sim2sim] WARNING override {line}")
for path in ENV_STRUCTURAL_DENYLIST:
if path in canonical_snapshot:
continue
target_value = _select_contract(target_cfg, path)
if target_value is not None:
denials.append(_asymmetric_line(path, target_value, source_present=False))
if denials:
message = (
"Cross-backend sim2sim contract mismatch between the trained policy and "
f"the target play config.\nSource run: {run_dir}\n"
"The following policy-defining fields differ and must be reconciled in "
"the target task YAML:\n " + "\n ".join(denials)
)
if strict:
raise CrossBackendIncompatibleError(message)
print(f"[sim2sim] WARNING (non-strict) {message}")
return target_cfg
def _looks_like_dim_mismatch(message: str) -> bool:
"""Return whether ``load_state_dict`` reported a parameter size mismatch."""
return "size mismatch for " in message.lower()
[docs]
@contextmanager
def policy_load_dim_guard(
*,
env_obs_dim: int | None = None,
env_action_dim: int | None = None,
algo_name: str | None = None,
) -> Iterator[None]:
"""Re-raise a tensor shape mismatch during checkpoint load as a sim2sim diagnostic.
Non-matching errors propagate unchanged, so a valid load is never blocked.
"""
try:
yield
except (RuntimeError, ValueError) as exc:
if not _looks_like_dim_mismatch(str(exc)):
raise
raise CrossBackendIncompatibleError(
"Trained policy checkpoint does not fit this play environment -- likely a "
"cross-backend sim2sim dimension mismatch.\n"
f" algo: {algo_name}\n"
f" env policy obs dim: {env_obs_dim}\n"
f" env action dim: {env_action_dim}\n"
"The checkpoint's tensor shapes do not match the env's observation/action "
"dimensions. Check the task's obs_groups_spec and action space across "
"backends; see resolve_sim2sim_config and run "
"`uv run scripts/audit_sim2sim_contracts.py`.\n"
f"Original load error:\n{exc}"
) from exc
[docs]
class Sim2SimConfigResolver:
"""Object facade over the module-level sim2sim contract API."""
ALLOWLIST = ALLOWLIST
WARNING_LIST = WARNING_LIST
DENYLIST = DENYLIST
ENV_STRUCTURAL_DENYLIST = ENV_STRUCTURAL_DENYLIST
[docs]
@staticmethod
def resolve(
source_run_dir: str | Path | None,
target_cfg: DictConfig,
*,
algo_name: str | None = None,
strict: bool = True,
) -> DictConfig | None:
"""See :func:`resolve_sim2sim_config`."""
return resolve_sim2sim_config(
source_run_dir, target_cfg, algo_name=algo_name, strict=strict
)
[docs]
@staticmethod
def load_dim_guard(
*,
env_obs_dim: int | None = None,
env_action_dim: int | None = None,
algo_name: str | None = None,
):
"""See :func:`policy_load_dim_guard`."""
return policy_load_dim_guard(
env_obs_dim=env_obs_dim, env_action_dim=env_action_dim, algo_name=algo_name
)