You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

PyTorch实现磁盘数据分批训练示例(类似Keras fit_generator)

在PyTorch中实现类似Keras fit_generator的磁盘分批加载CSV训练

刚好我之前做过类似的需求,完全懂你想要的那种「不用把全量CSV塞进内存,分批从磁盘读着训」的效果,就像Keras的fit_generator那样对吧?下面我一步步给你拆解实现,从自定义数据集到完整的训练循环都给你写清楚,保证适配你的需求:

1. 先写个能按需读CSV的自定义Dataset类

这是实现「不加载全量数据」的核心——我们只在需要某条样本的时候才从磁盘读取它,而不是一次性把整个CSV读进内存。这里我给你写两种实现,第一种是基础版,第二种是优化了读取速度的进阶版:

基础版(适合小数据集,逻辑简单)

import torch
from torch.utils.data import Dataset, DataLoader
import csv

class CSVDataSet(Dataset):
    def __init__(self, csv_path, label_col, transform=None):
        self.csv_path = csv_path
        self.label_col = label_col  # 指定哪一列是标签
        self.transform = transform  # 可选的预处理逻辑
        
        # 先统计总样本数,顺便拿到特征列名
        with open(csv_path, 'r', newline='') as f:
            reader = csv.DictReader(f)
            self.feature_cols = [col for col in reader.fieldnames if col != label_col]
            self.total_samples = sum(1 for _ in reader)  # 跳过表头数行数
    
    def __len__(self):
        return self.total_samples  # DataLoader需要知道总样本数
    
    def __getitem__(self, idx):
        # 读取指定索引的样本行
        with open(self.csv_path, 'r', newline='') as f:
            reader = csv.DictReader(f)
            for i, row in enumerate(reader):
                if i == idx:
                    # 把特征转成Tensor
                    features = torch.tensor([float(row[col]) for col in self.feature_cols], dtype=torch.float32)
                    # 把标签转成Tensor(根据任务调整 dtype,分类任务可以用long)
                    label = torch.tensor(float(row[self.label_col]), dtype=torch.float32)
                    
                    if self.transform:
                        features = self.transform(features)
                    
                    return features, label

进阶版(适合大数据集,读取速度更快)

基础版每次找样本都要从头遍历CSV,大数据集效率低。进阶版会预先记录每行数据在文件中的偏移量,读取时直接跳转到对应位置,速度快很多:

class CSVDataSet(Dataset):
    def __init__(self, csv_path, label_col, transform=None):
        self.csv_path = csv_path
        self.label_col = label_col
        self.transform = transform
        
        # 预先记录每行数据的文件偏移量(跳过表头)
        self.row_offsets = []
        with open(csv_path, 'r', newline='') as f:
            header = f.readline()  # 读表头
            self.row_offsets.append(f.tell())  # 第一行数据的起始位置
            # 遍历所有行,记录每个行的起始偏移
            while f.readline():
                self.row_offsets.append(f.tell())
        self.total_samples = len(self.row_offsets) - 1  # 最后一个偏移是文件末尾,所以减1
        
        # 拿到特征列名
        reader = csv.DictReader([header])
        self.feature_cols = [col for col in reader.fieldnames if col != label_col]
    
    def __len__(self):
        return self.total_samples
    
    def __getitem__(self, idx):
        with open(self.csv_path, 'r', newline='') as f:
            # 直接跳转到目标行的位置
            f.seek(self.row_offsets[idx])
            row = f.readline().strip()
            # 解析行数据成字典
            row_dict = dict(zip(self.feature_cols + [self.label_col], row.split(',')))
            
            features = torch.tensor([float(row_dict[col]) for col in self.feature_cols], dtype=torch.float32)
            label = torch.tensor(float(row_dict[self.label_col]), dtype=torch.float32)
            
            if self.transform:
                features = self.transform(features)
            
            return features, label

2. 用DataLoader把数据集包成批次加载器

有了自定义Dataset,就可以用PyTorch的DataLoader自动处理分批、打乱顺序、多进程读取这些事情,完全对标Keras的fit_generator:

# 训练集加载器:shuffle=True 表示每个epoch打乱样本顺序,num_workers用多进程加速读取
train_dataset = CSVDataSet(csv_path='train_data.csv', label_col='target')
train_loader = DataLoader(train_dataset, batch_size=32, shuffle=True, num_workers=4)

# 验证集加载器:验证时不需要打乱顺序,所以shuffle=False
val_dataset = CSVDataSet(csv_path='val_data.csv', label_col='target')
val_loader = DataLoader(val_dataset, batch_size=32, shuffle=False, num_workers=4)

3. 写一个类似fit_generator的训练验证循环

