使用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
相关产品推荐
相关产品推荐

