Coverage for python/lsst/multiprofit/plotting/plot_loglike.py: 93%
89 statements
« prev ^ index » next coverage.py v7.16.1, created at 2026-09-23 09:58 +0000
« prev ^ index » next coverage.py v7.16.1, created at 2026-09-23 09:58 +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/>.
22__all__ = ["plot_loglike"]
24import itertools
26import matplotlib.pyplot as plt
27import numpy as np
29import lsst.gauss2d.fit as g2f
31from ..utils import get_params_uniq
32from .config import linestyles_default
33from .errorvalues import ErrorValues
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.
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`.
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())
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))
74 n_params = len(params)
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=}")
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
83 n_loglikes = len(loglike_init)
84 labels = [channel.name for channel in model.data.channels]
85 labels.extend(["prior", "total"])
87 for param in params:
88 param.fixed = True
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]]
97 diff_init = 1e-4 * np.sign(loglike_grads[row])
98 diff = diff_init
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
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]
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)
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"])
147 subplot = axes[row][1]
148 subplot.plot(values, dlls)
149 subplot.axhline(0, color="k")
150 subplot.set_ylabel("dloglike/dx")
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")
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()
170 for param in params:
171 param.fixed = False
173 return fig, ax