Source code for optixstuff.disperser
"""Disperser hardware descriptors for integral field spectrographs.
optixstuff owns the descriptor: the interface plus a cheap closed-form
scalar/ETC face. The heavy render logic (building the forward operator) lives
in coronachrome. This mirrors the coronagraph split, so jaxedith and yield
tools can read IFS hardware info without importing the render engine.
"""
import abc
import equinox as eqx
import jax.numpy as jnp
from jax import Array
from optixstuff.optical_elements import AbstractOpticalElement, ConstantThroughput
[docs]
class AbstractDisperser(eqx.Module):
"""Interface for a dispersing IFS element (lenslet array, slicer, MSA).
Only the scalar/ETC face is defined here. Render geometry lives in
coronachrome and dispatches on the concrete descriptor type.
"""
[docs]
@abc.abstractmethod
def spectral_resolution(self, wavelength_nm):
"""Resolving power R = lambda / dlambda at the given wavelength."""
[docs]
@abc.abstractmethod
def spectral_sampling(self):
"""Detector pixels per resolution element."""
[docs]
@abc.abstractmethod
def n_pix_spread(self, wavelength_min_nm, wavelength_max_nm):
"""Detector pixels a single spaxel spectrum spans across a band."""
[docs]
@abc.abstractmethod
def throughput(self, wavelength_nm):
"""Disperser optical throughput in [0, 1] at the given wavelength."""
[docs]
def _polyval_deriv(coeffs, x):
"""Evaluate the derivative of a descending-order polynomial at x."""
n = coeffs.shape[0]
if n <= 1:
return jnp.zeros_like(jnp.asarray(x, dtype=float))
powers = jnp.arange(n - 1, 0, -1)
return jnp.polyval(coeffs[:-1] * powers, x)
[docs]
class LensletDisperser(AbstractDisperser):
"""Lenslet-array IFS disperser (CRISPY heritage).
Config only. The render geometry (IR build) is performed by coronachrome,
which reads these fields. Scalar/ETC methods derive from
``dispersion_coeffs`` + ``pix_per_reselt`` so the dispersion model is the
single source of truth.
``psflet_params[0]`` is the PSFlet core width in detector pixels (Gaussian
``sigma`` or Moffat ``alpha``); any trailing entries are dimensionless shape
parameters (e.g. Moffat ``beta``). ``psflet_ref_nm`` is the wavelength at
which that core width is specified: a diffraction-limited spot scales as
``lambda f / D``, so at fixed pixel scale coronachrome scales the core width
by ``lambda / psflet_ref_nm`` per wavelength (the shape parameters do not
scale).
``sky_pitch_arcsec`` is the lenslet pitch projected on sky, a plain angle
in arcseconds (the focal-plane cube a render layer consumes lives on a
fixed angular grid, so no reference wavelength is involved). It makes the
spatial sampling an instrument property: the render layer derives
focal-plane pixels per lenslet as the ratio of this pitch to the cube's
angular plate scale instead of taking a free knob. ``None`` means
unspecified.
``psflet_pack_path`` is a reference to a frozen PSFlet template pack file
(used with ``psflet_kind="template"``), the way ``throughput_element``
carries the blaze curve: config only, the render layer loads it.
"""
pitch_m: float
pixsize_m: float
angle_rad: float
lam_ref_nm: float
pix_per_reselt: float
dispersion_coeffs: Array = eqx.field(converter=jnp.asarray)
psflet_params: Array = eqx.field(converter=jnp.asarray)
psflet_ref_nm: float
grid_kind: str = eqx.field(static=True)
n_lenslets: int = eqx.field(static=True)
psflet_kind: str = eqx.field(static=True)
detector_shape: tuple[int, int] = eqx.field(static=True)
throughput_element: AbstractOpticalElement = ConstantThroughput(1.0)
sky_pitch_arcsec: float | None = None
psflet_pack_path: str | None = eqx.field(static=True, default=None)
[docs]
def _dispersion_px(self, wavelength_nm):
"""Spectral-axis detector offset [px] for the wavelength(s)."""
u = jnp.log(jnp.asarray(wavelength_nm, dtype=float) / self.lam_ref_nm)
return jnp.polyval(self.dispersion_coeffs, u)
[docs]
def spectral_resolution(self, wavelength_nm):
"""R = (local px per unit log-lambda) / pixels-per-resolution-element."""
u = jnp.log(jnp.asarray(wavelength_nm, dtype=float) / self.lam_ref_nm)
local = jnp.abs(_polyval_deriv(self.dispersion_coeffs, u))
return local / self.pix_per_reselt
[docs]
def spectral_sampling(self):
"""Detector pixels per resolution element."""
return self.pix_per_reselt
[docs]
def n_pix_spread(self, wavelength_min_nm, wavelength_max_nm):
"""Spectral trace length [px] across a band, plus a PSFlet-width margin."""
span = jnp.abs(
self._dispersion_px(wavelength_max_nm)
- self._dispersion_px(wavelength_min_nm)
)
return span + self.psflet_params[0]
[docs]
def throughput(self, wavelength_nm):
"""Disperser optical throughput in [0, 1], shaped like wavelength_nm.
Delegates to the composed throughput element (ConstantThroughput by
default; SpectralThroughput for a tabulated blaze/transmission curve) and
broadcasts to the wavelength shape so the output shape is canonical
regardless of the backing element.
"""
w = jnp.asarray(wavelength_nm, dtype=float)
return self.throughput_element.get_throughput(w) * jnp.ones_like(w)