Coverage for python/lsst/multiprofit/transforms.py: 24%

28 statements  

« prev     ^ index     » next       coverage.py v7.16.0, created at 2026-09-23 02:31 -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__ = ["get_logit_limited", "transforms_ref", "verify_transform_derivative"] 

23 

24 

25from collections.abc import Iterable 

26from typing import Any 

27 

28import numpy as np 

29 

30import lsst.gauss2d.fit as g2f 

31 

32from .limits import limits_ref 

33 

34 

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

39 

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. 

50 

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 ) 

68 

69 

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. 

79 

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. 

97 

98 Raises 

99 ------ 

100 RuntimeError 

101 Raised if the transform derivative doesn't match finite differences 

102 within the specified tolerances. 

103 

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. 

108 

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 ) 

136 

137 

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}