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

如何使PyTorch类Transformer模型在训练与推理阶段行为不同?

PyTorch Transformer训练/推理阶段缓存控制的规范方案

PyTorch的model.train()和model.eval()本质是切换模型的self.training布尔属性,完全可以基于这个机制来规范控制缓存的启用/禁用,替代零散的if-else判断。以下是两种常用的规范实现方案:

方案1:在forward方法中基于self.training直接控制

直接利用模型自带的self.training属性,在forward逻辑里分支处理训练和推理的缓存逻辑:

import torch
import torch.nn as nn

class CachedTransformer(nn.Module):
    def __init__(self, d_model=512, nhead=8, num_layers=6):
        super().__init__()
        self.transformer = nn.Transformer(d_model=d_model, nhead=nhead, num_encoder_layers=num_layers, num_decoder_layers=num_layers)
        # 用于存储推理阶段的缓存数据
        self.inference_cache = {}

    def forward(self, src, tgt):
        if self.training:
            # 训练阶段:禁用缓存,直接计算
            output = self.transformer(src, tgt)
            # 清空缓存,避免干扰后续推理
            self.inference_cache.clear()
            return output
        else:
            # 推理阶段:启用缓存加速
            if "prev_tgt" in self.inference_cache:
                # 拼接历史输入,复用已计算的缓存
                tgt = torch.cat([self.inference_cache["prev_tgt"], tgt], dim=0)
            output = self.transformer(src, tgt)
            # 更新缓存
            self.inference_cache["prev_tgt"] = tgt
            return output

调用model.train()后,self.training自动设为True,走无缓存的训练逻辑;调用model.eval()后设为False,自动启用缓存。

方案2:重写train/eval方法,显式绑定缓存状态

如果缓存逻辑复杂(比如需要维护多组缓存数据),可以重写模型的train()和eval()方法,在切换模式时同步处理缓存的初始化或清空:

import torch
import torch.nn as nn

class CachedTransformer(nn.Module):
    def __init__(self, d_model=512, nhead=8, num_layers=6):
        super().__init__()
        self.transformer = nn.Transformer(d_model=d_model, nhead=nhead, num_encoder_layers=num_layers, num_decoder_layers=num_layers)
        self.inference_cache = {}
        self._use_cache = False

    def train(self, mode=True):
        # 先调用父类的train方法,切换training状态
        super().train(mode)
        self._use_cache = False
        # 训练前强制清空缓存
        self.inference_cache.clear()

    def eval(self):
        # 先调用父类的eval方法,切换training状态
        super().eval()
        self._use_cache = True
        # 初始化推理所需的缓存结构
        self.inference_cache = {"prev_tgt": None, "encoder_memory": None}

    def forward(self, src, tgt):
        if self._use_cache:
            # 推理阶段:复用缓存
            if self.inference_cache["prev_tgt"] is not None:
                tgt = torch.cat([self.inference_cache["prev_tgt"], tgt], dim=0)
            # 假设Transformer返回encoder memory,用于后续解码复用
            output, encoder_memory = self.transformer(src, tgt, return_memory=True)
            self.inference_cache["prev_tgt"] = tgt
            self.inference_cache["encoder_memory"] = encoder_memory
            return output
        else:
            # 训练阶段:不使用缓存
            output = self.transformer(src, tgt)
            return output

这种方式把模式切换和缓存控制的逻辑解耦,代码结构更清晰,适合缓存逻辑复杂的场景。

额外注意事项

  • 训练阶段必须清空缓存,防止残留的缓存数据干扰训练计算,导致结果异常。
  • 如果使用PyTorch内置的nn.TransformerDecoder,它原生支持cache参数,可以直接在forward时根据self.training决定是否传入缓存,无需自行维护缓存容器。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.28 21:03:40