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

使用PyTorch Lightning 0.7.6如何获取全部训练批次的模型输出并保存

PyTorch Lightning 0.7.6 全量训练批次输出获取方案

方法1:直接通过训练阶段内置钩子实现(无需额外Callback)

0.7.6版本中,training_step 每次返回的所有值都会被框架自动收集,在 training_epoch_end 钩子中传入的 outputs 参数就是所有训练批次的输出集合,拿不到全量大概率是误用了钩子或变量名:

  • 第一步:在training_step中明确返回你需要保存的词向量,示例代码:
def training_step(self, batch, batch_idx):
    # 原有训练逻辑
    loss = ...
    word_emb = model(batch) # 你的模型输出词向量
    return {
        'loss': loss,
        'word_emb': word_emb # 必须显式返回需要收集的张量
    }
  • 第二步:在training_epoch_end中直接处理全量输出:
def training_epoch_end(self, outputs):
    # outputs是列表,每个元素对应一个batch的training_step返回值
    all_word_emb = []
    for batch_out in outputs:
        all_word_emb.append(batch_out['word_emb'].cpu().numpy())
    # 拼接所有批次结果保存到文件
    import numpy as np
    np.save('all_train_word_emb.npy', np.concatenate(all_word_emb, axis=0))
    # 原有epoch end的其他逻辑(比如打印指标等)
    return {'log': ...}

如果是多卡训练场景,需要先调用self.all_gather对张量做分布式聚合再保存。

方法2:通过Callback实现(适合解耦业务逻辑的场景)

如果不想把保存逻辑耦合在模型代码里,可以自定义Callback实现,0.7.6版本已经支持Callback的on_epoch_end钩子:

  • 自定义Callback示例:
from pytorch_lightning.callbacks import Callback
import numpy as np

class SaveEmbCallback(Callback):
    def on_epoch_end(self, trainer, pl_module):
        # 从trainer对象中获取当前epoch所有批次的输出
        outputs = trainer.train_loop.outputs
        all_word_emb = np.concatenate([x['word_emb'].cpu().numpy() for x in outputs], axis=0)
        np.save(f'epoch_{trainer.current_epoch}_word_emb.npy', all_word_emb)
  • 训练时传入Callback即可:
trainer = Trainer(
    callbacks=[SaveEmbCallback()],
    # 其他参数
)

常见问题排查

  • 如果还是只能拿到单批次输出,先检查是否把逻辑写在了training_step_end而非training_epoch_end,前者是单批次训练后的钩子,只会拿到单个批次的结果
  • 确认training_step没有漏返回需要收集的词向量字段
  • 0.7.6版本不支持自动将输出移动到CPU,必须显式调用.cpu()再保存,避免显存溢出

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.25 23:06:04