Coverage for python/lsst/images/_transforms/_transform.py: 53%

292 statements  

« prev     ^ index     » next       coverage.py v7.16.0, created at 2026-09-16 09:59 +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__ = ( 

15 "Transform", 

16 "TransformCompositionError", 

17 "TransformSerializationModel", 

18) 

19 

20import enum 

21import textwrap 

22from collections.abc import Iterable 

23from typing import TYPE_CHECKING, Any, ClassVar, TypeVar, assert_type, cast, final, overload 

24 

25import astropy.io.fits.header 

26import astropy.units as u 

27import numpy as np 

28import numpy.typing as npt 

29import pydantic 

30 

31from .._concrete_bounds import BoundsSerializationModel 

32from .._geom import XY, YX, Bounds, Box 

33from ..describe import DescribableMixin, DescribeOptions, FieldRole, Report, ReportField 

34from ..serialization import ArchiveReadError, ArchiveTree, InputArchive, InvalidParameterError, OutputArchive 

35from . import _ast as astshim 

36from ._frames import Frame, SerializableFrame, SkyFrame 

37 

38if TYPE_CHECKING: 

39 try: 

40 from lsst.afw.geom import TransformPoint2ToPoint2 as LegacyTransform 

41 except ImportError: 

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

43 

44# These pre-python-3.12 declaration are needed by Sphinx (probably the 

45# autodoc-typehints plugin. 

46I = TypeVar("I", bound=Frame) # noqa: E741 

47O = TypeVar("O", bound=Frame) # noqa: E741 

48P = TypeVar("P", bound=pydantic.BaseModel) 

49 

50 

51class TransformCompositionError(RuntimeError): 

52 """Exception raised when two transforms cannot be composed.""" 

53 

54 

55def _frame_label(frame: Frame) -> str: 

56 """Return a short name identifying a coordinate frame. 

57 

58 Most frames are pydantic models whose ``str`` spells out every field, 

59 which is far too much for a one-line summary, so the type name stands in 

60 for them. The sky frames are an enum, where the member is the name. 

61 """ 

62 if isinstance(frame, enum.Enum): 

63 return str(frame.value) 

64 return type(frame).__name__ 

65 

66 

67@final 

68class Transform[I: Frame, O: Frame](DescribableMixin): 

69 """A transform that maps two coordinate frames. 

70 

71 Parameters 

72 ---------- 

73 in_frame 

74 Input coordinate frame. 

75 out_frame 

76 Output coordinate frame. 

77 ast_mapping 

78 AST mapping that implements the transform. 

79 in_bounds 

80 Bounds of the input frame, defaulting to the input frame's 

81 bounding box. 

82 out_bounds 

83 Bounds of the output frame, defaulting to the output frame's 

84 bounding box. 

85 components 

86 Component transforms that this transform was composed from. 

87 

88 Notes 

89 ----- 

90 The `Transform` class constructor is considered a private implementation 

91 detail. Instead of using this, various factory methods are available: 

92 

93 - `from_fits_wcs` constructs a transform from a FITS WCS, as represented 

94 `astropy.wcs.WCS`; 

95 - `then` composes two transforms; 

96 - `identity` constructs a trivial transform that does nothing; 

97 - `affine` contructs an affine transform from a 2x2 or 3x3 matrix; 

98 - `inverted` returns the inverse of a transform; 

99 - `from_legacy` converts an `lsst.afw.geom.Transform` instance. 

100 

101 When applied to celestial coordinate systems, ``x=ra`` and ``y=dec``. 

102 `SkyProjection` provides a more natural interface for pixel-to-sky 

103 transforms. 

104 

105 `Transform` is conceptually immutable (the internal AST Mapping should 

106 never be modified in-place after construction), and hence does not need to 

107 be copied when any object that holds it is copied. 

108 """ 

109 

110 def __init__( 

111 self, 

112 in_frame: I, 

113 out_frame: O, 

114 ast_mapping: astshim.Mapping, 

115 in_bounds: Bounds | None = None, 

116 out_bounds: Bounds | None = None, 

117 components: Iterable[Transform[Any, Any]] = (), 

118 ) -> None: 

119 self._in_frame = in_frame 

120 self._out_frame = out_frame 

121 self._ast_mapping = ast_mapping 

122 self._in_bounds = in_bounds or getattr(in_frame, "bbox", None) 

123 self._out_bounds = out_bounds or getattr(out_frame, "bbox", None) 

124 self._components = list(components) 

125 

126 def __eq__(self, other: Any) -> bool: 

127 if self is other: 

128 # Short circuit for case where you are quickly checking 

129 # that the image WCS and variance WCS are the same object. 

130 return True 

131 if not isinstance(other, Transform): 

132 return NotImplemented 

