Skip to content
1 change: 1 addition & 0 deletions changelog/76.doc.rst
Original file line number Diff line number Diff line change
@@ -0,0 +1 @@
Updated docstrings and type annotations in `sunkit_dem.GenericModel`.
61 changes: 42 additions & 19 deletions sunkit_dem/base_model.py
Original file line number Diff line number Diff line change
@@ -1,5 +1,5 @@
"""
Base model class for DEM models
Base model class for differential emission measure models.
"""
from abc import ABC, abstractmethod

Expand All @@ -25,8 +25,8 @@ def defines_model_for(self):


class GenericModel(BaseModel):
"""
Base class for implementing a differential emission measure model
r"""
Base class for implementing a differential emission measure (DEM) model.

Parameters
----------
Expand All @@ -37,7 +37,7 @@ class GenericModel(BaseModel):
temperature_bin_edges : `~astropy.units.Quantity`
Edges of the temperature bins in which the DEM is computed. The
rightmost edge is included. The kernel is evaluated at the bin centers.
The bin widths must be equal in log10.
The bin widths must be equal in :math:`\log_{10}` space.
"""

_registry = dict()
Expand All @@ -56,7 +56,14 @@ def __init_subclass__(cls, **kwargs):
cls._registry[cls] = cls.defines_model_for

@u.quantity_input
def __init__(self, data, kernel, temperature_bin_edges: u.K, kernel_temperatures=None, **kwargs):
def __init__(
self,
data: ndcube.NDCollection,
kernel: dict[str, u.Quantity],
temperature_bin_edges: u.Quantity[u.K],
kernel_temperatures=None,
**kwargs,
):
self.temperature_bin_edges = temperature_bin_edges
self.data = data
self.kernel_temperatures = kernel_temperatures
Expand All @@ -65,27 +72,43 @@ def __init__(self, data, kernel, temperature_bin_edges: u.K, kernel_temperatures
self.kernel = kernel

@property
def _keys(self):
def _keys(self) -> list[str]:
# Internal reference for entries in kernel and data
# This ensures consistent ordering in kernel and data matrices
return sorted(list(self.kernel.keys()))

@property
@u.quantity_input
def temperature_bin_centers(self) -> u.K:
def temperature_bin_centers(self) -> u.Quantity[u.K]:
r"""
The temperature at the midpoint of each temperature bin.

Notes
-----
The center of each temperature bin is calculated in physical
space, not in :math:`log` space.
"""
return (self.temperature_bin_edges[1:] + self.temperature_bin_edges[:-1])/2

@property
@u.quantity_input
def temperature_bin_widths(self) -> u.K:
def temperature_bin_widths(self) -> u.Quantity[u.K]:
r"""
The width of each temperature bin.

Notes
-----
The widths of each bin are calculated in physical space, not in
:math:`\log` space.
"""
return np.diff(self.temperature_bin_edges)

@property
def data(self) -> ndcube.NDCollection:
return self._data

@data.setter
def data(self, data):
def data(self, data: ndcube.NDCollection):
"""
Check that input data is correctly formatted as an
`ndcube.NDCollection`
Expand All @@ -99,8 +122,8 @@ def data(self, data):
@property
def combined_mask(self):
"""
Combined mask of all members of ``data``. Will be True if any member is masked.
This is propagated to the final DEM result
Combined mask of all members of ``data``. Will be `True` if any member is masked.
This mask is propagated to the final DEM result.
"""
combined_mask = []
for k in self._keys:
Expand All @@ -111,33 +134,33 @@ def combined_mask(self):
return np.any(combined_mask, axis=0)

@property
def kernel(self):
def kernel(self) -> dict[str, u.Quantity]:
return self._kernel

@kernel.setter
def kernel(self, kernel):
def kernel(self, kernel: dict[str, u.Quantity]):
if len(kernel) != len(self.data):
raise ValueError('Number of kernels must be equal to length of wavelength dimension.')
if not all([v.shape == self.kernel_temperatures.shape for _, v in kernel.items()]):
raise ValueError('Temperature bin centers and kernels must have the same shape.')
self._kernel = kernel

@property
def data_matrix(self):
def data_matrix(self) -> u.Quantity:
return np.stack([self.data[k].data for k in self._keys])

@property
def uncertainty_matrix(self):
def uncertainty_matrix(self) -> u.Quantity:
uncertainties = [self.data[k].uncertainty for k in self._keys]
if any([_u is None for _u in uncertainties]):
return None
return np.stack([_u.array for _u in uncertainties])

@property
def kernel_matrix(self):
def kernel_matrix(self) -> u.Quantity:
return np.stack([self.kernel[k].value for k in self._keys])

def fit(self, *args, **kwargs):
def fit(self, *args, **kwargs) -> ndcube.NDCube:
r"""
Apply inversion procedure to data.

Expand All @@ -146,7 +169,7 @@ def fit(self, *args, **kwargs):
dem : `~ndcube.NDCube`
Differential emission measure as a function of temperature. The
temperature axis is evenly spaced in :math:`\log{T}`. The number
of dimensions depend on the input data.
of dimensions depends on the input data.
"""
dem_dict = self._model(*args, **kwargs)
wcs = self._make_dem_wcs()
Expand Down Expand Up @@ -178,6 +201,6 @@ def _make_dem_wcs(self):
compound_wcs = CompoundLowLevelWCS(data_wcs, temp_table_coord.wcs, mapping=mapping)
return compound_wcs

def _make_dem_meta(self):
def _make_dem_meta(self) -> dict[str, object]:
# Individual classes should override this if they want specific metadata
return {}