minimind
minimind
Oppentokenizer
BPE((Byte Pair Encoding))
基础词表由 256 个单字节组成,随后通过统计算法将高频相邻字节合并,直到词表大小达到 6400
embedding
输入为:(batch_size,seq_len,dim)
位置编码
RoPE
1 | def precompute_freqs_cis(dim: int, end: int = int(32 * 1024), rope_base: float = 1e6, |
- 每个维度的基础旋转频率$\theta_i = \text{rope_base}^{-\frac{2i}{\text{dim}}}$
- 把一个长度为
dim(比如 128 维)的向量,切分成一个个二维的旋转平面。 - 每 2 个维度组成一个平面,共享同一个旋转频率($\theta_i$)。
- 所以,128 维的向量,实际上只有 64 个独立的旋转频率。
rope_base默认1e61
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。
- 把一个长度为
- YaRN外推
1
2
3
4
5
6
7
8if 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_fast或beta_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
2ramp = 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
2t = 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
4freqs_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_factortorch.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
13def 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 内核里,线程可以精确控制
| 机制 | 全称 | 结构特点 | 优点 | 缺点 |
|---|---|---|---|---|
| MHA | Multi-Head Attention | 每个 Q 都有对应独立的 K 和 V。 | 表达能力最强,能学习到多种不同的上下文关系。 | 显存占用高,推理时 KV Cache 巨大,速度慢。 |
| MQA | Multi-Query Attention | 所有的 Q 共享同一组 K 和 V。 | 极大地减少了 KV Cache 显存,推理速度最快。 | 可能会损失一定的模型精度(表达能力下降)。 |
| GQA | Grouped-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
10def 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就是MQAAttention
输入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 | class Attention(nn.Module): |
- 线性层不加偏置:
- 高维空间的注意力计算中,模型更在乎的是向量之间的相对方向(点积结果),而不是绝对位置
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 | class FeedForward(nn.Module): |
传统的 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:(专家总数, 隐藏层维度)权重初始化
为什么:在神经网络中,如果权重一开始设置得太小,信号在经过多层传递后会越来越弱,导致梯度消失(Vanishing Gradients);如果设置得太大,信号会越来越强,导致梯度爆炸(Exploding Gradients)。1
2def reset_parameters(self) -> None:
init.kaiming_uniform_(self.weight, a=math.sqrt(5))- 为了保证经过多个神经层后方差不变,而每个神经层输出要经过激活函数,这里是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
18def 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
| 特性 | Sequence 级别 (seq_aux=True) | Batch 级别 (seq_aux=False) |
|---|---|---|
| 颗粒度 | 极细,每一条数据内部都要求专家公平分配 | 较粗,只要整个批次(Batch)整体公平即可 |
| 推理表现 | 更均衡。单次推理时负载预测性好,不容易出现卡顿 | 可能存在波动。某些特定话题的句子可能导致局部专家过载 |
| 学习难度 | 较高。约束太死,有时会违背数据本身的语义聚集性 | 较低。给了模型弹性,允许专家在特定领域进行深度专业化 |
| 适用场景 | 中小模型、对推理延迟极其敏感的任务、下游任务微调 | 超大规模预训练、追求极致吞吐量的大规模集群训练 |
Sequence 级别
每个batch(句子)里面计算损失
batch_size:句子数
seq_len:一个句子的token数
- 为什么做:
- 大语言模型在推理(生成)的时候,是一条一条 Sequence 生成的,如果我们只做全局 Batch 级别的均衡,在实际单条推理时依然会出现某个专家过载的情况。所以现代的 MoE 模型(比如 Mixtral)很多都会引入 seq_aux(Sequence 级别的辅助损失),这就需要保留 bsz 这个维度来逐条计算
1
2
3
4
5ce = 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)`:
- 大语言模型在推理(生成)的时候,是一条一条 Sequence 生成的,如果我们只做全局 Batch 级别的均衡,在实际单条推理时依然会出现某个专家过载的情况。所以现代的 MoE 模型(比如 Mixtral)很多都会引入 seq_aux(Sequence 级别的辅助损失),这就需要保留 bsz 这个维度来逐条计算
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$
- 如果一个专家刚好拿到了平均数量的 Token,它的
- 辅助损失计算
1
2aux_loss = (ce * scores_for_seq_aux.mean(dim=1)
).sum(dim=1).mean() * self.alphascores_for_seq_aux:每一个 Token 对所有专家的原始“匹配度得分”,就是scoresscores_for_seq_aux.mean(dim=1):计算一个句子里token分配给不同专家的概率$P_i$ - 梯度计算
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$有关
- 举个极端的例子:
假设门控网络彻底“偏科”,把 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$)去变得更加平均
- 为什么乘上专家总数 $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
6mask_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.alphamask_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就是一个形状为(总任务数, 专家总数)的二维矩阵- 这里的总任务数是所有
batchce = 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个数连续排
- (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刚好对应每一个expert对应一个FFN层
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())- 输入与门控权重相乘,维度上升
- 经过SiLU函数控制维度内的数,大的通过,小的阻止
- 只有门控通道觉得“重要”的信息,才会在 Up 通道中被保留或增强
- 再降维
- 经过
drop_out
对输入的Token进行特征提取,将结果赋值给y
- 对于没有分到token的专家,反向传播就没有梯度,多卡训练(DDP)同步时就会崩溃
- 在多卡训练(DDP)中,有一个非常严格的规定:每次前向传播,模型的所有参数都必须参与计算,从而在反向传播时产生梯度(哪怕梯度是 0)。 如果有任何一个参数没有参与计算图,DDP 在同步多张显卡的梯度时就会因为等不到那个参数的梯度,导致所有显卡互相死锁等待,程序直接崩溃卡死。
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_cacheflat_expert_weights:topk_weight.view(-1, 1),连续的top_k个数对应一个tokentoken_idxs = idxs // self.config.num_experts_per_tok:
因为idxs索引从0到batch_size seq_len top_k,我要找对应的token的索引,要整除top_kbincount: 统计每个专家要处理几个 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=0。start == 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=5。continue。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 0和Token 2,exp_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的下一个输出tokenself.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
7loss = 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_lenF.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
16class 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%




