Skip to content

Commit ec78a2d

Browse files
committed
refactor(render): carry ColorSpec to the leaves, drop loose intermediate vectors (#700)
The renderers unpacked color_spec into loose color_source_vector/color_vector up front, then mutated and passed them around — so our own helpers took the pair instead of the spec. Keep color_spec as the single carrier instead: each post-resolution mutation is a transform (with_color_vector/with_source_vector), make_palette is a ColorSpec method, and _render_centroids_as_points + _add_legend_and_colorbar take the spec. Loose arrays now appear only at the genuine leaves that consume them (ax.scatter/imshow, PatchCollection, _map_color_seg, the datashader funcs): those unpack color_spec.color_vector inline, and the datashader leaves' returns are re-wrapped into the spec. Mechanical, value-for-value identical; set-diff shows zero new failures and the non-visual suite is green.
1 parent 7f47459 commit ec78a2d

2 files changed

Lines changed: 64 additions & 66 deletions

File tree

src/spatialdata_plot/pl/_color.py

Lines changed: 14 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -702,6 +702,20 @@ def apply_transfunc(self, transfunc: Any) -> ColorSpec:
702702
return replace(self, color_vector=transfunc(self.color_vector))
703703
return self
704704

705+
def with_color_vector(self, color_vector: ArrayLike) -> ColorSpec:
706+
"""Return a copy with a replaced ``color_vector`` (a post-resolution rewrite: reprocess, compute)."""
707+
return replace(self, color_vector=color_vector)
708+
709+
def with_source_vector(self, source_vector: ArrayLike | pd.Series | None) -> ColorSpec:
710+
"""Return a copy with a replaced ``source_vector`` (e.g. ``remove_unused_categories``, ``compute``)."""
711+
return replace(self, source_vector=source_vector)
712+
713+
def make_palette(self) -> ListedColormap:
714+
"""Build a ``ListedColormap`` from the colors, dropping NaN categories when categorical."""
715+
if self.source_vector is None:
716+
return ListedColormap(dict.fromkeys(self.color_vector))
717+
return ListedColormap(dict.fromkeys(self.color_vector[~pd.Categorical(self.source_vector).isnull()]))
718+
705719
def groups_keep_mask(self, groups: Any, na_color: Color) -> Any | None:
706720
"""Rows whose category is in ``groups``, or None when the groups filter does not apply.
707721

src/spatialdata_plot/pl/render.py

Lines changed: 50 additions & 66 deletions
Original file line numberDiff line numberDiff line change
@@ -34,7 +34,6 @@
3434
from spatialdata_plot._logging import _log_context, logger
3535
from 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-
338327
def _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

Comments
 (0)