Skip to content

Commit 110470d

Browse files
committed
add voltage filter to plotting
1 parent bfa1483 commit 110470d

1 file changed

Lines changed: 98 additions & 27 deletions

File tree

src/ionique/plotting.py

Lines changed: 98 additions & 27 deletions
Original file line numberDiff line numberDiff line change
@@ -7,6 +7,7 @@
77
from ionique.core import AnySegment
88
from ionique.datatypes import SessionFileManager
99
import matplotlib.pyplot as plt
10+
1011
def qp_trace(seg:AnySegment|None = None, ranks=["vstepgap","event"],downsamples={"vstepgap":50,"event":1},fig_size=(6,5),ranks_kwargs={},fig_kwargs={},plot_voltage:Literal["same","split",None]=None):
1112
"""
1213
quickly plot a trace, or segment.
@@ -77,7 +78,7 @@ def qp_trace(seg:AnySegment|None = None, ranks=["vstepgap","event"],downsamples=
7778
LegendItem,
7879
Range1d,
7980
)
80-
from bokeh.palettes import Category10
81+
from bokeh.palettes import Category10, Category20
8182

8283
pn.extension()
8384

@@ -174,6 +175,13 @@ def compute_sampling_frequency(row: pd.Series) -> Optional[float]:
174175
# Main builder
175176
# -----------------------------
176177

178+
from bokeh.palettes import Category10, Category20
179+
from bokeh.plotting import figure
180+
from bokeh.models import (
181+
HoverTool, TapTool, Range1d, ColumnDataSource,
182+
CDSView, BooleanFilter
183+
)
184+
177185
def dashboard_event_inspection(df: pd.DataFrame):
178186
"""
179187
Interactive Jupyter dashboard (Panel + Bokeh) for exploring ionic current events in `df`.
@@ -194,17 +202,35 @@ def dashboard_event_inspection(df: pd.DataFrame):
194202
6) If 'Show sub-segments' is ON, plots cyclically colored segments from 'subevent_start'/'subevent_end';
195203
otherwise plots wrap (black) and current (blue).
196204
"""
205+
def _pick_palette(n):
206+
# Use Category10/20; repeat if needed
207+
if n <= 10:
208+
return Category10[10][:n]
209+
elif n <= 20 and hasattr(Category10, "20"): # some Bokeh installs have Category20
210+
return Category20[20][:n] # fallback if available
211+
else:
212+
base = Category10[10]
213+
reps = (n // 10) + 1
214+
return (base * reps)[:n]
215+
197216
if not isinstance(df, pd.DataFrame) or df.empty:
198217
return pn.pane.Markdown("❌ DataFrame is empty or invalid.")
199218

200219
# Column classifications
201220
numeric_cols = _numeric_non_array_columns(df)
202221
group_cols = ["(None)"] + _categorical_candidates(df)
222+
if "Voltage" in df.columns and "Voltage" not in group_cols:
223+
group_cols.append("Voltage")
203224

204225
# Widgets
205226
x_select = pn.widgets.Select(name="X axis", options=numeric_cols, value=numeric_cols[0] if numeric_cols else None)
206227
y_select = pn.widgets.Select(name="Y axis", options=numeric_cols, value=numeric_cols[1] if len(numeric_cols) > 1 else (numeric_cols[0] if numeric_cols else None))
207228
group_select = pn.widgets.Select(name="Color by (group)", options=group_cols, value="(None)")
229+
if "Voltage" in df.columns:
230+
_voltages = [str(v) for v in sorted(df["Voltage"].dropna().unique().tolist(), key=float)]
231+
else:
232+
_voltages = []
233+
voltage_filter = pn.widgets.Select(name="Show voltage", options=["(All)"] + _voltages, value="(All)")
208234

209235
x_log = pn.widgets.Checkbox(name="Log X", value=False)
210236
y_log = pn.widgets.Checkbox(name="Log Y", value=False)
@@ -217,43 +243,97 @@ def dashboard_event_inspection(df: pd.DataFrame):
217243
status = pn.pane.Markdown("")
218244

219245
# Prepare scatter data source
220-
scatter_source = ColumnDataSource(data=dict(index=np.arange(len(df)), x=np.zeros(len(df)), y=np.zeros(len(df)), color=["#1f77b4"]*len(df), size=[6]*len(df), alpha=[0.8]*len(df)))
221-
246+
# scatter_source = ColumnDataSource(data=dict(index=np.arange(len(df)), x=np.zeros(len(df)), y=np.zeros(len(df)), color=["#1f77b4"]*len(df), size=[6]*len(df), alpha=[0.8]*len(df)))
247+
scatter_source = ColumnDataSource(data=dict(
248+
index=np.arange(len(df)),
249+
x=np.zeros(len(df)),
250+
y=np.zeros(len(df)),
251+
color=["#1f77b4"] * len(df),
252+
size=[6] * len(df),
253+
alpha=[0.8] * len(df),
254+
legend=[""] * len(df),
255+
group_val=[""] * len(df),
256+
voltage_str=(df["Voltage"].astype(str).fillna("NA").values
257+
if "Voltage" in df.columns else ["NA"] * len(df)),
258+
))
222259
def _build_scatter_source():
223260
if x_select.value is None or y_select.value is None:
224261
return
225262
x = df[x_select.value].values
226263
y = df[y_select.value].values
227-
# Group colors
264+
265+
legend = [""] * len(df)
266+
colors = ["#1f77b4"] * len(df)
267+
228268
if group_select.value and group_select.value != "(None)":
229-
groups = df[group_select.value].astype(str).fillna("NA").values
230-
uniq = pd.unique(groups)
231-
palette = (Category10[10] if len(uniq) <= 10 else (Category10[10] * ((len(uniq)//10)+1)))
232-
color_map = {g: palette[i] for i, g in enumerate(uniq)}
269+
groups_raw = df[group_select.value]
270+
groups = groups_raw.astype(str).fillna("NA").values
271+
uniq = list(pd.unique(groups))
272+
palette = _pick_palette(len(uniq))
273+
color_map = {g: palette[i % len(palette)] for i, g in enumerate(uniq)}
233274
colors = [color_map[g] for g in groups]
275+
legend = groups
276+
277+
scatter_source.data.update(
278+
index=np.arange(len(df)),
279+
x=x,
280+
y=y,
281+
color=colors,
282+
size=[7] * len(df),
283+
alpha=[0.8] * len(df),
284+
legend=legend,
285+
group_val=legend,
286+
voltage_str=(df["Voltage"].astype(str).fillna("NA").values
287+
if "Voltage" in df.columns else scatter_source.data["voltage_str"])
288+
)
289+
290+
bool_filter = BooleanFilter(booleans=[True] * len(df))
291+
view = CDSView(filter=bool_filter)
292+
293+
294+
def _update_voltage_view():
295+
if not _voltages or voltage_filter.value == "(All)":
296+
mask = [True] * len(df)
234297
else:
235-
colors = ["#1f77b4"] * len(df)
236-
scatter_source.data.update(index=np.arange(len(df)), x=x, y=y, color=colors, size=[7]*len(df), alpha=[0.8]*len(df))
298+
allowed = {voltage_filter.value}
299+
vstr = scatter_source.data.get("voltage_str", ["NA"] * len(df))
300+
mask = [vs in allowed for vs in vstr]
301+
bool_filter.booleans = mask
237302

238-
_build_scatter_source()
239303

240-
# Scatter figure factory to support log toggles
304+
_update_voltage_view()
305+
306+
_build_scatter_source()
241307
scatter_pane = pn.pane.Bokeh(sizing_mode="stretch_both")
242308

309+
243310
def _make_scatter_figure():
244311
x_axis_type = "log" if x_log.value else "linear"
245312
y_axis_type = "log" if y_log.value else "linear"
246-
p = figure(height=350, sizing_mode="stretch_width", tools="pan,wheel_zoom,box_zoom,reset,tap,hover,save", x_axis_type=x_axis_type, y_axis_type=y_axis_type)
247-
r = p.circle(source=scatter_source, x="x", y="y", size="size", color="color", alpha="alpha", line_color=None)
313+
p = figure(height=350, sizing_mode="stretch_width",
314+
tools="pan,wheel_zoom,box_zoom,reset,tap,hover,save",
315+
x_axis_type=x_axis_type, y_axis_type=y_axis_type)
316+
r = p.circle(source=scatter_source, x="x", y="y",
317+
size="size", color="color", alpha="alpha", line_color=None,
318+
legend_field="legend") # NEW
319+
248320
p.add_tools(TapTool())
249321
hover = p.select_one(HoverTool)
250322
hover.tooltips = [
251323
("row", "@index"),
252324
(x_select.name, "@x"),
253325
(y_select.name, "@y"),
326+
("group", "@group_val")
254327
]
255328
p.title.text = "Event Scatter"
329+
330+
# Only show legend if grouping is active
331+
p.legend.visible = (group_select.value and group_select.value != "(None)")
332+
p.legend.location = "top_right"
333+
p.legend.click_policy = "hide"
334+
256335
scatter_pane.object = p
336+
257337
return p
258338

259339
p_scatter = _make_scatter_figure()
@@ -262,8 +342,8 @@ def _make_scatter_figure():
262342
event_fig = figure(height=300, sizing_mode="stretch_width", tools="pan,wheel_zoom,box_zoom,reset,save",output_backend='webgl')
263343
event_fig.title.text = "Blockade Event View"
264344
# glyph refs to update
265-
wrap_renderer = event_fig.line([], [], line_color="#000000", line_width=2, alpha=0.9)
266-
current_renderer = event_fig.line([], [], line_color="#1f77b4", line_width=2, alpha=0.9)
345+
wrap_renderer = event_fig.line([], [], line_color="#000000", line_width=1, alpha=0.9)
346+
current_renderer = event_fig.line([], [], line_color="#1f77b4", line_width=1, alpha=0.9)
267347
# Segments datasource (fast): one MultiLine glyph for all segments
268348
segments_source = ColumnDataSource(data=dict(xs=[], ys=[], color=[]))
269349
segments_renderer = event_fig.multi_line(xs='xs', ys='ys', line_color='color', line_width=3, alpha=0.95, source=segments_source)
@@ -395,15 +475,11 @@ def _on_spinner_change(event):
395475

396476
select_idx.param.watch(_on_spinner_change, "value")
397477

398-
def _on_clear_click(event):
399-
scatter_source.selected.indices = []
478+
clear_btn.on_click(lambda _ : setattr(scatter_source.selected, "indices", []))
400479

401-
clear_btn.on_click(_on_clear_click)
402-
403-
# React to control changes
404480
def _on_axis_change(event=None):
405481
_build_scatter_source()
406-
_make_scatter_figure()
482+
p = _make_scatter_figure()
407483
_highlight_selection()
408484

409485
for w in (x_select, y_select, group_select, x_log, y_log):
@@ -437,11 +513,6 @@ def _on_flags_change(event=None):
437513
return layout
438514

439515

440-
441-
442-
443-
444-
445516
def qp_scatter(**args):
446517
"""
447518
quickly generate a scatterplot from specified parameters. wrapper for seaborn scatterplot

0 commit comments

Comments
 (0)