基于条件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 # 修复说明文档
# 创建conda环境
conda create -n tooth_env python=3.9
conda activate tooth_env
# 安装依赖
pip install -r requirements.txt将数据集放置在项目根目录:
tooth_seg/
├── trainset_valset/ # 训练和验证集
│ ├── 0/ # 牙齿类别0 (FDI 11号牙)
│ ├── 1/ # 牙齿类别1 (FDI 12号牙)
│ ├── 2/ # 牙齿类别2 (FDI 13号牙)
│ └── 3/ # 牙齿类别3 (FDI 14号牙)
└── testset/ # 测试集
python train.py训练配置(可在config.py中修改):
- 批次大小: 4 (有效批次=8,梯度累积)
- 学习率: 1e-4
- 训练轮数: 50
- 图像大小: 512×512
- 混合精度: 启用 (USE_AMP=True)
python test.py --checkpoint checkpoints/best_model.pth输入: 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的优化配置:
- 降低批次大小: BATCH_SIZE = 4
- 梯度累积: ACCUMULATION_STEPS = 2
- 混合精度训练: USE_AMP = True
- 降低workers: NUM_WORKERS = 2
预期效果:
- 显存占用: 2.5-3GB
- 训练速度: 提升1.5-2倍
- 有效批次大小: 保持为8
# 进一步降低批次大小
BATCH_SIZE = 2
ACCUMULATION_STEPS = 4
# 或降低图像分辨率
IMAGE_SIZE = 384- IoU (Intersection over Union)
- Dice coefficient
- mPA (Mean Pixel Accuracy)
- Pixel Accuracy
- Precision/Recall/F1
主要修复:
- ✅ 移除
verbose参数(兼容新版PyTorch) - ✅ 优化显存占用配置
- ✅ 实现混合精度训练
- ✅ 修复tensor转numpy错误
主要依赖:
- PyTorch >= 1.12.0
- torchvision
- numpy
- opencv-python
- albumentations
- tqdm
完整列表见 requirements.txt
MIT License
牙齿分割项目组