Source code for tabascal.components.rfi_vis

import jax.numpy as jnp

from tabascal.distributed import psum_over_rfi, sharding_enabled
from tabascal.interferometry import calculate_rfi_vis_fine, calculate_rfi_vis_variable
from tabascal.components import Component
from ri_kernels.jax_api import RFIVisOp


[docs] class RiemannVis(Component): required_inputs = { "rfi_phase": ("n_rfi", "n_ant", "n_freq_fine", "n_time_fine"), "rfi_A": ("n_rfi", "n_ant", "n_freq_fine", "n_time_fine"), } output_shape = {"vis_rfi": ("n_bl", "n_freq", "n_time")} parameters = {}
[docs] def setup(self, config): """All validation and error-prone operations here""" try: self.a1 = config.a1 self.a2 = config.a2 self.n_int_time = config.n_int_time # self.n_int_freq = config.n_int_freq self.n_int_freq = config.args["rfi"]["freq_int_samples"] self.n_time = config.n_time self.n_bl = config.n_bl self.n_freq = config.n_freq # Validate dimensions self._set_outputs() # self._validate_dimensions() except Exception as e: raise RuntimeError(f"{self.__class__.__name__} setup failed: {e}")
# def _validate_dimensions(self): # """Ensure all setup operations completed successfully""" # assert hasattr(self, "")
[docs] def build_set_params(self): def set_params(params): return params return set_params
[docs] def build_constants(self): return {"a1": self.a1, "a2": self.a2}
[docs] def build_forward(self): """Return pure, JIT-compatible function""" prefix = self.prefix n_int_time = self.n_int_time n_int_freq = self.n_int_freq n_time = self.n_time n_bl = self.n_bl n_freq = self.n_freq def forward(params, state, constants): # Pure JAX operations only a1 = constants[f"{prefix}/a1"] a2 = constants[f"{prefix}/a2"] # Per-RFI-shard body (any leading RFI count); psum-ed across devices # under sharding. The fine->coarse mean runs before the cross-device # sum, so the collective is only coarse-grid sized (sum/mean commute). def local_vis(rfi_A, rfi_phase): vis_rfi_fine = calculate_rfi_vis_fine(rfi_A, rfi_phase, a1, a2) # vis_rfi_fine is shape (n_bl, n_freq_fine, n_time_fine) new_shape = (n_bl, n_freq, n_int_freq, n_time, n_int_time) return jnp.mean(jnp.reshape(vis_rfi_fine, new_shape), axis=(-3, -1)) vis_rfi = psum_over_rfi(local_vis)(state["rfi_A"], state["rfi_phase"]) # vis_rfi is shape (n_bl, n_freq, n_time) state = {**state, "vis_rfi": state["vis_rfi"] + vis_rfi} return state return forward
def _set_outputs(self): self.state_outputs = { "vis_rfi": jnp.zeros((self.n_bl, self.n_freq, self.n_time), dtype=complex), }
[docs] class RiemannVisFFI(Component): required_inputs = { "rfi_phase": ("n_rfi", "n_ant", "n_freq_fine", "n_time_fine"), "rfi_A": ("n_rfi", "n_ant", "n_freq_fine", "n_time_fine"), } output_shape = {"vis_rfi": ("n_bl", "n_freq", "n_time")} parameters = {}
[docs] def setup(self, config): """All validation and error-prone operations here""" try: self.a1 = config.a1 self.a2 = config.a2 self.n_int_time = config.n_int_time # self.n_int_freq = config.n_int_freq self.n_int_freq = config.args["rfi"]["freq_int_samples"] self.n_time = config.n_time self.n_bl = config.n_bl self.n_freq = config.n_freq self.n_ant = config.n_ant self.n_rfi = config.n_rfi # Validate dimensions self._set_outputs() # self._validate_dimensions() except Exception as e: raise RuntimeError(f"{self.__class__.__name__} setup failed: {e}")
# def _validate_dimensions(self): # """Ensure all setup operations completed successfully""" # assert hasattr(self, "")
[docs] def build_set_params(self): def set_params(params): return params return set_params
[docs] def build_forward(self): """Return pure, JIT-compatible function""" # Pre-compute everything possible n_int_time = self.n_int_time n_int_freq = self.n_int_freq n_time = self.n_time n_freq = self.n_freq n_ant = self.n_ant op = RFIVisOp(n_ant, self.a1, self.a2) def forward(params, state, constants): # Leading dim is -1, not n_rfi: under sharding the body below runs on # the per-device RFI shard, whose count is n_rfi / n_devices. The FFI # kernel itself runs unmodified per device inside shard_map (GSPMD # cannot partition a custom call); results are psum-ed across devices. def local_vis(rfi_A, rfi_phase): new_shape = (-1, n_ant, n_freq, n_int_freq, n_time, n_int_time) rfi_amp_fine = rfi_A.reshape(new_shape) rfi_phase_fine = rfi_phase.reshape(new_shape) # Transpose to (n_ant, n_freq, n_time, n_rfi_local, n_int_freq, n_int_time) rfi_amp_fine = jnp.transpose(rfi_amp_fine, (1, 2, 4, 0, 3, 5)) rfi_phase_fine = jnp.transpose(rfi_phase_fine, (1, 2, 4, 0, 3, 5)) return op.eval(rfi_amp_fine, rfi_phase_fine) vis_rfi = psum_over_rfi(local_vis)(state["rfi_A"], state["rfi_phase"]) state = {**state, "vis_rfi": state["vis_rfi"] + vis_rfi} return state return forward
def _set_outputs(self): self.state_outputs = { "vis_rfi": jnp.zeros((self.n_bl, self.n_freq, self.n_time), dtype=complex), }
[docs] class RiemannVisVariable(Component): required_inputs = { "rfi_phase": ("n_rfi", "n_ant", "n_freq_fine", "n_time_fine"), "rfi_A": ("n_rfi", "n_ant", "n_freq_fine", "n_time_fine"), } output_shape = {"vis_rfi": ("n_bl", "n_freq", "n_time")} parameters = {}
[docs] def setup(self, config): """All validation and error-prone operations here""" try: self.a1 = config.a1 self.a2 = config.a2 self.n_int_time = config.n_int_time self.n_int_freq = config.args["rfi"]["freq_int_samples"] self.n_rfi = config.n_rfi self.n_ant = config.n_ant self.n_time = config.n_time self.n_bl = config.n_bl self.n_freq = config.n_freq self.time_sample_idxs = config.time_sample_idxs self.time_strides = config.time_strides self._set_outputs() except Exception as e: raise RuntimeError(f"{self.__class__.__name__} setup failed: {e}")
[docs] def build_set_params(self): def set_params(params): return params return set_params
[docs] def build_constants(self): constants = {"a1": self.a1, "a2": self.a2} for i, idx in enumerate(self.time_sample_idxs): constants[f"time_sample_idxs_{i}"] = idx return constants
[docs] def build_forward(self): """Return pure, JIT-compatible function""" prefix = self.prefix n_int_time = self.n_int_time n_int_freq = self.n_int_freq n_ant = self.n_ant n_time = self.n_time n_bl = self.n_bl n_freq = self.n_freq n_groups = len(self.time_sample_idxs) time_strides = self.time_strides def calculate_grouped_rfi_vis(rfi_A, rfi_phase, a1, a2, constants): vis_rfi = jnp.empty((n_bl, n_freq, n_time), dtype=complex) for i, time_stride in zip(range(n_groups), time_strides): idx = constants[f"{prefix}/time_sample_idxs_{i}"] vis_rfi = vis_rfi.at[idx].set( calculate_rfi_vis_variable( rfi_A, rfi_phase, a1[idx], a2[idx], 1, time_stride ) ) return vis_rfi def forward(params, state, constants): # Pure JAX operations only a1 = constants[f"{prefix}/a1"] a2 = constants[f"{prefix}/a2"] # Leading dim -1: under sharding the body sees the per-device RFI # shard. Only replicated arrays (a1/a2, time_sample_idxs) are closed # over; the local sum over sources happens before the psum. def local_vis(rfi_A_flat, rfi_phase_flat): new_shape = ( -1, n_ant, n_freq, n_int_freq, n_time, n_int_time, ) # calculate_rfi_vis_variable expects the n_rfi axis on axis 1 and # reduces over it internally, so reshape to (n_rfi, n_ant, ...) and # swap to (n_ant, n_rfi, n_freq, n_int_freq, n_time, n_int_time). rfi_A = jnp.swapaxes(jnp.reshape(rfi_A_flat, new_shape), 0, 1) rfi_phase = jnp.swapaxes(jnp.reshape(rfi_phase_flat, new_shape), 0, 1) return calculate_grouped_rfi_vis(rfi_A, rfi_phase, a1, a2, constants) vis_rfi = psum_over_rfi(local_vis)(state["rfi_A"], state["rfi_phase"]) # vis_rfi is shape (n_bl, n_freq, n_time) state = {**state, "vis_rfi": state["vis_rfi"] + vis_rfi} return state return forward
def _set_outputs(self): self.state_outputs = { "vis_rfi": jnp.zeros((self.n_bl, self.n_freq, self.n_time), dtype=complex), }
[docs] class RiemannVisVariableFFI(Component): required_inputs = { "rfi_phase": ("n_rfi", "n_ant", "n_freq_fine", "n_time_fine"), "rfi_A": ("n_rfi", "n_ant", "n_freq_fine", "n_time_fine"), } output_shape = {"vis_rfi": ("n_bl", "n_freq", "n_time")} parameters = {}
[docs] def setup(self, config): """All validation and error-prone operations here""" try: self.a1 = config.a1 self.a2 = config.a2 self.n_int_time = config.n_int_time self.n_int_freq = config.args["rfi"]["freq_int_samples"] self.n_rfi = config.n_rfi self.n_ant = config.n_ant self.n_time = config.n_time self.n_bl = config.n_bl self.n_freq = config.n_freq self.time_sample_idxs = config.time_sample_idxs self.time_strides = config.time_strides self._set_outputs() except Exception as e: raise RuntimeError(f"{self.__class__.__name__} setup failed: {e}")
[docs] def build_set_params(self): def set_params(params): return params return set_params
[docs] def build_forward(self): """Return pure, JIT-compatible function""" # Pre-compute everything possible n_int_time = self.n_int_time n_int_freq = self.n_int_freq n_ant = self.n_ant n_time = self.n_time n_bl = self.n_bl n_freq = self.n_freq n_groups = len(self.time_sample_idxs) time_strides = self.time_strides time_sample_idxs = self.time_sample_idxs # Build one FFI operator per baseline group, each holding the precomputed # antenna-baseline indices for that group's subset of baselines. ops = [ RFIVisOp(n_ant, self.a1[idx], self.a2[idx]) for idx in time_sample_idxs ] if sharding_enabled(): print( "\n!!! WARNING !!! RiemannVisVariableFFI scales poorly " "across multiple devices. Consider using " "RiemannVisFFI instead for multi-device runs.\n" ) def calculate_grouped_rfi_vis(rfi_amp_fine, rfi_phase): vis_rfi = jnp.empty((n_bl, n_freq, n_time), dtype=complex) for i, time_stride in zip(range(n_groups), time_strides): idx = time_sample_idxs[i] # Subsample the integration-time axis by the group's stride, # mirroring calculate_rfi_vis_variable. The FFI kernel then # reduces over the remaining integration samples. t_idx = slice(time_stride // 2, None, time_stride) vis_rfi = vis_rfi.at[idx].set( ops[i].eval( rfi_amp_fine[..., t_idx], rfi_phase[..., t_idx], ) ) return vis_rfi def forward(params, state, constants): # Leading dim -1: under sharding the body sees the per-device RFI # shard. The FFI kernel reduces over the source axis itself, so the # local sum over sources happens before the psum. shard_map is also # what lets the custom call run at all -- GSPMD cannot partition it. def local_vis(rfi_A_flat, rfi_phase_flat): new_shape = ( -1, n_ant, n_freq, n_int_freq, n_time, n_int_time, ) rfi_amp_fine = jnp.reshape(rfi_A_flat, new_shape) rfi_phase = jnp.reshape(rfi_phase_flat, new_shape) # Transpose to (n_ant, n_freq, n_time, n_rfi, n_int_freq, n_int_time) rfi_amp_fine = jnp.transpose(rfi_amp_fine, (1, 2, 4, 0, 3, 5)) rfi_phase = jnp.transpose(rfi_phase, (1, 2, 4, 0, 3, 5)) return calculate_grouped_rfi_vis(rfi_amp_fine, rfi_phase) vis_rfi = psum_over_rfi(local_vis)(state["rfi_A"], state["rfi_phase"]) # vis_rfi is shape (n_bl, n_freq, n_time) state = {**state, "vis_rfi": state["vis_rfi"] + vis_rfi} return state return forward
def _set_outputs(self): self.state_outputs = { "vis_rfi": jnp.zeros((self.n_bl, self.n_freq, self.n_time), dtype=complex), }