| 21 | |
| 22 | |
| 23 | def generate_plot( |
| 24 | trajectories: list[Float[np.ndarray, "frame 3"]], |
| 25 | labels: list[str], |
| 26 | margin: float = 0.2, |
| 27 | ) -> Figure: |
| 28 | fig = plt.figure(figsize=plt.figaspect(1.0)) |
| 29 | ax = fig.add_subplot(1, 1, 1, projection="3d") |
| 30 | ax.set_proj_type("ortho") |
| 31 | for trajectory, label in zip(trajectories, labels): |
| 32 | xyz = rearrange(trajectory, "f xyz -> xyz f") |
| 33 | ax.plot3D(*xyz, label=label) |
| 34 | |
| 35 | # Set square axes. |
| 36 | points = np.concatenate(trajectories) |
| 37 | minima = points.min(axis=0) |
| 38 | maxima = points.max(axis=0) |
| 39 | span = (maxima - minima).max() * (1 + margin) |
| 40 | means = 0.5 * (maxima + minima) |
| 41 | starts = means - 0.5 * span |
| 42 | ends = means + 0.5 * span |
| 43 | ax.set_xlim(starts[0], ends[0]) |
| 44 | ax.set_ylim(starts[1], ends[1]) |
| 45 | ax.set_zlim(starts[2], ends[2]) |
| 46 | fig.legend() |
| 47 | return fig |
| 48 | |
| 49 | |
| 50 | @dataclass |