diff --git a/motrix_env_core/src/motrix_env_core/manager/__init__.py b/motrix_env_core/src/motrix_env_core/manager/__init__.py index 0808d9bd..b20208b7 100644 --- a/motrix_env_core/src/motrix_env_core/manager/__init__.py +++ b/motrix_env_core/src/motrix_env_core/manager/__init__.py @@ -10,6 +10,7 @@ from motrix_env_core.numba.kernel_data import SharedArray, kernel_data from motrix_env_core.numba.manager.actions import ( ActionCfg, + ActionState, ActionTerm, ManagerActionsCfg, ) @@ -38,6 +39,7 @@ __all__ = [ "ActionCfg", "ActionTerm", + "ActionState", "CommandCfg", "CommandTerm", "ManagerActionsCfg", diff --git a/motrix_env_core/src/motrix_env_core/mdp/action.py b/motrix_env_core/src/motrix_env_core/mdp/action.py new file mode 100644 index 00000000..90ed3ed3 --- /dev/null +++ b/motrix_env_core/src/motrix_env_core/mdp/action.py @@ -0,0 +1,215 @@ +# Copyright Motphys Technology Co., Ltd. 2025, 2026 +# SPDX-License-Identifier: Apache-2.0 + +"""Generic joint-position action terms for manager-based environments.""" + +import gymnasium as gym +import numpy as np + +from motrix_env_core.config import configclass +from motrix_env_core.config.scene import RobotCfg +from motrix_env_core.manager import ActionCfg, ActionState, ActionTerm, ManagerEnv, SharedArray, kernel_data +from motrix_env_core.mdp.action_space import joint_position_action_space_from_ctrl_ranges +from motrix_env_core.sim.model import ActuatorSpec, ActuatorType + + +@kernel_data +class JointPositionActionState(ActionState): + """Persistent joint-position action pipeline, history, and shared model data. + + The applied position target is ``target_raw * action_scales + + default_angles``, where ``target_raw`` is the raw or delayed policy + action. + + Attributes: + action_queue: Raw policy-action ring buffer ``(N, W, A)`` with + ``W = max(delay_hi + 1, 2)``. Current and previous actions occupy + slots ``action_ptr`` and ``(action_ptr - 1) % W`` respectively. + It is consumed by action + observations, the action-rate penalty, and delayed target lookup. + default_angles: Per-actuator default pose, ``(A,)``. + joint_lower: Joint position lower limits, ``(A,)``. + joint_upper: Joint position upper limits, ``(A,)``. + action_scales: Per-actuator target scaling, ``(A,)``. + delay_steps: Per-lane applied delay in control steps, resampled + uniformly from ``[delay_lo, delay_hi]`` at reset. + action_ptr: Ring-buffer write cursor as a one-element int array: the + column holding the latest raw action; advanced once per control + step. + delay_lo: Inclusive lower delay-range bound in control steps. + delay_hi: Inclusive upper delay-range bound; 0 disables the delay + path entirely. + """ + + default_angles: SharedArray + joint_lower: SharedArray + joint_upper: SharedArray + action_scales: SharedArray + delay_steps: np.ndarray + delay_lo: int + delay_hi: int + + +class JointPositionActionTerm(ActionTerm): + """Host runtime for joint-position targets and reset-time action delays.""" + + state: JointPositionActionState + + def __init__(self, space: gym.spaces.Box, state: JointPositionActionState) -> None: + super().__init__(space, state) + + def _process(self, actions: np.ndarray) -> np.ndarray: + """Map raw policy actions to targets using manager-owned action history.""" + state = self.state + ptr = int(state.action_ptr[0]) + target_raw = actions + if state.delay_hi > 0: + width = state.action_queue.shape[1] + delay_indices = (ptr - state.delay_steps) % width + target_raw = np.take_along_axis( + state.action_queue, + delay_indices[:, None, None], + axis=1, + )[:, 0, :] + return target_raw * state.action_scales + state.default_angles + + def _reset(self, env_ids: np.ndarray) -> None: + """Resample per-episode delay after the base term clears raw history.""" + state = self.state + if state.delay_hi > 0: + state.delay_steps[env_ids] = np.random.randint(state.delay_lo, state.delay_hi + 1, size=env_ids.size) + + +@configclass(kw_only=True) +class JointPositionActionCfg(ActionCfg): + """Joint-position targets from normalized policy actions. + + Attributes: + action_scale: Positive target-scaling factor in rad per unit action. + With ``action_scales_by_effort_limit_over_p_gain`` True each + actuator scales it by ``effort_limit / kp`` so torque-weak + actuators move more conservatively, instead of using the value + directly. + action_scales_by_effort_limit_over_p_gain: Select the uniform vs + effort/kp-normalized scaling mode. + action_delay_steps: Inclusive per-lane control-step delay range + ``(lo, hi)``, resampled at reset; ``(0, 0)`` disables delayed + target lookup but still retains current and previous raw actions. + """ + + action_scale: float = 0.25 + action_scales_by_effort_limit_over_p_gain: bool = False + action_delay_steps: tuple[int, int] = (0, 0) + + def __call__(self, env: ManagerEnv, actuators: tuple[ActuatorSpec, ...] | None) -> JointPositionActionTerm: + if actuators is None: + actuators = env.model.actuators + robot = env.cfg.scene.objs.robot + if not isinstance(robot, RobotCfg): + raise TypeError(f"scene robot must be RobotCfg, got {type(robot).__name__}") + if "default" not in robot.key_pose.poses: + raise ValueError("robot must define key pose 'default'") + kps = self._read_position_actuator_kps(env) + default_angles = self._resolve_default_pose( + env, + tuple(robot.resolve_name(name) for name in robot.key_pose.joint_names), + tuple(robot.key_pose.poses["default"]), + ) + action_scales = self._init_action_scales(env, kps) + joint_lower, joint_upper = env.model.others["robot_joint_position_limits"] + expected_model_shape = (env.num_actuators,) + if joint_lower.shape != expected_model_shape or joint_upper.shape != expected_model_shape: + raise ValueError( + "robot joint position limits must match the model actuator count: " + f"lower={joint_lower.shape}, upper={joint_upper.shape}, expected={expected_model_shape}." + ) + actuator_names = tuple(spec.name for spec in actuators) + all_names = tuple(spec.name for spec in env.model.actuators) + indices = np.asarray([all_names.index(name) for name in actuator_names], dtype=np.int64) + joint_lower = joint_lower[indices] + joint_upper = joint_upper[indices] + lo, hi = self.action_delay_steps + if lo < 0 or hi < lo: + raise ValueError(f"action_delay_steps must be 0 <= lo <= hi, got {self.action_delay_steps!r}") + # Keep current and previous actions even when delay is disabled. + width = max(hi + 1, 2) + action_queue = np.zeros((env.num_envs, width, len(actuators)), dtype=np.float32) + delay_steps = np.zeros(env.num_envs, dtype=np.int64) + action_ptr = np.zeros(1, dtype=np.int64) + state = JointPositionActionState( + action_queue=action_queue, + default_angles=default_angles[indices], + joint_lower=joint_lower, + joint_upper=joint_upper, + action_scales=action_scales[indices], + delay_steps=delay_steps, + action_ptr=action_ptr, + delay_lo=int(lo), + delay_hi=int(hi), + ) + ctrl_ranges = [] + for spec in actuators: + if spec.actuator_type is not ActuatorType.POSITION: + raise ValueError(f"actuator {spec.name!r} must be a position actuator, got {spec.actuator_type!r}") + if spec.ctrl_range is None: + raise ValueError(f"position actuator {spec.name!r} must define or inherit ctrl_range") + ctrl_ranges.append(spec.ctrl_range) + space = joint_position_action_space_from_ctrl_ranges( + np.asarray(ctrl_ranges, dtype=np.float32), state.default_angles, state.action_scales + ) + return JointPositionActionTerm(space, state) + + @staticmethod + def _resolve_default_pose( + env: ManagerEnv, + joint_names: tuple[str, ...], + joint_positions: tuple[float, ...], + ) -> np.ndarray: + if len(joint_names) != len(joint_positions): + raise ValueError( + f"default pose must contain one position per joint: {len(joint_names)} names, " + f"{len(joint_positions)} positions" + ) + actuator_joint_names = [] + for spec in env.model.actuators: + actuator_joint_names.append(spec.target_name) + positions = dict(zip(joint_names, joint_positions, strict=True)) + missing = sorted(set(actuator_joint_names).difference(positions)) + extra = sorted(set(positions).difference(actuator_joint_names)) + if missing or extra: + raise ValueError( + f"robot key pose 'default' must match actuator joint targets exactly: missing={missing}, extra={extra}" + ) + return np.asarray([positions[name] for name in actuator_joint_names], dtype=np.float32) + + @staticmethod + def _read_position_actuator_kps(env: ManagerEnv) -> np.ndarray: + return np.asarray(env.model.others["actuator_kp"], dtype=np.float32) + + @staticmethod + def _read_position_actuator_effort_limits(env: ManagerEnv) -> np.ndarray: + actuators = env.model.actuators + effort_limits = np.empty(len(actuators), dtype=np.float32) + for index, spec in enumerate(actuators): + if spec.force_range is None: + raise ValueError(f"actuator '{spec.name}' must define force_range") + force_range = np.asarray(spec.force_range, dtype=np.float32) + if force_range.shape != (2,) or not np.all(np.isfinite(force_range)): + raise ValueError(f"actuator '{spec.name}' force_range must contain two finite values") + effort_limit = float(np.max(np.abs(force_range))) + if effort_limit <= 0.0: + raise ValueError(f"actuator '{spec.name}' force_range must define a positive effort limit") + effort_limits[index] = effort_limit + return effort_limits + + def _init_action_scales(self, env: ManagerEnv, kps: np.ndarray) -> np.ndarray: + if not np.isfinite(self.action_scale) or self.action_scale <= 0.0: + raise ValueError(f"action_scale must be positive and finite, got {self.action_scale}") + if self.action_scales_by_effort_limit_over_p_gain: + effort = self._read_position_actuator_effort_limits(env) + safe_kp = np.where(kps == 0.0, 1.0, kps) + return np.where(kps == 0.0, 0.0, self.action_scale * effort / safe_kp).astype(np.float32) + return np.full(env.num_actuators, self.action_scale, dtype=np.float32) + + +__all__ = ["JointPositionActionState", "JointPositionActionTerm", "JointPositionActionCfg"] diff --git a/motrix_envs/src/motrix_envs/locomotion/action_space.py b/motrix_env_core/src/motrix_env_core/mdp/action_space.py similarity index 72% rename from motrix_envs/src/motrix_envs/locomotion/action_space.py rename to motrix_env_core/src/motrix_env_core/mdp/action_space.py index 38c82ebf..ca8c9528 100644 --- a/motrix_envs/src/motrix_envs/locomotion/action_space.py +++ b/motrix_env_core/src/motrix_env_core/mdp/action_space.py @@ -1,7 +1,7 @@ # Copyright Motphys Technology Co., Ltd. 2025, 2026 # SPDX-License-Identifier: Apache-2.0 -"""Action-space helpers shared by locomotion environments.""" +"""Action-space construction helpers for manager and direct environments.""" import gymnasium as gym import numpy as np @@ -19,6 +19,14 @@ def symmetric_residual_action_space( Bounds are divided by the per-actuator scale applied by the environment so the normalized action range maps exactly onto the control limits. Zero scales (inert actuators) get a zero bound instead of an infinite one. + + Args: + control_ranges: ``(2, A)`` lower/upper position-control limits. + default_values: Default (center) values, ``(A,)``. + action_scales: Per-actuator scales ``(A,)`` or one scalar for all. + + Returns: + Symmetric ``Box`` with bounds ``[-limit, +limit]`` per actuator. """ lower, upper = control_ranges scales = np.asarray(action_scales, dtype=np.float32) @@ -38,6 +46,14 @@ def asymmetric_residual_action_space( own lower/upper residual extent, so the action is not constrained to be zero-centered. This is useful when the default value sits far from the midpoint of the control range (e.g. heavily asymmetric joint limits). + + Args: + control_ranges: ``(2, A)`` lower/upper position-control limits. + default_values: Default (center) values, ``(A,)``. + action_scales: Per-actuator scales ``(A,)`` or one scalar for all. + + Returns: + Asymmetric ``Box`` with independent per-actuator low/high bounds. """ lower, upper = control_ranges scales = np.asarray(action_scales, dtype=np.float32) @@ -51,7 +67,7 @@ def asymmetric_residual_action_space( def joint_position_action_space( - actuators, + actuators: tuple, default_angles: np.ndarray, action_scales: float | np.ndarray, actuator_indices: np.ndarray | None = None, @@ -63,6 +79,16 @@ def joint_position_action_space( by the environment. Pass a scalar for uniform scales. Position actuators must declare ``ctrl_range`` directly or inherit it from their target joint during model construction. + + Args: + actuators: Actuator specs in canonical model order. + default_angles: Default angles for the selected actuators. + action_scales: Per-actuator scales or one scalar for all. + actuator_indices: Indices of the actuators to include; ``None`` + selects every actuator. + + Returns: + Symmetric ``Box`` with bounds ``[-limit, +limit]`` per actuator. """ if actuator_indices is None: actuator_indices = np.arange(len(actuators), dtype=np.int64) @@ -90,7 +116,19 @@ def joint_position_action_space_from_ctrl_ranges( action_scales: float | np.ndarray, actuator_indices: np.ndarray | None = None, ) -> gym.spaces.Box: - """Build symmetric joint-position action bounds from ``(num_actuators, 2)`` ctrl ranges.""" + """Build symmetric joint-position action bounds from ctrl ranges. + + Args: + ctrl_ranges: ``(num_actuators, 2)`` lower/upper control limits. + default_angles: Default angles, ``(num_actuators,)`` (after any + subsetting by ``actuator_indices``). + action_scales: Per-actuator scales or one scalar for all. + actuator_indices: Indices of the actuators to include; ``None`` + selects every row. + + Returns: + Symmetric ``Box`` with bounds ``[-limit, +limit]`` per actuator. + """ if actuator_indices is None: actuator_indices = np.arange(ctrl_ranges.shape[0], dtype=np.int64) default_angles = np.asarray(default_angles, dtype=np.float32) diff --git a/motrix_env_core/src/motrix_env_core/mdp/observations.py b/motrix_env_core/src/motrix_env_core/mdp/observations.py index b493653b..1a6888d6 100644 --- a/motrix_env_core/src/motrix_env_core/mdp/observations.py +++ b/motrix_env_core/src/motrix_env_core/mdp/observations.py @@ -58,7 +58,7 @@ def body_joint_vel_obs( def actions_obs(ctx: ManagerContext, out: np.ndarray, action_name: str) -> None: action_name = literally(action_name) action = ctx.actions[action_name] - out[:] = action.current + out[:] = action.current() @configclass(kw_only=True) @@ -68,8 +68,8 @@ class ActionsObsCfg(ObservationTermCfg): action_name: str = "joint_position" def __call__(self, ctx: BuildContext) -> ObsTerm: - action = ctx.action_terms[self.action_name] - return ObsTerm(action.current.shape[1], actions_obs, self.action_name) + action = ctx.action_terms[self.action_name].state + return ObsTerm(action.action_queue.shape[2], actions_obs, self.action_name) @dispatch diff --git a/motrix_env_core/src/motrix_env_core/mdp/rewards.py b/motrix_env_core/src/motrix_env_core/mdp/rewards.py index dde318a6..5504a4fa 100644 --- a/motrix_env_core/src/motrix_env_core/mdp/rewards.py +++ b/motrix_env_core/src/motrix_env_core/mdp/rewards.py @@ -38,7 +38,7 @@ def __call__(self, ctx) -> RewardTerm: def action_rate_reward(ctx: ManagerContext, action_name: str) -> float: action_name = literally(action_name) action = ctx.actions[action_name] - delta = action.current - action.previous + delta = action.current() - action.previous() return float(np.dot(delta, delta)) diff --git a/motrix_env_core/src/motrix_env_core/numba/fingerprint.py b/motrix_env_core/src/motrix_env_core/numba/fingerprint.py new file mode 100644 index 00000000..4db78dfa --- /dev/null +++ b/motrix_env_core/src/motrix_env_core/numba/fingerprint.py @@ -0,0 +1,48 @@ +# Copyright Motphys Technology Co., Ltd. 2025, 2026 +# SPDX-License-Identifier: Apache-2.0 + +"""Function dependency fingerprints shared by Numba compilation paths.""" + +import hashlib +import inspect +import marshal +from collections.abc import Callable +from typing import Any + +import numpy as np + + +def function_fingerprint(function: Callable[..., Any]) -> str: + """Hash function code, defaults, closures, and referenced global helpers. + + Dispatch entries and kernel-data methods use the same dependency traversal. + Numba dispatchers are unwrapped to their Python functions. Code objects include + constants and nested code, even when source is unavailable. Cycles are visited + once; cheaply representable global/closure constants are hashed by value. + Callers remain responsible for discovering member-method dependencies and + incorporating their fingerprints into the appropriate compilation cache key. + """ + hasher = hashlib.sha256() + seen: set[int] = set() + queue = [function] + while queue: + current = queue.pop() + current = getattr(current, "py_func", current) + if not inspect.isfunction(current) or id(current) in seen: + continue + seen.add(id(current)) + code = current.__code__ + hasher.update(f"{current.__module__}.{current.__qualname__}\0".encode()) + hasher.update(marshal.dumps(code)) + references = [(name, current.__globals__[name]) for name in code.co_names if name in current.__globals__] + references.extend((f"default[{i}]", value) for i, value in enumerate(current.__defaults__ or ())) + references.extend(sorted((current.__kwdefaults__ or {}).items())) + if current.__closure__: + references.extend(zip(code.co_freevars, (cell.cell_contents for cell in current.__closure__))) + for name, referenced in references: + helper = getattr(referenced, "py_func", referenced) + if inspect.isfunction(helper): + queue.append(helper) + elif isinstance(referenced, (int, float, bool, str, complex, np.generic, tuple, type(None))): + hasher.update(f"{name}={referenced!r}\0".encode()) + return hasher.hexdigest() diff --git a/motrix_env_core/src/motrix_env_core/numba/kernel_data/lowering.py b/motrix_env_core/src/motrix_env_core/numba/kernel_data/lowering.py index 079fc262..fb56f528 100644 --- a/motrix_env_core/src/motrix_env_core/numba/kernel_data/lowering.py +++ b/motrix_env_core/src/motrix_env_core/numba/kernel_data/lowering.py @@ -11,10 +11,12 @@ import numpy as np +from motrix_env_core.numba.fingerprint import function_fingerprint from motrix_env_core.numba.kernel_data.map import map_proxy +from motrix_env_core.numba.kernel_data.methods import dispatch_methods, register_proxy_method from motrix_env_core.numba.kernel_data.tree import KernelMapDef, LeafDef, TreeClassDef -_LOWERING_SCHEMA_VERSION = 1 +_LOWERING_SCHEMA_VERSION = 2 _LAYOUT_CACHE: dict[ tuple[str, bool], KernelRecordLayout, @@ -191,6 +193,21 @@ def proxy_symbol(proxy: type[tuple[Any, ...]]) -> str: return proxy.__name__ +def _tree_method_fingerprint(tree_def: TreeClassDef | KernelMapDef | LeafDef) -> Any: + """Include nested records and map values before consulting the layout cache.""" + if isinstance(tree_def, LeafDef): + return () + if isinstance(tree_def, KernelMapDef): + return tuple((entry.key, _tree_method_fingerprint(entry.tree_def)) for entry in tree_def.entries) + return ( + tuple( + (name, function_fingerprint(method)) + for name, method in sorted(dispatch_methods(tree_def.logical_type).items()) + ), + tuple((field.name, _tree_method_fingerprint(field.tree_def)) for field in tree_def.fields), + ) + + class KernelDataLowering: """Lower one logical ``TreeClassDef`` into a fixed flat Numba ABI layout.""" @@ -201,7 +218,8 @@ def lower( context: str, force_shared: bool = False, ) -> KernelDataLayout: - cache_key = (tree_def.fingerprint, force_shared) + method_key = hashlib.sha256(repr(_tree_method_fingerprint(tree_def)).encode()).hexdigest() + cache_key = (f"{tree_def.fingerprint}:{method_key}", force_shared) cached = _LAYOUT_CACHE.get(cache_key) if cached is not None: return cached @@ -254,7 +272,13 @@ def _lower_record( field_layouts.append(KernelFieldLayout(field.name, child)) fingerprint_fields.append((field.name, self._fingerprint_part(child))) identity = f"{tree_def.logical_type.__module__}.{tree_def.logical_type.__qualname__}" - raw_fingerprint = repr((_LOWERING_SCHEMA_VERSION, tree_def.fingerprint, identity, tuple(fingerprint_fields))) + method_fingerprints = tuple( + (name, function_fingerprint(method)) + for name, method in sorted(dispatch_methods(tree_def.logical_type).items()) + ) + raw_fingerprint = repr( + (_LOWERING_SCHEMA_VERSION, tree_def.fingerprint, identity, tuple(fingerprint_fields), method_fingerprints) + ) fingerprint = hashlib.sha256(raw_fingerprint.encode()).hexdigest() lowered_type = self._lowered_proxy( tree_def.logical_type, @@ -331,12 +355,14 @@ def _fingerprint_part(layout: KernelDataLayout) -> Any: return ( "map", layout.tree_def.path, + layout.fingerprint, tuple((entry.key, KernelDataLowering._fingerprint_part(entry.child)) for entry in layout.entries), ) assert isinstance(layout, KernelRecordLayout) return ( "record", f"{layout.logical_type.__module__}.{layout.logical_type.__qualname__}", + layout.fingerprint, tuple((field.name, KernelDataLowering._fingerprint_part(field.child)) for field in layout.fields), ) @@ -366,6 +392,8 @@ def _lowered_proxy( lowered.__module__ = logical_type.__module__ lowered.__qualname__ = name setattr(module, name, lowered) + for method_name, method in dispatch_methods(logical_type).items(): + register_proxy_method(lowered, method_name, method) return lowered @staticmethod diff --git a/motrix_env_core/src/motrix_env_core/numba/kernel_data/methods.py b/motrix_env_core/src/motrix_env_core/numba/kernel_data/methods.py new file mode 100644 index 00000000..6dcc4b7b --- /dev/null +++ b/motrix_env_core/src/motrix_env_core/numba/kernel_data/methods.py @@ -0,0 +1,51 @@ +# Copyright Motphys Technology Co., Ltd. 2025, 2026 +# SPDX-License-Identifier: Apache-2.0 + +from __future__ import annotations + +import inspect +from collections.abc import Callable +from typing import Any + +from numba import types +from numba.extending import overload_method + +from motrix_env_core.numba.manager.dispatch import _DISPATCH_MARKER + +_METHOD_OVERLOADS: dict[str, dict[type, Callable[..., Any]]] = {} + + +def dispatch_methods(logical_type: type) -> dict[str, Callable[..., Any]]: + """Resolve marked instance methods with normal Python MRO shadowing.""" + methods: dict[str, Callable[..., Any]] = {} + seen: set[str] = set() + for base in logical_type.__mro__: + for name, value in vars(base).items(): + if name in seen: + continue + seen.add(name) + if inspect.isfunction(value) and getattr(value, _DISPATCH_MARKER, False): + methods[name] = value + return methods + + +def register_proxy_method(proxy: type, name: str, method: Callable[..., Any]) -> None: + """Register one marked KernelData method for a lowered proxy type.""" + methods = _METHOD_OVERLOADS.get(name) + if methods is not None: + methods[proxy] = method + return + methods = {proxy: method} + _METHOD_OVERLOADS[name] = methods + + def getter(self: Any, *args: Any, **kwargs: Any) -> Any: + return methods.get(self.instance_class) + + # Register one getter per name: competing attribute templates for the same + # tuple method can hide later proxies. Homogeneous records (including + # single-field records) use NamedUniTuple. + for tuple_type in (types.NamedTuple, types.NamedUniTuple): + overload_method(tuple_type, name, strict=False, inline="always")(getter) + + +__all__ = ["dispatch_methods", "register_proxy_method"] diff --git a/motrix_env_core/src/motrix_env_core/numba/manager/actions.py b/motrix_env_core/src/motrix_env_core/numba/manager/actions.py index cc6edcfc..2513ae3b 100644 --- a/motrix_env_core/src/motrix_env_core/numba/manager/actions.py +++ b/motrix_env_core/src/motrix_env_core/numba/manager/actions.py @@ -11,26 +11,77 @@ import numpy as np from motrix_env_core.config import configclass +from motrix_env_core.numba.kernel_data import SharedArray, canonicalize_kernel_data, kernel_data +from motrix_env_core.numba.manager.dispatch import dispatch if TYPE_CHECKING: from motrix_env_core.numba.manager.env import ManagerEnv from motrix_env_core.sim.model import ActuatorSpec +@kernel_data +class ActionState: + """Manager-owned raw-action history shared by every action type. + + Attributes: + action_queue: Raw policy actions, host ``(N, W, A)`` and kernel + lane ``(W, A)``; ``W >= 2`` retains current and previous actions. + action_ptr: Shared one-element integer cursor holding the current + slot. Only the manager advances it; partial reset leaves it intact. + """ + + action_queue: np.ndarray + action_ptr: SharedArray + + @dispatch + def current(self) -> np.ndarray: + """Return current raw actions: host ``(N, A)``, kernel lane ``(A,)``. + + The returned view shares history storage and does not advance the cursor. + """ + return self.action_queue[..., self.action_ptr[0], :] + + @dispatch + def previous(self) -> np.ndarray: + """Return previous raw actions: host ``(N, A)``, kernel lane ``(A,)``. + + The returned view wraps around the ring without copying data. + """ + ptr = (self.action_ptr[0] - 1) % self.action_queue.shape[-2] + return self.action_queue[..., ptr, :] + + class ActionTerm(abc.ABC): - """Environment-local KernelData action pipeline with persistent runtime state.""" + """Extensible host-side action pipeline with kernel-bound state.""" - @abc.abstractmethod - def action_space(self, env: ManagerEnv, actuators: tuple[ActuatorSpec, ...] | None) -> gym.spaces.Box: - """Return the unbatched action space.""" + def __init__(self, action_space: gym.spaces.Box, state: ActionState) -> None: + if not isinstance(action_space, gym.spaces.Box): + raise TypeError("ActionTerm action_space must be a gym.spaces.Box.") + if not isinstance(state, ActionState): + raise TypeError("ActionTerm state must be an ActionState.") + self.action_space = action_space + self.state = canonicalize_kernel_data(state, context="ActionTerm state") - @abc.abstractmethod def process(self, actions: np.ndarray) -> np.ndarray | None: - """Process one action batch and return route-local actuator controls.""" + """Record raw actions and process one batched action slice.""" + state = self.state + ptr = (int(state.action_ptr[0]) + 1) % state.action_queue.shape[1] + state.action_ptr[0] = ptr + state.action_queue[:, ptr] = actions + return self._process(actions) @abc.abstractmethod + def _process(self, actions: np.ndarray) -> np.ndarray | None: + """Process one batched action slice into route-local controls.""" + def reset(self, env_ids: np.ndarray) -> None: - """Reset persistent action state for selected environments.""" + """Clear raw history and reset action-specific state for selected rows.""" + self.state.action_queue[env_ids] = 0.0 + self._reset(env_ids) + + @abc.abstractmethod + def _reset(self, env_ids: np.ndarray) -> None: + """Reset action-specific state for selected environment rows.""" @configclass(kw_only=True) @@ -38,10 +89,11 @@ class ActionCfg(abc.ABC): """Configuration that creates one environment-local action term.""" actuator_names: tuple[str, ...] | None = None + """Actuator route: ``None`` selects all actuators; ``()`` selects none.""" @abc.abstractmethod def __call__(self, env: ManagerEnv, actuators: tuple[ActuatorSpec, ...] | None) -> ActionTerm: - """Create the concrete KernelData action term.""" + """Assemble the action space, persistent state, and host callbacks.""" @configclass diff --git a/motrix_env_core/src/motrix_env_core/numba/manager/compiler/compiler.py b/motrix_env_core/src/motrix_env_core/numba/manager/compiler/compiler.py index 333c4eea..e4b8de34 100644 --- a/motrix_env_core/src/motrix_env_core/numba/manager/compiler/compiler.py +++ b/motrix_env_core/src/motrix_env_core/numba/manager/compiler/compiler.py @@ -18,6 +18,7 @@ import numpy as np from numba.extending import register_jitable +from motrix_env_core.numba.fingerprint import function_fingerprint from motrix_env_core.numba.kernel import clone_kernel_value from motrix_env_core.numba.kernel_data import ( KernelDataLayout, @@ -44,7 +45,6 @@ ) from motrix_env_core.numba.manager.compiler.codegen import KernelSourceGenerator from motrix_env_core.numba.manager.compiler.fingerprint import ( - function_fingerprint, plan_key, type_name, ) @@ -392,7 +392,7 @@ def _sim_input_layout(self) -> tuple[SimInputLayout, ...]: def _resolve_runtime_terms(self) -> None: for name, action_term in self._env.action_terms.items(): - _, tree_def = flatten_kernel_data(action_term) + _, tree_def = flatten_kernel_data(action_term.state) self._shared_parts.append( ( "action", @@ -422,7 +422,7 @@ def _resolve_manager_context( reward_terms: dict[str, RewardTerm], termination_terms: dict[str, TerminationTerm], ) -> ResolvedManagerContext: - actions = dict(self._env.action_terms) + actions = {name: term.state for name, term in self._env.action_terms.items()} sim = {key: binding.value for key, binding in self._sim_inputs.items()} commands = dict(self._env.command_terms) value = ManagerContext( diff --git a/motrix_env_core/src/motrix_env_core/numba/manager/compiler/fingerprint.py b/motrix_env_core/src/motrix_env_core/numba/manager/compiler/fingerprint.py index 1f60aba5..cbf5222b 100644 --- a/motrix_env_core/src/motrix_env_core/numba/manager/compiler/fingerprint.py +++ b/motrix_env_core/src/motrix_env_core/numba/manager/compiler/fingerprint.py @@ -1,16 +1,11 @@ # Copyright Motphys Technology Co., Ltd. 2025, 2026 # SPDX-License-Identifier: Apache-2.0 -"""Plan-key and term-source fingerprinting for the manager kernel compiler.""" +"""Plan-key helpers for the manager kernel compiler.""" import hashlib -import inspect -from collections.abc import Callable from typing import Any -import numba -import numpy as np - def plan_key(parts: list[Any]) -> str: return hashlib.sha256(repr(parts).encode()).hexdigest() @@ -20,46 +15,4 @@ def type_name(value_type: type[Any]) -> str: return f"{value_type.__module__}.{value_type.__qualname__}" -def function_fingerprint(function: Callable[..., Any]) -> str: - """Fingerprint a dispatch entry and every function it transitively references. - - The fused kernel inlines the dispatch body plus any ``@njit(inline="always")`` - helpers it calls from its module globals. Hashing only the entry's own source - would leave helper edits invisible to the plan key, so stale compiled kernels - would be silently reused (issue #54). Referenced functions are resolved through - ``co_names`` and hashed recursively; module-level non-function constants are - hashed by value when cheaply representable. - """ - hasher = hashlib.sha256() - seen: set[int] = set() - queue = [function] - while queue: - current = queue.pop() - if id(current) in seen: - continue - seen.add(id(current)) - if isinstance(current, numba.core.dispatcher.Dispatcher): - current = current.py_func - if not inspect.isfunction(current): - continue - code = current.__code__ - try: - source = inspect.getsource(current).encode() - except (OSError, TypeError): - source = code.co_code - hasher.update(f"{current.__module__}.{current.__qualname__}\0".encode()) - hasher.update(source) - hasher.update(b"\0") - module_globals = getattr(current, "__globals__", {}) - for name in code.co_names: - if name not in module_globals: - continue - referenced = module_globals[name] - if isinstance(referenced, (numba.core.dispatcher.Dispatcher,)) or inspect.isfunction(referenced): - queue.append(referenced) - elif isinstance(referenced, (int, float, bool, str, complex, np.generic)): - hasher.update(f"{name}={referenced!r}\0".encode()) - return hasher.hexdigest() - - -__all__ = ["function_fingerprint", "plan_key", "type_name"] +__all__ = ["plan_key", "type_name"] diff --git a/motrix_env_core/src/motrix_env_core/numba/manager/env.py b/motrix_env_core/src/motrix_env_core/numba/manager/env.py index 480368ad..ff729257 100644 --- a/motrix_env_core/src/motrix_env_core/numba/manager/env.py +++ b/motrix_env_core/src/motrix_env_core/numba/manager/env.py @@ -381,18 +381,22 @@ def __init__(self, cfg: EnvCfgType, num_envs: int = 1, backend: str | None = Non self._action_actuators = self._resolve_action_actuators() self._action_writes = self.sim.compile_writes( { - name: CtrlTargetsWrite(None if action_cfg.actuator_names == () else action_cfg.actuator_names) + name: CtrlTargetsWrite(action_cfg.actuator_names) for name, action_cfg in self._action_cfgs.items() - if action_cfg.actuator_names is not None + if action_cfg.actuator_names != () } ) self._action_terms: dict[str, ActionTerm] = { - name: canonicalize_kernel_data( - action_cfg(self, self._action_actuators[name]), - context=f"Manager action {name!r} __call__()", - ) - for name, action_cfg in self._action_cfgs.items() + name: action_cfg(self, self._action_actuators[name]) for name, action_cfg in self._action_cfgs.items() } + for name, term in self._action_terms.items(): + if not isinstance(term, ActionTerm): + raise TypeError(f"Manager action {name!r} must return an ActionTerm.") + queue = term.state.action_queue + if queue.ndim != 3 or queue.shape[0] != self.num_envs or queue.shape[1] < 2: + raise ValueError(f"Manager action {name!r} state.action_queue must have shape (N, W, A), W >= 2.") + if queue.shape[2] != term.action_space.shape[0]: + raise ValueError(f"Manager action {name!r} action dimension must match action space.") self._command_cfgs = cfg.command_cfgs() self._command_terms: dict[str, CommandTerm] = { name: canonicalize_kernel_data( @@ -554,7 +558,7 @@ def _resolve_action_actuators(self) -> dict[str, tuple[ActuatorSpec, ...] | None owners = {} by_name = {spec.name: spec for spec in self.model.actuators} for term_name, cfg in self._action_cfgs.items(): - if cfg.actuator_names is None: + if cfg.actuator_names == (): routes[term_name] = None continue names = cfg.actuator_names or tuple(by_name) @@ -586,7 +590,7 @@ def _build_action_space(self) -> tuple[gym.spaces.Box, dict[str, slice]]: action_slices = {} offset = 0 for name, term in self._action_terms.items(): - space = term.action_space(self, self._action_actuators[name]) + space = term.action_space if not isinstance(space, gym.spaces.Box): raise TypeError(f"Manager action {name!r} must produce a gym.spaces.Box.") if len(space.shape) != 1 or space.shape[0] <= 0: diff --git a/motrix_env_core/tests/test_action_state.py b/motrix_env_core/tests/test_action_state.py new file mode 100644 index 00000000..49cd5ab1 --- /dev/null +++ b/motrix_env_core/tests/test_action_state.py @@ -0,0 +1,98 @@ +# Copyright Motphys Technology Co., Ltd. 2025, 2026 +# SPDX-License-Identifier: Apache-2.0 + +"""Host accessors for manager-owned raw-action history.""" + +import gymnasium as gym +import numpy as np +import pytest + +from motrix_env_core.mdp.action import JointPositionActionState, JointPositionActionTerm +from motrix_env_core.numba.kernel_data import canonicalize_kernel_data, kernel_data +from motrix_env_core.numba.manager.actions import ActionState + + +@kernel_data +class _ActionState(ActionState): + scale: float + + +@pytest.mark.parametrize("ptr", [0, 1, 2]) +def test_action_state_accessors_return_history_views(ptr: int) -> None: + queue = np.arange(18, dtype=np.float32).reshape(2, 3, 3) + state = canonicalize_kernel_data( + _ActionState(action_queue=queue, action_ptr=np.asarray([ptr], dtype=np.int64), scale=1.0), + context="action state test", + ) + history = state.action_queue.copy() + pointer = state.action_ptr.copy() + + current = state.current() + previous = state.previous() + + np.testing.assert_array_equal(current, history[:, ptr]) + np.testing.assert_array_equal(previous, history[:, (ptr - 1) % 3]) + assert current.shape == previous.shape == (2, 3) + assert np.shares_memory(current, state.action_queue) + assert np.shares_memory(previous, state.action_queue) + np.testing.assert_array_equal(state.action_queue, history) + np.testing.assert_array_equal(state.action_ptr, pointer) + + current[0, 0] = -1.0 + assert state.action_queue[0, ptr, 0] == -1.0 + + +def test_joint_position_action_applies_per_lane_delay() -> None: + state = JointPositionActionState( + action_queue=np.asarray( + [ + [[10.0], [20.0], [30.0], [40.0]], + [[1.0], [2.0], [3.0], [4.0]], + ], + dtype=np.float32, + ), + action_ptr=np.asarray([2], dtype=np.int64), + default_angles=np.asarray([100.0], dtype=np.float32), + joint_lower=np.asarray([-100.0], dtype=np.float32), + joint_upper=np.asarray([100.0], dtype=np.float32), + action_scales=np.asarray([2.0], dtype=np.float32), + delay_steps=np.asarray([1, 2], dtype=np.int64), + delay_lo=0, + delay_hi=2, + ) + term = JointPositionActionTerm(gym.spaces.Box(-np.inf, np.inf, shape=(1,)), state) + + np.testing.assert_array_equal(term._process(np.asarray([[0.0], [0.0]], dtype=np.float32)), [[140.0], [102.0]]) + + +def test_joint_position_action_resamples_delay_on_reset(monkeypatch) -> None: + state = JointPositionActionState( + action_queue=np.zeros((3, 2, 1), dtype=np.float32), + action_ptr=np.asarray([0], dtype=np.int64), + default_angles=np.zeros(1, dtype=np.float32), + joint_lower=np.zeros(1, dtype=np.float32), + joint_upper=np.ones(1, dtype=np.float32), + action_scales=np.ones(1, dtype=np.float32), + delay_steps=np.zeros(3, dtype=np.int64), + delay_lo=1, + delay_hi=3, + ) + term = JointPositionActionTerm(gym.spaces.Box(-np.inf, np.inf, shape=(1,)), state) + monkeypatch.setattr(np.random, "randint", lambda lo, hi, size: np.full(size, lo + 1, dtype=np.int64)) + + term._reset(np.asarray([0, 2], dtype=np.int64)) + + np.testing.assert_array_equal(term.state.delay_steps, [2, 0, 2]) + + +def test_action_state_accessors_follow_cursor_changes() -> None: + state = ActionState( + action_queue=np.arange(12, dtype=np.float32).reshape(2, 2, 3), + action_ptr=np.zeros(1, dtype=np.int64), + ) + initial_current = state.current().copy() + initial_previous = state.previous().copy() + state.action_ptr[0] = 1 + + np.testing.assert_array_equal(state.current(), initial_previous) + np.testing.assert_array_equal(state.previous(), initial_current) diff --git a/motrix_env_core/tests/test_function_fingerprint.py b/motrix_env_core/tests/test_function_fingerprint.py new file mode 100644 index 00000000..7bbda8e0 --- /dev/null +++ b/motrix_env_core/tests/test_function_fingerprint.py @@ -0,0 +1,69 @@ +# Copyright Motphys Technology Co., Ltd. 2025, 2026 +# SPDX-License-Identifier: Apache-2.0 + +import numba + +from motrix_env_core.numba.fingerprint import function_fingerprint + + +def test_source_unavailable_constants_change_fingerprint(): + namespace = {"__name__": "fingerprint_test"} + exec("def entry(): return 1", namespace) + first = function_fingerprint(namespace["entry"]) + exec("def entry(): return 2", namespace) + assert function_fingerprint(namespace["entry"]) != first + + +def test_helper_changes_and_cycles(): + namespace = {"__name__": "fingerprint_test"} + exec("def helper(): return entry() + 1\ndef entry(): return helper()", namespace) + entry = namespace["entry"] + first = function_fingerprint(entry) + assert function_fingerprint(entry) == first + exec("def helper(): return entry() + 2", namespace) + assert function_fingerprint(entry) != first + + +def test_closure_constants_and_helpers(): + def make(value): + def entry(): + return value + + return entry + + assert function_fingerprint(make(1)) != function_fingerprint(make(2)) + + def helper(): + return 1 + + entry = make(helper) + first = function_fingerprint(entry) + helper.__code__ = (lambda: 2).__code__ + assert function_fingerprint(entry) != first + + +def test_defaults_keyword_defaults_and_helper_dependencies(): + def entry(value=1, *, scale=2): + return value * scale + + first = function_fingerprint(entry) + entry.__defaults__ = (3,) + second = function_fingerprint(entry) + assert second != first + entry.__kwdefaults__ = {"scale": 4} + assert function_fingerprint(entry) != second + + def helper(): + return 1 + + entry.__defaults__ = (helper,) + first = function_fingerprint(entry) + helper.__code__ = (lambda: 2).__code__ + assert function_fingerprint(entry) != first + + +def test_numba_dispatcher_matches_python_function(): + def helper(value): + return value + 1 + + assert function_fingerprint(numba.njit(helper)) == function_fingerprint(helper) diff --git a/motrix_env_core/tests/test_kernel_data_methods.py b/motrix_env_core/tests/test_kernel_data_methods.py new file mode 100644 index 00000000..5403b88a --- /dev/null +++ b/motrix_env_core/tests/test_kernel_data_methods.py @@ -0,0 +1,275 @@ +# Copyright Motphys Technology Co., Ltd. 2025, 2026 +# SPDX-License-Identifier: Apache-2.0 + +import json +import os +import subprocess +import sys +from pathlib import Path + +import numba +import numpy as np +import pytest + +from motrix_env_core.numba.kernel_data import ( + KernelDataLowering, + KernelDataScope, + SharedArray, + flatten_kernel_data, + iter_layout_leaves, + kernel_data, + rebuild_lowered, +) +from motrix_env_core.numba.manager.dispatch import dispatch + + +@kernel_data +class _MethodData: + values: np.ndarray + bias: SharedArray + scale: np.float32 + + @dispatch + def total(self): + return self.values.sum() * self.scale + self.bias[0] + + @dispatch + def scaled(self): + return self.values * self.scale + + @dispatch + def combined(self, factor): + return self.total() * factor + + def host_only(self): + return self.values.sum() + + +@kernel_data +class _InheritedMethodData(_MethodData): + extra: np.float32 + + @dispatch + def with_extra(self): + return self.total() + self.extra + + +@kernel_data +class _OverriddenMethodData(_InheritedMethodData): + @dispatch + def total(self): + return self.values.sum() * self.scale - self.bias[0] + + +def _lower_lane(value, env_id): + """Use the same leaf scopes as manager lowering, retaining lane views.""" + leaves, tree_def = flatten_kernel_data(value) + layout = KernelDataLowering().lower(tree_def, context="kernel data methods test") + lane_leaves = tuple( + leaves[leaf.slot_index][env_id] if leaf.scope is KernelDataScope.PER_ENV else leaves[leaf.slot_index] + for leaf in iter_layout_leaves(layout) + ) + return layout, rebuild_lowered(layout, lane_leaves) + + +def _read_total(value): + # Kept at module scope so the subprocess test has a stable disk-cache locator. + return value.total() + + +def _method_value(data_type=_MethodData): + args = ( + np.asarray([[1.0, 2.0], [4.0, 8.0]], dtype=np.float32), + np.asarray([0.5], dtype=np.float32), + np.float32(2.0), + ) + return data_type(*args) if data_type is _MethodData else data_type(*args, np.float32(3.0)) + + +def test_marked_methods_compile_on_lowered_lanes_and_preserve_host_methods() -> None: + value = _method_value() + original_method = _MethodData.scaled + layout, lane = _lower_lane(value, 1) + assert layout.logical_type is _MethodData + assert np.shares_memory(lane.values, value.values) + assert lane.bias is value.bias + + read_total = numba.njit(_read_total) + + @numba.njit + def scaled(kernel_value): + return kernel_value.scaled() + + assert read_total(lane) == value.values[1].sum() * value.scale + value.bias[0] + np.testing.assert_array_equal(scaled(lane), value.values[1] * value.scale) + assert read_total.nopython_signatures + assert scaled.nopython_signatures + + # Lowering must not replace host Python functions or their batched behavior. + assert _MethodData.scaled is original_method + assert not isinstance(_MethodData.scaled, numba.core.registry.CPUDispatcher) + np.testing.assert_array_equal(value.scaled(), value.values * value.scale) + assert value.total() == value.values.sum() * value.scale + value.bias[0] + assert value.host_only() == value.values.sum() + + value.values[1, 0] = 7.0 + assert read_total(lane) == value.values[1].sum() * value.scale + value.bias[0] + + +@pytest.mark.parametrize("data_type", [_InheritedMethodData, _OverriddenMethodData]) +def test_marked_methods_follow_inheritance_and_overrides(data_type) -> None: + value = _method_value(data_type) + _, lane = _lower_lane(value, 0) + + @numba.njit + def read_methods(kernel_value): + return kernel_value.total(), kernel_value.with_extra(), kernel_value.scaled() + + total, with_extra, scaled = read_methods(lane) + expected = value.values[0].sum() * value.scale + expected += -value.bias[0] if data_type is _OverriddenMethodData else value.bias[0] + assert total == expected + assert with_extra == expected + value.extra + np.testing.assert_array_equal(scaled, value.values[0] * value.scale) + assert read_methods.nopython_signatures + + +@pytest.mark.parametrize("data_type", [_MethodData, _OverriddenMethodData]) +def test_marked_method_can_call_another_marked_method(data_type) -> None: + value = _method_value(data_type) + _, lane = _lower_lane(value, 1) + + @numba.njit + def combined(kernel_value, factor): + return kernel_value.combined(factor) + + expected = value.values[1].sum() * value.scale + expected += -value.bias[0] if data_type is _OverriddenMethodData else value.bias[0] + assert combined(lane, np.float32(4.0)) == expected * np.float32(4.0) + assert combined.nopython_signatures + + +def test_unmarked_methods_remain_host_only() -> None: + value = _method_value() + _, lane = _lower_lane(value, 0) + + @numba.njit + def host_only(kernel_value): + return kernel_value.host_only() + + with pytest.raises(numba.TypingError, match="host_only"): + host_only(lane) + assert value.host_only() == value.values.sum() + + +def test_marked_method_cannot_call_an_unmarked_method() -> None: + @kernel_data + class _CallsHostMethod(_MethodData): + @dispatch + def call_host(self): + return self.host_only() + + base = _method_value() + value = _CallsHostMethod(base.values, base.bias, base.scale) + _, lane = _lower_lane(value, 0) + + @numba.njit + def call_host(kernel_value): + return kernel_value.call_host() + + with pytest.raises(numba.TypingError, match="host_only"): + call_host(lane) + assert value.call_host() == value.values.sum() + + +def test_unmarked_override_does_not_inherit_a_marked_base_method() -> None: + @kernel_data + class _HostOverride(_MethodData): + def total(self): + return self.values.sum() - self.bias[0] + + base = _method_value() + value = _HostOverride(base.values, base.bias, base.scale) + _, lane = _lower_lane(value, 0) + read_total = numba.njit(_read_total) + with pytest.raises(numba.TypingError, match="total"): + read_total(lane) + assert value.total() == value.values.sum() - value.bias[0] + + +def test_marked_method_implementation_changes_layout_fingerprint(monkeypatch) -> None: + @kernel_data + class _MutableMethod: + values: np.ndarray + + @dispatch + def total(self): + return self.values.sum() + 1.0 + + value = _MutableMethod(np.asarray([[2.0, 3.0]], dtype=np.float32)) + leaves, tree_def = flatten_kernel_data(value) + lowering = KernelDataLowering() + first = lowering.lower(tree_def, context="before method replacement") + first_lane = rebuild_lowered(first, (leaves[0][0],)) + + @dispatch + def replacement(self): + return self.values.sum() + 2.0 + + monkeypatch.setattr(_MutableMethod, "total", replacement) + _, changed_tree = flatten_kernel_data(value) + second = lowering.lower(changed_tree, context="after method replacement") + second_lane = rebuild_lowered(second, (leaves[0][0],)) + assert first.fingerprint != second.fingerprint + assert first.lowered_type is not second.lowered_type + read_total = numba.njit(_read_total) + assert read_total(first_lane) == 6.0 + assert read_total(second_lane) == 7.0 + assert value.total() == 7.0 + + +def test_marked_methods_load_numba_disk_cache_in_a_second_process(tmp_path) -> None: + # Import the real test module rather than duplicating its classes in a script: + # dynamically generated proxies need a stable importable module for unpickling. + script = """ +import importlib +import json +import numba + +module = importlib.import_module('test_kernel_data_methods') +value = module._method_value() +layout, lane = module._lower_lane(value, 1) +read_total = numba.njit(cache=True)(module._read_total) +result = read_total(lane) +print(json.dumps({ + 'result': result, + 'fingerprint': layout.fingerprint, + 'hits': sum(read_total.stats.cache_hits.values()), + 'misses': sum(read_total.stats.cache_misses.values()), +})) +""" + env = os.environ.copy() + env["NUMBA_CACHE_DIR"] = str(tmp_path / "numba-cache") + env["PYTHONPATH"] = os.pathsep.join(filter(None, (str(Path(__file__).parent), env.get("PYTHONPATH")))) + env["NUMBA_NUM_THREADS"] = "1" + + def run_process(): + completed = subprocess.run( + [sys.executable, "-c", script], + env=env, + capture_output=True, + text=True, + timeout=120, + check=False, + ) + assert completed.returncode == 0, completed.stdout + completed.stderr + return json.loads(completed.stdout.splitlines()[-1]) + + cold = run_process() + warm = run_process() + assert cold["result"] == warm["result"] == 24.5 + assert cold["fingerprint"] == warm["fingerprint"] + assert cold["misses"] == 1 + assert cold["hits"] == 0 + assert warm["misses"] == 0 + assert warm["hits"] == 1 diff --git a/motrix_env_core/tests/test_manager_sim_backend.py b/motrix_env_core/tests/test_manager_sim_backend.py index b5bcc2e4..11a831bd 100644 --- a/motrix_env_core/tests/test_manager_sim_backend.py +++ b/motrix_env_core/tests/test_manager_sim_backend.py @@ -13,7 +13,7 @@ from motrix_env_core.config import configclass from motrix_env_core.config.scene import SceneCfg from motrix_env_core.numba.kernel_data import kernel_data -from motrix_env_core.numba.manager.actions import ActionCfg, ActionTerm, ManagerActionsCfg +from motrix_env_core.numba.manager.actions import ActionCfg, ActionState, ActionTerm, ManagerActionsCfg from motrix_env_core.numba.manager.context import ManagerContext from motrix_env_core.numba.manager.dispatch import dispatch from motrix_env_core.numba.manager.env import ManagerBasedEnvCfg, ManagerEnv @@ -200,19 +200,25 @@ def write_compiler(self): @kernel_data -class _CtrlAction(ActionTerm): +class _CtrlActionState(ActionState): source: np.ndarray - def action_space(self, env, actuator_indices) -> gym.spaces.Box: - del env, actuator_indices - return gym.spaces.Box(-1.0, 1.0, (1,), dtype=np.float32) - def process(self, actions: np.ndarray) -> np.ndarray: - self.source[...] = actions +class _CtrlActionTerm(ActionTerm): + def __init__(self, action_space: gym.spaces.Box, state: _CtrlActionState) -> None: + super().__init__(action_space, state) + + def _process(self, actions: np.ndarray) -> np.ndarray: + self.state.source[...] = actions return actions * np.float32(2.0) - def reset(self, env_ids: np.ndarray) -> None: - self.source[env_ids] = 0.0 + def _reset(self, env_ids: np.ndarray) -> None: + self.state.source[env_ids] = 0.0 + + +def _ctrl_action_space(env, actuator_indices) -> gym.spaces.Box: + del env, actuator_indices + return gym.spaces.Box(-1.0, 1.0, (1,), dtype=np.float32) @configclass(kw_only=True) @@ -220,8 +226,14 @@ class _CtrlActionCfg(ActionCfg): actuator_names: tuple[str, ...] | None = ("routed",) def __call__(self, env, actuator_indices): - del actuator_indices - return _CtrlAction(np.zeros((env.num_envs, 1), dtype=np.float32)) + return _CtrlActionTerm( + _ctrl_action_space(env, actuator_indices), + _CtrlActionState( + np.zeros((env.num_envs, 2, 1), dtype=np.float32), + np.zeros(1, dtype=np.int64), + np.zeros((env.num_envs, 1), dtype=np.float32), + ), + ) @dispatch diff --git a/motrix_env_core/tests/test_numba_manager.py b/motrix_env_core/tests/test_numba_manager.py index b18d5607..60a2f805 100644 --- a/motrix_env_core/tests/test_numba_manager.py +++ b/motrix_env_core/tests/test_numba_manager.py @@ -19,6 +19,7 @@ from motrix_env_core.config.scene import SceneCfg # noqa: E402 from motrix_env_core.manager import ( # noqa: E402 ActionCfg, + ActionState, ActionTerm, CommandCfg, CommandTerm, @@ -54,40 +55,50 @@ @kernel_data -class _TestAction(ActionTerm): +class _TestActionState(ActionState): reset_flags: np.ndarray source: np.ndarray - def action_space(self, env: ManagerEnv, actuator_indices: np.ndarray | None) -> gym.spaces.Box: - del env, actuator_indices - return gym.spaces.Box(-1.0, 1.0, (1,), dtype=np.float32) - - def process(self, actions: np.ndarray) -> None: - self.source[...] = actions - - def reset(self, env_ids: np.ndarray) -> None: - self.reset_flags.fill(False) - self.reset_flags[env_ids] = True - @kernel_data -class _SecondAction(ActionTerm): +class _SecondActionState(ActionState): reset_flags: np.ndarray source: np.ndarray - def action_space(self, env: ManagerEnv, actuator_indices: np.ndarray | None) -> gym.spaces.Box: - del env, actuator_indices - return gym.spaces.Box( - np.asarray([-2.0, -3.0], dtype=np.float32), - np.asarray([2.0, 3.0], dtype=np.float32), - ) - def process(self, actions: np.ndarray) -> None: - self.source[...] = actions +class _TestActionTerm(ActionTerm): + def __init__(self, action_space: gym.spaces.Box, state: _TestActionState) -> None: + super().__init__(action_space, state) + + def _process(self, actions: np.ndarray) -> None: + self.state.source[...] = actions + + def _reset(self, env_ids: np.ndarray) -> None: + self.state.reset_flags.fill(False) + self.state.reset_flags[env_ids] = True + - def reset(self, env_ids: np.ndarray) -> None: - self.reset_flags.fill(False) - self.reset_flags[env_ids] = True +class _SecondActionTerm(ActionTerm): + def __init__(self, action_space: gym.spaces.Box, state: _SecondActionState) -> None: + super().__init__(action_space, state) + + def _process(self, actions: np.ndarray) -> None: + self.state.source[...] = actions + + def _reset(self, env_ids: np.ndarray) -> None: + self.state.reset_flags.fill(False) + self.state.reset_flags[env_ids] = True + + +def _test_action_space(_: ManagerEnv, __: np.ndarray | None) -> gym.spaces.Box: + return gym.spaces.Box(-1.0, 1.0, (1,), dtype=np.float32) + + +def _second_action_space(_: ManagerEnv, __: np.ndarray | None) -> gym.spaces.Box: + return gym.spaces.Box( + np.asarray([-2.0, -3.0], dtype=np.float32), + np.asarray([2.0, 3.0], dtype=np.float32), + ) @kernel_data @@ -98,7 +109,7 @@ class _ObservationParams: @dispatch def _injected_observation(ctx: ManagerContext, out: np.ndarray, params: _ObservationParams) -> None: - action: _TestAction = ctx.actions["test"] + action: _TestActionState = ctx.actions["test"] command: _CounterCommand = ctx.commands["counter"] out[0] = action.source[0] * params.scale + params.offset[0] out[1] = command.command[0] @@ -136,7 +147,7 @@ def __call__(self, env: ManagerEnv) -> ObsTerm: @dispatch def _lane_observation(ctx: ManagerContext, out: np.ndarray) -> None: - action: _TestAction = ctx.actions["test"] + action: _TestActionState = ctx.actions["test"] out[0] = ctx.env_id + np.float32(action.reset_flags[0]) * 0.0 @@ -163,7 +174,7 @@ class _LaneObservationsCfg(ManagerObservationsCfg): @dispatch def _injected_reward(ctx: ManagerContext) -> float: - action: _TestAction = ctx.actions["test"] + action: _TestActionState = ctx.actions["test"] return action.source[0] @@ -179,7 +190,7 @@ def _injected_termination( ctx: ManagerContext, threshold: np.float32, ) -> bool: - action: _TestAction = ctx.actions["test"] + action: _TestActionState = ctx.actions["test"] ctx.metrics["source_at_termination"][0] = action.source[0] return action.source[0] >= threshold @@ -199,21 +210,33 @@ def __call__(self, env: ManagerEnv) -> TerminationTerm: @configclass(kw_only=True) class _TestActionCfg(ActionCfg): - def __call__(self, env: ManagerEnv, actuator_indices: np.ndarray | None) -> _TestAction: - del actuator_indices - return _TestAction( - np.zeros((env.num_envs, 1), dtype=bool), - np.zeros((env.num_envs, 1), dtype=np.float32), + actuator_names: tuple[str, ...] = () + + def __call__(self, env: ManagerEnv, actuator_indices: np.ndarray | None) -> ActionTerm: + return _TestActionTerm( + _test_action_space(env, actuator_indices), + _TestActionState( + np.zeros((env.num_envs, 2, 1), dtype=np.float32), + np.zeros(1, dtype=np.int64), + np.zeros((env.num_envs, 1), dtype=bool), + np.zeros((env.num_envs, 1), dtype=np.float32), + ), ) @configclass(kw_only=True) class _SecondActionCfg(ActionCfg): - def __call__(self, env: ManagerEnv, actuator_indices: np.ndarray | None) -> _SecondAction: - del actuator_indices - return _SecondAction( - np.zeros((env.num_envs, 1), dtype=bool), - np.zeros((env.num_envs, 2), dtype=np.float32), + actuator_names: tuple[str, ...] = () + + def __call__(self, env: ManagerEnv, actuator_indices: np.ndarray | None) -> ActionTerm: + return _SecondActionTerm( + _second_action_space(env, actuator_indices), + _SecondActionState( + np.zeros((env.num_envs, 2, 2), dtype=np.float32), + np.zeros(1, dtype=np.int64), + np.zeros((env.num_envs, 1), dtype=bool), + np.zeros((env.num_envs, 2), dtype=np.float32), + ), ) @@ -233,7 +256,7 @@ class _CounterCommand(CommandTerm): @dispatch def update(self, ctx: ManagerContext) -> None: - action: _TestAction = ctx.actions["test"] + action: _TestActionState = ctx.actions["test"] self.double[0] = 2.0 * action.source[0] def reset(self, ctx) -> None: @@ -396,9 +419,9 @@ def _counter_command(env: _ManagerEnv) -> _CounterCommand: def test_action_and_command_terms_are_environment_owned() -> None: env = _ManagerEnv(num_envs=3) - assert isinstance(env.action_terms["test"], _TestAction) + assert isinstance(env.action_terms["test"], ActionTerm) assert env.action_terms is env._action_terms - assert env.action_terms["test"].reset_flags.shape == (3, 1) + assert env.action_terms["test"].state.reset_flags.shape == (3, 1) assert _counter_command(env) is env.command_terms["counter"] assert not hasattr(env, "value_manager") assert not hasattr(env.cfg, "values") @@ -528,8 +551,8 @@ def test_manager_cfg_accepts_dict_groups_and_empty_commands() -> None: def test_multiple_action_terms_concatenate_spaces_and_receive_ordered_slices() -> None: @dispatch def _multiple_action_observation(ctx: ManagerContext, out: np.ndarray) -> None: - first: _TestAction = ctx.actions["test"] - second: _SecondAction = ctx.actions["second"] + first: _TestActionState = ctx.actions["test"] + second: _SecondActionState = ctx.actions["second"] out[0] = first.source[0] out[1:] = second.source @@ -568,8 +591,8 @@ class _MultipleActionsCfg(ManagerActionsCfg): assert env.action_space.shape == (3,) np.testing.assert_array_equal(env.action_space.low, [-1.0, -2.0, -3.0]) np.testing.assert_array_equal(env.action_space.high, [1.0, 2.0, 3.0]) - np.testing.assert_array_equal(env.action_terms["test"].source, actions[:, :1]) - np.testing.assert_array_equal(env.action_terms["second"].source, actions[:, 1:]) + np.testing.assert_array_equal(env.action_terms["test"].state.source, actions[:, :1]) + np.testing.assert_array_equal(env.action_terms["second"].state.source, actions[:, 1:]) env._refresh_sim_reads() env._execute_observe_kernel(env._kernel_inputs) np.testing.assert_array_equal(state.obs.policy, actions) @@ -580,8 +603,8 @@ class _MultipleActionsCfg(ManagerActionsCfg): state.terminated[:] = [False, True] env._reset_done_envs() - np.testing.assert_array_equal(env.action_terms["test"].reset_flags[:, 0], [False, True]) - np.testing.assert_array_equal(env.action_terms["second"].reset_flags[:, 0], [False, True]) + np.testing.assert_array_equal(env.action_terms["test"].state.reset_flags[:, 0], [False, True]) + np.testing.assert_array_equal(env.action_terms["second"].state.reset_flags[:, 0], [False, True]) def test_multiple_action_terms_validate_total_action_shape() -> None: @@ -659,18 +682,20 @@ def test_manager_context_is_injected_once_and_reused_across_all_term_kinds() -> state = env.init_state() action = env.action_terms["test"] - assert isinstance(action, _TestAction) + assert isinstance(action, ActionTerm) command = _counter_command(env) context_sources = [slot for slot in env.manager_layout.inputs if "manager_context" in slot.source] - assert len(context_sources) == 10 - assert [slot.scope.value for slot in context_sources] == ["per_env"] * 8 + ["shared", "per_env"] + assert len(context_sources) == 12 + scopes = {slot.source: slot.scope.value for slot in context_sources} + assert scopes["manager_context.actions.test.action_queue"] == "per_env" + assert scopes["manager_context.actions.test.action_ptr"] == "shared" assert [term.output_slice for term in env.manager_layout.observations["policy"].terms] == [ slice(0, 2), slice(2, 3), slice(3, 4), ] - action.source[:, 0] = [0.25, 0.75] + action.state.source[:, 0] = [0.25, 0.75] env.compute_transition(state) env.compute_observation(state) @@ -719,8 +744,8 @@ def test_command_evaluation_updates_state_only_in_evaluate_kernel() -> None: state = env.init_state() assert env._compiled_manager_program is not None action = env.action_terms["test"] - assert isinstance(action, _TestAction) - action.source[:, 0] = [1.0, 2.0] + assert isinstance(action, ActionTerm) + action.state.source[:, 0] = [1.0, 2.0] derived = _counter_command(env).double derived.fill(-1.0) @@ -748,7 +773,7 @@ def test_command_term_updates_lifecycle_and_selected_reset_ids() -> None: env._reset_done_envs() np.testing.assert_array_equal(command[:, 0], [1.0, -1.0, 3.0]) - np.testing.assert_array_equal(np.flatnonzero(action.reset_flags[:, 0]), [1]) + np.testing.assert_array_equal(np.flatnonzero(action.state.reset_flags[:, 0]), [1]) assert tuple(term_field.name for term_field in fields(env.cfg.commands)) == ("counter",) @@ -781,8 +806,8 @@ def test_post_reset_observation_does_not_run_command_evaluation() -> None: state = env.init_state() assert env._compiled_manager_program is not None action = env.action_terms["test"] - assert isinstance(action, _TestAction) - action.source[:, 0] = [0.25, 0.5, 0.75] + assert isinstance(action, ActionTerm) + action.state.source[:, 0] = [0.25, 0.5, 0.75] state.terminated[:] = [False, True, False] env._refresh_sim_reads() env._reset_done_envs() @@ -905,8 +930,8 @@ def test_numeric_values_reuse_compiled_plan_and_remain_environment_local() -> No assert compiled.task.evaluate_kernel is env._compiled_manager_program.task.evaluate_kernel np.testing.assert_allclose(compiled.task.reward_weights, [2.0]) action = env.action_terms["test"] - assert isinstance(action, _TestAction) - action.source[:] = 0.25 + assert isinstance(action, ActionTerm) + action.state.source[:] = 0.25 assert env._kernel_outputs is not None inputs = compiled.read_plan.read(env) compiled.task.observe_kernel(inputs, env._kernel_outputs) @@ -937,8 +962,8 @@ def test_scalar_term_args_and_ctrl_dt_share_one_compiled_plan() -> None: for env, scale, threshold in ((first, 3.0, 0.5), (second, 5.0, 0.1)): action = env.action_terms["test"] - assert isinstance(action, _TestAction) - action.source[:] = 0.25 + assert isinstance(action, ActionTerm) + action.state.source[:] = 0.25 state = env._state assert state is not None env.compute_observation(state) @@ -952,8 +977,8 @@ def test_scalar_term_args_and_ctrl_dt_share_one_compiled_plan() -> None: third.init_state() assert third.manager_layout.plan_keys == first.manager_layout.plan_keys action = third.action_terms["test"] - assert isinstance(action, _TestAction) - action.source[:] = 0.5 + assert isinstance(action, ActionTerm) + action.state.source[:] = 0.5 third.compute_transition(third._state) np.testing.assert_allclose(third._state.reward, 2.0 * 0.5 * 0.02) @@ -1064,7 +1089,7 @@ def __call__(self, env: ManagerEnv) -> RewardTerm: def test_dispatch_term_parameter_names_are_not_constrained() -> None: @dispatch def _renamed_reward(context: ManagerContext) -> float: - action: _TestAction = context.actions["test"] + action: _TestActionState = context.actions["test"] return action.source[0] * 2.0 @configclass(kw_only=True) diff --git a/motrix_env_core/tests/test_sim_write_dispatch.py b/motrix_env_core/tests/test_sim_write_dispatch.py index aa01b0da..4d5a4fee 100644 --- a/motrix_env_core/tests/test_sim_write_dispatch.py +++ b/motrix_env_core/tests/test_sim_write_dispatch.py @@ -141,8 +141,8 @@ def test_sim_write_compiler_dispatches_each_write_to_its_typed_compiler() -> Non KinematicBodyRotationWrite(("body",)).compile_with(compiler, "write") ActuatorKpWrite(("actuator",)).compile_with(compiler, "write") ActuatorDampingWrite(("actuator",)).compile_with(compiler, "write") - BodyMassWrite(("link",)).compile_with(compiler, "write") - BodyComWrite(("link",)).compile_with(compiler, "write") + BodyMassWrite(("body",)).compile_with(compiler, "write") + BodyComWrite(("body",)).compile_with(compiler, "write") GeomFrictionWrite(("geom",)).compile_with(compiler, "write") assert compiler.dispatched == [ diff --git a/motrix_envs/src/motrix_envs/locomotion/ball_balance/microduck.py b/motrix_envs/src/motrix_envs/locomotion/ball_balance/microduck.py index bba699c4..7bf15bd8 100644 --- a/motrix_envs/src/motrix_envs/locomotion/ball_balance/microduck.py +++ b/motrix_envs/src/motrix_envs/locomotion/ball_balance/microduck.py @@ -22,6 +22,9 @@ ManagerTerminationsCfg, SimQueriesCfg, ) +from motrix_env_core.mdp.action import ( + JointPositionActionCfg, +) from motrix_env_core.mdp.observations import ( ActionsObsCfg, BodyAngularVelocityObsCfg, @@ -67,10 +70,6 @@ BadOrientationTerminationCfg, BallEscapedTerminationCfg, ) -from motrix_envs.locomotion.wbt.mdp.action import ( - WbtControlCfg, - WbtJointPositionActionCfg, -) from motrix_envs.locomotion.wbt.mdp.observations import ( DofPosRelObsCfg, DofVelObsCfg, @@ -110,8 +109,9 @@ class BallBalanceSceneObjsCfg(StandardSceneObjsCfg): @configclass class ActionsCfg(ManagerActionsCfg): - joint_position: WbtJointPositionActionCfg = WbtJointPositionActionCfg( - control=WbtControlCfg(action_scale=0.5, action_scales_by_effort_limit_over_p_gain=False), + joint_position: JointPositionActionCfg = JointPositionActionCfg( + action_scale=0.5, + action_scales_by_effort_limit_over_p_gain=False, ) diff --git a/motrix_envs/src/motrix_envs/locomotion/humanoid/cfg.py b/motrix_envs/src/motrix_envs/locomotion/humanoid/cfg.py index 052bdc90..15613647 100644 --- a/motrix_envs/src/motrix_envs/locomotion/humanoid/cfg.py +++ b/motrix_envs/src/motrix_envs/locomotion/humanoid/cfg.py @@ -31,6 +31,7 @@ ManagerTerminationsCfg, SimQueriesCfg, ) +from motrix_env_core.mdp.action import JointPositionActionCfg from motrix_env_core.mdp.observations import ( ActionsObsCfg, BodyAngularVelocityObsCfg, @@ -66,7 +67,6 @@ PenaltyOrientationRewardCfg, PoseRewardCfg, ) -from motrix_envs.locomotion.wbt.mdp.action import WbtControlCfg, WbtJointPositionActionCfg from motrix_envs.robot import HumanoidRobotCfg @@ -94,9 +94,7 @@ class HumanoidWalkSceneCfg(StandardSceneCfg): class WalkActionsCfg(ManagerActionsCfg): """Position action term shared with the WBT task family.""" - joint_position: WbtJointPositionActionCfg = WbtJointPositionActionCfg( - control=WbtControlCfg(action_scale=0.5, action_scales_by_effort_limit_over_p_gain=False) - ) + joint_position: JointPositionActionCfg = JointPositionActionCfg(action_scale=0.5) @configclass diff --git a/motrix_envs/src/motrix_envs/locomotion/humanoid/walk_manager_mdp/rewards.py b/motrix_envs/src/motrix_envs/locomotion/humanoid/walk_manager_mdp/rewards.py index beab31a8..b8a7370c 100644 --- a/motrix_envs/src/motrix_envs/locomotion/humanoid/walk_manager_mdp/rewards.py +++ b/motrix_envs/src/motrix_envs/locomotion/humanoid/walk_manager_mdp/rewards.py @@ -80,7 +80,7 @@ def __call__(self, ctx) -> RewardTerm: @dispatch def penalty_action_rate_reward(ctx: ManagerContext) -> float: action = ctx.actions["joint_position"] - delta = action.current - action.previous + delta = action.current() - action.previous() walk: WalkCommand = ctx.commands["walk"] return float(np.dot(delta, delta)) * walk.penalty_scale[0] diff --git a/motrix_envs/src/motrix_envs/locomotion/quadruped/walk_np.py b/motrix_envs/src/motrix_envs/locomotion/quadruped/walk_np.py index 3a252459..a307c355 100644 --- a/motrix_envs/src/motrix_envs/locomotion/quadruped/walk_np.py +++ b/motrix_envs/src/motrix_envs/locomotion/quadruped/walk_np.py @@ -9,6 +9,7 @@ from motrix_env_core.array.env import ArrayEnvState, NpObs from motrix_env_core.base import ObsSpace from motrix_env_core.direct.env import DirectEnv +from motrix_env_core.mdp.action_space import asymmetric_residual_action_space from motrix_env_core.sim import ( ActuatorCtrlQuery, ActuatorKdQuery, @@ -37,7 +38,6 @@ CtrlTargetsWrite, GeomFrictionWrite, ) -from motrix_envs.locomotion.action_space import asymmetric_residual_action_space from motrix_envs.locomotion.quadruped.cfg import QuadrupedWalkEnvCfg from motrix_envs.locomotion.quadruped.velocity_command import RandomPlanarVelocityBinding from motrix_envs.robot import QuadrupedRobotCfg diff --git a/motrix_envs/src/motrix_envs/locomotion/wbt/cfg.py b/motrix_envs/src/motrix_envs/locomotion/wbt/cfg.py index 5a057caf..079c7dba 100644 --- a/motrix_envs/src/motrix_envs/locomotion/wbt/cfg.py +++ b/motrix_envs/src/motrix_envs/locomotion/wbt/cfg.py @@ -22,6 +22,7 @@ ManagerTerminationsCfg, SimQueriesCfg, ) +from motrix_env_core.mdp.action import JointPositionActionCfg from motrix_env_core.mdp.observations import ( ActionsObsCfg, BodyAngularVelocityObsCfg, @@ -39,9 +40,6 @@ JointPositionQuery, JointVelocityQuery, ) -from motrix_envs.locomotion.wbt.mdp.action import ( - WbtJointPositionActionCfg, -) from motrix_envs.locomotion.wbt.mdp.command import ( WbtMotionCommandCfg, ) @@ -84,7 +82,9 @@ class ActionsCfg(ManagerActionsCfg): """Typed action terms for WBT.""" - joint_position: WbtJointPositionActionCfg = WbtJointPositionActionCfg() + joint_position: JointPositionActionCfg = JointPositionActionCfg( + action_scales_by_effort_limit_over_p_gain=True, + ) @configclass diff --git a/motrix_envs/src/motrix_envs/locomotion/wbt/dex_evt.py b/motrix_envs/src/motrix_envs/locomotion/wbt/dex_evt.py index 7ed1e49b..08c470dc 100644 --- a/motrix_envs/src/motrix_envs/locomotion/wbt/dex_evt.py +++ b/motrix_envs/src/motrix_envs/locomotion/wbt/dex_evt.py @@ -11,14 +11,13 @@ from motrix_env_core.config import configclass from motrix_env_core.config.scene import FlatTerrainCfg, SystemCameraCfg from motrix_env_core.manager import ManagerEnv +from motrix_env_core.mdp.action import ( + JointPositionActionCfg, +) from motrix_env_core.mdp.rewards import ActionRateRewardCfg from motrix_env_core.sim import BodyLinkNetContactForceQuery from motrix_envs.config.scene import StandardSceneCfg, StandardSceneObjsCfg from motrix_envs.locomotion.wbt.cfg import ActionsCfg, CommandsCfg, RewardsCfg, TerminationsCfg, WbtEnvCfg -from motrix_envs.locomotion.wbt.mdp.action import ( - WbtControlCfg, - WbtJointPositionActionCfg, -) from motrix_envs.locomotion.wbt.mdp.command import ( WbtMotionCommandCfg, ) @@ -56,8 +55,9 @@ class DexEvtWbtEnvCfg(WbtEnvCfg): motion_files: InitVar[tuple[str, ...] | None] = None commands: CommandsCfg = CommandsCfg(motion=WbtMotionCommandCfg()) actions: ActionsCfg = ActionsCfg( - joint_position=WbtJointPositionActionCfg( - control=WbtControlCfg(action_scale=1.0, action_scales_by_effort_limit_over_p_gain=False), + joint_position=JointPositionActionCfg( + action_scale=1.0, + action_scales_by_effort_limit_over_p_gain=False, ), ) sim: SimCfg = SimCfg(dt=0.005, solver_iterations=6, solver_tolerance=0.0001) diff --git a/motrix_envs/src/motrix_envs/locomotion/wbt/k1.py b/motrix_envs/src/motrix_envs/locomotion/wbt/k1.py index 3f9692b4..26b7ea44 100644 --- a/motrix_envs/src/motrix_envs/locomotion/wbt/k1.py +++ b/motrix_envs/src/motrix_envs/locomotion/wbt/k1.py @@ -11,14 +11,13 @@ from motrix_env_core.config import configclass from motrix_env_core.config.scene import SystemCameraCfg from motrix_env_core.manager import ManagerEnv +from motrix_env_core.mdp.action import ( + JointPositionActionCfg, +) from motrix_env_core.mdp.rewards import ActionRateRewardCfg from motrix_env_core.sim import BodyLinkNetContactForceQuery from motrix_envs.config.scene import StandardSceneCfg, StandardSceneObjsCfg from motrix_envs.locomotion.wbt.cfg import ActionsCfg, CommandsCfg, TerminationsCfg, WbtEnvCfg -from motrix_envs.locomotion.wbt.mdp.action import ( - WbtControlCfg, - WbtJointPositionActionCfg, -) from motrix_envs.locomotion.wbt.mdp.command import ( WbtMotionCommandCfg, ) @@ -52,11 +51,12 @@ class K1WbtEnvCfg(WbtEnvCfg): motion_files: InitVar[tuple[str, ...] | None] = None commands: CommandsCfg = CommandsCfg(motion=WbtMotionCommandCfg()) actions: ActionsCfg = ActionsCfg( - joint_position=WbtJointPositionActionCfg( + joint_position=JointPositionActionCfg( # K1 motion clips span large arm/leg offsets from the walk handoff pose. # Direct position scaling keeps the full joint range reachable; # the MJCF actuator forceranges still enforce K1 torque limits. - control=WbtControlCfg(action_scale=1.0, action_scales_by_effort_limit_over_p_gain=False), + action_scale=1.0, + action_scales_by_effort_limit_over_p_gain=False, ), ) sim: SimCfg = SimCfg(dt=0.005, solver_iterations=6, solver_tolerance=1e-4) diff --git a/motrix_envs/src/motrix_envs/locomotion/wbt/mdp/action.py b/motrix_envs/src/motrix_envs/locomotion/wbt/mdp/action.py deleted file mode 100644 index d3261d99..00000000 --- a/motrix_envs/src/motrix_envs/locomotion/wbt/mdp/action.py +++ /dev/null @@ -1,153 +0,0 @@ -# Copyright Motphys Technology Co., Ltd. 2025, 2026 -# SPDX-License-Identifier: Apache-2.0 - -"""Action terms for manager-based whole-body tracking.""" - -import gymnasium as gym -import numpy as np - -from motrix_env_core.config import configclass -from motrix_env_core.config.scene import RobotCfg -from motrix_env_core.manager import ActionCfg, ActionTerm, ManagerEnv, SharedArray, kernel_data -from motrix_env_core.sim.model import ActuatorSpec, ActuatorType -from motrix_envs.locomotion.action_space import joint_position_action_space_from_ctrl_ranges - - -@configclass -class WbtControlCfg: - """Position-action scaling for WBT actuators.""" - - action_scale: float = 0.25 - action_scales_by_effort_limit_over_p_gain: bool = True - - -@kernel_data -class WbtJointPositionAction(ActionTerm): - """Persistent WBT action pipeline, history, and shared model data.""" - - current: np.ndarray - previous: np.ndarray - default_angles: SharedArray - joint_lower: SharedArray - joint_upper: SharedArray - action_scales: SharedArray - - def action_space(self, env: ManagerEnv, actuators: tuple[ActuatorSpec, ...] | None) -> gym.spaces.Box: - assert actuators is not None - ctrl_ranges = [] - for spec in actuators: - if spec.actuator_type is not ActuatorType.POSITION: - raise ValueError(f"actuator {spec.name!r} must be a position actuator, got {spec.actuator_type!r}") - if spec.ctrl_range is None: - raise ValueError(f"position actuator {spec.name!r} must define or inherit ctrl_range") - ctrl_ranges.append(spec.ctrl_range) - return joint_position_action_space_from_ctrl_ranges( - np.asarray(ctrl_ranges, dtype=np.float32), - self.default_angles, - self.action_scales, - ) - - def process(self, actions: np.ndarray) -> np.ndarray: - np.copyto(self.previous, self.current) - np.copyto(self.current, actions, casting="unsafe") - return self.current * self.action_scales + self.default_angles - - def reset(self, env_ids: np.ndarray) -> None: - self.current[env_ids] = 0.0 - self.previous[env_ids] = 0.0 - - -@configclass(kw_only=True) -class WbtJointPositionActionCfg(ActionCfg): - control: WbtControlCfg = WbtControlCfg() - actuator_names: tuple[str, ...] = () - - def __call__(self, env: ManagerEnv, actuators: tuple[ActuatorSpec, ...] | None) -> ActionTerm: - assert actuators is not None - robot = env.cfg.scene.objs.robot - if not isinstance(robot, RobotCfg): - raise TypeError(f"WBT scene robot must be RobotCfg, got {type(robot).__name__}") - if "default" not in robot.key_pose.poses: - raise ValueError("WBT robot must define key pose 'default'") - kps = self._read_position_actuator_kps(env) - default_angles = self._resolve_default_pose( - env, - tuple(robot.resolve_name(name) for name in robot.key_pose.joint_names), - tuple(robot.key_pose.poses["default"]), - ) - action_scales = self._init_action_scales(env, kps) - joint_lower, joint_upper = env.model.others["robot_joint_position_limits"] - expected_joint_shape = (len(actuators),) - if joint_lower.shape != expected_joint_shape or joint_upper.shape != expected_joint_shape: - raise ValueError( - "WBT robot joint position limits must match robot_dof_pos: " - f"lower={joint_lower.shape}, upper={joint_upper.shape}, dof_pos={expected_joint_shape}." - ) - actuator_names = tuple(spec.name for spec in actuators) - all_names = tuple(spec.name for spec in env.model.actuators) - indices = np.asarray([all_names.index(name) for name in actuator_names], dtype=np.int64) - shape = (env.num_envs, len(actuators)) - return WbtJointPositionAction( - current=np.zeros(shape, dtype=np.float32), - previous=np.zeros(shape, dtype=np.float32), - default_angles=default_angles[indices], - joint_lower=joint_lower, - joint_upper=joint_upper, - action_scales=action_scales[indices], - ) - - @staticmethod - def _resolve_default_pose( - env: ManagerEnv, - joint_names: tuple[str, ...], - joint_positions: tuple[float, ...], - ) -> np.ndarray: - if len(joint_names) != len(joint_positions): - raise ValueError( - f"default pose must contain one position per joint: {len(joint_names)} names, " - f"{len(joint_positions)} positions" - ) - actuator_joint_names = [] - for spec in env.model.actuators: - actuator_joint_names.append(spec.target_name) - positions = dict(zip(joint_names, joint_positions, strict=True)) - missing = sorted(set(actuator_joint_names).difference(positions)) - extra = sorted(set(positions).difference(actuator_joint_names)) - if missing or extra: - raise ValueError( - "WBT robot key pose 'default' must match actuator joint targets exactly: " - f"missing={missing}, extra={extra}" - ) - return np.asarray([positions[name] for name in actuator_joint_names], dtype=np.float32) - - @staticmethod - def _read_position_actuator_kps(env: ManagerEnv) -> np.ndarray: - return np.asarray(env.model.others["actuator_kp"], dtype=np.float32) - - @staticmethod - def _read_position_actuator_effort_limits(env: ManagerEnv) -> np.ndarray: - actuators = env.model.actuators - effort_limits = np.empty(len(actuators), dtype=np.float32) - for index, spec in enumerate(actuators): - if spec.force_range is None: - raise ValueError(f"WBT actuator '{spec.name}' must define force_range") - force_range = np.asarray(spec.force_range, dtype=np.float32) - if force_range.shape != (2,) or not np.all(np.isfinite(force_range)): - raise ValueError(f"WBT actuator '{spec.name}' force_range must contain two finite values") - effort_limit = float(np.max(np.abs(force_range))) - if effort_limit <= 0.0: - raise ValueError(f"WBT actuator '{spec.name}' force_range must define a positive effort limit") - effort_limits[index] = effort_limit - return effort_limits - - def _init_action_scales(self, env: ManagerEnv, kps: np.ndarray) -> np.ndarray: - if not np.isfinite(self.control.action_scale) or self.control.action_scale <= 0.0: - raise ValueError(f"action_scale must be positive and finite, got {self.control.action_scale}") - if self.control.action_scales_by_effort_limit_over_p_gain: - effort = self._read_position_actuator_effort_limits(env) - safe_kp = np.where(kps == 0.0, 1.0, kps) - return np.where(kps == 0.0, 0.0, self.control.action_scale * effort / safe_kp).astype(np.float32) - return np.full(env.num_actuators, self.control.action_scale, dtype=np.float32) - - -__all__ = ["WbtControlCfg", "WbtJointPositionAction", "WbtJointPositionActionCfg"] diff --git a/motrix_envs/src/motrix_envs/locomotion/wbt/mdp/reset.py b/motrix_envs/src/motrix_envs/locomotion/wbt/mdp/reset.py index 13321c7a..425eb1c2 100644 --- a/motrix_envs/src/motrix_envs/locomotion/wbt/mdp/reset.py +++ b/motrix_envs/src/motrix_envs/locomotion/wbt/mdp/reset.py @@ -15,6 +15,7 @@ ResetTerm, ResetTermCfg, ) +from motrix_env_core.mdp.action import JointPositionActionState from motrix_env_core.numba.kernel_data import Map from motrix_env_core.numba.manager.dispatch import dispatch from motrix_env_core.numba.math import quaternion as numba_quaternion @@ -26,7 +27,6 @@ JointPositionWrite, JointVelocityWrite, ) -from motrix_envs.locomotion.wbt.mdp.action import WbtJointPositionAction from motrix_envs.locomotion.wbt.mdp.command import WbtMotionCommand @@ -170,7 +170,7 @@ def _reset_body_dof_pos(ctx: ManagerContext, sim_writes: Map[np.ndarray], noise_ position = sim_writes["position"] velocity = sim_writes["velocity"] motion: WbtMotionCommand = ctx.commands["motion"] - action: WbtJointPositionAction = ctx.actions["joint_position"] + action: JointPositionActionState = ctx.actions["joint_position"] step = motion.steps[0] position[:] = motion.clip.joint_pos[step] velocity[:] = motion.clip.joint_vel[step] diff --git a/motrix_envs/src/motrix_envs/locomotion/wbt/mdp/rewards.py b/motrix_envs/src/motrix_envs/locomotion/wbt/mdp/rewards.py index 642db966..c2c25688 100644 --- a/motrix_envs/src/motrix_envs/locomotion/wbt/mdp/rewards.py +++ b/motrix_envs/src/motrix_envs/locomotion/wbt/mdp/rewards.py @@ -11,9 +11,9 @@ from motrix_env_core.config import configclass from motrix_env_core.manager import ManagerContext, RewardTerm, RewardTermCfg from motrix_env_core.manager.math.quaternion import rotation_distance +from motrix_env_core.mdp.action import JointPositionActionState from motrix_env_core.numba.kernel_data import SharedArray, kernel_data from motrix_env_core.numba.manager.dispatch import dispatch -from motrix_envs.locomotion.wbt.mdp.action import WbtJointPositionAction from motrix_envs.locomotion.wbt.mdp.command import WbtMotionCommand @@ -190,7 +190,7 @@ class DofLimitRewardCfg(RewardTermCfg): cap: float def __call__(self, ctx) -> RewardTerm: - action = cast(WbtJointPositionAction, ctx.action_terms["joint_position"]) + action = cast(JointPositionActionState, ctx.action_terms["joint_position"].state) params = DofLimitParams( midpoint=(action.joint_lower + action.joint_upper) * 0.5, half_range=(action.joint_upper - action.joint_lower) * 0.5, diff --git a/motrix_envs/src/motrix_envs/locomotion/wbt/mdp/terminations.py b/motrix_envs/src/motrix_envs/locomotion/wbt/mdp/terminations.py index 656e5d5b..4dcc7f19 100644 --- a/motrix_envs/src/motrix_envs/locomotion/wbt/mdp/terminations.py +++ b/motrix_envs/src/motrix_envs/locomotion/wbt/mdp/terminations.py @@ -9,8 +9,8 @@ from motrix_env_core.config import configclass from motrix_env_core.manager import ManagerContext, TerminationTerm, TerminationTermCfg +from motrix_env_core.mdp.action import JointPositionActionState from motrix_env_core.numba.manager.dispatch import dispatch -from motrix_envs.locomotion.wbt.mdp.action import WbtJointPositionAction from motrix_envs.locomotion.wbt.mdp.command import WbtMotionCommand @@ -143,7 +143,7 @@ def __call__(self, ctx) -> TerminationTerm: @dispatch def bad_dof_position_termination(ctx: ManagerContext, threshold: np.float32) -> bool: dof_pos = ctx.sim["robot_dof_pos"] - action: WbtJointPositionAction = ctx.actions["joint_position"] + action: JointPositionActionState = ctx.actions["joint_position"] error = 0.0 finite = True for joint_id in range(dof_pos.shape[0]): diff --git a/motrix_envs/tests/test_action_space.py b/motrix_envs/tests/test_action_space.py index 6be0ffdf..508c6984 100644 --- a/motrix_envs/tests/test_action_space.py +++ b/motrix_envs/tests/test_action_space.py @@ -7,8 +7,8 @@ import pytest from motrix_env_core import registry +from motrix_env_core.mdp.action_space import joint_position_action_space from motrix_env_core.sim.model import ActuatorType -from motrix_envs.locomotion.action_space import joint_position_action_space def test_joint_position_action_space_uses_symmetric_position_ranges(): diff --git a/motrix_envs/tests/test_mdp_obs.py b/motrix_envs/tests/test_mdp_obs.py index 762659e4..09257300 100644 --- a/motrix_envs/tests/test_mdp_obs.py +++ b/motrix_envs/tests/test_mdp_obs.py @@ -7,6 +7,7 @@ import pytest from motrix_env_core.config.scene import KeyPoseCfg, ModelFileCfg, RobotCfg +from motrix_env_core.mdp.action import JointPositionActionState from motrix_env_core.mdp.observations import ( ActionsObsCfg, BodyAngularVelocityObsCfg, @@ -25,7 +26,6 @@ LinkLinearVelocityQuery, LinkQuaternionQuery, ) -from motrix_envs.locomotion.wbt.mdp.action import WbtJointPositionAction from motrix_envs.locomotion.wbt.mdp.observations import ( DofPosRelObsCfg, ) @@ -102,22 +102,34 @@ def _env() -> SimpleNamespace: def test_actions_observation_reads_current_actions() -> None: - action = WbtJointPositionAction( - current=np.zeros((1, 3), dtype=np.float32), - previous=np.zeros((1, 3), dtype=np.float32), + action = JointPositionActionState( + action_queue=np.zeros((1, 2, 3), dtype=np.float32), default_angles=np.zeros(3, dtype=np.float32), joint_lower=np.zeros(3, dtype=np.float32), joint_upper=np.zeros(3, dtype=np.float32), action_scales=np.ones(3, dtype=np.float32), + delay_steps=np.zeros(1, dtype=np.int64), + action_ptr=np.zeros(1, dtype=np.int64), + delay_lo=0, + delay_hi=0, ) - env = SimpleNamespace(action_terms={"joint_position": action}) + # Host configuration sees an ActionTerm; kernel context sees its state. + env = SimpleNamespace(action_terms={"joint_position": SimpleNamespace(state=action)}) cfg = ActionsObsCfg() term = cfg.__call__(env) current = np.asarray([1.0, 2.0, 3.0], dtype=np.float32) out = np.empty(term.size, dtype=np.float32) assert term.size == 3 - ctx = SimpleNamespace(actions={"joint_position": SimpleNamespace(current=current)}) + # Manager kernels receive a lane view: queue shape is (W, A), not (N, W, A). + queue = np.zeros((2, 3), dtype=np.float32) + queue[1] = current + lane_state = SimpleNamespace( + current=lambda: queue[1], + action_queue=queue, + action_ptr=np.asarray([1], dtype=np.int64), + ) + ctx = SimpleNamespace(actions={"joint_position": lane_state}) term.dispatch(ctx, out, *term.args) np.testing.assert_array_equal(out, current) diff --git a/motrix_envs/tests/test_wbt_numba.py b/motrix_envs/tests/test_wbt_numba.py index ffccd185..1a15de9e 100644 --- a/motrix_envs/tests/test_wbt_numba.py +++ b/motrix_envs/tests/test_wbt_numba.py @@ -16,6 +16,10 @@ ManagerEnv, ManagerResetCfg, ) +from motrix_env_core.mdp.action import ( # noqa: E402 + JointPositionActionCfg, + JointPositionActionState, +) from motrix_env_core.mdp.observations import ( # noqa: E402 BodyAngularVelocityObsCfg, UniformNoiseCfg, @@ -33,10 +37,6 @@ from motrix_envs.locomotion.wbt.g1.common import MOTION_DIR as _G1_MOTION_DIR # noqa: E402 from motrix_envs.locomotion.wbt.g1.common import G1WbtEnvCfg # noqa: E402 from motrix_envs.locomotion.wbt.k1 import K1WbtEnvCfg # noqa: E402 -from motrix_envs.locomotion.wbt.mdp.action import ( # noqa: E402 - WbtJointPositionAction, - WbtJointPositionActionCfg, -) from motrix_envs.locomotion.wbt.mdp.command import ( # noqa: E402 WbtMotionCommand, WbtMotionCommandCfg, @@ -196,7 +196,7 @@ def test_numba_wbt_read_plan_reuses_preallocated_arrays() -> None: for first_value, second_value in zip(first, second, strict=True): assert first_value is second_value or np.shares_memory(first_value, second_value) or first_value.size == 0 - action = env.action_terms["joint_position"] + action = env.action_terms["joint_position"].state actions = np.full((env.num_envs, *env.action_space.shape), 0.25, dtype=np.float32) env.apply_action(actions, state) partial = env._compiled_manager_program.read_plan.read( @@ -204,15 +204,15 @@ def test_numba_wbt_read_plan_reuses_preallocated_arrays() -> None: np.asarray([1, 3], dtype=np.int64), ) assert partial is first - assert any(value is action.current for value in first) - np.testing.assert_array_equal(action.current, actions) + assert any(value is action.action_queue for value in first) + np.testing.assert_array_equal(action.action_queue[:, action.action_ptr[0]], actions) state.terminated[:] = [False, True, False, True] env._reset_done_envs() env._refresh_sim_reads() assert env._kernel_inputs is first - np.testing.assert_array_equal(action.current[[0, 2]], 0.25) - np.testing.assert_array_equal(action.current[[1, 3]], 0.0) + np.testing.assert_array_equal(action.action_queue[[0, 2], action.action_ptr[0]], 0.25) + np.testing.assert_array_equal(action.action_queue[[1, 3], action.action_ptr[0]], 0.0) def test_numba_wbt_step_preserves_previous_actor_and_critic_observations() -> None: @@ -361,15 +361,15 @@ def test_numba_wbt_clip_wrap_rematerializes_sim_only() -> None: env = _make_numba_env(_single_file_cfg(start_at_timestep_zero_prob=1.0), num_envs=2, seed=11) env.init_state() motion = _motion_command(env) - action = env.action_terms["joint_position"] - assert isinstance(action, WbtJointPositionAction) + action = env.action_terms["joint_position"].state + assert isinstance(action, JointPositionActionState) num_frames = motion.clip.joint_pos.shape[0] # 0.25 is exactly representable in float32, so equality checks stay exact. actions = np.full((env.num_envs, *env.action_space.shape), 0.25, dtype=np.float32) # Prime persistent action state and the episode counter. env.step(actions) - np.testing.assert_array_equal(action.current, 0.25) + np.testing.assert_array_equal(action.action_queue[:, action.action_ptr[0]], 0.25) episode_steps_before_wrap = env.state.episode_steps.copy() # Force every lane to wrap the clip on the next transition. @@ -384,9 +384,11 @@ def test_numba_wbt_clip_wrap_rematerializes_sim_only() -> None: np.testing.assert_array_equal(state.truncated, False) np.testing.assert_array_equal(state.episode_steps, episode_steps_before_wrap + 1) # Sim-only: the persistent action state (which an action reset would - # zero) keeps the processed values of this step. - np.testing.assert_array_equal(action.current, 0.25) - np.testing.assert_array_equal(action.previous, 0.25) + # zero) keeps the raw policy actions of this step. + np.testing.assert_array_equal(action.action_queue[:, action.action_ptr[0]], 0.25) + np.testing.assert_array_equal( + action.action_queue[:, (action.action_ptr[0] - 1) % action.action_queue.shape[1]], 0.25 + ) # The request flag is cleared before the next physics step: a follow-up # step neither rematerializes the lane nor rewinds its frame. @@ -447,17 +449,17 @@ def test_numba_wbt_masked_reset_preserves_bound_buffer_identity() -> None: buffers = env._kernel_buffers assert buffers is not None env._refresh_sim_reads() - action_value = env.action_terms["joint_position"] + action_value = env.action_terms["joint_position"].state motion = _motion_command(env) identities = { - "current_actions": id(action_value.current), - "last_actions": id(action_value.previous), + "action_queue": id(action_value.action_queue), + "action_ptr": id(action_value.action_ptr), "reward_terms": id(buffers[0]), "termination_masks": id(buffers[2]), "target_body_position_relative": id(motion.target_body_position_relative), "sim_inputs": tuple(id(env.sim_data[key]) for key in env.sim_data.keys), } - env.apply_action(np.ones_like(action_value.current), state) + env.apply_action(np.ones_like(action_value.action_queue[:, 0]), state) motion.steps[:, 0] = [1, 2, 3] robot_dof_pos = env.sim_data["robot_dof_pos"] non_reset_dof_pos = robot_dof_pos[1].copy() @@ -473,15 +475,17 @@ def test_numba_wbt_masked_reset_preserves_bound_buffer_identity() -> None: assert env._task_program is not None assert env._task_program.reset_kernel.nopython_signatures - assert id(action_value.current) == identities["current_actions"] - assert id(action_value.previous) == identities["last_actions"] + assert id(action_value.action_queue) == identities["action_queue"] + assert id(action_value.action_ptr) == identities["action_ptr"] assert id(buffers[0]) == identities["reward_terms"] assert id(buffers[2]) == identities["termination_masks"] assert id(motion.target_body_position_relative) == identities["target_body_position_relative"] assert tuple(id(env.sim_data[key]) for key in env.sim_data.keys) == identities["sim_inputs"] - np.testing.assert_array_equal(action_value.current[[0, 2]], 0.0) - np.testing.assert_array_equal(action_value.current[1], 1.0) - np.testing.assert_array_equal(action_value.previous[[0, 2]], 0.0) + ptr = int(action_value.action_ptr[0]) + previous_ptr = (ptr - 1) % action_value.action_queue.shape[1] + np.testing.assert_array_equal(action_value.action_queue[[0, 2], ptr], 0.0) + np.testing.assert_array_equal(action_value.action_queue[1, ptr], 1.0) + np.testing.assert_array_equal(action_value.action_queue[[0, 2], previous_ptr], 0.0) np.testing.assert_array_equal(motion.steps[1, 0], 2) np.testing.assert_allclose(robot_dof_pos[[0, 2]], motion.clip.joint_pos[motion.steps[[0, 2], 0]]) np.testing.assert_array_equal(robot_dof_pos[1], non_reset_dof_pos) @@ -491,38 +495,48 @@ def test_numba_wbt_masked_reset_preserves_bound_buffer_identity() -> None: np.testing.assert_array_equal(state.terminated, terminated) -def test_numba_wbt_action_term_owns_rolls_and_resets_action_buffers() -> None: +def test_numba_wbt_manager_rolls_and_resets_action_buffers() -> None: env = _make_numba_env(_deterministic_manager_cfg(), num_envs=2) state = env.init_state() - value = env.action_terms["joint_position"] - assert value is env._action_terms["joint_position"] + term = env.action_terms["joint_position"] + assert term is env._action_terms["joint_position"] + value = term.state - env.apply_action(np.ones_like(value.current), state) - np.testing.assert_array_equal(value.current, 1.0) - np.testing.assert_array_equal(value.previous, 0.0) + actions = np.ones((env.num_envs, *env.action_space.shape), dtype=np.float32) + env.apply_action(actions, state) + ptr = int(value.action_ptr[0]) + previous_ptr = (ptr - 1) % value.action_queue.shape[1] + np.testing.assert_array_equal(value.action_queue[:, ptr], 1.0) + np.testing.assert_array_equal(value.action_queue[:, previous_ptr], 0.0) - env.apply_action(np.full_like(value.current, 2.0), state) - np.testing.assert_array_equal(value.current, 2.0) - np.testing.assert_array_equal(value.previous, 1.0) + actions.fill(2.0) + env.apply_action(actions, state) + ptr = int(value.action_ptr[0]) + previous_ptr = (ptr - 1) % value.action_queue.shape[1] + np.testing.assert_array_equal(value.action_queue[:, ptr], 2.0) + np.testing.assert_array_equal(value.action_queue[:, previous_ptr], 1.0) state.terminated[:] = True + ptr_before_reset = value.action_ptr.copy() env._reset_done_envs() - np.testing.assert_array_equal(value.current, 0.0) - np.testing.assert_array_equal(value.previous, 0.0) + np.testing.assert_array_equal(value.action_queue, 0.0) + np.testing.assert_array_equal(value.action_ptr, ptr_before_reset) def test_numba_wbt_action_owns_shared_writable_model_data() -> None: env = _make_numba_env(_deterministic_manager_cfg(), num_envs=1) state = SimpleNamespace() actions = np.full((1, *env.action_space.shape), 0.25, dtype=np.float32) - value = env.action_terms["joint_position"] + value = env.action_terms["joint_position"].state env.apply_action(actions, state) - np.testing.assert_array_equal(value.current, actions) - np.testing.assert_array_equal(value.previous, 0.0) - assert value.current.flags.writeable - assert value.previous.flags.writeable + ptr = int(value.action_ptr[0]) + previous_ptr = (ptr - 1) % value.action_queue.shape[1] + np.testing.assert_array_equal(value.action_queue[:, ptr], actions) + np.testing.assert_array_equal(value.action_queue[:, previous_ptr], 0.0) + assert value.action_queue.flags.writeable + assert value.action_ptr.flags.writeable assert all( array.flags.writeable for array in ( @@ -566,7 +580,7 @@ def test_numba_wbt_registry_uses_generic_manager_env() -> None: assert motion_command_cfg.kernel_size == 1 assert motion_command_cfg.kernel_lambda == pytest.approx(0.8) action_cfg = manager_cfg.actions.joint_position - assert isinstance(action_cfg, WbtJointPositionActionCfg) + assert isinstance(action_cfg, JointPositionActionCfg) assert not hasattr(manager_cfg, "values") assert motion_command_cfg.motion_files == (str(_G1_MOTION_DIR / "dance"),) tracked_body_pos = env.sim_data.query("tracked_body_pos") @@ -578,7 +592,7 @@ def test_numba_wbt_registry_uses_generic_manager_env() -> None: assert not hasattr(motion_command_cfg, "robot") assert not hasattr(env, "value_manager") assert isinstance(manager_cfg.actions, ActionsCfg) - assert isinstance(manager_cfg.actions.joint_position, WbtJointPositionActionCfg) + assert isinstance(manager_cfg.actions.joint_position, JointPositionActionCfg) assert isinstance(manager_cfg.sim_reset, ManagerResetCfg) assert isinstance(manager_cfg.sim_reset.body_pos, BodyPosResetCfg) assert isinstance(manager_cfg.sim_reset.body_rot, BodyRotResetCfg) @@ -702,8 +716,8 @@ def test_numba_wbt_manager_builds_for_all_wbt_presets(env_name: str) -> None: assert env.cfg.queries.data["robot_dof_pos"].joints == env.cfg.commands.motion.joint_names assert env.sim_data.query("robot_dof_pos").joints == env.cfg.commands.motion.joint_names env.init_state() - action = env.action_terms["joint_position"] - assert isinstance(action, WbtJointPositionAction) + action = env.action_terms["joint_position"].state + assert isinstance(action, JointPositionActionState) joint_lower, joint_upper = env.model.others["robot_joint_position_limits"] np.testing.assert_array_equal(action.joint_lower, joint_lower) np.testing.assert_array_equal(action.joint_upper, joint_upper) @@ -888,8 +902,8 @@ def test_numba_wbt_multi_clip_wrap_rematerializes_sim_only(tmp_path) -> None: env = _make_numba_env(cfg, num_envs=2, seed=11) env.init_state() motion = _motion_command(env) - action = env.action_terms["joint_position"] - assert isinstance(action, WbtJointPositionAction) + action = env.action_terms["joint_position"].state + assert isinstance(action, JointPositionActionState) total = motion.clip.joint_pos.shape[0] actions = np.full((env.num_envs, *env.action_space.shape), 0.25, dtype=np.float32) @@ -914,8 +928,10 @@ def test_numba_wbt_multi_clip_wrap_rematerializes_sim_only(tmp_path) -> None: np.testing.assert_array_equal(state.terminated, False) np.testing.assert_array_equal(state.truncated, False) np.testing.assert_array_equal(state.episode_steps, episode_steps_before_wrap + 2) - np.testing.assert_array_equal(action.current, 0.25) - np.testing.assert_array_equal(action.previous, 0.25) + np.testing.assert_array_equal(action.action_queue[:, action.action_ptr[0]], 0.25) + np.testing.assert_array_equal( + action.action_queue[:, (action.action_ptr[0] - 1) % action.action_queue.shape[1]], 0.25 + ) steps_after_wrap = motion.steps[:, 0].copy() env.step(actions) @@ -951,8 +967,8 @@ def test_numba_wbt_multi_clip_sequential_crosses_boundaries_and_loops(tmp_path) env = _make_numba_env(cfg, num_envs=2, seed=11) env.init_state() motion = _motion_command(env) - action = env.action_terms["joint_position"] - assert isinstance(action, WbtJointPositionAction) + action = env.action_terms["joint_position"].state + assert isinstance(action, JointPositionActionState) total = motion.clip.joint_pos.shape[0] actions = np.full((env.num_envs, *env.action_space.shape), 0.25, dtype=np.float32) @@ -979,9 +995,11 @@ def test_numba_wbt_multi_clip_sequential_crosses_boundaries_and_loops(tmp_path) np.testing.assert_array_equal(state.truncated, False) np.testing.assert_array_equal(state.episode_steps, episode_steps_before_boundary + 2) # Rematerialization is sim-only: the persistent action state keeps the - # processed values of this step. - np.testing.assert_array_equal(action.current, 0.25) - np.testing.assert_array_equal(action.previous, 0.25) + # raw policy actions of this step. + np.testing.assert_array_equal(action.action_queue[:, action.action_ptr[0]], 0.25) + np.testing.assert_array_equal( + action.action_queue[:, (action.action_ptr[0] - 1) % action.action_queue.shape[1]], 0.25 + ) # The corpus-final frame loops back to the corpus head frame. motion.steps[:, 0] = total - 1 diff --git a/motrix_rl/tests/test_fastsac_ipc_ring.py b/motrix_rl/tests/test_fastsac_ipc_ring.py index 4eb14573..1d40eae5 100644 --- a/motrix_rl/tests/test_fastsac_ipc_ring.py +++ b/motrix_rl/tests/test_fastsac_ipc_ring.py @@ -120,6 +120,7 @@ def test_ipc_ring_backpressure_bounds_in_flight(): while receiver.has_next(): k, _ = receiver.read_span() receiver.commit_reads(k) + torch.cuda.synchronize() assert not owner.is_full() assert owner.push(*_batch(CAPACITY, seed=3)) diff --git a/wiki/design/manager/runtime.md b/wiki/design/manager/runtime.md index 4b7d24f1..02171510 100644 --- a/wiki/design/manager/runtime.md +++ b/wiki/design/manager/runtime.md @@ -89,10 +89,18 @@ reset 只处理指定的原始 `env_ids`: Term 不要求继承统一实现基类,但必须遵循对应 protocol: ```python +@kernel_data +class BaseActionState: + action_queue: np.ndarray # host (N, W, A), lane (W, A); W >= 2 + action_ptr: SharedArray # one shared cursor, advanced once per control step + + +@dataclass(frozen=True) class ActionTerm: - def action_space(env, actuator_indices): ... - def process(actions): ... - def reset(env_ids): ... + action_space: gym.spaces.Box + state: BaseActionState + process: Callable # process(state, actions_batch) -> controls_batch | None + reset: Callable # reset(state, env_ids) -> None class ManagerContext: @@ -143,8 +151,17 @@ buffer,每次 physics step 前清零)即可请求 sim-only reset:该 lane 该机制。host `reset(ResetContext)` 只在 episode reset 时准备跨 lane 的共享数据(如采样分布);`on_transition()` 每步折叠 统计。 -Action term 的 host `process()` 不直接写 simulator state;Manager 根据 actuator route 合并其输出。Term 不应调用其他 term 的行为方法, -跨 term 依赖应通过 `ManagerContext` 的数据 store 表达。 +ActionTerm 是 host 侧描述对象:静态 action space、canonicalized state 和批量 process/reset callback。 +只有 `term.state` 进入 kernel ABI,`ctx.actions[name]` 直接取得 lane-scoped state;host 通过 +`env.action_terms[name].state` 访问同一份 backing。callback 不进入 kernel data,不要求 `@dispatch`。 + +Manager 在调用 `process(state, actions)` 前推进共享 cursor 一次并写入 raw policy actions;处理回调只计算 +route-local controls,不重复推进 history,也不直接写 simulator state。每个 state 的 queue 至少保留两帧, +当前和上一帧分别为 `queue[ptr]` 与 `queue[(ptr - 1) % W]`,observation 和 action-rate reward 均读取 raw history。 +延迟 action 可增加 queue 宽度。episode reset 时 Manager 先清空选中 lanes 的整个 queue,再调用 +`reset(state, env_ids)`;共享 cursor 和其他 lanes 的 history 不变,sim-only reset 不清 action history。 +Manager 根据 actuator route 合并 controls 后写入 backend。Term 不应调用其他 term 的行为方法,跨 term +依赖通过 `ManagerContext` 的数据 store 表达。 Observation terms use the host-side ``ObservationTerm(size, dispatch, *args)`` form. ``size`` is the fixed output width, and additional arguments are passed positionally; simple scalar/array arguments do not require a separate Args class. Reward terms use the analogous ``RewardTerm(dispatch, *args)`` form and return one numeric scalar per environment. Termination terms use ``TerminationTerm(dispatch, *args, metrics=...)``; metrics are optional per-environment outputs exposed to the dispatch as a static ``Map[np.ndarray]``. The compiler validates and lowers these values into the static kernel ABI, then emits ``dispatch(ctx, out, *args)`` for observations, ``dispatch(ctx, *args)`` for rewards, and ``dispatch(ctx, metrics, *args)`` when termination metrics are present. diff --git a/wiki/design/manager/task-authoring.md b/wiki/design/manager/task-authoring.md index b4eb135d..28e4ea5c 100644 --- a/wiki/design/manager/task-authoring.md +++ b/wiki/design/manager/task-authoring.md @@ -126,16 +126,24 @@ Term 不持有 backend handle,也不在运行期创建或执行 simulator quer `ActionCfg.actuator_names` 声明 simulator control route,空 tuple 表示所有 actuator。Manager 在构造期解析并校验未知、重复和重叠 route。 -Action term 提供: +ActionCfg 在构造期返回冻结的 host 描述对象: ```python -action_space(env, actuator_indices) -> gym.spaces.Box -process(actions) -> np.ndarray -reset(env_ids) -> None +ActionTerm( + action_space=space, + state=state, # @kernel_data BaseActionState 子类 + process=process_action, # process_action(state, actions_batch) -> controls | None + reset=reset_action, # reset_action(state, env_ids) -> None +) ``` -`process()` 只处理分配给自己的 policy action slice,更新自己的 history,并返回 route-local controls。Manager 负责将各 term 输出合并并写入 -simulator;compiled term 可以读取 action data,但不能调用 action term 的 host 方法。 +`BaseActionState` 声明 `action_queue` 和 `action_ptr`,配置为每个环境分配独立 history。 +Manager 在 process 回调之前推进共享 cursor 一次并写入 raw action slice;回调只处理该 slice, +返回 route-local controls,不重复维护 history。episode reset 时 Manager 先清目标 lanes 的整个 queue, +再调用 reset 回调,cursor 保持不变。Manager 将各 term 输出合并并写入 simulator。 + +Host 通过 `env.action_terms[name].state` 访问状态;compiled term 通过 `ctx.actions[name]` 直接读取 +lane-scoped state,不能调用 process/reset callback。具体状态可增加 delay、目标缩放和模型参数字段。 ## 6. Observation、Reward、Termination