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

如何在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一致

注意事项

  1. 初始化触发条件:所有Lazy模块必须通过第一次前向传播触发参数初始化,不能仅通过调用权重初始化方法完成。
  2. 维度一致性:确保编码器输入src和解码器输入tgt的特征维度一致,否则初始化会抛出断言错误。
  3. 模型保存:保存模型前必须完成至少一次前向传播,否则保存的模型会缺少参数。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.30 04:38:17