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

Python实现单Epoch训练多个Pickle图像数据集的高效方法

实现每个Epoch遍历全部Pickle训练文件的高效方案

你现在的问题是把大型图像数据集拆成了3个Pickle文件,但训练代码每个Epoch只训练其中一个,要改成每个Epoch都跑完所有三个文件的高效方案,下面是优化后的代码和逻辑:

一、优化数据读取逻辑

先把重复的加载代码改成循环,把所有训练对和标签存到列表里,后续遍历更方便:

import pickle
import numpy as np
import os

path = "你的实际路径/"  # 替换成你的数据集路径

train_pairs = []
train_labels = []

# 循环加载3份Pickle数据
for i in range(1, 4):
    # 加载训练对并处理形状
    with open(f"{path}TrainPairs{i}.pickle", "rb") as pickle_in:
        trainPixel = pickle.load(pickle_in)
        trainPixel = np.asarray(trainPixel)
        tr_pairs = trainPixel.reshape(trainPixel.shape[0], trainPixel.shape[1], 71, 71, 1)
        train_pairs.append(tr_pairs)
    
    # 加载对应标签
    with open(f"{path}TrainLabels{i}.pickle", "rb") as pickle_lb:
        tr_y = pickle.load(pickle_lb)
        tr_y = np.array(tr_y)
        train_labels.append(tr_y)

二、修改训练循环实现全数据遍历

调整训练逻辑,每个Epoch内依次遍历所有Pickle文件训练,同时保留断点续训功能:

# 训练前先加载预训练权重(仅加载一次)
weight_save_dir = f"{path}Saved_Weights"
weight_path = f"{weight_save_dir}/weights.ckpt"

if os.path.exists(weight_save_dir):
    print("加载预训练权重...")
    model.load_weights(weight_path)
else:
    print("无预训练权重,从头开始训练")
    os.makedirs(weight_save_dir, exist_ok=True)  # 提前创建权重保存目录

training_loss = []
test_loss = []
total_epochs = 50

for epoch in range(total_epochs):
    print(f"\n=== 第 {epoch+1}/{total_epochs} 个Epoch 开始 ===")
    
    # 逐个训练所有Pickle数据
    for batch_idx, (pairs, labels) in enumerate(zip(train_pairs, train_labels), 1):
        print(f"正在训练第 {batch_idx} 份Pickle数据")
        history = model.fit(
            [pairs[:, 0], pairs[:, 1]], 
            labels,
            batch_size=128,
            epochs=1,  # 当前子训练只跑1轮
            shuffle=True,
            validation_data=([te_pairs[:, 0], te_pairs[:, 1]], te_y)
        )
        
        # 累加损失记录
        training_loss.extend(history.history['loss'])
        test_loss.extend(history.history['val_loss'])
    
    # 每个Epoch结束后保存模型和权重
    base_network.save(f"{path}my_model")
    model.save_weights(weight_path)
    print(f"第 {epoch+1} 个Epoch结束,权重已保存")

关键优化点说明

  • 循环加载数据:避免重复代码,后续新增Pickle文件只需修改循环范围
  • 单Epoch多批次训练:每个Epoch内遍历所有数据,保证模型每次都能看到全部训练样本
  • 权重仅初始化加载一次:原代码每个Epoch都加载权重是冗余操作,现在只在训练开始时加载一次,减少开销
  • 损失正确累加:用extend替代原代码的+=,确保每个子训练步骤的损失都被完整记录

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.13 09:25:26