| """ |
| 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 |
|
|