Skip to content

Commit fc4399d

Browse files
sueogluÖykü Süoglu
andauthored
Add ep.pl.timeseries() to visualize variables over time (#994)
* timeseries() plotting function with holoviews --------- Co-authored-by: Öykü Süoglu <oeykue.sueoglu@helmholtz-munich.de>
1 parent c213b85 commit fc4399d

5 files changed

Lines changed: 281 additions & 0 deletions

File tree

55.9 KB
Loading

docs/api/plot_index.md

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -26,6 +26,7 @@ For most tools and for some preprocessing functions, you will find a plotting fu
2626
plot.ranking
2727
plot.dendrogram
2828
plot.catplot
29+
plot.timeseries
2930
```
3031

3132
## Quality Control and missing values

ehrapy/plot/__init__.py

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -43,6 +43,7 @@
4343
violin,
4444
)
4545
from ehrapy.plot._survival_analysis import cox_ph_forestplot, kaplan_meier, ols
46+
from ehrapy.plot._timeseries import timeseries
4647
from ehrapy.plot.causal_inference._dowhy import causal_effect
4748
from ehrapy.plot.feature_ranking._feature_importances import rank_features_supervised
4849

ehrapy/plot/_timeseries.py

Lines changed: 164 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,164 @@
1+
from __future__ import annotations
2+
3+
from typing import TYPE_CHECKING, Any
4+
5+
import holoviews as hv
6+
import numpy as np
7+
import pandas as pd
8+
9+
if TYPE_CHECKING:
10+
from collections.abc import Sequence
11+
12+
from ehrdata import EHRData
13+
14+
15+
def timeseries(
16+
edata: EHRData,
17+
*,
18+
obs_names: str | int | Sequence[str | int] | None = None,
19+
var_names: str | Sequence[str] | None = None,
20+
tem_names: Any | Sequence[Any] | slice | None = None,
21+
layer: str = "tem_data",
22+
overlay: bool = False,
23+
xlabel: str | None = None,
24+
ylabel: str | None = None,
25+
width: int | None = 600,
26+
height: int | None = 400,
27+
title: str | None = None,
28+
) -> hv.Overlay | hv.Layout:
29+
"""Plot time series from a 3D EHRData object.
30+
31+
Selection logic:
32+
obs_names, var_names, tem_names select labels from `edata.obs_names`, `edata.var_names`, `edata.tem.index`.
33+
Use :class:`slice` (e.g. ``slice(0, 5)``) for positional selection along the axes.
34+
35+
Args:
36+
edata: Central data object.
37+
obs_names: Unique observation identifier(s) to plot.
38+
var_names: Variable name or list of variable names in `edata.var_names` to plot.
39+
tem_names: Time indices to plot.
40+
layer: layer to use for time series data.
41+
overlay: Whether to overlay multiple observations in a single plot (True) or create subplots (False).
42+
xlabel: The x-axis label text.
43+
ylabel: The y-axis label text.
44+
width: Plot width in pixels.
45+
height: Plot height in pixels.
46+
title: Set the title of the plot.
47+
48+
Returns:
49+
HoloViews Overlay (if overlay=True) or Layout (if overlay=False) object representing the time series plot(s).
50+
51+
Examples:
52+
>>> import ehrapy as ep
53+
>>> import ehrdata as ed
54+
>>> edata = ed.dt.ehrdata_blobs(n_variables=10, n_observations=5, base_timepoints=100)
55+
>>> ep.pl.timeseries(edata, obs_names="1", var_names=["feature_1", "feature_2"], tem_names=slice(0, 10))
56+
57+
.. image:: /_static/docstring_previews/timeseries_plot.png
58+
"""
59+
opts_dict: dict[str, Any] = {}
60+
if width is not None:
61+
opts_dict["width"] = width
62+
if height is not None:
63+
opts_dict["height"] = height
64+
if xlabel is not None:
65+
opts_dict["xlabel"] = xlabel
66+
if ylabel is not None:
67+
opts_dict["ylabel"] = ylabel
68+
opts_dict["shared_axes"] = True
69+
opts_dict["legend_position"] = "right"
70+
71+
if layer not in edata.layers:
72+
raise KeyError(f"Layer {layer!r} not found in edata.layers. Available layers: {list(edata.layers)}")
73+
mtx = np.asarray(edata.layers[layer])
74+
if mtx.ndim != 3:
75+
raise ValueError(f"Layer {layer!r} must be 3D (n_obs, n_vars, n_time), got shape {mtx.shape}.")
76+
77+
obs_pos, obs_labels = _resolve_axis(pd.Index(edata.obs_names), obs_names, "obs_names")
78+
var_pos, var_labels = _resolve_axis(pd.Index(edata.var_names), var_names, "var_names")
79+
tem_pos, tem_labels = _resolve_axis(pd.Index(edata.tem.index), tem_names, "tem_names")
80+
81+
if obs_pos.size == 0:
82+
raise ValueError("No observations selected (obs_names resolved to empty).")
83+
if var_pos.size == 0:
84+
raise ValueError("No variables selected (var_names resolved to empty).")
85+
if tem_pos.size == 0:
86+
raise ValueError("No timepoints selected (tem_names resolved to empty).")
87+
88+
mtx = mtx[np.ix_(obs_pos, var_pos, tem_pos)]
89+
timepoints = np.asarray(tem_labels)
90+
91+
if overlay:
92+
if len(var_labels) != 1:
93+
raise ValueError("When overlay=True, only a single var_name can be plotted at a time.")
94+
95+
k = str(var_labels[0])
96+
y = np.asarray(mtx[:, 0, :], dtype=float)
97+
n_obs, n_time = y.shape
98+
99+
df = pd.DataFrame(
100+
{
101+
"time": np.tile(timepoints, n_obs),
102+
"value": y.ravel(order="C"),
103+
"series": np.repeat([str(x) for x in obs_labels], n_time),
104+
"variable": k,
105+
}
106+
)
107+
108+
curves = [
109+
hv.Curve(g, kdims="time", vdims="value", label=series) for series, g in df.groupby("series", sort=False)
110+
]
111+
plot = hv.Overlay(curves)
112+
113+
plot_title = title if title is not None else f"Time series for variable {k}"
114+
plot = plot.relabel(plot_title).opts(**opts_dict)
115+
116+
return plot
117+
118+
# overlay=False: one panel per observation; within each panel overlay variables
119+
panels = []
120+
for obs_i, obs_label in enumerate(obs_labels):
121+
curves = []
122+
for var_i, var_label in enumerate(var_labels):
123+
y = np.asarray(mtx[obs_i, var_i, :], dtype=float)
124+
g = pd.DataFrame({"time": timepoints, "value": y})
125+
curves.append(hv.Curve(g, kdims="time", vdims="value", label=str(var_label)))
126+
127+
panel = hv.Overlay(curves)
128+
129+
panel_title = (
130+
title if (title is not None and len(obs_labels) == 1) else f"Time series for observation {obs_label}"
131+
)
132+
133+
panel = panel.relabel(panel_title).opts(**opts_dict)
134+
panels.append(panel)
135+
136+
layout = hv.Layout(panels).cols(1)
137+
return layout
138+
139+
140+
def _resolve_axis(index: pd.Index, names: Any, axis: str) -> tuple[np.ndarray, pd.Index]:
141+
n = len(index)
142+
143+
if names is None:
144+
pos = np.arange(n, dtype=int)
145+
return pos, index.take(pos)
146+
147+
if isinstance(names, slice):
148+
pos = np.arange(n, dtype=int)[names]
149+
return pos, index.take(pos)
150+
151+
if isinstance(names, (str, int, np.integer)):
152+
names_list = [names]
153+
else:
154+
names_list = list(names)
155+
156+
names_list = list(dict.fromkeys(names_list))
157+
158+
pos = index.get_indexer(names_list)
159+
if (pos < 0).any():
160+
missing = [names_list[i] for i, p in enumerate(pos) if p < 0]
161+
raise KeyError(f"{', '.join(str(x) for x in missing)} not found in edata.{axis}")
162+
163+
pos = pos.astype(int, copy=False)
164+
return pos, index.take(pos)

tests/plot/test_timeseries.py

Lines changed: 115 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,115 @@
1+
from pathlib import Path
2+
3+
import ehrdata as ed
4+
import holoviews as hv
5+
6+
hv.extension("bokeh")
7+
import pytest
8+
from ehrdata.core.constants import DEFAULT_TEM_LAYER_NAME
9+
10+
import ehrapy as ep
11+
12+
CURRENT_DIR = Path(__file__).parent
13+
14+
15+
def test_timeseries(edata_blob_small):
16+
edata = edata_blob_small
17+
18+
plot = ep.pl.timeseries(edata, obs_names="1", layer=DEFAULT_TEM_LAYER_NAME)
19+
assert plot is not None
20+
assert isinstance(plot, hv.Layout)
21+
22+
23+
def test_timeseries_multiple_obs(edata_blob_small):
24+
edata = edata_blob_small
25+
26+
plot = ep.pl.timeseries(
27+
edata,
28+
obs_names=["3", "4"],
29+
var_names=["feature_1", "feature_2", "feature_3"],
30+
layer=DEFAULT_TEM_LAYER_NAME,
31+
)
32+
33+
assert plot is not None
34+
assert isinstance(plot, hv.Layout)
35+
36+
37+
def test_timeseries_overlay(edata_blob_small):
38+
edata = edata_blob_small
39+
40+
plot = ep.pl.timeseries(
41+
edata,
42+
obs_names=["3", "4", "5"],
43+
var_names="feature_1",
44+
layer=DEFAULT_TEM_LAYER_NAME,
45+
overlay=True,
46+
)
47+
assert plot is not None
48+
assert isinstance(plot, hv.Overlay)
49+
50+
51+
def test_timeseries_subset_time(edata_blob_small):
52+
edata = edata_blob_small
53+
54+
plot_1 = ep.pl.timeseries(
55+
edata,
56+
obs_names=["3", "4"],
57+
var_names=["feature_1", "feature_2", "feature_3"],
58+
tem_names=slice(0, 5),
59+
layer=DEFAULT_TEM_LAYER_NAME,
60+
)
61+
62+
assert plot_1 is not None
63+
assert isinstance(plot_1, hv.Layout)
64+
65+
66+
def test_timeseries_list(edata_blob_small):
67+
edata = edata_blob_small
68+
69+
plot = ep.pl.timeseries(
70+
edata,
71+
obs_names=["3", "4"],
72+
var_names=["feature_1", "feature_2", "feature_3"],
73+
tem_names=["0", "1", "2"],
74+
layer=DEFAULT_TEM_LAYER_NAME,
75+
)
76+
77+
assert plot is not None
78+
assert isinstance(plot, hv.Layout)
79+
80+
81+
def test_timeseries_error_cases(mar_edata, edata_blob_small):
82+
edata_2d_layer = mar_edata.X
83+
edata_2d = ed.EHRData(shape=(100, 10), layers={"X": edata_2d_layer})
84+
85+
with pytest.raises(ValueError, match="Layer 'X' must be 3D"):
86+
ep.pl.timeseries(
87+
edata_2d,
88+
obs_names="0",
89+
var_names="feature_1",
90+
layer="X",
91+
)
92+
93+
with pytest.raises(KeyError, match="Layer 'unknown_layer' not found in edata.layers"):
94+
ep.pl.timeseries(
95+
edata_blob_small,
96+
obs_names="0",
97+
var_names="feature_0",
98+
layer="unknown_layer",
99+
)
100+
101+
with pytest.raises(KeyError, match="unknown_feature not found in edata.var_names"):
102+
ep.pl.timeseries(
103+
edata_blob_small,
104+
obs_names="0",
105+
var_names="unknown_feature",
106+
layer=DEFAULT_TEM_LAYER_NAME,
107+
)
108+
with pytest.raises(ValueError, match="When overlay=True, only a single var_name can be plotted at a time"):
109+
ep.pl.timeseries(
110+
edata_blob_small,
111+
obs_names=["0", "1"],
112+
var_names=["feature_1", "feature_2"],
113+
layer=DEFAULT_TEM_LAYER_NAME,
114+
overlay=True,
115+
)

0 commit comments

Comments
 (0)