|
| 1 | +"""Collection of Trajectory objects sharing the same Model.""" |
| 2 | + |
| 3 | +import src.constants as cn # type: ignore |
| 4 | +from src.plot_options import PlotOptions # type: ignore |
| 5 | +from src.trajectory import Trajectory # type: ignore |
| 6 | + |
| 7 | +import numpy as np # type: ignore |
| 8 | +from typing import List |
| 9 | + |
| 10 | + |
| 11 | +class TrajectoryCollection(object): |
| 12 | + """An ordered, non-overlapping collection of Trajectory sharing one Model. |
| 13 | +
|
| 14 | + Use split() to partition a single Trajectory into a TrajectoryCollection. |
| 15 | + """ |
| 16 | + |
| 17 | + def __init__(self, trajectories: List[Trajectory]) -> None: |
| 18 | + """ |
| 19 | + Parameters |
| 20 | + ---------- |
| 21 | + trajectories : List[Trajectory] |
| 22 | + All must share the same Model. Sorted by Trajectory.__lt__; |
| 23 | + overlapping time ranges raise ValueError. |
| 24 | + """ |
| 25 | + # Error checking |
| 26 | + if not trajectories: |
| 27 | + raise ValueError("trajectories must not be empty.") |
| 28 | + self.trajectories = sorted(trajectories) |
| 29 | + self.model = self.trajectories[0].model |
| 30 | + for traj in trajectories[1:]: |
| 31 | + if traj.model != self.model: |
| 32 | + raise ValueError( |
| 33 | + "All trajectories must share the same Model.") |
| 34 | + self.start_time = self.trajectories[0].start_time |
| 35 | + self.end_time = self.trajectories[-1].end_time |
| 36 | + |
| 37 | + def isConsecutive(self) -> bool: |
| 38 | + """True iff trajectories are consecutive (no gaps or overlaps).""" |
| 39 | + for t1, t2 in zip(self.trajectories, self.trajectories[1:]): |
| 40 | + if not np.isclose(t1.end_time, t2.start_time): |
| 41 | + return False |
| 42 | + return True |
| 43 | + |
| 44 | + def __eq__(self, other: object) -> bool: |
| 45 | + if not isinstance(other, TrajectoryCollection): |
| 46 | + raise ValueError("Can only compare TrajectoryCollection to another ") |
| 47 | + if len(self.trajectories) != len(other.trajectories): |
| 48 | + return False |
| 49 | + return all(t1 == t2 for t1, t2 in zip(self.trajectories, other.trajectories)) |
| 50 | + |
| 51 | + def plotTimecourse(self, **kwargs) -> PlotOptions: |
| 52 | + """Plot the pieced-together timecourse with dashed vertical separators. |
| 53 | +
|
| 54 | + Each trajectory's timecourse is drawn continuously; a vertical dashed |
| 55 | + line marks the boundary between adjacent trajectories. One legend |
| 56 | + entry per species. |
| 57 | +
|
| 58 | + Parameters |
| 59 | + ---------- |
| 60 | + **kwargs |
| 61 | + Passed to PlotOptions. Supported keys: ax, fig, title, xlabel, |
| 62 | + ylabel, legend, xlim, ylim, model_name. |
| 63 | +
|
| 64 | + Returns |
| 65 | + ------- |
| 66 | + PlotOptions |
| 67 | + """ |
| 68 | + plot_options = PlotOptions(**kwargs) |
| 69 | + ax = plot_options.ax |
| 70 | + for i, name in enumerate(self.model.species_names): |
| 71 | + color = f"C{i}" |
| 72 | + for j, traj in enumerate(self.trajectories): |
| 73 | + label = name if j == 0 else None |
| 74 | + ax.plot( |
| 75 | + traj.timecourse_df.index, |
| 76 | + traj.timecourse_df[name], |
| 77 | + color=color, |
| 78 | + label=label, |
| 79 | + ) |
| 80 | + for traj in self.trajectories[:-1]: |
| 81 | + ax.axvline(x=traj.end_time, color="black", linestyle="--", |
| 82 | + linewidth=0.8) |
| 83 | + plot_options.apply() |
| 84 | + return plot_options |
| 85 | + |
| 86 | + @classmethod |
| 87 | + def split(cls, |
| 88 | + trajectory: Trajectory, |
| 89 | + timepoints: List[float]) -> "TrajectoryCollection": |
| 90 | + """Partition a Trajectory at the given split times. |
| 91 | +
|
| 92 | + Each split time becomes the shared end/start boundary between adjacent |
| 93 | + sub-trajectories. Times not exactly in timepoint_arr snap to the |
| 94 | + nearest available timepoint. Split times that snap to start_time or |
| 95 | + end_time are ignored. |
| 96 | +
|
| 97 | + Parameters |
| 98 | + ---------- |
| 99 | + trajectory : Trajectory |
| 100 | + timepoints : List[float] |
| 101 | +
|
| 102 | + Returns |
| 103 | + ------- |
| 104 | + TrajectoryCollection |
| 105 | + """ |
| 106 | + tp_arr = trajectory.timepoint_arr |
| 107 | + snapped = [] |
| 108 | + for t in timepoints: |
| 109 | + idx = int(np.argmin(np.abs(tp_arr - t))) |
| 110 | + snapped.append(float(tp_arr[idx])) |
| 111 | + snapped = sorted({ |
| 112 | + t for t in snapped |
| 113 | + if trajectory.start_time < t < trajectory.end_time |
| 114 | + }) |
| 115 | + boundaries = [trajectory.start_time] + snapped + [trajectory.end_time] |
| 116 | + sub_trajectories = [ |
| 117 | + trajectory.makeSubmodel(boundaries[i], boundaries[i + 1]) |
| 118 | + for i in range(len(boundaries) - 1) |
| 119 | + ] |
| 120 | + return cls(sub_trajectories) |
0 commit comments