diff --git a/CMD.md b/CMD.md index 2ad0487..e52fa8e 100644 --- a/CMD.md +++ b/CMD.md @@ -23,11 +23,18 @@ python robogauge/scripts/run.py \ --headless ``` -# Radar Plot +# Radar/Bar Plot 将Multi Run结果绘制在雷达图中 ```bash -python robogauge/utils/radar_plot.py \ +# 可选--out保存到图片 +python robogauge/utils/visualize_results.py \ aggregated_results_1.yaml \ aggregated_results_2.yaml \ - --out logs/go2_flat_vs_moe_flat.png + --out logs/1.jpg + +# 绘制下全部 *.yaml, 可选range控制radar, bar y轴显示范围 +python robogauge/utils/visualize_results.py \ + /home/xfy/Coding/robot_gauge/mytest/results \ + --range 0.35 1.0 \ + --out logs/1.jpg ``` \ No newline at end of file diff --git a/UPDATE.md b/UPDATE.md index 81c733d..edc902d 100644 --- a/UPDATE.md +++ b/UPDATE.md @@ -1,8 +1,9 @@ # UPDATE -TODO: 在模型崩溃时也记录下最后的gauge信息 +TODO: 即使模型崩溃也要继续测完后续的goals ## 20251221 ### v0.1.12 1. 在模型崩溃时也记录下最后的gauge信息, 修改single/multi pipeline逻辑 +2. 添加slope全等级地形, 关系式: 斜率$0.1+0.47d$,角度范围$5.7\sim29.7^\circ$, 发现29.7度在2.4摩擦系数下可以稳定通过 ## 20251220 ### v0.1.11 1. 加入`os.environ["OMP_NUM_THREADS"] = "2"; os.environ["MKL_NUM_THREADS"] = "2"`避免并行时cpu线程爆炸, `--multi`模式能稳定提高速度了 diff --git a/resources/models/go2/go2_cts_83501.pt b/resources/models/go2/go2_cts_83501.pt deleted file mode 100644 index cc975c3..0000000 Binary files a/resources/models/go2/go2_cts_83501.pt and /dev/null differ diff --git a/resources/models/go2/go2_cts_cmd-1,1_38k.pt b/resources/models/go2/go2_cts_cmd-1,1_38k.pt deleted file mode 100644 index e14df59..0000000 Binary files a/resources/models/go2/go2_cts_cmd-1,1_38k.pt and /dev/null differ diff --git a/resources/models/go2/go2_cts_max2_100k.pt b/resources/models/go2/go2_cts_max2_100k.pt new file mode 100644 index 0000000..05ec4a1 Binary files /dev/null and b/resources/models/go2/go2_cts_max2_100k.pt differ diff --git a/resources/models/go2/go2_moe_cts_8experts_2cmd_100k.pt b/resources/models/go2/go2_moe_cts_8experts_2cmd_100k.pt new file mode 100644 index 0000000..5762923 Binary files /dev/null and b/resources/models/go2/go2_moe_cts_8experts_2cmd_100k.pt differ diff --git a/resources/terrains/slope/slope_1.xml b/resources/terrains/slope/slope_1.xml new file mode 100644 index 0000000..f8c90ac --- /dev/null +++ b/resources/terrains/slope/slope_1.xml @@ -0,0 +1,19 @@ + + + + + + + + + + + + + + + + + + + diff --git a/resources/terrains/slope/slope_10.xml b/resources/terrains/slope/slope_10.xml new file mode 100644 index 0000000..b61559d --- /dev/null +++ b/resources/terrains/slope/slope_10.xml @@ -0,0 +1,19 @@ + + + + + + + + + + + + + + + + + + + diff --git a/resources/terrains/slope/slope_2.xml b/resources/terrains/slope/slope_2.xml new file mode 100644 index 0000000..8e8b7a0 --- /dev/null +++ b/resources/terrains/slope/slope_2.xml @@ -0,0 +1,19 @@ + + + + + + + + + + + + + + + + + + + diff --git a/resources/terrains/slope/slope_3.xml b/resources/terrains/slope/slope_3.xml new file mode 100644 index 0000000..eae16cd --- /dev/null +++ b/resources/terrains/slope/slope_3.xml @@ -0,0 +1,19 @@ + + + + + + + + + + + + + + + + + + + diff --git a/resources/terrains/slope/slope_4.xml b/resources/terrains/slope/slope_4.xml new file mode 100644 index 0000000..ad3fd52 --- /dev/null +++ b/resources/terrains/slope/slope_4.xml @@ -0,0 +1,19 @@ + + + + + + + + + + + + + + + + + + + diff --git a/resources/terrains/slope/slope_5.xml b/resources/terrains/slope/slope_5.xml new file mode 100644 index 0000000..180f4cf --- /dev/null +++ b/resources/terrains/slope/slope_5.xml @@ -0,0 +1,19 @@ + + + + + + + + + + + + + + + + + + + diff --git a/resources/terrains/slope/slope_6.xml b/resources/terrains/slope/slope_6.xml new file mode 100644 index 0000000..d53ac48 --- /dev/null +++ b/resources/terrains/slope/slope_6.xml @@ -0,0 +1,19 @@ + + + + + + + + + + + + + + + + + + + diff --git a/resources/terrains/slope/slope_7.xml b/resources/terrains/slope/slope_7.xml new file mode 100644 index 0000000..f3bdf46 --- /dev/null +++ b/resources/terrains/slope/slope_7.xml @@ -0,0 +1,19 @@ + + + + + + + + + + + + + + + + + + + diff --git a/resources/terrains/slope/slope_8.xml b/resources/terrains/slope/slope_8.xml new file mode 100644 index 0000000..c9d9170 --- /dev/null +++ b/resources/terrains/slope/slope_8.xml @@ -0,0 +1,19 @@ + + + + + + + + + + + + + + + + + + + diff --git a/resources/terrains/slope/slope_9.xml b/resources/terrains/slope/slope_9.xml new file mode 100644 index 0000000..e23df0f --- /dev/null +++ b/resources/terrains/slope/slope_9.xml @@ -0,0 +1,19 @@ + + + + + + + + + + + + + + + + + + + diff --git a/resources/terrains/slope_5.xml b/resources/terrains/slope_5.xml new file mode 100644 index 0000000..51fc78b --- /dev/null +++ b/resources/terrains/slope_5.xml @@ -0,0 +1,23 @@ + + + + + + + + + + + + + + + + + + + + + + + \ No newline at end of file diff --git a/resources/terrains/slope_5_old.xml b/resources/terrains/slope_5_old.xml new file mode 100644 index 0000000..d89c807 --- /dev/null +++ b/resources/terrains/slope_5_old.xml @@ -0,0 +1,23 @@ + + + + + + + + + + + + + + + + + + + + + + + \ No newline at end of file diff --git a/robogauge/tasks/__init__.py b/robogauge/tasks/__init__.py index 1cf91b2..2074fa3 100644 --- a/robogauge/tasks/__init__.py +++ b/robogauge/tasks/__init__.py @@ -4,8 +4,10 @@ from robogauge.tasks.robots import RobotConfig, Go2Config, Go2MoEConfig from robogauge.tasks.pipeline import BasePipeline from robogauge.tasks.gauge import BaseGaugeConfig -from robogauge.tasks.custom.go2_flat_task import Go2FlatGaugeConfig, Go2FlatConfig, Go2MoEFlatConfig, Go2MoEFlatMujocoConfig +from robogauge.tasks.custom.go2 import * task_register.register('base', BasePipeline, MujocoConfig, BaseGaugeConfig, RobotConfig) -task_register.register('go2_flat', BasePipeline, MujocoConfig, Go2FlatGaugeConfig, Go2FlatConfig) -task_register.register('go2_moe_flat', BasePipeline, Go2MoEFlatMujocoConfig, Go2FlatGaugeConfig, Go2MoEFlatConfig) +task_register.register('go2_flat', BasePipeline, MujocoConfig, Go2FlatGaugeConfig, Go2Config) +task_register.register('go2_moe_flat', BasePipeline, MujocoConfig, Go2FlatGaugeConfig, Go2MoEConfig) +task_register.register('go2_slope', BasePipeline, Go2SlopeMujocoConfig, Go2SlopeGaugeConfig, Go2Config) +task_register.register('go2_moe_slope', BasePipeline, Go2SlopeMujocoConfig, Go2SlopeGaugeConfig, Go2MoEConfig) \ No newline at end of file diff --git a/robogauge/tasks/custom/go2/__init__.py b/robogauge/tasks/custom/go2/__init__.py new file mode 100644 index 0000000..3d67a3f --- /dev/null +++ b/robogauge/tasks/custom/go2/__init__.py @@ -0,0 +1,10 @@ +# -*- coding: utf-8 -*- +''' +@File : __init__.py +@Time : 2025/12/21 17:10:59 +@Author : wty-yy +@Version : 1.0 +@Blog : https://wty-yy.github.io/ +''' +from .go2_flat_task import Go2FlatGaugeConfig +from .go2_slope_task import Go2SlopeGaugeConfig, Go2SlopeMujocoConfig diff --git a/robogauge/tasks/custom/go2_flat_task.py b/robogauge/tasks/custom/go2/go2_flat_task.py similarity index 66% rename from robogauge/tasks/custom/go2_flat_task.py rename to robogauge/tasks/custom/go2/go2_flat_task.py index af53a8e..c07aaa1 100644 --- a/robogauge/tasks/custom/go2_flat_task.py +++ b/robogauge/tasks/custom/go2/go2_flat_task.py @@ -37,24 +37,3 @@ class Go2FlatGaugeConfig(FlatGaugeConfig): ang_vel_yaw = 1.5 # +/- rad/s max_cmd_duration = 10.0 # [s] maximum duration to reach the target position reach_threshold = 0.1 # [m] distance threshold to consider the target reached - -class Go2FlatConfig(Go2Config): - class commands(Go2Config.commands): - lin_vel_x = [-1.8, 1.8] # min max [m/s] - lin_vel_y = [-1.8, 1.8] # min max [m/s] - ang_vel_yaw = [-2.0, 2.0] # min max [rad/s] - -class Go2MoEFlatConfig(Go2MoEConfig): - class commands(Go2Config.commands): - lin_vel_x = [-1.8, 1.8] # min max [m/s] - lin_vel_y = [-1.8, 1.8] # min max [m/s] - ang_vel_yaw = [-2.0, 2.0] # min max [rad/s] - - class control(Go2Config.control): - model_path = "{ROBOGAUGE_ROOT_DIR}/resources/models/go2/go2_moe_cts_124k.pt" - # model_path = "/home/xfy/Coding/kaiwu2025/rob_finals/sim2real/models/v6-2_106503/kaiwu_script_v6-2_106503.pt" - -class Go2MoEFlatMujocoConfig(MujocoConfig): - class domain_rand(MujocoConfig.domain_rand): - base_mass = 0.0 - friction = 1.0 diff --git a/robogauge/tasks/custom/go2/go2_slope_task.py b/robogauge/tasks/custom/go2/go2_slope_task.py new file mode 100644 index 0000000..8dff561 --- /dev/null +++ b/robogauge/tasks/custom/go2/go2_slope_task.py @@ -0,0 +1,24 @@ +# -*- coding: utf-8 -*- +''' +@File : go2_slope_task.py +@Time : 2025/12/21 17:06:27 +@Author : wty-yy +@Version : 1.0 +@Blog : https://wty-yy.github.io/ +@Desc : Go2 Slope Task Configuration +''' +from robogauge.tasks.robots import Go2Config, Go2MoEConfig +from robogauge.tasks.gauge import SlopeGaugeConfig +from robogauge.tasks.simulator.mujoco_config import MujocoConfig + +class Go2SlopeGaugeConfig(SlopeGaugeConfig): + class metrics(SlopeGaugeConfig.metrics): + class dof_limits(SlopeGaugeConfig.metrics.dof_limits): + enabled = True + soft_dof_limit_ratio = 0.7 + dof_names = ['hip', 'thigh'] # List of DOF names to monitor, None for all + +class Go2SlopeMujocoConfig(MujocoConfig): + class domain_rand(MujocoConfig.domain_rand): + action_delay = True + friction = 2.4 diff --git a/robogauge/tasks/gauge/__init__.py b/robogauge/tasks/gauge/__init__.py index c40b81b..1ce2ec5 100644 --- a/robogauge/tasks/gauge/__init__.py +++ b/robogauge/tasks/gauge/__init__.py @@ -1,3 +1,4 @@ from .base_gauge import BaseGauge from .base_gauge_config import BaseGaugeConfig from .gauge_configs.flat_gauge_config import FlatGaugeConfig +from .gauge_configs.slope_gauge_config import SlopeGaugeConfig diff --git a/robogauge/tasks/gauge/base_gauge_config.py b/robogauge/tasks/gauge/base_gauge_config.py index e3c39fa..1ccffd1 100644 --- a/robogauge/tasks/gauge/base_gauge_config.py +++ b/robogauge/tasks/gauge/base_gauge_config.py @@ -20,18 +20,19 @@ class BaseGaugeConfig(Config): class goals: class max_velocity: # goal with maximum velocity - enabled = True + enabled = False cmd_duration = 5.0 # [s] duration for each velocity command class diagonal_velocity: # goal with diagonal velocity changes - enabled = True + enabled = False cmd_duration = 6.0 # [s] duration for a pair of diagonal velocity commands class target_pos_velocity: # goal to reach a target position by velocity command - enabled = True + enabled = False target_pos = [5, 0, 0] # x y z [m], target position in the environment, used for target position goal lin_vel_x = 1.0 # +/- m/s - ang_vel_yaw = 1.0 # +/- rad/s + lin_vel_y = 1.0 # +/- m/s + ang_vel_yaw = 1.5 # +/- rad/s max_cmd_duration = 10.0 # [s] maximum duration to reach the target position reach_threshold = 0.1 diff --git a/robogauge/tasks/gauge/gauge_configs/flat_gauge_config.py b/robogauge/tasks/gauge/gauge_configs/flat_gauge_config.py index f1aa193..e74f918 100644 --- a/robogauge/tasks/gauge/gauge_configs/flat_gauge_config.py +++ b/robogauge/tasks/gauge/gauge_configs/flat_gauge_config.py @@ -16,40 +16,3 @@ class FlatGaugeConfig(BaseGaugeConfig): terrain_name = "flat_0" # {type}_{level} terrain_xml = '{ROBOGAUGE_ROOT_DIR}/resources/terrains/flat.xml' terrain_spawn_pos = [0, 0, 0] # x y z [m], robot freejoint spawn position on the terrain - - class goals(BaseGaugeConfig.goals): - class max_velocity: # goal with maximum velocity - enabled = True - move_duration = 5.0 # [s] duration for each velocity command - end_stance = True # whether to end with zero velocity command - standce_duration = 2.0 # [s] duration for the ending stance command - - class diagonal_velocity: # goal with diagonal velocity changes - enabled = True - cmd_duration = 6.0 # [s] duration for a pair of diagonal velocity commands - - class target_pos_velocity: # goal to reach a target position by velocity command, config target at assets.target_pos - enabled = True - target_pos = [5, 0, 0] # x y z [m], target position in the environment, used for target position goal - lin_vel_x = 1.0 # +/- m/s - ang_vel_yaw = 1.0 # +/- rad/s - max_cmd_duration = 10.0 # [s] maximum duration to reach the target position - reach_threshold = 0.1 - - class metrics(BaseGaugeConfig.metrics): - metric_dt = 0.1 # [s] frequency to compute metrics - class dof_limits: - enabled = True - soft_dof_limit_ratio = 0.9 - dof_names = None # List of DOF names to monitor, None for all - - class visualization: - enabled = True - dof_torque = True - dof_pos = True - - class lin_vel_err: - enabled = True - - class ang_vel_err: - enabled = True diff --git a/robogauge/tasks/gauge/gauge_configs/slope_gauge_config.py b/robogauge/tasks/gauge/gauge_configs/slope_gauge_config.py new file mode 100644 index 0000000..3a59d22 --- /dev/null +++ b/robogauge/tasks/gauge/gauge_configs/slope_gauge_config.py @@ -0,0 +1,28 @@ +# -*- coding: utf-8 -*- +''' +@File : slope_gauge_config.py +@Time : 2025/12/21 17:02:23 +@Author : wty-yy +@Version : 1.0 +@Blog : https://wty-yy.github.io/ +@Desc : Slope Gauge Configuration +''' +from robogauge.tasks.gauge.base_gauge_config import BaseGaugeConfig + +class SlopeGaugeConfig(BaseGaugeConfig): + gauge_class = 'BaseGauge' + + class assets(BaseGaugeConfig.assets): + terrain_name = "slope_5" # {type}_{level} + terrain_xml = '{ROBOGAUGE_ROOT_DIR}/resources/terrains/slope/slope_10.xml' + terrain_spawn_pos = [0, 0, 0] # x y z [m], robot freejoint spawn position on the terrain + + class goals: + class target_pos_velocity: # goal to reach a target position by velocity command + enabled = True + target_pos = [4, 0, 2.0] # x y z [m], target position in the environment, used for target position goal + lin_vel_x = 1.0 # +/- m/s + lin_vel_y = 1.0 # +/- m/s + ang_vel_yaw = 1.5 # +/- rad/s + max_cmd_duration = 20.0 # [s] maximum duration to reach the target position + reach_threshold = 0.1 diff --git a/robogauge/tasks/pipeline/base_pipeline.py b/robogauge/tasks/pipeline/base_pipeline.py index c2970c1..235cdd2 100644 --- a/robogauge/tasks/pipeline/base_pipeline.py +++ b/robogauge/tasks/pipeline/base_pipeline.py @@ -9,6 +9,7 @@ ''' import yaml import random +import traceback import numpy as np from pathlib import Path from copy import deepcopy @@ -104,7 +105,7 @@ class BasePipeline: if self.gauge.is_reset(sim_data): sim_data = self.reset_sim(sim_data) except Exception as e: - logger.error(f"❌ Pipeline execution failed with error: {e}") + logger.error(f"❌ Pipeline execution failed with error: {e},\n{traceback.format_exc()}") self.gauge.switch_to_next_goal() # save current goal metrics self.gauge.save_results() return logger.log_dir, e diff --git a/robogauge/tasks/robots/base_robot_config.py b/robogauge/tasks/robots/base_robot_config.py index 30d6414..9fa7664 100644 --- a/robogauge/tasks/robots/base_robot_config.py +++ b/robogauge/tasks/robots/base_robot_config.py @@ -21,7 +21,7 @@ class RobotConfig(Config): class control: device = 'cpu' # torch script model path - model_path = "{ROBOGAUGE_ROOT_DIR}/resources/models/go2/go2_cts_83501.pt" + model_path = "{ROBOGAUGE_ROOT_DIR}/resources/models/go2/go2_cts_max2_100k.pt" control_dt = 0.02 # 50 Hz control_type = 'P' # Position control support_goal: Literal['velocity', 'position'] = 'velocity' diff --git a/robogauge/tasks/robots/go2/go2_config.py b/robogauge/tasks/robots/go2/go2_config.py index 04562a3..f927675 100644 --- a/robogauge/tasks/robots/go2/go2_config.py +++ b/robogauge/tasks/robots/go2/go2_config.py @@ -20,7 +20,7 @@ class Go2Config(RobotConfig): class control(RobotConfig.control): device = 'cpu' - model_path = "{ROBOGAUGE_ROOT_DIR}/resources/models/go2/go2_cts_83501.pt" + model_path = "{ROBOGAUGE_ROOT_DIR}/resources/models/go2/go2_cts_max2_100k.pt" # model_path = "{ROBOGAUGE_ROOT_DIR}/resources/models/go2/go2_cts_cmd-1,1_38k.pt" control_dt = 0.02 # 50 Hz control_type = 'P' # Position control @@ -47,8 +47,8 @@ class Go2Config(RobotConfig): cmd = [2.0, 2.0, 0.25] class commands(RobotConfig.commands): - lin_vel_x = [-1.5, 1.5] # min max [m/s] - lin_vel_y = [-1, 1] # min max [m/s] + lin_vel_x = [-2.0, 2.0] # min max [m/s] + lin_vel_y = [-1.0, 1.0] # min max [m/s] lin_vel_z = None # min max [m/s] ang_vel_roll = None # min max [rad/s] ang_vel_pitch = None # min max [rad/s] diff --git a/robogauge/tasks/robots/go2/go2_moe.py b/robogauge/tasks/robots/go2/go2_moe.py index 67015ac..b75254d 100644 --- a/robogauge/tasks/robots/go2/go2_moe.py +++ b/robogauge/tasks/robots/go2/go2_moe.py @@ -15,7 +15,12 @@ from robogauge.tasks.robots.go2.go2 import Go2 class Go2MoE(Go2): def get_action(self, obs: np.ndarray): obs_tensor = torch.tensor(obs, dtype=torch.float32).unsqueeze(0).to(self.device) - action, weights = self.model(obs_tensor) + action, results = self.model(obs_tensor) + if isinstance(results, tuple): + weights, latent = results + latent = latent.detach().cpu().numpy().squeeze(0) + else: + weights = results action = action.detach().cpu().numpy().squeeze(0)[self.model2mj_idx] weights = weights.detach().cpu().numpy().squeeze(0) self.last_action = action diff --git a/robogauge/utils/radar_plot.py b/robogauge/utils/radar_plot.py deleted file mode 100644 index 80dbb9f..0000000 --- a/robogauge/utils/radar_plot.py +++ /dev/null @@ -1,183 +0,0 @@ -""" -python robogauge/utils/radar_plot.py \ - /home/xfy/Coding/robot_gauge/logs/go2_moe_flat_debug_multi/20251218-22-49-29_run_multi/aggregated_results.yaml \ - /home/xfy/Coding/robot_gauge/logs/go2_flat_debug_multi/20251218-22-59-04_run_multi/aggregated_results.yaml \ - --out logs/go2_flat_vs_moe_flat.png -""" -import matplotlib.pyplot as plt -config = { - "font.family": 'serif', # 衬线字体 - "figure.figsize": (6, 6), # 图像大小 - "font.size": 14, # 字号大小 - "mathtext.fontset": 'cm', # 渲染数学公式字体 - 'axes.unicode_minus': False # 显示负号 -} -plt.rcParams.update(config) - -import numpy as np -import yaml -import argparse -import os -import sys - -# 设置字体,尝试匹配参考图的衬线体风格 (如果系统没有会回退到默认) -plt.rcParams['font.family'] = 'serif' -plt.rcParams['font.serif'] = ['Times New Roman', 'DejaVu Serif', 'serif'] - -def parse_value_string(val_str): - if isinstance(val_str, (int, float)): - return float(val_str) - if isinstance(val_str, str): - if '±' in val_str: - return float(val_str.split('±')[0].strip()) - return float(val_str) - return 0.0 - -def load_data(file_paths): - all_data = [] - - # 指标键名 - metric_keys = [ - 'lin_vel_err', - 'ang_vel_err', - 'orientation_stability', - 'dof_limits', - 'torque_smoothness', - 'dof_power' - ] - - # 标签 (增加换行以避免拥挤) - labels_map = { - 'lin_vel_err': 'Lin Vel\nAccuracy', - 'ang_vel_err': 'Ang Vel\nAccuracy', - 'dof_limits': 'Joint Limits\nMargin', - 'dof_power': 'Energy\nEfficiency', - 'orientation_stability': 'Orientation\nStability', - 'torque_smoothness': 'Torque\nSmoothness', - } - - for path in file_paths: - if not os.path.exists(path): - continue - - with open(path, 'r', encoding='utf-8') as f: - content = yaml.safe_load(f) - - raw_path = content.get('model_path', 'Unknown_Model') - # 简化图例名称:只取文件名,去掉 .pt - model_name = os.path.basename(raw_path).replace('.pt', '') - - # 如果名称过长,可以考虑进一步截断,例如: - # if len(model_name) > 20: model_name = model_name[:10] + "..." + model_name[-5:] - - values = [] - for k in metric_keys: - if k in content: - raw_val = content[k]['mean'] - values.append(parse_value_string(raw_val)) - else: - values.append(0.0) - - all_data.append({'name': model_name, 'values': values}) - - return all_data, [labels_map[k] for k in metric_keys] - -def plot_radar(data_list, labels, output_file=None): - if not data_list: - print("No data to plot.") - return - - num_vars = len(labels) - angles = np.linspace(0, 2 * np.pi, num_vars, endpoint=False).tolist() - angles += angles[:1] # 闭合 - - # --- 颜色设置 --- - # 使用参考图类似的配色 (深蓝、浅蓝、绿等) - # 或者使用 'tab10', 'Set2' 等 - colors = plt.cm.get_cmap("tab10", len(data_list)) - - # 创建画布,稍微宽一点以便放图例 - fig, ax = plt.subplots(figsize=(10, 8), subplot_kw=dict(polar=True)) - - # --- 核心修改:调整布局 --- - # left=0.1, bottom=0.1, top=0.9 是为了给标题留空 - # right=0.75 是关键!这意味着图表只占画布左边 75% 的宽度,右边 25% 留给图例 - plt.subplots_adjust(left=0.05, right=0.75, top=0.9, bottom=0.1) - - # 设置方向 - ax.set_theta_offset(np.pi / 2) - ax.set_theta_direction(-1) - - # --- 绘制标签 --- - plt.xticks(angles[:-1], labels, color='#444444', size=13) - - # 标签对齐优化 - for label, angle in zip(ax.get_xticklabels(), angles[:-1]): - if angle in (0, np.pi): - label.set_horizontalalignment('center') - elif 0 < angle < np.pi: - label.set_horizontalalignment('left') - else: - label.set_horizontalalignment('right') - - # --- 绘制刻度 --- - ax.set_rlabel_position(0) - # 字体稍微调淡一点,不要抢眼 - plt.yticks([0.25, 0.50, 0.75, 1.00], ["0.25", "0.50", "0.75", "1.00"], - color="grey", size=10) - plt.ylim(0, 1.05) - - # 网格线:点状虚线,稍微粗一点 - ax.grid(True, color='gray', linestyle=':', linewidth=1.5, alpha=0.5) - ax.spines['polar'].set_visible(False) - - # --- 绘制数据 --- - # 加粗线条以匹配 bsuite 风格 - linewidth = 3.0 - - for idx, item in enumerate(data_list): - values = item['values'] - name = item['name'] - values_closed = values + values[:1] - - color = colors(idx) - - ax.plot(angles, values_closed, linewidth=linewidth, linestyle='-', label=name, color=color) - ax.fill(angles, values_closed, color=color, alpha=0.2) # 填充透明度低一点 - - # --- 核心修改:图例位置 --- - # bbox_to_anchor=(1.1, 0.2) 的意思是: - # 锚点位于坐标轴右侧(1.1倍宽位置),垂直方向在底部(0.2倍高位置) - # loc='upper left' 意思是图例的左上角对齐这个锚点 - legend = plt.legend( - loc='upper left', - bbox_to_anchor=(1.0, 0.1), # 调整这里的 0.3 可以上下移动图例 - title="Models", - title_fontsize=16, - fontsize=12, - frameon=False, # 无边框 - labelspacing=0.8 # 图例行间距 - ) - - # 设置图例标题对齐方式 (左对齐) - legend._legend_box.align = "left" - - plt.title('Multi-Model Performance Comparison', size=18, y=1.08, color='#333333') - - if output_file: - plt.savefig(output_file, dpi=300, bbox_inches='tight') # bbox_inches='tight' 会自动裁剪白边 - print(f"Plot saved to {output_file}") - else: - plt.show() - -if __name__ == "__main__": - parser = argparse.ArgumentParser() - parser.add_argument('files', metavar='F', type=str, nargs='+', help='YAML files') - parser.add_argument('--out', type=str, default=None, help='Output file') - - # 调试用(如果你直接运行脚本,请取消注释并填入你的文件名) - # sys.argv = ['plot.py', 'aggregated_results.yaml', 'aggregated_results2.yaml', '--out', 'fixed_radar.png'] - - args = parser.parse_args() - data, metrics_labels = load_data(args.files) - plot_radar(data, metrics_labels, output_file=args.out) diff --git a/robogauge/utils/visualize_results.py b/robogauge/utils/visualize_results.py new file mode 100644 index 0000000..9ee2aa3 --- /dev/null +++ b/robogauge/utils/visualize_results.py @@ -0,0 +1,170 @@ +import matplotlib.pyplot as plt +from pathlib import Path +import numpy as np +import yaml +import argparse +import os + +# --- 配置 Matplotlib 样式 --- +config = { + "font.family": 'serif', + "figure.figsize": (8, 6), + "font.size": 12, + "mathtext.fontset": 'cm', + 'axes.unicode_minus': False +} +plt.rcParams.update(config) +plt.rcParams['font.serif'] = ['Times New Roman', 'DejaVu Serif', 'serif'] + +def parse_value_string(val_str): + if isinstance(val_str, (int, float)): + return float(val_str) + if isinstance(val_str, str): + if '±' in val_str: + return float(val_str.split('±')[0].strip()) + return float(val_str) + return 0.0 + +def load_data(file_paths): + all_data = [] + metric_keys = ['lin_vel_err', 'ang_vel_err', 'orientation_stability', 'dof_limits', 'torque_smoothness', 'dof_power'] + labels_map = { + 'lin_vel_err': 'Lin Vel\nAccuracy', + 'ang_vel_err': 'Ang Vel\nAccuracy', + 'dof_limits': 'Joint Limits\nMargin', + 'dof_power': 'Energy\nEfficiency', + 'orientation_stability': 'Orientation\nStability', + 'torque_smoothness': 'Torque\nSmoothness', + } + + for path in file_paths: + if not os.path.exists(path): + print(f"[ERROR] File not found: {path}") + continue + with open(path, 'r', encoding='utf-8') as f: + content = yaml.safe_load(f) + + raw_path = content.get('model_path', 'Unknown_Model') + model_name = os.path.basename(raw_path).replace('.pt', '') + + values = [] + for k in metric_keys: + if k in content: + raw_val = content[k]['mean@50'] + values.append(parse_value_string(raw_val)) + else: + values.append(0.0) + all_data.append({'name': model_name, 'values': values}) + + return all_data, [labels_map[k] for k in metric_keys] + +def plot_radar(data_list, labels, output_file=None, r_range=(0, 1.05)): + """绘制雷达图""" + if not data_list: + return + + num_vars = len(labels) + angles = np.linspace(0, 2 * np.pi, num_vars, endpoint=False).tolist() + angles += angles[:1] + + colors = plt.cm.get_cmap("tab10", len(data_list)) + fig, ax = plt.subplots(figsize=(10, 8), subplot_kw=dict(polar=True)) + plt.subplots_adjust(left=0.05, right=0.75, top=0.9, bottom=0.1) + + ax.set_theta_offset(np.pi / 2) + ax.set_theta_direction(-1) + + plt.xticks(angles[:-1], labels, color='#444444', size=11) + + for label, angle in zip(ax.get_xticklabels(), angles[:-1]): + if angle in (0, np.pi): label.set_horizontalalignment('center') + elif 0 < angle < np.pi: label.set_horizontalalignment('left') + else: label.set_horizontalalignment('right') + + r_min, r_max = r_range + ax.set_ylim(r_min, r_max) + + ticks = np.linspace(r_min, r_max, 5) + plt.yticks(ticks, [f"{t:.2f}" for t in ticks], color="grey", size=10) + + ax.set_rlabel_position(180) + ax.grid(True, color='gray', linestyle=':', linewidth=1.5, alpha=0.5) + ax.spines['polar'].set_visible(False) + + linewidth = 2.5 + for idx, item in enumerate(data_list): + values_closed = item['values'] + [item['values'][0]] + color = colors(idx) + ax.plot(angles, values_closed, linewidth=linewidth, label=item['name'], color=color) + ax.fill(angles, values_closed, color=color, alpha=0.15) + + plt.legend(loc='lower left', bbox_to_anchor=(0.82, 0.0), title="Models", frameon=False) + plt.title('Performance Comparison (Radar)', size=16, y=1.06) + + if output_file: + plt.savefig(output_file, dpi=300, bbox_inches='tight') + print(f"Radar plot saved to {output_file}") + +def plot_bar(data_list, labels, output_file=None, h_range=(0, 1.05)): + """绘制柱状图""" + if not data_list: + return + + num_models = len(data_list) + num_metrics = len(labels) + + # 设置柱状图参数 + x = np.arange(num_metrics) + width = 0.8 / num_models # 自动计算柱子宽度 + + fig, ax = plt.subplots(figsize=(12, 6)) + colors = plt.cm.get_cmap("tab10", num_models) + + for idx, item in enumerate(data_list): + # 计算每个模型柱子的偏移量 + offset = (idx - (num_models - 1) / 2) * width + ax.bar(x + offset, item['values'], width, label=item['name'], color=colors(idx), alpha=0.8) + + ax.set_ylabel('Score / Value') + ax.set_title('Performance Comparison (Bar)', size=16) + ax.set_xticks(x) + # 将标签中的换行符去掉或处理,使其在柱状图中更美观 + clean_labels = [l.replace('\n', ' ') for l in labels] + ax.set_xticklabels(clean_labels, rotation=15, ha='right') + ax.legend(title="Models", bbox_to_anchor=(1.05, 1), loc='upper left', frameon=False) + ax.grid(axis='y', linestyle='--', alpha=0.7) + if h_range: + ax.set_ylim(h_range) + + plt.tight_layout() + + if output_file: + plt.savefig(output_file, dpi=300, bbox_inches='tight') + print(f"Bar plot saved to {output_file}") + +if __name__ == "__main__": + parser = argparse.ArgumentParser() + parser.add_argument('files_or_dir', metavar='F', type=str, nargs='+', help='YAML files or directory') + parser.add_argument('--out', type=str, help='Base name for output files') + parser.add_argument('--range', type=float, nargs=2, default=[0.0, 1.0], help='Axis range for radar and bar plots') + args = parser.parse_args() + + # 获取文件列表 + files = args.files_or_dir + if Path(args.files_or_dir[0]).is_dir(): + dir_path = Path(args.files_or_dir[0]) + files = [str(p) for p in dir_path.glob('*.yaml')] + + data, metrics_labels = load_data(files) + + # 处理输出文件名 + radar_out, bar_out = None, None + if args.out: + out_path = Path(args.out) + radar_out = str(out_path.with_name(f"{out_path.stem}_radar{out_path.suffix}")) + bar_out = str(out_path.with_name(f"{out_path.stem}_bar{out_path.suffix}")) + + # 分别调用绘图函数 + plot_radar(data, metrics_labels, output_file=radar_out, r_range=args.range) + plot_bar(data, metrics_labels, output_file=bar_out, h_range=args.range) + plt.show()