minimind

tokenizer

BPE((Byte Pair Encoding))

基础词表由 256 个单字节组成,随后通过统计算法将高频相邻字节合并,直到词表大小达到 6400

embedding

输入为:(batch_size,seq_len,dim)

位置编码

RoPE

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
def precompute_freqs_cis(dim: int, end: int = int(32 * 1024), rope_base: float = 1e6,
                         rope_scaling: Optional[dict] = None):
    """
    预计算用于RoPE(Rotary Position Embedding)的频率余弦和正弦值。
    参数:
        dim (int): 嵌入维度,必须为偶数。
        end (int): 序列的最大长度,默认为32*1024,推理阶段。
        rope_base (float): RoPE的基础频率参数,默认为1e6。
        rope_scaling (Optional[dict]): 可选的缩放配置字典,包含以下键:
            - "original_max_position_embeddings" (int): 原始最大位置嵌入数,默认为2048,预训练的最大上下文长度。
            - "factor" (float): 扩展因子,默认为16。
            - "beta_fast" (float): 快速衰减参数,默认为32.0。
            - "beta_slow" (float): 慢速衰减参数,默认为1.0,说明模型见到的角度在一个周期内
            - "attention_factor" (float): 注意力因子,默认为1.0。
    返回:
        Tuple[Tensor, Tensor]: 包含两个张量:
            - freqs_cos: 频率余弦值张量,形状为[end, dim]。
            - freqs_sin: 频率正弦值张量,形状为[end, dim]。
    """
  1. 每个维度的基础旋转频率$\theta_i = \text{rope_base}^{-\frac{2i}{\text{dim}}}$
    • 把一个长度为 dim(比如 128 维)的向量,切分成一个个二维的旋转平面
    • 每 2 个维度组成一个平面,共享同一个旋转频率($\theta_i$)。
    • 所以,128 维的向量,实际上只有 64 个独立的旋转频率
    • rope_base默认1e6
      1
      		  freqs, attn_factor = 1.0 / (rope_base ** (torch.arange(0, dim, 2)[: (dim // 2)].float() / dim)), 1.0
    • torch.arange(0, dim, 2):生成从 $0$ 到 dim-2 的偶数序列(例如 dim=128 时,生成 [0, 2, 4, ..., 126])。
    • [: (dim // 2)]:截取前一半的元素。这其实是一个安全保护措施,因为上面的 arange 已经生成了 dim/2 个元素。
    • .float() / dim:计算 $2i/d$。
    • `rope_base **:计算 $\text{base}^{2i/d}$。rope_base` 默认通常是 $10000$ 或 $1000000$。
    • 1.0 / ...:取倒数,得到最终的频率向量 freqs。形状为 [dim // 2]。同时初始化 attn_factor(注意力缩放因子)为 1.0。
  2. YaRN外推
    1
    2
    3
    4
    5
    6
    7
    8
    	    if rope_scaling is not None:
            orig_max, factor, beta_fast, beta_slow, attn_factor = (
                rope_scaling.get("original_max_position_embeddings", 2048),
                rope_scaling.get("factor", 16),
                rope_scaling.get("beta_fast", 32.0),
                rope_scaling.get("beta_slow", 1.0),
                rope_scaling.get("attention_factor", 1.0)
            )
    • 如果传入了 rope_scaling 字典,说明需要进行上下文长度外推。这里使用 dictionary.get() 安全地提取参数,如果没传,则使用 YaRN 的默认推荐值。
    • factor要扩展的倍数(比如把 2k 上下文扩展到 32k,$s=16$)。
    • attn_factor 是由于 YaRN 改变了注意力分布,需要引入的一个温度放大因子来恢复注意力熵。
      1
      2
      3
      4
      5
      6
              # 当我们要推理的长度 (end) 大于模型原本见过的长度 (orig_max) 时,才触发 YaRN 外推机制
              # YaRN高频外推,低频内插,维度越低,频率越高,波长越小,在原文本中对应的周期越大
              if end / orig_max > 1.0:
                  # 找到波长边界对应的维度索引
                  def inv_dim(b):
                      return (dim * math.log(orig_max / (b * 2 * math.pi))) / (2 * math.log(rope_base))
      推导:
  • 在 YaRN 中,波长 $\lambda$ 是通过原始上下文长度 $L$ (orig_max) 和一个比例系数 $b$ (beta_fastbeta_slow) 来定义的。YaRN 认为如果一个频率在整个上下文长度 $L$ 中只完成了极少数的几个周期(即 $b$ 个周期),那么它就属于低频(全局信息);如果完成了非常多个周期,就属于高频(局部信息)。因此,阈值波长可以表示为:
  • RoPE 中第 $d$ 个维度($d \in [0, D/2)$)对应的旋转波长公式为:

_其中,$D$ 是总维度 dim,$\text{base}$ 是 rope_base
将 $\lambda = \frac{L}{b}$ 代入上面的公式:

  • $D$ $\rightarrow$ dim
  • $L$ $\rightarrow$ orig_max
  • $\text{base}$ $\rightarrow$ rope_base
  • $\ln$ $\rightarrow$ math.log
    1
    low, high = max(math.floor(inv_dim(beta_fast)), 0), min(math.ceil(inv_dim(beta_slow)), dim // 2 - 1)
    计算混合区的起始索引 low 和结束索引 high,索引范围[0,dim//2-1]
  • 索引 $< \text{low}$ 的是高频区(代表局部信息,不需要缩放,直接外推)。
  • 索引 $> \text{high}$ 的是低频区(代表全局信息,需要按 factor 进行线性内插缩放)。
    1
    2
                ramp = torch.clamp((torch.arange(
                    dim // 2, device=freqs.device).float() - low) / max(high - low, 0.001), 0, 1)
    torch.arange(dim // 2, device=freqs.device).float():生成当前所有复数维度的索引,即一个向量 $i = [0, 1, 2, …, \frac{d}{2}-1]$
    ( ... - low) / max(high - low, 0.001):计算 $\frac{i - \text{low}}{\text{high} - \text{low}}$
    torch.clamp( ... , 0, 1):将上一步算出的值强行限制在 $[0, 1]$ 闭区间内
  • 如果值 $< 0$(也就是 $i < \text{low}$ 的高频区),强制变成 $0$。
  • 如果值 $> 1$(也就是 $i > \text{high}$ 的低频区),强制变成 $1$。
  • 如果是 $0$ 到 $1$ 之间的值(即 $\text{low} \le i \le \text{high}$ 的混合区),保持原样。
    1
              freqs = freqs * (1 - ramp + ramp / factor)
  • 在高频区,模型见到的角度更多,即使文本变长,角度还是那些,频率保持不变
  • 低频区,模型见到的角度少,文本变长,角度范围变大,需要降低频率,使模型见到的角度范围可控,实现内插
  • 混合区频率的缩放比例从 1 平滑过渡到 $1 / \text{factor}$,避免注意力机制产生突变
    1
    2
       t = torch.arange(end, device=freqs.device)
       freqs = torch.outer(t, freqs).float()
  • t:生成位置序列向量 $[0, 1, 2, …, \text{end}-1]$,代表每一个 Token 的绝对位置 $m$。
  • torch.outer(t, freqs):计算外积。将位置标量向量 $m$ 与频率向量 $\theta$ 相乘。
  • freqs 现在是一个形状为 [end, dim // 2] 的矩阵,矩阵里的每一个元素是 $m \cdot \theta_i$
    1
    2
    3
    4
        freqs_cos = torch.cat(
            [torch.cos(freqs), torch.cos(freqs)], dim=-1) * attn_factor
        freqs_sin = torch.cat(
            [torch.sin(freqs), torch.sin(freqs)], dim=-1) * attn_factor
  • torch.cat([..., ...], dim=-1):将矩阵在最后一个维度复制拼接一次。
  • 原本的形状是 [end, dim // 2],拼接后变成了 [end, dim]
  • 这种拼接方式意味着在后续的注意力计算中,查询(Query)和键(Key)的实部和虚部是前半部分和后半部分对应的,即格式为 $[x_0, x_1, …, x_n, x_0, x_1, …, x_n]$。(有些模型的 RoPE 实现是交叉排列的,如 $[x_0, y_0, x_1, y_1…]$,这段代码显然采用的是前后拼接排列)。
  • attn_factor:低频压缩后在高维相邻token的相邻角度变小了,注意力分布变得平坦,在预计算时给Q,K都乘上了 attn_factor
  • Q 放大了 attn_factor 倍。
  • K 放大了 attn_factor 倍。
  • 那么它们的点积结果(注意力原始得分),就会被放大 $\text{attn_factor}^2$ 倍,在softmax时指数倍扩大,结果自然也集中了
    1
    2
    3
    4
    5
    6
    7
    8
    9
    10
    11
    12
    13
    def apply_rotary_pos_emb(q, k, cos, sin, position_ids=None, unsqueeze_dim=1):
        def rotate_half(x):
            return torch.cat((-x[..., x.shape[-1] // 2:], x[..., : x.shape[-1] // 2]), dim=-1)
        # unsqueeze 的作用是在指定位置插入一个大小为 1 的新维度。如果 unsqueeze_dim=1,cos 就会从 [seq_len, head_dim] 变成 [seq_len, 1, head_dim]
        # q 的形状:                [batch_size,  seq_len,  num_heads,  head_dim]
        # 相乘广播,cos 的形状:     [         1,  seq_len,          1,  head_dim]
        # 不管是第 1 个注意力头,还是第 32 个注意力头,它们在同一个词的位置上,使用的是完全相同的旋转角度
        # 在最后一个维度 `head_dim` 上发生实质性的两两数值相乘
        q_embed = (q * cos.unsqueeze(unsqueeze_dim)) + \
            (rotate_half(q) * sin.unsqueeze(unsqueeze_dim))
        k_embed = (k * cos.unsqueeze(unsqueeze_dim)) + \
            (rotate_half(k) * sin.unsqueeze(unsqueeze_dim))
        return q_embed, k_embed
  • x[..., x.shape[-1] // 2:]:切片取出后半部分,即 $[y_1, y_2]$。
  • -x[...]:加个负号,变成了 $[-y_1, -y_2]$。
  • x[..., : x.shape[-1] // 2]:切片取出前半部分,即 $[x_1, x_2]$。
  • torch.cat((...), dim=-1):把取负的后半部分,拼在前半部分的前面。
  • $\text{rotate_half}(x) = [-y_1, -y_2, x_1, x_2]$
  • 在RoPE原文中是对维度分量相邻的两两分组,每组对应一个旋转矩阵,只是不同组旋转矩阵频率不同,组成一个大的分块矩阵,有很多零元素,直接用矩阵相乘空间复杂度$O(d^2)$
  • 论文中$[x_1, y_1, x_2, y_2]$
  • 改用逐元素相乘
  • 假设 $q = [x_1, y_1, x_2, y_2, x_3, y_3, …]$ 我们要把它变成 rotate(q) = $[-y_1, x_1, -y_2, x_2, -y_3, x_3, …]$,在 PyTorch 等框架中,为了从原向量中把所有的 $y$ 取出来加负号,把所有的 $x$ 取出来,你需要按步长(stride)去切片。
  • 这种操作在内存中是跳跃访问的(不连续)。在 GPU 这种高度依赖内存带宽和连续读取的硬件上,跳跃访问会破坏缓存一致性(Cache Coherence),导致计算变慢。
  • 在手写的 CUDA 内核里,线程可以精确控制
    • 一个 GPU 线程直接把相邻的 $(x_1, y_1)$ 一次性从显存读进寄存器。
    • 在寄存器(超高速缓存)里,线程自己完成 $-y_1\sin\theta_1 + x_1\cos\theta_1$ 的数学计算。
    • 计算完后,再把结果连续地写回显存。

      GQA

机制全称结构特点优点缺点
MHAMulti-Head Attention每个 Q 都有对应独立的 K 和 V。表达能力最强,能学习到多种不同的上下文关系。显存占用高,推理时 KV Cache 巨大,速度慢。
MQAMulti-Query Attention所有的 Q 共享同一组 K 和 V。极大地减少了 KV Cache 显存,推理速度最快。可能会损失一定的模型精度(表达能力下降)。
GQAGrouped-Query Attention将 Q 分组,每一组 Q 共享一组 K 和 V。性能与效果的折中。速度接近 MQA,精度接近 MHA。需要手动调整分组(Group)参数。
  • MHA: 1 个 $Q$ 对应 1 个 $K, V$(一对一)。
  • MQA: 所有 $Q$ 对应 1 个 $K, V$(多对一)。
  • GQA: 一组 $Q$ 对应 1 个 $K, V$(多对少)。
  • Q 的头数很多(例如 32 个)。
  • K 和 V 的头数很少(例如 8 个)。
  • 这就意味着,每 4 个 Q 头,必须共享同一个 K 头和 V 头

    1
    2
    3
    4
    5
    6
    7
    8
    9
    10
    def repeat_kv(x: torch.Tensor, n_rep: int) -> torch.Tensor:
        """torch.repeat_interleave(x, dim=2, repeats=n_rep)"""
        bs, slen, num_key_value_heads, head_dim = x.shape
        if n_rep == 1:
            return x
        # (batch_size,seq_len,num_heads,head_dim)
        return (
            x[:, :, :, None, :].expand(bs, slen, num_key_value_heads, n_rep, head_dim).reshape(
                bs, slen, num_key_value_heads * n_rep, head_dim)
        )
  • None 的作用是在指定位置插入一个大小为 1 的新维度。
  • .expand(bs, slen, num_key_value_heads, n_rep, head_dim):把刚才那个大小为 1 的维度,扩充成了 n_rep (也就是 4),零内存扩充
  • 如果num_key_value_heads=1就是MQA

    Attention


    输入x(经过embedding)经过Q,K,V后分头,K,V的头少,Q的头多,得到xq(batch_size,seq_len,num_q_heads,head_dim)
    xk,xv(batch_size,seq_len,num_q_heads,head_dim)
    对xq,xk进行位置编码
    考虑kv_cache,拼在xk,xv的seq_len上(前面),变成(batch_size,seq_len_1,num_q_heads,head_dim)
    将xk,xv头复制,和xq头一样多
    attention公式,xq,xk相乘,除以$\sqrt{d_k}$,考虑因果编码,填充编码,经过softmax后经过dropout与xv相乘
    结果再通过线性层,再经过dropout

>

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
class Attention(nn.Module):
    def __init__(self, args: MiniMindConfig):
        super().__init__()
        # 变量 = 值A if 条件 else 值B
        # %取余数
        # //整除,返回整数部分
        self.num_key_value_heads = args.num_attention_heads if args.num_key_value_heads is None else args.num_key_value_heads
        # Query 的头数,必须能够被 K,V 的头数完美整除
        assert args.num_attention_heads % self.num_key_value_heads == 0
        self.n_local_heads = args.num_attention_heads
        self.n_local_kv_heads = self.num_key_value_heads
        self.n_rep = self.n_local_heads // self.n_local_kv_heads  # 头复制次数
        self.head_dim = args.hidden_size // args.num_attention_heads
        self.q_proj = nn.Linear(
            args.hidden_size, args.num_attention_heads * self.head_dim, bias=False)
        self.k_proj = nn.Linear(
            args.hidden_size, self.num_key_value_heads * self.head_dim, bias=False)
        self.v_proj = nn.Linear(
            args.hidden_size, self.num_key_value_heads * self.head_dim, bias=False)
        self.o_proj = nn.Linear(
            args.num_attention_heads * self.head_dim, args.hidden_size, bias=False)
        self.attn_dropout = nn.Dropout(args.dropout)
        self.resid_dropout = nn.Dropout(args.dropout)
        self.dropout = args.dropout
        self.flash = hasattr(
            torch.nn.functional, 'scaled_dot_product_attention') and args.flash_attn
        # print("WARNING: using slow attention. Flash Attention requires PyTorch >= 2.0")
    def forward(self,
                x: torch.Tensor,
                # 修改为接收cos和sin
                position_embeddings: Tuple[torch.Tensor, torch.Tensor],
                past_key_value: Optional[Tuple[torch.Tensor,
                                               torch.Tensor]] = None,
                use_cache=False,
                attention_mask: Optional[torch.Tensor] = None):
        bsz, seq_len, _ = x.shape
        xq, xk, xv = self.q_proj(x), self.k_proj(x), self.v_proj(x)
        xq = xq.view(bsz, seq_len, self.n_local_heads, self.head_dim)
        xk = xk.view(bsz, seq_len, self.n_local_kv_heads, self.head_dim)
        xv = xv.view(bsz, seq_len, self.n_local_kv_heads, self.head_dim)

        cos, sin = position_embeddings
        xq, xk = apply_rotary_pos_emb(xq, xk, cos, sin)
        # xq:(batch_size,seq_len,num_heads,head_dim)
        # kv_cache实现
        # 将过去算好的历史 K/V 与当前新词的 K/V 在序列长度维度(seq_len)拼接
        # past_key_value是元组(xk,xv)
        # (batch_size,seq_len_1,num_heads,head_dim)
        if past_key_value is not None:
            xk = torch.cat([past_key_value[0], xk], dim=1)
            xv = torch.cat([past_key_value[1], xv], dim=1)
        past_kv = (xk, xv) if use_cache else None

        xq, xk, xv = (
            xq.transpose(1, 2),  # (batch_size,num_heads,seq_len,head_dim)
            repeat_kv(xk, self.n_rep).transpose(1, 2),#(batch_size,num_heads,seq_len_1,head_dim)
            repeat_kv(xv, self.n_rep).transpose(1, 2)
        )

        if self.flash and (seq_len > 1) and (past_key_value is None) and (attention_mask is None or torch.all(attention_mask == 1)):
            output = F.scaled_dot_product_attention(
                xq, xk, xv, dropout_p=self.dropout if self.training else 0.0, is_causal=True)
        # 经典 Attention 手动实现
        else:
            # 计算 QK^T / sqrt(d),(batch_size,num_heads,seq_len,seq_len_1)
            # 行(Rows)代表 Query:每一行代表当前正在“思考”的那个词
            # 列(Cols)代表 Key:每一列代表被参考的历史词
            scores = (xq @ xk.transpose(-2, -1)) / math.sqrt(self.head_dim)
            # torch.full((seq_len, seq_len), float("-inf")),创建一个形状为 [seq_len, seq_len] 的正方形矩阵,里面所有的坑位都填上负无穷大
            # triu 代表 Triangular Upper(上三角)。它会保留矩阵的右上部分,把剩下的(左下部分)全部变成0
            # diagonal=1:表示从主对角线往右偏移 1 行开始保留。这样,主对角线(词看自己)也被排除在外,变成了0
            # 使因果编码矩阵
            # 最后scores加上因果编码矩阵
            scores[:, :, :, -seq_len:] += torch.triu(torch.full(
                (seq_len, seq_len), float("-inf"), device=scores.device), diagonal=1)
            # Padding Mask
            # 在大模型处理批次(Batch)数据时,句子的长度往往不一
            # 为了能把不同长度句子塞进同一个矩阵并行计算,我们会把短句子补上无意义的占位符
            # 模型在计算注意力时,不应该把注意力浪费在这些“占位符”上。我们需要强行让这些位置的权重变成 0
            # 原始 attention_mask:通常是 [batch_size, seq_len],里面的值是 1(真词)或 0(填充词)
            # 将真词设置为0,填充词设置为无穷
            if attention_mask is not None:
                # 两次 unsqueeze 后:形状变成了 [batch_size, 1, 1, seq_len]
                extended_attention_mask = attention_mask.unsqueeze(
                    1).unsqueeze(2)
                extended_attention_mask = (
                    1.0 - extended_attention_mask) * -1e9
                scores = scores + extended_attention_mask
            # xq 和 scores 通常是以 FP16,Softmax 内部的指数运算非常容易导致数值溢出(超出 FP16 能表达的最大值 65504),因为要先算指数,再归一化
            # .float(),把 scores 强行提升到 FP32
            scores = F.softmax(scores.float(), dim=-1).type_as(xq)
            scores = self.attn_dropout(scores)
            output = scores @ xv  # (batch_size,num_heads,seq_len,head_dim)
        output = output.transpose(1, 2).reshape(
            bsz, seq_len, -1)  # (batch_size,seq_len,dim)
        output = self.resid_dropout(self.o_proj(output))
        return output, past_kv

  • 线性层不加偏置:
  • 高维空间的注意力计算中,模型更在乎的是向量之间的相对方向(点积结果),而不是绝对位置
  • RMSNorm的数值稳定性:

    • 如果此时线性层有偏置 $b$: 偏置会将激活值 $x$ 整体向某个方向平移(不再以零为中心)。因为 RMSNorm 不减去均值,这个非零的偏置会直接导致分母(均方根)变大。
    • 稳定性灾难: 当分母被人为撑大后,整个激活值会被严重压缩(缩放)。在反向传播时,这种偏移和压缩会导致梯度更新极不稳定。
    • 带偏置的情况: 在训练初期,梯度的剧烈更新可能会导致偏置 $b$ 发生较大的跳变。这意味着整个特征空间被强行“挪动”了一个位置。下一层的网络不得不去努力适应这种突如其来的空间平移,导致收敛缓慢甚至震荡。
    • 不带偏置的情况: 纯矩阵乘法 $Wa$ 只改变向量的方向和长度,不改变空间的基准原点。特征分布被锚定在原点附近,每一层只需要关注向量之间的相对角度和特征放缩,这极大地降低了优化难度,让损失函数(Loss)的下降曲线更加平滑稳定。
  • 这里的 self.q_proj(x) 在 PyTorch 底层执行的是:
  • $x$ 的形状: [batch_size, seq_len, hidden_size]
  • $W_Q$ 的形状: [hidden_size, num_heads * head_dim]
  • 相乘后的结果 xq: [batch_size, seq_len, num_heads * head_dim]

    QKV计算

操作结果对训练的影响
不缩放点积方差大,Softmax 容易进入饱和区梯度消失,模型无法收敛
缩放 ($\frac{1}{\sqrt{d_k}}$)点积方差保持为 1,Softmax 处于敏感区梯度充足,训练稳定高效

因果编码

Query:正在思考哪个词
Key:参考的历史词
思考的词时,前面的是历史词,模型在生成第 $t$ 个词时,只能看到第 $1$ 到第 $t$ 个词
所以因果编码矩阵右上角为副无穷,这样做softmax概率为0

Padding Mask

屏蔽无效的占位符

Softmax

每一行代表一个词,每一列代表它对某个位置的“关注度”。所有分值相加等于 1
与V相乘,每个词都在保持自己位置的同时,吸收了来自全场其他词的信息

KV cache

K和V的seq_len增加,使词吸收过去的信息

FFN


1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
class FeedForward(nn.Module):
    def __init__(self, config: MiniMindConfig):
        super().__init__()
        if config.intermediate_size is None:
            intermediate_size = int(config.hidden_size * 8 / 3)
            config.intermediate_size = 64 * \
                ((intermediate_size + 64 - 1) // 64)
        self.gate_proj = nn.Linear(
            config.hidden_size, config.intermediate_size, bias=False)
        self.down_proj = nn.Linear(
            config.intermediate_size, config.hidden_size, bias=False)
        self.up_proj = nn.Linear(
            config.hidden_size, config.intermediate_size, bias=False)
        self.dropout = nn.Dropout(config.dropout)
        self.act_fn = ACT2FN[config.hidden_act]

    def forward(self, x):
        return self.dropout(self.down_proj(self.act_fn(self.gate_proj(x)) * self.up_proj(x)))

传统的 Transformer 使用的是简单的全连接层:$FFN(x) = \text{ReLU}(xW_1)W_2$
Swish函数:$f(x) = x \cdot \sigma(\beta x)$
$\sigma(z) = \frac{1}{1 + e^{-z}}$

  • 当 $x$ 很大时(正向强信号):$\sigma(x)$ 趋近于 $1$。此时 $f(x) \approx x$,信号几乎全量通过。
  • 当 $x \approx 0$ 时:$\sigma(x) = 0.5$。信号被削减一半。
  • 当 $x$ 很小时(负向弱信号):$\sigma(x)$ 趋近于 $0$。此时 $f(x) \approx 0$,信号被拦截(关门)
    self.act_fn(self.gate_proj(x)) * self.up_proj(x))只有门控通道觉得“重要”的信息,才会在 Up 通道中被保留或增强
    $\text{SiLU}(\cdot)$:这是 Swish 的一种特例($\beta=1$),公式为 $f(x) = x \cdot \sigma(x)$。
    down_proj (Down Proj): 负责将处理后的高维特征投影回原始维度
    dropout :防止过拟合

    为什么8/3

    传统的 Transformer :$2 \times (d \times 4d) = \mathbf{8d^2}$
    SwiGLU架构:$3 \times (d \times Kd) = 8d^2$
  • 为了不增加计算成本和参数量并且使用 SwiGLU 能够显著提升模型在大规模数据上的收敛速度和最终性能

    MoE

路由(routing)策略

Token Choice

对于输入的每一个 Token,路由器(Router)会计算它与所有专家之间的匹配得分,然后这个 Token 会主动选择得分最高的 top_k 个专家
缺点:负载不均衡

  • 某些明星专家会被大多数 Token 选中,导致计算瓶颈;而某些专家则无人问津,得不到训练(专家退化)。
  • 为了解决这个问题,通常需要额外添加辅助损失函数(Auxiliary Loss)来强迫流量均衡

    Expert Choice

专家来挑选 Token。每个专家都有一个固定的“容量”(Capacity),它会从当前的 Batch 中挑选最契合自己的前 $M$ 个 Token
缺点:

  • Token 遗漏:可能会有倒霉的 Token 没被任何一个专家相中,导致信息丢失。
  • 实时预测难:在推理(Inference)阶段,很难确定一个专家该拿多少 Token,这在处理流式任务时比较麻烦。

    MoEGate

    采用Token choice
    算的是不同token对应的专家,以及专家的权重(可能归一化),辅助损失函数
    每一轮都会变

  • 每一个专家对应一个FFN
  • weight:(专家总数, 隐藏层维度)

    权重初始化

    1
    2
        def reset_parameters(self) -> None:
            init.kaiming_uniform_(self.weight, a=math.sqrt(5))
    为什么:在神经网络中,如果权重一开始设置得太小,信号在经过多层传递后会越来越弱,导致梯度消失(Vanishing Gradients);如果设置得太大,信号会越来越强,导致梯度爆炸(Exploding Gradients)
  • 为了保证经过多个神经层后方差不变,而每个神经层输出要经过激活函数,这里是SiLU。

    基本原理

    一个神经元接收上一层的输入信号时,做的是这样的计算:

这里的 $n$ 就是输入维度。
假设输入 $x$ 和权重 $w$ 的均值都是 0,并且彼此独立。根据概率论公式,输出 $y$ 的方差可以近似表示为:

因为有 $n$ 个输入,我们可以把公式简化为:

如果我们固定权重的方差 $Var(w)$ 不变,那么输入维度 $n$ 越大,最后加出来的输出方差 $Var(y)$ 就会成倍地暴增!如果每过一层,方差都被放大 1000 倍,不出三层,数值就会大到计算机无法表示,这就是梯度爆炸
保持信号在网络中传播时,方差始终保持稳定(既不爆炸,也不消失)。 也就是说,我们希望输出的方差等于输入的方差:

因为输出还要激活,通过ReLU,方差还要减半,所以设置为$\frac{2}{n}$

  • 采用均匀分布,找边界
  • 代码中的$\sqrt5$是假设的LeakyReLU分布

    • 当 $y > 0$ 时,$f(y) = y$
    • 当 $y \le 0$ 时,$f(y) = ay$刚好得到$\text{bound} = \sqrt{\frac{1}{n}}$

      每一个Token对所有专家的原始匹配得分

      1
      2
      3
      4
      5
      6
      7
      8
      9
      10
      11
      12
      13
      14
      15
      16
      17
      18
      	    def forward(self, hidden_states):
              bsz, seq_len, h = hidden_states.shape
              # (batch_size* seq_len, dim)
              hidden_states = hidden_states.view(-1, h)
              # logits = hidden_states @ weight.T,(batch_size* seq_len,专家总数)
              # 每一个 Token 对所有专家的原始“匹配度得分”。
              logits = F.linear(hidden_states, self.weight, None)
              if self.scoring_func == 'softmax':
                  scores = logits.softmax(dim=-1)
              else:
                  raise NotImplementedError(
                      f'insupportable scoring function for MoE gating: {self.scoring_func}')
              topk_weight, topk_idx = torch.topk(
                  scores, k=self.top_k, dim=-1, sorted=False)
              # 如果 Top-K 大于 1 且配置要求归一化,则对选出的概率进行缩放,使它们之和为 1
              if self.top_k > 1 and self.norm_topk_prob:
                  denominator = topk_weight.sum(dim=-1, keepdim=True) + 1e-20
                  topk_weight = topk_weight / denominator
    1. 总token:batch_size*seq_len
    2. topk_weight:每个 Token 对应的 Top-K 专家的概率权重,形状 (bsz * seq_len, top_k)
    3. topk_idx:对应的专家索引,形状 (bsz * seq_len, top_k)

      辅助损失计算

      只在训练时计算

特性Sequence 级别 (seq_aux=True)Batch 级别 (seq_aux=False)
颗粒度极细,每一条数据内部都要求专家公平分配较粗,只要整个批次(Batch)整体公平即可
推理表现更均衡。单次推理时负载预测性好,不容易出现卡顿可能存在波动。某些特定话题的句子可能导致局部专家过载
学习难度较高。约束太死,有时会违背数据本身的语义聚集性较低。给了模型弹性,允许专家在特定领域进行深度专业化
适用场景中小模型、对推理延迟极其敏感的任务、下游任务微调超大规模预训练、追求极致吞吐量的大规模集群训练

Sequence 级别

每个batch(句子)里面计算损失
batch_size:句子数
seq_len:一个句子的token数

  1. 为什么做:
    • 大语言模型在推理(生成)的时候,是一条一条 Sequence 生成的,如果我们只做全局 Batch 级别的均衡,在实际单条推理时依然会出现某个专家过载的情况。所以现代的 MoE 模型(比如 Mixtral)很多都会引入 seq_aux(Sequence 级别的辅助损失),这就需要保留 bsz 这个维度来逐条计算
      1
      2
      3
      4
      5
      	                ce = torch.zeros(bsz, self.n_routed_experts,
                                       device=hidden_states.device)
                      ce.scatter_add_(1, topk_idx_for_aux_loss,
                                      torch.ones(bsz, seq_len * aux_topk, device=hidden_states.device)).div_(
                          seq_len * aux_topk / self.n_routed_experts)
      scatter_add_(1, ...):沿着第 2 个维度(列的方向,也就是专家的方向)去寻找位置并累加”
      topk_idx_for_aux_loss:(batchsize, seq_len top_k)
      执行完毕后,ce 里装的就是*
      每个专家实际分到的 Token 数量
      `.div
      (seq_len * aux_topk / self.n_routed_experts)`:
  • seq_len * aux_topk:一个句子里总共要分配多少次任务
  • seq_len * aux_topk / self.n_routed_experts:理想情况每个专家分配多少个任务
  • $\frac{\text{Count}}{N}$ 就是当前专家分到 Token 的真实频率 $f_i$
  • 乘以专家总数 $E$,是为了让这个值的期望基准线变成 1.0
    • 如果一个专家刚好拿到了平均数量的 Token,它的 ce 值就会是 $1.0$
    • 如果它拿得太多,ce 就会大于 $1.0$;如果它一直在“摸鱼”,ce 就会接近 $0$
  1. 辅助损失计算
    1
    2
    	                aux_loss = (ce * scores_for_seq_aux.mean(dim=1)
                                ).sum(dim=1).mean() * self.alpha
    scores_for_seq_aux:每一个 Token 对所有专家的原始“匹配度得分”,就是scores
    scores_for_seq_aux.mean(dim=1):计算一个句子里token分配给不同专家的概率$P_i$
  2. 梯度计算
    scores ➔ 挑出最大的 K 个 ➔ 得到 topk_idx ➔ 算出现实分配比例 $f_i$(即 ce
  • scores 经过 torch.topk 变成 topk_idx 时:数据类型从浮点数(Float)变成了整数索引(Long)
    • 顺着 topk_idx 往下算出来的 $f_i$(也就是 ce),在 PyTorch 眼里就失去了梯度追踪的资格。PyTorch 会把它当成一个纯粹的常数
  • scores 算出来的平均概率 $P_i$,全程都是平滑的浮点数运算(Softmax 和 Mean),所以它保留了完整的梯度信息
  • 门控网络的权重 self.weight只和$P_i$有关
  1. 举个极端的例子:
    假设门控网络彻底“偏科”,把 100 个 Token 全给了专家 A。
  • 前向传播时scores 里 A 的得分最高,导致 topk_idx 全是 A,最后算出来 $f_i$(专家 A 的接单比例)是惊人的 8.0(假设有 8 个专家)。同时,A 的 $P_i$ 也是 0.9
  • 计算 Loss 时:$Loss = 8.0 \times 0.9 = 7.2$(一个巨大的惩罚)。
  • 反向传播时:因为 8.0 是个不可导的死数字,模型无法通过改变接单量来降低 Loss(已经分配完了),模型唯一能做的,就是狠狠把专家 A 的预测概率 $P_i$ 往下调,迫使门控网络下次给其他专家打高分。
  • 前向计算上,现实($f_i$)确实是由理想($P_i$ 的前身 scores)决定的; 但在梯度回传上,现实($f_i$)化身成了一个冷酷的常数权重,反过来倒逼理想($P_i$)去变得更加平均
  1. 为什么乘上专家总数 $E$
  • 最完美均衡的情况:每个专家分到的活一样多($fi = \frac{1}{E}$),概率也一样($P_i = \frac{1}{E}$)。
    此时 $\sum f_i P_i = E \times (\frac{1}{E} \times \frac{1}{E}) = \frac{1}{E}$。
    再乘以公式前面的 $E$:$L
    {\text{aux}} = E \cdot \frac{1}{E} = 1$。
  • 最极端的崩溃情况:所有活全给了 1 个专家($f1 = 1, P_1 = 1$,其他全为 0)。
    此时 $\sum f_i P_i = 1 \times 1 = 1$。
    再乘以公式前面的 $E$:$L
    {\text{aux}} = E \cdot 1 = E$。
  • 乘以 $E$ 之后,这个辅助损失的值就被完美地限制在了 $[1, E]$ 的区间内
    • 最佳辅助损失不论怎么改变专家个数都是1
    • 可以用同一套 alpha 参数训练不同规模的模型

      Batch 级别

      1
      2
      3
      4
      5
      6
                      mask_ce = F.one_hot(topk_idx_for_aux_loss.view(-1), num_classes=self.n_routed_experts)
                      # 每个专家干了多少活
                      ce = mask_ce.float().mean(0)
                      Pi = scores_for_aux.mean(0)
                      fi = ce * self.n_routed_experts
                      aux_loss = (Pi * fi).sum() * self.alpha
      mask_ce = F.one_hot(topk_idx_for_aux_loss.view(-1)
      假设我们有 2 个句子(bsz=2),每个句子只有 2 个 Token,只选 1 个专家(top_k=1)。 它原本长这样:
      [[2, 0]: 第 1 个句子的两个 Token 分别选了专家 2 和 0
      [1, 2]]:第 2 个句子的两个 Token 分别选了专家 1 和 2
      刚才展平的数组 [2, 0, 1, 2] 扔进 F.one_hot 后,会发生什么?
  • 第一个任务分给了专家 2 ➔ 变成选票 [0, 0, 1]
  • 第二个任务分给了专家 0 ➔ 变成选票 [1, 0, 0]
  • 第三个任务分给了专家 1 ➔ 变成选票 [0, 1, 0]
  • 第四个任务分给了专家 2 ➔ 变成选票 [0, 0, 1]
  • 最终生成的 mask_ce 就是一个形状为 (总任务数, 专家总数) 的二维矩阵
    • 这里的总任务数是所有batch
      ce = mask_ce.float().mean(0)
  • 第 0 列(专家 0)的平均值:(0 + 1 + 0 + 0) / 4 = 0.25
  • 第 1 列(专家 1)的平均值:(0 + 0 + 1 + 0) / 4 = 0.25
  • 第 2 列(专家 2)的平均值:(1 + 0 + 0 + 1) / 4 = 0.5
    得到每个专家实际抢到 Token 的全局频率

    MOE FeedForward

  • 输入x变化到:(batch_size * seq_len, hidden_dim)
    • 这个二维矩阵,行可以看作分了batch_size组,每个组里有seq_len个数连续排
      flat_topk_idx = topk_idx.view(-1)
  • (batch_size seq_len top_k)
  • 连续的top_k个数对应一个token的top_k个专家索引

算的是x中所有token经过不同专家处理并将结果根据token对应的专家权重叠加后的x

训练

x = x.repeat_interleave(self.config.num_experts_per_tok, dim=0)

  • 把x沿batch维度复制,挨个连续复制(交错复制)到top_k
    • flat_topk_idx刚好对应
      1
      2
      3
      4
      5
      6
      7
      8
      9
      10
                  for i, expert in enumerate(self.experts):
                      # x[...]:把属于第 i 号专家的所有 Token 提取出来
                      expert_out = expert(x[flat_topk_idx == i])
                      # 如果这个专家分到了 token
                      if expert_out.shape[0] > 0:
                          y[flat_topk_idx == i] = expert_out.to(y.dtype)
                      else:
                          # 数学上相当于加了 0,不影响结果;但对于 autograd(自动求导)来说,它证明该专家参与了前向传播,避免分布式训练时梯度同步死锁
                          y[flat_topk_idx == i] = expert_out.to(
                              y.dtype) + 0 * sum(p.sum() for p in expert.parameters())
      每一个expert对应一个FFN层
  • 输入与门控权重相乘,维度上升
  • 经过SiLU函数控制维度内的数,大的通过,小的阻止
  • 只有门控通道觉得“重要”的信息,才会在 Up 通道中被保留或增强
  • 再降维
  • 经过drop_out
    对输入的Token进行特征提取,将结果赋值给y
  1. 对于没有分到token的专家,反向传播就没有梯度,多卡训练(DDP)同步时就会崩溃
    • 在多卡训练(DDP)中,有一个非常严格的规定:每次前向传播,模型的所有参数都必须参与计算,从而在反向传播时产生梯度(哪怕梯度是 0)。 如果有任何一个参数没有参与计算图,DDP 在同步多张显卡的梯度时就会因为等不到那个参数的梯度,导致所有显卡互相死锁等待,程序直接崩溃卡死。
  2. for p in expert.parameters():将专家的所有权重都参与了计算
    1
    2
    3
    4
    5
    6
                # * 是 Python 的解包(Unpack)操作。它把 (bsz * seq_len, top_k) 这个元组拆开,变成了两个独立的参数 bsz * seq_len, top_k
                # y.view(*topk_weight.shape, -1:(batch_size*seq_len,top_k,dim)
                # .sum(dim=1):每个token将所有专家的报告相加
                y = (y.view(*topk_weight.shape, -1) *topk_weight.unsqueeze(-1)).sum(dim=1)
                # 将结果恢复成最开始的 3D 形状 [batch_size, seq_len, hidden_dim]
                y = y.view(*orig_shape)
    y.view(*topk_weight.shape, -1):
  • 这个y是二维矩阵,行可以看作分了batch_size组,每个组里有seq_len个token,每个token有top_k个连续排,top_k个连续排的经过view后变到了另一个维度,(batch_size* seq_len,top_k,dim)
  • topk_weight:每个 Token 对应的 Top-K 专家的概率权重,形状 (bsz * seq_len, top_k)
  • 乘法 (*):利用广播机制,给 Token A 的第一份报告乘上 1号专家的权重,第二份报告乘上 3号专家的权重,一个token的一个专家的所有维度都乘同一个数
  • .sum(dim=1):沿着第 1 维度(即那 2 份专家报告的维度)进行求和。
    • Token A 的两份加权报告被压扁成了一份,变成了 128 维。
    • Token B 的两份加权报告也被压扁成了一份,变成了 128
  • 这里的y是有所有专家的所有权重组成的,之所以能实现,就是做了repeat_interleave

    推理

    推理时,不计算梯度,使用 repeat_interleave 复制 Token 会非常浪费显存。这段代码利用了排序 (argsort) 和累计和 (cumsum),实现了极其高效的 Token 聚集分发计算
    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
        @torch.no_grad()
        def moe_infer(self, x, flat_expert_indices, flat_expert_weights):
            expert_cache = torch.zeros_like(x)
            # 这里的 flat_expert_indices,其实就是我们在前向传播(forward)里刚刚聊过的那个 flat_topk_idx
            # argsort: 返回对数组进行升序排序的索引。例如 [30, 10, 20] 返回 [1, 2, 0]
            idxs = flat_expert_indices.argsort() # 返回不同专家(从小到大)的索引
            token_idxs = idxs // self.config.num_experts_per_tok# 每个专家该去原始数据 x 里提取哪几行
            # bincount —— 统计每个专家要处理几个 Token
            # cumsum(0) 的作用是计算累积和(前缀和),返回数组,数组里的每一个数字,刚好就是每个专家在长队伍里的结束位置(不包含该位置),0是维度
            tokens_per_expert = flat_expert_indices.bincount().cpu().numpy().cumsum(0)
            # 当tokens_per_expert = [6, 15, 20, 26],tokens_per_expert.shape[0]即为专家数量(此时为4)
            # 且token_idxs = [3, 7, 19, 21, 24, 25,  4,  5,  6, 10, 11, 12...] 时
            # 意味token_idxs[:6] -> [3, 7, 19, 21, 24, 25]这6个位置属于专家0处理的token(每个token有可能被多个专家处理,这取决于num_experts_per_tok)
            # 接下来9个位置token_idxs[6:15] -> [4,  5,  6, 10, 11, 12...]属于专家1处理的token...依此类推
            for i, end_idx in enumerate(tokens_per_expert):
                start_idx = 0 if i == 0 else tokens_per_expert[i - 1]
                if start_idx == end_idx:
                    continue
                expert = self.experts[i]
                exp_token_idx = token_idxs[start_idx:end_idx]
                expert_tokens = x[exp_token_idx]
                expert_out = expert(expert_tokens).to(expert_cache.dtype)
                expert_out.mul_(flat_expert_weights[idxs[start_idx:end_idx]])
                expert_cache.scatter_add_(0, exp_token_idx.view(-1, 1).repeat(1, x.shape[-1]), expert_out)

            return expert_cache
    flat_expert_weightstopk_weight.view(-1, 1),连续的top_k个数对应一个token
    token_idxs = idxs // self.config.num_experts_per_tok
    因为idxs索引从0到batch_size seq_len top_k,我要找对应的token的索引,要整除top_k
    bincount: 统计每个专家要处理几个 Token,返回一个一维数组,长度小于等于专家数,小于是因为可能后面的专家分配的token数为0
    .cpu().numpy()
  • tokens_per_expert 留在 GPU 上。当 Python 执行 for 循环时,每次循环都需要单独向显卡要一个数字(end_idx)。显卡每次都要停下手头的大规模并行计算,把这一个微小的数字打包,跨越主板的 PCIE 通道传给 CPU。这种极其频繁的“单点通讯”,会造成严重的同步阻塞(Host-Device Sync),让推理速度慢得令人发指
  • cumsum(0) :在维度0计算累积和(前缀和),因为是一维数组
  • [start_idx:end_idx]:区间左开右闭
  • expert_out:(专家选择的Token数,dim)
  • start_idx:end_idx:该专家对应得token索引
    • 在token_idxs中提取的是在真实的x中token索引
    • 在idxs中提取的是x交错复制后的索引,与flat_expert_weights相对应,得到该专家在不同token的权重
      举例
  • Token 0 分配给 -> 专家 1、专家 3
  • Token 1 分配给 -> 专家 2、专家 5
  • Token 2 分配给 -> 专家 1、专家 2
    flat_expert_indices:[1, 3, 2, 5, 1, 2],与flat_expert_weights的权重对应
    argsortidxs:[0, 4, 2, 5, 1, 3]
    提取真实 Token IDtoken_idxs:[0, 2, 1, 2, 0, 1]
    .bincount():[0, 2, 2, 1, 0, 1]
  • 0 号专家:没出现过,分到了 0
  • 1 号专家:出现了 2 次,分到了 2
  • 2 号专家:出现了 2 次,分到了 2
  • 3 号专家:出现了 1 次,分到了 1
  • 4 号专家:没出现过,分到了 0
  • 5 号专家:出现了 1 次,分到了 1
    cumsum(0) 的作用是计算累积和(前缀和):[0, 2, 4, 5, 5, 6]
  • i=0 (0号专家)start=0, end=0start == end,直接 continue(它没分到 Token,休息)。
  • i=1 (1号专家)start=0, end=2。去队伍里切片 [0:2]。 拿到 token_idxs[0:2],也就是 [0, 2]。(1号专家成功拿到了属于它的 Token 0 和 Token 2!)
  • i=2 (2号专家)start=2, end=4。去队伍里切片 [2:4]。 拿到 token_idxs[2:4],也就是 [1, 2]。(2号专家成功拿到了 Token 1 和 Token 2!)
  • i=3 (3号专家)start=4, end=5。去队伍里切片 [4:5]。 拿到 token_idxs[4:5],也就是 [0]
  • i=4 (4号专家)start=5, end=5continue
  • i=5 (5号专家)start=5, end=6。去队伍里切片 [5:6]。 拿到 token_idxs[5:6],也就是 [1]
  • .mul_():注意这个下划线 _。在 PyTorch 中,带下划线的方法代表“原地操作(In-place)”。它不会在显存里新建一个矩阵来存结果,而是直接在 expert_out 原有的内存上把权重乘进去,把旧数据覆盖掉。
    • 得到不同token经过该专家的提取后并附上不同token对该专家的权重
  • expert_cache.scatter_add_(0, exp_token_idx.view(-1, 1).repeat(1, x.shape[-1]), expert_out)
    • 假设当前专家处理了 Token 0Token 2exp_token_idx 原本是 [0, 2]
    • .view(-1, 1) 将它竖起来,变成列向量
    • .repeat(1, x.shape[-1]) 将它在横向(特征维度)上复制
    • expert_cache.scatter_add_(0, 索引矩阵, expert_out)
    • expert_cache (初始状态):
      • [ [0.0, 0.0, 0.0], <— 等待接收 Token 0 的结果
      • [0.0, 0.0, 0.0], <— 等待接收 Token 1 的结果
      • [0.0, 0.0, 0.0] <— 等待接收 Token 2 的结果 ]
    • expert_out (1号专家算出的特征):
      • [ [1.1, 1.2, 1.3], <— 这是它给 Token 0 算的
      • [8.1, 8.2, 8.3] <— 这是它给 Token 2 算的 ]
    • 索引矩阵:
      • [ [0, 0, 0], <— 告诉 PyTorch: 第一行的数据,全给我扔到 0 号位置
      • [2, 2, 2] <— 告诉 PyTorch: 第二行的数据,全给我扔到 2 号位置 ]
    • dim=0,把索引矩阵里的数字当作行号(token的ID),将 expertout的数放到expert_cache对应token的位置,实现将不同专家的结果考虑专家的权重相加到同一token上
      `x.scatter_add
      (维度a, 索引矩阵index, A)`用到了两次,
  • 一次在算每个专家分配到的token数,x为ce(batch_size, 专家总数)
    • 维度为1,对应一个batch里的所有专家
    • 索引矩阵为专家索引((batch_size, seq_len * top_k))
    • A全为1,(bsz, seq_len * top_k),将A对应的1加在ce上
  • 一次将专家处理结果加在一起,x为expert_cache(batch_size * seq_len,dim)
    • 维度为0,对应不同token
    • 索引矩阵为(专家处理的token数(值为token的位置),dim)
    • A为(专家处理的token数,dim)
  • 总体逻辑:
  • A[i][j](第 i 行,第 j 列)
  • 查阅 index[i][j] 里的数字
  • dim = 0 意味着行坐标(第 0 维)由 index 说了算,列坐标保持不变。 投递目标:x[ index[i][j] ][ j ] += A[i][j]
  • dim = 1 意味着列坐标(第 1 维)由 index 说了算,行坐标保持不变。 投递目标:x[ i ][ index[i][j] ] += A[i][j]

MiniMind Block


输入x经过层归一化经attention层后与x相加,实现残差连接,结果经过层归一化经FFN再残差连接

MiniMind Model


输入索引(经过分词)经过embedding再dropout经过K层transformer层
经过层归一化
返回:

  • y(batch_size,seq_len,dim)
    将K层的transformer层的辅助损失加在一起,得到总的辅助损失

    MiniMind For CausalLM

    将y经过线性层(dim,vocab_size)映射到词表维度,再经过softmax选择该token的下一个输出token
    self.model.embed_tokens.weight = self.lm_head.weight:
  • nn.Linear 的权重矩阵在内存中是转置存放的
  • 输入词嵌入层(Token Embedding)的权重和输出分类层(LM Head)的权重绑定为同一个。这背后的直觉是:理解一个词的特征(Embedding)和生成一个词的特征(LM Head)本质上是同一件事
    logits_to_keep
  • 1:推理
  • 0:训练
  • 选取y的[:, -logits_to_keep:, :]

    训练

    1
    2
    3
    4
    5
    6
    7
            loss = None
            if labels is not None:
                # contiguous() 的作用是在内存中把数据重新排列成连续的块,防止后面的 view() 变形操作报错。
                shift_logits = logits[..., :-1, :].contiguous()
                shift_labels = labels[..., 1:].contiguous()
                loss = F.cross_entropy(
                    shift_logits.view(-1, shift_logits.size(-1)), shift_labels.view(-1), ignore_index=-100)
    做自回归
    logits:(模型的预测输出)[Batch大小, seq_len, vocab_size]
    labels:(标准答案)[Batch大小, seq_len
    F.cross_entropy:包含softmax得到概率和取负对数
    对于我们要预测的真实正确类别 $c$,模型计算出的概率 $p_c$ 的解析式为:
  • $x_c$:模型给正确答案打出的原始得分(logit)。
  • $K$:词表的总大小(比如 100,000 种可能的词)。
  • 分母:把词表里所有 100,000 个词的得分取自然指数($\exp$)后求和。这是为了把得分压缩到 0~1 之间并做归一化。
  • $p_c$在0~1之间,越接近0,损失越大
  • 合在一块

    RMSNorm

    1
    2
    3
    4
    5
    6
    7
    8
    9
    10
    11
    12
    13
    14
    15
    16
    class RMSNorm(torch.nn.Module):
        def __init__(self, dim: int, eps: float = 1e-5):
            super().__init__()
            self.eps = eps
            # nn.Parameter特殊的包装器,告诉 PyTorch:“这个张量不是普通的常量,而是模型的一部分。”
            self.weight = nn.Parameter(torch.ones(dim))
        def _norm(self, x):
            # x:[batch_size, seq_len, dim]
            # keepdim=True:保证结果的形状依然是 (batch_size, seq_len, 1),如果不写:(batch_size, seq_len)
            return x * torch.rsqrt(x.pow(2).mean(-1, keepdim=True) + self.eps)
        def forward(self, x):
            # x.float():将输入x提升到高精度的 float32
            # .type_as(x):核心步骤。检查原始输入x是什么类型(比如 torch.float16),然后把 _norm 的结果压回那个类型
            # .weight做广播运算
            # 将当前的张量(Tensor)转换为与目标张量 x 相同的数据类型(dtype)和设备(Device)
            return self.weight * self._norm(x.float()).type_as(x)

其中 $g_i$ 是可学习的缩放参数(即代码中的 self.weight),$\epsilon$ 是为了防止除以零的极小值

学习率计算

学习率从设置的最大值,沿着一条平滑的余弦曲线,慢慢降到最大值的 10%