Skip to content
Merged
325 changes: 293 additions & 32 deletions docs/examples/rmsynth_3d_stokes_i.ipynb

Large diffs are not rendered by default.

107 changes: 84 additions & 23 deletions rm_lite/tools_1d/rmsynth.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,8 +3,9 @@
from __future__ import annotations

import time
from typing import Literal, NamedTuple
from typing import Literal, NamedTuple, get_args

import dask.array as da
import numpy as np
import polars as pl
from numpy.typing import NDArray
Expand All @@ -16,13 +17,17 @@
StokesIFitOptions,
coefficient_errors,
coefficient_names,
field_spectral_index,
)
from rm_lite.utils.logging import logger
from rm_lite.utils.synthesis import (
FDFOptions,
LamSq0Mode,
NoiseWeightType,
StokesData,
StokesIWeighting,
WeightType,
apply_weight_type,
compute_rmsynth_params,
compute_theoretical_noise,
create_fractional_spectra,
Expand All @@ -32,6 +37,7 @@
get_rmsf_nufft,
lambda2_to_freq,
rmsynth_nufft,
stokes_i_template,
)


Expand Down Expand Up @@ -114,6 +120,25 @@ def _stokes_i_terms(
)


def _spectrum_alpha(
stokes_data: StokesData,
fdf_options: FDFOptions,
fit_options: StokesIFitOptions,
) -> float:
"""The power-law index for the weights, pinned or fitted to this spectrum."""
if fdf_options.stokes_i_weight_alpha != "auto":
return float(fdf_options.stokes_i_weight_alpha)
model = stokes_data.stokes_i_model_arr
spectrum = model if model is not None else stokes_data.stokes_i_arr
assert spectrum is not None
return field_spectral_index(
da.from_array(np.asarray(spectrum)[:, np.newaxis, np.newaxis]),
None if model is not None else stokes_data.stokes_i_error_arr,
stokes_data.freq_arr_hz,
fit_options,
)


