Coverage for python/lsst/multiprofit/plotting/plot_loglike.py: 93%

89 statements  

« prev     ^ index     » next       coverage.py v7.16.0, created at 2026-09-14 09:28 +0000

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__ = ["plot_loglike"] 

23 

24import itertools 

25 

26import matplotlib.pyplot as plt 

27import numpy as np 

28 

29import lsst.gauss2d.fit as g2f 

30 

31from ..utils import get_params_uniq 

32from .config import linestyles_default 

33from .errorvalues import ErrorValues 

34 

35 

36def plot_loglike( 

37 model: g2f.ModelD, 

38 params: list[g2f.ParameterD] | None = None, 

39 n_values: int = 15, 

40 errors: dict[str, ErrorValues] | None = None, 

41 values_reference: np.ndarray | None = None, 

42) -> tuple[plt.Figure, plt.Axes]: 

43 """Plot the loglikehood and derivatives vs free parameter values around 

44 best-fit values. 

45 

46 Parameters 

47 ---------- 

48 model 

49 The model to evaluate. 

50 params 

51 Free parameters to plot marginal loglikelihood for. 

52 n_values 

53 The number of evaluations to make on either side of each param value. 

54 errors 

55 A dict keyed by label of uncertainties to plot. Values must be the same 

56 length as `params`. 

57 values_reference 

58 Reference values to plot (e.g. true parameter values). Must be the same 

59 length as `params`. 

60 

61 Returns 

62 ------- 

63 fig, ax 

64 Matplotlib figure and axis handles, as returned by plt.subplots. 

65 """ 

66 if errors is None: 66 ↛ 67line 66 didn't jump to line 67 because the condition on line 66 was never true

67 errors = {} 

68 loglike_grads = np.array(model.compute_loglike_grad()) 

69 loglike_init = np.array(model.evaluate()) 

70 

71 if params is None: 71 ↛ 74line 71 didn't jump to line 74 because the condition on line 71 was always true

72 params = tuple(get_params_uniq(model, fixed=False)) 

73 

74 n_params = len(params) 

75 

76 if values_reference is not None and len(values_reference) != n_params: 76 ↛ 77line 76 didn't jump to line 77 because the condition on line 76 was never true

77 raise ValueError(f"{len(values_reference)=} != {n_params=}") 

78 

79 n_rows = n_params 

80 fig, ax = plt.subplots(nrows=n_rows, ncols=2, figsize=(10, 3 * n_rows)) 

81 axes = [ax] if (n_rows == 1) else ax 

82 

83 n_loglikes = len(loglike_init) 

84 labels = [channel.name for channel in model.data.channels] 

85 labels.extend(["prior", "total"]) 

86 

87 for param in params: 

88 param.fixed = True 

89 

90 for row, param in enumerate(params): 

91 value_init = param.value 

92 param.fixed = False 

93 values = [value_init] 

94 loglikes = [loglike_init * 0] 

95 dlls = [loglike_grads[row]] 

96 

97 diff_init = 1e-4 * np.sign(loglike_grads[row]) 

98 diff = diff_init 

99 

100 # TODO: This entire scheme should be improved/replaced 

101 # It sometimes takes excessively large steps 

102 # Option: Try to fit a curve once there are a couple of points 

103 # on each side of the peak 

104 idx_prev = -1 

105 for idx in range(2 * n_values): 

106 try: 

107 param.value_transformed += diff 

108 loglikes_new = np.array(model.evaluate()) - loglike_init 

109 dloglike_actual = np.sum(loglikes_new) - np.sum(loglikes[idx_prev]) 

110 values.append(param.value) 

111 loglikes.append(loglikes_new) 

112 dloglike_actual_abs = np.abs(dloglike_actual) 

113 if dloglike_actual_abs > 1: 

114 diff /= dloglike_actual_abs 

115 elif dloglike_actual_abs < 0.5: 

116 diff /= np.clip(dloglike_actual_abs, 0.2, 0.5) 

117 dlls.append(model.compute_loglike_grad()[0]) 

118 if idx == n_values: 

119 diff = -diff_init 

120 param.value = value_init 

121 idx_prev = 0 

122 else: 

123 idx_prev = -1 

124 except RuntimeError: 

125 break 

126 param.value = value_init 

127 param.fixed = True 

128 

129 subplot = axes[row][0] 

130 sorted = np.argsort(values) 

131 values = np.array(values)[sorted] 

132 loglikes = [loglikes[idx] for idx in sorted] 

133 dlls = np.array(dlls)[sorted] 

134 

135 for idx in range(n_loglikes): 

136 subplot.plot(values, [loglike[idx] for loglike in loglikes], label=labels[idx]) 

137 subplot.plot(values, np.sum(loglikes, axis=1), label=labels[-1]) 

138 vline_kwargs = dict(ymin=np.min(loglikes) - 1, ymax=np.max(loglikes) + 1, color="k") 

139 subplot.vlines(value_init, **vline_kwargs) 

140 

141 suffix = f" {param.label}" if param.label else "" 

142 subplot.legend() 

143 subplot.set_title(f"{param.name}{suffix}") 

144 subplot.set_ylabel("loglike") 

145 subplot.set_ylim(vline_kwargs["ymin"], vline_kwargs["ymax"]) 

146 

147 subplot = axes[row][1] 

148 subplot.plot(values, dlls) 

149 subplot.axhline(0, color="k") 

150 subplot.set_ylabel("dloglike/dx") 

151 

152 vline_kwargs = dict(ymin=np.min(dlls), ymax=np.max(dlls)) 

153 subplot.vlines(value_init, **vline_kwargs, color="k", label="fit") 

154 if values_reference is not None: 154 ↛ 157line 154 didn't jump to line 157 because the condition on line 154 was always true

155 subplot.vlines(values_reference[row], **vline_kwargs, color="b", label="ref") 

156 

157 cycler_linestyle = itertools.cycle(linestyles_default) 

158 for name_error, valerr in errors.items(): 

159 linestyle = valerr.kwargs_plot.pop("linestyle", next(cycler_linestyle)) 

160 for idx_ax in range(2): 

161 axes[row][idx_ax].vlines( 

162 [value_init - valerr.values[row], value_init + valerr.values[row]], 

163 linestyles=[linestyle, linestyle], 

164 label=name_error if (idx_ax == 1) else None, 

165 **valerr.kwargs_plot, 

166 **vline_kwargs, 

167 ) 

168 subplot.legend() 

169 

170 for param in params: 

171 param.fixed = False 

172 

173 return fig, ax