欢迎光临
我们一直在努力

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 16),KV 缓存内存与带宽降 8 倍,
// 精度损失很小(同组查询头看到同一组键值)。
// 说明:M1 朴素实现——逐位置读键算点积、惰性 softmax(减最大值防溢出);
// 为后续 AVX-512 优化(M2)提供教学对照。

#pragma once

// 引入基础类型(浮点/向量/索引)
#include "公共/基础定义.h"
// 引入 KV 缓存(注意力读取历史键值)
#include "推理/KV缓存.h"

// GQA注意力朴素:单个查询头的缩放点积注意力
// 参数:查询 = 单个查询头的向量(形状 [头维度],形状 [256]);
// 缓存 = 键值缓存(本层全部 KV 头的历史键值);
// 层 = 注意力层层号;查询头/KV头 = 当前查询头与其所属 KV 头号;
// 头维度 = 每头向量维数(qwen35moe 为 256);
// 上下文长度 = 参与注意力计算的位置数(= 查询位置 + 1,因果掩码边界);
// 缩放 = 1/√头维度(qwen35moe 为 0.0625);输出 = 结果向量(长度 头维度)
// 因果掩码说明:生成时查询位置由调用方决定,本函数对位置 p = 0..上下文长度-1
// 全部计入(调用方传 上下文长度 = 查询位置 + 1,即只含 ≤ 查询位置 的键,
// 后位键被掩码在外)。若 上下文长度 > 缓存当前长度 抛运行错误。
void GQA注意力朴素(const float* 查询, const KV缓存& 缓存, size_t 层, size_t 查询头,
size_t KV头, size_t 头维度, size_t 上下文长度, 浮点 缩放, float* 输出);

// GQA注意力朴素全部头:循环全部查询头,每头映射到所属 KV 头后调用单头版本
// 参数:查询全部 = 全部查询头的拼接向量(形状 [查询头数×头维度]);
// 缓存 = 键值缓存;层 = 注意力层层号;
// 查询头数 = Q 头数(qwen35moe 为 16);KV头数 = KV 头数(qwen35moe 为 2);
// 头维度 = 每头向量维数;查询位置 = 当前生成位置(决定因果掩码边界);
// 缩放 = 1/√头维度;输出全部 = 全部查询头的输出(形状 [查询头数×头维度])
// 说明:查询头 h → KV 头 h/(查询头数/KV头数);查询头数须为 KV头数 的整数倍,
// 否则抛运行错误
void GQA注意力朴素全部头(const float* 查询全部, const KV缓存& 缓存, 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)
#include <cmath>
// 引入标准头:float 最大值(防御校验用)
#include <limits>

// GQA注意力朴素:单个查询头的缩放点积注意力
// 实现:防御校验 → 逐位置算分数并记录最大值 → softmax(减最大值)→ 加权求和值
void GQA注意力朴素(const float* 查询, const KV缓存& 缓存, 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 头后调用单头版本
// 实现:防御校验分组关系 → 按组大小映射 → 逐头调用单头版本
void GQA注意力朴素全部头(const float* 查询全部, const KV缓存& 缓存, 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 = 8)
const size_t 组大小 = 查询头数 / KV头数;
// 因果掩码边界:上下文长度 = 查询位置 + 1(只含 ≤ 查询位置 的键,后位被掩码)
const size_t 上下文长度 = 查询位置 + 1;
// 逐查询头:h → KV 头 h/组大小,调用单头版本
for (size_t 查询头 = 0; 查询头 < 查询头数; ++查询头) {
const size_t KV头 = 查询头 / 组大小;
GQA注意力朴素(查询全部 + 查询头 * 头维度, 缓存,, 查询头, KV头,
头维度, 上下文长度, 缩放, 输出全部 + 查询头 * 头维度);
}
}

赞(0)
未经允许不得转载:171主机测评 » Attention:GQA注意力朴素
分享到: 更多 (0)

评论 抢沙发

  • 昵称 (必填)
  • 邮箱 (必填)
  • 网址