Generative model of Euclid Q1 galaxy images, developed during an internship at CosmoStat (CEA).
Two stages, trained on 64×64 postage stamps:
- Autoencoder: reconstructs the galaxy. Its output is convolved with the PSF (via
jax-galsim) before being compared with the observed image. - Normalizing flow: fitted on the latent space of the frozen autoencoder, to sample new galaxies.
The modeling code in pshear/ builds on prior work by Benjamin Rémy (CosmoStat). Training runs were done on the Jean Zay supercomputer (IDRIS/CNRS).
Two Hugging Face datasets. Each sample contains a science image, its PSF (full and partial), a noise map and a mask.
| Dataset | Galaxies | Used by |
|---|---|---|
| euclid-Q1-VF | ~50k | train_test.py, train_test_partial.py |
| Euclid-Q1-postage-stamps | ~260k | train_partial_parallel.py, train_flow.py, verification.py, evaluate_residuals.py |
The model was developed on euclid-Q1-VF (just over 50k galaxies). It was then also trained on Euclid-Q1-postage-stamps (about 260k galaxies), which seems to give better results.
pshear/ Core library (JAX / Equinox)
galaxy.py Galaxy autoencoder with PSF convolution, and its losses
nn/ Autoencoder, flow and network blocks
utils.py Checkpoint save/load, W&B checkpoint download
experiments/
train_partial_parallel.py Autoencoder training, multi-GPU (main script)
train_test_partial.py Autoencoder training, partial PSF, single GPU
train_test.py Autoencoder training, full PSF, single GPU
lr_range_test.py Learning-rate range test for train_partial_parallel.py
train_flow.py Flow training on the frozen autoencoder's latents
verification.py PQMass test: generated vs real distribution
evaluate_residuals.py Residual diagnostics to compare autoencoder checkpoints
download_wandb_weights.py Pre-downloads W&B checkpoints (for offline compute nodes)
test/ Environment and multi-GPU sanity checks
JAX is installed separately, with the CUDA build that matches the machine:
pip install -U "jax[cuda12]"
pip install -r requirements.txt
pip install "pqm>=0.6" # only for verification.py
Scripts are run from the repository root. Their hyperparameters are in the CONFIG dict at the top of each file. Runs are logged to Weights & Biases, and outputs go to $SCRATCH/pshear/cosmos/runs/ if $SCRATCH is set, ./runs/ otherwise.
python test/test_multi_gpu.py # check that JAX sees the GPUs
python -m experiments.train_partial_parallel # 1. train the autoencoder
python -m experiments.train_flow # 2. train the flow (set ae_run_dir / ae_epoch in CONFIG)
python -m experiments.verification # 3. check the generated samples
train_test.pyuses the full PSF. Deconvolving with it is ill-posed, so the loss adds a total-variation term to suppress pixelization artifacts.train_test_partial.pyandtrain_partial_parallel.pyuse the partial PSF, also provided in the dataset. No regularization term is needed.
train_partial_parallel.py trains on all the GPUs of one node from a single process (Mesh + shard_map). Model and optimizer state are replicated, batch_size is the global batch split across GPUs, and gradients are averaged with pmean. Submit it with one task for the whole node (e.g. --ntasks=1 --gres=gpu:4).
verification.pyuses PQMass to test whether the flow's samples follow the real data distribution, both in latent space and in image space. Real-vs-real calibration tests serve as the reference. Figures are written toPQM_results/.evaluate_residuals.pycompares several autoencoder checkpoints on the same test images: noise floor, residuals binned by signal-to-noise, and how many latent dimensions are used.
download_wandb_weights.py downloads the autoencoder and/or flow checkpoints of a W&B run into wandb_weights/<run_id>/epoch_<n>/. Run it on a node with network access. On an offline compute node, fetch_wandb_checkpoint then reads this cache without calling W&B.
python download_wandb_weights.py # runs set in the CONFIG block
python download_wandb_weights.py --only flow --flow-run-id <id> --flow-epoch <n>
python download_wandb_weights.py --cache-dir <dir> # other destination
The galaxy-morphometrics repository reads the same wandb_weights/ layout, so --cache-dir can point directly at its checkpoint directory.
Some Weights & Biases and Hugging Face identifiers are hard-coded and must be changed by anyone else using this repository:
- W&B entity and run IDs (
vincentb03-imt-atlantique):WANDB_ENTITYand the run IDs at the top ofdownload_wandb_weights.pyandverification.py. - W&B project names:
wandb_projectin theCONFIGoftrain_partial_parallel.pyandlr_range_test.py, andproject=inwandb.initintrain_test.py,train_test_partial.pyandtrain_flow.py. - Hugging Face datasets (
VincentB03/...): theload_datasetcalls in the training scripts,DATASET_NAMEinverification.pyandevaluate_residuals.py, andtest/test_requirements.py. Change them if the datasets are moved to another account.