Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
4 changes: 0 additions & 4 deletions motrix_env_core/src/motrix_env_core/sim/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -58,8 +58,6 @@
BodyLinearVelocityWrite,
BodyPositionWrite,
BodyRotationWrite,
DofPositionWrite,
DofVelocityWrite,
JointPositionWrite,
JointVelocityWrite,
)
Expand Down Expand Up @@ -89,9 +87,7 @@
"BodyMassQuery",
"DofPositionLimitsQuery",
"DofPositionQuery",
"DofPositionWrite",
"DofVelocityQuery",
"DofVelocityWrite",
"GeomFrictionQuery",
"GeomLinearVelocityQuery",
"GeomSpec",
Expand Down
79 changes: 50 additions & 29 deletions motrix_env_core/src/motrix_env_core/sim/write.py
Original file line number Diff line number Diff line change
Expand Up @@ -31,22 +31,6 @@ def compile_with(self, compiler: SimWriteCompiler, name: str) -> None:
"""Record this declaration on the compiler through its typed hook."""


@dataclass(frozen=True)
class DofPositionWrite(SimWrite):
"""Complete canonical DOF position write."""

def compile_with(self, compiler: SimWriteCompiler, name: str) -> None:
compiler.compile_dof_position(name, self)


@dataclass(frozen=True)
class DofVelocityWrite(SimWrite):
"""Complete canonical DOF velocity write."""

def compile_with(self, compiler: SimWriteCompiler, name: str) -> None:
compiler.compile_dof_velocity(name, self)


@dataclass(frozen=True)
class BodyJointPositionWrite(SimWrite):
"""One body's articulated-joint position write."""
Expand Down Expand Up @@ -87,6 +71,29 @@ def compile_with(self, compiler: SimWriteCompiler, name: str) -> None:
compiler.compile_joint_velocity(name, self)


@dataclass(frozen=True)
class JointQuaternionWrite(SimWrite):
"""Local orientation quaternions (xyzw) of declared ball joints: ``(N, J, 4)``.

The backend normalizes each quaternion on write.
"""

joints: tuple[str, ...]

def compile_with(self, compiler: SimWriteCompiler, name: str) -> None:
compiler.compile_joint_quaternion(name, self)


@dataclass(frozen=True)
class JointAngularVelocityWrite(SimWrite):
"""Local angular velocities of declared ball joints: ``(N, J, 3)``."""

joints: tuple[str, ...]

def compile_with(self, compiler: SimWriteCompiler, name: str) -> None:
compiler.compile_joint_angular_velocity(name, self)


@dataclass(frozen=True)
class CtrlTargetsWrite(SimWrite):
"""Actuator ctrl targets in declared name order.
Expand Down Expand Up @@ -143,13 +150,23 @@ def compile_with(self, compiler: SimWriteCompiler, name: str) -> None:


@dataclass(frozen=True)
class MocapPoseWrite(SimWrite):
"""Mocap body poses in declared order: ``(N, B, 7)`` float32."""
class KinematicBodyPositionWrite(SimWrite):
"""Kinematic body world positions in declared order: ``(N, B, 3)`` float32."""

bodies: tuple[str, ...]

def compile_with(self, compiler: SimWriteCompiler, name: str) -> None:
compiler.compile_kinematic_body_position(name, self)


@dataclass(frozen=True)
class KinematicBodyRotationWrite(SimWrite):
"""Kinematic body world quaternions in declared order: ``(N, B, 4)`` float32."""

bodies: tuple[str, ...]

def compile_with(self, compiler: SimWriteCompiler, name: str) -> None:
compiler.compile_mocap_pose(name, self)
compiler.compile_kinematic_body_rotation(name, self)


@dataclass(frozen=True)
Expand Down Expand Up @@ -239,14 +256,6 @@ def _begin_compile(self) -> None:
def _build_program(self, *, reset: bool, forward_kinematics: bool) -> WriteProgram:
"""Assemble the program from the ops recorded during dispatch."""

@abc.abstractmethod
def compile_dof_position(self, name: str, write: DofPositionWrite) -> None:
"""Record a complete canonical DOF position write."""

@abc.abstractmethod
def compile_dof_velocity(self, name: str, write: DofVelocityWrite) -> None:
"""Record a complete canonical DOF velocity write."""

@abc.abstractmethod
def compile_body_joint_position(self, name: str, write: BodyJointPositionWrite) -> None:
"""Record one body's articulated DOF position write."""
Expand All @@ -263,6 +272,14 @@ def compile_joint_position(self, name: str, write: JointPositionWrite) -> None:
def compile_joint_velocity(self, name: str, write: JointVelocityWrite) -> None:
"""Record named one-DOF joint velocity writes."""

