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

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
9 changes: 7 additions & 2 deletions configs/task/g1-walk-flat/motrix.fastsac.yaml
Original file line number Diff line number Diff line change
Expand Up @@ -14,6 +14,11 @@ checkpoint:
interval: 1000
algo:
agent:
alpha_init: 0.01
target_entropy_ratio: -0.1
# Holosoma-alpha parity (alpha_init 0.001, target_entropy_ratio 0.0).
# UTD 4 instead of holosoma's 8: the utd16 sweep measured no tracking
# gain from more updates, at half the learner cost.
alpha_init: 0.001
target_entropy_ratio: 0.0
num_updates: 4
trainer:
num_learning_iterations: 30000
46 changes: 45 additions & 1 deletion motrix_env_core/src/motrix_env_core/mdp/terminations.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,20 +3,37 @@

"""Reusable termination terms for manager-based environments."""

import math

import numpy as np

from motrix_env_core.config import configclass
from motrix_env_core.manager import ManagerContext, TerminationTerm, TerminationTermCfg
from motrix_env_core.numba.manager.context import BuildContext
from motrix_env_core.numba.manager.dispatch import dispatch
from motrix_env_core.sim import GeomPairCollidingQuery
from motrix_env_core.sim import BodyJointVelocityQuery, GeomPairCollidingQuery


@dispatch
def colliding_termination(ctx: ManagerContext, colliding: np.ndarray) -> bool:
return bool(colliding.any())


@dispatch
def bad_dof_velocity_termination(ctx: ManagerContext, dof_vel: np.ndarray, threshold: np.float32) -> bool:
error = 0.0
finite = True
for joint_id in range(dof_vel.shape[0]):
velocity = dof_vel[joint_id]
if math.isfinite(velocity):
error = max(error, abs(velocity))
else:
finite = False
error = math.inf
ctx.metrics["dof_vel_abs_max"][0] = error
return (not finite) or error > threshold


@configclass(kw_only=True)
class CollidingTerminationCfg(TerminationTermCfg):
"""Terminate when any of ``termination_geoms`` contacts ``ground_geom``.
Expand All @@ -35,7 +52,34 @@ def __call__(self, ctx: BuildContext) -> TerminationTerm:
return TerminationTerm(colliding_termination, query)


@configclass(kw_only=True)
class BadDofVelocityTerminationCfg(TerminationTermCfg):
"""Terminate when one body's joint speed magnitude exceeds ``threshold``.

