Swin / DAT / DAT++ Nowcasting on SEVIR

Controlled comparison of three attention mechanisms for radar nowcasting on the SEVIR dataset, using a shared UNet decoder to isolate the effect of the encoder's attention type.

Variants

Variant Encoder Attention Dims Depths Pretrained Params (approx)
Swin torchvision swin_t Shifted Window 96/192/384/768 2/2/6/2 ImageNet-1K (auto) ~30M
DAT Custom (same skeleton) Local + Deformable 96/192/384/768 2/2/6/2 From scratch ~31M
DAT++ Official LeapLabTHU DAT Local/Deformable + NATTEN 64/128/256/512 2/4/18/2 ImageNet-1K (manual) ~28M

All three share:

  • Decoder: UNetDecoder with skip connections (GroupNorm + SiLU)
  • Loss: CombinedLoss = (1-alpha) * B-MSE + alpha * SSIM
  • Metrics: CSI/POD/FAR at 5 thresholds + CSI_avg + MSE + MAE
  • Output: Raw logits (no sigmoid), clamp to [0,1] at eval time

Quick Start

git clone https://huggingface.co/huilinsigehigh/dat-swin-sevir-nowcast
cd dat-swin-sevir-nowcast
pip install -r requirements.txt
# For DAT++ only: pip install natten -f https://shi-labs.com/natten/wheels

# Smoke test (5 epochs, verify fix works)
python train_compare.py --variants swin --pretrained \
  --exp_name smoke --epochs 5 --batch_size 4 \
  --lr 1e-4 --warmup_epochs 2 --weight_decay 1e-5

# Full training (100 epochs)
python train_compare.py --variants swin dat dat_pp --pretrained \
  --epochs 100 --batch_size 8 --lr 1e-4 --warmup_epochs 5

# Pass criteria for smoke (5 ep):
#   - train_loss monotonically decreasing (no epoch-2 rebound)
#   - CSI_avg >= 0.20
#   - POD_light < 0.95

SEVIR Data

The training script auto-detects SEVIR data from these paths (first match wins):

  1. /root/autodl-tmp/sevir (seetacloud data disk)
  2. /root/data/sevir
  3. /data/datasets/sevir
  4. X:\datasets\sevir (Windows mapped drive)
  5. C:\Users\97290\Desktop\datasets\sevir (local dev)

Expected structure:

sevir/
β”œβ”€β”€ CATALOG.csv
└── data/
    └── vil/
        └── *.h5

Pretrained Weights

Weight Source Size Download
Swin-T ImageNet-1K torchvision ~110MB Auto-downloaded by torchvision.models.swin_t(weights=IMAGENET1K_V1)
DAT-T++ ImageNet-1K LeapLabTHU/DAT ~92MB Included in pretrained/ via LFS. Original: OneDrive / TsinghuaCloud
DAT-T (original) N/A N/A Not needed -- DAT variant trains from scratch

Known Issues and Fixes Applied

This repo contains fixes for a training collapse bug. See docs/DIAGNOSIS_AND_FIX.md for full details.

Root cause: sigmoid output + bf16 autocast + lr=2e-4 no warmup + 50x BMSE weighting caused the model to saturate into an "everywhere rain" dead zone (pred mean=0.82 vs target mean=0.05, Pearson r=-0.50).

Fixes applied:

  1. Removed torch.sigmoid() from all model forward methods (raw logits output)
  2. Disabled bf16 autocast (FP32 training)
  3. lr=1e-4 with 5-epoch linear warmup + cosine annealing
  4. weight_decay=1e-5 (was 1e-4)

See docs/SATURATION_CONFIRMED.md for diagnostic evidence.

Deployment

For RTX6000 server deployment instructions, see docs/RTX6000_DEPLOYMENT.md.

Citation

If you use this code, please cite the original DAT papers:

@article{xia2023dat,
    title={DAT++: Spatially Dynamic Vision Transformer with Deformable Attention},
    author={Zhuofan Xia and Xuran Pan and Shiji Song and Li Erran Li and Gao Huang},
    year={2023},
    journal={arXiv preprint arXiv:2309.01430},
}

@InProceedings{Xia_2022_CVPR,
    author    = {Xia, Zhuofan and Pan, Xuran and Song, Shiji and Li, Li Erran and Huang, Gao},
    title     = {Vision Transformer With Deformable Attention},
    booktitle = {CVPR},
    year      = {2022},
    pages     = {4794-4803}
}

License

Apache 2.0. See LICENSE and NOTICE.

Downloads last month

-

Downloads are not tracked for this model. How to track
Inference Providers NEW
This model isn't deployed by any Inference Provider. πŸ™‹ Ask for provider support

Paper for huilinsigehigh/dat-swin-sevir-nowcast