Skip to content

Commit 0abd740

Browse files
TrajectoryCollection
1 parent 054aec2 commit 0abd740

3 files changed

Lines changed: 475 additions & 10 deletions

File tree

‎docs/model_based_design.md‎

Lines changed: 10 additions & 10 deletions
Original file line numberDiff line numberDiff line change
@@ -63,21 +63,21 @@ This class does linear prediction and evaluations of these predictions. It is co
6363
This class represents a collection of ``Trajectory`` with the same
6464
``Model``. Methods include:
6565

66-
* ``plotTimecourse`` pieces together the timecourse for each ``DynamicModel``, and plots it with a vertical dashed line separating each ``DynamicModel``.
67-
* ``__eq__`` which checks if it's the same as another ``DynamicModelCollection`` by comparing each ``DynamicModel``.
68-
* ``__lt`` if the last timepoint of the first Trajector equals the first timepoint of the second trajectory.
69-
* ``split`` takes as input timepoints to create multiple Trajectory objects. A timepoint specifies the last time for the preceeding Trajectory and the first time for the next Trajectory.
66+
* Constructor takes a list of ``Tracjectory`` all of which have the same ``Model``. Error checking is done to ensure that timepoints do not overlap. The list is sorted using ``Trajectory.__lt__``. The constructor does not verify that adjacent ``Trajectory`` overlap their end_time and start_time.
67+
* ``plotTimecourse`` pieces together the timecourse for each ``DynamicModel``, and plots it with a vertical dashed line separating each ``Trajectory``. The arguments to this method are the kwards used by PlotOptions. Internally, the method uses PlotOptions. It should return PlotOptions.
68+
* ``__eq__`` which checks if it's the same as another ``TrajectoryCollection`` by comparing each ``Trajectory``.
69+
* ``split`` is a class method that takes as input timepoints to create multiple Trajectory objects. Its signature is ``split(cls, trajectory: Trajectory, timepoints: List[float])``.
70+
* A timepoint specifies the last time for the preceeding Trajectory and the first time for the next Trajectory.
71+
* Returns ``TrajectoryCollection``.
72+
* If a split time falls between existing timepoints (i.e., not exactly in timepoint_arr), it snaps to the nearest timepoint.
7073

7174
### ``MultipleLinearPredictor``
7275

7376
This class performs piece-wise linear prediction. The times at which
74-
there is a partition of the linear model is a "split point". Using slicing, we can easily construct the DynamicModels for a collection of split points. (Of course, all will have the same StaticModel.) If there are n split points, then there are n + 1 linear models.
75-
Note that the old Trajectory.sequentialPartition / nonsequentialPartition aren't mentioned beause their implementation is deferred.
77+
there is a partition of the linear model is a "split point".
7678

77-
* When splitting a ``DynamicModel``, slicing is used, not simulation.
78-
* ``split`` can be called with specific split times or without any split time specified. A split time t1 specifies the time at which the previous DynamicModel ends and the second timepoint of the new Dynamics model.
79-
**Does this belong here?**
80-
* ``predict``
79+
* Constructor has the arguments ``TrajectoryCollection`` and ``num_step``, the number of steps ahead for which prediction is done.
80+
* ``predict`` uses LinearPredictor.predict for each Trajectory.
8181
* ``score``
8282
* ``plotPrediction`` The plot shows predicted and actual (simulated) values with vertical dashed lines to indicate regions for submodels.
8383

‎src/trajectory_collection.py‎

Lines changed: 120 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,120 @@
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

Comments
 (0)