Coverage for python/lsst/multiprofit/plotting/plot_sersicmix_interp.py: 9%
100 statements
« prev ^ index » next coverage.py v7.16.0, created at 2026-09-19 02:18 -0700
« prev ^ index » next coverage.py v7.16.0, created at 2026-09-19 02:18 -0700
1# This file is part of multiprofit.
2#
3# Developed for the LSST Data Management System.
4# This product includes software developed by the LSST Project
5# (https://www.lsst.org).
6# See the COPYRIGHT file at the top-level directory of this distribution
7# for details of code ownership.
8#
9# This program is free software: you can redistribute it and/or modify
10# it under the terms of the GNU General Public License as published by
11# the Free Software Foundation, either version 3 of the License, or
12# (at your option) any later version.
13#
14# This program is distributed in the hope that it will be useful,
15# but WITHOUT ANY WARRANTY; without even the implied warranty of
16# MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the
17# GNU General Public License for more details.
18#
19# You should have received a copy of the GNU General Public License
20# along with this program. If not, see <https://www.gnu.org/licenses/>.
22__all__ = ["Interpolator", "plot_sersicmix_interp"]
24import sys
25from typing import Any
27import matplotlib as mpl
28import matplotlib.pyplot as plt
29import numpy as np
31import lsst.gauss2d.fit as g2f
33from .types import FigureAxes
35_has_py_13_plus = sys.version_info >= (3, 13, 0)
36if _has_py_13_plus: 36 ↛ 39line 36 didn't jump to line 39 because the condition on line 36 was always true
37 Interpolator: type = g2f.SersicMixInterpolator | tuple[type, dict[str, Any]]
38else:
39 from typing import TypeAlias
41 Interpolator: TypeAlias = g2f.SersicMixInterpolator | tuple[type, dict[str, Any]] # noqa: UP040
44def plot_sersicmix_interp(
45 interps: dict[str, tuple[Interpolator, str | tuple]], n_ser: np.ndarray, **kwargs: Any
46) -> FigureAxes:
47 """Plot Gaussian mixture Sersic profile interpolated values.
49 Parameters
50 ----------
51 interps
52 Dict of interpolators by name.
53 n_ser
54 Array of Sersic index values to plot interpolated quantities for.
55 **kwargs
56 Keyword arguments to pass to matplotlib.pyplot.subplots.
58 Returns
59 -------
60 figure
61 The resulting figure.
63 Examples
64 --------
65 plot_sersicmix_interp
66 """
67 orders = {
68 name: interp.order
69 for name, (interp, _) in interps.items()
70 if isinstance(interp, g2f.SersicMixInterpolator)
71 }
72 order = set(orders.values())
73 if not len(order) == 1:
74 raise ValueError(f"len(set({orders})) != 1; all interpolators must have the same order")
75 order = tuple(order)[0]
77 cmap = mpl.cm.get_cmap("tab20b")
78 colors_ord = [None] * order
79 for i_ord in range(order):
80 colors_ord[i_ord] = cmap(i_ord / (order - 1.0))
82 n_ser_min = np.min(n_ser)
83 n_ser_max = np.max(n_ser)
84 knots = g2f.sersic_mix_knots(order=order)
85 n_knots = len(knots)
86 integrals_knots = np.empty((n_knots, order))
87 sigmas_knots = np.empty((n_knots, order))
88 n_ser_knots = np.empty(n_knots)
90 i_knot_first = None
91 i_knot_last = n_knots
92 for i_knot, knot in enumerate(knots):
93 if i_knot_first is None:
94 if knot.sersicindex > n_ser_min:
95 i_knot_first = i_knot
96 else:
97 continue
98 if knot.sersicindex > n_ser_max:
99 i_knot_last = i_knot
100 break
101 n_ser_knots[i_knot] = knot.sersicindex
102 for i_ord in range(order):
103 values = knot.values[i_ord]
104 integrals_knots[i_knot, i_ord] = values.integral
105 sigmas_knots[i_knot, i_ord] = values.sigma
106 range_knots = range(i_knot_first, i_knot_last)
107 integrals_knots = integrals_knots[range_knots, :]
108 sigmas_knots = sigmas_knots[range_knots, :]
109 n_ser_knots = n_ser_knots[range_knots]
111 n_values = len(n_ser)
112 integrals, dintegrals, sigmas, dsigmas = (
113 {name: np.empty((n_values, order)) for name in interps} for _ in range(4)
114 )
116 for name, (interp, _) in interps.items():
117 if not isinstance(interp, g2f.SersicMixInterpolator):
118 kwargs = interp[1] if interp[1] is not None else {}
119 interp = interp[0]
120 x = [knot.sersicindex for knot in knots]
121 for i_ord in range(order):
122 integrals_i = np.empty(n_knots, dtype=float)
123 sigmas_i = np.empty(n_knots, dtype=float)
124 for i_knot, knot in enumerate(knots):
125 integrals_i[i_knot] = knot.values[i_ord].integral
126 sigmas_i[i_knot] = knot.values[i_ord].sigma
127 interp_int = interp(x, integrals_i, **kwargs)
128 dinterp_int = interp_int.derivative()
129 interp_sigma = interp(x, sigmas_i, **kwargs)
130 dinterp_sigma = interp_sigma.derivative()
131 for i_val, value in enumerate(n_ser):
132 integrals[name][i_val, i_ord] = interp_int(value)
133 sigmas[name][i_val, i_ord] = interp_sigma(value)
134 dintegrals[name][i_val, i_ord] = dinterp_int(value)
135 dsigmas[name][i_val, i_ord] = dinterp_sigma(value)
137 for i_val, value in enumerate(n_ser):
138 for name, (interp, _) in interps.items():
139 if isinstance(interp, g2f.SersicMixInterpolator):
140 values = interp.integralsizes(value)
141 derivs = interp.integralsizes_derivs(value)
142 for i_ord in range(order):
143 integrals[name][i_val, i_ord] = values[i_ord].integral
144 sigmas[name][i_val, i_ord] = values[i_ord].sigma
145 dintegrals[name][i_val, i_ord] = derivs[i_ord].integral
146 dsigmas[name][i_val, i_ord] = derivs[i_ord].sigma
148 fig, axes = plt.subplots(2, 2, **kwargs)
149 for idx_row, (yv, yd, yk, y_label) in (
150 (0, (integrals, dintegrals, integrals_knots, "integral")),
151 (1, (sigmas, dsigmas, sigmas_knots, "sigma")),
152 ):
153 is_label_row = idx_row == 1
154 for idx_col, y_i, y_prefix in ((0, yv, ""), (1, yd, "d")):
155 is_label_col = idx_col == 0
156 make_label = is_label_col and is_label_row
157 axis = axes[idx_row, idx_col]
158 if is_label_col:
159 for i_ord in range(order):
160 axis.plot(
161 n_ser_knots,
162 yk[:, i_ord],
163 "kx",
164 label="knots" if make_label and (i_ord == 0) else None,
165 )
166 for name, (_, lstyle) in interps.items():
167 for i_ord in range(order):
168 label = f"{name}" if make_label and (i_ord == 0) else None
169 axis.plot(n_ser, y_i[name][:, i_ord], c=colors_ord[i_ord], label=label, linestyle=lstyle)
170 axis.set_xlim((n_ser_min, n_ser_max))
171 axis.set_ylabel(f"{y_prefix}{y_label}")
172 if make_label:
173 axis.legend(loc="upper left")
174 return fig, axes