from __future__ import annotations from dataclasses import dataclass from pathlib import Path import numpy as np @dataclass(frozen=True) class EpisodeTrajectory: source_path: Path times_s: np.ndarray target_poses: np.ndarray initial_joints: np.ndarray def _normalized_quaternion(values: np.ndarray) -> np.ndarray: quaternion = np.asarray(values, dtype=float) if quaternion.shape != (4,) or not np.isfinite(quaternion).all(): raise ValueError("quaternion must contain 4 finite values") norm = float(np.linalg.norm(quaternion)) if norm <= 1e-12: raise ValueError("quaternion norm must be positive") return quaternion / norm def _slerp_quaternion( start: np.ndarray, end: np.ndarray, fraction: float, ) -> np.ndarray: first = _normalized_quaternion(start) second = _normalized_quaternion(end) dot = float(np.dot(first, second)) if dot < 0.0: second = -second dot = -dot dot = float(np.clip(dot, -1.0, 1.0)) if dot > 1.0 - 1e-8: return _normalized_quaternion( first + float(fraction) * (second - first) ) angle = float(np.arccos(dot)) sine = float(np.sin(angle)) return _normalized_quaternion( np.sin((1.0 - fraction) * angle) / sine * first + np.sin(fraction * angle) / sine * second ) def resample_trajectory( trajectory: EpisodeTrajectory, sample_rate_hz: float, ) -> EpisodeTrajectory: if not np.isfinite(sample_rate_hz) or sample_rate_hz <= 0.0: raise ValueError("sample_rate_hz must be finite and positive") source_times = np.asarray(trajectory.times_s, dtype=float) poses = np.asarray(trajectory.target_poses, dtype=float) if source_times.ndim != 1 or poses.shape != (source_times.size, 7): raise ValueError("trajectory must contain N timestamps and N x 7 poses") if source_times.size < 2 or np.any(np.diff(source_times) <= 0.0): raise ValueError("trajectory timestamps must be strictly increasing") duration = float(source_times[-1] - source_times[0]) count = int(round(duration * sample_rate_hz)) + 1 target_times = np.linspace(source_times[0], source_times[-1], count) target_poses = np.empty((count, 7), dtype=float) for axis in range(3): target_poses[:, axis] = np.interp( target_times, source_times, poses[:, axis], ) for index, timestamp in enumerate(target_times): right = int(np.searchsorted(source_times, timestamp, side="right")) right = min(max(right, 1), source_times.size - 1) left = right - 1 interval = source_times[right] - source_times[left] fraction = float((timestamp - source_times[left]) / interval) target_poses[index, 3:] = _slerp_quaternion( poses[left, 3:], poses[right, 3:], fraction ) target_poses[0] = poses[0] target_poses[-1] = poses[-1] return EpisodeTrajectory( source_path=trajectory.source_path, times_s=target_times - target_times[0], target_poses=target_poses, initial_joints=np.asarray(trajectory.initial_joints, dtype=float).copy(), )