MST-plus-plus:基于 Transformer 的光谱重建算法项目

"MST++: Multi-stage Spectral-wise Transformer for Efficient Spectral Reconstruction" (CVPRW 2022) & (Winner of NTIRE 2022 Spectral Recovery Challenge) and a toolbox for spectral reconstruction

Branch1Tags0
FilesLast commitLast update
4 years ago
2 years ago
4 years ago
4 years ago
3 years ago
1 year ago
3 years ago
2 years ago
2 years ago
9 months ago
4 years ago

MST++: Multi-stage Spectral-wise Transformer for Efficient Spectral Reconstruction (CVPRW 2022)

winner arXiv zhihu mst

蔡元浩,林静,林祖迪,王昊谦,张宇伦,Hanspeter Pfister,Radu Timofte,Luc Van Gool

前两位作者对本工作贡献相同

新闻

  • 2024.03.21 : 我们的方法 RetinexformerMST++(NTIRE 2022 光谱重建挑战赛冠军)在 NTIRE 2024 低光增强挑战赛 中排名第二。代码、预训练模型、训练日志和增强结果将在 Retinexformer 的代码库 中发布。敬请期待!🚀
  • 2024.02.15 : NTIRE 2024 低光增强挑战赛 开始。欢迎使用我们的 RetinexformerMST++(NTIRE 2022 光谱重建挑战赛冠军)参与此次挑战!🏆
  • 2023.11.02 : 我们的 MST++ 已被收录至 Awesome-Transformer-Attention 集合中。💫
  • 2022.10.24 : 我们提供了参数和 FLOPS 评估函数。欢迎查看和使用。
  • 2022.10.23 : 我们提供了一些可视化工具函数。欢迎查看和使用。
  • 2022.04.17 : 我们的论文已被 CVPRW 2022 接收,代码和模型已发布。🚀
  • 2022.04.02 : 我们获得了 NTIRE 2022 RGB 到光谱重建挑战赛的第一名。🏆
480 nm 520 nm 580 nm 660 nm

摘要: 现有的领先光谱重建(SR)方法侧重于设计更深或更宽的卷积神经网络(CNNs)来学习从 RGB 图像到其高光谱图像(HSI)的端到端映射。这些基于 CNN 的方法取得了令人印象深刻的重建性能,但在捕捉长程依赖关系和自相似先验方面存在局限性。为了解决这个问题,我们提出了一种新颖的基于 Transformer 的方法,即多阶段光谱注意力 Transformer(MST++),用于高效光谱重建。具体而言,我们利用基于高光谱图像空间稀疏但光谱自相似特性的光谱多头自注意力(S-MSA)来构建基本单元——光谱注意力块(SAB)。然后,SAB 构建单阶段光谱注意力 Transformer(SST),该 SST 利用 U 形结构提取多分辨率上下文信息。最后,我们的 MST++ 由多个 SST 级联而成,从粗到精逐步提高重建质量。综合实验表明,我们的 MST++ 显著优于其他最先进的方法。在 NTIRE 2022 光谱重建挑战赛中,我们的方法获得了第一名。


网络架构

MST示意图

我们的MST++主要基于我们已被CVPR 2022接收的工作MST

与现有先进方法的对比

本仓库是一个基准和工具箱,包含11种用于光谱重建的图像复原算法。

未来我们将继续扩充模型库。

支持的算法:

对比图

在NTIRE 2022 HSI数据集 - 验证集上的结果

