Skip to content

About

基于条件Attention U-Net的牙齿分割项目

Resources

Stars

3 stars

Watchers

0 watching

Forks

Repository files navigation

牙齿分割项目 - Conditional Attention U-Net

基于条件Attention U-Net的牙齿分割深度学习项目。

📋 项目简介

本项目实现了一个条件Attention U-Net模型,用于牙齿分割任务。模型接受牙齿ID作为条件输入,结合attention机制进行精准的牙齿实例分割。

主要特性

  • ✅ 条件输入: 使用tooth_id作为条件引导分割
  • ✅ Attention机制: 在U-Net的skip connection上应用attention gate
  • ✅ 混合精度训练: 支持AMP训练,降低显存占用,加速训练
  • ✅ 梯度累积: 在小显存GPU上实现大批次训练效果
  • ✅ 完整评估指标: 支持IoU、Dice、mPA、mIoU等多种指标
  • ✅ 数据增强: 使用Albumentations进行高效数据增强

🏗️ 项目结构

tooth_seg/
├── config.py                           # 配置文件
├── dataset.py                          # 数据集加载
├── transforms.py                       # 数据增强
├── models/
│   ├── __init__.py
│   └── attention_unet_conditional.py   # 条件Attention U-Net模型
├── losses.py                           # 损失函数(BCE + Dice)
├── metrics.py                          # 评估指标
├── train.py                            # 训练脚本
├── test.py                             # 测试脚本
├── visualize.py                        # 可视化工具
├── utils.py                            # 工具函数
├── requirements.txt                    # 依赖包
└── FIXES_AND_OPTIMIZATIONS.md         # 修复说明文档

🚀 快速开始

1. 环境配置

# 创建conda环境
conda create -n tooth_env python=3.9
conda activate tooth_env

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

2. 数据准备

将数据集放置在项目根目录:

tooth_seg/
├── trainset_valset/     # 训练和验证集
│   ├── 0/              # 牙齿类别0 (FDI 11号牙)
│   ├── 1/              # 牙齿类别1 (FDI 12号牙)
│   ├── 2/              # 牙齿类别2 (FDI 13号牙)
│   └── 3/              # 牙齿类别3 (FDI 14号牙)
└── testset/            # 测试集

3. 训练模型

python train.py

训练配置(可在config.py中修改):

  • 批次大小: 4 (有效批次=8,梯度累积)
  • 学习率: 1e-4
  • 训练轮数: 50
  • 图像大小: 512×512
  • 混合精度: 启用 (USE_AMP=True)

4. 测试模型

python test.py --checkpoint checkpoints/best_model.pth

📊 模型架构

Conditional Attention U-Net

输入: RGB图像 (3, 512, 512) + Tooth ID
         ↓
    [Tooth Embedding]
         ↓
    [Encoder Path]
    ├── Conv Block 1 (64) + Tooth Condition
    ├── Conv Block 2 (128) + Tooth Condition
    ├── Conv Block 3 (256) + Tooth Condition
    └── Conv Block 4 (512) + Tooth Condition
         ↓
    [Bottleneck] (512)
         ↓
    [Decoder Path with Attention]
    ├── Up Block 4 (512) + Attention Gate
    ├── Up Block 3 (256) + Attention Gate
    ├── Up Block 2 (128) + Attention Gate
    └── Up Block 1 (64) + Attention Gate
         ↓
输出: Segmentation Mask (1, 512, 512)

参数量: ~21M

📈 训练技巧

内存优化

针对6GB显存GPU的优化配置:

  1. 降低批次大小: BATCH_SIZE = 4
  2. 梯度累积: ACCUMULATION_STEPS = 2
  3. 混合精度训练: USE_AMP = True
  4. 降低workers: NUM_WORKERS = 2

预期效果:

  • 显存占用: 2.5-3GB
  • 训练速度: 提升1.5-2倍
  • 有效批次大小: 保持为8

如果仍遇到OOM

# 进一步降低批次大小
BATCH_SIZE = 2
ACCUMULATION_STEPS = 4

# 或降低图像分辨率
IMAGE_SIZE = 384

📝 评估指标

  • IoU (Intersection over Union)
  • Dice coefficient
  • mPA (Mean Pixel Accuracy)
  • Pixel Accuracy
  • Precision/Recall/F1

🔧 已知问题修复

详见 FIXES_AND_OPTIMIZATIONS.md

主要修复:

  1. ✅ 移除verbose参数(兼容新版PyTorch)
  2. ✅ 优化显存占用配置
  3. ✅ 实现混合精度训练
  4. ✅ 修复tensor转numpy错误

📚 依赖包

主要依赖:

  • PyTorch >= 1.12.0
  • torchvision
  • numpy
  • opencv-python
  • albumentations
  • tqdm

完整列表见 requirements.txt

📄 License

MIT License

👥 作者

牙齿分割项目组

🙏 致谢

About

基于条件Attention U-Net的牙齿分割项目

Resources

Stars

3 stars

Watchers

0 watching

Forks

Releases

Packages

Contributors

Languages