3434from spatialdata_plot ._logging import _log_context , logger
3535from spatialdata_plot .pl ._color import (
3636 ColorSpec ,
37- ColorType ,
3837 _get_colors_for_categorical_obs ,
3938 _get_linear_colormap ,
4039 _map_color_seg ,
@@ -325,16 +324,6 @@ def _should_request_colorbar(
325324 return bool (auto_condition )
326325
327326
328- def _make_palette (
329- color_source_vector : pd .Series | None ,
330- color_vector : Any ,
331- ) -> ListedColormap :
332- """Build a ListedColormap from a color vector, filtering out NaN entries when categorical."""
333- if color_source_vector is None :
334- return ListedColormap (dict .fromkeys (color_vector ))
335- return ListedColormap (dict .fromkeys (color_vector [~ pd .Categorical (color_source_vector ).isnull ()]))
336-
337-
338327def _add_legend_and_colorbar (
339328 ax : matplotlib .axes .SubplotBase ,
340329 cax : ScalarMappable | None ,
@@ -371,7 +360,7 @@ def _add_legend_and_colorbar(
371360 return
372361
373362 if palette is None and fill_has_decorations :
374- palette = _make_palette ( color_source_vector , color_vector )
363+ palette = color_spec . make_palette ( )
375364
376365 if color_source_vector is not None and hasattr (color_source_vector , "remove_unused_categories" ):
377366 color_source_vector = color_source_vector .remove_unused_categories ()
@@ -633,7 +622,6 @@ def _render_shapes(
633622 table_layer = table_layer ,
634623 coordinate_system = coordinate_system ,
635624 )
636- colortype = color_spec .colortype
637625
638626 col_for_outline_color = render_params .col_for_outline_color
639627 outline_table_name = render_params .outline_table_name
@@ -692,49 +680,45 @@ def _render_shapes(
692680 outline_color_spec = outline_color_spec .filter (keep )
693681
694682 color_spec = color_spec .apply_transfunc (render_params .transfunc )
695- color_source_vector , color_vector = color_spec .source_vector , color_spec .color_vector
696- outline_color_source_vector , outline_color_vector = (
697- (outline_color_spec .source_vector , outline_color_spec .color_vector )
698- if outline_color_spec is not None
699- else (None , None )
700- )
701683
702684 norm = render_params .cmap_params .fresh_norm ()
703685
704- if len (color_vector ) == 0 :
705- color_vector = [render_params .cmap_params .na_color .get_hex_with_alpha ()]
686+ if len (color_spec . color_vector ) == 0 :
687+ color_spec = color_spec . with_color_vector ( [render_params .cmap_params .na_color .get_hex_with_alpha ()])
706688
707689 # continuous case: leave NaNs as NaNs; utils maps them to na_color during draw
708690 if color_spec .is_continuous :
709- _series = color_vector if isinstance (color_vector , pd .Series ) else pd .Series (color_vector )
691+ cv = color_spec .color_vector
692+ _series = cv if isinstance (cv , pd .Series ) else pd .Series (cv )
710693
711694 try :
712- color_vector = np .asarray (_series , dtype = float )
695+ cv = np .asarray (_series , dtype = float )
713696 except (TypeError , ValueError ):
714697 nan_count = int (_series .isna ().sum ())
715698 if nan_count :
716699 logger .warning (
717700 f"Found { nan_count } NaN values in color data. "
718701 "These observations will be colored with the 'na_color'."
719702 )
720- color_vector = _series .to_numpy ()
703+ cv = _series .to_numpy ()
721704 else :
722- if np .isnan (color_vector ).any ():
723- nan_count = int (np .isnan (color_vector ).sum ())
705+ if np .isnan (cv ).any ():
706+ nan_count = int (np .isnan (cv ).sum ())
724707 logger .warning (
725708 f"Found { nan_count } NaN values in color data. "
726709 "These observations will be colored with the 'na_color'."
727710 )
711+ color_spec = color_spec .with_color_vector (cv )
728712
729- palette = _make_palette ( color_source_vector , color_vector )
713+ palette = color_spec . make_palette ( )
730714
731715 has_valid_color = (
732- len (set (color_vector )) != 1
733- or list (set (color_vector ))[0 ] != render_params .cmap_params .na_color .get_hex_with_alpha ()
716+ len (set (color_spec . color_vector )) != 1
717+ or list (set (color_spec . color_vector ))[0 ] != render_params .cmap_params .na_color .get_hex_with_alpha ()
734718 )
735- if has_valid_color and color_source_vector is not None and col_for_color is not None :
719+ if has_valid_color and color_spec . is_categorical and col_for_color is not None :
736720 # necessary in case different shapes elements are annotated with one table
737- color_source_vector = color_source_vector . remove_unused_categories ()
721+ color_spec = color_spec . with_source_vector ( color_spec . source_vector . remove_unused_categories () )
738722
739723 shapes = gpd .GeoDataFrame (shapes , geometry = "geometry" )
740724
@@ -749,9 +733,7 @@ def _render_shapes(
749733 render_params ,
750734 x = xy [:, 0 ],
751735 y = xy [:, 1 ],
752- color_vector = color_vector ,
753- color_source_vector = color_source_vector ,
754- colortype = colortype ,
736+ color_spec = color_spec ,
755737 norm = norm ,
756738 na_color = render_params .cmap_params .na_color ,
757739 adata = table ,
@@ -834,6 +816,10 @@ def _render_shapes(
834816
835817 cvs = ds .Canvas (plot_width = plot_width , plot_height = plot_height , x_range = x_ext , y_range = y_ext )
836818
819+ # datashader consumes raw arrays; unpack the carrier for this leaf and re-wrap its result
820+ color_vector = color_spec .color_vector
821+ color_source_vector = color_spec .source_vector
822+
837823 # in case we are coloring by a column in table
838824 if col_for_color is not None and col_for_color not in transformed_element .columns :
839825 # Ensure color vector length matches the number of shapes
@@ -875,6 +861,7 @@ def _render_shapes(
875861 default_reduction = _default_reduction ,
876862 kind = "shapes" ,
877863 )
864+ color_spec = color_spec .with_color_vector (color_vector )
878865
879866 _render_ds_outlines (
880867 cvs ,
@@ -885,8 +872,8 @@ def _render_shapes(
885872 factor ,
886873 x_min = x_ext [0 ],
887874 y_min = y_ext [0 ],
888- outline_color_vector = outline_color_vector ,
889- outline_color_source_vector = outline_color_source_vector ,
875+ outline_color_vector = outline_color_spec . color_vector if outline_color_spec is not None else None ,
876+ outline_color_source_vector = outline_color_spec . source_vector if outline_color_spec is not None else None ,
890877 )
891878
892879 _cax = _render_ds_image (
@@ -988,7 +975,7 @@ def _render_shapes(
988975
989976 if color_spec .is_continuous :
990977 # Colorbar range from the same resolved norm the fill pixels use.
991- used_norm = _resolve_continuous_norm (color_vector , render_params .cmap_params )
978+ used_norm = _resolve_continuous_norm (color_spec . color_vector , render_params .cmap_params )
992979 _cax .set_clim (vmin = used_norm .vmin , vmax = used_norm .vmax )
993980
994981 _add_legend_and_colorbar (
@@ -997,7 +984,7 @@ def _render_shapes(
997984 fig_params = fig_params ,
998985 adata = table ,
999986 col_for_color = col_for_color ,
1000- color_spec = ColorSpec ( colortype , color_source_vector , color_vector ) ,
987+ color_spec = color_spec ,
1001988 palette = palette ,
1002989 alpha = render_params .fill_alpha ,
1003990 na_color = render_params .cmap_params .na_color ,
@@ -1077,9 +1064,7 @@ def _render_centroids_as_points(
10771064 * ,
10781065 x : Any ,
10791066 y : Any ,
1080- color_vector : Any ,
1081- color_source_vector : pd .Series | None ,
1082- colortype : ColorType ,
1067+ color_spec : ColorSpec ,
10831068 norm : Normalize | None ,
10841069 na_color : Any ,
10851070 adata : AnnData | None ,
@@ -1100,12 +1085,12 @@ def _render_centroids_as_points(
11001085 method = _resolve_as_points_method (render_params , n = len (x ), allow_datashader = allow_datashader )
11011086 if method == "datashader" :
11021087 df = pd .DataFrame ({"x" : x , "y" : y })
1103- cax , color_vector , color_source_vector = _datashader_points (
1088+ cax , cv , csv = _datashader_points (
11041089 ax ,
11051090 df ,
11061091 col_for_color = col_for_color ,
1107- color_vector = color_vector ,
1108- color_source_vector = color_source_vector ,
1092+ color_vector = color_spec . color_vector ,
1093+ color_source_vector = color_spec . source_vector ,
11091094 norm = norm ,
11101095 cmap_params = render_params .cmap_params ,
11111096 alpha = render_params .fill_alpha ,
@@ -1121,12 +1106,13 @@ def _render_centroids_as_points(
11211106 as_markers = True ,
11221107 axes_extent = axes_extent ,
11231108 )
1109+ color_spec = color_spec .with_source_vector (csv ).with_color_vector (cv )
11241110 else :
11251111 cax = _scatter_points (
11261112 ax ,
11271113 x ,
11281114 y ,
1129- color_vector ,
1115+ color_spec . color_vector ,
11301116 size = render_params .size ,
11311117 cmap = render_params .cmap_params .cmap ,
11321118 norm = norm ,
@@ -1140,7 +1126,7 @@ def _render_centroids_as_points(
11401126 fig_params = fig_params ,
11411127 adata = adata ,
11421128 col_for_color = col_for_color ,
1143- color_spec = ColorSpec ( colortype , color_source_vector , color_vector ) ,
1129+ color_spec = color_spec ,
11441130 palette = palette ,
11451131 alpha = render_params .fill_alpha ,
11461132 na_color = na_color ,
@@ -1405,7 +1391,6 @@ def _render_points(
14051391 coordinate_system = coordinate_system ,
14061392 preloaded_color_data = _preloaded ,
14071393 )
1408- colortype = color_spec .colortype
14091394
14101395 if added_color_from_table and col_for_color is not None :
14111396 _reparse_points (
@@ -1431,7 +1416,6 @@ def _render_points(
14311416 _reparse_points (sdata_filt , element , points , transformation_in_cs , coordinate_system , col_for_color )
14321417
14331418 color_spec = color_spec .apply_transfunc (render_params .transfunc )
1434- color_source_vector , color_vector = color_spec .source_vector , color_spec .color_vector
14351419
14361420 trans , trans_data = _prepare_transformation (sdata .points [element ], coordinate_system , ax )
14371421
@@ -1464,12 +1448,12 @@ def _render_points(
14641448 # any other elements on the axes.
14651449 return
14661450
1467- cax , color_vector , color_source_vector = _datashader_points (
1451+ cax , cv , csv = _datashader_points (
14681452 ax ,
14691453 transformed_element ,
14701454 col_for_color = col_for_color ,
1471- color_vector = color_vector ,
1472- color_source_vector = color_source_vector ,
1455+ color_vector = color_spec . color_vector ,
1456+ color_source_vector = color_spec . source_vector ,
14731457 norm = norm ,
14741458 cmap_params = render_params .cmap_params ,
14751459 alpha = render_params .alpha ,
@@ -1481,6 +1465,7 @@ def _render_points(
14811465 fig_params = fig_params ,
14821466 default_reduction = _default_reduction ,
14831467 )
1468+ color_spec = color_spec .with_source_vector (csv ).with_color_vector (cv )
14841469
14851470 elif method == "matplotlib" :
14861471 # update axis limits if plot was empty before (necessary if datashader comes after)
@@ -1489,7 +1474,7 @@ def _render_points(
14891474 ax ,
14901475 adata [:, 0 ].X .flatten (),
14911476 adata [:, 1 ].X .flatten (),
1492- color_vector ,
1477+ color_spec . color_vector ,
14931478 size = render_params .size ,
14941479 cmap = render_params .cmap_params .cmap ,
14951480 norm = norm ,
@@ -1509,7 +1494,7 @@ def _render_points(
15091494 fig_params = fig_params ,
15101495 adata = adata ,
15111496 col_for_color = col_for_color ,
1512- color_spec = ColorSpec ( colortype , color_source_vector , color_vector ) ,
1497+ color_spec = color_spec ,
15131498 palette = None ,
15141499 alpha = render_params .alpha ,
15151500 na_color = render_params .cmap_params .na_color ,
@@ -2234,7 +2219,6 @@ def _render_labels(
22342219 render_type = "labels" ,
22352220 coordinate_system = coordinate_system ,
22362221 )
2237- colortype = color_spec .colortype
22382222
22392223 # Outline color lookup must run BEFORE any masking so the returned vector aligns to
22402224 # the original instance_id. The same masks applied to fill below are then applied
@@ -2288,7 +2272,7 @@ def _render_labels(
22882272 outline_color_spec = outline_color_spec .filter (keep_vec )
22892273
22902274 color_spec = color_spec .apply_transfunc (render_params .transfunc )
2291- color_source_vector , color_vector = color_spec . source_vector , color_spec . color_vector
2275+ # outline still feeds the _map_color_seg leaf as loose arrays
22922276 outline_color_source_vector , outline_color_vector = (
22932277 (outline_color_spec .source_vector , outline_color_spec .color_vector )
22942278 if outline_color_spec is not None
@@ -2320,10 +2304,10 @@ def _render_labels(
23202304 point_color_vector = np .random .default_rng (42 ).random ((len (point_ids ), 3 ))
23212305 point_color_source_vector = None
23222306 allow_datashader = False
2323- elif len (color_vector ) == len (instance_id ):
2307+ elif len (color_spec . color_vector ) == len (instance_id ):
23242308 # data-driven colour is per-instance
2325- point_color_vector = np .asarray (color_vector )[keep ]
2326- point_color_source_vector = None if color_source_vector is None else color_source_vector [keep ]
2309+ point_color_vector = np .asarray (color_spec . color_vector )[keep ]
2310+ point_color_source_vector = None if color_spec . source_vector is None else color_spec . source_vector [keep ]
23272311 else :
23282312 # literal colour / user-set na_color -> one colour per centroid
23292313 point_color_vector = np .full (len (point_ids ), na_color .get_hex_with_alpha ())
@@ -2335,9 +2319,7 @@ def _render_labels(
23352319 render_params ,
23362320 x = xy [:, 0 ],
23372321 y = xy [:, 1 ],
2338- color_vector = point_color_vector ,
2339- color_source_vector = point_color_source_vector ,
2340- colortype = colortype ,
2322+ color_spec = ColorSpec (color_spec .colortype , point_color_source_vector , point_color_vector ),
23412323 norm = render_params .cmap_params .fresh_norm (), # ax.scatter autoscales in place; don't mutate the shared norm
23422324 na_color = na_color ,
23432325 adata = table if table_name is not None else None ,
@@ -2360,8 +2342,8 @@ def _draw_labels(
23602342 labels = _map_color_seg (
23612343 seg = label .values ,
23622344 cell_id = instance_id ,
2363- color_vector = color_vector ,
2364- color_source_vector = color_source_vector ,
2345+ color_vector = color_spec . color_vector ,
2346+ color_source_vector = color_spec . source_vector ,
23652347 cmap_params = render_params .cmap_params ,
23662348 seg_erosionpx = seg_erosionpx ,
23672349 seg_boundaries = seg_boundaries ,
@@ -2378,7 +2360,7 @@ def _draw_labels(
23782360 cmap = None if color_spec .is_categorical else render_params .cmap_params .cmap ,
23792361 norm = None
23802362 if color_spec .is_categorical
2381- else _resolve_continuous_norm (color_vector , render_params .cmap_params ),
2363+ else _resolve_continuous_norm (color_spec . color_vector , render_params .cmap_params ),
23822364 alpha = alpha ,
23832365 origin = "lower" ,
23842366 zorder = render_params .zorder ,
@@ -2446,7 +2428,7 @@ def _draw_labels(
24462428 fig_params = fig_params ,
24472429 adata = table ,
24482430 col_for_color = col_for_color ,
2449- color_spec = ColorSpec ( colortype , color_source_vector , color_vector ) ,
2431+ color_spec = color_spec ,
24502432 palette = palette ,
24512433 alpha = alpha_to_decorate_ax ,
24522434 na_color = render_params .cmap_params .na_color ,
@@ -2457,7 +2439,9 @@ def _draw_labels(
24572439 outline_col_for_color = col_for_outline_color ,
24582440 outline_color_spec = outline_color_spec ,
24592441 outline_cmap_params = render_params .cmap_params ,
2460- na_in_legend_override = (legend_params .na_in_legend if groups is None else len (groups ) == len (set (color_vector ))),
2442+ na_in_legend_override = (
2443+ legend_params .na_in_legend if groups is None else len (groups ) == len (set (color_spec .color_vector ))
2444+ ),
24612445 )
24622446
24632447
0 commit comments