77from ionique .core import AnySegment
88from ionique .datatypes import SessionFileManager
99import matplotlib .pyplot as plt
10+
1011def 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
8283pn .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+
177185def 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-
445516def qp_scatter (** args ):
446517 """
447518 quickly generate a scatterplot from specified parameters. wrapper for seaborn scatterplot
0 commit comments