| """ |
| 注意力机制模块 — 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, |
| key: torch.Tensor, |
| value: torch.Tensor, |
| key_padding_mask: Optional[torch.BoolTensor] = None, |
| attn_mask: Optional[torch.Tensor] = None, |
| ) -> 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") |
|
|