Attention:GQA注意力朴素
GQA注意力朴素.h// GQA注意力朴素.h —— 分组查询注意力GQA朴素标量实现声明// 用途10 个全注意力层使用。Q 16 头、KV 2 头分组比 8:1头维度 256。// 数学纯文本// attn softmax(Q·Kᵀ / √d_k) · V// 每组 n_q/n_kv 个查询头共享 1 个 KV 头查询头 h 对应 KV 头 h/(n_q/n_kv)。// GQA 相比 MHA 大幅减少 KV 头数2 vs 16KV 缓存内存与带宽降 8 倍// 精度损失很小同组查询头看到同一组键值。// 说明M1 朴素实现——逐位置读键算点积、惰性 softmax减最大值防溢出// 为后续 AVX-512 优化M2提供教学对照。#pragmaonce// 引入基础类型浮点/向量/索引#include公共/基础定义.h// 引入 KV 缓存注意力读取历史键值#include推理/KV缓存.h// GQA注意力朴素单个查询头的缩放点积注意力// 参数查询 单个查询头的向量形状 [头维度]形状 [256]// 缓存 键值缓存本层全部 KV 头的历史键值// 层 注意力层层号查询头/KV头 当前查询头与其所属 KV 头号// 头维度 每头向量维数qwen35moe 为 256// 上下文长度 参与注意力计算的位置数 查询位置 1因果掩码边界// 缩放 1/√头维度qwen35moe 为 0.0625输出 结果向量长度 头维度// 因果掩码说明生成时查询位置由调用方决定本函数对位置 p 0..上下文长度-1// 全部计入调用方传 上下文长度 查询位置 1即只含 ≤ 查询位置 的键// 后位键被掩码在外。若 上下文长度 缓存当前长度 抛运行错误。voidGQA注意力朴素(constfloat*查询,constKV缓存缓存,size_t 层,size_t 查询头,size_t KV头,size_t 头维度,size_t 上下文长度,浮点 缩放,float*输出);// GQA注意力朴素全部头循环全部查询头每头映射到所属 KV 头后调用单头版本// 参数查询全部 全部查询头的拼接向量形状 [查询头数×头维度]// 缓存 键值缓存层 注意力层层号// 查询头数 Q 头数qwen35moe 为 16KV头数 KV 头数qwen35moe 为 2// 头维度 每头向量维数查询位置 当前生成位置决定因果掩码边界// 缩放 1/√头维度输出全部 全部查询头的输出形状 [查询头数×头维度]// 说明查询头 h → KV 头 h/(查询头数/KV头数)查询头数须为 KV头数 的整数倍// 否则抛运行错误voidGQA注意力朴素全部头(constfloat*查询全部,constKV缓存缓存,size_t 层,size_t 查询头数,size_t KV头数,size_t 头维度,size_t 查询位置,浮点 缩放,float*输出全部);GQA注意力朴素.cpp// GQA注意力朴素.cpp —— 分组查询注意力GQA朴素标量实现// 数学纯文本// 缩放点积注意力attn softmax(Q·Kᵀ / √d_k) · V// 第一步分数 s_p 缩放 × 查询·键_p缩放 1/√头维度防点积过大使 softmax 饱和// 第二步softmax_p exp(s_p − max_s) / Σ exp(s_j − max_s)减最大值防溢出// 第三步输出 Σ_p softmax_p × 值_p// GQA 分组每组 查询头数/KV头数 个查询头共享 1 个 KV 头// 查询头 h → KV 头 h/组大小。单头版本直接用调用方传入的 KV头// 全部头版本负责把每个查询头映射到所属 KV 头。#include内核/注意力/GQA注意力朴素.h// 引入标准头指数函数softmax#includecmath// 引入标准头float 最大值防御校验用#includelimits// GQA注意力朴素单个查询头的缩放点积注意力// 实现防御校验 → 逐位置算分数并记录最大值 → softmax减最大值→ 加权求和值voidGQA注意力朴素(constfloat*查询,constKV缓存缓存,size_t 层,size_t 查询头,size_t KV头,size_t 头维度,size_t 上下文长度,浮点 缩放,float*输出){// 查询头 参数仅由全部头版本用于分组映射单头版本计算只依赖 KV头// 显式消用避免未使用参数警告分组映射见 GQA注意力朴素全部头(void)查询头;// 防御上下文长度不得超过缓存已写长度未写槽位读出的 0 值会导致错误分数if(上下文长度缓存.获取当前长度()){抛出运行错误(上下文长度超出缓存当前长度);}// 分数缓冲保存每个可见位置的缩放点积双精度累加点积减少舍入误差向量浮点分数(上下文长度);// 键临时缓冲逐位置读键复用避免反复分配向量浮点键(头维度);// 第一步对位置 p 0..上下文长度-1 计算 分数 缩放 × 查询·键[层][KV头][p]浮点 最大分数-std::numeric_limits浮点::infinity();for(size_t 位置0;位置上下文长度;位置){缓存.读取键(层,KV头,位置,键.data());长浮点 点积0.0;for(size_t 维号0;维号头维度;维号){点积static_cast长浮点(查询[维号])*static_cast长浮点(键[维号]);}分数[位置]static_cast浮点(点积*static_cast长浮点(缩放));// 记录最大值softmax 数值稳定用减最大值后 exp 不溢出if(分数[位置]最大分数){最大分数分数[位置];}}// 第二步softmax 分子 exp(s_p − max_s)并累加权重和分母长浮点 权重和0.0;for(size_t 位置0;位置上下文长度;位置){分数[位置]static_cast浮点(std::exp(static_cast长浮点(分数[位置])-static_cast长浮点(最大分数)));权重和static_cast长浮点(分数[位置]);}// 第三步输出 Σ_p softmax_p × 值[层][KV头][p]双精度累加减少误差向量浮点值(头维度);for(size_t 维号0;维号头维度;维号){输出[维号]0.0f;}for(size_t 位置0;位置上下文长度;位置){// softmax 权重 分子 / 权重和sum 正常化const浮点 权重static_cast浮点(static_cast长浮点(分数[位置])/权重和);缓存.读取值(层,KV头,位置,值.data());for(size_t 维号0;维号头维度;维号){输出[维号]权重*值[维号];}}}// GQA注意力朴素全部头循环全部查询头每头映射到所属 KV 头后调用单头版本// 实现防御校验分组关系 → 按组大小映射 → 逐头调用单头版本voidGQA注意力朴素全部头(constfloat*查询全部,constKV缓存缓存,size_t 层,size_t 查询头数,size_t KV头数,size_t 头维度,size_t 查询位置,浮点 缩放,float*输出全部){// 防御GQA 分组要求 查询头数 是 KV头数 的整数倍否则无法均匀分组if(查询头数0||KV头数0||查询头数%KV头数!0){抛出运行错误(查询头数必须是 KV 头数的整数倍);}// 组大小每组查询头数 查询头数 / KV头数qwen35moe 为 16/2 8constsize_t 组大小查询头数/KV头数;// 因果掩码边界上下文长度 查询位置 1只含 ≤ 查询位置 的键后位被掩码constsize_t 上下文长度查询位置1;// 逐查询头h → KV 头 h/组大小调用单头版本for(size_t 查询头0;查询头查询头数;查询头){constsize_t KV头查询头/组大小;GQA注意力朴素(查询全部查询头*头维度,缓存,层,查询头,KV头,头维度,上下文长度,缩放,输出全部查询头*头维度);}}