Skip to content

Commit 5d06317

Browse files
authored
Merge pull request #20 from singjc/add/3d_peakmap
Add/3d peakmap
2 parents 3343937 + 798fa3e commit 5d06317

7 files changed

Lines changed: 88200 additions & 251 deletions

File tree

‎nbs/PeakMap.ipynb‎

Lines changed: 580 additions & 22 deletions
Large diffs are not rendered by default.

‎nbs/pyopenms_viz_tutorial.ipynb‎

Lines changed: 87280 additions & 123 deletions
Large diffs are not rendered by default.

‎pyopenms_viz/_bokeh/core.py‎

Lines changed: 38 additions & 32 deletions
Original file line numberDiff line numberDiff line change
@@ -210,7 +210,7 @@ class BOKEHLinePlot(BOKEHPlot, LinePlot):
210210

211211
@classmethod
212212
@APPEND_PLOT_DOC
213-
def plot(cls, fig, data, x, y, by: str | None = None, **kwargs):
213+
def plot(cls, fig, data, x, y, by: str | None = None, plot_3d=False, **kwargs):
214214
"""
215215
Plot a line plot
216216
"""
@@ -219,7 +219,7 @@ def plot(cls, fig, data, x, y, by: str | None = None, **kwargs):
219219
if by is None:
220220
source = ColumnDataSource(data)
221221
if color_gen is not None:
222-
kwargs["line_color"] = next(color_gen)
222+
kwargs["line_color"] = color_gen if isinstance(color_gen, str) else next(color_gen)
223223
line = fig.line(x=x, y=y, source=source, **kwargs)
224224

225225
return fig, None
@@ -229,7 +229,7 @@ def plot(cls, fig, data, x, y, by: str | None = None, **kwargs):
229229
for group, df in data.groupby(by):
230230
source = ColumnDataSource(df)
231231
if color_gen is not None:
232-
kwargs["line_color"] = next(color_gen)
232+
kwargs["line_color"] = color_gen if isinstance(color_gen, str) else next(color_gen)
233233
line = fig.line(x=x, y=y, source=source, **kwargs)
234234
legend_items.append((group, [line]))
235235

@@ -245,30 +245,17 @@ class BOKEHVLinePlot(BOKEHPlot, VLinePlot):
245245

246246
@classmethod
247247
@APPEND_PLOT_DOC
248-
def plot(cls, fig, data, x, y, by: str | None = None, **kwargs):
248+
def plot(cls, fig, data, x, y, by: str | None = None, plot_3d=False, **kwargs):
249249
"""
250250
Plot a set of vertical lines
251251
"""
252252
color_gen = kwargs.pop("line_color", None)
253253
if color_gen is None:
254254
color_gen = ColorGenerator()
255255
data["line_color"] = [next(color_gen) for _ in range(len(data))]
256-
if by is None:
257-
source = ColumnDataSource(data)
258-
line = fig.segment(
259-
x0=x,
260-
y0=0,
261-
x1=x,
262-
y1=y,
263-
source=source,
264-
line_color="line_color",
265-
**kwargs,
266-
)
267-
return fig, None
268-
else:
269-
legend_items = []
270-
for group, df in data.groupby(by):
271-
source = ColumnDataSource(df)
256+
if not plot_3d:
257+
if by is None:
258+
source = ColumnDataSource(data)
272259
line = fig.segment(
273260
x0=x,
274261
y0=0,
@@ -278,11 +265,27 @@ def plot(cls, fig, data, x, y, by: str | None = None, **kwargs):
278265
line_color="line_color",
279266
**kwargs,
280267
)
281-
legend_items.append((group, [line]))
282-
283-
legend = Legend(items=legend_items)
284-
285-
return fig, legend
268+
return fig, None
269+
else:
270+
legend_items = []
271+
for group, df in data.groupby(by):
272+
source = ColumnDataSource(df)
273+
line = fig.segment(
274+
x0=x,
275+
y0=0,
276+
x1=x,
277+
y1=y,
278+
source=source,
279+
line_color="line_color",
280+
**kwargs,
281+
)
282+
legend_items.append((group, [line]))
283+
284+
legend = Legend(items=legend_items)
285+
286+
return fig, legend
287+
else:
288+
raise NotImplementedError("3D Vline plots are not supported in Bokeh")
286289

