feat(training): release V0.8 自调参 Agent
This commit is contained in:
@@ -106,6 +106,7 @@ def unitree_go2_rough_env_cfg(
|
||||
cfg.rewards["body_ang_vel"].params["asset_cfg"].body_names = ("base_link",)
|
||||
cfg.rewards["foot_clearance"].params["asset_cfg"].site_names = site_names
|
||||
cfg.rewards["foot_slip"].params["asset_cfg"].site_names = site_names
|
||||
cfg.metrics["slip_velocity"].params["asset_cfg"].site_names = site_names
|
||||
|
||||
cfg.terminations["illegal_contact"] = TerminationTermCfg(
|
||||
func=mdp.illegal_contact,
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
from mjlab.envs.mdp import * # noqa: F401, F403
|
||||
|
||||
from .curriculums import * # noqa: F403
|
||||
from .metrics import * # noqa: F403
|
||||
from .observations import * # noqa: F403
|
||||
from .rewards import * # noqa: F403
|
||||
from .terminations import * # noqa: F403
|
||||
|
||||
@@ -0,0 +1,88 @@
|
||||
"""Reward-weight-independent quality metrics for velocity tasks."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
import torch
|
||||
|
||||
from mjlab.entity import Entity
|
||||
from mjlab.managers.scene_entity_config import SceneEntityCfg
|
||||
from mjlab.sensor import ContactSensor
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from mjlab.envs import ManagerBasedRlEnv
|
||||
|
||||
_DEFAULT_ASSET_CFG = SceneEntityCfg("robot")
|
||||
|
||||
|
||||
def linear_velocity_rmse(
|
||||
env: ManagerBasedRlEnv,
|
||||
command_name: str,
|
||||
asset_cfg: SceneEntityCfg = _DEFAULT_ASSET_CFG,
|
||||
) -> torch.Tensor:
|
||||
"""Per-step commanded-vs-actual base linear velocity RMSE in body frame."""
|
||||
asset: Entity = env.scene[asset_cfg.name]
|
||||
command = env.command_manager.get_command(command_name)
|
||||
assert command is not None
|
||||
actual = asset.data.root_link_lin_vel_b
|
||||
error = torch.cat((command[:, :2] - actual[:, :2], -actual[:, 2:3]), dim=1)
|
||||
return torch.sqrt(torch.mean(torch.square(error), dim=1))
|
||||
|
||||
|
||||
def angular_velocity_rmse(
|
||||
env: ManagerBasedRlEnv,
|
||||
command_name: str,
|
||||
asset_cfg: SceneEntityCfg = _DEFAULT_ASSET_CFG,
|
||||
) -> torch.Tensor:
|
||||
"""Per-step commanded-vs-actual base angular velocity RMSE in body frame."""
|
||||
asset: Entity = env.scene[asset_cfg.name]
|
||||
command = env.command_manager.get_command(command_name)
|
||||
assert command is not None
|
||||
actual = asset.data.root_link_ang_vel_b
|
||||
desired = torch.zeros_like(actual)
|
||||
desired[:, 2] = command[:, 2]
|
||||
return torch.sqrt(torch.mean(torch.square(desired - actual), dim=1))
|
||||
|
||||
|
||||
def orientation_error(
|
||||
env: ManagerBasedRlEnv,
|
||||
asset_cfg: SceneEntityCfg = _DEFAULT_ASSET_CFG,
|
||||
) -> torch.Tensor:
|
||||
"""Magnitude of projected gravity in the base x/y plane; zero is upright."""
|
||||
asset: Entity = env.scene[asset_cfg.name]
|
||||
return torch.linalg.vector_norm(asset.data.projected_gravity_b[:, :2], dim=1)
|
||||
|
||||
|
||||
def fall_indicator(env: ManagerBasedRlEnv) -> torch.Tensor:
|
||||
"""One on non-timeout terminal steps, otherwise zero."""
|
||||
return env.termination_manager.terminated.float()
|
||||
|
||||
|
||||
def slip_velocity(
|
||||
env: ManagerBasedRlEnv,
|
||||
sensor_name: str,
|
||||
asset_cfg: SceneEntityCfg = _DEFAULT_ASSET_CFG,
|
||||
) -> torch.Tensor:
|
||||
"""Mean x/y velocity of feet currently touching the ground."""
|
||||
asset: Entity = env.scene[asset_cfg.name]
|
||||
sensor: ContactSensor = env.scene[sensor_name]
|
||||
assert sensor.data.found is not None
|
||||
contact = (sensor.data.found > 0).float()
|
||||
speed = torch.linalg.vector_norm(
|
||||
asset.data.site_lin_vel_w[:, asset_cfg.site_ids, :2], dim=-1
|
||||
)
|
||||
count = torch.clamp(torch.sum(contact, dim=1), min=1.0)
|
||||
return torch.sum(speed * contact, dim=1) / count
|
||||
|
||||
|
||||
def mechanical_power(
|
||||
env: ManagerBasedRlEnv,
|
||||
asset_cfg: SceneEntityCfg = _DEFAULT_ASSET_CFG,
|
||||
) -> torch.Tensor:
|
||||
"""Positive actuator mechanical power in watts (regeneration is ignored)."""
|
||||
asset: Entity = env.scene[asset_cfg.name]
|
||||
torque = asset.data.actuator_force[:, asset_cfg.actuator_ids]
|
||||
velocity = asset.data.joint_vel[:, asset_cfg.joint_ids]
|
||||
count = min(torque.shape[1], velocity.shape[1])
|
||||
return torch.sum(torch.clamp(torque[:, :count] * velocity[:, :count], min=0.0), dim=1)
|
||||
@@ -140,8 +140,27 @@ def make_velocity_env_cfg() -> ManagerBasedRlEnvCfg:
|
||||
##
|
||||
|
||||
metrics = {
|
||||
"mean_action_acc": MetricsTermCfg(
|
||||
func=mdp.mean_action_acc,
|
||||
"linear_velocity_rmse": MetricsTermCfg(
|
||||
func=mdp.linear_velocity_rmse,
|
||||
params={"command_name": "twist"},
|
||||
),
|
||||
"angular_velocity_rmse": MetricsTermCfg(
|
||||
func=mdp.angular_velocity_rmse,
|
||||
params={"command_name": "twist"},
|
||||
),
|
||||
"mean_action_acc": MetricsTermCfg(func=mdp.mean_action_acc),
|
||||
"orientation_error": MetricsTermCfg(func=mdp.orientation_error),
|
||||
"fall_indicator": MetricsTermCfg(func=mdp.fall_indicator),
|
||||
"slip_velocity": MetricsTermCfg(
|
||||
func=mdp.slip_velocity,
|
||||
params={
|
||||
"sensor_name": "feet_ground_contact",
|
||||
"asset_cfg": SceneEntityCfg("robot", site_names=()), # Set per-robot.
|
||||
},
|
||||
),
|
||||
"mechanical_power": MetricsTermCfg(
|
||||
func=mdp.mechanical_power,
|
||||
params={"asset_cfg": SceneEntityCfg("robot", joint_names=(".*",))},
|
||||
),
|
||||
}
|
||||
|
||||
@@ -298,6 +317,11 @@ def make_velocity_env_cfg() -> ManagerBasedRlEnvCfg:
|
||||
params={"sensor_name": "robot/root_angmom"},
|
||||
),
|
||||
"is_terminated": RewardTermCfg(func=mdp.is_terminated, weight=-200.0),
|
||||
"electrical_power": RewardTermCfg(
|
||||
func=mdp.electrical_power_cost,
|
||||
weight=0.0,
|
||||
params={"asset_cfg": SceneEntityCfg("robot", joint_names=(".*",))},
|
||||
),
|
||||
"joint_acc_l2": RewardTermCfg(func=mdp.joint_acc_l2, weight=-2.5e-7),
|
||||
"joint_pos_limits": RewardTermCfg(func=mdp.joint_pos_limits, weight=-10.0),
|
||||
"action_rate_l2": RewardTermCfg(func=mdp.action_rate_l2, weight=-0.05),
|
||||
|
||||
Reference in New Issue
Block a user