Coverage for python/lsst/multiprofit/transforms.py: 24%
28 statements
« prev ^ index » next coverage.py v7.15.4, created at 2026-08-26 02:06 -0700
« prev ^ index » next coverage.py v7.15.4, created at 2026-08-26 02:06 -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__ = ["get_logit_limited", "transforms_ref", "verify_transform_derivative"]
25from collections.abc import Iterable
26from typing import Any
28import numpy as np
30import lsst.gauss2d.fit as g2f
32from .limits import limits_ref
35def get_logit_limited(
36 lower: float, upper: float, factor: float = 1.0, name: str | None = None
37) -> g2f.LogitLimitedTransformD:
38 """Get a logit transform stretched to span a different range than [0,1].
40 Parameters
41 ----------
42 lower
43 The lower limit of the range to span.
44 upper
45 The upper limit of the range to span.
46 factor
47 A multiplicative factor to apply to the transformed result.
48 name
49 A descriptive name for the transform.
51 Returns
52 -------
53 transform
54 A modified logit transform as specified.
55 """
56 return g2f.LogitLimitedTransformD(
57 limits=g2f.LimitsD(
58 min=lower,
59 max=upper,
60 name=(
61 name
62 if name is not None
63 else f"LogitLimitedTransformD(min={lower}, max={upper}, factor={factor})"
64 ),
65 ),
66 factor=factor,
67 )
70def verify_transform_derivative(
71 transform: g2f.TransformD,
72 value_transformed: float,
73 derivative: float | None = None,
74 abs_max: float = 1e6,
75 dx_ratios: Iterable[float] | None = None,
76 **kwargs: Any,
77) -> None:
78 """Verify that the derivative of a transform class is correct.
80 Parameters
81 ----------
82 transform
83 The transform to verify.
84 value_transformed
85 The un-transformed value at which to verify the transform.
86 derivative
87 The nominal derivative at value_transformed.
88 Must equal transform.derivative(value_transformed).
89 abs_max
90 The x value to skip verification if np.abs(derivative) > x.
91 dx_ratios
92 Iterable of signed ratios to set dx for finite differencing, where
93 dx = value*ratio (untransformed).
94 **kwargs
95 Keyword arguments to pass to np.isclose when comparing derivatives to
96 finite differences.
98 Raises
99 ------
100 RuntimeError
101 Raised if the transform derivative doesn't match finite differences
102 within the specified tolerances.
104 Notes
105 -----
106 derivative should only be specified if it has previously been computed for
107 the exact value_transformed, to avoid re-computing it unnecessarily.
109 Default dx_ratios are [1e-4, 1e-6, 1e-8, 1e-10, 1e-12, 1e-14].
110 Verification will test all ratios until at least one passes.
111 """
112 value = transform.reverse(value_transformed)
113 if derivative is None:
114 derivative = transform.derivative(value)
115 if np.abs(derivative) > abs_max:
116 # Skip testing finite differencing if the derivative is very large
117 # This might happen e.g. near the limits of the transformation
118 # TODO: Check if finite differencing can be improved for large values
119 return
120 if dx_ratios is None:
121 dx_ratios = [1e-4, 1e-6, 1e-8, 1e-10, 1e-12, 1e-14]
122 elif not (len(dx_ratios) > 0):
123 raise ValueError(f"{dx_ratios=} must not be empty")
124 for ratio in dx_ratios:
125 dx = value * ratio
126 fin_diff = (transform.forward(value + dx) - value_transformed) / dx
127 if not np.isfinite(fin_diff):
128 fin_diff = -(transform.forward(value - dx) - value_transformed) / dx
129 is_close = np.isclose(derivative, fin_diff, **kwargs)
130 if is_close:
131 return
132 raise RuntimeError(
133 f"{transform} derivative={derivative:.8e} != last "
134 f"finite diff.={fin_diff:8e} with {dx=} and dx_abs_max={abs_max}"
135 )
138transforms_ref = {
139 "none": g2f.UnitTransformD(),
140 "log": g2f.LogTransformD(),
141 "log10": g2f.Log10TransformD(),
142 "inverse": g2f.InverseTransformD(),
143 "logit": g2f.LogitTransformD(),
144 "logit_fluxfrac": get_logit_limited(
145 limits_ref["fluxfrac"].min,
146 limits_ref["fluxfrac"].max,
147 name=f"ref_logit_fluxfrac[{limits_ref['fluxfrac'].min}, {limits_ref['fluxfrac'].max}]",
148 ),
149 "logit_rho": get_logit_limited(
150 limits_ref["rho"].min,
151 limits_ref["rho"].max,
152 name=f"ref_logit_rho[{limits_ref['rho'].min}, {limits_ref['rho'].max}]",
153 ),
154 "logit_axrat": get_logit_limited(1e-4, 1, name="ref_logit_axrat[1e-4, 1]"),
155 "logit_axrat_prior": get_logit_limited(-0.0001, 1.1, name="ref_logit_axrat_prior[-0.0001, 1.1]"),
156 "logit_sersic": get_logit_limited(0.49, 6.01, name="ref_logit_sersic[0.49, 6.01]"),
157}