PyTorch implementation of MSAtt-TransUNet, a multi-source attention-enhanced Transformer U-Net for sea surface salinity (SSS) retrieval in the Northwest Pacific.
This project uses multi-source satellite and auxiliary data to reconstruct daily gridded SSS fields. The current release provides the core training and inference code used in the manuscript.
- Task: SSS retrieval from remote sensing data
- Study area: Northwest Pacific (0–50°N, 100–150°E)
- Main inputs: SMAP brightness temperatures, SST, wind, and precipitation
- Target: Gridded sea surface salinity
- Framework: TransUNet-style Vision Transformer + U-Net decoder
.
├── train.py
├── predict.py
├── train_eval.py
├── criterion.py
├── mydataset.py
├── transforms.py
├── models/
├── mlp/
└── vit_checkpoint/
Install dependencies with:
pip install -r requirements.txtor
conda env create -f environment.yml
conda activate msatt-transunetRaw datasets are not re-hosted in this repository. Please download them from the original providers and organize them into your local data directory before training/inference.
-
SMAP L2B CAP Sea Surface Salinity / Brightness Temperature
https://podaac.jpl.nasa.gov/dataset/SMAP_JPL_L2B_SSS_CAP_V5 -
NOAA OISST v2.1
https://www.ncei.noaa.gov/products/optimum-interpolation-sst -
GPM IMERG precipitation
https://gpm.nasa.gov/data/imerg -
Copernicus Marine Multi Observation Global Ocean Sea Surface Salinity (used as label/reference)
https://data.marine.copernicus.eu/product/MULTIOBS_GLO_PHY_S_SURFACE_MYNRT_015_013/services
-
HYCOM + NCODA Global 1/12° Analysis
https://www.hycom.org/data/glbu0pt08 -
Argo GDAC profile data
https://argo.ucsd.edu/data/data-from-gdacs/
Example:
python train.py \
--input-root /path/to/data_root \
--mask-path /path/to/mask.npy \
--device cuda \
--batch-size 32 \
--epochs 100 \
--img_size 224 \
--vit_name R50-ViT-B_16The prediction script can export daily NetCDF files and summary metrics.
- Please replace local paths with your own paths before running the code.
- Raw public datasets should be downloaded from their official sources.
- Mean/std normalization values should be set according to your processed dataset.
- Pretrained ViT weights should be placed in
vit_checkpoint/if required by your setup.
If you use this code, please cite the corresponding manuscript.
This implementation is built upon the TransUNet-style hybrid CNN-Transformer framework and is adapted for multi-source ocean remote sensing inputs.
- Image_Segmentation: https://github.com/LeeJunHyun/Image_Segmentation
- TransUNet: https://github.com/Beckschen/TransUNet