marionette-js / tests /scripts /dance_plot.py
RemiFabre
tests/scripts: dance playback comparison (python_native vs marionette_style)
7ea69d7
Raw
History Blame Contribute Delete
5.43 kB
"""Render comparison plots for dance_compare.py CSV output.
For one dance, generate a figure showing the same channels across
both methods (python_native vs marionette_style), with N runs of
each overlaid. Channels: head Y translation, head pitch (derived
from rotation matrix), antenna_left, antenna_right.
"""
import csv
import glob
import sys
from pathlib import Path
import matplotlib
matplotlib.use("Agg")
import matplotlib.pyplot as plt
import numpy as np
DANCE = sys.argv[1] if len(sys.argv) > 1 else "head_tilt_roll"
CSV_DIR = Path("/Users/remi/Downloads/dance-compare")
OUT_PNG = Path(f"/Users/remi/Downloads/dance-compare-{DANCE}.png")
def load(path):
rows = []
with open(path) as f:
for r in csv.DictReader(f):
rows.append({k: float(v) if v else None for k, v in r.items()})
return rows
def head_rpy(rows, prefix):
"""Extract (time, roll_deg, pitch_deg, yaw_deg) sequences from the
rotation submatrix of the stored 4x4 head transform.
Flat row-major indices:
[0 1 2 3 ] [R00 R01 R02 X ]
[4 5 6 7 ] [R10 R11 R12 Y ]
[8 9 10 11] [R20 R21 R22 Z ]
[12 13 14 15] [0 0 0 1 ]
Standard ZYX Euler decomposition:
roll = atan2( R21, R22 ) around X
pitch = atan2(-R20, sqrt(R21² + R22²)) around Y
yaw = atan2( R10, R00 ) around Z
"""
t = np.array([r["t_s"] for r in rows])
R00 = np.array([r[f"{prefix}h0"] for r in rows])
R10 = np.array([r[f"{prefix}h4"] for r in rows])
R20 = np.array([r[f"{prefix}h8"] for r in rows])
R21 = np.array([r[f"{prefix}h9"] for r in rows])
R22 = np.array([r[f"{prefix}h10"] for r in rows])
roll = np.degrees(np.arctan2(R21, R22))
pitch = np.degrees(np.arctan2(-R20, np.sqrt(R21 * R21 + R22 * R22)))
yaw = np.degrees(np.arctan2(R10, R00))
return t, roll, pitch, yaw
def antennas(rows, prefix):
t = np.array([r["t_s"] for r in rows])
l = np.array([r[f"{prefix}ant_l_deg"] for r in rows])
r = np.array([r[f"{prefix}ant_r_deg"] for r in rows])
return t, l, r
def plot_method_column(axes, csvs, method_label, color):
"""Plot all runs of one method into axes[0..4]. Channels:
roll, pitch, yaw, left antenna, right antenna."""
cmd_drawn = False
for path in csvs:
rows = load(path)
if not rows: continue
if not cmd_drawn:
t_c, roll_c, pitch_c, yaw_c = head_rpy(rows, "cmd_")
t_ca, l_c, r_c = antennas(rows, "cmd_")
axes[0].plot(t_c, roll_c, color="#444", linestyle="--", linewidth=1.0, label="commanded")
axes[1].plot(t_c, pitch_c, color="#444", linestyle="--", linewidth=1.0, label="commanded")
axes[2].plot(t_c, yaw_c, color="#444", linestyle="--", linewidth=1.0, label="commanded")
axes[3].plot(t_ca, l_c, color="#444", linestyle="--", linewidth=1.0, label="commanded")
axes[4].plot(t_ca, r_c, color="#444", linestyle="--", linewidth=1.0, label="commanded")
cmd_drawn = True
t_a, roll_a, pitch_a, yaw_a = head_rpy(rows, "act_")
t_aa, l_a, r_a = antennas(rows, "act_")
axes[0].plot(t_a, roll_a, color=color, alpha=0.55, linewidth=0.9)
axes[1].plot(t_a, pitch_a, color=color, alpha=0.55, linewidth=0.9)
axes[2].plot(t_a, yaw_a, color=color, alpha=0.55, linewidth=0.9)
axes[3].plot(t_aa, l_a, color=color, alpha=0.55, linewidth=0.9)
axes[4].plot(t_aa, r_a, color=color, alpha=0.55, linewidth=0.9)
def main():
python_csvs = sorted(glob.glob(str(CSV_DIR / f"{DANCE}-python_native-run*.csv")))
mariontte_csvs = sorted(glob.glob(str(CSV_DIR / f"{DANCE}-marionette_style-run*.csv")))
print(f"python_native runs: {len(python_csvs)}")
print(f"marionette_style runs: {len(mariontte_csvs)}")
fig, axes = plt.subplots(5, 2, figsize=(13, 12), sharex=True)
fig.suptitle(f"Dance playback comparison — {DANCE}\n"
f"Same motion, two streaming strategies, {len(python_csvs)} runs each",
fontsize=13, fontweight="bold")
plot_method_column(axes[:, 0], python_csvs, "python_native", "#1f77b4")
plot_method_column(axes[:, 1], mariontte_csvs, "marionette_style", "#2ca02c")
titles = ["Head roll (deg)", "Head pitch (deg)", "Head yaw (deg)", "Left antenna (deg)", "Right antenna (deg)"]
for row, t in enumerate(titles):
axes[row, 0].set_ylabel(t, fontsize=10)
for col in (0, 1):
axes[row, col].grid(True, alpha=0.3)
axes[row, col].tick_params(labelsize=8)
axes[0, 0].set_title("Python native (3 calls/tick @ 100 Hz, no lead comp)", fontsize=10)
axes[0, 1].set_title("Marionette style (combined set_target @ 50 Hz, lead 90/205 ms)", fontsize=10)
for col in (0, 1):
axes[-1, col].set_xlabel("t (s)", fontsize=10)
axes[0, col].legend(fontsize=8, loc="upper right")
# Share y-axis per channel for fair visual comparison
for row in range(5):
ymin = min(axes[row, 0].get_ylim()[0], axes[row, 1].get_ylim()[0])
ymax = max(axes[row, 0].get_ylim()[1], axes[row, 1].get_ylim()[1])
axes[row, 0].set_ylim(ymin, ymax)
axes[row, 1].set_ylim(ymin, ymax)
fig.tight_layout(rect=(0, 0, 1, 0.95))
fig.savefig(OUT_PNG, dpi=140)
print(f"Wrote {OUT_PNG}")
if __name__ == "__main__":
main()