RNN/LSTM 的局限性:
Transformer 的优势:
数学公式:
Attention(Q, K, V) = softmax(\frac{QK^T}{\sqrt{d_k}})V
Python 实现:
import torch import torch.nn as nn import torch.nn.functional as F class ScaledDotProductAttention(nn.Module): def __init__(self, d_k): super().__init__() self.d_k = d_k def forward(self, Q, K, V, mask=None): # Q, K, V: (batch_size, seq_len, d_k) scores = torch.matmul(Q, K.transpose(-2, -1)) / torch.sqrt(torch.tensor(self.d_k, dtype=torch.float32)) if mask is not None: scores = scores.masked_fill(mask == 0, -1e9) attention_weights = F.softmax(scores, dim=-1) output = torch.matmul(attention_weights, V) return output, attention_weights
原理:并行计算多个注意力,捕捉不同特征。
class MultiHeadAttention(nn.Module): def __init__(self, d_model, num_heads): super().__init__() self.d_model = d_model self.num_heads = num_heads self.d_k = d_model // num_heads self.W_q = nn.Linear(d_model, d_model) self.W_k = nn.Linear(d_model, d_model) self.W_v = nn.Linear(d_model, d_model) self.W_o = nn.Linear(d_model, d_model) self.attention = ScaledDotProductAttention(self.d_k) def forward(self, query, key, value, mask=None): batch_size = query.size(0) # 线性变换并分割头 Q = self.W_q(query).view(batch_size, -1, self.num_heads, self.d_k).transpose(1, 2) K = self.W_k(key).view(batch_size, -1, self.num_heads, self.d_k).transpose(1, 2) V = self.W_v(value).view(batch_size, -1, self.num_heads, self.d_k).transpose(1, 2) # 注意力计算 attn_output, attn_weights = self.attention(Q, K, V, mask) # 拼接多头 attn_output = attn_output.transpose(1, 2).contiguous().view(batch_size, -1, self.d_model) # 输出投影 output = self.W_o(attn_output) return output, attn_weights
class TransformerEncoderLayer(nn.Module): def __init__(self, d_model, num_heads, d_ff, dropout=0.1): super().__init__() self.self_attn = MultiHeadAttention(d_model, num_heads) self.feed_forward = nn.Sequential( nn.Linear(d_model, d_ff), nn.ReLU(), nn.Linear(d_ff, d_model) ) self.norm1 = nn.LayerNorm(d_model) self.norm2 = nn.LayerNorm(d_model) self.dropout = nn.Dropout(dropout) def forward(self, x, mask=None): # Self-Attention + Residual + Norm attn_output, _ = self.self_attn(x, x, x, mask) x = self.norm1(x + self.dropout(attn_output)) # Feed-Forward + Residual + Norm ff_output = self.feed_forward(x) x = self.norm2(x + self.dropout(ff_output)) return x
class TransformerDecoderLayer(nn.Module): def __init__(self, d_model, num_heads, d_ff, dropout=0.1): super().__init__() self.self_attn = MultiHeadAttention(d_model, num_heads) self.cross_attn = MultiHeadAttention(d_model, num_heads) self.feed_forward = nn.Sequential( nn.Linear(d_model, d_ff), nn.ReLU(), nn.Linear(d_ff, d_model) ) self.norm1 = nn.LayerNorm(d_model) self.norm2 = nn.LayerNorm(d_model) self.norm3 = nn.LayerNorm(d_model) self.dropout = nn.Dropout(dropout) def forward(self, x, enc_output, self_mask=None, cross_mask=None): # Self-Attention(带因果掩码) attn_output, _ = self.self_attn(x, x, x, self_mask) x = self.norm1(x + self.dropout(attn_output)) # Cross-Attention attn_output, _ = self.cross_attn(x, enc_output, enc_output, cross_mask) x = self.norm2(x + self.dropout(attn_output)) # Feed-Forward ff_output = self.feed_forward(x) x = self.norm3(x + self.dropout(ff_output)) return x
原因:Self-Attention 没有位置信息。
Sinusoidal 位置编码:
class PositionalEncoding(nn.Module): def __init__(self, d_model, max_len=5000): super().__init__() pe = torch.zeros(max_len, d_model) position = torch.arange(0, max_len, dtype=torch.float).unsqueeze(1) div_term = torch.exp(torch.arange(0, d_model, 2).float() * (-math.log(10000.0) / d_model)) pe[:, 0::2] = torch.sin(position * div_term) pe[:, 1::2] = torch.cos(position * div_term) pe = pe.unsqueeze(0) self.register_buffer('pe', pe) def forward(self, x): return x + self.pe[:, :x.size(1)]
class FeedForward(nn.Module): def __init__(self, d_model, d_ff, dropout=0.1): super().__init__() self.linear1 = nn.Linear(d_model, d_ff) self.linear2 = nn.Linear(d_ff, d_model) self.dropout = nn.Dropout(dropout) def forward(self, x): return self.linear2(self.dropout(F.relu(self.linear1(x))))
因果掩码实现:
def generate_causal_mask(seq_len): mask = torch.tril(torch.ones(seq_len, seq_len)) return mask.unsqueeze(0).unsqueeze(0) # (1, 1, seq_len, seq_len)
Warmup + Cosine Decay:
def get_lr_scheduler(optimizer, warmup_steps, total_steps): def lr_lambda(current_step): if current_step < warmup_steps: return float(current_step) / float(max(1, warmup_steps)) progress = float(current_step - warmup_steps) / float(max(1, total_steps - warmup_steps)) return max(0.0, 0.5 * (1.0 + math.cos(math.pi * progress))) return torch.optim.lr_scheduler.LambdaLR(optimizer, lr_lambda)
# 模拟大 batch size accumulation_steps = 4 for i, batch in enumerate(dataloader): loss = model(batch) loss = loss / accumulation_steps loss.backward() if (i + 1) % accumulation_steps == 0: optimizer.step() optimizer.zero_grad()
from torch.cuda.amp import autocast, GradScaler scaler = GradScaler() for batch in dataloader: with autocast(): loss = model(batch) scaler.scale(loss).backward() scaler.step(optimizer) scaler.update() optimizer.zero_grad()
def generate_text(model, prompt, max_length=100, temperature=0.7): model.eval() input_ids = tokenizer.encode(prompt, return_tensors='pt') with torch.no_grad(): for _ in range(max_length): outputs = model(input_ids) next_token_logits = outputs.logits[:, -1, :] # Temperature scaling next_token_logits = next_token_logits / temperature # Top-k sampling top_k = 50 top_k_logits, top_k_indices = torch.topk(next_token_logits, top_k) probabilities = F.softmax(top_k_logits, dim=-1) next_token = torch.multinomial(probabilities, num_samples=1) input_ids = torch.cat([input_ids, top_k_indices.gather(-1, next_token)], dim=-1) return tokenizer.decode(input_ids[0], skip_special_tokens=True)
class TextClassifier(nn.Module): def __init__(self, encoder, num_classes): super().__init__() self.encoder = encoder self.classifier = nn.Linear(encoder.config.hidden_size, num_classes) def forward(self, input_ids, attention_mask): outputs = self.encoder(input_ids, attention_mask=attention_mask) pooled_output = outputs.pooler_output logits = self.classifier(pooled_output) return logits
优势:
优势:
优势:
Transformer 彻底改变了 NLP:
掌握 Transformer 是理解现代 LLM 的基础!