Coverage for python/lsst/multiprofit/fitting/fit_catalog.py: 93%

109 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__ = ["CatalogExposureABC", "CatalogFitterConfig", "ColumnInfo"] 

23 

24from abc import ABC, abstractmethod 

25from collections.abc import Iterable 

26from typing import ClassVar 

27 

28import astropy.units as u 

29import numpy as np 

30import pydantic 

31from astropy.table import Table 

32 

33import lsst.pex.config as pexConfig 

34 

35from ..componentconfig import GaussianComponentConfig, SersicComponentConfig 

36from ..modeller import ModelFitConfig 

37from ..utils import frozen_arbitrary_allowed_config 

38 

39 

40class CatalogExposureABC(ABC): 

41 """Interface for catalog-exposure pairs.""" 

42 

43 # TODO: add get_exposure (with Any return type?) 

44 

45 @abstractmethod 

46 def get_catalog(self) -> Iterable: 

47 """Return a row-iterable catalog covering an exposure.""" 

48 

49 

50class ColumnInfo(pydantic.BaseModel): 

51 """Metadata for a column in a catalog.""" 

52 

53 model_config: ClassVar[pydantic.ConfigDict] = frozen_arbitrary_allowed_config 

54 

55 dtype: str = pydantic.Field(title="Column data type name (numpy or otherwise)") 

56 key: str = pydantic.Field(title="Column key (name)") 

57 description: str = pydantic.Field("", title="Column description") 

58 unit: u.UnitBase | None = pydantic.Field(None, title="Column unit (astropy)") 

59 

60 

61class CatalogFitterConfig(pexConfig.Config): 

62 """Configuration for generic MultiProFit fitting tasks.""" 

63 

64 column_id = pexConfig.Field[str](default="id", doc="Catalog index column key") 

65 compute_errors = pexConfig.ChoiceField[str]( 

66 default="INV_HESSIAN_BESTFIT", 

67 doc="Whether/how to compute sqrt(variances) of each free parameter", 

68 allowed={ 

69 "NONE": "no errors computed", 

70 "INV_HESSIAN": "inverse hessian using noisy image as data", 

71 "INV_HESSIAN_BESTFIT": "inverse hessian using best-fit model as data", 

72 }, 

73 ) 

74 compute_errors_from_jacobian = pexConfig.Field[bool]( 

75 default=True, 

76 doc="Whether to estimate the Hessian from the Jacobian first, with finite differencing as a backup", 

77 ) 

78 compute_errors_no_covar = pexConfig.Field[bool]( 

79 default=True, 

80 doc="Whether to compute parameter errors independently, ignoring covariances", 

81 ) 

82 config_fit = pexConfig.ConfigField[ModelFitConfig](default=ModelFitConfig, doc="Fitter configuration") 

83 fit_centroid = pexConfig.Field[bool](default=True, doc="Fit centroid parameters") 

84 fit_linear_init = pexConfig.Field[bool](default=True, doc="Fit linear parameters after initialization") 

85 fit_linear_final = pexConfig.Field[bool](default=True, doc="Fit linear parameters after optimization") 

86 float_fill_value = pexConfig.Field[float]( 

87 default=np.nan, doc="Fill value for float fields when creating the output table." 

88 ) 

89 flag_errors = pexConfig.DictField( 

90 default={}, 

91 keytype=str, 

92 itemtype=str, 

93 doc="Flag column names to set, keyed by name of exception to catch", 

94 ) 

95 integer_fill_value = pexConfig.Field[int]( 

96 default=-1, doc="Fill value for integer fields when creating the output table." 

97 ) 

98 naming_scheme = pexConfig.ChoiceField[str]( 

99 doc="Naming scheme for column names", 

100 allowed={ 

101 "default": "snake_case with {component_name}[_{band}]_{parameter}[_err]", 

102 "camel": "CamelCase with {component_name}[_{band}]_{parameter}[Err]", 

103 "lsst": "snake_case with [{band}_]{component_name}_{parameter}[Err]", 

104 }, 

105 default="default", 

106 ) 

107 prefix_column = pexConfig.Field[str](default="mpf_", doc="Column name prefix") 

108 suffix_error = pexConfig.Field[str]( 

109 default="_err", 

110 doc="Default suffix for error columns. Can be overridden by naming_scheme.", 

111 ) 

112 

