Coverage for python/lsst/images/cells/_aperture_corrections.py: 70%

149 statements  

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

1# This file is part of lsst-images. 

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# Use of this source code is governed by a 3-clause BSD-style 

10# license that can be found in the LICENSE file. 

11 

12from __future__ import annotations 

13 

14__all__ = ("CellApertureCorrectionMapSerializationModel", "CellField") 

15 

16from collections.abc import Mapping 

17from typing import TYPE_CHECKING, Any, ClassVar, final 

18 

19import astropy.table 

20import astropy.units 

21import numpy as np 

22import pydantic 

23 

24from .._cell_grid import CellGridBounds, CellIJ 

25from .._geom import BoundsError, Box 

26from .._image import Image 

27from ..fields import BaseField 

28from ..serialization import ( 

29 ArchiveReadError, 

30 ArchiveTree, 

31 InputArchive, 

32 InvalidParameterError, 

33 OutputArchive, 

34 TableModel, 

35) 

36 

37if TYPE_CHECKING: 

38 try: 

39 from lsst.afw.image import ApCorrMap as LegacyApCorrMap 

40 from lsst.cell_coadds import StitchedApertureCorrection as LegacyStitchedApertureCorrection 

41 except ImportError: 

42 type LegacyApCorrMap = Any # type: ignore[no-redef] 

43 type LegacyStitchedApertureCorrection = Any # type: ignore[no-redef] 

44 

45 

46@final 

47class CellField(BaseField): 

48 """A piecewise 2-d function on a cell-coadd grid. 

49 

50 Parameters 

51 ---------- 

52 bounds 

53 Description of the cell grid and any missing cells. Array entries for 

54 missing cells should be NaN. 

55 array 

56 A 2-d array of cell values with shape 

57 ``bounds.subgrid_size.as_tuple()``. 

58 unit 

59 Units of the field values, or `None` if dimensionless. 

60 

61 Notes 

62 ----- 

63 `CellField` is not directly serializable and is not included in the 

64 ``Field`` union type alias as a result. A `~collections.abc.Mapping` of 

65 `CellField` is instead serializable via 

66 `CellApertureCorrectionMapSerializationModel`. 

67 """ 

68 

69 def __init__( 

70 self, bounds: CellGridBounds, array: np.ndarray, unit: astropy.units.UnitBase | None = None 

71 ) -> None: 

72 self._array = array 

73 self._bounds = bounds 

74 self._unit = unit 

75 if self._array.shape != self._bounds.subgrid_size.as_tuple(): 75 ↛ 76line 75 didn't jump to line 76 because the condition on line 75 was never true

76 raise ValueError( 

77 f"Array shape ({self._array.shape}) differs from subgrid size ({self._bounds.subgrid_size})." 

78 ) 

79 

80 __hash__ = None # type: ignore[assignment] 

81 

82 @property 

83 def bounds(self) -> CellGridBounds: 

84 return self._bounds 

85 

86 @property 

87 def unit(self) -> astropy.units.UnitBase | None: 

88 return self._unit 

89 

90 @property 

91 def is_constant(self) -> bool: 

92 indices = iter(self._bounds.cell_indices()) 

93 try: 

94 first = self.value_in_cell(next(indices)) 

95 except StopIteration: 

96 return True 

97 for other_index in indices: 

98 if self.value_in_cell(other_index) != first: 

99 return False 

100 return True 

101 

102 def value_in_cell(self, key: CellIJ) -> float: 

103 """Return the value of the field in the cell with the given index. 

104 

105 Parameters 

106 ---------- 

107 key 

108 Index of the cell to evaluate. 

109 """ 

110 if key in self._bounds.missing: 

111 raise BoundsError(f"Cell {key} is missing for this field.") 

112 if not self._bounds.contains_cell(key): 

113 raise BoundsError(f"Cell {key} is out of bounds for this field.") 

114 index = key - self._bounds.subgrid_start 

115 try: 

116 return self._array[index.i, index.j] 

117 except IndexError: 

118 raise BoundsError(f"Cell {key} is out of bounds for this field.") from None 

119 

120 def quantity_in_cell(self, key: CellIJ) -> astropy.units.Quantity: 

121 """Return the quantity (value with units) of the field in the cell 

122 with the given index. 

123 

124 Parameters 

125 ---------- 

126 key 

127 Index of the cell to evaluate. 

128 """ 

129 return astropy.units.Quantity(self.value_in_cell(key), self._unit) 

130 