Training hygiene for harsh-contact tasks: a physics near-blowup leaves the
lane with finite-but-absurd joint speeds; terminating resets the lane so
those states do not keep feeding the learner. The threshold must stay far
above any healthy gait speed.
"""

body: str = "robot"
threshold: float = 100.0
Comment thread
wlgys8 marked this conversation as resolved.

def __call__(self, ctx: BuildContext) -> TerminationTerm:
if self.threshold <= 0.0:
raise ValueError(f"BadDofVelocityTerminationCfg.threshold must be positive, got {self.threshold!r}")
body = ctx.model.bodies[self.body]
return TerminationTerm(
bad_dof_velocity_termination,
BodyJointVelocityQuery(body=body.base_link_name),
np.float32(self.threshold),
metric_names=("dof_vel_abs_max",),
)


__all__ = [
"BadDofVelocityTerminationCfg",
"CollidingTerminationCfg",
"bad_dof_velocity_termination",
"colliding_termination",
]
10 changes: 6 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 @@ -8,15 +8,16 @@
ActuatorKpQuery,
ActuatorSpec,
ActuatorType,
BodyCenterOfMassQuery,
BodyJointPositionLimitsQuery,
BodyMassQuery,
BodyMassesQuery,
BodyModel,
DofPositionLimitsQuery,
GeomFrictionQuery,
GeomSpec,
GeomSpecsQuery,
HeightFieldDataQuery,
LinkCenterOfMassQuery,
LinkMassQuery,
ModelQuery,
SimModel,
SimModelCompiler,
Expand Down Expand Up @@ -75,7 +76,7 @@
"BatchLinkNetContactForceQuery",
"BatchLinkPositionQuery",
"BatchLinkQuaternionQuery",
"BodyCenterOfMassQuery",
"LinkCenterOfMassQuery",
"BodyAngularVelocityWrite",
"BodyJointPositionWrite",
"BodyJointVelocityWrite",
Expand All @@ -84,7 +85,8 @@
"BodyRotationWrite",
"BodyJointPositionLimitsQuery",
"BodyLinkNetContactForceQuery",
"BodyMassQuery",
"LinkMassQuery",
"BodyMassesQuery",
"DofPositionLimitsQuery",
"DofPositionQuery",
"DofVelocityQuery",
Expand Down
49 changes: 43 additions & 6 deletions motrix_env_core/src/motrix_env_core/sim/model.py
Original file line number Diff line number Diff line change
Expand Up @@ -7,6 +7,11 @@
...) define what every backend must produce as ``env.model``; the query
classes declare environment-owned metadata lookups; the compiler base wires
the two together. Runtime behavior lives in ``sim.backend``.

Naming convention for query targets: ``Link*`` types address a single rigid
link; ``Body*`` types address one ``BodyCfg`` body tree by its root's name
and read the whole link subtree in tree order. How a backend maps these two
namespaces onto its own model representation is backend-private.
"""

from __future__ import annotations
Expand Down Expand Up @@ -215,23 +220,41 @@ def compile_with(self, compiler: SimModelCompiler, *, key: str) -> None:


@dataclass(frozen=True)
class BodyMassQuery(ModelQuery):
class LinkMassQuery(ModelQuery):
"""Scalar ``float`` nominal mass of one named link."""

name: str

def compile_with(self, compiler: SimModelCompiler, *, key: str) -> None:
compiler.compile_body_mass(key, self.name)
compiler.compile_link_mass(key, self.name)


@dataclass(frozen=True)
class BodyMassesQuery(ModelQuery):
"""``(L,)`` float32 nominal masses of one body tree's links.

``body`` names the body tree root (a ``BodyCfg`` base link works, since the
root link and the tree share one name); this is a different namespace from
:class:`LinkMassQuery`, whose ``name`` addresses a single link. ``links``
selects and orders the result: ``None`` returns every link of the tree in
tree order, otherwise exactly the named links in the given order.
"""

body: str
links: tuple[str, ...] | None = None

def compile_with(self, compiler: SimModelCompiler, *, key: str) -> None:
compiler.compile_body_masses(key, self.body, self.links)


@dataclass(frozen=True)
class BodyCenterOfMassQuery(ModelQuery):
class LinkCenterOfMassQuery(ModelQuery):
"""``(3,)`` float32 nominal center-of-mass offset of one named link."""

name: str

def compile_with(self, compiler: SimModelCompiler, *, key: str) -> None:
compiler.compile_body_center_of_mass(key, self.name)
compiler.compile_link_center_of_mass(key, self.name)


@dataclass(frozen=True)
Expand Down Expand Up @@ -334,7 +357,7 @@ def compile_actuator_kd(self, key: str, actuator_names: tuple[str, ...] | None)
"""

@abc.abstractmethod
def compile_body_mass(self, key: str, body: str) -> None:
def compile_link_mass(self, key: str, body: str) -> None:
"""Compile the nominal mass of one body link.

Args:
Expand All @@ -343,7 +366,21 @@ def compile_body_mass(self, key: str, body: str) -> None:
"""

@abc.abstractmethod
def compile_body_center_of_mass(self, key: str, body: str) -> None:
def compile_body_masses(self, key: str, body: str, links: tuple[str, ...] | None) -> None:
"""Compile the nominal masses of one body tree's links.