接下来就是写训练循环,重复指定轮数,每轮遍历训练批次,然后在验证集上评估效果:

def train_model(model, train_loader, val_loader, criterion, optimizer, num_epochs):
    for epoch in range(num_epochs):
        # 训练模式:打开Dropout、BatchNorm等训练专属层
        model.train()
        train_total_loss = 0.0
        
        for batch_idx, (features, labels) in enumerate(train_loader):
            # 前向传播计算输出
            outputs = model(features)
            # 计算损失(这里假设是回归任务,标签要加维度匹配输出;分类任务自行调整)
            loss = criterion(outputs, labels.unsqueeze(1))
            
            # 反向传播+优化
            optimizer.zero_grad()
            loss.backward()
            optimizer.step()
            
            # 累计损失
            train_total_loss += loss.item() * features.size(0)
            
            # 每100个批次打印一次进度
            if batch_idx % 100 == 0:
                print(f'Epoch {epoch+1}/{num_epochs}, Batch {batch_idx}, 当前批次损失: {loss.item():.4f}')
        
        # 计算训练集平均损失
        avg_train_loss = train_total_loss / len(train_loader.dataset)
        
        # 验证模式:关闭Dropout等,不计算梯度
        model.eval()
        val_total_loss = 0.0
        correct = 0
        total = 0
        
        with torch.no_grad():  # 验证时不需要计算梯度,节省内存和时间
            for features, labels in val_loader:
                outputs = model(features)
                loss = criterion(outputs, labels.unsqueeze(1))
                val_total_loss += loss.item() * features.size(0)
                
                # 分类任务可以加准确率计算(回归任务注释掉这段)
                _, predicted = torch.max(outputs.data, 1)
                total += labels.size(0)
                correct += (predicted == labels).sum().item()
        
        avg_val_loss = val_total_loss / len(val_loader.dataset)
        val_acc = correct / total if total > 0 else 0.0
        
        # 打印本轮训练验证结果
        print(f'\n=== Epoch {epoch+1}/{num_epochs} ===')
        print(f'训练集平均损失: {avg_train_loss:.4f}, 验证集平均损失: {avg_val_loss:.4f}')
        if val_acc > 0:
            print(f'验证集准确率: {val_acc:.4f}')
        print('-' * 60)

然后你只需要初始化自己的模型、损失函数和优化器,调用这个函数就行:

# 示例模型(替换成你自己的模型)
class SimpleModel(torch.nn.Module):
    def __init__(self, input_size):
        super(SimpleModel, self).__init__()
        self.fc1 = torch.nn.Linear(input_size, 64)
        self.fc2 = torch.nn.Linear(64, 1)  # 回归任务输出1个值,分类任务改成类别数
    
    def forward(self, x):
        x = torch.relu(self.fc1(x))
        x = self.fc2(x)
        return x

# 初始化模型
input_size = len(train_dataset.feature_cols)  # 特征列的数量
model = SimpleModel(input_size)

# 损失函数和优化器(根据任务调整:回归用MSELoss,分类用CrossEntropyLoss)
criterion = torch.nn.MSELoss()
optimizer = torch.optim.Adam(model.parameters(), lr=0.001)

# 开始训练,比如训10轮
train_model(model, train_loader, val_loader, criterion, optimizer, num_epochs=10)

4. 关于自定义TestLoader的说明

你提到需要自定义TestLoader,其实逻辑和验证集完全一致:创建测试集的CSVDataSet实例,用DataLoader封装即可。如果需要测试时的特殊逻辑(比如预处理、结果保存),可以在Dataset里加transform,或者在遍历TestLoader时处理:

# 自定义测试集加载器
test_dataset = CSVDataSet(csv_path='test_data.csv', label_col='target')
test_loader = DataLoader(test_dataset, batch_size=32, shuffle=False, num_workers=2)

# 测试循环(比如保存预测结果)
model.eval()
test_predictions = []
test_labels = []

with torch.no_grad():
    for features, labels in test_loader:
        outputs = model(features)
        test_predictions.extend(outputs.cpu().numpy().flatten())
        test_labels.extend(labels.cpu().numpy().flatten())

# 后续可以把结果保存成CSV或者做其他分析

几个关键注意点

  • 多进程读取:num_workers设置大于0时,每个worker进程会独立打开CSV文件,我们的Dataset实现是安全的,不会出现文件指针冲突的问题。
  • 内存占用:整个流程不会把全量数据加载到内存,每次只读取当前批次的样本,完全适合超大CSV文件。
  • 预处理扩展:可以通过transform参数给Dataset加各种预处理逻辑(比如标准化、归一化),用法和PyTorch自带的图像数据集完全一致。

内容的提问来源于stack exchange,提问作者B_Miner

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.05.19 07:30:36