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)