这是一个基于 PyTorch 的深度学习训练框架,提供了一个抽象的 Trainer 类,可以方便地扩展用于监督学习和无监督学习任务。
- 支持监督和无监督学习:根据数据集参数自动判断训练模式
- 自动数据分割:将数据自动分割为训练集、验证集和测试集
- 完整的训练流程:包含训练、验证、评估和模型保存
- 日志记录:详细的训练日志和性能指标记录
- 可视化:训练过程的可视化图表生成
- 模型检查点:自动保存最佳模型和定期保存检查点
- 设备自动检测:自动检测并使用 CUDA(如果可用)
- 灵活的数据集支持:支持自定义数据集类
- 学习率调度:支持学习率调度器
- 批量和周期更新模式:支持按批次或周期更新学习率
Trainer (抽象基类)
├── __init__() - 初始化模型、数据集和设备
├── build_dataset() - 构建数据集(抽象方法)
├── split_data() - 数据分割(抽象方法)
├── iter_train() - 单批次训练(抽象方法)
├── iter_val() - 单批次验证(抽象方法)
├── evaluate() - 模型评估(抽象方法)
├── train() - 主要训练循环
└── plot_history() - 训练历史可视化# 建议将代码拆分为多个模块
├── base_trainer.py # 基础训练器类
├── data_utils.py # 数据处理工具
├── logger.py # 日志配置
├── metrics.py # 评估指标
├── visualizer.py # 可视化工具
└── utils.py # 通用工具函数import torch
import torch.nn as nn
from Trainer import Trainer
# 定义你的模型
class MyModel(nn.Module):
def __init__(self):
super().__init__()
self.layer = nn.Linear(10, 1)
def forward(self, x):
return self.layer(x)
# 创建模型和数据
model = MyModel()
data = torch.randn(1000, 10) # 1000个样本,每个样本10个特征
labels = torch.randn(1000, 1) # 监督学习需要标签
# 实例化训练器
trainer = MyTrainer(model, [data, labels]) # 监督学习
trainer = MyTrainer(model, [data]) # 无监督学习你需要继承 Trainer 类并实现抽象方法:
class MyTrainer(Trainer):
def build_dataset(self):
"""构建数据集"""
super().build_dataset()
# 可以在这里添加自定义数据集构建逻辑
# 可以在这里构建一个损失函数计算方式
def split_data(self, random_seed=42, rate=(0.7, 0.2, 0.1)):
"""数据分割"""
super().split_data(random_seed, rate)
# 可以在这里添加自定义数据分割逻辑
def iter_train(self, data=None, label=None) -> torch.Tensor:
"""单批次训练"""
# 实现你的训练逻辑
predictions = self.model(data)
loss = nn.MSELoss()(predictions, label)
return loss
def iter_val(self, data=None, label=None) -> torch.Tensor:
"""单批次验证"""
# 实现你的验证逻辑
with torch.no_grad():
predictions = self.model(data)
loss = nn.MSELoss()(predictions, label)
return loss
def evaluate(self, data=None, label=None) -> torch.Tensor:
"""模型评估"""
# 实现你的评估逻辑
return self.iter_val(data, label)# 创建优化器
optimizer = torch.optim.Adam(model.parameters(), lr=0.001)
# 可选:创建学习率调度器
scheduler = torch.optim.lr_scheduler.StepLR(optimizer, step_size=30, gamma=0.1)
# 开始训练
trainer.train(
epochs=100,
batch_size=32,
eval_epoch=10,
save_dir="logs",
optimizer_instance=optimizer,
scheduler_instance=scheduler,
update_mode="epoch"
)训练过程中会生成以下文件:
logs/
├── trainer.log # 训练日志
├── weights/
│ ├── best_model.pth # 最佳模型
│ └── model_{epoch}.pth # 定期保存的检查点
└── train_images/
├── history_{epoch}.png # 训练历史图表
└── finished.png # 最终训练图表