287290
def _add_annotations(
288291
self,
@@ -312,7 +315,7 @@ class BOKEHScatterPlot(BOKEHPlot, ScatterPlot):
312315

313316
@classmethod
314317
@APPEND_PLOT_DOC
315-
def plot(cls, fig, data, x, y, by: str | None = None, **kwargs):
318+
def plot(cls, fig, data, x, y, by: str | None = None, plot_3d=False, **kwargs):
316319
"""
317320
Plot a scatter plot
318321
"""
@@ -466,16 +469,19 @@ class BOKEHPeakMapPlot(BOKEH_MSPlot, PeakMapPlot):
466469
"""
467470

468471
def create_main_plot(self, x, y, z, class_kwargs, other_kwargs):
469-
scatterPlot = self.get_scatter_renderer(self.data, x, y, **class_kwargs)
472+
if not self.plot_3d:
473+
scatterPlot = self.get_scatter_renderer(self.data, x, y, **class_kwargs)
470474

471-
self.fig = scatterPlot.generate(z=z, **other_kwargs)
475+
self.fig = scatterPlot.generate(z=z, **other_kwargs)
472476

473-
if self.annotation_data is not None:
474-
self._add_box_boundaries(self.annotation_data)
477+
if self.annotation_data is not None:
478+
self._add_box_boundaries(self.annotation_data)
475479

476-
tooltips, _ = self._create_tooltips({self.xlabel: x, self.ylabel: y, "intensity": z})
480+
tooltips, _ = self._create_tooltips({self.xlabel: x, self.ylabel: y, "intensity": z})
477481

478-
self._add_tooltips(self.fig, tooltips)
482+
self._add_tooltips(self.fig, tooltips)
483+
else:
484+
raise NotImplementedError("3D PeakMap plots are not supported in Bokeh")
479485

480486
def create_x_axis_plot(self, x, z, class_kwargs):
481487
x_fig = super().create_x_axis_plot(x, z, class_kwargs)

‎pyopenms_viz/_config.py‎

Lines changed: 4 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -178,6 +178,7 @@ def default_legend_factory():
178178
title: str = "1D Plot"
179179
xlabel: str = "X-axis"
180180
ylabel: str = "Y-axis"
181+
zlabel: str = "Z-axis"
181182
x_axis_location: str = "below"
182183
y_axis_location: str = "left"
183184
min_border: str = 0
@@ -231,6 +232,7 @@ def set_plot_labels(self):
231232
"title": "PeakMap",
232233
"xlabel": "Retention Time",
233234
"ylabel": "mass-to-charge",
235+
"zlabel": "Intensity",
234236
},
235237
# Add more plot types as needed
236238
}
@@ -239,6 +241,8 @@ def set_plot_labels(self):
239241
self.title = plot_configs[self.kind]["title"]
240242
self.xlabel = plot_configs[self.kind]["xlabel"]
241243
self.ylabel = plot_configs[self.kind]["ylabel"]
244+
if self.kind == "peakmap":
245+
self.zlabel = plot_configs[self.kind]["zlabel"]
242246

243247
if self.relative_intensity and "Intensity" in self.ylabel:
244248
self.ylabel = "Relative " + self.ylabel

‎pyopenms_viz/_core.py‎

Lines changed: 44 additions & 10 deletions
Original file line numberDiff line numberDiff line change
@@ -108,6 +108,7 @@ def __init__(
108108
z: str | None = None,
109109
kind=None,
110110
by: str | None = None,
111+
plot_3d: bool = False,
111112
relative_intensity: bool = False,
112113
subplots: bool | None = None,
113114
sharex: bool | None = None,
@@ -120,6 +121,7 @@ def __init__(
120121
title: str | None = None,
121122
xlabel: str | None = None,
122123
ylabel: str | None = None,
124+
zlabel: str | None = None,
123125
x_axis_location: str | None = None,
124126
y_axis_location: str | None = None,
125127
line_type: str | None = None,
@@ -136,6 +138,7 @@ def __init__(
136138
self.data = data.copy()
137139
self.kind = kind
138140
self.by = by
141+
self.plot_3d = plot_3d
139142
self.relative_intensity = relative_intensity
140143

141144
# Plotting attributes
@@ -150,6 +153,7 @@ def __init__(
150153
self.title = title
151154
self.xlabel = xlabel
152155
self.ylabel = ylabel
156+
self.zlabel = zlabel
153157
self.x_axis_location = x_axis_location
154158
self.y_axis_location = y_axis_location
155159
self.line_type = line_type
@@ -298,7 +302,7 @@ def _make_plot(self, fig, **kwargs) -> None:
298302
tooltips = kwargs.pop("tooltips", None)
299303
custom_hover_data = kwargs.pop("custom_hover_data", None)
300304

301-
newlines, legend = self.plot(fig, self.data, self.x, self.y, self.by, **kwargs)
305+
newlines, legend = self.plot(fig, self.data, self.x, self.y, self.by, self.plot_3d, **kwargs)
302306

303307
if legend is not None:
304308
self._add_legend(newlines, legend)
@@ -308,7 +312,7 @@ def _make_plot(self, fig, **kwargs) -> None:
308312
self._add_tooltips(newlines, tooltips, custom_hover_data)
309313

310314
@abstractmethod
311-
def plot(cls, fig, data, x, y, by: str | None = None, **kwargs):
315+
def plot(cls, fig, data, x, y, by: str | None = None, plot_3d: bool = False, **kwargs):
312316
"""
313317
Create the plot
314318
"""
@@ -491,7 +495,11 @@ def plot(self, data, x, y, **kwargs):
491495
"""
492496
Create the plot
493497
"""
494-
color_gen = ColorGenerator()
498+
if 'line_color' not in kwargs:
499+
color_gen = ColorGenerator()
500+
else:
501+
color_gen = kwargs['line_color']
502+
495503
tooltip_entries = {"retention time": x, "intensity": y}
496504
if "Annotation" in self.data.columns:
497505
tooltip_entries["annotation"] = "Annotation"
@@ -795,6 +803,7 @@ def __init__(
795803
num_x_bins: int = 50,
796804
num_y_bins: int = 50,
797805
z_log_scale: bool = False,
806+
# plot_3d: bool = False,
798807
**kwargs,
799808
) -> None:
800809
# Copy data since it will be modified
@@ -826,13 +835,23 @@ def __init__(
826835
):
827836
data[x] = cut(data[x], bins=num_x_bins)
828837
data[y] = cut(data[y], bins=num_y_bins)
829-
830-
# Group by x and y bins and calculate the mean intensity within each bin
831-
data = (
832-
data.groupby([x, y], observed=True)
833-
.agg({z: "mean"})
834-
.reset_index()
835-
)
838+
by = kwargs.pop("by", None)
839+
if by is not None:
840+
# Group by x, y and by columns and calculate the mean intensity within each bin
841+
data = (
842+
data.groupby([x, y, by], observed=True)
843+
.agg({z: "mean"})
844+
.reset_index()
845+
)
846+
# Add by back to kwargs
847+
kwargs["by"] = by
848+
else:
849+
# Group by x and y bins and calculate the mean intensity within each bin
850+
data = (
851+
data.groupby([x, y], observed=True)
852+
.agg({z: "mean"})
853+
.reset_index()
854+
)
836855
data[x] = data[x].apply(lambda interval: interval.mid).astype(float)
837856
data[y] = data[y].apply(lambda interval: interval.mid).astype(float)
838857
data = data.fillna(0)
@@ -846,7 +865,11 @@ def __init__(
846865

847866
super().__init__(data, x, y, z=z, **kwargs)
848867

868+
# if not plot_3d:
849869
self.plot(x, y, z, **kwargs)
870+
# else:
871+
# self.plot_3d(x, y, z, **kwargs)
872+
850873
if self.show_plot:
851874
self.show()
852875

@@ -873,6 +896,12 @@ def plot(self, x, y, z, **kwargs):
873896
y_fig = self.create_y_axis_plot(y, z, class_kwargs_copy)
874897

875898
self.combine_plots(x_fig, y_fig)
899+
900+
# def plot_3d(self, x, y, z, **kwargs):
901+
# class_kwargs, other_kwargs = self._separate_class_kwargs(**kwargs)
902+
903+
# self.create_main_plot_3d(x, y, z, class_kwargs, other_kwargs)
904+
# pass
876905

877906
@staticmethod
878907
def _integrate_data_along_dim(
@@ -896,6 +925,10 @@ def create_main_plot(self, x, y, z, class_kwargs, other_kwargs):
896925
# by default the main plot with marginals is plotted the same way as the main plot unless otherwise specified
897926
def create_main_plot_marginals(self, x, y, z, class_kwargs, other_kwargs):
898927
self.create_main_plot(x, y, z, class_kwargs, other_kwargs)
928+
929+
# @abstractmethod
930+
# def create_main_plot_3d(self, x, y, z, class_kwargs, other_kwargs):
931+
# pass
899932

900933
@abstractmethod
901934
def create_x_axis_plot(self, x, z, class_kwargs) -> "figure":
@@ -1006,6 +1039,7 @@ def __call__(self, *args: Any, **kwargs: Any) -> Any:
10061039
# Call the plot method of the selected backend
10071040
if "backend" in kwargs:
10081041
kwargs.pop("backend")
1042+
10091043
return plot_backend.plot(self._parent, x=x, y=y, kind=kind, **kwargs)
10101044

10111045
@staticmethod

0 commit comments

Comments
 (0)