方法 参数数量 (M) 计算量 (G) MRAE RMSE PSNR 模型库
HSCNN+ 4.65 304.45 0.3814 0.0588 26.36 Google Drive / 百度网盘
HRNet 31.70 163.81 0.3476 0.0550 26.89 Google Drive / 百度网盘
EDSR 2.42 158.32 0.3277 0.0437 28.29 Google Drive / 百度网盘
AWAN 4.04 270.61 0.2500 0.0367 31.22 Google Drive / 百度网盘
HDNet 2.66 173.81 0.2048 0.0317 32.13 Google Drive / 百度网盘
HINet 5.21 31.04 0.2032 0.0303 32.51 Google Drive / 百度网盘
MIRNet 3.75 42.95 0.1890 0.0274 33.29 Google Drive / 百度网盘
Restormer 15.11 93.77 0.1833 0.0274 33.40 Google Drive / 百度网盘
MPRNet 3.62 101.59 0.1817 0.0270 33.50 Google Drive / 百度网盘
MST-L 2.45 32.07 0.1772 0.0256 33.90 Google Drive / 百度网盘
MST++ 1.62 23.05 0.1645 0.0248 34.32 Google Drive / 百度网盘

我们的MST++在显著优于其他方法的同时,所需的参数数量和计算量也更低。

注:百度网盘的提取码为mst1

