File size: 3,070 Bytes
c1a46f7
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
"""
位置编码模块 — Person B 负责实现

包含:
1. SinusoidalPositionalEncoding: 经典正弦余弦位置编码
2. RotaryPositionalEmbedding (RoPE): 旋转位置编码

RoPE 是当前最前沿的位置编码方案,被 LLaMA, GPT-NeoX 等大模型广泛采用。
优势: 更好地捕获相对位置信息,支持长度外推。
"""

from __future__ import annotations

import math

import torch
import torch.nn as nn


class SinusoidalPositionalEncoding(nn.Module):
    """
    经典正弦余弦位置编码 (Vaswani et al., 2017)。

    PE(pos, 2i)   = sin(pos / 10000^(2i/d_model))
    PE(pos, 2i+1) = cos(pos / 10000^(2i/d_model))

    TODO [Person B]: 实现以下内容:
    1. 在 __init__ 中预计算位置编码矩阵 [max_seq_len, d_model]
    2. 注册为 buffer (不参与梯度更新)
    3. forward(x): 返回 x + pe[:, :x.size(1), :]
    """

    def __init__(self, d_model: int, max_seq_len: int = 5000, dropout: float = 0.1):
        super().__init__()
        raise NotImplementedError("TODO: Person B 实现 SinusoidalPositionalEncoding.__init__")

    def forward(self, x: torch.Tensor) -> torch.Tensor:
        """
        Args:
            x: [B, L, D]
        Returns:
            x + positional_encoding: [B, L, D]
        """
        raise NotImplementedError("TODO: Person B 实现 SinusoidalPositionalEncoding.forward")


class RotaryPositionalEmbedding(nn.Module):
    """
    旋转位置编码 (RoPE) — Su et al., 2021

    核心思想: 通过旋转变换将位置信息编码到 Q, K 向量中。
    q' = q * cos(θ) + rotate_half(q) * sin(θ)
    k' = k * cos(θ) + rotate_half(k) * sin(θ)

    TODO [Person B]: 实现以下内容:

    __init__:
    1. 计算频率: inv_freq = 1.0 / (base ** (torch.arange(0, dim, 2) / dim))
    2. 注册为 buffer

    _compute_rope(seq_len):
    1. 计算 position indices: t = [0, 1, ..., seq_len-1]
    2. 计算 freqs = torch.outer(t, inv_freq)
    3. 构建 cos_cached, sin_cached

    apply_rotary_pos_emb(q, k):
    1. 对 q 和 k 应用旋转变换
    2. 返回旋转后的 q', k'

    参考: https://arxiv.org/abs/2104.09864
    """

    def __init__(self, dim: int, max_seq_len: int = 2048, base: float = 10000.0):
        super().__init__()
        raise NotImplementedError("TODO: Person B 实现 RotaryPositionalEmbedding.__init__")

    def _compute_rope(self, seq_len: int, device: torch.device):
        raise NotImplementedError("TODO: Person B 实现 _compute_rope")

    @staticmethod
    def _rotate_half(x: torch.Tensor) -> torch.Tensor:
        """将 x 的后半部分取反并与前半部分交换。"""
        x1, x2 = x.chunk(2, dim=-1)
        return torch.cat((-x2, x1), dim=-1)

    def apply_rotary_pos_emb(
        self,
        q: torch.Tensor,  # [B, nhead, L, d_k]
        k: torch.Tensor,  # [B, nhead, L, d_k]
    ) -> tuple[torch.Tensor, torch.Tensor]:
        """
        对 Q, K 应用 RoPE。

        Returns:
            (q_rotated, k_rotated)
        """
        raise NotImplementedError("TODO: Person B 实现 apply_rotary_pos_emb")