feat: 添加逆运动学对比实验入口

This commit is contained in:
2026-08-24 18:07:22 +08:00
parent 6df791b00f
commit d7ff769109
2 changed files with 87 additions and 10 deletions
@@ -6,7 +6,7 @@
**Architecture:** 新增一个仅供离线实验使用的脚本,复用现有 `PlacoIkSolver`、双臂 URDF、右臂 YAML 参数和关节命令限速函数。脚本只读加载 episode,将 30 Hz 位姿重采样到 90 Hz,然后让三种方法从同一初始关节状态独立复放,最后统一计算指标并输出 CSV、JSON、SVG、PNG 和 Markdown。
**Tech Stack:** Python 3.11、NumPy 2.2.6、h5py 3.16.0、Matplotlib 3.10.9、PyYAML、Placo 0.9.4、pytest、ROS2 Humble/colcon。
**Tech Stack:** Python 3.10、NumPy 2.2.6、h5py 3.16.0、Matplotlib 3.10.9、PyYAML、Placo 0.9.4、pytest、ROS2 Humble/colcon。
---
+86 -9
View File
@@ -1,8 +1,10 @@
from __future__ import annotations
import argparse
import csv
import json
import math
import subprocess
import time
from dataclasses import asdict, dataclass, replace
from pathlib import Path
@@ -48,6 +50,16 @@ plt.rcParams.update(
}
)
SOURCE_ROOT = Path(__file__).resolve().parents[2]
DEFAULT_EPISODE = Path("/home/robot/ACT_Data/tomato_pick/episode_0.hdf5")
DEFAULT_URDF = (
SOURCE_ROOT / "xr_rm_teleop" / "models" / "dual_rm75" / "Dual_arm.urdf"
)
DEFAULT_CONFIG = (
SOURCE_ROOT / "xr_rm_bringup" / "config" / "right_arm_rm75.yaml"
)
DEFAULT_OUTPUT = SOURCE_ROOT / "output" / "ik_comparison" / "episode_0"
@dataclass(frozen=True)
class EpisodeTrajectory:
@@ -633,13 +645,23 @@ def evaluate_methods(
max_speed=speed,
max_acceleration=acceleration,
)
max_difference = float(
np.max(np.abs(repeated.joints - result.joints))
)
success_equal = np.array_equal(
repeated.success, result.success
)
if not np.allclose(
repeated.joints,
result.joints,
atol=1e-10,
atol=1e-6,
rtol=0.0,
):
raise RuntimeError(f"{name} replay is not deterministic")
) or not success_equal:
raise RuntimeError(
f"{name} replay is not deterministic: "
f"max_joint_difference={max_difference:.3e} rad, "
f"success_equal={success_equal}"
)
timed.append(repeated.solve_durations_ms)
timing_matrix = np.stack(timed)
result = replace(
@@ -698,7 +720,9 @@ def plot_tracking_error(
linewidth=1.15,
label=label,
)
failed = ~result.success
failed = np.flatnonzero(~result.success)
if failed.size:
failed = failed[::max(1, failed.size // 80)]
axis.scatter(
result.times_s[failed],
series[failed],
@@ -708,10 +732,19 @@ def plot_tracking_error(
linewidths=0.7,
zorder=3,
)
axes[0].set_ylabel("位置误差 (mm)")
axes[1].set_ylabel("姿态误差 (°)")
axes[0].set_yscale("log")
axes[1].set_yscale("log")
axes[0].set_ylabel("位置误差 (mm,对数坐标)")
axes[1].set_ylabel("姿态误差 (°,对数坐标)")
axes[1].set_xlabel("时间 (s)")
axes[0].legend(frameon=False, ncol=3, loc="upper right")
axes[0].legend(
frameon=True,
facecolor="white",
framealpha=0.9,
edgecolor="none",
ncol=3,
loc="upper right",
)
_style_axes(axes)
_save_figure(fig, output_dir, "figure_2_11_tracking_error")
@@ -944,8 +977,8 @@ DLS 扫描的固定阻尼候选为 {", ".join(map(str, DLS_DAMPING_CANDIDATES))}
| --- | ---: | ---: | ---: | ---: | ---: | ---: |
{chr(10).join(rows)}
图 2-11 三种逆运动学方法的末端位置与姿态跟踪误差。曲线来自统一时间轴,叉号表示该
周期数值求解失败并保持上一安全关节状态。
图 2-11 三种逆运动学方法的末端位置与姿态跟踪误差。纵轴采用对数坐标以同时显示不同
数量级的误差,叉号稀疏标记数值求解失败并保持上一安全关节状态的周期
图 2-12 三种逆运动学方法的最大关节速度与最小归一化关节安全裕度。红色虚线表示
180°/s 输出速度上限,裕度越大表示离关节位置边界越远。
@@ -1000,3 +1033,47 @@ def write_outputs(
source_path,
sample_rate_hz,
)
def main() -> None:
parser = argparse.ArgumentParser(
description="离线比较 RM75 伪逆、DLS 和当前优化 QP"
)
parser.add_argument("--episode", type=Path, default=DEFAULT_EPISODE)
parser.add_argument("--urdf", type=Path, default=DEFAULT_URDF)
parser.add_argument("--config", type=Path, default=DEFAULT_CONFIG)
parser.add_argument("--output-dir", type=Path, default=DEFAULT_OUTPUT)
parser.add_argument("--timing-repeats", type=int, default=10)
args = parser.parse_args()
if args.timing_repeats <= 0:
parser.error("--timing-repeats must be positive")
source = load_episode(args.episode)
trajectory = resample_trajectory(source, 90.0)
config = load_right_config(args.config)
results, summaries, damping = evaluate_methods(
trajectory,
args.urdf,
config,
args.timing_repeats,
)
commit = subprocess.run(
["git", "rev-parse", "HEAD"],
cwd=SOURCE_ROOT,
check=True,
capture_output=True,
text=True,
).stdout.strip()
write_outputs(
args.output_dir,
results,
summaries,
selected_damping=damping,
source_path=source.source_path,
git_commit=commit,
)
print(f"results written to {args.output_dir}")
if __name__ == "__main__":
main()