131 def _evaluate( 

132 self, *, x: np.ndarray, y: np.ndarray, quantity: bool 

133 ) -> np.ndarray | astropy.units.Quantity: 

134 # This implementation is optimized for the case where there are many 

135 # more evaluation points than cells. We could switch to an 

136 # implementation that zip-broadcast-iterates over x and y when that is 

137 # not the case, but that feels like a premature optimization right now. 

138 result = np.full(np.broadcast_shapes(y.shape, x.shape), np.nan, dtype=np.float64) 

139 for cell_index in self._bounds.cell_indices(): 

140 cell_bbox = self._bounds.grid.bbox_of(cell_index) 

141 result[cell_bbox.contains(x=x, y=y)] = self.value_in_cell(cell_index) 

142 if quantity: 

143 return astropy.units.Quantity(result, self._unit) 

144 return result 

145 

146 def render(self, bbox: Box | None = None, *, dtype: np.typing.DTypeLike | None = None) -> Image: 

147 if bbox is None: 

148 bbox = self._bounds.bbox 

149 bounds = self._bounds 

150 else: 

151 bounds = self._bounds[bbox] 

152 result = Image(np.nan, bbox=bbox, dtype=dtype, unit=self._unit) 

153 for cell_index in bounds.cell_indices(): 

154 cell_bbox = self._bounds.grid.bbox_of(cell_index).intersection(bbox) 

155 result[cell_bbox].array = self.value_in_cell(cell_index) 

156 return result 

157 

158 def _multiply_constant( 

159 self, factor: float | astropy.units.Quantity | astropy.units.UnitBase 

160 ) -> CellField: 

161 factor, unit = self._handle_factor_units(factor) 

162 return CellField(self._bounds, self._array * factor, unit=unit) 

163 

164 @staticmethod 

165 def from_legacy_aperture_correction( 

166 legacy: LegacyStitchedApertureCorrection, bounds: CellGridBounds 

167 ) -> CellField: 

168 """Convert from a legacy `lsst.cell_coadds.StitchedApertureCorrection`. 

169 

170 Parameters 

171 ---------- 

172 legacy 

173 Legacy field to convert. 

174 bounds 

175 The grid and bounds of the returned field. 

176 """ 

177 array = np.full(bounds.subgrid_size.as_tuple(), np.nan, dtype=np.float64) 

178 new_missing: set[CellIJ] = set() 

179 for cell_index in bounds.cell_indices(): 

180 array_index = cell_index - bounds.subgrid_start 

181 try: 

182 value = legacy.gc[cell_index.to_legacy()] 

183 except KeyError: 

184 value = np.nan 

185 new_missing.add(cell_index) 

186 array[array_index.i, array_index.j] = value 

187 if new_missing: 

188 bounds = CellGridBounds( 

189 grid=bounds.grid, missing=frozenset(bounds.missing | new_missing), bbox=bounds.bbox 

190 ) 

191 return CellField(bounds, array) 

192 

193 def to_legacy_aperture_correction(self) -> LegacyStitchedApertureCorrection: 

194 """Convert to a legacy 

195 `lsst.cell_coadds.StitchedApertureCorrection`. 

196 """ 

197 from lsst.cell_coadds import GridContainer, StitchedApertureCorrection 

198 

199 grid = self.bounds.grid.to_legacy() 

200 gc = GridContainer[float](grid.shape) 

201 for cell_index in self.bounds.cell_indices(): 

202 gc[cell_index.to_legacy()] = self.value_in_cell(cell_index) 

203 return StitchedApertureCorrection(grid, gc) 

204 

205 

206class CellApertureCorrectionMapSerializationModel(ArchiveTree): 

207 """A serialization model for a `~collections.abc.Mapping` of `CellField`, 

208 which is used to represent aperture corrections for cell-based coadds. 

209 """ 

210 

211 SCHEMA_NAME: ClassVar[str] = "cell_aperture_correction_map" 

212 SCHEMA_VERSION: ClassVar[str] = "1.0.0" 

213 MIN_READ_VERSION: ClassVar[int] = 1 

214 PUBLIC_TYPE: ClassVar[type] = dict 

215 

216 table: TableModel = pydantic.Field( 

217 description="Table with one row for each cell and different photometry algorithms in columns." 

218 ) 

219 bounds: CellGridBounds = pydantic.Field( 

220 description=( 

221 "Description of the cell grid and any missing cells. Array entries for " 

222 "missing cells should be NaN." 

223 ), 

224 ) 

225 

226 @staticmethod 

