WaveNeXt / README.md
umpaoflumpia's picture
Upload README.md with huggingface_hub
52f13c3 verified
|
Raw
History Blame Contribute Delete
5.71 kB
---
license: cc-by-nc-4.0
base_model: facebook/convnextv2-base-22k-224
tags:
- sar
- remote-sensing
- sar-to-optical
- image-to-image
- gan
- convnextv2
- wavelet
- onnx
pipeline_tag: image-to-image
library_name: pytorch
model-index:
- name: WaveNeXt
results:
- task:
type: image-to-image
name: SAR-to-optical translation
dataset:
type: sen1-2
name: SEN1-2 (5-scene subset)
metrics:
- type: psnr
value: 18.54
name: PSNR
- type: ssim
value: 0.432
name: SSIM
- type: fid
value: 58.5
name: FID
- type: lpips
value: 0.241
name: LPIPS
---
# WaveNeXt (Base)
**Wavelet + ConvNeXt** generator that translates single-channel Sentinel-1 **SAR**
amplitude into 3-channel Sentinel-2-like **optical** imagery at 256×256.
![SAR input · Baseline · HF-D · ground truth](hfd-showcase.jpg)
## Model details
| Field | Value |
|---|---|
| **Task** | Conditional SAR → optical image translation |
| **Architecture** | Haar-wavelet stem → ConvNeXt V2-Base backbone → inverse-Haar head (~98 M params) |
| **Finetuned from** | [`facebook/convnextv2-base-22k-224`](https://huggingface.co/facebook/convnextv2-base-22k-224) |
| **Resolution** | 256 × 256, input/output in `[-1, 1]` |
| **Formats** | `safetensors` (PyTorch) · `model.onnx` (fp32, opset 17) |
| **License** | CC-BY-NC-4.0 |
| **Repository** | [github.com/Tiruum/sar2opt_light](https://github.com/Tiruum/sar2opt_light) |
### How it works
- **Wavelet I/O** — a fixed orthonormal 2-level Haar transform replaces the patch-embed
stem, and an inverse-Haar head reconstructs the optical image, so the network predicts
wavelet sub-bands rather than raw pixels.
- **ConvNeXt V2-Base** backbone transfers ImageNet-22k features into the data-scarce SAR domain.
- **High-frequency discriminator (HF-D)** — an adversarial critic on the residual
`x − gaussian_blur(x)` drives coherent fine detail. It is used **only during training**
and adds no inference cost; these weights are the generator alone.
Full architecture and design notes: [`ARCHITECTURE.md`](https://github.com/Tiruum/sar2opt_light/blob/master/src/models/wavenext/ARCHITECTURE.md).
## Intended uses & limitations
**Intended use** — research on SAR→optical translation, despeckling, and high-frequency
detail synthesis for remote-sensing imagery.
**Limitations** — trained on a representative 5-scene subset of SEN1-2 (scenes
`5, 45, 52, 84, 100`); performance on regions, seasons, or sensors outside that
distribution is unverified. Outputs are plausible reconstructions, not measurements —
do not use for quantitative geophysical analysis. Non-commercial use only.
## Usage
### ONNX (no PyTorch / transformers)
```python
import numpy as np, onnxruntime as ort
from huggingface_hub import hf_hub_download
onnx_path = hf_hub_download("umpaoflumpia/WaveNeXt", "model.onnx")
sess = ort.InferenceSession(onnx_path, providers=["CPUExecutionProvider"])
sar = np.random.randn(1, 1, 256, 256).astype("float32") # SAR in [-1, 1]
optical = sess.run(None, {"sar": sar})[0] # [1, 3, 256, 256] in [-1, 1]
```
The batch axis is dynamic (`[N,1,256,256]`). Swap the provider for
`CUDAExecutionProvider`, `TensorrtExecutionProvider`, `CoreMLExecutionProvider`, or
`DmlExecutionProvider` to match your hardware.
### PyTorch
```python
import torch
from huggingface_hub import hf_hub_download
from safetensors.torch import load_file
from omegaconf import OmegaConf
from src.models.wavenext.gen import WaveNeXtGenerator # from the source repository
weights = hf_hub_download("umpaoflumpia/WaveNeXt", "generator.safetensors")
cfg = OmegaConf.load("src/models/wavenext/config.yaml")
g = WaveNeXtGenerator(cfg).eval()
g.load_state_dict(load_file(weights))
sar = torch.randn(1, 1, 256, 256) # SAR in [-1, 1]
with torch.no_grad():
optical = g(sar) # [1, 3, 256, 256] in [-1, 1]
```
Map either output to display range with `(x + 1) / 2`.
## Training
- **Data** — [SEN1-2](https://mediatum.ub.tum.de/1436631), paired Sentinel-1/Sentinel-2
patches; a fixed, representative 5-scene split (`5, 45, 52, 84, 100`).
- **Backbone** — ConvNeXt V2-Base, ImageNet-22k pretrained; Haar stem and inverse-Haar
head are fixed (non-learnable).
- **Objective** — LSGAN + feature matching + HF-D adversarial + MS-SSIM + per-band Haar L1
+ LPIPS + focal frequency loss + PatchNCE (no pixel-space L1).
- **Schedule** — AdamW/Adam (`lr 2e-4`), bf16 mixed precision, EMA (decay 0.999),
200 epochs with a linear LR decay tail.
Reproduce from the [source repository](https://github.com/Tiruum/sar2opt_light):
`python -m src.models.wavenext.train`.
## Evaluation
SEN1-2 held-out validation:
| Variant | PSNR ↑ | SSIM ↑ | FID ↓ | LPIPS ↓ |
|--------|--------|--------|-------|---------|
| **WaveNeXt Base** (this model) | **18.54** | **0.432** | **58.5** | **0.241** |
| WaveNeXt Tiny | 17.28 | 0.369 | 73.0 | 0.311 |
The figure above contrasts the baseline (HF-D disabled) with HF-D on a held-out crop:
HF-D recovers coherent high-frequency structure the baseline blurs away.
## Acknowledgements
Built on [ConvNeXt V2](https://huggingface.co/facebook/convnextv2-base-22k-224) (Meta AI)
and trained on [SEN1-2](https://mediatum.ub.tum.de/1436631) (TU Munich).
## License
**CC-BY-NC-4.0** — non-commercial. The weights are derived from ConvNeXt V2 (Meta,
CC-BY-NC-4.0) and trained on SEN1-2 (research use); those terms are inherited. Please
attribute WaveNeXt, ConvNeXt V2, and SEN1-2 in derivative work.