如何在PyTorch中实现LazyTransformer及相关Lazy化编码解码层?
自定义PyTorch LazyTransformer及相关层实现
PyTorch的Lazy系列层允许我们在不提前指定输入特征维度的情况下初始化模块,官方未提供Transformer相关的Lazy实现,我们可以基于LazyModuleMixin和Transformer核心逻辑自定义实现。
核心思路
利用PyTorch的LazyModuleMixin特性,在模块第一次前向传播时,根据输入张量的形状自动推断特征维度d_model,并完成所有子模块(多头注意力、层归一化、前馈网络等)的初始化。
实现LazyTransformerEncoderLayer
对应官方TransformerEncoderLayer的延迟初始化版本,无需提前指定d_model:
import torch import torch.nn as nn import torch.nn.functional as F class LazyTransformerEncoderLayer(nn.Module, nn.LazyModuleMixin): def __init__(self, nhead, dim_feedforward=2048, dropout=0.1, activation=F.relu, layer_norm_eps=1e-5, batch_first=False, norm_first=False, device=None, dtype=None): factory_kwargs = {'device': device, 'dtype': dtype} super().__init__() self.nhead = nhead self.dim_feedforward = dim_feedforward self.dropout = dropout self.activation = activation self.layer_norm_eps = layer_norm_eps self.batch_first = batch_first self.norm_first = norm_first # 延迟初始化的子模块,初始设为None self.self_attn = None self.norm1 = None self.norm2 = None self.dropout1 = nn.Dropout(dropout, **factory_kwargs) self.dropout2 = nn.Dropout(dropout, **factory_kwargs) self.linear1 = None self.linear2 = None self.dropout3 = nn.Dropout(dropout, **factory_kwargs) def initialize_parameters(self, input): # 从输入推断特征维度d_model d_model = input.size(-1) if self.batch_first else input.size(1) # 初始化多头注意力 self.self_attn = nn.MultiheadAttention(d_model, self.nhead, dropout=self.dropout, batch_first=self.batch_first, **self.factory_kwargs) # 初始化层归一化 self.norm1 = nn.LazyLayerNorm(d_model, eps=self.layer_norm_eps, **self.factory_kwargs) self.norm2 = nn.LazyLayerNorm(d_model, eps=self.layer_norm_eps, **self.factory_kwargs) # 初始化前馈网络 self.linear1 = nn.LazyLinear(self.dim_feedforward, **self.factory_kwargs) self.linear2 = nn.LazyLinear(d_model, **self.factory_kwargs) # 同步设备和数据类型 for module in [self.self_attn, self.norm1, self.norm2, self.linear1, self.linear2]: module.to(input.device, input.dtype) def forward(self, src, src_mask=None, src_key_padding_mask=None): # 第一次前向传播时完成参数初始化 if self.self_attn is None: self.initialize_parameters(src) x = src if self.norm_first: x = x + self._sa_block(self.norm1(x), src_mask, src_key_padding_mask) x = x + self._ff_block(self.norm2(x)) else: x = self.norm1(x + self._sa_block(x, src_mask, src_key_padding_mask)) x = self.norm2(x + self._ff_block(x)) return x # 自注意力计算块 def _sa_block(self, x, attn_mask, key_padding_mask): x = self.self_attn(x, x, x, attn_mask=attn_mask, key_padding_mask=key_padding_mask, need_weights=False)[0] return self.dropout1(x) # 前馈网络计算块 def _ff_block(self, x): x = self.linear2(self.dropout3(self.activation(self.linear1(x)))) return self.dropout2(x)
实现LazyTransformerDecoderLayer
对应官方TransformerDecoderLayer的延迟初始化版本,支持自注意力和交叉注意力的延迟初始化:
class LazyTransformerDecoderLayer(nn.Module, nn.LazyModuleMixin): def __init__(self, nhead, dim_feedforward=2048, dropout=0.1, activation=F.relu, layer_norm_eps=1e-5, batch_first=False, norm_first=False, device=None, dtype=None): factory_kwargs = {'device': device, 'dtype': dtype} super().__init__() self.nhead = nhead self.dim_feedforward = dim_feedforward self.dropout = dropout self.activation = activation self.layer_norm_eps = layer_norm_eps self.batch_first = batch_first self.norm_first = norm_first # 延迟初始化的子模块 self.self_attn = None self.multihead_attn = None self.norm1 = None self.norm2 = None self.norm3 = None self.dropout1 = nn.Dropout(dropout, **factory_kwargs) self.dropout2 = nn.Dropout(dropout, **factory_kwargs) self.dropout3 = nn.Dropout(dropout, **factory_kwargs) self.linear1 = None self.linear2 = None self.dropout4 = nn.Dropout(dropout, **factory_kwargs) def initialize_parameters(self, tgt, memory): # 推断特征维度,确保tgt和memory的特征维度一致 d_model_tgt = tgt.size(-1) if self.batch_first else tgt.size(1) d_model_mem = memory.size(-1) if self.batch_first else memory.size(1) assert d_model_tgt == d_model_mem, "输入序列与记忆序列的特征维度必须一致" d_model = d_model_tgt # 初始化自注意力和交叉注意力 self.self_attn = nn.MultiheadAttention(d_model, self.nhead, dropout=self.dropout, batch_first=self.batch_first, **self.factory_kwargs) self.multihead_attn = nn.MultiheadAttention(d_model, self.nhead, dropout=self.dropout, batch_first=self.batch_first, **self.factory_kwargs) # 初始化层归一化 self.norm1 = nn.LazyLayerNorm(d_model, eps=self.layer_norm_eps, **self.factory_kwargs) self.norm2 = nn.LazyLayerNorm(d_model, eps=self.layer_norm_eps, **self.factory_kwargs) self.norm3 = nn.LazyLayerNorm(d_model, eps=self.layer_norm_eps, **self.factory_kwargs) # 初始化前馈网络 self.linear1 = nn.LazyLinear(self.dim_feedforward, **self.factory_kwargs) self.linear2 = nn.LazyLinear(d_model, **self.factory_kwargs) # 同步设备和数据类型 for module in [self.self_attn, self.multihead_attn, self.norm1, self.norm2, self.norm3, self.linear1, self.linear2]: module.to(tgt.device, tgt.dtype) def forward(self, tgt, memory, tgt_mask=None, memory_mask=None, tgt_key_padding_mask=None, memory_key_padding_mask=None): # 第一次前向传播时完成参数初始化 if self.self_attn is None: self.initialize_parameters(tgt, memory) x = tgt if self.norm_first: x = x + self._sa_block(self.norm1(x), tgt_mask, tgt_key_padding_mask) x = x + self._mha_block(self.norm2(x), memory, memory_mask, memory_key_padding_mask) x = x + self._ff_block(self.norm3(x)) else: x = self.norm1(x + self._sa_block(x, tgt_mask, tgt_key_padding_mask)) x = self.norm2(x + self._mha_block(x, memory, memory_mask, memory_key_padding_mask)) x = self.norm3(x + self._ff_block(x)) return x # 自注意力计算块 def _sa_block(self, x, attn_mask, key_padding_mask): x = self.self_attn(x, x, x, attn_mask=attn_mask, key_padding_mask=key_padding_mask, need_weights=False)[0] return self.dropout1(x) # 交叉注意力计算块 def _mha_block(self, x, mem, attn_mask, key_padding_mask): x = self.multihead_attn(x, mem, mem, attn_mask=attn_mask, key_padding_mask=key_padding_mask, need_weights=False)[0] return self.dropout2(x) # 前馈网络计算块 def _ff_block(self, x): x = self.linear2(self.dropout4(self.activation(self.linear1(x)))) return self.dropout3(x)
实现LazyTransformer
基于上面的Lazy编码器/解码器层,组合成完整的Transformer模型:
class LazyTransformer(nn.Module, nn.LazyModuleMixin): def __init__(self, nhead, num_encoder_layers=6, num_decoder_layers=6, dim_feedforward=2048, dropout=0.1, activation=F.relu, layer_norm_eps=1e-5, batch_first=False, norm_first=False, device=None, dtype=None): factory_kwargs = {'device': device, 'dtype': dtype} super().__init__() self.nhead = nhead self.num_encoder_layers = num_encoder_layers self.num_decoder_layers = num_decoder_layers self.dim_feedforward = dim_feedforward self.dropout = dropout self.activation = activation self.layer_norm_eps = layer_norm_eps self.batch_first = batch_first self.norm_first = norm_first # 延迟初始化的编码器和解码器 self.encoder = None self.decoder = None def initialize_parameters(self, src, tgt): # 推断特征维度 d_model = src.size(-1) if self.batch_first else src.size(1) # 构建编码器层 encoder_layers = [LazyTransformerEncoderLayer( self.nhead, self.dim_feedforward, self.dropout, self.activation, self.layer_norm_eps, self.batch_first, self.norm_first, **self.factory_kwargs ) for _ in range(self.num_encoder_layers)] self.encoder = nn.TransformerEncoder(nn.ModuleList(encoder_layers), num_layers=self.num_encoder_layers) # 构建解码器层 decoder_layers = [LazyTransformerDecoderLayer( self.nhead, self.dim_feedforward, self.dropout, self.activation, self.layer_norm_eps, self.batch_first, self.norm_first, **self.factory_kwargs ) for _ in range(self.num_decoder_layers)] self.decoder = nn.TransformerDecoder(nn.ModuleList(decoder_layers), num_layers=self.num_decoder_layers) # 同步设备和数据类型 self.encoder.to(src.device, src.dtype) self.decoder.to(src.device, src.dtype) def forward(self, src, tgt, src_mask=None, tgt_mask=None, memory_mask=None, src_key_padding_mask=None, tgt_key_padding_mask=None, memory_key_padding_mask=None): # 第一次前向传播时完成参数初始化 if self.encoder is None: self.initialize_parameters(src, tgt) memory = self.encoder(src, mask=src_mask, src_key_padding_mask=src_key_padding_mask) output = self.decoder(tgt, memory, tgt_mask=tgt_mask, memory_mask=memory_mask, tgt_key_padding_mask=tgt_key_padding_mask, memory_key_padding_mask=memory_key_padding_mask) return output
使用示例
无需提前指定d_model,第一次前向传播时自动完成初始化:
# 生成测试输入 batch_size = 2 seq_len_src = 10 seq_len_tgt = 8 feat_dim = 512 # 仅用于生成测试数据,实际使用时无需提前指定 src = torch.randn(batch_size, seq_len_src, feat_dim) # batch_first=True模式 tgt = torch.randn(batch_size, seq_len_tgt, feat_dim) # 初始化LazyTransformer model = LazyTransformer(nhead=8, batch_first=True) output = model(src, tgt) print(f"输入src形状: {src.shape}") print(f"输入tgt形状: {tgt.shape}") print(f"输出形状: {output.shape}") # 输出形状与tgt一致
注意事项
- 初始化触发条件:所有Lazy模块必须通过第一次前向传播触发参数初始化,不能仅通过调用权重初始化方法完成。
- 维度一致性:确保编码器输入
src和解码器输入tgt的特征维度一致,否则初始化会抛出断言错误。 - 模型保存:保存模型前必须完成至少一次前向传播,否则保存的模型会缺少参数。
内容的提问来源于stack exchange,提问作者Rylan Schaeffer
相关产品推荐
相关产品推荐