227 def serialize( 

228 aperture_correction_map: Mapping[str, CellField], archive: OutputArchive[Any] 

229 ) -> CellApertureCorrectionMapSerializationModel | None: 

230 if not aperture_correction_map: 230 ↛ 231line 230 didn't jump to line 231 because the condition on line 230 was never true

231 return None 

232 bounds = next(iter(aperture_correction_map.values())).bounds 

233 for field in aperture_correction_map.values(): 

234 if field.bounds.grid != bounds.grid: 

235 raise ValueError("Cell aperture corrections do not have a consistent grid.") 

236 if field.bounds.bbox != bounds.bbox: 

237 raise ValueError("Cell aperture corrections do not have a consistent bounding box.") 

238 if not all(field.bounds == bounds for field in aperture_correction_map.values()): 

239 # In rare cases some field can have more missing cells from the 

240 # others. This isn't worth a full denormalization into full 

241 # per-field bounds (especially since that's a schema change). 

242 # Instead, the shared, serialized 'bounds' only marks a cell as 

243 # missing if it's absent from all fields (as is the case when 

244 # there's no data), and we restore the per-field missing cells in 

245 # deserialized. 

246 always_missing: set[CellIJ] = set(bounds.missing) 

247 for field in aperture_correction_map.values(): 

248 always_missing &= field.bounds.missing 

249 bounds = CellGridBounds(grid=bounds.grid, missing=frozenset(always_missing), bbox=bounds.bbox) 

250 if any(field.unit is not None for field in aperture_correction_map.values()): 250 ↛ 251line 250 didn't jump to line 251 because the condition on line 250 was never true

251 raise ValueError("Aperture corrections should be dimensionless.") 

252 table = astropy.table.Table( 

253 rows=[cell_index.as_tuple() for cell_index in bounds.cell_indices()], names=["cell_i", "cell_j"] 

254 ) 

255 good_cell_mask = np.ones(bounds.subgrid_size.as_tuple(), dtype=bool) 

256 for cell_index in bounds.missing: 

257 array_index = cell_index - bounds.subgrid_start 

258 good_cell_mask[array_index.i, array_index.j] = False 

259 for name, field in aperture_correction_map.items(): 

260 table.add_column(field._array[good_cell_mask], name=name, copy=False) 

261 return CellApertureCorrectionMapSerializationModel( 

262 table=archive.add_table(table, name="table"), bounds=bounds 

263 ) 

264 

265 def deserialize(self, archive: InputArchive[Any], **kwargs: Any) -> dict[str, CellField]: 

266 if kwargs: 266 ↛ 267line 266 didn't jump to line 267 because the condition on line 266 was never true

267 raise InvalidParameterError( 

268 f"Unrecognized parameters for cell aperture correction map: {set(kwargs.keys())}." 

269 ) 

270 good_cell_mask = np.zeros(self.bounds.subgrid_size.as_tuple(), dtype=bool) 

271 table = archive.get_table(self.table) 

272 for tbl_ij, cell_index in zip( 

273 table["cell_i", "cell_j"].iterrows(), self.bounds.cell_indices(), strict=True 

274 ): 

275 if cell_index.as_tuple() != tbl_ij: 275 ↛ 276line 275 didn't jump to line 276 because the condition on line 275 was never true

276 raise ArchiveReadError( 

277 "Inconsistency between serialized aperture correction bounds and table." 

278 ) 

279 array_index = cell_index - self.bounds.subgrid_start 

280 good_cell_mask[array_index.i, array_index.j] = True 

281 result: dict[str, CellField] = {} 

282 for name, column in table.columns.items(): 

283 if name in ("cell_i", "cell_j"): 

284 continue 

285 extra_missing_mask = np.isnan(column) 

286 # If there are any NaN entries in this column, make a custom 

287 # CellGridBounds with additional missing cells for this field. 

288 bounds = self.bounds 

289 if extra_missing_mask.any(): 

290 extra_missing = frozenset( 

291 {CellIJ(i=i, j=j) for i, j in table["cell_i", "cell_j"][extra_missing_mask].iterrows()} 

292 ) 

293 bounds = CellGridBounds( 

294 grid=self.bounds.grid, missing=self.bounds.missing | extra_missing, bbox=self.bounds.bbox 

295 ) 

296 array = np.full(self.bounds.subgrid_size.as_tuple(), np.nan, dtype=np.float64) 

297 array[good_cell_mask] = column 

298 result[name] = CellField(bounds, array) 

299 return result