| """
|
| Transformer Decoder 模块 — Person B 负责实现
|
|
|
| 包含:
|
| - TransformerDecoderLayer: 单层解码器
|
| - TransformerDecoder: 多层解码器堆叠
|
|
|
| 架构 (Pre-LayerNorm):
|
| x → LN → Masked Self-Attention → Residual
|
| → LN → Cross-Attention → Residual
|
| → LN → FFN → Residual
|
| """
|
|
|
| from __future__ import annotations
|
|
|
| import copy
|
|
|
| import torch
|
| import torch.nn as nn
|
| from typing import Optional
|
|
|
| from easytranslate.model.attention import MultiHeadAttention, FlashMultiHeadAttention
|
|
|
|
|
| class TransformerDecoderLayer(nn.Module):
|
| """
|
| 单层 Transformer Decoder。
|
|
|
| 架构 (Pre-LayerNorm):
|
| x → LN → Masked Self-Attention → Residual
|
| → LN → Cross-Attention → Residual
|
| → LN → FFN → Residual
|
| """
|
|
|
| def __init__(
|
| self,
|
| d_model: int = 512,
|
| nhead: int = 8,
|
| dim_feedforward: int = 2048,
|
| dropout: float = 0.1,
|
| activation: str = "gelu",
|
| use_flash_attention: bool = True,
|
| use_rotary_embedding: bool = True,
|
| pre_norm: bool = True,
|
| ):
|
| super().__init__()
|
| self.pre_norm = pre_norm
|
| self.d_model = d_model
|
|
|
| attn_cls = FlashMultiHeadAttention if use_flash_attention else MultiHeadAttention
|
|
|
|
|
| self.self_attn = attn_cls(
|
| d_model=d_model,
|
| nhead=nhead,
|
| dropout=dropout,
|
| use_rotary_embedding=use_rotary_embedding,
|
| )
|
|
|
|
|
| self.multihead_attn = attn_cls(
|
| d_model=d_model,
|
| nhead=nhead,
|
| dropout=dropout,
|
| use_rotary_embedding=False,
|
| )
|
|
|
|
|
| self.linear1 = nn.Linear(d_model, dim_feedforward)
|
| self.activation = nn.GELU() if activation == "gelu" else nn.ReLU()
|
| self.dropout = nn.Dropout(p=dropout)
|
| self.linear2 = nn.Linear(dim_feedforward, d_model)
|
|
|
|
|
| self.norm1 = nn.LayerNorm(d_model)
|
| self.norm2 = nn.LayerNorm(d_model)
|
| self.norm3 = nn.LayerNorm(d_model)
|
|
|
|
|
| self.dropout1 = nn.Dropout(p=dropout)
|
| self.dropout2 = nn.Dropout(p=dropout)
|
| self.dropout3 = nn.Dropout(p=dropout)
|
|
|
| def forward(
|
| self,
|
| tgt: torch.Tensor,
|
| memory: torch.Tensor,
|
| tgt_mask: Optional[torch.Tensor] = None,
|
| memory_key_padding_mask: Optional[torch.BoolTensor] = None,
|
| tgt_key_padding_mask: Optional[torch.BoolTensor] = None,
|
| ) -> torch.Tensor:
|
| """Pre-LayerNorm 前向传播。"""
|
|
|
| is_causal = tgt_mask is not None
|
|
|
| if self.pre_norm:
|
|
|
| residual = tgt
|
| tgt = self.norm1(tgt)
|
| tgt = self.self_attn(
|
| tgt, tgt, tgt,
|
| key_padding_mask=tgt_key_padding_mask,
|
| is_causal=is_causal,
|
| )
|
| tgt = residual + self.dropout1(tgt)
|
|
|
|
|
| residual = tgt
|
| tgt = self.norm2(tgt)
|
| tgt = self.multihead_attn(
|
| tgt, memory, memory,
|
| key_padding_mask=memory_key_padding_mask,
|
| )
|
| tgt = residual + self.dropout2(tgt)
|
|
|
|
|
| residual = tgt
|
| tgt = self.norm3(tgt)
|
| tgt = self.linear2(self.dropout(self.activation(self.linear1(tgt))))
|
| tgt = residual + self.dropout3(tgt)
|
| else:
|
|
|
| residual = tgt
|
| tgt = self.self_attn(
|
| tgt, tgt, tgt,
|
| key_padding_mask=tgt_key_padding_mask,
|
| is_causal=is_causal,
|
| )
|
| tgt = self.norm1(residual + self.dropout1(tgt))
|
|
|
| residual = tgt
|
| tgt = self.multihead_attn(
|
| tgt, memory, memory,
|
| key_padding_mask=memory_key_padding_mask,
|
| )
|
| tgt = self.norm2(residual + self.dropout2(tgt))
|
|
|
| residual = tgt
|
| tgt = self.linear2(self.dropout(self.activation(self.linear1(tgt))))
|
| tgt = self.norm3(residual + self.dropout3(tgt))
|
|
|
| return tgt
|
|
|
|
|
| class TransformerDecoder(nn.Module):
|
| """
|
| 多层 Transformer Decoder。
|
| """
|
|
|
| def __init__(self, decoder_layer: TransformerDecoderLayer, num_layers: int):
|
| super().__init__()
|
| self.layers = nn.ModuleList(
|
| [copy.deepcopy(decoder_layer) for _ in range(num_layers)]
|
| )
|
| self.num_layers = num_layers
|
| self.norm = nn.LayerNorm(decoder_layer.d_model)
|
|
|
| def forward(
|
| self,
|
| tgt: torch.Tensor,
|
| memory: torch.Tensor,
|
| tgt_mask: Optional[torch.Tensor] = None,
|
| memory_key_padding_mask: Optional[torch.BoolTensor] = None,
|
| tgt_key_padding_mask: Optional[torch.BoolTensor] = None,
|
| ) -> torch.Tensor:
|
| output = tgt
|
| for layer in self.layers:
|
| output = layer(
|
| output,
|
| memory,
|
| tgt_mask=tgt_mask,
|
| memory_key_padding_mask=memory_key_padding_mask,
|
| tgt_key_padding_mask=tgt_key_padding_mask,
|
| )
|
| output = self.norm(output)
|
| return output
|
|
|