This commit is contained in:
wty-yy
2025-12-06 17:11:40 +08:00
parent a05d047b39
commit 9d509829a5
19 changed files with 364 additions and 60 deletions

View File

@@ -1,6 +1,7 @@
from robogauge.utils.logger import logger
from robogauge.tasks.robots import RobotConfig
from robogauge.tasks.simulator.sim_data import SimData
from robogauge.tasks.gauge.goals.base_goal import GoalData
class BaseMetric:
""" Base class for all metric functions. """
@@ -9,7 +10,7 @@ class BaseMetric:
def __init__(self, robot_cfg: RobotConfig, **kwargs):
self.robot_cfg = robot_cfg
def __call__(self, sim_data: SimData) -> float:
def __call__(self, sim_data: SimData, goal_data: GoalData) -> float:
value = 0.0
logger.log(value, self.name, step=sim_data.n_step)
return value

View File

@@ -1,8 +1,7 @@
import numpy as np
from robogauge.tasks.robots import RobotConfig
from robogauge.tasks.simulator.sim_data import SimData
from robogauge.tasks.gauge.metrics.base_metric import BaseMetric
from robogauge.tasks.gauge.metrics.base_metric import BaseMetric, GoalData, SimData
from robogauge.utils.logger import logger
@@ -21,7 +20,7 @@ class DofLimitsMetric(BaseMetric):
self.soft_dof_limit_ratio = soft_dof_limit_ratio
self.calc_dof_names = dof_names
def __call__(self, sim_data: SimData) -> float:
def __call__(self, sim_data: SimData, goal_data: GoalData) -> float:
values = []
for i in range(len(sim_data.proprio.joint.limits)):
lower_limit = sim_data.proprio.joint.limits[i, 0]

View File

@@ -1,11 +1,10 @@
import numpy as np
from robogauge.tasks.robots import RobotConfig
from robogauge.tasks.simulator.sim_data import SimData
from robogauge.tasks.gauge.metrics.base_metric import BaseMetric
from robogauge.tasks.gauge.goals.base_goal import GoalData
from robogauge.tasks.gauge.metrics.base_metric import BaseMetric, SimData, GoalData
from robogauge.utils.logger import logger
from robogauge.utils.helpers import class_to_dict
class LinVelErrMetric(BaseMetric):
@@ -15,13 +14,11 @@ class LinVelErrMetric(BaseMetric):
def __init__(self, robot_cfg: RobotConfig, **kwargs):
super().__init__(robot_cfg)
max_ranges = []
cfg_commands = class_to_dict(robot_cfg.commands)
for name in ['lin_vel_x', 'lin_vel_y', 'lin_vel_z']:
max_ranges.append(
max(
abs(getattr(robot_cfg.commands, name)[0]),
abs(getattr(robot_cfg.commands, name)[1])
)
)
cmds = cfg_commands.get(name)
if cmds is not None:
max_ranges.append(max(abs(cmds[0]), abs(cmds[1])))
self.norm_vel = np.linalg.norm(max_ranges)
def __call__(self, sim_data: SimData, goal_data: GoalData) -> float:
@@ -41,16 +38,11 @@ class AngVelErrMetric(BaseMetric):
def __init__(self, robot_cfg: RobotConfig, **kwargs):
super().__init__(robot_cfg)
max_ranges = []
cfg_commands = class_to_dict(robot_cfg.commands)
for name in ['ang_vel_roll', 'ang_vel_pitch', 'ang_vel_yaw']:
cmd_range = getattr(robot_cfg.commands, name)
if cmd_range is None:
cmd_range = [0, 0]
max_ranges.append(
max(
abs(cmd_range[0]),
abs(cmd_range[1])
)
)
cmds = cfg_commands.get(name)
if cmds is not None:
max_ranges.append(max(abs(cmds[0]), abs(cmds[1])))
self.norm_vel = np.linalg.norm(max_ranges)
def __call__(self, sim_data: SimData, goal_data: GoalData) -> float:

View File

@@ -1,6 +1,5 @@
from robogauge.tasks.robots import RobotConfig
from robogauge.tasks.simulator.sim_data import SimData
from robogauge.tasks.gauge.metrics.base_metric import BaseMetric
from robogauge.tasks.gauge.metrics.base_metric import BaseMetric, SimData, GoalData
from robogauge.utils.logger import logger
@@ -18,7 +17,7 @@ class VisualizationMetric(BaseMetric):
self.dof_force = dof_force
self.dof_pos = dof_pos
def __call__(self, sim_data: SimData) -> float:
def __call__(self, sim_data: SimData, goal_data: GoalData) -> float:
for i in range(len(sim_data.proprio.joint.force)):
name = sim_data.proprio.joint.names[i]
if self.dof_force: