ARTICLE DETAIL

资讯详情

深耕郑州网站建设与运营推广的一线实战洞察。

因果注意力(Causal Attention)完整原理与从零 NumPy 实现

因果注意力(Causal Attention)完整原理与从零 NumPy 实现 一、概述因果注意力Causal Attention是自回归大模型的核心注意力机制是 GPT、Qwen、LLaMA 等生成式模型的标配也广泛用于时间序列预测、语音生成等时序任务。普通双向注意力Bidi Attention能看到整句所有 Token存在未来信息泄露问题而因果注意力严格遵循时序因果约束核心规则如下每个 Token 只能看见自身及前文内容绝对无法看见未来 Token。通俗理解模型逐词生成文本时只能依据已生成的内容续写不能偷看后续文本完美贴合真实的自回归生成逻辑保证模型学习到合规的时序生成规律。二、核心原理因果上三角掩码2.1 标准注意力计算公式在传统缩放点积注意力基础上加入因果掩码得到因果注意力完整计算公式Attention(Q,K,V)softmax(QK⊤dkMcausal)V Attention(Q,K,V) \text{softmax}\left(\frac{QK^\top}{\sqrt{d_k}} M_{causal}\right)VAttention(Q,K,V)softmax(dk​​QK⊤​Mcausal​)V参数释义Q、K、VQ、K、VQ、K、V分别为查询矩阵、键矩阵、值矩阵dk\sqrt{d_k}dk​​维度缩放因子避免点积结果数值过大导致梯度消失McausalM_{causal}Mcausal​因果掩码矩阵未来 Token 位置填充负无穷历史及当前 Token 位置置 0由于 Softmax 函数对负无穷输入的输出结果无限趋近于 0因此可以彻底屏蔽未来 Token 的注意力权重实现时序约束。2.2 掩码矩阵核心规则因果注意力采用上三角屏蔽机制通过np.triu(k1)生成掩码矩阵统一语义标准1 屏蔽未来位置0 可见当前/历史位置以序列长度为 4 的场景为例标准因果掩码矩阵如下[[0,1,1,1],[0,0,1,1],[0,0,0,1],[0,0,0,0]]矩阵解读矩阵第 i 行、第 j 列代表当前第 i 个 Token 对第 j 个 Token 的可见性若 j i 判定为未来 Token全部屏蔽。三、从零 NumPy 完整实现工业级稳定版本实现包含数值稳定 Softmax、动态数据类型适配、维度广播、数值异常兜底逻辑完全对标官方标准可直接用于学习与验证。importnumpyasnpdefsoftmax(x):数值稳定版 Softmax防止指数溢出、数值爆炸max_xnp.max(x,axis-1,keepdimsTrue)exp_xnp.exp(x-max_x)returnexp_x/np.sum(exp_x,axis-1,keepdimsTrue)defcausal_self_attention(Q,K,V): 因果自注意力机制自回归模型专用 :param Q: 查询矩阵 (batch, seq_len, dim) :param K: 键矩阵 (batch, seq_len, dim) :param V: 值矩阵 (batch, seq_len, dim) :return: attn_output(注意力输出结果), attn_weight(注意力权重矩阵) batch,seq_len,dimQ.shape# 1. 计算缩放点积注意力得分attn_scorenp.matmul(Q,K.transpose(0,2,1))/np.sqrt(dim)# 2. 构造因果掩码适配批量维度广播masknp.triu(np.ones((seq_len,seq_len)),k1)maskmask[np.newaxis,:,:]# 3. 未来位置填充对应数据类型极小值屏蔽未来信息min_valnp.finfo(Q.dtype).minattn_scorenp.where(mask1,min_val,attn_score)# 4. 计算归一化注意力权重attn_weightsoftmax(attn_score)# 数值异常兜底规避 NaN 异常attn_weightnp.nan_to_num(attn_weight,nan0.0)# 5. 权重加权求和输出注意力结果outputnp.matmul(attn_weight,V)returnoutput,attn_weight四、测试代码验证因果掩码有效性if__name____main__:batch_size2seq_length4dim8# 随机初始化 QKV 矩阵Qnp.random.randn(batch_size,seq_length,dim)Knp.random.randn(batch_size,seq_length,dim)Vnp.random.randn(batch_size,seq_length,dim)out,weightcausal_self_attention(Q,K,V)print(输出特征 Shape:,out.shape)print(\n 因果注意力权重矩阵首条样本)print(np.round(weight[0],4))五、实验验证结论运行上述代码可直观验证因果注意力的约束效果核心特征如下注意力权重矩阵严格下三角有效数值正常分布矩阵所有上三角位置权重无限趋近于 0完全被屏蔽彻底实现「仅可见前文、不可见未来」的时序因果约束六、核心应用场景因果注意力专为时序生成任务设计杜绝未来信息泄露主流应用场景大语言模型生成GPT、LLaMA、Qwen 等自回归生成模型核心主干机制时序预测任务金融数据、气象数据、流量数据时序预测语音/音频生成WaveNet 等时序音频生成模型流式序列建模所有需要逐帧迭代生成、禁止未来信息渗透的任务七、面试核心速记总结核心构成因果注意力 标准自注意力 上三角因果掩码核心作用杜绝未来信息泄露适配自回归时序生成逻辑实现关键通过np.triu(k1)生成未来位置掩码填充负无穷屏蔽权重核心差异双向注意力可查看全局 Token因果注意力仅单向查看历史 Token备注个人水平有限有问题随时联系~
返回列表