Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
14 changes: 12 additions & 2 deletions src/cmap/_colormap.py
Original file line number Diff line number Diff line change
Expand Up @@ -903,8 +903,8 @@ class ColorStops(Sequence[ColorStop]):
The array must be an (N, 5) array, where the first column is the position
(0-1) and the remaining columns are the color (RGBA, 0-1).
lut_func : callable, optional
A callable that takes a single argument (an (N, 1) array of positions) and
returns an (N, 4) array of colors. This will be used to generate the LUT
A callable that takes a single argument (an (N,) array of positions) and
returns an (N, 3) RGB or (N, 4) RGBA array. This will generate the LUT
instead of the stops array. If provided, the stops argument will be ignored.
interpolation : str, optional
Interpolation mode. Must be one of 'linear' (or `True`) or 'nearest' (or
Expand Down Expand Up @@ -1028,6 +1028,16 @@ def color_array(self) -> np.ndarray:
"""Return an (N, 4) array of RGBA values."""
return self._stops[:, 1:]

@property
def lut_func(self) -> LutCallable | None:
"""Callable that generates this colormap's colors, or None if defined by stops.

Evaluating it gives exact colors at arbitrary positions, unlike `stops` and
`color_array`, which hold a 256-point sampling. It is the stored callable,
so its output is not clipped and may be RGB rather than RGBA.
"""
return self._lut_func

def __len__(self) -> int:
return len(self._stops)

Expand Down
17 changes: 17 additions & 0 deletions tests/test_colormap.py
Original file line number Diff line number Diff line change
Expand Up @@ -8,6 +8,7 @@

from cmap import Color, Colormap
from cmap._colormap import ColorStop, ColorStops, _fill_stops
from cmap.data.matlab import prism

DATA = [
[0.0, 1.0, 0.0, 0.0, 1.0],
Expand Down Expand Up @@ -90,6 +91,22 @@ def test_colorstops() -> None:
assert reversed(cmap.color_stops) == ColorStops.parse(["b", "m", "r"])


def test_colorstops_lut_func() -> None:
def f(x: np.ndarray) -> np.ndarray:
return np.stack([x, 2 * x, -x], axis=-1)

cs = ColorStops(lut_func=f)
assert cs.lut_func is f
assert Colormap("prism").color_stops.lut_func is prism
assert Colormap("viridis").color_stops.lut_func is None
x = np.array([0.1234])
reversed_func = cs.reversed().lut_func
assert reversed_func is not None
npt.assert_allclose(reversed_func(x), f(1 - x), rtol=0, atol=1e-12)
with pytest.raises(AttributeError):
cs.lut_func = f


def test_colorstops_reversed_does_not_mutate_source() -> None:
stops = ColorStops.parse(["red", "green", "blue"])
before = np.asarray(stops).copy()
Expand Down
Loading