133 if self._ast_mapping != other._ast_mapping: 

134 return False 

135 if self._in_bounds != other._in_bounds: 

136 return False 

137 if self._out_bounds != other._out_bounds: 

138 return False 

139 if self._in_frame != other._in_frame: 

140 return False 

141 if self._out_frame != other._out_frame: 

142 return False 

143 if self._components != other._components: 

144 return False 

145 return True 

146 

147 @staticmethod 

148 def from_fits_wcs( 

149 fits_wcs: astropy.wcs.WCS, 

150 in_frame: I, 

151 out_frame: O, 

152 in_bounds: Bounds | None = None, 

153 out_bounds: Bounds | None = None, 

154 x0: int = 0, 

155 y0: int = 0, 

156 ) -> Transform[I, O]: 

157 """Construct a transform from a FITS WCS. 

158 

159 Parameters 

160 ---------- 

161 fits_wcs 

162 FITS WCS to convert. 

163 in_frame 

164 Coordinate frame for input points to the forward transform. 

165 out_frame 

166 Coordinate frame for output points from the forward transform. 

167 in_bounds 

168 The region that bounds valid input points. 

169 out_bounds 

170 The region that bounds valid output points. 

171 x0 

172 Logical coordinate of the first column in the array this WCS 

173 relates to world coordinates. 

174 y0 

175 Logical coordinate of the first column in the array this WCS 

176 relates to world coordinates. 

177 

178 Notes 

179 ----- 

180 The ``x0`` and ``y0`` parameters reflect the fact that for FITS, the 

181 first row and column are always labeled ``(1, 1)``, while in Astropy 

182 and most other Python libraries, they are ``(0, 0)``. The `types` in 

183 this package (e.g. `Image`, `Mask`) allow them to be any pair of 

184 integers. 

185 

186 See Also 

187 -------- 

188 SkyProjection.from_fits_wcs 

189 """ 

190 ast_stream = astshim.StringStream(fits_wcs.to_header_string(relax=True)) 

191 ast_fits_chan = astshim.FitsChan(ast_stream, "Encoding=FITS-WCS, SipReplace=0, IWC=1") 

192 ast_frame_set = ast_fits_chan.read() 

193 _prepend_ast_shift(ast_frame_set, x=x0 - 1.0, y=y0 - 1.0, ast_domain="PIXEL") 

194 return Transform( 

195 in_frame, 

196 out_frame, 

197 ast_frame_set, 

198 in_bounds=in_bounds, 

199 out_bounds=out_bounds, 

200 ) 

201 

202 @staticmethod 

203 def identity(frame: I) -> Transform[I, I]: 

204 """Construct a trivial transform that maps a frame to itelf. 

205 

206 Parameters 

207 ---------- 

208 frame 

209 Frame used for both input and output points. 

210 """ 

211 return Transform(frame, frame, astshim.UnitMap(2)) 

212 

213 @staticmethod 

214 def affine(in_frame: I, out_frame: O, matrix: np.ndarray) -> Transform[I, O]: 

215 """Construct an affine transform from a matrix. 

216 

217 Parameters 

218 ---------- 

219 in_frame 

220 Coordinate frame for input points to the forward transform. 

221 out_frame 

222 Coordinate frame for output points from the forward transform. 

223 matrix 

224 Matrix of coefficients, either a 2x2 linear transform or a 3x3 

225 augmented affine transform, with a shift embedded in the third 

226 column and ``[0, 0, 1]`` the third row. 

227 """ 

228 if matrix.shape == (2, 2): 

229 return Transform(in_frame, out_frame, astshim.MatrixMap(matrix.copy())) 

230 elif matrix.shape == (3, 3): 230 ↛ 237line 230 didn't jump to line 237 because the condition on line 230 was always true

231 linear = astshim.MatrixMap(matrix[:2, :2].copy()) 

232 shift = astshim.ShiftMap(matrix[:2, 2]) 

233 if not np.array_equal(matrix[2, :], np.array([0.0, 0.0, 1.0])): 233 ↛ 234line 233 didn't jump to line 234 because the condition on line 233 was never true

234 raise ValueError("3x3 affine transform array must have [0, 0, 1] in its last row.") 

235 return Transform(in_frame, out_frame, linear.then(shift)) 

236 else: 

237 raise ValueError("Affine transform array must be 2x2 or 3x3.") 

238 

239 @property 

240 def in_frame(self) -> I: 

241 """Coordinate frame for input points.""" 

242 return self._in_frame 

243 

244 @property 

245 def out_frame(self) -> O: 

246 """Coordinate frame for output points.""" 

247 return self._out_frame 

248 

249 @property 

250 def in_bounds(self) -> Bounds | None: 

251 """The region that bounds valid input points (`Bounds` | `None`).""" 

