Retrieval-augmented semantic segmentation with a mask-similarity-trained retriever

Weights bundle for the preprint "Retrieval-augmented semantic segmentation with a mask-similarity-trained retriever" by Jianan Ren, D. A. Payunen, Ya. Matyushenko, T. M. Tatarnikova.

Code, training scripts, and paper source: github.com/Yeetmq/retrieval-augmented-segmentation.

Headline result

On the Pascal VOC 2012 official val split (1449 images, 21 classes), fp32 re-evaluation on the same best.ckpt:

Model mIoU pixel-acc
UNet baseline 0.4655 0.8672
UNet + retrieval 0.5648 0.8894
Δ +0.0993 +0.0222

Improvement observed on all 21 classes simultaneously. Largest per-class gains: bird +0.25, bottle +0.19, dog +0.16, cow +0.16, cat +0.15.

What's in this repo

baseline/
  best.ckpt              # 390 MB — UNet baseline, best epoch 48
  hparams.yaml
  training_metrics.csv   # per-epoch CSVLogger log (50 epochs)
retrieval/
  best.ckpt              # 510 MB — UNet + cross-attention fusion, epoch 49
  hparams.yaml
  training_metrics.csv
retriever/
  best.ckpt              #  39 MB — MobileNet-v3-Large + SupCon head
  memory_bank.pt         # 1.3 GB — 1587 L2-normalized support embeddings + masks
  hparams.yaml
  training_metrics.csv
eval/
  comparison.csv         # fp32 re-eval baseline vs retrieval on VOC val
  per_class_iou.csv      # per-class IoU for both models

Total: ~2.1 GB.

Method (one paragraph)

Segmentation models systematically underperform on rare or atypical cases. This project augments a plain UNet with an inference-time retrieval path: a MobileNet-v3 retriever encodes the query image into a 256-D L2-normalized embedding, the top-k nearest labeled examples are fetched from a memory bank built from a held-out support split, their masks are re-encoded by a small MaskEncoder, and the resulting multi-scale features are fused into the UNet decoder via cross-attention at scales /16, /8, /4. The main technical contribution is a mask-similarity retriever loss L_ret = 0.8·L_Boundary-IoU + 0.2·L_CS-Jaccard — a soft-target SupCon in which pair relevance is defined by mask geometry rather than class labels, aligning retrieval with what cross-attention actually consumes.

Usage

# 1. Get the code
git clone https://github.com/Yeetmq/retrieval-augmented-segmentation.git
cd retrieval-augmented-segmentation
pip install -e .

# 2. Download this bundle
pip install huggingface_hub
huggingface-cli download yeetmq/retrieval-augmented-segmentation \
    --local-dir weights_release

# 3. Rebuild the runs/ layout the scripts expect
mkdir -p runs/baseline runs/retrieval runs/retriever
cp weights_release/baseline/best.ckpt      runs/baseline/best.ckpt
cp weights_release/retrieval/best.ckpt     runs/retrieval/best.ckpt
cp weights_release/retriever/best.ckpt     runs/retriever/best.ckpt
cp weights_release/retriever/memory_bank.pt runs/retriever/memory_bank.pt

# 4. Re-evaluate on Pascal VOC val (needs the dataset — see prepare_data.py)
python Experiments/pascal_voc/prepare_data.py
python Experiments/pascal_voc/build_h5.py
python Experiments/pascal_voc/build_splits.py
python Experiments/pascal_voc/compare_results.py

Precision note

training_metrics.csv values were logged live under precision="16-mixed". The headline numbers above and in eval/comparison.csv are an fp32 re-evaluation of the same best.ckpt on the same val split. The gap is small (e.g. retrieval mIoU 0.5648 fp32 vs 0.5634 fp16-mixed live) but non-zero due to the precision regime.

Training setup (short)

  • Data: Pascal VOC 2012 + SBD (trainaug 10 582 images). Stratified 15 % split → support (1587) for the memory bank; the remaining 8995 → query_train. Official VOC val (1449) held out.
  • Baseline UNet: 50 epochs, batch 16, AdamW (lr 3e-4, wd 1e-4), CosineAnnealingLR, Dice + Focal (0.5/0.5). No ImageNet pretrain.
  • Retriever: 80 epochs on the support split, batch 32, MobileNet-v3-Large (ImageNet-pretrained), embedding 256-D, mask-similarity SupCon (Ï„=0.1, α=0.2, boundary erosion d=2).
  • Retrieval segmenter: same optimization as baseline, batch 12 (fusion memory), fusion at scales {20, 40, 80}, 4 attention heads. Retriever frozen; UNet + MaskEncoder + 3 fusion modules trained.
  • Hardware: one NVIDIA RTX 5090 (32 GB). End-to-end ~24 h.

Citation

@misc{ren2026retrievalseg,
  title  = {Retrieval-augmented semantic segmentation with a mask-similarity-trained retriever},
  author = {Ren, Jianan and Payunen, D. A. and Matyushenko, Ya. and Tatarnikova, T. M.},
  year   = {2026},
  note   = {Preprint},
}

License

MIT.

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

Evaluation results

  • mIoU (fp32 re-eval, 21 classes) on Pascal VOC 2012 + SBD
    self-reported
    0.565
  • Pixel accuracy on Pascal VOC 2012 + SBD
    self-reported
    0.889