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

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/>. 

21 

22__all__ = ["Interpolator", "plot_sersicmix_interp"] 

23 

24import sys 

25from typing import Any 

26 

27import matplotlib as mpl 

28import matplotlib.pyplot as plt 

29import numpy as np 

30 

31import lsst.gauss2d.fit as g2f 

32 

33from .types import FigureAxes 

34 

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 

40 

41 Interpolator: TypeAlias = g2f.SersicMixInterpolator | tuple[type, dict[str, Any]] # noqa: UP040 

42 

43 

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. 

48 

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. 

57 

58 Returns 

59 ------- 

60 figure 

61 The resulting figure. 

62 

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] 

76 

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)) 

81 

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) 

89 

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] 

110 

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 ) 

115 

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) 

136 

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 

147 

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