| |
| """ |
| Plot comparison between real and simulated observations using Plotly. |
| """ |
|
|
| import argparse |
| import pickle |
| import numpy as np |
| import plotly.graph_objects as go |
| from plotly.subplots import make_subplots |
| from pathlib import Path |
|
|
|
|
| def load_observations(pkl_path: str): |
| """Load observations from pickle file.""" |
| with open(pkl_path, 'rb') as f: |
| data = pickle.load(f) |
|
|
| if isinstance(data, dict): |
| if 'observations' in data and 'timestamps' in data: |
| observations = data['observations'] |
| timestamps = data['timestamps'] |
| else: |
| raise ValueError("Dictionary must contain 'observations' and 'timestamps' keys") |
| elif isinstance(data, list): |
| if len(data) == 0: |
| raise ValueError("Empty data list") |
|
|
| if isinstance(data[0], dict) and 'timestamp' in data[0] and 'observation' in data[0]: |
| timestamps = [item['timestamp'] for item in data] |
| observations = [item['observation'] for item in data] |
| elif isinstance(data[0], tuple): |
| timestamps = [item[0] for item in data] |
| observations = [item[1] for item in data] |
| else: |
| observations = data |
| timestamps = [i * 0.02 for i in range(len(observations))] |
| else: |
| raise ValueError(f"Unsupported data format: {type(data)}") |
|
|
| return np.array(observations), np.array(timestamps) |
|
|
|
|
| def plot_comparison(real_obs, real_ts, sim_obs=None, sim_ts=None): |
| """ |
| Plot comparison between real and simulated observations using Plotly. |
| If sim_obs is None, only plots real data. |
| """ |
|
|
| |
| joint_names = [ |
| 'L_hip_yaw', 'L_hip_roll', 'L_hip_pitch', 'L_knee', 'L_ankle', |
| 'neck_pitch', 'head_pitch', 'head_yaw', 'head_roll', |
| 'R_hip_yaw', 'R_hip_roll', 'R_hip_pitch', 'R_knee', 'R_ankle' |
| ] |
|
|
| obs_dim = real_obs.shape[1] if sim_obs is None else min(real_obs.shape[1], sim_obs.shape[1]) |
|
|
| |
| base_ang_vel_start = 0 |
| gravity_start = 3 |
| joint_pos_start = 6 |
| joint_vel_start = 20 |
| action_start = 34 |
|
|
| |
| subplot_titles = [] |
|
|
| |
| subplot_titles.extend(['<b>BASE ANG VEL</b><br>ω_x', 'ω_y', 'ω_z', '']) |
|
|
| |
| subplot_titles.extend(['<b>Raw Accelero</b><br>g_x', 'g_y', 'g_z', '']) |
|
|
| |
| subplot_titles.append(f'<b>JOINT POSITIONS</b><br>{joint_names[0]}') |
| subplot_titles.extend(joint_names[1:14]) |
| subplot_titles.extend(['', '']) |
|
|
| |
| subplot_titles.append(f'<b>JOINT VELOCITIES</b><br>{joint_names[0]}') |
| subplot_titles.extend(joint_names[1:14]) |
| subplot_titles.extend(['', '']) |
|
|
| |
| subplot_titles.append(f'<b>ACTIONS</b><br>{joint_names[0]}') |
| subplot_titles.extend(joint_names[1:14]) |
| subplot_titles.extend(['', '']) |
|
|
| num_rows = 14 |
| fig = make_subplots( |
| rows=num_rows, cols=4, |
| subplot_titles=subplot_titles, |
| vertical_spacing=0.02, |
| horizontal_spacing=0.05, |
| row_heights=[1]*num_rows, |
| ) |
|
|
| plot_idx = 0 |
|
|
| |
| command_data = [] |
|
|
| def add_traces(row, col, real_data, sim_data=None, y_range=None): |
| """Helper to add real and sim traces to a subplot.""" |
| fig.add_trace( |
| go.Scatter(x=real_ts, y=real_data, name='Real', |
| line=dict(color='blue', width=1.5), |
| showlegend=(plot_idx == 0)), |
| row=row, col=col |
| ) |
| if sim_data is not None: |
| fig.add_trace( |
| go.Scatter(x=sim_ts, y=sim_data, name='Sim', |
| line=dict(color='red', width=1.5, dash='dash'), |
| showlegend=(plot_idx == 0)), |
| row=row, col=col |
| ) |
| if y_range: |
| fig.update_yaxes(range=y_range, row=row, col=col) |
|
|
| base_ang_vel_data = [] |
| gravity_data = [] |
| joint_pos_data = [] |
| joint_vel_data = [] |
| action_data = [] |
|
|
| |
| for i in range(3): |
| row, col = divmod(plot_idx, 4) |
| row += 1 |
| col += 1 |
| base_ang_vel_data.append(real_obs[:, base_ang_vel_start+i]) |
| if sim_obs is not None: |
| base_ang_vel_data.append(sim_obs[:, base_ang_vel_start+i]) |
| add_traces(row, col, real_obs[:, base_ang_vel_start+i], None if sim_obs is None else sim_obs[:, base_ang_vel_start+i]) |
| fig.update_yaxes(title_text='rad/s', row=row, col=col) |
| plot_idx += 1 |
|
|
| |
| plot_idx += 1 |
|
|
| |
| for i in range(3): |
| row, col = divmod(plot_idx, 4) |
| row += 1 |
| col += 1 |
| gravity_data.append(real_obs[:, gravity_start+i]) |
| if sim_obs is not None: |
| gravity_data.append(sim_obs[:, gravity_start+i]) |
| add_traces(row, col, real_obs[:, gravity_start+i], None if sim_obs is None else sim_obs[:, gravity_start+i]) |
| fig.update_yaxes(title_text='g', row=row, col=col) |
| plot_idx += 1 |
|
|
| |
| plot_idx += 1 |
|
|
| |
| for i in range(14): |
| row, col = divmod(plot_idx, 4) |
| row += 1 |
| col += 1 |
| if joint_pos_start + i < obs_dim: |
| joint_pos_data.append(real_obs[:, joint_pos_start+i]) |
| if sim_obs is not None: |
| joint_pos_data.append(sim_obs[:, joint_pos_start+i]) |
| add_traces(row, col, real_obs[:, joint_pos_start+i], None if sim_obs is None else sim_obs[:, joint_pos_start+i]) |
| fig.update_yaxes(title_text='rad', row=row, col=col) |
| plot_idx += 1 |
|
|
| |
| plot_idx += 2 |
|
|
| |
| for i in range(14): |
| row, col = divmod(plot_idx, 4) |
| row += 1 |
| col += 1 |
| if joint_vel_start + i < obs_dim: |
| joint_vel_data.append(real_obs[:, joint_vel_start+i]) |
| if sim_obs is not None: |
| joint_vel_data.append(sim_obs[:, joint_vel_start+i]) |
| add_traces(row, col, real_obs[:, joint_vel_start+i], None if sim_obs is None else sim_obs[:, joint_vel_start+i]) |
| fig.update_yaxes(title_text='rad/s', row=row, col=col) |
| plot_idx += 1 |
|
|
| |
| plot_idx += 2 |
|
|
| |
| for i in range(14): |
| row, col = divmod(plot_idx, 4) |
| row += 1 |
| col += 1 |
| if action_start + i < obs_dim: |
| action_data.append(real_obs[:, action_start+i]) |
| if sim_obs is not None: |
| action_data.append(sim_obs[:, action_start+i]) |
| add_traces(row, col, real_obs[:, action_start+i], None if sim_obs is None else sim_obs[:, action_start+i]) |
| fig.update_yaxes(title_text='action', row=row, col=col) |
| fig.update_xaxes(title_text='Time (s)', row=row, col=col) |
| plot_idx += 1 |
|
|
| |
| def compute_range(data_list): |
| if not data_list: |
| return None |
| all_data = np.concatenate([d.flatten() for d in data_list]) |
| y_min, y_max = np.min(all_data), np.max(all_data) |
| margin = (y_max - y_min) * 0.1 |
| return [y_min - margin, y_max + margin] |
|
|
| base_ang_vel_range = compute_range(base_ang_vel_data) |
| gravity_range = compute_range(gravity_data) |
| joint_pos_range = compute_range(joint_pos_data) |
| joint_vel_range = compute_range(joint_vel_data) |
| action_range = compute_range(action_data) |
|
|
| |
| plot_idx = 0 |
|
|
| for i in range(3): |
| row, col = divmod(plot_idx, 4) |
| fig.update_yaxes(range=base_ang_vel_range, row=row+1, col=col+1) |
| plot_idx += 1 |
| plot_idx += 1 |
|
|
| for i in range(3): |
| row, col = divmod(plot_idx, 4) |
| fig.update_yaxes(range=gravity_range, row=row+1, col=col+1) |
| plot_idx += 1 |
| plot_idx += 1 |
|
|
| for i in range(14): |
| row, col = divmod(plot_idx, 4) |
| fig.update_yaxes(range=joint_pos_range, row=row+1, col=col+1) |
| plot_idx += 1 |
| plot_idx += 2 |
|
|
| for i in range(14): |
| row, col = divmod(plot_idx, 4) |
| fig.update_yaxes(range=joint_vel_range, row=row+1, col=col+1) |
| plot_idx += 1 |
| plot_idx += 2 |
|
|
| for i in range(14): |
| row, col = divmod(plot_idx, 4) |
| fig.update_yaxes(range=action_range, row=row+1, col=col+1) |
| plot_idx += 1 |
|
|
| |
| title = 'Real vs Simulated Observations Comparison' if sim_obs is not None else 'Real Robot Observations' |
|
|
| fig.update_layout( |
| title_text=title, |
| title_font_size=24, |
| height=4600, |
| width=1600, |
| showlegend=True, |
| legend=dict(x=0.85, y=0.99, bgcolor='rgba(255,255,255,0.8)'), |
| hovermode='x unified' |
| ) |
|
|
| fig.show() |
|
|
|
|
| def main(): |
| parser = argparse.ArgumentParser( |
| description="Compare real and simulated observations (Plotly version)" |
| ) |
| parser.add_argument("real_pkl", type=str, |
| help="Path to .pkl file with real robot observations") |
| parser.add_argument("sim_pkl", type=str, nargs='?', default=None, |
| help="Path to .pkl file with simulated observations (optional)") |
| args = parser.parse_args() |
|
|
| |
| if not Path(args.real_pkl).exists(): |
| print(f"Error: {args.real_pkl} not found") |
| return 1 |
|
|
| |
| print(f"Loading real observations from {args.real_pkl}...") |
| real_obs, real_ts = load_observations(args.real_pkl) |
| print(f"Loaded {len(real_obs)} real observations (shape: {real_obs.shape})") |
|
|
| if args.sim_pkl: |
| if not Path(args.sim_pkl).exists(): |
| print(f"Error: {args.sim_pkl} not found") |
| return 1 |
| print(f"Loading simulated observations from {args.sim_pkl}...") |
| sim_obs, sim_ts = load_observations(args.sim_pkl) |
| print(f"Loaded {len(sim_obs)} simulated observations (shape: {sim_obs.shape})") |
| else: |
| print("No sim data provided, plotting real data only") |
| sim_obs, sim_ts = None, None |
|
|
| |
| print(f"\nGenerating interactive comparison plots...") |
| plot_comparison(real_obs, real_ts, sim_obs, sim_ts) |
|
|
| return 0 |
|
|
|
|
| if __name__ == "__main__": |
| exit(main()) |
|
|