""" 注意力机制模块 — Person B 负责实现 包含: 1. MultiHeadAttention: 标准多头注意力 (支持 RoPE) 2. FlashMultiHeadAttention: Flash Attention 2 加速版本 技术要点: - Scaled Dot-Product Attention - 支持 key_padding_mask 和 attn_mask - Flash Attention 2 使用 torch.nn.functional.scaled_dot_product_attention - RoPE 旋转位置编码集成 """ from __future__ import annotations import math from typing import Optional import torch import torch.nn as nn import torch.nn.functional as F class MultiHeadAttention(nn.Module): """ 标准多头注意力机制。 TODO [Person B]: 实现以下内容: __init__: 1. Q, K, V 线性投影: nn.Linear(d_model, d_model) 2. 输出投影: nn.Linear(d_model, d_model) 3. Dropout forward(query, key, value, key_padding_mask=None, attn_mask=None): 1. 线性投影 Q, K, V 2. reshape 为 [B, nhead, L, d_k] 3. (可选) 应用 RoPE 旋转位置编码 4. 计算 attention scores: QK^T / sqrt(d_k) 5. 应用 masks (padding mask + causal mask) 6. Softmax + Dropout 7. 加权求和 V 8. reshape 回 [B, L, d_model] 9. 输出投影 """ def __init__( self, d_model: int = 512, nhead: int = 8, dropout: float = 0.1, use_rotary_embedding: bool = False, ): super().__init__() assert d_model % nhead == 0, "d_model 必须能被 nhead 整除" self.d_model = d_model self.nhead = nhead self.d_k = d_model // nhead raise NotImplementedError("TODO: Person B 实现 MultiHeadAttention.__init__") def forward( self, query: torch.Tensor, # [B, L_q, D] key: torch.Tensor, # [B, L_k, D] value: torch.Tensor, # [B, L_v, D] key_padding_mask: Optional[torch.BoolTensor] = None, # [B, L_k] attn_mask: Optional[torch.Tensor] = None, # [L_q, L_k] ) -> torch.Tensor: raise NotImplementedError("TODO: Person B 实现 MultiHeadAttention.forward") class FlashMultiHeadAttention(nn.Module): """ Flash Attention 2 加速的多头注意力。 TODO [Person B]: 使用 PyTorch 2.0+ 的 F.scaled_dot_product_attention 实现: 1. 与 MultiHeadAttention 结构相同 2. 在 forward 中使用 F.scaled_dot_product_attention(Q, K, V, attn_mask, dropout, is_causal) 3. 会自动选择最优的 attention kernel (Flash Attention / Memory-Efficient Attention) 注意: - 需要 PyTorch >= 2.0 - is_causal=True 时自动生成因果掩码,不需要手动传入 attn_mask """ def __init__( self, d_model: int = 512, nhead: int = 8, dropout: float = 0.1, use_rotary_embedding: bool = False, ): super().__init__() raise NotImplementedError("TODO: Person B 实现 FlashMultiHeadAttention.__init__") def forward( self, query: torch.Tensor, key: torch.Tensor, value: torch.Tensor, key_padding_mask: Optional[torch.BoolTensor] = None, is_causal: bool = False, ) -> torch.Tensor: raise NotImplementedError("TODO: Person B 实现 FlashMultiHeadAttention.forward")