252 return self._in_bounds 

253 

254 @property 

255 def out_bounds(self) -> Bounds | None: 

256 """The region that bounds valid output points (`Bounds` | `None`).""" 

257 return self._out_bounds 

258 

259 def show(self, simplified: bool = False, comments: bool = False) -> str: 

260 """Return the AST native representation of the transform. 

261 

262 Parameters 

263 ---------- 

264 simplified 

265 Whether to ask AST to simplify the mapping before showing it. 

266 This will make it much more likely that two equivalent transforms 

267 have the same `show` result. If the internal mapping is actually 

268 a frame set (as needed to round-trip legacy 

269 `lsst.afw.geom.SkyWcs` objects), this will also just show the 

270 mapping with no frame set information. 

271 comments 

272 Whether to include descriptive comments. 

273 """ 

274 ast_mapping = self._ast_mapping 

275 if simplified: 

276 if isinstance(ast_mapping, astshim.FrameSet): 276 ↛ 277line 276 didn't jump to line 277 because the condition on line 276 was never true

277 ast_mapping = ast_mapping.getMapping() 

278 ast_mapping = ast_mapping.simplified() 

279 return ast_mapping.show(comments) 

280 

281 def _describe(self, options: DescribeOptions = DescribeOptions(), /) -> Report: 

282 """Return a `Report` describing this transform. 

283 

284 Parameters 

285 ---------- 

286 options : `DescribeOptions`, optional 

287 Rendering options. `DescribeOptions.brief` reports only the input 

288 and output frames, skipping the bounds and the (potentially very 

289 long) AST mapping dump. 

290 """ 

291 fields = [ 

292 ReportField(label="in_frame", value=self.in_frame, role=FieldRole.DERIVED), 

293 ReportField(label="out_frame", value=self.out_frame, role=FieldRole.DERIVED), 

294 ] 

295 if not options.brief: 295 ↛ 305line 295 didn't jump to line 305 because the condition on line 295 was always true

296 fields += [ 

297 ReportField(label="in_bounds", value=self.in_bounds, role=FieldRole.DERIVED), 

298 ReportField(label="out_bounds", value=self.out_bounds, role=FieldRole.DERIVED), 

299 ReportField( 

300 label="mapping", 

301 value=self.show(simplified=True), 

302 role=FieldRole.DERIVED, 

303 ), 

304 ] 

305 return Report( 

306 type_name="Transform", 

307 summary=f"{_frame_label(self.in_frame)} → {_frame_label(self.out_frame)}", 

308 fields=fields, 

309 ) 

310 

311 @overload 

312 def apply_forward(self, point: XY[int | float] | YX[int | float], /) -> XY[float]: ... 

313 

314 @overload 

315 def apply_forward(self, point: XY[npt.ArrayLike] | YX[npt.ArrayLike], /) -> XY[np.ndarray]: ... 

316 

317 @overload 

318 def apply_forward(self, /, *, x: int | float, y: int | float) -> XY[float]: ... 

319 

320 @overload 

321 def apply_forward(self, /, *, x: npt.ArrayLike, y: npt.ArrayLike) -> XY[np.ndarray]: ... 

322 

323 def apply_forward( 

324 self, point: XY[Any] | YX[Any] | None = None, /, *, x: Any = None, y: Any = None 

325 ) -> XY[float] | XY[np.ndarray]: 

326 """Apply the forward transform to one or more points. 

327 

328 Parameters 

329 ---------- 

330 point 

331 An `XY` or `YX` coordinate pair to transform. Mutually exclusive 

332 with ``x`` and ``y``. 

333 x : `float` | array-like 

334 ``x`` values of the points to transform, as a scalar or any 

335 array-like. Results are broadcast against ``y``. 

336 Mutually exclusive with ``point``. 

337 y : `float` | array-like 

338 ``y`` values of the points to transform, as a scalar or any 

339 array-like. Results are broadcast against ``x``. 

340 Mutually exclusive with ``point``. 

341 

342 Returns 

343 ------- 

344 `XY` [`float` | `numpy.ndarray`] 

345 The transformed point or points. A scalar input pair returns 

346 `XY` of `float`; array-like inputs return `XY` of 

347 `numpy.ndarray` with the broadcast shape of ``x`` and ``y``. 

348 """ 

349 match point: 

350 case None: 

351 if x is None or y is None: 

352 raise TypeError("Pass either a point or both x= and y= to 'apply_forward'.") 

353 case XY() | YX(): 353 ↛ 357line 353 didn't jump to line 357 because the pattern on line 353 always matched

354 if x is not None or y is not None: 

355 raise TypeError("'apply_forward' point argument is mutually exclusive with x= and y=.") 

356 x, y = point.x, point.y 

