From 1055133df3b64df713805586102adb35b93837be Mon Sep 17 00:00:00 2001 From: AdamField118 Date: Mon, 27 Jul 2026 12:26:50 -0400 Subject: [PATCH 1/2] Add DES_PSFEx (PSFEx model reader) to jax_galsim.des Port galsim.des.DES_PSFEx so an empirical PSFEx model (*_psfcat.psf) can be read and evaluated entirely within JAX-GalSim, removing the need to fall back to reference GalSim just to build the PSF. The PSFEx file read (read()) is host-side FITS I/O, unchanged from galsim. getPSFArray -- the per-position evaluation of the interpolated PSF -- is reimplemented in JAX (the [1, x, x^2, ...] powers are built with a cumulative product instead of galsim's in-place np.empty loop, which JAX forbids), so it can be jitted, vmapped over image positions, and differentiated with respect to the image position. getPSF wraps the array in an InterpolatedImage exactly as galsim does (Lanczos(3), scale = PSF_SAMP), and applies the WCS if given. The GalSim config framework registration (des_psfex input type / DES_PSFEx object type) is intentionally not ported, since JAX-GalSim has no config system; this is noted in the lax_description. Tests compare getPSFArray and the drawn effective-PSF image against reference GalSim (to single-precision, since PSFEx bases are float32) using the existing DECam_00154912_12_psfcat.psf test file, and check that getPSFArray is jittable, vmappable, and differentiable in the image position. --- jax_galsim/__init__.py | 1 + jax_galsim/des/__init__.py | 1 + jax_galsim/des/des_psfex.py | 173 ++++++++++++++++++++++++++++++++ tests/jax/test_des_psfex_jax.py | 99 ++++++++++++++++++ 4 files changed, 274 insertions(+) create mode 100644 jax_galsim/des/__init__.py create mode 100644 jax_galsim/des/des_psfex.py create mode 100644 tests/jax/test_des_psfex_jax.py diff --git a/jax_galsim/__init__.py b/jax_galsim/__init__.py index 80cbe04b..d548c48f 100644 --- a/jax_galsim/__init__.py +++ b/jax_galsim/__init__.py @@ -102,6 +102,7 @@ from . import bessel from . import fits from . import integ +from . import des # this one is specific to jax_galsim from . import core diff --git a/jax_galsim/des/__init__.py b/jax_galsim/des/__init__.py new file mode 100644 index 00000000..4fc9e9b5 --- /dev/null +++ b/jax_galsim/des/__init__.py @@ -0,0 +1 @@ +from .des_psfex import DES_PSFEx diff --git a/jax_galsim/des/des_psfex.py b/jax_galsim/des/des_psfex.py new file mode 100644 index 00000000..c9cf6b24 --- /dev/null +++ b/jax_galsim/des/des_psfex.py @@ -0,0 +1,173 @@ +# This is a JAX port of galsim.des.des_psfex (galsim/des/des_psfex.py). +# The reading of the PSFEx file is unchanged host-side I/O; the per-position +# PSF evaluation (getPSFArray) is reimplemented in JAX so it can be jitted, +# vmapped, and differentiated with respect to the image position. +import os + +import galsim as _galsim +import galsim.des # noqa: F401 (populates _galsim.des for @implements below) +import jax.numpy as jnp +import numpy as np + +from jax_galsim._pyfits import pyfits +from jax_galsim.core.utils import implements +from jax_galsim.errors import GalSimIncompatibleValuesError +from jax_galsim.fits import FitsHeader +from jax_galsim.image import Image +from jax_galsim.interpolant import Lanczos +from jax_galsim.interpolatedimage import InterpolatedImage +from jax_galsim.wcs import readFromFitsHeader + +LAX_DES_PSFEX = """\ +The JAX-GalSim version of ``DES_PSFEx`` does not register itself with the +GalSim config framework (the ``des_psfex`` input type and ``DES_PSFEx`` object +type are not available), since JAX-GalSim does not implement config processing. +The ``getPSFArray`` method is implemented in JAX, so it may be jitted, vmapped, +and differentiated with respect to the ``image_pos`` coordinates. +""" + + +@implements(_galsim.des.DES_PSFEx, lax_description=LAX_DES_PSFEX, module="galsim.des") +class DES_PSFEx: + _req_params = {"file_name": str} + _opt_params = {"dir": str, "image_file_name": str} + _single_params = [] + _takes_rng = False + + def __init__(self, file_name, image_file_name=None, wcs=None, dir=None): + if dir: + if not isinstance(file_name, str): + raise TypeError("file_name must be a string") + file_name = os.path.join(dir, file_name) + if image_file_name is not None: + image_file_name = os.path.join(dir, image_file_name) + self.file_name = file_name + if image_file_name: + if wcs is not None: + raise GalSimIncompatibleValuesError( + "Cannot provide both image_file_name and wcs", + image_file_name=image_file_name, + wcs=wcs, + ) + header = FitsHeader(file_name=image_file_name) + wcs, origin = readFromFitsHeader(header) + self.wcs = wcs + elif wcs: + self.wcs = wcs + else: + self.wcs = None + self.read() + + def read(self): + if isinstance(self.file_name, str): + hdu_list = pyfits.open(self.file_name) + hdu = hdu_list[1] + else: + hdu = self.file_name + hdu_list = None + pol_naxis = hdu.header["POLNAXIS"] + + pol_name1 = hdu.header["POLNAME1"] + pol_name2 = hdu.header["POLNAME2"] + + pol_zero1 = hdu.header["POLZERO1"] + pol_zero2 = hdu.header["POLZERO2"] + pol_scal1 = hdu.header["POLSCAL1"] + pol_scal2 = hdu.header["POLSCAL2"] + + pol_ngrp = hdu.header["POLNGRP"] + pol_group1 = hdu.header["POLGRP1"] + pol_group2 = hdu.header["POLGRP2"] + pol_deg = hdu.header["POLDEG1"] + + psf_naxis = hdu.header["PSFNAXIS"] + psf_axis1 = hdu.header["PSFAXIS1"] + psf_axis2 = hdu.header["PSFAXIS2"] + psf_axis3 = hdu.header["PSFAXIS3"] + psf_samp = hdu.header["PSF_SAMP"] + + basis = hdu.data.field("PSF_MASK")[0] + + if hdu_list: + hdu_list.close() + + try: + assert pol_naxis == 2 + assert pol_name1.startswith("X") and pol_name1.endswith("IMAGE") + assert pol_name2.startswith("Y") and pol_name2.endswith("IMAGE") + assert pol_ngrp == 1 + assert pol_group1 == 1 + assert pol_group2 == 1 + assert psf_naxis == 3 + assert psf_axis3 == ((pol_deg + 1) * (pol_deg + 2)) // 2 + assert basis.shape[0] == psf_axis3 + assert basis.shape[1] == psf_axis2 + assert basis.shape[2] == psf_axis1 + except AssertionError as e: + raise OSError("PSFEx file %s is not as expected.\n%r" % (self.file_name, e)) + + # Static configuration is kept as plain Python/NumPy; only the basis + # cube (which is combined with traced coefficients) is moved onto a JAX + # array so getPSFArray traces cleanly. PSFEx stores the cube as + # big-endian float32, which JAX will not accept, so cast to native + # float32 first (galsim casts the combined array to float32 anyway). + self.basis = jnp.asarray(np.ascontiguousarray(basis, dtype=np.float32)) + self.fit_order = int(pol_deg) + self.fit_size = int(psf_axis3) + self.x_zero = pol_zero1 + self.y_zero = pol_zero2 + self.x_scale = pol_scal1 + self.y_scale = pol_scal2 + self.sample_scale = psf_samp + + @implements(_galsim.des.DES_PSFEx.getSampleScale) + def getSampleScale(self): + return self.sample_scale + + @implements(_galsim.des.DES_PSFEx.getLocalWCS) + def getLocalWCS(self, image_pos): + if self.wcs: + return self.wcs.local(image_pos) + else: + return None + + @implements(_galsim.des.DES_PSFEx.getPSF) + def getPSF(self, image_pos, gsparams=None): + im = Image(self.getPSFArray(image_pos)) + psf = InterpolatedImage( + im, + scale=self.sample_scale, + flux=1, + x_interpolant=Lanczos(3), + gsparams=gsparams, + ) + if self.wcs: + psf = self.wcs.toWorld(psf, image_pos=image_pos) + return psf + + @implements(_galsim.des.DES_PSFEx.getPSFArray) + def getPSFArray(self, image_pos): + xto = self._powers((image_pos.x - self.x_zero) / self.x_scale) + yto = self._powers((image_pos.y - self.y_zero) / self.y_scale) + order = self.fit_order + # order is a static Python int, so this comprehension is unrolled at + # trace time; it mirrors galsim's ordering of the polynomial terms. + P = jnp.stack( + [ + xto[nx] * yto[ny] + for ny in range(order + 1) + for nx in range(order + 1 - ny) + ] + ) + return jnp.tensordot(P, self.basis, (0, 0)).astype(jnp.float32) + + def _powers(self, x): + # JAX-safe replacement for galsim's ``np.empty`` + in-place loop: build + # [1, x, x**2, ..., x**order] via a cumulative product (same recurrence + # as galsim, but without an in-place update, which JAX forbids). + return jnp.concatenate( + [ + jnp.ones((1,), dtype=jnp.result_type(float)), + jnp.cumprod(jnp.full((self.fit_order,), x)), + ] + ) diff --git a/tests/jax/test_des_psfex_jax.py b/tests/jax/test_des_psfex_jax.py new file mode 100644 index 00000000..48ef61a0 --- /dev/null +++ b/tests/jax/test_des_psfex_jax.py @@ -0,0 +1,99 @@ +import os + +import galsim as _galsim +import jax +import jax.numpy as jnp +import numpy as np +import pytest +from galsim.utilities import timer + +import jax_galsim as galsim + +DES_DATA_DIR = os.path.join( + os.path.dirname(__file__), "..", "GalSim", "tests", "des_data" +) +PSFEX_FILE = "DECam_00154912_12_psfcat.psf" + +# A few positions spread across the DECam chip. +POSITIONS = [(100.0, 100.0), (456.0, 789.0), (1024.0, 2048.0), (1700.0, 3500.0)] + + +def _have_des_data(): + return os.path.isfile(os.path.join(DES_DATA_DIR, PSFEX_FILE)) + + +requires_des_data = pytest.mark.skipif( + not _have_des_data(), + reason="DES test data (tests/GalSim submodule) not available", +) + + +@requires_des_data +@timer +def test_des_psfex_getPSFArray_vs_galsim(): + """The interpolated PSF array should match reference GalSim.""" + ref = _galsim.des.DES_PSFEx(PSFEX_FILE, dir=DES_DATA_DIR) + jgs = galsim.des.DES_PSFEx(PSFEX_FILE, dir=DES_DATA_DIR) + + assert jgs.fit_order == ref.fit_order + assert jgs.fit_size == ref.fit_size + np.testing.assert_allclose(jgs.sample_scale, ref.sample_scale) + + for x, y in POSITIONS: + a = np.asarray(jgs.getPSFArray(galsim.PositionD(x, y))) + b = ref.getPSFArray(_galsim.PositionD(x, y)) + # float32 interpolation, so compare at ~single precision. + np.testing.assert_allclose(a, b, rtol=0, atol=1e-6) + + +@requires_des_data +@timer +def test_des_psfex_getPSF_drawn_image_vs_galsim(): + """The effective-PSF image (drawn with no_pixel) should match GalSim.""" + ref = _galsim.des.DES_PSFEx(PSFEX_FILE, dir=DES_DATA_DIR) + jgs = galsim.des.DES_PSFEx(PSFEX_FILE, dir=DES_DATA_DIR) + + for x, y in POSITIONS: + # PSFEx PSFs already include the pixel, so draw with method='no_pixel'. + gimg = ref.getPSF(_galsim.PositionD(x, y)).drawImage( + nx=25, ny=25, scale=0.2, method="no_pixel" + ) + jimg = jgs.getPSF(galsim.PositionD(x, y)).drawImage( + nx=25, ny=25, scale=0.2, method="no_pixel" + ) + np.testing.assert_allclose( + np.asarray(jimg.array), gimg.array, rtol=0, atol=1e-6 + ) + + +@requires_des_data +@timer +def test_des_psfex_is_jittable_vmappable_differentiable(): + """getPSFArray should support jit, vmap, and grad over the image position.""" + jgs = galsim.des.DES_PSFEx(PSFEX_FILE, dir=DES_DATA_DIR) + + def psf_sum(x, y): + return jnp.sum(jgs.getPSFArray(galsim.PositionD(x, y))) + + # jit + jitted = jax.jit(lambda x, y: jgs.getPSFArray(galsim.PositionD(x, y))) + arr = jitted(456.0, 789.0) + ref = jgs.getPSFArray(galsim.PositionD(456.0, 789.0)) + np.testing.assert_allclose(np.asarray(arr), np.asarray(ref), rtol=0, atol=1e-6) + + # vmap over a batch of positions + xs = jnp.array([p[0] for p in POSITIONS]) + ys = jnp.array([p[1] for p in POSITIONS]) + batched = jax.vmap(lambda x, y: jgs.getPSFArray(galsim.PositionD(x, y)))(xs, ys) + assert batched.shape[0] == len(POSITIONS) + + # grad w.r.t. position must be finite (the PSF is differentiable in position) + gx = jax.grad(psf_sum, argnums=0)(456.0, 789.0) + assert np.isfinite(float(gx)) + + +if __name__ == "__main__": + test_des_psfex_getPSFArray_vs_galsim() + test_des_psfex_getPSF_drawn_image_vs_galsim() + test_des_psfex_is_jittable_vmappable_differentiable() + print("all DES_PSFEx tests passed") From 5d5441e723924d00c1d72159212f3f2679f68f3c Mon Sep 17 00:00:00 2001 From: AdamField118 Date: Tue, 28 Jul 2026 16:16:11 -0400 Subject: [PATCH 2/2] pytree, wcs trace, basis numpy --- jax_galsim/des/des_psfex.py | 72 ++++++++++++++++++++++++++++----- tests/jax/test_des_psfex_jax.py | 29 +++++++++++++ 2 files changed, 91 insertions(+), 10 deletions(-) diff --git a/jax_galsim/des/des_psfex.py b/jax_galsim/des/des_psfex.py index c9cf6b24..5678ebdd 100644 --- a/jax_galsim/des/des_psfex.py +++ b/jax_galsim/des/des_psfex.py @@ -8,9 +8,10 @@ import galsim.des # noqa: F401 (populates _galsim.des for @implements below) import jax.numpy as jnp import numpy as np +from jax.tree_util import register_pytree_node_class from jax_galsim._pyfits import pyfits -from jax_galsim.core.utils import implements +from jax_galsim.core.utils import ensure_hashable, implements from jax_galsim.errors import GalSimIncompatibleValuesError from jax_galsim.fits import FitsHeader from jax_galsim.image import Image @@ -22,12 +23,17 @@ The JAX-GalSim version of ``DES_PSFEx`` does not register itself with the GalSim config framework (the ``des_psfex`` input type and ``DES_PSFEx`` object type are not available), since JAX-GalSim does not implement config processing. -The ``getPSFArray`` method is implemented in JAX, so it may be jitted, vmapped, -and differentiated with respect to the ``image_pos`` coordinates. + +As a PyTree, all of the data read from the PSFEx file (the PCA basis and the +polynomial fit parameters) is static auxiliary data stored as NumPy arrays, +and only the ``wcs`` is traced. Autodiff is therefore supported with respect +to the ``image_pos`` argument (via ``getPSFArray``) and the ``wcs``, but not +with respect to the fixed PSFEx calibration data. """ @implements(_galsim.des.DES_PSFEx, lax_description=LAX_DES_PSFEX, module="galsim.des") +@register_pytree_node_class class DES_PSFEx: _req_params = {"file_name": str} _opt_params = {"dir": str, "image_file_name": str} @@ -106,12 +112,11 @@ def read(self): except AssertionError as e: raise OSError("PSFEx file %s is not as expected.\n%r" % (self.file_name, e)) - # Static configuration is kept as plain Python/NumPy; only the basis - # cube (which is combined with traced coefficients) is moved onto a JAX - # array so getPSFArray traces cleanly. PSFEx stores the cube as - # big-endian float32, which JAX will not accept, so cast to native - # float32 first (galsim casts the combined array to float32 anyway). - self.basis = jnp.asarray(np.ascontiguousarray(basis, dtype=np.float32)) + # All data read from the PSFEx file is fixed calibration and is kept as + # static NumPy (it becomes auxiliary PyTree data, not traced leaves). + # PSFEx stores the cube as big-endian float32; cast to native float32 + # (galsim casts the combined array to float32 anyway). + self.basis = np.ascontiguousarray(basis, dtype=np.float32) self.fit_order = int(pol_deg) self.fit_size = int(psf_axis3) self.x_zero = pol_zero1 @@ -159,7 +164,8 @@ def getPSFArray(self, image_pos): for nx in range(order + 1 - ny) ] ) - return jnp.tensordot(P, self.basis, (0, 0)).astype(jnp.float32) + # basis is static NumPy; jnp folds it in as a compile-time constant. + return jnp.tensordot(P, jnp.asarray(self.basis), (0, 0)).astype(jnp.float32) def _powers(self, x): # JAX-safe replacement for galsim's ``np.empty`` + in-place loop: build @@ -171,3 +177,49 @@ def _powers(self, x): jnp.cumprod(jnp.full((self.fit_order,), x)), ] ) + + def tree_flatten(self): + """Flatten into traced children and static auxiliary data. + + Only ``wcs`` is traced. All of the data read from the PSFEx file is + auxiliary; the basis array is passed through ``ensure_hashable`` so the + PyTree metadata stays hashable when instances are used as arguments to + transformed functions. + """ + children = (self.wcs,) + aux_data = { + "file_name": self.file_name, + "basis": ensure_hashable(jnp.asarray(self.basis)), + "basis_shape": self.basis.shape, + "fit_order": self.fit_order, + "fit_size": self.fit_size, + "x_zero": ensure_hashable(self.x_zero), + "y_zero": ensure_hashable(self.y_zero), + "x_scale": ensure_hashable(self.x_scale), + "y_scale": ensure_hashable(self.y_scale), + "sample_scale": ensure_hashable(self.sample_scale), + } + return children, aux_data + + @classmethod + def tree_unflatten(cls, aux_data, children): + """Rebuild an instance without re-reading the file. + + ``__init__`` opens the PSFEx file, so (following ``CelestialCoord`` / + ``Image``) we construct via ``object.__new__`` and restore attributes + directly from the flattened representation. + """ + obj = object.__new__(cls) + (obj.wcs,) = children + obj.file_name = aux_data["file_name"] + obj.basis = np.asarray(aux_data["basis"], dtype=np.float32).reshape( + aux_data["basis_shape"] + ) + obj.fit_order = aux_data["fit_order"] + obj.fit_size = aux_data["fit_size"] + obj.x_zero = aux_data["x_zero"] + obj.y_zero = aux_data["y_zero"] + obj.x_scale = aux_data["x_scale"] + obj.y_scale = aux_data["y_scale"] + obj.sample_scale = aux_data["sample_scale"] + return obj diff --git a/tests/jax/test_des_psfex_jax.py b/tests/jax/test_des_psfex_jax.py index 48ef61a0..c60cd089 100644 --- a/tests/jax/test_des_psfex_jax.py +++ b/tests/jax/test_des_psfex_jax.py @@ -66,6 +66,34 @@ def test_des_psfex_getPSF_drawn_image_vs_galsim(): ) +@requires_des_data +@timer +def test_des_psfex_pytree_roundtrip_and_traced_arg(): + """DES_PSFEx is a registered PyTree: it round-trips through flatten/ + unflatten (without re-reading the file) and can be passed as an argument to + a transformed function, including two distinct-but-equal instances.""" + jgs = galsim.des.DES_PSFEx(PSFEX_FILE, dir=DES_DATA_DIR) + + leaves, treedef = jax.tree_util.tree_flatten(jgs) + rebuilt = jax.tree_util.tree_unflatten(treedef, leaves) + np.testing.assert_array_equal(np.asarray(rebuilt.basis), np.asarray(jgs.basis)) + for x, y in POSITIONS: + np.testing.assert_allclose( + np.asarray(rebuilt.getPSFArray(galsim.PositionD(x, y))), + np.asarray(jgs.getPSFArray(galsim.PositionD(x, y))), + rtol=0, + atol=1e-6, + ) + + # Pass the object itself as a jitted argument. Using two distinct instances + # exercises the hashability of the (auxiliary) PSFEx data in the treedef. + f = jax.jit(lambda obj, x, y: obj.getPSFArray(galsim.PositionD(x, y))) + jgs2 = galsim.des.DES_PSFEx(PSFEX_FILE, dir=DES_DATA_DIR) + a1 = f(jgs, 456.0, 789.0) + a2 = f(jgs2, 456.0, 789.0) + np.testing.assert_allclose(np.asarray(a1), np.asarray(a2), rtol=0, atol=1e-6) + + @requires_des_data @timer def test_des_psfex_is_jittable_vmappable_differentiable(): @@ -95,5 +123,6 @@ def psf_sum(x, y): if __name__ == "__main__": test_des_psfex_getPSFArray_vs_galsim() test_des_psfex_getPSF_drawn_image_vs_galsim() + test_des_psfex_pytree_roundtrip_and_traced_arg() test_des_psfex_is_jittable_vmappable_differentiable() print("all DES_PSFEx tests passed")