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

如何在PyTorch中保存所有批次的MLP预测结果?

问题分析与解决方法

问题根源

你当前代码只保存了最后一个批次的预测结果,因为pred = x_proj.detach().cpu().numpy()仅获取了循环结束时x_proj的值(也就是最后一批的7个样本预测),之前所有批次的结果都被覆盖了,所以不管怎么调整batch_size,最终只能得到最后一批的输出。

解决代码

要保存所有样本的预测结果,需要用一个列表逐批次收集预测结果,最后合并成完整的结果再保存。以下是修改后的代码:

import numpy as np
import pandas as pd
import torch

length = len(X) // batch_size
print(length)

# 初始化列表,用于存储所有批次的预测结果
all_preds = []

for epoch in range(total_epoch):
    loss_sum = 0.0
    for batch_idx in range(length + 1):
        # 分批次取数据
        if batch_idx != length:
            x_ = X[batch_idx*batch_size : (batch_idx+1)*batch_size]
            y_ = Y[batch_idx*batch_size : (batch_idx+1)*batch_size]
        else:
            x_ = X[batch_idx*batch_size : ]
            y_ = Y[batch_idx*batch_size : ]

        # 模型前向传播
        x_proj = projector(x_)
        # 计算损失与反向传播
        loss = torch.nn.MSELoss()(x_proj, y_)
        loss_sum += loss.item()
        optimizer.zero_grad()
        loss.backward()
        optimizer.step()
        
        # 收集当前批次的预测结果:转CPU→去计算图→转numpy
        batch_pred = x_proj.detach().cpu().numpy()
        all_preds.append(batch_pred)
        
        if batch_idx % print_freq == 0:
            print(f"Epoch: {epoch}, Iteration: {batch_idx}, Loss: {loss.item():.4f}")

    print(f"Epoch: {epoch}, Loss Avg: {loss_sum/length: .4f}")
    # 可选:打印当前epoch已收集的预测样本总数
    print(f"Collected predictions count: {sum(len(batch) for batch in all_preds)}")

# 合并所有批次的结果为一个完整数组(形状为(98, 输出维度))
all_preds_np = np.concatenate(all_preds, axis=0)
# 保存到CSV文件
pd.DataFrame(all_preds_np).to_csv('prediction_all.csv', index=False)

关键修改点

  • 新增all_preds = []列表,用于逐批次存储每个batch的预测结果,避免被覆盖。
  • 在每个batch处理完成后,将x_proj转换为numpy格式并添加到列表中。
  • 训练结束后,用np.concatenate将列表中所有小批次数组合并成一个大数组,对应全部98个样本的预测结果。
  • 如果需要保存每个epoch的预测结果,可以在每个epoch循环的末尾添加合并与保存代码,同时在epoch开始时清空all_preds:
    # 在每个epoch结束时保存当前epoch的预测
    epoch_preds_np = np.concatenate(all_preds, axis=0)
    pd.DataFrame(epoch_preds_np).to_csv(f'prediction_epoch_{epoch}.csv', index=False)
    # 清空列表,准备下一个epoch的收集
    all_preds = []
    

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.27 10:42:49