357 case _: 

358 raise TypeError(f"Unexpected positional argument type: {type(point)!r}.") 

359 return _standardize_xy( 

360 _ast_apply( 

361 self._ast_mapping.applyForward, 

362 x=self._in_frame.standardize_x(x), 

363 y=self._in_frame.standardize_y(y), 

364 ), 

365 self._out_frame, 

366 ) 

367 

368 @overload 

369 def apply_inverse(self, point: XY[int | float] | YX[int | float], /) -> XY[float]: ... 

370 

371 @overload 

372 def apply_inverse(self, point: XY[npt.ArrayLike] | YX[npt.ArrayLike], /) -> XY[np.ndarray]: ... 

373 

374 @overload 

375 def apply_inverse(self, /, *, x: int | float, y: int | float) -> XY[float]: ... 

376 

377 @overload 

378 def apply_inverse(self, /, *, x: npt.ArrayLike, y: npt.ArrayLike) -> XY[np.ndarray]: ... 

379 

380 def apply_inverse( 

381 self, point: XY[Any] | YX[Any] | None = None, /, *, x: Any = None, y: Any = None 

382 ) -> XY[float] | XY[np.ndarray]: 

383 """Apply the inverse transform to one or more points. 

384 

385 Parameters 

386 ---------- 

387 point 

388 An `XY` or `YX` coordinate pair to transform. Mutually exclusive 

389 with ``x`` and ``y``. 

390 x : `float` | array-like 

391 ``x`` values of the points to transform, as a scalar or any 

392 array-like. Results are broadcast against ``y``. 

393 Mutually exclusive with ``point``. 

394 y : `float` | array-like 

395 ``y`` values of the points to transform, as a scalar or any 

396 array-like. Results are broadcast against ``x``. 

397 Mutually exclusive with ``point``. 

398 

399 Returns 

400 ------- 

401 `XY` [`float` | `numpy.ndarray`] 

402 The transformed point or points. A scalar input pair returns 

403 `XY` of `float`; array-like inputs return `XY` of 

404 `numpy.ndarray` with the broadcast shape of ``x`` and ``y``. 

405 """ 

406 match point: 

407 case None: 

408 if x is None or y is None: 408 ↛ 409line 408 didn't jump to line 409 because the condition on line 408 was never true

409 raise TypeError("Pass either a point or both x= and y= to 'apply_inverse'.") 

410 case XY() | YX(): 410 ↛ 414line 410 didn't jump to line 414 because the pattern on line 410 always matched

411 if x is not None or y is not None: 411 ↛ 412line 411 didn't jump to line 412 because the condition on line 411 was never true

412 raise TypeError("'apply_inverse' point argument is mutually exclusive with x= and y=.") 

413 x, y = point.x, point.y 

414 case _: 

415 raise TypeError(f"Unexpected positional argument type: {type(point)!r}.") 

416 return _standardize_xy( 

417 _ast_apply( 

418 self._ast_mapping.applyInverse, 

419 x=self._out_frame.standardize_x(x), 

420 y=self._out_frame.standardize_y(y), 

421 ), 

422 self._in_frame, 

423 ) 

424 

425 @overload 

426 def apply_forward_q(self, point: XY[u.Quantity] | YX[u.Quantity], /) -> XY[u.Quantity]: ... 

427 

428 @overload 

429 def apply_forward_q(self, /, *, x: u.Quantity, y: u.Quantity) -> XY[u.Quantity]: ... 

430 

431 def apply_forward_q( 

432 self, point: XY[u.Quantity] | YX[u.Quantity] | None = None, /, *, x: Any = None, y: Any = None 

433 ) -> XY[u.Quantity]: 

434 """Apply the forward transform to one or more unit-aware points. 

435 

436 Parameters 

437 ---------- 

438 point 

439 An `XY` or `YX` coordinate pair of `~astropy.units.Quantity` to 

440 transform. Mutually exclusive with ``x`` and ``y``. 

441 x 

442 ``x`` values of the points to transform. 

443 Mutually exclusive with ``point``. 

444 y 

445 ``y`` values of the points to transform. 

446 Mutually exclusive with ``point``. 

447 

448 Returns 

449 ------- 

450 `XY` [`astropy.units.Quantity`] 

451 The transformed point or points. 

452 """ 

453 match point: 

454 case None: 

455 if x is None or y is None: 455 ↛ 456line 455 didn't jump to line 456 because the condition on line 455 was never true

456 raise TypeError("Pass either a point or both x= and y= to 'apply_forward_q'.") 

457 case XY() | YX(): 457 ↛ 461line 457 didn't jump to line 461 because the pattern on line 457 always matched

458 if x is not None or y is not None: 458 ↛ 459line 458 didn't jump to line 459 because the condition on line 458 was never true

