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

