sdfjliom's picture
model (#4)
4d62693
Raw
History Blame Contribute Delete
5.62 kB
"""
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
# 1. Masked Self-Attention
self.self_attn = attn_cls(
d_model=d_model,
nhead=nhead,
dropout=dropout,
use_rotary_embedding=use_rotary_embedding,
)
# 2. Cross-Attention (decoder queries encoder memory)
self.multihead_attn = attn_cls(
d_model=d_model,
nhead=nhead,
dropout=dropout,
use_rotary_embedding=False, # cross-attention 不使用 RoPE
)
# 3. Feed-Forward Network
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)
# 4. LayerNorms
self.norm1 = nn.LayerNorm(d_model)
self.norm2 = nn.LayerNorm(d_model)
self.norm3 = nn.LayerNorm(d_model)
# 5. Dropouts for residuals
self.dropout1 = nn.Dropout(p=dropout)
self.dropout2 = nn.Dropout(p=dropout)
self.dropout3 = nn.Dropout(p=dropout)
def forward(
self,
tgt: torch.Tensor, # [B, T, D]
memory: torch.Tensor, # [B, S, D] (encoder output)
tgt_mask: Optional[torch.Tensor] = None, # [T, T] causal mask
memory_key_padding_mask: Optional[torch.BoolTensor] = None, # [B, S]
tgt_key_padding_mask: Optional[torch.BoolTensor] = None, # [B, T]
) -> torch.Tensor:
"""Pre-LayerNorm 前向传播。"""
# 适配 Flash Attention: 使用 is_causal 替代显式 causal mask
is_causal = tgt_mask is not None
if self.pre_norm:
# 1. Masked Self-Attention
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)
# 2. Cross-Attention
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)
# 3. FFN
residual = tgt
tgt = self.norm3(tgt)
tgt = self.linear2(self.dropout(self.activation(self.linear1(tgt))))
tgt = residual + self.dropout3(tgt)
else:
# Post-LayerNorm (备用)
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