Skip to content

Latest commit

 

History

1 Commit

Folders and files

NameName
Last commit message
Last commit date
 
 
 
 

Repository files navigation

Deep Learning Trainer Framework

这是一个基于 PyTorch 的深度学习训练框架,提供了一个抽象的 Trainer 类,可以方便地扩展用于监督学习和无监督学习任务。

功能特性

核心功能

  • 支持监督和无监督学习:根据数据集参数自动判断训练模式
  • 自动数据分割:将数据自动分割为训练集、验证集和测试集
  • 完整的训练流程:包含训练、验证、评估和模型保存
  • 日志记录:详细的训练日志和性能指标记录
  • 可视化:训练过程的可视化图表生成
  • 模型检查点:自动保存最佳模型和定期保存检查点

技术特性

  • 设备自动检测:自动检测并使用 CUDA(如果可用)
  • 灵活的数据集支持:支持自定义数据集类
  • 学习率调度:支持学习率调度器
  • 批量和周期更新模式:支持按批次或周期更新学习率

代码结构重写说明

原始代码结构

Trainer (抽象基类)
├── __init__() - 初始化模型数据集和设备
├── build_dataset() - 构建数据集抽象方法)
├── split_data() - 数据分割抽象方法)
├── iter_train() - 单批次训练抽象方法)
├── iter_val() - 单批次验证抽象方法)
├── evaluate() - 模型评估抽象方法)
├── train() - 主要训练循环
└── plot_history() - 训练历史可视化

重写建议

1. 模块化重构

# 建议将代码拆分为多个模块
├── base_trainer.py      # 基础训练器类
├── data_utils.py        # 数据处理工具
├── logger.py           # 日志配置
├── metrics.py          # 评估指标
├── visualizer.py       # 可视化工具
└── utils.py            # 通用工具函数

使用方法

1. 基础使用

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])  # 无监督学习

2. 自定义训练器

你需要继承 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)

3. 开始训练

# 创建优化器
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        # 最终训练图表

About

This is a training framework.you need to rewrite some class method to finish the training task.

Resources

Stars

0 stars

Watchers

0 watching

Forks

Releases

Packages

Contributors

Languages