459 raise TypeError("'apply_forward_q' point argument is mutually exclusive with x= and y=.") 

460 x, y = point.x, point.y 

461 case _: 

462 raise TypeError(f"Unexpected positional argument type: {type(point)!r}.") 

463 xy = self.apply_forward(x=x.to_value(self._in_frame.unit), y=y.to_value(self._in_frame.unit)) 

464 return XY(xy.x * self._out_frame.unit, xy.y * self._out_frame.unit) 

465 

466 @overload 

467 def apply_inverse_q(self, point: XY[u.Quantity] | YX[u.Quantity], /) -> XY[u.Quantity]: ... 

468 

469 @overload 

470 def apply_inverse_q(self, /, *, x: u.Quantity, y: u.Quantity) -> XY[u.Quantity]: ... 

471 

472 def apply_inverse_q( 

473 self, point: XY[u.Quantity] | YX[u.Quantity] | None = None, /, *, x: Any = None, y: Any = None 

474 ) -> XY[u.Quantity]: 

475 """Apply the inverse transform to one or more unit-aware points. 

476 

477 Parameters 

478 ---------- 

479 point 

480 An `XY` or `YX` coordinate pair of `~astropy.units.Quantity` to 

481 transform. Mutually exclusive with ``x`` and ``y``. 

482 x 

483 ``x`` values of the points to transform. 

484 Mutually exclusive with ``point``. 

485 y 

486 ``y`` values of the points to transform. 

487 Mutually exclusive with ``point``. 

488 

489 Returns 

490 ------- 

491 `XY` [`astropy.units.Quantity`] 

492 The transformed point or points. 

493 """ 

494 match point: 

495 case None: 

496 if x is None or y is None: 496 ↛ 497line 496 didn't jump to line 497 because the condition on line 496 was never true

497 raise TypeError("Pass either a point or both x= and y= to 'apply_inverse_q'.") 

498 case XY() | YX(): 498 ↛ 502line 498 didn't jump to line 502 because the pattern on line 498 always matched

499 if x is not None or y is not None: 499 ↛ 500line 499 didn't jump to line 500 because the condition on line 499 was never true

500 raise TypeError("'apply_inverse_q' point argument is mutually exclusive with x= and y=.") 

501 x, y = point.x, point.y 

502 case _: 

503 raise TypeError(f"Unexpected positional argument type: {type(point)!r}.") 

504 xy = self.apply_inverse(x=x.to_value(self._out_frame.unit), y=y.to_value(self._out_frame.unit)) 

505 return XY(xy.x * self._in_frame.unit, xy.y * self._in_frame.unit) 

506 

507 def decompose(self) -> list[Transform[Any, Any]]: 

508 """Deconstruct a composed transform into its constituent parts. 

509 

510 Notes 

511 ----- 

512 Most transforms will just return a single-element list holding 

513 ``self``. Identity transform will return an empty list, and 

514 transforms composed with `then` will return the original transforms. 

515 Transforms constructed by `FrameSet` may or may not be decomposable. 

516 """ 

517 if not self._components: 517 ↛ 523line 517 didn't jump to line 523 because the condition on line 517 was always true

518 if self.in_frame == self._out_frame: 518 ↛ 521line 518 didn't jump to line 521 because the condition on line 518 was always true

519 return [] 

520 else: 

521 return [self] 

522 else: 

523 return list(self._components) 

524 

525 def inverted(self) -> Transform[O, I]: 

526 """Return the inverse of this transform.""" 

527 return Transform[O, I]( 

528 self._out_frame, 

529 self._in_frame, 

530 self._ast_mapping.inverted(), 

531 in_bounds=self.out_bounds, 

532 out_bounds=self.in_bounds, 

533 components=[t.inverted() for t in reversed(self._components)], 

534 ) 

535 

536 def then[F: Frame](self, next: Transform[O, F], remember_components: bool = True) -> Transform[I, F]: 

537 """Compose two transforms into another. 

538 

539 Parameters 

540 ---------- 

541 next 

542 Another transform to apply after ``self``. 

543 remember_components 

544 If `True`, the returned composed transform will remember ``self`` 

545 and ``other`` so they can be returned by `decompose`. 

546 """ 

547 if self._out_frame != next._in_frame: 

548 raise TransformCompositionError( 

549 "Cannot compose transforms that do not share a common intermediate frame: " 

550 f"{self._out_frame} != {next._in_frame}." 

551 ) 

552 components = self.decompose() + next.decompose() if remember_components else () 

553 return Transform( 

554 self._in_frame, 

555 next._out_frame, 

556 self._ast_mapping.then(next._ast_mapping), 

557 in_bounds=self.in_bounds, 

558 out_bounds=next.out_bounds, 

559 components=components, 

560 ) 

561 

562 def as_fits_wcs(self, bbox: Box) -> astropy.wcs.WCS | None: 

563 """Return a FITS WCS representation of this transform, if possible. 

564 

565 Parameters 

566 ---------- 

567 bbox 

568 Bounding box of the array the FITS WCS will describe. This 

569 transform object is assumed to work on the same coordinate system 

570 in which ``bbox`` is defined, while the FITS WCS will consider the 

571 first row and column in that box to be ``(0, 0)`` (in Astropy 

572 interfaces) or ``(1, 1)`` (in the FITS representation itself). 

573 

574 Notes 

575 ----- 

576 This method assumes the transform maps pixel coordinates to world 

577 coordinates. 

578 

579 Not all transforms can be represented exactly; when a FITS 

580 represention is not possible, `None` is returned. When the returned 

581 WCS is not `None`, it will have the same functional form, but it may 

582 not evaluate identically due to small implementation differences in 

583 the order of floating-point operations. 

584 """ 

585 ast_frame_set = self._get_ast_frame_set() 

586 _prepend_ast_shift(ast_frame_set, x=1.0 - bbox.x.start, y=1.0 - bbox.y.start, ast_domain="GRID") 

587 ast_stream = astshim.StringStream() 

588 ast_fits_chan = astshim.FitsChan( 

589 ast_stream, "Encoding=FITS-WCS, CDMatrix=1, FitsAxisOrder=<copy>, FitsTol=0.0001" 

590 ) 

591 ast_fits_chan.setFitsI("NAXIS1", bbox.x.size) 

592 ast_fits_chan.setFitsI("NAXIS2", bbox.y.size) 

593 n_writes = ast_fits_chan.write(ast_frame_set) 

594 if not n_writes: 

595 return None 

596 header = astropy.io.fits.Header(astropy.io.fits.Card.fromstring(c) for c in ast_fits_chan) 

597 try: 

598 return astropy.wcs.WCS(header) 

599 except (KeyError, astropy.wcs.InconsistentAxisTypesError): 

600 # AST wrote cards (so ``n_writes`` was nonzero) but they do not 

601 # form a valid, complete FITS WCS (e.g. a partial header with no 

602 # primary CTYPE). That means the transform is not exactly 

603 # representable as a FITS WCS, so return None instead of crashing. 

604 # This was fixed upstream in AST on DM-55992 (which also addresses 

605 # a segfault that can occur more rarely here), but we want this 

606 # workaround until that's released. 

607 return None 

608 

609 def serialize[P: pydantic.BaseModel]( 

610 self, archive: OutputArchive[P], *, use_frame_sets: bool = False 

611 ) -> TransformSerializationModel[P]: 

612 """Serialize a transform to an archive. 

613 

614 Parameters 

615 ---------- 

616 archive 

617 Archive to serialize to. 

618 use_frame_sets 

619 If `True`, decompose the transform and try to reference component 

620 mappings that were already serialized into a `FrameSet` in the 

621 archive. Note that if multiple transforms exist between a pair of 

622 frames (e.g. a `SkyProjection` and its FITS approximation), this 

623 may cause the wrong one to be saved. When this option is used, the 

624 frame set must be saved before the transform, and it must be 

625 deserialized before the transform as well. 

626 

627 Returns 

628 ------- 

629 `TransformSerializationModel` 

630 Serialized form of the transform. 

631 """ 

632 model = TransformSerializationModel[P]() 

633 if use_frame_sets: 633 ↛ 634line 633 didn't jump to line 634 because the condition on line 633 was never true

634 for link in self.decompose(): 

635 model.frames.append(link.in_frame.serialize()) 

636 model.bounds.append(link.in_bounds.serialize() if link.in_bounds is not None else None) 

637 for frame_set, pointer in archive.iter_frame_sets(): 

638 if link.in_frame in frame_set and link.out_frame in frame_set: 

639 model.mappings.append(pointer) 

640 break 

641 else: 

642 model.mappings.append(MappingSerializationModel(ast=link._ast_mapping.show())) 

643 else: 

644 model.frames.append(self.in_frame.serialize()) 

645 model.bounds.append(self.in_bounds.serialize() if self.in_bounds is not None else None) 

646 model.mappings.append(MappingSerializationModel(ast=self._ast_mapping.show())) 

647 model.frames.append(self.out_frame.serialize()) 

648 model.bounds.append(self.out_bounds.serialize() if self.out_bounds is not None else None) 

649 return model 

650 

651 @staticmethod 

652 def _get_archive_tree_type[P: pydantic.BaseModel]( 

653 pointer_type: type[P], 

654 ) -> type[TransformSerializationModel[P]]: 

655 """Return the serialization model type for this object for an archive 

656 type that uses the given pointer type. 

657 """ 

