MCPcopy Create free account
hub / github.com/dcharatan/flowmap / generate_plot

Function generate_plot

flowmap/visualization/visualizer_trajectory.py:23–47  ·  view source on GitHub ↗
(
    trajectories: list[Float[np.ndarray, "frame 3"]],
    labels: list[str],
    margin: float = 0.2,
)

Source from the content-addressed store, hash-verified

21
22
23def 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

Callers 1

visualizeMethod · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected