mwa_jaxbeam
mwa_jaxbeam is a pure JAX implementation of the Murchison Widefield Array (MWA) Average Embedded Element (AEE) beam model.
The package reproduces the 137 MHz AEE model distributed with pyuvdata while providing a differentiable implementation suitable for optimization and inference problems. It is designed for applications such as satellite-based beam calibration, where the beam model must be evaluated millions of times and differentiated with respect to instrument parameters.
[!WARNING] This is experimental software. At present, only the 137 MHz MWA AEE beam model with zenith pointing is implemented.
Features
- Pure JAX implementation with automatic differentiation.
- Numerically validated against the MWA AEE implementation in
pyuvdata. - Fast evaluation of the full complex tile Jones matrix.
- Support for arbitrary real or complex dipole excitations.
- Efficient JIT compilation for large sky grids.
Installation
pip install mwa-jaxbeam
or
git clone https://github.com/amanchokshi/mwa_jaxbeam.git
cd mwa_jaxbeam
uv sync
or
pip install -e .
Example
import jax.numpy as jnp
from mwa_jaxbeam.aee_137mhz import jones
az_rad = jnp.deg2rad(jnp.linspace(0.0, 360.0, 361))
za_rad = jnp.deg2rad(jnp.linspace(0.0, 90.0, 91))
az_grid, za_grid = jnp.meshgrid(
az_rad,
za_rad,
)
beam = jones(
az_rad=az_grid,
za_rad=za_grid,
)
The returned Jones matrix has shape
(2, 2, *broadcast_shape)
where the first axis contains the sky-vector components (φ, θ) and the second axis contains the tile feeds (X, Y).
Validation
The implementation has been validated against the MWA AEE model distributed with pyuvdata by independently reproducing each stage of the calculation
Performance
On an Apple M4 Pro, evaluation of the full tile Jones matrix for a typical satellite beam-calibration workload (6144 sky directions) takes approximately
- 0.5 ms per forward evaluation (~2000 evaluations/s),
- 0.86 ms for a forward evaluation plus reverse-mode gradient (~1160 evaluations/s).
Performance scales approximately linearly with the number of sky directions after JIT compilation.
Status
The current implementation supports
- the 137 MHz MWA Average Embedded Element beam model,
- arbitrary complex dipole excitations,
- automatic differentiation through the complete beam calculation.
Future work includes interpolation across frequency and support for arbitrary beamformer delay settings.
Citation
If you use this package in published work, please cite the relevant MWA beam-model and pyuvdata publications, together with this repository.
Release files for mwa-jaxbeam 0.2.0
For a detailed explanation of source distributions (sdists) and built distributions (wheels), please see the package formats documentation.
Source distribution (sdist)
| File | Size | Uploaded | |
|---|---|---|---|
| mwa_jaxbeam-0.2.0.tar.gz | 1.3 MB | Details |
Built distribution (wheel)
| File | Interpreter | ABI | Platform | Reset |
|---|---|---|---|---|
| mwa_jaxbeam-0.2.0-py3-none-any.whl | Python 3 | none | any | Details |
Total release size: 1.5 MB
Release files / mwa_jaxbeam-0.2.0.tar.gz
| Download URL | mwa_jaxbeam-0.2.0.tar.gz |
|---|---|
| Size | 1.3 MB |
| Tags | Source |
|
SHA-256 checksum How to use checksums |
2ed9978b54cc9a1630fd27c89ca079f4e3a01334089e7d47388ffa6536cba4ed
|
|
BLAKE2b-256 checksum How to use checksums |
84e1a5d53b80af8e81fb5bf5c5c02b4a71f10a7138dba4d065447692e5924c24
|
| Upload date | |
|
Uploaded using Trusted Publishing? What is trusted publishing? |
No |
| Uploaded via |
twine/7.0.0 CPython/3.14.3
|
Release files / mwa_jaxbeam-0.2.0-py3-none-any.whl
| Download URL | mwa_jaxbeam-0.2.0-py3-none-any.whl |
|---|---|
| Size | 170.1 kB |
| Tags | Python 3 |
|
SHA-256 checksum How to use checksums |
28462e93111d2efb226ad1d2ab119b3280966c2c2c2ede19edbb086ae5530dd7
|
|
BLAKE2b-256 checksum How to use checksums |
3ca6386d7bdd313e425e1f950b5e828db026a92d515805722f4428b974d5b416
|
| Upload date | |
|
Uploaded using Trusted Publishing? What is trusted publishing? |
No |
| Uploaded via |
twine/7.0.0 CPython/3.14.3
|