658 return TransformSerializationModel[pointer_type] # type: ignore 

659 

660 @staticmethod 

661 def from_legacy( 

662 legacy: LegacyTransform, 

663 in_frame: I, 

664 out_frame: O, 

665 in_bounds: Bounds | None = None, 

666 out_bounds: Bounds | None = None, 

667 ) -> Transform[I, O]: 

668 """Construct a transform from a legacy `lsst.afw.geom.Transform`. 

669 

670 Parameters 

671 ---------- 

672 legacy : `lsst.afw.geom.Transform` 

673 Legacy transform object. 

674 in_frame 

675 Coordinate frame for input points to the forward transform. 

676 out_frame 

677 Coordinate frame for output points from the forward transform. 

678 in_bounds 

679 The region that bounds valid input points. 

680 out_bounds 

681 The region that bounds valid output points. 

682 """ 

683 return Transform( 

684 in_frame, 

685 out_frame, 

686 legacy.getMapping(), 

687 in_bounds=in_bounds, 

688 out_bounds=out_bounds, 

689 ) 

690 

691 def to_legacy(self) -> LegacyTransform: 

692 """Convert to a legacy `lsst.afw.geom.TransformPoint2ToPoint2` 

693 instance. 

694 """ 

695 from lsst.afw.geom import TransformPoint2ToPoint2 as LegacyTransform 

696 

697 return LegacyTransform(self._ast_mapping, False) 

698 

699 def _get_ast_frame_set(self) -> Any: 

700 ast_frame_set = astshim.FrameSet(_make_ast_frame(self._in_frame)) 

701 ast_frame_set.addFrame(astshim.FrameSet.BASE, self._ast_mapping, _make_ast_frame(self._out_frame)) 

702 return ast_frame_set 

703 

704 

705def _ast_apply(method: Any, *, x: Any, y: Any) -> XY[float] | XY[np.ndarray]: 

706 # TODO: add bounds argument and check inputs 

707 xa = np.asarray(x) 

708 ya = np.asarray(y) 

709 broadcast_shape = np.broadcast(xa, ya).shape 

710 scalar = not broadcast_shape 

711 xb, yb = np.broadcast_arrays(xa, ya) 

712 xy_in = np.vstack([xb.ravel(), yb.ravel()]).astype(np.float64) 

713 xy_out = method(xy_in) 

714 if scalar: 

715 return XY(float(xy_out[0, 0]), float(xy_out[1, 0])) 

716 return XY(xy_out[0].reshape(broadcast_shape), xy_out[1].reshape(broadcast_shape)) 

717 

718 

719def _prepend_ast_shift(ast_frame_set: Any, x: float, y: float, ast_domain: str) -> None: 

720 ast_output_frame_id = ast_frame_set.current 

721 ast_frame_set.addFrame( 

722 astshim.FrameSet.BASE, 

723 astshim.ShiftMap([x, y]), 

724 astshim.Frame(2, f"Domain={ast_domain}"), 

725 ) 

726 ast_frame_set.base = ast_frame_set.current 

727 ast_frame_set.current = ast_output_frame_id 

728 

729 

730def _make_ast_frame(frame: Frame) -> Any: 

731 if frame is SkyFrame.ICRS: 

732 return astshim.SkyFrame("") 

733 ast_frame = astshim.Frame(2, f"Ident={frame._ast_ident}") 

734 if frame.unit is not None: 734 ↛ 738line 734 didn't jump to line 738 because the condition on line 734 was always true

735 fits_unit = frame.unit.to_string(format="fits") 

736 ast_frame.setUnit(1, fits_unit) 

737 ast_frame.setUnit(2, fits_unit) 

738 ast_frame.setLabel(1, "x") 

739 ast_frame.setLabel(2, "y") 

740 return ast_frame 

741 

742 

743def _standardize_xy(xy: XY[Any], frame: Frame) -> XY[Any]: 

744 return XY(x=frame.standardize_x(xy.x), y=frame.standardize_y(xy.y)) 

745 

746 

747class MappingSerializationModel(pydantic.BaseModel): 

748 """Serialization model for an AST Mapping.""" 

749 

750 ast: str = pydantic.Field(description="A serialized Starlink AST Mapping, using the AST native encoding.") 

751 

752 

753class TransformSerializationModel[P: pydantic.BaseModel](ArchiveTree): 

754 """Serialization model for coordinate transforms.""" 

755 

756 SCHEMA_NAME: ClassVar[str] = "transform" 

757 SCHEMA_VERSION: ClassVar[str] = "1.0.0" 

758 MIN_READ_VERSION: ClassVar[int] = 1 

759 PUBLIC_TYPE: ClassVar[type] = Transform 

760 