Unlike :meth:`compile_link_mass` (a single link name), ``body`` names
the body tree root. ``links=None`` covers the whole link subtree in
tree order; an explicit tuple selects and orders exactly those links.

Args:
key: Logical key under which the result is stored.
body: Name of the body tree whose link masses are read.
links: Link-name subset to read, or ``None`` for the full tree.
"""

@abc.abstractmethod
def compile_link_center_of_mass(self, key: str, body: str) -> None:
"""Compile the nominal center-of-mass offset of one body link.

Args:
Expand Down
16 changes: 10 additions & 6 deletions motrix_env_core/src/motrix_env_core/sim/write.py
Original file line number Diff line number Diff line change
Expand Up @@ -12,6 +12,10 @@
Reset-before-write and forward-kinematics behavior are fixed by
:meth:`SimWriteCompiler.compile`; execution only selects rows. Target names
are validated at compile time and fail loudly.

Naming convention for write targets: ``Link*`` ops address single links;
``Body*`` ops address one ``BodyCfg`` body tree by its root's name and
cover the whole link subtree in tree order.
"""

from __future__ import annotations
Expand Down Expand Up @@ -190,23 +194,23 @@ def compile_with(self, compiler: SimWriteCompiler, name: str) -> None:


@dataclass(frozen=True)
class BodyMassWrite(SimWrite):
class LinkMassWrite(SimWrite):
"""Link mass overrides in declared order: ``(N, L)``."""

links: tuple[str, ...]

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


@dataclass(frozen=True)
class BodyComWrite(SimWrite):
class LinkComWrite(SimWrite):
"""Link center-of-mass overrides in declared order: ``(N, L, 3)``."""

links: tuple[str, ...]

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


@dataclass(frozen=True)
Expand Down Expand Up @@ -317,11 +321,11 @@ def compile_actuator_damping(self, name: str, write: ActuatorDampingWrite) -> No
"""Record actuator damping overrides."""

@abc.abstractmethod
def compile_body_mass(self, name: str, write: BodyMassWrite) -> None:
def compile_link_mass(self, name: str, write: LinkMassWrite) -> None:
"""Record body mass overrides."""

@abc.abstractmethod
def compile_body_com(self, name: str, write: BodyComWrite) -> None:
def compile_link_com(self, name: str, write: LinkComWrite) -> None:
"""Record body center-of-mass overrides."""

@abc.abstractmethod
Expand Down
16 changes: 10 additions & 6 deletions motrix_env_core/tests/test_model_query_dispatch.py
Original file line number Diff line number Diff line change
Expand Up @@ -7,13 +7,13 @@
from motrix_env_core.sim import (
ActuatorKdQuery,
ActuatorKpQuery,
BodyCenterOfMassQuery,
BodyJointPositionLimitsQuery,
BodyMassQuery,
DofPositionLimitsQuery,
GeomFrictionQuery,
GeomSpecsQuery,
HeightFieldDataQuery,
LinkCenterOfMassQuery,
LinkMassQuery,
SimModelCompiler,
)
from motrix_env_core.sim.model import SimModel
Expand Down Expand Up @@ -50,14 +50,18 @@ def compile_actuator_kd(self, key, actuator_names) -> None:
del actuator_names
self.dispatched[key] = "kd"

def compile_body_mass(self, key, body) -> None:
def compile_link_mass(self, key, body) -> None:
del body
self.dispatched[key] = "mass"

def compile_body_center_of_mass(self, key, body) -> None:
def compile_link_center_of_mass(self, key, body) -> None:
del body
self.dispatched[key] = "com"

def compile_body_masses(self, key, body, links) -> None:
del body, links
self.dispatched[key] = "masses"

def compile_geom_friction(self, key, geom) -> None:
del geom
self.dispatched[key] = "friction"
Expand All @@ -75,8 +79,8 @@ def test_model_queries_dispatch_to_typed_compiler_methods() -> None:
"dof_limits": DofPositionLimitsQuery(),
"kp": ActuatorKpQuery(names=("first", "second")),
"kd": ActuatorKdQuery(names=None),
"mass": BodyMassQuery(name="body"),
"com": BodyCenterOfMassQuery(name="body"),
"mass": LinkMassQuery(name="body"),
"com": LinkCenterOfMassQuery(name="body"),
"friction": GeomFrictionQuery(name="geom"),
"heightfield": HeightFieldDataQuery(geom="floor"),
}
Expand Down
12 changes: 6 additions & 6 deletions motrix_env_core/tests/test_sim_write_dispatch.py
Original file line number Diff line number Diff line change
Expand Up @@ -9,11 +9,9 @@
ActuatorDampingWrite,
ActuatorKpWrite,
BodyAngularVelocityWrite,
BodyComWrite,
BodyJointPositionWrite,
BodyJointVelocityWrite,
BodyLinearVelocityWrite,
BodyMassWrite,
BodyPositionWrite,
BodyRotationWrite,
CtrlTargetsWrite,
Expand All @@ -24,6 +22,8 @@
JointVelocityWrite,
KinematicBodyPositionWrite,
KinematicBodyRotationWrite,
LinkComWrite,
LinkMassWrite,
SimWriteCompiler,
WriteProgram,
)
Expand Down Expand Up @@ -110,11 +110,11 @@ def compile_actuator_damping(self, name, write) -> None:
del name, write
self.dispatched.append("damping")

def compile_body_mass(self, name, write) -> None:
def compile_link_mass(self, name, write) -> None:
del name, write
self.dispatched.append("mass")

def compile_body_com(self, name, write) -> None:
def compile_link_com(self, name, write) -> None:
del name, write
self.dispatched.append("com")

Expand All @@ -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(("body",)).compile_with(compiler, "write")
BodyComWrite(("body",)).compile_with(compiler, "write")
LinkMassWrite(("body",)).compile_with(compiler, "write")
LinkComWrite(("body",)).compile_with(compiler, "write")
GeomFrictionWrite(("geom",)).compile_with(compiler, "write")

assert compiler.dispatched == [
Expand Down
21 changes: 19 additions & 2 deletions motrix_env_motrixsim/src/motrix_env_motrixsim/runtime.py
Original file line number Diff line number Diff line change
Expand Up @@ -97,10 +97,27 @@ def compile_actuator_kp(self, key: str, actuator_names: tuple[str, ...] | None)
def compile_actuator_kd(self, key: str, actuator_names: tuple[str, ...] | None) -> None:
self._others[key] = _nominal_actuator_kd(self._model, actuator_names)

def compile_body_mass(self, key: str, body: str) -> None:
def compile_link_mass(self, key: str, body: str) -> None:
self._others[key] = float(_named_link(self._model, body).mass)

def compile_body_center_of_mass(self, key: str, body: str) -> None:
def compile_body_masses(self, key: str, body: str, links: tuple[str, ...] | None) -> None:
# ``body`` names the body tree root (unlike compile_link_mass, which
# addresses a single link); ``links=None`` covers the whole subtree in
# tree order, otherwise exactly the named links in the given order.
scene_body = self._model.get_body(body)
if scene_body is None:
raise KeyError(f"Unknown body {body!r}.")
body_links = {link.name: link for link in scene_body.links}
if links is None:
masses = [link.mass for link in scene_body.links]
else:
try:
masses = [body_links[name].mass for name in links]
except KeyError as exc:
raise KeyError(f"Link {exc.args[0]!r} is not part of body {body!r}.") from None
self._others[key] = np.asarray(masses, dtype=np.float32)

def compile_link_center_of_mass(self, key: str, body: str) -> None:
self._others[key] = np.asarray(_named_link(self._model, body).center_of_mass, dtype=np.float32).reshape(3)

def compile_geom_friction(self, key: str, geom: str) -> None:
Expand Down
Loading
Loading