113 _format_flux = { 

114 "default": "{label}{band}_flux", 

115 "lsst": "{band}_{label}Flux", 

116 "camel": "{label}{band}Flux", 

117 } 

118 _key_cen = {"default": "_cen", "lsst": "_cen", "camel": "Cen"} 

119 _key_reff = {"default": f"_{SersicComponentConfig._size_label}", "lsst": "_reff", "camel": "Reff"} 

120 _key_rho = {"default": "_rho", "lsst": "_rho", "camel": "Rho"} 

121 _key_sigma = {"default": f"_{GaussianComponentConfig._size_label}", "lsst": "_sigma", "camel": "Sigma"} 

122 _key_sersicindex = {"default": "_sersic_index", "lsst": "sersic_index", "camel": "SersicIndex"} 

123 _suffix_dec = {"default": "_dec", "lsst": "_dec", "camel": "Dec"} 

124 _suffix_ra = {"default": "_ra", "lsst": "_ra", "camel": "Ra"} 

125 _suffix_ra_dec_cov = {"default": "_ra_dec_cov", "lsst": "_ra_dec_Cov", "camel": "RaDecCov"} 

126 _suffix_x = {"default": "_x", "lsst": "_x", "camel": "X"} 

127 _suffix_y = {"default": "_y", "lsst": "_y", "camel": "Y"} 

128 

129 def _get_label(self, format_name: str, values: dict[str, str]) -> str: 

130 """Get the label for part of a column name for a given format. 

131 

132 Parameters 

133 ---------- 

134 format_name 

135 The name of the format to get the label for. 

136 values 

137 The values of the name by format. 

138 

139 Returns 

140 ------- 

141 label 

142 The formatted label, if specified for that format, else the 

143 value for the default format. 

144 """ 

145 return values.get(format_name, values["default"]) 

146 

147 def get_key_cen(self) -> str: 

148 """Get the key for centroid columns.""" 

149 return self._get_label(self.naming_scheme, self._key_cen) 

150 

151 def get_key_flux(self, band: str, label: str = "") -> str: 

152 """Get the key for a flux column. 

153 

154 Parameters 

155 ---------- 

156 band 

157 The band of the flux column. 

158 label 

159 A label for this flux, e.g. a component name. 

160 

161 Returns 

162 ------- 

163 key_flux 

164 The flux column key. 

165 """ 

166 return self._get_label(self.naming_scheme, self._format_flux).format(band=band, label=label) 

167 

168 def get_key_reff(self) -> str: 

169 """Get the key for Sersic effective radius columns.""" 

170 return self._get_label(self.naming_scheme, self._key_reff) 

171 

172 def get_key_rho(self) -> str: 

173 """Get the key for ellipse rho columns.""" 

174 return self._get_label(self.naming_scheme, self._key_rho) 

175 

176 def get_key_sersicindex(self) -> str: 

177 """Get the key for Sersic index columns.""" 

178 return self._get_label(self.naming_scheme, self._key_sersicindex) 

179 

180 def get_key_sigma(self) -> str: 

181 """Get the key for Gaussian sigma columns.""" 

182 return self._get_label(self.naming_scheme, self._key_sigma) 

183 

184 def get_key_size(self, label_size: str) -> str: 

185 """Get the key for a size column by its label. 

186 

187 Parameters 

188 ---------- 

189 label_size 

190 The label of the size, usually specified in a ComponentConfig. 

191 

192 Returns 

193 ------- 

194 key_size 

195 The size column key. 

196 """ 

197 if label_size == GaussianComponentConfig._size_label: 

198 return self._get_label(self.naming_scheme, self._key_sigma) 

199 elif label_size == SersicComponentConfig._size_label: 199 ↛ 201line 199 didn't jump to line 201 because the condition on line 199 was always true

200 return self._get_label(self.naming_scheme, self._key_reff) 

201 return label_size 

202 

203 def get_prefixed_label(self, label: str, prefix: str) -> str: 

204 """Get a prefixed label with redundant underscores removed. 

205 

206 Parameters 

207 ---------- 

208 label 

209 The label to format. 

210 prefix 

211 The prefix to prepend. 

212 

213 Returns 

214 ------- 

215 label_prefixed 

216 The prefixed label, with redundant underscores removed. 

217 """ 

218 if label.startswith("_") and ((prefix == "") or (prefix[-1] == "_")): 

219 return f"{prefix}{label[1:]}" 