1. 创建环境:

  • Python 3(建议使用 Anaconda

  • NVIDIA GPU + CUDA

  • Python 包:

    cd MST-plus-plus
    pip install -r requirements.txt
    

2. 数据准备:

  • 从 NTIRE 2022 光谱重建挑战赛的 竞赛网站 下载训练光谱图像(Google Drive / 百度网盘,提取码:mst1)、训练 RGB 图像(Google Drive / 百度网盘)、验证光谱图像(Google Drive / 百度网盘)、验证 RGB 图像(Google Drive / 百度网盘)以及测试 RGB 图像(Google Drive / 百度网盘)。

  • 将训练光谱图像和验证光谱图像放置到 /MST-plus-plus/dataset/Train_Spec/ 路径下。

  • 将训练 RGB 图像和验证 RGB 图像放置到 /MST-plus-plus/dataset/Train_RGB/ 路径下。

  • 将测试 RGB 图像放置到 /MST-plus-plus/dataset/Test_RGB/ 路径下。

  • 完成后,本仓库的文件结构如下所示:

    |--MST-plus-plus
        |--test_challenge_code
        |--test_develop_code
        |--train_code  
        |--dataset 
            |--Train_Spec
                |--ARAD_1K_0001.mat
                |--ARAD_1K_0002.mat
                : 
                |--ARAD_1K_0950.mat
      	|--Train_RGB
                |--ARAD_1K_0001.jpg
                |--ARAD_1K_0002.jpg
                : 
                |--ARAD_1K_0950.jpg
            |--Test_RGB
                |--ARAD_1K_0951.jpg
                |--ARAD_1K_0952.jpg
                : 
                |--ARAD_1K_1000.jpg
            |--split_txt
                |--train_list.txt
                |--valid_list.txt
    

    注意: 如果你使用自定义数据集进行训练,请使用 train_code/utils.py 中的 Loss_MRAE_custom 函数,以避免出现 Nan 问题。

3. 在验证集上的评估:

(1) 从 (Google Drive / 百度网盘,提取码:mst1) 下载预训练模型库,并将其放置到 /MST-plus-plus/test_develop_code/model_zoo/ 路径下。

(2) 运行以下命令在验证集的 RGB 图像上测试模型。

cd /MST-plus-plus/test_develop_code/

# test MST++
python test.py --data_root ../dataset/  --method mst_plus_plus --pretrained_model_path ./model_zoo/mst_plus_plus.pth --outf ./exp/mst_plus_plus/  --gpu_id 0

# test MST-L
python test.py --data_root ../dataset/  --method mst --pretrained_model_path ./model_zoo/mst.pth --outf ./exp/mst/  --gpu_id 0

# test MIRNet
python test.py --data_root ../dataset/  --method mirnet --pretrained_model_path ./model_zoo/mirnet.pth --outf ./exp/mirnet/  --gpu_id 0

# test HINet
python test.py --data_root ../dataset/  --method hinet --pretrained_model_path ./model_zoo/hinet.pth --outf ./exp/hinet/  --gpu_id 0

# test MPRNet
python test.py --data_root ../dataset/  --method mprnet --pretrained_model_path ./model_zoo/mprnet.pth --outf ./exp/mprnet/  --gpu_id 0

# test Restormer
python test.py --data_root ../dataset/  --method restormer --pretrained_model_path ./model_zoo/restormer.pth --outf ./exp/restormer/  --gpu_id 0

# test EDSR
python test.py --data_root ../dataset/  --method edsr --pretrained_model_path ./model_zoo/edsr.pth --outf ./exp/edsr/  --gpu_id 0

# test HDNet
python test.py --data_root ../dataset/  --method hdnet --pretrained_model_path ./model_zoo/hdnet.pth --outf ./exp/hdnet/  --gpu_id 0

# test HRNet
python test.py --data_root ../dataset/  --method hrnet --pretrained_model_path ./model_zoo/hrnet.pth --outf ./exp/hrnet/  --gpu_id 0

# test HSCNN+
python test.py --data_root ../dataset/  --method hscnn_plus --pretrained_model_path ./model_zoo/hscnn_plus.pth --outf ./exp/hscnn_plus/  --gpu_id 0

# test AWAN
python test.py --data_root ../dataset/  --method awan --pretrained_model_path ./model_zoo/awan.pth --outf ./exp/awan/  --gpu_id 0

结果将以 mat 格式保存至 /MST-plus-plus/test_develop_code/exp/ 目录,并会打印评估指标(包括 MRAE、RMSE、PSNR)。

  • 评估模型的参数量与 FLOPS

我们已在 test_develop_code/utils.py 中提供了 my_summary() 函数,请使用该函数评估模型(尤其是 Transformers)的参数量与计算复杂度。

from utils import my_summary
my_summary(MST_Plus_Plus(), 256, 256, 3, 1)

4. 测试集评估:

(1) 从 (Google Drive / 百度网盘,提取码:mst1) 下载预训练模型库,并将其放置到 /MST-plus-plus/test_challenge_code/model_zoo/ 目录下。

(2) 运行以下命令在测试 RGB 图像上对模型进行测试。

cd /MST-plus-plus/test_challenge_code/

# test MST++
python test.py --data_root ../dataset/  --method mst_plus_plus --pretrained_model_path ./model_zoo/mst_plus_plus.pth --outf ./exp/mst_plus_plus/  --gpu_id 0

# test MST-L
python test.py --data_root ../dataset/  --method mst --pretrained_model_path ./model_zoo/mst.pth --outf ./exp/mst/  --gpu_id 0

# test MIRNet
python test.py --data_root ../dataset/  --method mirnet --pretrained_model_path ./model_zoo/mirnet.pth --outf ./exp/mirnet/  --gpu_id 0

# test HINet
python test.py --data_root ../dataset/  --method hinet --pretrained_model_path ./model_zoo/hinet.pth --outf ./exp/hinet/  --gpu_id 0

# test MPRNet
python test.py --data_root ../dataset/  --method mprnet --pretrained_model_path ./model_zoo/mprnet.pth --outf ./exp/mprnet/  --gpu_id 0

# test Restormer
python test.py --data_root ../dataset/  --method restormer --pretrained_model_path ./model_zoo/restormer.pth --outf ./exp/restormer/  --gpu_id 0

# test EDSR
python test.py --data_root ../dataset/  --method edsr --pretrained_model_path ./model_zoo/edsr.pth --outf ./exp/edsr/  --gpu_id 0

# test HDNet
python test.py --data_root ../dataset/  --method hdnet --pretrained_model_path ./model_zoo/hdnet.pth --outf ./exp/hdnet/  --gpu_id 0

# test HRNet
python test.py --data_root ../dataset/  --method hrnet --pretrained_model_path ./model_zoo/hrnet.pth --outf ./exp/hrnet/  --gpu_id 0

# test HSCNN+
python test.py --data_root ../dataset/  --method hscnn_plus --pretrained_model_path ./model_zoo/hscnn_plus.pth --outf ./exp/hscnn_plus/  --gpu_id 0

结果和 submission.zip 将保存在 /MST-plus-plus/test_challenge_code/exp/ 目录下。

5. 训练

要训练模型,请运行

cd /MST-plus-plus/train_code/

# train MST++
python train.py --method mst_plus_plus  --batch_size 20 --end_epoch 300 --init_lr 4e-4 --outf ./exp/mst_plus_plus/ --data_root ../dataset/  --patch_size 128 --stride 8  --gpu_id 0

# train MST-L
python train.py --method mst  --batch_size 20 --end_epoch 300 --init_lr 4e-4 --outf ./exp/mst/ --data_root ../dataset/  --patch_size 128 --stride 8  --gpu_id 0

# train MIRNet
python train.py --method mirnet  --batch_size 20 --end_epoch 300 --init_lr 4e-4 --outf ./exp/mirnet/ --data_root ../dataset/  --patch_size 128 --stride 8  --gpu_id 0

# train HINet
python train.py --method hinet  --batch_size 20 --end_epoch 300 --init_lr 2e-4 --outf ./exp/hinet/ --data_root ../dataset/  --patch_size 128 --stride 8  --gpu_id 0

# train MPRNet
python train.py --method mprnet  --batch_size 20 --end_epoch 300 --init_lr 2e-4 --outf ./exp/mprnet/ --data_root ../dataset/  --patch_size 128 --stride 8  --gpu_id 0

# train Restormer
python train.py --method restormer  --batch_size 20 --end_epoch 300 --init_lr 2e-4 --outf ./exp/restormer/ --data_root ../dataset/  --patch_size 128 --stride 8  --gpu_id 0

# train EDSR
python train.py --method edsr  --batch_size 20 --end_epoch 300 --init_lr 1e-4 --outf ./exp/edsr/ --data_root ../dataset/  --patch_size 128 --stride 8  --gpu_id 0

# train HDNet
python train.py --method hdnet  --batch_size 20 --end_epoch 300 --init_lr 4e-4 --outf ./exp/hdnet/ --data_root ../dataset/  --patch_size 128 --stride 8  --gpu_id 0

# train HRNet
python train.py --method hrnet  --batch_size 20 --end_epoch 300 --init_lr 1e-4 --outf ./exp/hrnet/ --data_root ../dataset/  --patch_size 128 --stride 8  --gpu_id 0

# train HSCNN+
python train.py --method hscnn_plus  --batch_size 20 --end_epoch 300 --init_lr 2e-4 --outf ./exp/hscnn_plus/ --data_root ../dataset/  --patch_size 128 --stride 8  --gpu_id 0

# train AWAN
python train.py --method awan  --batch_size 20 --end_epoch 300 --init_lr 1e-4 --outf ./exp/awan/ --data_root ../dataset/  --patch_size 128 --stride 8  --gpu_id 0

训练日志和模型将保存在 /MST-plus-plus/train_code/exp/ 目录下。

6. 预测

(1) 从 (Google Drive / 百度网盘,提取码:mst1) 下载预训练模型库,并将其放置到 /MST-plus-plus/predict_code/model_zoo/ 目录下。

(2) 运行以下命令重建您自己的 RGB 图像。

cd /MST-plus-plus/predict_code/

# reconstruct by MST++
python test.py --rgb_path ./demo/ARAD_1K_0912.jpg  --method mst_plus_plus --pretrained_model_path ./model_zoo/mst_plus_plus.pth --outf ./exp/mst_plus_plus/  --gpu_id 0

# reconstruct by MST-L
python test.py --rgb_path ./demo/ARAD_1K_0912.jpg  --method mst --pretrained_model_path ./model_zoo/mst.pth --outf ./exp/mst/  --gpu_id 0

# reconstruct by MIRNet
python test.py --rgb_path ./demo/ARAD_1K_0912.jpg  --method mirnet --pretrained_model_path ./model_zoo/mirnet.pth --outf ./exp/mirnet/  --gpu_id 0

# reconstruct by HINet
python test.py --rgb_path ./demo/ARAD_1K_0912.jpg  --method hinet --pretrained_model_path ./model_zoo/hinet.pth --outf ./exp/hinet/  --gpu_id 0

# reconstruct by MPRNet
python test.py --rgb_path ./demo/ARAD_1K_0912.jpg  --method mprnet --pretrained_model_path ./model_zoo/mprnet.pth --outf ./exp/mprnet/  --gpu_id 0

# reconstruct by Restormer
python test.py --rgb_path ./demo/ARAD_1K_0912.jpg  --method restormer --pretrained_model_path ./model_zoo/restormer.pth --outf ./exp/restormer/  --gpu_id 0

# reconstruct by EDSR
python test.py --rgb_path ./demo/ARAD_1K_0912.jpg --method edsr --pretrained_model_path ./model_zoo/edsr.pth --outf ./exp/edsr/  --gpu_id 0

# reconstruct by HDNet
python test.py --rgb_path ./demo/ARAD_1K_0912.jpg  --method hdnet --pretrained_model_path ./model_zoo/hdnet.pth --outf ./exp/hdnet/  --gpu_id 0

# reconstruct by HRNet
python test.py --rgb_path ./demo/ARAD_1K_0912.jpg  --method hrnet --pretrained_model_path ./model_zoo/hrnet.pth --outf ./exp/hrnet/  --gpu_id 0

# reconstruct by HSCNN+
python test.py --rgb_path ./demo/ARAD_1K_0912.jpg  --method hscnn_plus --pretrained_model_path ./model_zoo/hscnn_plus.pth --outf ./exp/hscnn_plus/  --gpu_id 0

你可以将 ./demo/ARAD_1K_0912.jpg 替换为你的 RGB 图像路径。重建结果将保存至 /MST-plus-plus/predict_code/exp/ 目录下。

7. 可视化

  • 将重建的 HSI 放入 visualization/simulation_results/results/ 目录中。

  • 生成重建 HSI 的 RGB 图像

cd visualization/
Run show_simulation.m

引用

如果本仓库对您有所帮助,请考虑引用我们的研究成果:


# MST
@inproceedings{mst,
  title={Mask-guided Spectral-wise Transformer for Efficient Hyperspectral Image Reconstruction},
  author={Yuanhao Cai and Jing Lin and Xiaowan Hu and Haoqian Wang and Xin Yuan and Yulun Zhang and Radu Timofte and Luc Van Gool},
  booktitle={CVPR},
  year={2022}
}


# MST++
@inproceedings{mst_pp,
  title={MST++: Multi-stage Spectral-wise Transformer for Efficient Spectral Reconstruction},
  author={Yuanhao Cai and Jing Lin and Zudi Lin and Haoqian Wang and Yulun Zhang and Hanspeter Pfister and Radu Timofte and Luc Van Gool},
  booktitle={CVPRW},
  year={2022}
}


# HDNet
@inproceedings{hdnet,
  title={HDNet: High-resolution Dual-domain Learning for Spectral Compressive Imaging},
  author={Xiaowan Hu and Yuanhao Cai and Jing Lin and  Haoqian Wang and Xin Yuan and Yulun Zhang and Radu Timofte and Luc Van Gool},
  booktitle={CVPR},
  year={2022}
}

Introduction

"MST++: Multi-stage Spectral-wise Transformer for Efficient Spectral Reconstruction" (CVPRW 2022) & (Winner of NTIRE 2022 Spectral Recovery Challenge) and a toolbox for spectral reconstruction

Customize your domain