@abc.abstractmethod
def compile_joint_quaternion(self, name: str, write: JointQuaternionWrite) -> None:
"""Record ball-joint local orientation writes."""

@abc.abstractmethod
def compile_joint_angular_velocity(self, name: str, write: JointAngularVelocityWrite) -> None:
"""Record ball-joint local angular velocity writes."""

@abc.abstractmethod
def compile_ctrl_targets(self, name: str, write: CtrlTargetsWrite) -> None:
"""Record actuator control targets."""
Expand All @@ -284,8 +301,12 @@ def compile_body_angular_velocity(self, name: str, write: BodyAngularVelocityWri
"""Record floating-body world angular velocity writes."""

@abc.abstractmethod
def compile_mocap_pose(self, name: str, write: MocapPoseWrite) -> None:
"""Record mocap-body pose writes."""
def compile_kinematic_body_position(self, name: str, write: KinematicBodyPositionWrite) -> None:
"""Record kinematic-body world position writes."""

@abc.abstractmethod
def compile_kinematic_body_rotation(self, name: str, write: KinematicBodyRotationWrite) -> None:
"""Record kinematic-body world rotation writes."""

@abc.abstractmethod
def compile_actuator_kp(self, name: str, write: ActuatorKpWrite) -> None:
Expand Down
13 changes: 8 additions & 5 deletions motrix_env_core/tests/test_direct_env_sim_backend.py
Original file line number Diff line number Diff line change
Expand Up @@ -17,15 +17,14 @@
from motrix_env_core.sim import (
ActuatorCtrlQuery,
DofPositionQuery,
DofPositionWrite,
DofVelocityQuery,
ModelQuery,
PhysicsReadProgram,
)
from motrix_env_core.sim.backend import SimBackend
from motrix_env_core.sim.model import ActuatorSpec, ActuatorType, SimModel
from motrix_env_core.sim.registry import register_sim_backend
from motrix_env_core.sim.write import CtrlTargetsWrite, DofVelocityWrite, WriteProgram
from motrix_env_core.sim.write import CtrlTargetsWrite, JointPositionWrite, JointVelocityWrite, WriteProgram


def _core_model() -> SimModel:
Expand All @@ -47,9 +46,9 @@ def __init__(self, backend: "_FakeBackend", writes, reset: bool) -> None:
for name, write in writes.items():
if isinstance(write, CtrlTargetsWrite):
self._buffers[name] = np.zeros((backend.num_envs, backend.num_actuators), dtype=np.float32)
elif isinstance(write, DofPositionWrite):
elif isinstance(write, JointPositionWrite):
self._buffers[name] = np.zeros_like(backend.dof_pos)
elif isinstance(write, DofVelocityWrite):
elif isinstance(write, JointVelocityWrite):
self._buffers[name] = np.zeros_like(backend.dof_vel)

def buffer(self, name: str) -> np.ndarray:
Expand Down Expand Up @@ -189,7 +188,11 @@ def __init__(self, cfg: _FakeDirectCfg, num_envs: int, backend: str | None = Non
)
self._ctrl_writes = self.sim.compile_writes({"ctrl": CtrlTargetsWrite()})
self._reset_program = self.sim.compile_writes(
{"state_position": DofPositionWrite(), "state_velocity": DofVelocityWrite()}, reset=True
{
"state_position": JointPositionWrite(("j0", "j1")),
"state_velocity": JointVelocityWrite(("j0", "j1")),
},
reset=True,
)
self._action_space = gym.spaces.Box(-1.0, 1.0, (2,), dtype=np.float32)
self._observation_space = gym.spaces.Box(-np.inf, np.inf, (6,), dtype=np.float32)
Expand Down
6 changes: 1 addition & 5 deletions motrix_env_core/tests/test_manager_sim_backend.py
Original file line number Diff line number Diff line change
Expand Up @@ -38,7 +38,7 @@
from motrix_env_core.sim.backend import SimBackend
from motrix_env_core.sim.model import ActuatorSpec, ActuatorType, SimModel
from motrix_env_core.sim.registry import register_sim_backend
from motrix_env_core.sim.write import CtrlTargetsWrite, DofPositionWrite, DofVelocityWrite, WriteProgram
from motrix_env_core.sim.write import CtrlTargetsWrite, WriteProgram

_ACTUATORS = (
ActuatorSpec(
Expand Down Expand Up @@ -82,10 +82,6 @@ def __init__(self, backend: "_FakeBackend", writes, reset: bool) -> None:
)
self._routes[name] = route
self._buffers[name] = np.zeros((backend.num_envs, len(route)), dtype=np.float32)
elif isinstance(write, DofPositionWrite):
self._buffers[name] = np.zeros_like(backend.dof_pos)
elif isinstance(write, DofVelocityWrite):
self._buffers[name] = np.zeros_like(backend.dof_vel)

def buffer(self, name: str) -> np.ndarray:
return self._buffers[name]
Expand Down
11 changes: 5 additions & 6 deletions motrix_env_core/tests/test_numba_manager.py
Original file line number Diff line number Diff line change
Expand Up @@ -50,7 +50,7 @@
from motrix_env_core.numba.manager.observations import create_observation_groups
from motrix_env_core.numba.manager.rewards import create_reward_terms
from motrix_env_core.numba.manager.terminations import TerminationManager
from motrix_env_core.sim import DofPositionWrite
from motrix_env_core.sim.write import CtrlTargetsWrite


@kernel_data
Expand Down Expand Up @@ -352,14 +352,13 @@ def physics_step(self) -> None:

@dispatch
def _recording_reset(ctx: ManagerContext, sim_writes: Map[np.ndarray]) -> None:
dof_pos = sim_writes["dof_pos"]
dof_pos[:] = 0.0
sim_writes["ctrl"][:] = 0.0


@dispatch
def _noop_reset(ctx: ManagerContext, sim_writes: Map[np.ndarray]) -> None:
_ = ctx
sim_writes["dof_pos"][:] = 0.0
sim_writes["ctrl"][:] = 0.0


@configclass(kw_only=True)
Expand All @@ -368,14 +367,14 @@ class _DescriptorResetTermCfg(ResetTermCfg):

def __call__(self, env: ManagerEnv) -> ResetTerm:
del env
return ResetTerm(_noop_reset, writes={"dof_pos": DofPositionWrite()})
return ResetTerm(_noop_reset, writes={"ctrl": CtrlTargetsWrite()})


@configclass
class _RecordingResetTermCfg(ResetTermCfg):
def __call__(self, env: ManagerEnv) -> ResetTerm:
del env
return ResetTerm(_recording_reset, writes={"dof_pos": DofPositionWrite()})
return ResetTerm(_recording_reset, writes={"ctrl": CtrlTargetsWrite()})


@configclass
Expand Down
49 changes: 28 additions & 21 deletions motrix_env_core/tests/test_sim_write_dispatch.py
Original file line number Diff line number Diff line change
Expand Up @@ -17,12 +17,13 @@
BodyPositionWrite,
BodyRotationWrite,
CtrlTargetsWrite,
DofPositionWrite,
DofVelocityWrite,
GeomFrictionWrite,
JointAngularVelocityWrite,
JointPositionWrite,
JointQuaternionWrite,
JointVelocityWrite,
MocapPoseWrite,
KinematicBodyPositionWrite,
KinematicBodyRotationWrite,
SimWriteCompiler,
WriteProgram,
)
Expand All @@ -49,14 +50,6 @@ def _build_program(self, *, reset: bool, forward_kinematics: bool) -> WriteProgr
del reset, forward_kinematics
return _RecordingProgram()

def compile_dof_position(self, name, write) -> None:
del name, write
self.dispatched.append("dof_position")

def compile_dof_velocity(self, name, write) -> None:
del name, write
self.dispatched.append("dof_velocity")

def compile_body_joint_position(self, name, write) -> None:
del name, write
self.dispatched.append("body_dof_position")
Expand All @@ -73,6 +66,14 @@ def compile_joint_velocity(self, name, write) -> None:
del name, write
self.dispatched.append("joint_velocity")

def compile_joint_quaternion(self, name, write) -> None:
del name, write
self.dispatched.append("joint_quaternion")

def compile_joint_angular_velocity(self, name, write) -> None:
del name, write
self.dispatched.append("joint_angular_velocity")

def compile_ctrl_targets(self, name, write) -> None:
del name, write
self.dispatched.append("ctrl")
Expand All @@ -93,9 +94,13 @@ def compile_body_angular_velocity(self, name, write) -> None:
del name, write
self.dispatched.append("body_angular_velocity")

def compile_mocap_pose(self, name, write) -> None:
def compile_kinematic_body_position(self, name, write) -> None:
del name, write
self.dispatched.append("mocap_position")

def compile_kinematic_body_rotation(self, name, write) -> None:
del name, write
self.dispatched.append("mocap")
self.dispatched.append("mocap_rotation")

def compile_actuator_kp(self, name, write) -> None:
del name, write
Expand All @@ -121,37 +126,39 @@ def compile_geom_friction(self, name, write) -> None:
def test_sim_write_compiler_dispatches_each_write_to_its_typed_compiler() -> None:
compiler = _DispatchCompiler()

DofPositionWrite().compile_with(compiler, "write")
DofVelocityWrite().compile_with(compiler, "write")
BodyJointPositionWrite("body").compile_with(compiler, "write")
BodyJointVelocityWrite("body").compile_with(compiler, "write")
JointPositionWrite(("joint",)).compile_with(compiler, "write")
JointVelocityWrite(("joint",)).compile_with(compiler, "write")
JointQuaternionWrite(("joint",)).compile_with(compiler, "write")
JointAngularVelocityWrite(("joint",)).compile_with(compiler, "write")
CtrlTargetsWrite().compile_with(compiler, "write")
BodyPositionWrite(("body",)).compile_with(compiler, "write")
BodyRotationWrite(("body",)).compile_with(compiler, "write")
BodyLinearVelocityWrite(("body",)).compile_with(compiler, "write")
BodyAngularVelocityWrite(("body",)).compile_with(compiler, "write")
MocapPoseWrite(("body",)).compile_with(compiler, "write")
KinematicBodyPositionWrite(("body",)).compile_with(compiler, "write")
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")
GeomFrictionWrite(("geom",)).compile_with(compiler, "write")

assert compiler.dispatched == [
"dof_position",
"dof_velocity",
"body_dof_position",
"body_dof_velocity",
"joint_position",
"joint_velocity",
"joint_quaternion",
"joint_angular_velocity",
"ctrl",
"body_position",
"body_rotation",
"body_linear_velocity",
"body_angular_velocity",
"mocap",
"mocap_position",
"mocap_rotation",
"kp",
"damping",
"mass",
Expand All @@ -164,10 +171,10 @@ def test_compile_dispatches_every_write_in_order_and_builds_one_program() -> Non
compiler = _DispatchCompiler()

program = compiler.compile(
{"a": DofPositionWrite(), "b": CtrlTargetsWrite(), "c": MocapPoseWrite(("body",))},
{"a": JointQuaternionWrite(("joint",)), "b": CtrlTargetsWrite(), "c": KinematicBodyRotationWrite(("body",))},
reset=True,
forward_kinematics=False,
)

assert isinstance(program, WriteProgram)
assert compiler.dispatched == ["dof_position", "ctrl", "mocap"]
assert compiler.dispatched == ["joint_quaternion", "ctrl", "mocap_rotation"]
2 changes: 1 addition & 1 deletion motrix_env_motrixsim/pyproject.toml
Original file line number Diff line number Diff line change
Expand Up @@ -12,7 +12,7 @@ readme = "README.md"
license = "Apache-2.0"
dependencies = [
"motrix-env-core",
"motrixsim==0.10.1.dev123478",
"motrixsim==0.10.1",
"numpy>=1.26",
]

Expand Down
2 changes: 1 addition & 1 deletion motrix_env_motrixsim/src/motrix_env_motrixsim/runtime.py
Original file line number Diff line number Diff line change
Expand Up @@ -314,7 +314,7 @@ def __init__(self, scene: SceneCfg, sim: SimCfg, num_envs: int) -> None:
self._data: mtx.SceneData = mtx.SceneData(self._model, batch=[num_envs])
self._num_envs = num_envs
self._model_compiler = MotrixSimModelCompiler(self._model)
self._write_compiler = MotrixSimWriteCompiler(self._model, self._data, self._masked_rows)
self._write_compiler = MotrixSimWriteCompiler(self._model, self._data)

@property
def model_compiler(self) -> SimModelCompiler:
Expand Down
Loading
Loading