220 return f"{prefix}{label}" 

221 

222 def get_suffix_dec(self) -> str: 

223 """Get the suffix for declination columns.""" 

224 return self._get_label(self.naming_scheme, self._suffix_dec) 

225 

226 def get_suffix_ra(self) -> str: 

227 """Get the suffix for right ascension columns.""" 

228 return self._get_label(self.naming_scheme, self._suffix_ra) 

229 

230 def get_suffix_ra_dec_cov(self) -> str: 

231 """Get the suffix for right ascension columns.""" 

232 return self._get_label(self.naming_scheme, self._suffix_ra_dec_cov) 

233 

234 def get_suffix_x(self) -> str: 

235 """Get the suffix for x-axis columns.""" 

236 return self._get_label(self.naming_scheme, self._suffix_x) 

237 

238 def get_suffix_y(self) -> str: 

239 """Get the suffix for y-axis columns.""" 

240 return self._get_label(self.naming_scheme, self._suffix_y) 

241 

242 def make_catalog(self, n_rows: int, **kwargs): 

243 """Make a catalog with default-initialized column values. 

244 

245 Parameters 

246 ---------- 

247 n_rows 

248 The number of rows to create. 

249 **kwargs 

250 Keyword arguments to pass to self.schema. 

251 

252 Returns 

253 ------- 

254 catalog 

255 The initialized catalog. 

256 columns 

257 The columns as returned by self.schema. 

258 """ 

259 columns = self.schema(**kwargs) 

260 keys = [column.key for column in columns] 

261 prefix = self.prefix_column 

262 

263 idx_flag_first = keys.index("unknown_flag") 

264 idx_flag_last = idx_flag_first + len(self.flag_errors) 

265 dtypes = [(f"{prefix if col.key != self.column_id else ''}{col.key}", col.dtype) for col in columns] 

266 

267 results = Table(np.empty(n_rows, dtype=dtypes)) 

268 for colname in results.colnames: 

269 column = results[colname] 

270 dtype = column.dtype 

271 if ( 

272 value := ( 

273 self.float_fill_value 

274 if np.issubdtype(dtype, np.floating) 

275 else (self.integer_fill_value if np.issubdtype(dtype, np.integer) else None) 

276 ) 

277 ) is not None: 

278 column[:] = value 

279 

280 # Set nan-default flags to False instead 

281 errors = [] 

282 for flag in columns[idx_flag_first : (idx_flag_last + 1)]: 

283 column = results[f"{prefix}{flag.key}"] 

284 column[:] = False 

285 if not ((column.dtype == bool) and column.name.endswith("_flag")): 285 ↛ 286line 285 didn't jump to line 286 because the condition on line 285 was never true

286 errors.append(f"{column.name=} should end with _flag and {column.dtype=} must be bool") 

287 if errors: 287 ↛ 288line 287 didn't jump to line 288 because the condition on line 287 was never true

288 errors.append(f"These may be logic errors in {self=}") 

289 raise RuntimeError("\n".join(errors)) 

290 results.meta["config"] = self.toDict() 

291 

292 return results, columns 

293 

294 def schema( 

295 self, 

296 bands: list[str] | None = None, 

297 ) -> list[ColumnInfo]: 

298 """Return the schema as an ordered list of columns. 

299 

300 Parameters 

301 ---------- 

302 bands 

303 A list of band names to prefix band-dependent columns with. 

304 Band prefixes should not be used if None. 

305 

306 Returns 

307 ------- 

308 schema 

309 An ordered list of ColumnInfo instances. 

310 """ 

311 schema = [ 

312 ColumnInfo(key=self.column_id, dtype="i8"), 

313 ColumnInfo(key="n_iter", dtype="i4"), 

314 ColumnInfo(key="time_eval", dtype="f8", unit=u.s), 

315 ColumnInfo(key="time_fit", dtype="f8", unit=u.s), 

316 ColumnInfo(key="time_full", dtype="f8", unit=u.s), 

317 ColumnInfo(key="chisq_reduced", dtype="f8"), 

318 ColumnInfo(key="unknown_flag", dtype="bool"), 

319 ] 

320 schema.extend([ColumnInfo(key=key, dtype="bool") for key in self.flag_errors.keys()]) 

321 # Subclasses should always write out centroids even if not fitting 

322 # They are helpful for reconstructing models 

323 return schema