761 frames: list[SerializableFrame] = pydantic.Field( 

762 default_factory=list, 

763 description=textwrap.dedent( 

764 """ 

765 List of frames that this transform passes through. 

766 

767 All transforms include at least two frames (the endpoints). Others 

768 intermediate frames may be included to facilitate data-sharing 

769 between transforms. 

770 """ 

771 ), 

772 ) 

773 

774 bounds: list[BoundsSerializationModel | None] = pydantic.Field( 

775 default_factory=list, 

776 description=textwrap.dedent( 

777 """ 

778 List of the bounds of the ``frames`` for this transform. 

779 

780 This always has the same number of elements as ``frames``. 

781 """ 

782 ), 

783 ) 

784 

785 mappings: list[P | MappingSerializationModel] = pydantic.Field( 

786 default_factory=list, 

787 description=textwrap.dedent( 

788 """ 

789 The actual mappings between frames, or archive pointers to 

790 serialized FrameSet objects from which they can be obtained. 

791 

792 This always has one fewer element than ``frames``. 

793 """ 

794 ), 

795 ) 

796 

797 def deserialize(self, archive: InputArchive[P], **kwargs: Any) -> Transform[Any, Any]: 

798 """Deserialize a transform from an archive. 

799 

800 Parameters 

801 ---------- 

802 archive 

803 Archive to read from. 

804 **kwargs 

805 Unsupported keyword arguments are accepted only to provide better 

806 error messages (raising `serialization.InvalidParameterError`). 

807 """ 

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

809 raise InvalidParameterError(f"Unrecognized parameters for Transform: {set(kwargs.keys())}.") 

810 if len(self.frames) != len(self.bounds): 810 ↛ 811line 810 didn't jump to line 811 because the condition on line 810 was never true

811 raise ArchiveReadError( 

812 f"Inconsistent lengths for 'frames' ({len(self.frames)}) and 'bounds' ({len(self.bounds)})." 

813 ) 

814 if len(self.frames) != len(self.mappings) + 1: 814 ↛ 815line 814 didn't jump to line 815 because the condition on line 814 was never true

815 raise ArchiveReadError( 

816 f"Inconsistent lengths for 'frames' ({len(self.frames)}) and " 

817 f"'mappings' ({len(self.mappings)}; should be one less)." 

818 ) 

819 # We can't just compose onto an identity Transform if we want to 

820 # preserve the FrameSet-ness of any of these mappings. 

821 transform: Transform | None = None 

822 for n, mapping in enumerate(self.mappings): 

823 match mapping: 

824 case MappingSerializationModel(ast=serialized_mapping): 824 ↛ 835line 824 didn't jump to line 835 because the pattern on line 824 always matched

825 ast_mapping = astshim.Mapping.fromString(serialized_mapping) 

826 in_bounds = self.bounds[n] 

827 out_bounds = self.bounds[n + 1] 

828 new_transform = Transform( 

829 self.frames[n].deserialize(), 

830 self.frames[n + 1].deserialize(), 

831 ast_mapping, 

832 in_bounds.deserialize() if in_bounds is not None else None, 

833 out_bounds.deserialize() if out_bounds is not None else None, 

834 ) 

835 case reference: 

836 frame_set = archive.get_frame_set(reference) 

837 new_transform = frame_set[self.frames[n].deserialize(), self.frames[n + 1].deserialize()] 

838 if transform is None: 838 ↛ 841line 838 didn't jump to line 841 because the condition on line 838 was always true

839 transform = new_transform 

840 else: 

841 transform = transform.then(new_transform) 

842 if transform is None: 842 ↛ 843line 842 didn't jump to line 843 because the condition on line 842 was never true

843 transform = Transform.identity(self.frames[0].deserialize()) 

844 return transform 

845 

846 

847if TYPE_CHECKING: 

848 

849 def _test_types() -> None: 

850 t = cast(Transform, None) 

851 arr = np.zeros(3) 

852 

853 # Scalar inputs → XY[float] 

854 assert_type(t.apply_forward(x=1.0, y=2.0), XY[float]) 

855 assert_type(t.apply_inverse(x=1.0, y=2.0), XY[float]) 

856 

857 # Array inputs → XY[np.ndarray] 

858 assert_type(t.apply_forward(x=arr, y=arr), XY[np.ndarray]) 

859 assert_type(t.apply_inverse(x=arr, y=arr), XY[np.ndarray]) 

860 

861 # Array-like (list) inputs → XY[np.ndarray] 

862 assert_type(t.apply_forward(x=[1.0, 2.0], y=[3.0, 4.0]), XY[np.ndarray]) 

863 assert_type(t.apply_inverse(x=[1.0, 2.0], y=[3.0, 4.0]), XY[np.ndarray])