Skip to content
 
 

Latest commit

 

History

34 Commits

Folders and files

NameName
Last commit message
Last commit date
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 

Repository files navigation

基于扩散模型的中国古代书法字符修复

北京交通大学 计算机科学与技术学院
研究生课程《数字图像处理》大作业

姓名:柯劲帆
学号:25120323


📋 项目简介

本项目实现了一个基于去噪扩散概率模型 (DDPM) 的中国古代书法字符修复系统。该系统能够自动修复因年代久远而产生腐蚀、破损的古代拓片和手稿中的汉字图像。

主要特点

  • 🎨 扩散模型架构:采用基于 ConvNeXt 的 UNet 作为去噪网络
  • 高效训练:支持多 GPU 分布式训练,混合精度加速
  • 📊 全面评估:支持 PSNR、SSIM、LPIPS、FID 等多种图像质量评估指标
  • 🔧 灵活配置:通过 YAML 配置文件管理训练参数

🏗️ 项目结构

.
├── code/                       # 源代码目录
│   ├── train.py               # 训练入口
│   ├── evaluate.py            # 评估脚本
│   ├── model.py               # 模型初始化
│   ├── diffusion.py           # 扩散模型核心实现
│   ├── unet_convnext.py       # UNet 网络结构
│   ├── dataset.py             # 数据集加载
│   ├── trainer.py             # 训练器
│   ├── metrics.py             # 评估指标
│   └── args.py                # 参数解析
├── configs/                    # 配置文件目录
│   └── config_1.yaml          # 训练配置
├── scripts/                    # 启动脚本
│   ├── train.sh               # 训练脚本
│   └── evaluate.sh            # 评估脚本
├── data/                       # 数据目录
│   ├── processed_dataset/     # 预处理后的数据集
│   ├── raw_dataset/           # 原始数据集
│   ├── output/                # 模型输出(检查点)
│   └── evaluation/            # 评估结果
├── logs/                       # 日志目录
├── requirements.txt           # Python 依赖
└── README.md                  # 项目说明

🔧 环境配置

依赖安装

# 创建虚拟环境(推荐)
conda create -n diffusion python=3.10
conda activate diffusion

# 安装 PyTorch(根据 CUDA 版本选择)
pip install torch torchvision --index-url https://download.pytorch.org/whl/cu118

# 安装其他依赖
pip install -r requirements.txt

依赖列表

  • opencv-python - 图像处理
  • Pillow - 图像读取
  • transformers[torch] - Hugging Face 工具
  • datasets - 数据集管理
  • accelerate - 分布式训练加速
  • einops - 张量操作
  • safetensors - 模型存储
  • tqdm - 进度条
  • tensorboard - 训练可视化
  • scikit-image - 图像处理
  • lpips - 感知损失计算
  • pytorch-fid - FID 指标计算
  • torchmetrics[image] - 图像评估指标

🚀 使用方法

训练模型

  1. 准备数据集:将预处理后的数据放入 data/processed_dataset/data/ 目录,包含以下子目录:

    • real/ - 完整的真实字符图像(Ground Truth)
    • eroded/ - 腐蚀/破损的字符图像(输入条件)
    • mask/ - 可选的掩码图像
  2. 修改配置:根据需要修改 configs/config_1.yaml 中的参数:

    # 模型参数
    time_steps: 1000              # 扩散步数
    
    # 数据参数
    image_size: 224               # 图像尺寸
    input_dir: data/processed_dataset/data
    
    # 训练参数
    num_train_epochs: 20          # 训练轮数
    per_device_train_batch_size: 16
    learning_rate: 1.0e-4
    save_steps: 50                # 检查点保存间隔
  3. 启动训练

    # 方式一:使用脚本启动(推荐,支持多 GPU)
    bash scripts/train.sh
    
    # 方式二:直接启动(单 GPU)
    python code/train.py --config configs/config_1.yaml
  4. 查看训练日志

    # 使用 TensorBoard 可视化
    tensorboard --logdir data/output/config_1/runs

评估模型

  1. 运行评估脚本

    # 方式一:使用脚本启动
    bash scripts/evaluate.sh
    
    # 方式二:手动指定参数
    python code/evaluate.py \
        --checkpoint data/output/config_1/checkpoint_241/model.safetensors \
        --time_steps 1000 \
        --lpips_model_path data/evaluation_model/alexnet-owt-7be5be79.pth \
        --test_dir data/raw_dataset/dataset/dataset/testset \
        --output_dir data/evaluation/config_1 \
        --image_size 224 \
        --device cuda \
        --batch_size 64
  2. 评估指标说明

    • PSNR (Peak Signal-to-Noise Ratio):峰值信噪比,值越高越好
    • SSIM (Structural Similarity Index):结构相似性,值越高越好
    • LPIPS (Learned Perceptual Image Patch Similarity):感知相似度,值越低越好
    • FID (Fréchet Inception Distance):分布距离,值越低越好

📊 模型架构

扩散模型 (DDPM)

本项目采用去噪扩散概率模型进行图像修复:

  • 前向过程:逐步向图像添加高斯噪声
  • 反向过程:学习去除噪声,恢复原始图像
  • 噪声调度:使用余弦噪声调度(Cosine Schedule)

网络结构

  • 骨干网络:基于 ConvNeXt 的 UNet
  • 输入通道:2(条件图像 + 噪声图像,灰度)
  • 输出通道:1(预测噪声,灰度)
  • 特征维度:64, 128, 256, 512(4层下采样)

📁 数据集

本项目使用 ARMCD (Chinese Ancient Rubbing and Manuscript Character Dataset) 数据集,包含:

  • 15,553 张来自真实古代拓片和手稿的单字图像
  • 来源于 42 个拓片和手稿文献
  • 时间跨度:公元 200 年至 1800 年
  • 涵盖超过 200 位书法家的作品

详细数据集信息请参阅 Dataset.md

📝 配置文件说明

configs/config_1.yaml 主要参数:

参数 说明 默认值
time_steps 扩散时间步数 1000
image_size 输入图像尺寸 224
num_train_epochs 训练轮数 20
per_device_train_batch_size 每个 GPU 的批次大小 16
gradient_accumulation_steps 梯度累积步数 16
learning_rate 学习率 1e-4
ema_decay EMA 衰减率 0.9999
save_steps 检查点保存间隔 50
mixed_precision 混合精度训练 fp16

📄 许可证

本项目仅供学术研究使用。


About

bjtu graduated Digital Image Processing homework

Resources

Stars

0 stars

Watchers

0 watching

Forks

Releases

Packages

Contributors

Languages