def run_rmsynth(
freq_arr_hz: NDArray[np.float64],
complex_pol_arr: NDArray[np.complexfloating],
Expand All @@ -134,6 +159,8 @@ def run_rmsynth(
fit_order: int = 2,
stokes_i_robust_loss: RobustLoss = "cauchy",
stokes_i_f_scale: float = 3.0,
stokes_i_weighting: StokesIWeighting | None = "global",
stokes_i_weight_alpha: float | Literal["auto"] = "auto",
ignore_stokes_i: bool = False,
moment_threshold_snr: float = 5.0,
) -> RMSynth1DResults:
Expand Down Expand Up @@ -170,6 +197,14 @@ def run_rmsynth(
Defaults to "cauchy".
stokes_i_f_scale (float, optional): How far out, in sigma, before a
channel is downweighted. Flat from 1 to 10. Defaults to 3.0.
stokes_i_weighting ("global", "per_pixel", None, optional): How the
inverse-variance weights follow the Stokes I division, see
`rm_lite.tools_3d.rmsynth.rmsynth_3d`. "global" weights by a power law
of `stokes_i_weight_alpha`, "per_pixel" by this spectrum's own model,
None by 1/sigma^2. Defaults to "global".
stokes_i_weight_alpha (float | "auto", optional): Spectral index of the
power law. "auto" fits this spectrum, so pin it to give a catalogue
of spectra one RMSF. Defaults to "auto".
moment_threshold_snr (float, optional): SNR cut (times the theoretical FDF noise) applied to FDF amplitudes before computing the Faraday moments. Defaults to 5.0.

Returns:
Expand Down Expand Up @@ -200,10 +235,11 @@ def run_rmsynth(
lam_sq_0_m2=lam_sq_0_m2,
do_fit_rmsf=do_fit_rmsf,
do_fit_rmsf_real=do_fit_rmsf_real,
stokes_i_weighting=stokes_i_weighting,
stokes_i_weight_alpha=stokes_i_weight_alpha,
)
# snr_cut=None: the 1D fractional fit has only one spectrum, so an SNR cut
# would just silently drop fractional polarisation rather than fall back to
# a flat per-pixel model as it does in 3D.
# would just silently drop its fractional polarisation.
fit_options = StokesIFitOptions(
fit_order=fit_order,
fit_function=fit_function,
Expand Down Expand Up @@ -254,10 +290,23 @@ def _run_rmsynth(
stokes_i_terms (pl.DataFrame): Fitted Stokes I model terms
"""

weighting = (
fdf_options.stokes_i_weighting
if not ignore_stokes_i and fdf_options.weight_type in get_args(NoiseWeightType)
else None
)
noise_error = stokes_data.complex_pol_error
weight_error = noise_error
if weighting is not None:
weight_alpha = _spectrum_alpha(stokes_data, fdf_options, fit_options)
weight_error = noise_error / stokes_i_template(
stokes_data.freq_arr_hz, weight_alpha
)

rmsynth_params = compute_rmsynth_params(
freq_arr_hz=stokes_data.freq_arr_hz,
complex_pol_arr=stokes_data.complex_pol_arr,
complex_pol_error=stokes_data.complex_pol_error,
complex_pol_error=weight_error,
fdf_options=fdf_options,
)

Expand All @@ -276,15 +325,43 @@ def _run_rmsynth(
no_nan_idx = fractional_stokes_data.no_nan_idx
fit_result = fractional_stokes_data.fit_result

# Compute after any fractional spectra have been created
tick = time.time()

# Perform RM-synthesis on the spectrum
all_flagged = (~no_nan_idx).all()
if all_flagged:
msg = "All channels have been masked!"
logger.warning(msg)

stokes_i_reference_flux = np.nan
if not ignore_stokes_i:
assert stokes_data.stokes_i_model_arr is not None
assert stokes_data.freq_arr_hz.shape == stokes_data.stokes_i_model_arr.shape
if not all_flagged:
stokes_i_model = interpolate.interp1d(
stokes_data.freq_arr_hz[no_nan_idx],
stokes_data.stokes_i_model_arr[no_nan_idx],
)

stokes_i_reference_flux = float(stokes_i_model(ref_freq_hz))
else:
logger.warning("Using mean as reference flux")
stokes_i_reference_flux = float(np.nanmean(stokes_data.stokes_i_model_arr))

if weighting == "per_pixel" and stokes_data.stokes_i_model_arr is not None:
# This spectrum's own model, in place of the power law used up to here.
shape = stokes_data.stokes_i_model_arr / stokes_i_reference_flux
rmsynth_params = rmsynth_params._replace(
weight_arr=apply_weight_type(
lambda_sq_arr_m2=rmsynth_params.lambda_sq_arr_m2,
real_qu_error=np.abs(noise_error.real + noise_error.imag) / 2.0 / shape,
channel_mask=~np.isfinite(stokes_data.complex_pol_arr),
fdf_options=fdf_options,
cell_m2=rmsynth_params.cell_m2,
)
)

# Compute after any fractional spectra have been created
tick = time.time()

fdf_dirty_arr = rmsynth_nufft(
complex_pol_arr=stokes_data.complex_pol_arr[no_nan_idx],
lambda_sq_arr_m2=rmsynth_params.lambda_sq_arr_m2[no_nan_idx],
Expand Down Expand Up @@ -314,19 +391,6 @@ def _run_rmsynth(
)

if not ignore_stokes_i:
assert stokes_data.stokes_i_model_arr is not None
assert stokes_data.freq_arr_hz.shape == stokes_data.stokes_i_model_arr.shape
if not all_flagged:
stokes_i_model = interpolate.interp1d(
stokes_data.freq_arr_hz[no_nan_idx],
stokes_data.stokes_i_model_arr[no_nan_idx],
)

stokes_i_reference_flux = float(stokes_i_model(ref_freq_hz))
else:
logger.warning("Using mean as reference flux")
stokes_i_reference_flux = float(np.nanmean(stokes_data.stokes_i_model_arr))

fdf_dirty_arr *= stokes_i_reference_flux

theoretical_noise = theoretical_noise._replace(
Expand All @@ -335,9 +399,6 @@ def _run_rmsynth(
fdf_u_noise=theoretical_noise.fdf_u_noise * stokes_i_reference_flux,
)

else:
stokes_i_reference_flux = np.nan

# Measure the parameters of the dirty FDF
# Use the theoretical noise to calculate uncertainties
logger.warning(
Expand Down
Loading
Loading