已关闭
flash_attn支持qk_head_dim=192、v_head_dim=128的组合 #5068
jiang-lirui创建于  26 天前关闭于  26 天前
jiang-lirui成员
26 天前 创建

flash_attn 支持 qk≠v(QK 192/V 128):head_dim_v 设计调研(2026-09)

0. 结论摘要

  1. 算子名核实:attention/flash_attn/ 实现的算子是 FlashAttn(aclnnFlashAttn,非 FIA);attention/flash_attn_metadata/ 是 FlashAttnMetadata(aclnnFlashAttnMetadata,AICPU 算子,无 q/k/v 张量输入,只有标量 attr)。FIA(fused_infer_attention_score)是另一个目录,仅作机制参考。
  2. kernel 侧 d/dV 已是独立模板参数:FAType 的 dBaseSize/dVBaseSize、TilingData 的 dSize/dSizeV、cube bmm1 的 K 轴(用 dSize)与 bmm2 的 N 轴(用 dSizeV)、vector 侧全部宽度(dVBaseSize/dSizeV)天然分离。QK=192/V=128 时 vector 侧(vec2 rescale、copy-out、FD)零改动;cube 侧 bmm1 的 dBaseSize=192 走 MatmulK 拆 K(128+64,D=256 同机制已上线)、bmm2 走 dVBaseSize=128 现有分支。
  3. 方案核心(v2):主算子不加参数——vHeadDim 已从 value.shape 末维推导(flash_attn_tiling_info_parser.cpp:451-460),checker 的 qkHeadDim == vHeadDim 相等性强校验(common_checker.cpp:251-256)改为组合白名单 {(64,64),(128,128),(256,256),(192,128)},拒绝集合除新增 (192,128) 外与现状完全一致。仅 metadata 算子加 head_dim_v(Int,OPTIONAL,默认 -1;-1 → headDimV = headDim_,行为不变)。tiling 侧新增 config 6/7 = (s1=128, s2=128, d=192, dv=128)/(s1=64, s2=256, d=192, dv=128),Config 位宽 3-bit→4-bit(range 0-5→0-15),config 0-5 的 tiling key 值与模板实例完全不变(零行为变化红线成立)。
  4. 最强仓内先例:FIA 的 tiling key Config 已是 10-bit 表驱动,config 表含 (192,128)、(576,512) 等 qk≠v 组合(fia_tiling_nonquant_gqa.cpp:63-64,69);flash_attn 的 inferDTemplateType 枚举已含 Aligned192=192(flash_attn_common_def.h:70);arch35FA 已定义 MLA_QKD_SIZE=192/MLA_VD_SIZE=128 常量(fa_tiling_info.h:311-315)——机制与枚举全部就绪,只差接线。
  5. 最大风险:(a) tiling key 变体数 192→256(+1/3),全量编译时间上涨(-j16 约束);(b) metadata AICPU 算子的 AdjustSinnerAndSouter 目前传 qk 维 headDim_(主算子传 vHeadDim)——qk=192 不落入 fa_adjust_sinner_souter.h:65/79 任何分档分支(<=128/==256 均不命中,停在预置默认 64/128),当主算子 ≤128 档的 32/256 子条件触发时两者切分不一致(当前因 qk==v 被掩盖);(c) infershape 的 PA layout 分支未读 v 的 D 维(回退 qk 维),qk≠v+PA 时输出 shape 会算错;(d) csrc 的 out 尺寸用 qk 维(§7 R4,存量 bug)。
  6. -1 哨兵有完整仓内先例(v2 后仅 metadata 需要):metadata checker 显式定义 NONE_VALUE = -1(flash_attn_metadata_check.h:37,:99-104 判 == NONE_VALUE || 合法区间,batchSize/maxSeqlenQ/Kv 同款);主算子的 win_left/max_seqlen 系列先例(flash_attn_def.cpp:104-115)仅作三段式范式参考。

1. API 现状

1.1 flash_attn 主算子(FlashAttn)

  • 接口签名:op_api/aclnn_flash_attn.h:75-84(aclnnFlashAttnGetWorkspaceSize,21 个输入/attr 参数 + 2 输出);执行段 aclnn_flash_attn.h:94。inner 层签名 aclnn_flash_attn_inner.h:21-28。无任何 head_dim 相关标量参数(v2 口径下也无需新增)。
  • 算子定义:op_host/flash_attn_def.cpp:25-147。输入 11 个(q/k/v 必选 + 8 个可选),输出 2 个(attn_out 必选 + softmax_lse 可选),attr 10 个(索引 0-9):softmax_scale/mask_mode/win_left/win_right/max_seqlen_q/max_seqlen_kv/layout_q/layout_kv/layout_out/return_softmax_lse(flash_attn_def.cpp:98-127)。注册仅 ascend950(:142)。
  • head_dim 推导来源(无标量,全从张量 shape 来):
    • qkHeadDim = query 的 D 维:op_host/flash_attn_tiling_info_parser.cpp:306-316(GetQkHeadDim → queryShape_->GetShapeD());
    • vHeadDim = value 的 D 维:flash_attn_tiling_info_parser.cpp:451-460(GetValueHeadDim → valueShape_->GetShapeD());
    • 二者经 GenerateAxisInfo 填入 FaTilingInfo(flash_attn_tiling_info_parser.cpp:618-619;结构定义 fa_tiling_info.h:204-205)。
  • op_api 层无 head_dim 校验(flash_attn.cpp/aclnn_flash_attn.cpp grep 无命中);l0op::FlashAttn(op_api/flash_attn.cpp:21-86)按位置透传 attr(OP_ATTR 宏 :66-67,78-79)。
  • infershape:op_host/flash_attn_infershape.cpp:99-136 从 q shape 推 headDim;L138-145 读 v 的末维得 headDimV(BSND/BNSD→GetDim(3)、TND→GetDim(2);L143 的 "PA_ND" 分支不是合法 layout_kv 值,等于死分支),输出 attn_out 末维 = headDimV(:155,161,166)。注意:PA_BBND/PA_BNBD/PA_NZ 布局未覆盖,headDimV 回退为 qk 的 headDim——qk≠v + PA 场景输出 shape 会错(详见 §4.3)。

1.2 flash_attn_metadata 算子(FlashAttnMetadata)

  • 签名:op_host/op_api/aclnn_flash_attn_metadata.h:20-25——4 个可选张量输入(cu_seqlens_q/kv、seqused_q/kv)+ 15 个标量 + 1 个 metadata 输出,无 q/k/v 张量。AICPU 实现(op_kernel_aicpu/flash_attn_metadata_aicpu.json,engine DNN_VM_AICPU)。
  • 必选 attr:num_heads_q/num_heads_kv/head_dim/soc_version/aic_core_num/aiv_core_num(读取于 op_kernel_aicpu/flash_attn_metadata_aicpu.cpp:59-63);可选 attr:batch_size/max_seqlen_q/max_seqlen_kv/mask_mode/win_left/win_right/layout_q/layout_kv/layout_out(:66-74)。L0 层 attr 名单见 op_host/op_api/l0_flash_attn_metadata.cpp:45-51。
  • 它计算/输出什么:读 seqlens 张量 → AdjustSinnerAndSouter 定 sOuter/sInner(flash_attn_metadata_aicpu.cpp:201-203)→ load_balance SectionStreamK 均衡切分(:241-244)→ 生成 metadata(GenMetadata :246-255):header(sectionNum/mBaseSize/s2BaseSize/isFD/aicNum/aivNum/outputLayout,SetMetadataHead :257-270)+ 每 section 每 AIC 核的任务区间(bn/m/s2 的 start/end,:272-295)+ 每 section 每 AIV 核的 FD 任务(:297-312)。mBaseSize/s2BaseSize 就是 kernel 的基本块参数(消费端 op_kernel/arch35/flash_attn_tiling_data.h:24-47 的 metadata index 常量)。
  • head_dim 在其中的作用:① 传给 AdjustSinnerAndSouter 的第一个参数(vHeadDim 语义位,flash_attn_metadata_aicpu.cpp:201);② 填 load_balance baseInfo 的 headDimQk/headDimV(:221-222,当前两者填同一个值);③ checker 白名单 {64,128,256}(op_host/flash_attn_metadata_check.h:112-114)。load_balance 侧 headDimQk 参与 s1Cost、headDimQk+headDimV 参与 s2vCost 权重(attention/common/op_kernel/load_balance/section_stream_k/section_stream_k_impl.h:327-335;字段定义 attention/common/op_kernel/load_balance/base_info.h:211-212)。
  • head_dim_v 的必要性在 metadata 算子是硬需求:它没有 v 张量可推 D 维,qk≠v 时 load balance 与 sOuter/sInner 档位选择都需要显式 v 维输入。

2. 现有 qk≠v 特例的确切机制(FIA,回源码核实)

任务提示中 doc 26 §1.4 的结论(FIA 有 Prefill MLA 192/128、Decode MLA 576/512)已逐条核实,机制如下(全部在 attention/fused_infer_attention_score/):

  1. rope 分离探测:parser GetRopeMode(op_host/fused_infer_attention_score_tiling_info_parser.cpp:1223-1238)——arch35 下:query_rope 与 key_rope 张量同时传入 → RopeMode::ROPE_SPLIT;否则若 qkHeadDim_ == 192 && vHeadDim_ == 128(硬编码!)→ RopeMode::ROPE_COMBINE;其余 → NO_ROPE。
  2. MLA 模式判定:GetMlaMode(同文件 :1265-1279)——ROPE_SPLIT + qk=512 → ROPE_SPLIT_D512(Decode MLA:q/k/v 512 + rope 64 分离,QK 总维 576/V 512);ROPE_COMBINE → ROPE_COMBINE_D128(Prefill MLA:q/k 合并存 128 维 + rope 64 = 192、v 128);ROPE_SPLIT + qk=128 → ROPE_SPLIT_D128(Prefill MLA 的分离形态:q/k/v 128 + rope 64)。
  3. isQKVDDifferent 白名单:fiaInfo.isQKVDDifferent = (qkHeadDim_ != vHeadDim_) && !(qkHeadDim_ == 192U && vHeadDim_ == 128U)(:1806)——192/128 被硬编码豁免,其他 qk≠v 组合全部视为"D 不等长"。checker 对 isQKVDDifferent 强制 qkHeadDim <= 128 && vHeadDim <= 128(op_host/checkers/common_checker.cpp:606-617)。
  4. tiling 模板路由:按 IsCapable() 优先级探测注册(op_host/arch35/fia_tiling_nonquant_mla.cpp:40-115,MLA 模板仅认 mlaMode==ROPE_SPLIT_D512 且 qk==512&&rope==64&&v==512&&n2==1,:86-99;isQKVDDifferent 直接踢出,:70)。GQA 模板(fia_tiling_nonquant_gqa.cpp:55-92)的 config 映射表已含 (SOUTER_128, SINNER_128, 192, 128)、(…, 256, 128)、(…, 128, 64)、(…, 64, 128) 等 qk≠v 四元组;Prefill MLA 分离形态(qk=128+rope64)在 UpdateTilingKeyConfig 特判映射到 Config_..._DAligned192_DVAligned128(:595-601)。MLA D512 路径固定 config = Config_..._DAligned576_DVAligned512(fia_tiling_nonquant_mla.cpp:451-455)。
  5. tiling key 编码:FIA 的 Config 是 10-bit、range 0-1023(op_kernel/arch35/fused_infer_attention_score_template_tiling_key.h:40-61),config 4/9 即 (192,128)/(576,512)。

结论:FIA 用"rope 张量分离 + (192,128) 硬编码豁免"支持 MLA,没有通用 head_dim_v 参数;qk≠v 进入 config 表的机制((s1,s2,d,dv) 四元组查表)是 flash_attn 可直接借鉴的成建制成例。flash_attn 的做法最简单:无 rope 分离输入,qk≠v 的唯一入口就是 v 张量自身的 shape(vHeadDim 已按 value 末维解析,§1.1),主算子无需任何新参数。


3. -1 哨兵值先例

v2 口径:-1 哨兵仅 metadata 算子的 head_dim_v 需要(主算子无新 attr)。仓内先例(v1 调研的全部先例保留作三段式范式参考):

参数 def 默认 -1 parser 归一化 消费
win_left/win_right flash_attn_def.cpp:104-109(.Int(-1)) flash_attn_tiling_info_parser.cpp:252-258(GetMaskParams,注释"不传或传入-1代表正无穷") arch35/flash_attn_tiling.cpp:124-129(-1→MASK_MODE_INT_MAX);SetFATilingData :326-331 同款
max_seqlen_q/kv flash_attn_def.cpp:110-115 flash_attn_tiling_info_parser.cpp:711-712(nullptr→-1) fa_adjust_sinner_souter.h:54-59(-1→MAX_SEQ_LEN_DEFAULT)
(metadata)batchSize/maxSeqlenQ/Kv aclnn 层默认 — flash_attn_metadata_check.h:37(NONE_VALUE = -1 常量,:99-104 判 == NONE_VALUE || 合法区间)——head_dim_v 直接复用该模式
(姊妹算子)quant_flash_attn win/max 系列 quant_flash_attn_def.cpp:115-118 同构 —

metadata 的 head_dim_v 沿用 NONE_VALUE 同款模式:aclnn/l0 层默认 -1 → checker 判 == NONE_VALUE || 白名单 → aicpu 归一化 headDimV = (-1 == attr) ? headDim_ : attr。主算子不注册任何 attr,v1 方案关注的"attr 索引 10 与 ATTR_IDX_DETERMINISTIC(flash_attn_infershape.cpp:35)死预留冲突"问题随参数取消而消失(不动该死预留;UT csv 的 deterministic 列是 TilingContextFaker::DeterministicInfo,tests/ut/framework_normal/common/tiling_context_faker.cpp:70-72,不占 attr 索引)。


4. checker 层现状与改动

4.1 现状的强制点

  • op_host/checkers/common_checker.cpp:240-250:supportedHeadDims = {64, 128, 256},qkHeadDim 与 vHeadDim 各自校验,报错文案 "The value of axis D of query and key can only be 64/128/256"(:242-244)与 "axis D of value"(:248-249)。
  • common_checker.cpp:251-256:强制 qkHeadDim == vHeadDim,报错 "The value of axis D of query/key must be equal to the value of axis D of value"——这是本方案要解除的核心检查。
  • 调用链:fa_checker.cpp:85-108(FAChecker::Process:CheckSinglePara → CheckParaExistence → CheckFeature → CheckMultiPara);CheckMultiPara 内部顺序 common_checker.cpp:430-454(Layout → dtype → headnum → dtype 一致性 → CheckShapeConsistency:444 → CheckAxis:447 → CheckAttnOutShape:450)。
  • shape 比对已按维分离:K 用 qkHeadDim、V 用 vHeadDim(连续布局 CheckKVShapeForContinuous,common_checker.cpp:334-344;PA 布局 CheckKVShapeForPageAttention :366-376);attn_out 用 vHeadDim(CheckAttnOutShape :409)。即 checker 的 shape 比对层天然支持 qk≠v,只需放开 CheckAxis 的相等性检查与白名单。
  • 其它 checker(mask/sinks/LSE/metadata/paged/seq_len)无 D 维校验(grep 核实无 qkHeadDim/vHeadDim 命中)。

4.2 放开 qk≠v 后的校验逻辑(v2:无新参数,纯 shape 驱动)

主算子不引入任何 attr,vHeadDim 推导来源不变(value.shape 的 D 维)。CheckAxis(common_checker.cpp:240-256)改为:

  1. qkHeadDim 白名单 {64, 128, 256} → {64, 128, 192, 256}(同步报错文案);
  2. vHeadDim 白名单保持 {64, 128, 256}(首期;(192,192)、(256,128) 等组合待 config 位段扩容后再开,见 §5.3);
  3. 相等性强校验 → 组合白名单:{(64,64), (128,128), (256,256), (192,128)};不在表内报 "The (head_dim of query/key, head_dim of value) combination is not supported, supported: (64,64), (128,128), (256,256), (192,128)";
  4. 拒绝集合等价性论证:现状拒绝 = {qk≠v} ∪ {qk∉{64,128,256} 或 v∉{64,128,256}};新拒绝 = 组合不在白名单。两者除新增放行的 (192,128) 外完全一致(例:(192,192) 旧因 192 不在白名单被拒、新因组合不在表内被拒)——旧调用无任何行为变化,且不存在"参数与 shape 不一致"的新防呆需求(无参数即无错配)。

4.3 infershape 的 PA 分支缺口(必须同 PR 修复)

flash_attn_infershape.cpp:139-145 只覆盖 BSND/BNSD/TND(及死的 "PA_ND" 分支),PA_BBND/PA_BNBD/PA_NZ 下 headDimV 回退为 qk 的 headDim(:138)。qk≠v + PA 布局时 attn_out 末维会按 qk 维下发(:155,161,166)→ 输出 shape 错误。修复:三种 PA 布局的 v 张量 D 维均在 GetDim(3)(PA_BBND (Bn,Bs,N2,D)、PA_BNBD (Bn,N2,Bs,D)、PA_NZ (Bn,N2,Bs/16,D,16)),补 else if (layoutKvStr == "PA_BBND" || "PA_BNBD" || "PA_NZ") && vShape->GetDimNum() >= 4 分支读 GetDim(3)。qk==v 时该改动无行为差异。


5. tiling 层

5.1 TilingData / do_tiling 中 head_dim 相关现状

  • TilingData(op_kernel/arch35/flash_attn_tiling_data.h:49-76):dSize/dSizeV/dSizeRope 三个 uint32 运行时字段。dSizeRope 是死字段(全目录 grep 仅定义与 OP_LOGD 打印,无填充无消费)。
  • SetFATilingData(arch35/flash_attn_tiling.cpp:301-347):dSize = qkHeadDim、dSizeV = vHeadDim(:310-311)——已是分离字段,零改动。
  • SplitPolicy(:115-136)→ AdjustSinnerAndSouter(vHeadDim, ...)(:130):传的是 v 维。分档逻辑 fa_adjust_sinner_souter.h:60-88:vHeadDim <= 128 → 默认 sOuter=64/sInner=128(短 Q 长 KV 且窗口宽时 32/256);== 256 → 64/128 或 32/256。qk=192/v=128 落 ≤128 档,与 D=128 的切分完全一致(192 不进 Adjust 分档,无新分支)。
  • CalcWorkspaceSize(:215-262):dSize = vHeadDim(:219)分桶 dVBasicBlock;L240 dnFlag_ && dSize > DSIZE_192 对 v=128 不触发(bmm2/vec2 split workspace=0,与 D=128 一致);FD workspace faTmpAttenGmSize = coreNum*2*mSize*dSize(v)(:254),与 kernel 侧 dSizeV 一致。全函数零改动。
  • kernel 侧 InitConstInfo:constInfo_.dSize/dSizeV 直读 TilingData(op_kernel/arch35/flash_attn_kernel_dn.h:168-169、_nd.h:168-169);dBasicBlock = Align64Func(dSizeV)(两文件 :204,v=128 → 128,当前无消费者)。

5.2 tiling key 编码与 config 空位

  • 编码(compile-ops skill 核实,源自 CANN template_argument.h 的 FastEncodeTilingKeyDirect):key = inputLayout_index | (kvLayout_index << 8) | (hasAttenMask << 16) | (config << 17)——参数从左到右放低位到高位,config 在 bit17 起,其 key 值 = config 索引值本身。
  • op_kernel/arch35/flash_attn_template_tiling_key.h:40-67:InOutLayoutType 8-bit(0-3)、KvLayoutType 8-bit(0-3)、HasAttenMask 1-bit、Config 3-bit,range 0-5(ASCENDC_TPL_3_BW 宏为 flash_attn 自定义,:34 于 common_def.h)。SEL 展开列表 :69-75。
  • ConfigValue 表(op_kernel/utils/flash_attn_common_def.h:148-166):六个 config 的 (s1, s2, d, dv) 四元组,d 与 dv 恒相等;config 宏 :169-178。
  • 空位分析:3-bit 剩 6/7 两个值。为容纳 (192,128) 需 2 个 config(sOuter 64/32 各一),恰好占满;为后续 (192,192)、(256,128) 等组合留演进空间,应一步到位扩 4-bit、range 0-15(FIA 先例 10-bit;bsa_select_block_mask 等算子已用 ASCENDC_TPL_4_BW,bsa_select_block_mask_tiling_key.h:19-21)。扩位不改变 config 0-5 的 key 值(值=索引,位宽只是上限),hasAttenMask 仍在 bit16,config 0-5 的 key 逐 bit 不变。
  • UpdateTilingKeyConfig(arch35/flash_attn_tiling.cpp:138-158):只认 qkHeadDim == 64/128/256(L147/149/151),其余报错 "can only be 64/128/256"(:153-157)——QK=192 的宿主硬卡点。
  • kernel 入口(op_kernel/flash_attn.cpp:126-129):ConfigValue[config] 推 dTemplateType/dVTemplateType;EnableSoftmaxDn(:85-92)只认 config 0/2(无 mask 时 Dn 路径)。

5.3 新增 config 方案

新增两个 config(且仅当 qk=192 && v=128 时选中):

config s1(mBase) s2(sInner) d(D) dv(DV) 对应 sOuter 形态 说明
6 128 128 192 128 64/128 无 mask 走 Dn、有 mask 走 Nd
7 64 256 192 128 32/256 短 Q 长 KV;仅 Nd
  • 资源占用论证(对照 config 4/5 已上线):L1_Q = mBase×dBase×2B(config 6: 128×192×2=48KB < config 4 的 128×256×2=64KB;config 7: 64×192×2=24KB < config 5 的 32KB);L1_KV:Dn 路径单槽 64KB 上限内 K tile = s2Aligned×192×2B(config 6: 128×192×2=48KB ≤ 64KB ✓);Nd 路径 L1_KV_LARGE_BUF = (s2BaseSize==256 && dBaseSize>128)(flash_attn_block_cube_nd.h:102-104)对 dBaseSize=192 判定为 true → 128KB 槽,K tile = 256×192×2=96KB ≤ 128KB ✓;UB_MM_RES = mBase/2×max(s2,dv)×4B(config 6: 64×128×4=32KB、config 7: 32×256×4=32KB,均 < config 4 的 64KB)。全部 ≤ 现有 D=256 档,L1/UB/L0C 无新风险(UB 布局定性结论引 plan 16 §三/§五)。
  • 变体数:4(layout)×4(kv)×2(mask)×6 = 192 → ×8 = 256(+64,+1/3)。全量编译时间线性上涨(compile-ops skill 的 -j16 约束下需评估);--tiling_key 精编可单验 config 6/7 的新 key。
  • 后续组合演进:4-bit 剩 8-15 共 8 个空位,(192,192)、(256,128)、(128,64) 等按需求逐个加表项即可;若超过 15 再评估 FIA 式查表重构(本调研不推荐本期做,避免 tiling key 语义漂移)。

5.4 兼容性红线论证(主算子零参数,旧场景零行为变化)

  1. API/attr 层(v2 最强保证):主算子不加参数、不动 aclnn 签名、不注册 attr——aclnn_flash_attn.h / flash_attn_def.cpp / parser 均零改动,旧调用在 API 层与 master 完全一致。
  2. checker 层:组合白名单除新增 (192,128) 外与旧拒绝集合完全一致(§4.2 第 4 点);qk 白名单加 192 只放宽不收紧。
  3. tiling 层:UpdateTilingKeyConfig 只新增 else if (qkHeadDim == 192 && vHeadDim == 128) 分支,64/128/256 分支一字不动 → 旧场景 config 0-5 不变;SplitPolicy/CalcWorkspaceSize/SetFATilingData 零改动。
  4. tiling key 数值:Config 扩 4-bit 后,config 0-5 的 key 值 = 索引值(bit17 起拼接),逐 bit 不变;已发布的 tiling key(如 README L433 示例 2279866368 = config 2)与已编译二进制/缓存完全兼容。6/7 是纯新增 key。
  5. kernel 层:config 6/7 是新增模板实例,0-5 的实例化集合不变;EnableSoftmaxDn 对旧 config 行为不变。
  6. metadata 算子:head_dim_v 默认 -1 → aicpu 归一化 headDimV = headDim_,除 §10-#6 的 v 维修正(R2)外路径不变;R2 修正本身在 qk==v 时是无操作(两值相等)。
  7. 唯一全局代价:全量编译变体 192→256(新增能力的编译时长,不是行为变化)。既有场景(q==v、qk∈{64,128,256})的 tiling key、TilingData 内容、kernel 实例、数值结果与 master 完全一致。

6. kernel 层

6.1 HEAD_DIM 模板实例化现状

  • 模板链:flash_attn.cpp:126-129(ConfigValue → d/dVTemplateType)→ FAType(op_kernel/utils/flash_attn_type.h:103-121,dBaseSize/dVBaseSize 为独立模板参数)→ FlashAttentionNoQuantGqaKernelDn/Nd + Block/Vec(flash_attn.cpp:141-174)。入口零改动(config 6/7 经既有 ConfigValue 表自动路由)。
  • 现有档位:ConfigValue 只实例化 d=dv∈{64,128,256}(flash_attn_common_def.h:148-166);但 inferDTemplateType 枚举已含 Aligned192=192/576 等(:61-75),公共 regbase 枚举同(doc 28 §2.1 已核实)。

6.2 V 的加载 / UB buffer / bmm2 / 输出(区分 HEAD_DIM_V 的点)

位置 现状 qk=192/v=128 时
V GM 寻址 InitKVBuffer(..., constInfo_.dSizeV, valueGm_, ...)(flash_attn_block_cube_dn.h:143-144、_nd.h 同构) v=128 原值 ✓ 零改动
V L1 拷入 CopyValueSlice(..., 0, constInfo_.dSizeV, ...)(dn :392、nd :416),Nd2Nz pad 语义(attention/common/op_kernel/memcopy/attn_copy_gm_to_l1.h:26-72) 128 对齐良好 ✓ 零改动
bmm2 N 维 nLoops = (dSizeV+127)/128、尾块 tileN(dn :396-398、nd :420-422);MMParam{mBase, tileN, s2, ...};vL1Offset = n*128/16*s2Aligned*16(:402-405) dV=128 → nLoops=1、tileN=128,与 config 2/3 的 bmm2 完全同路径 ✓
bmm2 matmul 模板 dVBaseSize > 128 走 MatmulFull、否则 s2BaseSize==128 走 MatmulFull<128, dVBaseSize, 128> / MatmulBase(dn :408-419、nd :433-444) dVBaseSize=128 走现有分支 ✓ 零改动
Fixpipe 写 UB dstStride=(dVBaseSize+15)>>4<<4(dn :379、nd :403) 128 ✓ 零改动
vec2 rescale/o 累加 dTemplateAlign64 = Align64Func(dVBaseSize)(flash_attn_block_vec_dn.h:67、_nd.h:69);FlashUpdateNew/LastDivNew 按该宽度(vec_dn :519-527、vec_nd :635-643);UB_VEC2_RES_BUF_BYTES = mBase/2×dTemplateAlign64×4B(vec_dn :103、vec_nd :112) 128 ✓ 零改动
输出写出 Bmm2DataCopyOutTrans:FaUbTensor{colCount=dTemplateAlign64} + dDealSize=dSizeV(vec_dn :431-437、vec_nd :547-553);InitAttenOutBuffer(..., dSizeV, ...)(vec_dn :160) 128 ✓ 零改动
FD(flashdecode) dSizeV_Align_ = Align(dSizeV, 64)(flash_attn_block_vec_flashdecode.h:154);CopyAccumOutIn 的 rightPadding/dstStride(:250-259);ReduceFinalRes_VF/DealInvalidRows 用 dSizeV_Align_(:313-335) 128 ✓ 零改动
输出总量 attenOutTotalSize = t*n2*g*dSizeV(vec_dn :271、vec_nd :297) ✓ 零改动

6.3 QK=192 侧的 kernel 机制(bmm1)

  • bmm1 的 K 轴用 constInfo_.dSize(QK 维):Dn MakeMMParam(s2, m, dSize, false, true)(flash_attn_block_cube_dn.h:324-325)、Nd(flash_attn_block_cube_nd.h:341-342)。
  • dBaseSize > 128 时走 MatmulK<K=128 循环>(dn :326-332、nd :343-350;MatmulK 实现按 kLoops=(singleK+127)/128、尾块 tileK=singleK%128 拆 K,attention/common/op_kernel/matmul.h:811+,L872-882 同构的 MatmulKbias 可见拆分公式)——192 = 128+64,D=256(128+128)同机制已上线;Mmad k=64 满足 16 元素分形(doc 28 §3.1 的粒度结论)。
  • Q/K L1 拷入 CopyQuerySlice/CopyKeySlice(..., dSize, ...)(dn :311,318):Nd2Nz 的 dValue=192,行 stride 16 对齐(:234,249)——192%16=0 无 pad 语义依赖,与 256 同(不同于 head_dim 72 场景,无需 L1 零初始化)。
  • UnInitCrossCoreSync 的 dBaseSize <= 128 分支(dn :211-216):192 走 >128,与 256 同。
  • pse/quant 路径:flash_attn 仅非量化(flash_attn_def.cpp:14 注释"仅非量化"),无 pse/quant 输入 → 无 HEAD_DIM_V 区分点。sinks 作用于 S2 轴(SinkConstInfo 仅 bool),LSE 每 head 一个 float——均与 D 维解耦。
  • arch35 分工:AIC(__DAV_C310_CUBE__,flash_attn.cpp:150-154/165-169)跑 CubeBlock(bmm1/bmm2/L1),AIV 跑 VecFaBlock(vec1/vec2/copy-out)+ VecFdBlock(FD)——config 6/7 下 d 与 dV 模板分别只影响各自 block,分工点不变。

6.4 kernel 层唯一改动点

EnableSoftmaxDn(flash_attn.cpp:85-92):(config == 0) || (config == 2) → 追加 || (config == 6)(无 mask + qk=192/v=128 走 Dn;Dn 的 K L1 tile=48KB≤64KB 单槽可行)。若首期求稳可不动(192/128 全走 Nd),但 Dn 是无 mask 主路径(MLA prefill 无 mask 常见),建议同 PR 落地。可选:cube 侧加 static_assert 看护 L1_KV tile 上限(引 plan 16 的 UB 边界验证方法)。


7. torch_extension / python 层(v2:主算子链路仅修 csrc 尺寸 bug)

  • 主算子 flash_attn:schema/fake/dispatch/fallback 全部零改动(无新参数);fake 的 d_size 已从 v.size(2/3) 推导(:132,138,144),qk≠v 天然正确。
  • csrc(torch_extension/csrc/flash_attn.cpp:37-98):FlashAttn 不加形参;唯一必改点是 :53-67 的 dSize 从 q 的 shape 推导改为 v 的 shape(TND v.size(2)、BSND/BNSD v.size(3))——attention_out 尺寸(:79-85)当前用 qk 维,qk≠v 时 out shape 错(R4,与 head_dim_v 参数无关的存量 bug)。
  • metadata flash_attn_metadata:schema 加 int? head_dim_v=None;fake(:82-103)签名加 head_dim_v(仅透传);dispatch(:177-233)形参 + :215-233 位置传参列表加 head_dim_v(默认 -1 归一化复用 :202-209 的 xxx = -1 if xxx is None else xxx 模式);fallback(:239-276)同步;csrc FlashAttnMetadata(:25-35)加 headDimV 透传。
  • graph_convert(graph_convert_flash_attn.py:49-101):GE 转换器是 raise AssertionError("GE not supported!") 的占位(:100-101),仅 metadata 签名加 head_dim_v 保持与 schema 一致(注意该文件签名带 deterministic 而 schema 无——历史漂移,本任务顺带对齐但不扩 scope)。
  • 测试侧 python(详 §9)。

8. graph 模式风险

  • aclgraph/torchair 调用链:pytest 的 --graph_mode(tests/pytests/test_flash_attn.py:86,301,390)走 FlashAttnGraphNetwork(tests/pytests/core/backends/npu.py:35-74,位置传参,forward 签名 :38-60)→ torchair aclgraph 后端(npu.py:205-216 的 _get_graph_net,开 _aclnn_static_shape_kernel)。v2 主算子无新参数 → GraphNetwork.forward 签名与 net(...) 调用点零改动(v1 方案的"位置参数错位"风险整体消除);仅 metadata 的 python 生成链需同步(§9.1)。
  • 静态图 tiling:kernel 侧已有静态图适配(flash_attn.cpp:132-137 注释:静态图下 tiling 是编译期常量数组,GET_TILING_DATA_PTR_WITH_STRUCT 双展开)。qk≠v 是新 shape 场景,无 attr 依赖问题;新 config 6/7 的 tiling key 必须全量编译进 kernel json(否则 graph 模式报 Cannot find tilingKey,compile-ops skill §故障表 L203)——禁止用 --tiling_key 精编包交付。
  • graph 模式的 metadata 生成:tests/assets/impl/graph.py:47 head_dim = int(q.shape[-1])(用 qk 维)→ qk≠v 时应改 int(v.shape[-1])(AdjustSinnerAndSouter 的 v 维语义,见 §1.2/§10 风险 R2)。
  • shape 变化对 aclgraph 的影响:qk≠v 是新输入 shape → 新 tiling key → 新二进制,属正常编译矩阵扩展,无机制性风险。

9. 测试与文档

9.1 仓内 pytest(attention/flash_attn/tests/)

  • 参数化键 'D'/'DV' 已存在:tests/pytests/core/case_loader.py:55(c.setdefault("DV", c.get("D")))、tests/assets/gen_cases.py:115;tests/pytests/core/data.py:262(v shape 用 params.get("DV", params["D"]));golden tests/pytests/core/backends/cpu_golden.py:113(d_v = kwargs.get("DV", d))——golden 与输入生成已支持 D≠DV。
  • 用例组织:tests/pytests/test_cases/functional_nc.py 等字典参数化("D": [128], "DV": [128],:23 起);readme.md:282 已文档化 "DV: value head dim,默认 = D"。
  • 改动点:
    • functional_nc.py 加 (D=192, DV=128) 用例组(无 mask/Dn、有 mask/Nd、TND varlen、PA_BBND/PA_BNBD decode + FD、LSE 开关、bf16+fp16;PA_NZ 亦应覆盖——192%16=0,NZ 的 d1=192/16=12 合法,与 72 场景不同);
    • 负例:(D=256, DV=128)(首期未支持的组合报错);(D=192, DV=192)(报不支持)——v2 无主算子参数,不存在"参数与 shape 不一致"类负例;
    • 等价性:主算子 (D, DV)=(128,128) 现有用例即回归基线(无 -1/显式之分);metadata 层加 head_dim_v=-1(缺省)与显式 =D 的 metadata 输出一致性用例;
    • data.py:121 的 meta_kwargs["head_dim"](从 q shape 推)+ 新增 head_dim_v(从 v shape 推,缺省 -1)——metadata 参数组与 kernel 参数组对齐;
    • graph 模式:npu.py GraphNetwork 零改动;tests/assets/impl/graph.py:47 改 v 维。
  • UT:tests/ut/op_host/arch35/test_flash_attn_tiling.csv(v_shape 独立列,:1 表头)加 192/128 行,expectTilingKey 填 config 6/7 的新 key——v2 主算子无新 attr,csv 无需新增 head_dim_v 列,test_flash_attn_tiling.cpp/TilingContextFaker 的 attr 装配零改动。tests/ut/op_host/test_flash_attn_shape_infershape.csv 加 PA 布局 qk≠v 行(看护 §4.3 修复)。

9.2 testkit(/workspace/ops-transformer-testkit)

  • golden 已支持 v_dim(golden/flash_attn_golden.py:130-147,"DV != D works",v 维从 v tensor shape 运行时推导)✓ 零改动;
  • backends/flash_attn.py:90,131,155 的 d = case.get("D", 128) 用于 metadata 生成——需加 dv = case.get("DV", d) 并在 metadata 调用处传 head_dim_v=dv(v2:metadata 走新参数而非改 head_dim 语义);用例文件(test_cases/flash_attn/{nodrop_cases,dropmask_cases,pa_decode_baseline,pa_decode_random_kv}.py)加 D/DV=192/128 条目。

9.3 README(attention/flash_attn/README.md)

  • §基本块 config 表(:348-352,注意现状漏列 config 5,顺带修正)加 config 6/7 行与 (D, DV) 说明;:363 的 D→config 映射文字加 "D=192 且 DV=128 → config6/7";
  • §TilingKey(:429-433)补 config 4-bit 说明与新 key 示例;
  • §Dn/Nd 路由表(:444-452)补 config 6 进 Dn;
  • 参数说明/支持规格表补 (qk, v) 组合列表(主算子无新参数,qk≠v 由 v.shape 表达);metadata 算子部分补 head_dim_v(默认 -1 语义、"必须等于后续 flash_attn 调用中 value 末维"的契约);tests/pytests/readme.md:282 附近补 metadata head_dim_v 透传说明。

10. 修改方案(分层清单,格式对齐 doc 29;v2 修订版)

层次顺序:torch_extension → op_api/aclnn → 算子定义/checker → infershape → metadata 算子 → doTiling → tilingKey → tilingData → kernel → tests → README。
v2 口径:主算子 flash_attn 零新参数(aclnn 头/实现、def、parser 全部不动);head_dim_v 仅存在于 metadata 算子链(aclnn → l0 → checker → aicpu)。

# 层 文件(attention/ 下) 改动 性质 量
1 torch csrc flash_attn/torch_extension/csrc/flash_attn.cpp FlashAttn 的 out 尺寸 dSize 改从 v.shape 推(L53-67,R4);FlashAttnMetadata(L25-35)加 headDimV 形参透传 修改 ~20 行
2 torch python flash_attn/torch_extension/flash_attn.py 仅 metadata:schema/fake/dispatch/fallback 加 head_dim_v(默认 -1 归一化);主算子路径零改动 修改 ~30 行
3 graph_convert flash_attn/torch_extension/graph_convert_flash_attn.py 仅 metadata 占位签名加 head_dim_v 修改 ~2 行(勘误 v14:核实该文件只有 FlashAttn 主算子占位转换器、无 metadata 函数,本条目前提不成立,实际零改动为唯一自洽做法)
4 aclnn(metadata) flash_attn_metadata/op_host/op_api/aclnn_flash_attn_metadata.h/.cpp、l0_flash_attn_metadata.cpp 签名加 int64_t headDimV;attr 名单(L43-51)加 "head_dim_v"——主算子 aclnn/inner/l0op 零改动 修改 ~25 行
5 metadata checker flash_attn_metadata/op_host/flash_attn_metadata_check.h head_dim 白名单 {64,128,256}→{64,128,192,256}(L112-114);新增 head_dim_v:==NONE_VALUE(-1) 或 ∈{64,128,256};非 -1 时 (head_dim, head_dim_v) 组合白名单(与主算子 §4.2 一致) 修改 ~30 行
6 metadata aicpu flash_attn_metadata/op_kernel_aicpu/flash_attn_metadata_aicpu.cpp/.h 读可选 attr "head_dim_v";归一化 headDimV = (-1==attr) ? headDim_ : attr;L201 AdjustSinnerAndSouter 改传 headDimV(R2);L221-222 baseInfo.headDimV 填真值;examples/test_aclnn_flash_attn_metadata.cpp 同步 修改 ~25 行
7 主算子 checker flash_attn/op_host/checkers/common_checker.cpp CheckAxis:qk 白名单加 192;qkHeadDim != vHeadDim 相等性 → 组合白名单(§4.2) 修改 ~25 行
8 infershape flash_attn/op_host/flash_attn_infershape.cpp 补 PA_BBND/PA_BNBD/PA_NZ 的 headDimV=GetDim(3)(L139-145,R3) 修改 ~12 行
9 doTiling flash_attn/op_host/arch35/flash_attn_tiling.cpp UpdateTilingKeyConfig 加 else if (qkHeadDim == 192 && vHeadDim == 128) → config 6/7(sOuterFactor 二选一);报错文案改为组合列表 修改 ~15 行
10 tilingKey flash_attn/op_kernel/arch35/flash_attn_template_tiling_key.h(Config 3-bit→4-bit、range 0-15;SEL 加 6,7)、op_kernel/utils/flash_attn_common_def.h(ConfigValue 加两项 + 宏 6/7) 修改 ~25 行
11 tilingData flash_attn/op_kernel/arch35/flash_attn_tiling_data.h 零改动(dSize/dSizeV 已有;dSizeRope 死字段不动) — 0
12 kernel 入口 flash_attn/op_kernel/flash_attn.cpp EnableSoftmaxDn 加 || (config == 6)(L91) 修改 ~2 行
13 block_cube/vector flash_attn/op_kernel/arch35/flash_attn_block_cube_{dn,nd}.h、flash_attn_block_vec_{dn,nd,flashdecode}.h、flash_attn_kernel_{dn,nd}.h 零改动(§6.2/§6.3 逐点核实;可选 L1 tile static_assert 看护) — 0(可选 ~10 行)
14 tests 仓内 pytests(§9.1)+ UT csv + testkit(§9.2)+ examples 用例 + data.py/graph.py 链路 + csv 行 新增/修改 ~280 行
15 README flash_attn/README.md、tests/pytests/readme.md §9.3 修改 ~35 行

工作量估计:代码 ~200 行(1 人日)+ 测试 ~280 行(1.5 人日)+ 全量编译回归(256 变体、-j16、graph mode 用例,1-2 天)≈ 4 人日(v1 估 5 人日;主算子参数链——aclnn 签名×4、def、parser attr 读取/归一化、attr 索引、UT csv 新列——整体免除)。

风险表:

风险 等级 守护
R1 tiling key 变体 +33% 编译时长 中 编译基线对比;--tiling_key 精编做开发期快速验证(不交付);参照 compile-ops skill 的 -j16/超时约束
R2 metadata 的 AdjustSinnerAndSouter 传参不一致(aicpu 现传 qk 维 headDim_,主算子传 vHeadDim) 高(隐蔽正确性) #6 强制改传归一化 headDimV;qk=192 不落入 fa_adjust_sinner_souter.h:65/79 任何分档(≤128/==256 均不命中,停在预置默认 64/128),当主算子 ≤128 档的 32/256 子条件(:74-77)触发时两者切分不一致——必须语义对齐并用 metadata+kernel 联调用例验证(PA/FD 场景)
R3 infershape PA 分支 headDimV 回退 qk 维 高 #8 与 #7 同 PR;UT infershape csv 加 PA qk≠v 行看护
R4 csrc out 尺寸仍用 qk 维 高 #1 同 PR;D=192/DV=128 用例必现(out shape 错会被 golden 对比抓到)
R5 metadata head_dim_v 与实际 v.shape 不一致(metadata 无张量可交叉校验) 中 checker 只能校验值域/组合;靠 README 契约 + examples 示例约束"head_dim_v 必须等于后续 flash_attn 调用中 value 的末维";联调用例覆盖
R6 graph 模式 低(v2 缩小) 主算子无新参数,GraphNetwork 签名/调用零改动;仅 metadata python 生成链(data.py/graph.py)同步,graph_mode 用例进回归清单
R7 (192,192)/(256,128) 等组合被 shape 放开但无 config 中 checker 组合白名单兜底报"不支持";config 表后续按需扩(4-bit 剩 8 空位)
R8 旧场景回归退化 低 红线论证 §5.4(v2:主算子零参数,API 层与 master 完全一致);全量 pytest 回归 + 既有 tiling key 抽样比对

开放问题(不阻塞):

  1. (192,192)、(256,128)、(128,64) 等组合的 config 排期(4-bit 剩 8 个空位够用 4 个 (d,dv) 组合 ×2 形态;超过 15 再评估 FIA 式查表重构)。
  2. dSizeRope 死字段是否顺势清理(与本需求无关,建议不动)。
  3. metadata 算子 head_dim 参数的语义文档化(现状用户传 QK 维;qk≠v 后文档明确 "head_dim 为 QK 维,v 维由 head_dim_v 指定,缺省等于 head_dim")。
  4. graph_convert_flash_attn.py 签名带 deterministic 而 schema 无——历史漂移是否顺带对齐(建议仅对齐不扩 scope)。
  5. 主算子 checker 报错文案从 "must be equal" 改为组合列表——对外文案变化(低风险,一般无程序匹配文案)。

参考来源(本地一手代码,路径省略前缀 /workspace/ops-transformer-fa-qk192v128/)

  1. attention/flash_attn/op_api/aclnn_flash_attn.h(L75-84 主算子签名;L94 执行段)
  2. attention/flash_attn/op_api/aclnn_flash_attn_inner.h(L21-28 inner 签名)
  3. attention/flash_attn/op_api/flash_attn.cpp(L21-86 l0op 透传;L66-67/78-79 OP_ATTR)
  4. attention/flash_attn/op_host/flash_attn_def.cpp(L25-147 OpDef;L98-127 attr;L142 ascend950)
  5. attention/flash_attn/op_host/flash_attn_infershape.cpp(L99-145 headDim/headDimV 推导;L147-171 输出 shape;L35 死预留索引)
  6. attention/flash_attn/op_host/fa_tiling_info.h(L149-186 FAParaInfo;L204-205 qkHeadDim/vHeadDim;L311-315 MLA 常量;L340-345 DSIZE 系列)
  7. attention/flash_attn/op_host/flash_attn_tiling_info_parser.cpp(L213-250 GetAttrParaInfo;L252-258 win -1 归一化;L306-316 GetQkHeadDim;L451-460 GetValueHeadDim;L618-619 GenerateAxisInfo;L711-712 maxSeq -1)
  8. attention/flash_attn/op_host/flash_attn_tiling.cpp(L40-66 TilingFlashAttn 主入口;L89-92 注册)
  9. attention/flash_attn/op_host/arch35/flash_attn_tiling.cpp(L115-136 SplitPolicy;L130 传 vHeadDim;L138-158 UpdateTilingKeyConfig;L215-262 CalcWorkspaceSize;L301-347 SetFATilingData)
  10. attention/flash_attn/op_host/checkers/common_checker.cpp(L240-256 D 白名单与相等性强校验;L328-346/349-379/403-415 shape 比对;L430-454 CheckMultiPara)
  11. attention/flash_attn/op_host/checkers/fa_checker.cpp(L37-46 注册表;L85-108 Process)
  12. attention/flash_attn/op_host/fa_adjust_sinner_souter.h(L50-88 vHeadDim 分档——仅 <=128(L65)与 ==256(L79)两分支;L54-59 maxSeq -1 归一化)
  13. attention/flash_attn/op_kernel/arch35/flash_attn_template_tiling_key.h(L40-67 参数声明与 Config 3-bit)
  14. attention/flash_attn/op_kernel/utils/flash_attn_common_def.h(L61-75 inferDTemplateType 含 Aligned192;L134-178 ConfigParams/ConfigValue/宏;L243-293 CommonConstInfo)
  15. attention/flash_attn/op_kernel/utils/flash_attn_type.h(L103-121 FAType 的 d/dV 独立模板参数)
  16. attention/flash_attn/op_kernel/arch35/flash_attn_tiling_data.h(L49-76 dSize/dSizeV/dSizeRope;L24-47 metadata index)
  17. attention/flash_attn/op_kernel/flash_attn.cpp(L85-92 EnableSoftmaxDn;L126-129 模板推导;L132-137 静态图 tiling;L141-174 Dn/Nd 分发)
  18. attention/flash_attn/op_kernel/arch35/flash_attn_block_cube_dn.h(L93-103 buffer 常量;L143-144 V 用 dSizeV;L306-348 bmm1;L372-385/387-432 bmm2)
  19. attention/flash_attn/op_kernel/arch35/flash_attn_block_cube_nd.h(L102-104 L1_KV_LARGE_BUF;L323-371 bmm1;L394-457 bmm2)
  20. attention/flash_attn/op_kernel/arch35/flash_attn_block_vec_dn.h(L67 dTemplateAlign64;L103 UB_VEC2;L160/271 dSizeV;L431-437 copy-out;L519-527 FlashUpdate)
  21. attention/flash_attn/op_kernel/arch35/flash_attn_block_vec_nd.h(L69/112/176/297/538/547-553/635-643)
  22. attention/flash_attn/op_kernel/arch35/flash_attn_block_vec_flashdecode.h(L154 dSizeV_Align_;L250-259 CopyAccumOutIn;L313-335 FD 归约)
  23. attention/flash_attn/op_kernel/arch35/flash_attn_kernel_dn.h / _nd.h(L168-169 dSize/dSizeV 直读;L204 dBasicBlock)
  24. attention/common/op_kernel/matmul.h(L811+ MatmulK 切 K;L872-882 kLoops/tailK 拆分公式)
  25. attention/common/op_kernel/memcopy/attn_copy_gm_to_l1.h(L26-72 Nd2Nz)
  26. attention/flash_attn_metadata/op_host/op_api/aclnn_flash_attn_metadata.h(L20-28 签名)
  27. attention/flash_attn_metadata/op_host/op_api/l0_flash_attn_metadata.cpp(L43-51 attr 名单)
  28. attention/flash_attn_metadata/op_host/flash_attn_metadata_check.h(L37 NONE_VALUE;L99-104 -1 哨兵;L112-114 headDim 白名单)
  29. attention/flash_attn_metadata/op_kernel_aicpu/flash_attn_metadata_aicpu.cpp(L59-74 attr 读取;L201-203 AdjustSinnerAndSouter 传 headDim_;L214-239 InitBaseInfo;L221-222 headDimQk/headDimV 同值)
  30. attention/flash_attn_metadata/op_kernel_aicpu/flash_attn_metadata_aicpu.json(AICPU 注册)
  31. attention/common/op_kernel/load_balance/base_info.h(L211-212 headDimQk/headDimV 字段)
  32. attention/common/op_kernel/load_balance/section_stream_k/section_stream_k_impl.h(L327-335 cost 模型)
  33. attention/fused_infer_attention_score/op_host/fused_infer_attention_score_tiling_info_parser.cpp(L1223-1238 GetRopeMode 硬编码 192/128;L1265-1279 GetMlaMode;L1806 isQKVDDifferent 豁免)
  34. attention/fused_infer_attention_score/op_host/checkers/common_checker.cpp(L530-560/606-617 isQKVDDifferent 校验)
  35. attention/fused_infer_attention_score/op_host/arch35/fia_tiling_nonquant_gqa.cpp(L55-92 config 表含 (192,128) 等;L588-602 UpdateTilingKeyConfig)
  36. attention/fused_infer_attention_score/op_host/arch35/fia_tiling_nonquant_mla.cpp(L40-115 IsCapable;L451-455 576/512 config)
  37. attention/fused_infer_attention_score/op_kernel/arch35/fused_infer_attention_score_template_tiling_key.h(L40-61 Config 10-bit 表)
  38. attention/quant_flash_attn/op_host/quant_flash_attn_def.cpp(L115-118 -1 默认 attr 先例)
  39. attention/flash_attn/torch_extension/flash_attn.py(L58-74 schema;L105-169 fake;L177-233/279-330 dispatch)
  40. attention/flash_attn/torch_extension/csrc/flash_attn.cpp(L25-35 metadata;L37-98 FlashAttn;L53-67 dSize 来源;L79-85 out 尺寸)
  41. attention/flash_attn/torch_extension/graph_convert_flash_attn.py(L49-101 GE 占位转换器)
  42. attention/flash_attn/tests/pytests/core/case_loader.py(L55 DV 默认 D)、core/data.py(L121 metadata head_dim;L262 v shape DV)、core/backends/cpu_golden.py(L113 d_v)、core/backends/npu.py(L35-74 GraphNetwork;L205-276 graph 调用)
  43. attention/flash_attn/tests/assets/impl/graph.py(L47 head_dim 取 q 维)
  44. attention/flash_attn/tests/pytests/test_cases/functional_nc.py(L23 起 D/DV 键)、tests/pytests/readme.md(L282 DV 文档)
  45. attention/flash_attn/tests/ut/op_host/arch35/test_flash_attn_tiling.csv(表头含 v_shape/deterministic)、tests/ut/framework_normal/common/tiling_context_faker.cpp(L70-72 deterministic 非 attr)
  46. attention/flash_attn/README.md(L348-352 config 表;L363 D→config;L429-433 tiling key;L444-452 Dn/Nd 路由)
  47. /workspace/ops-transformer-testkit/golden/flash_attn_golden.py(L130-147 v_dim)、backends/flash_attn.py(L90/131/155 D 取值)
  48. /workspace/.opencode/skills/compile-ops/SKILL.md(L209-215 tiling key 编码规则;-j16 约束;--tiling_key 精编)
  49. /workspace/my-skills/docs/26_flash_attn_head_dim_combinations_research.md(§1.4/§6 FIA MLA 结论——已回源码核实)、16/17 号 plan(UB 布局与性能结论)

修订记录

  • 2026-09-03: 初始创建。基于 ops-transformer-fa-qk192v128(master@570ff443b)代码调研完成。核心结论:flash_attn(FlashAttn,非 FIA)与其 metadata 算子(FlashAttnMetadata,AICPU)新增 head_dim_v(Int,默认 -1,attr 索引 10);-1 时保留 qk==v 强校验、config 0-5 与 tiling key 逐 bit 不变(零行为变化红线成立,唯一代价是变体 192→256 的编译时长 +1/3);QK=192/V=128 走新增 config 6/7((128,128,192,128)/(64,256,192,128)),Config 3-bit→4-bit;kernel 侧 FAType 的 d/dV 模板、cube bmm1 的 MatmulK 切 K(192=128+64,D=256 同机制)、vector 全链路按 dV 维——block_cube/block_vector/tilingData 三层零改动;最大风险是 metadata 的 AdjustSinnerAndSouter 传参(现传 qk 维需改 v 维)、infershape PA 分支 headDimV 回退、csrc out 尺寸用 qk 维三处隐蔽错误,须与 checker 同 PR 修复。方案分层清单共 16 项、约 5 人日。
  • 2026-09-03 (v2): 按设计评审修订——主算子 flash_attn 不加参数(vHeadDim 已从 value.shape 末维解析,flash_attn_tiling_info_parser.cpp:451-460 核实;qk≠v 入口即 v 的 shape),checker 由相等性校验改为组合白名单 {(64,64),(128,128),(256,256),(192,128)}(拒绝集合除新增 (192,128) 外与现状一致);head_dim_v 参数仅保留在 metadata 算子(默认 -1,NONE_VALUE 先例)。免除项:主算子 aclnn 签名×4/def/parser attr/attr 索引 10/UT csv 新列/GraphNetwork 签名(v1 的主算子防呆与 graph 位置传参错位风险随之消失);新增 R5'(metadata 无张量可交叉校验 head_dim_v 与真实 v.shape 的一致性,靠文档契约);R2 细节修正(qk=192 不落入 fa_adjust_sinner_souter.h:65/79 任何分档,停在预置默认 64/128,与主算子 ≤128 档 32/256 子条件触发时切分不一致)。分层清单 16→15 项、工作量 5→4 人日。
  • 2026-09-04 (v3 实施记录): v2 方案已全部实施并提交(ops-transformer 分支 flash_attn_qk192_v128:ffd0db552 主算子 checker/infershape/UpdateTilingKeyConfig/tilingKey 4-bit/ConfigValue/EnableSoftmaxDn;5c1be9869 metadata aclnn/l0/checker/aicpu head_dim_v 含 R2 修复;7d950ca3c torch_extension csrc v 维 out 修复+python 透传;2b2521dff 仓内 pytest/UT csv;6e55bdab2 README/接口文档;testkit 同名分支 9d56b05 backends head_dim_v 透传 + qk192_v128_cases.py 6 用例)。tiling key 编码实测更正:本仓构建链为裸位拼 key = layout | (kv<<8) | (mask<<16) | (config<<17),无 README L433 示例(0x87E40000)的高位前缀——config 2 旧 key 实测 262144,新 key:config 6 无 mask=0xC0000、config 6+mask=0xD0000、config 7=0xE0000;快速编译 8 key 全部注册,R8 红线(旧 key 逐 bit 不变)在编译产物层验证成立;UT csv 新增行 expectTilingKey 用 UINT64_MAX 跳过规避环境差异。
  • 2026-09-04 (v4 板上验证记录): 全量编译 256 变体(config 0-7 × 32,-j16 约 3.5 分钟;旧 key 262144 在位、config 6/7 各 32)+ torch_extension 构建安装完成。新增 8 例 192/128 板上精度全部 PASS(连续 Dn 0.4874%/Nd 0.0856%、config 7 decode 0.0000%、PA_BBND 0.3036%、PA_BNBD+FD 0.0000%、TND varlen 0.4877%、PA_NZ 0.0887%、GQA+window 0.2930%;LSE 全 PASS)。存量回归 18/24 PASS;5 个巨型用例(75348×10240 等)单进程 CPU golden OOM 跑不完(环境内存限制,与本改动无关)。发现并修复一处实施期新 bug:PA_NZ 的 v 布局实为 (Bn,N2,D/16,Bs,16)(权威来源 attention/common/op_host/fia_tiling_shape.cpp:44 PA_NZ={Bn,N,D1,Bs,D0}),csrc/infershape 初版误取 dim3(=block_size),被存量用例 000010(BS=128≠D=256)暴露(000018 因 BS=DV=128 巧合掩盖);修复 D=dim2×dim4 后 000004(BS=384≠D)由 tiling 失败转 PASS 0.0930%。000010 定性为存量问题:纯 master(570ff443b detached worktree + master whl)对照跑出完全相同数字(FailElems 297/30720=0.9668%,MaxAbsErr 0.037069,LSE PASS),本分支与其 bit 级一致、非回归(首次对照的 segfault 是新旧 whl 与 op 包 ABI 错配假象,装 master whl 后即为该精度 FAIL)。bf16 边缘用例调优结论:生成数据 seed=42 固定、结果确定性;±5 range 反而使输出均值抵消变小、相对误差放大(0.49%→1.8%),回退 ±10;000015 改 fp16。三仓状态:ops-transformer 分支 flash_attn_qk192_v128 已 push(7 commit,至 a0a4d3e1f);testkit 同名分支已 push(9d56b05);遗留:巨型存量用例的环境内验证、graph_mode 用例。
  • 2026-09-04 (v5 testkit 板上验证): ops-transformer-testkit-qk192v128(分支 flash_attn_qk192_v128)跑 qk192_v128_cases 6 用例 --precision-only --gen-device npu --one-by-one:6/6 PASS(golden tiled 模式;bf16 atol/rtol=0.78%、fp16 rtol=0.5%;精度 99.95%–100%,max_abs ≤0.0156)。该链路验证了用户侧完整调用:torch_extension python → flash_attn_metadata(head_dim_v=128) → AICPU metadata(含 R2 修复路径)→ flash_attn 主算子 config 6/7。至此 testkit 侧验证完成(v4 遗留仅剩巨型存量用例的环境内验证与 graph_mode 用例)。
  • 2026-09-04 (v6 rebase master): master 前移 570ff443b→33d8246e0(20+ 提交,含 flash_attn kernel 重构 9 文件、D=256 优化 README、paged_attention_checker 修改;metadata/torch_extension/tests 零改动)。分支 rebase 到新 master:6 个 commit 中 5 个代码 commit 全部干净自动合并(EnableSoftmaxDn +config6、UpdateTilingKeyConfig 192/128 分支、Config 4-bit、checker 组合白名单均保留),仅 README 文档冲突 4 处(与 877e0b5cf 的 D=256 优化说明交叠),按"master 新文案 + 我的 6/7 增量"合并。rebase 后全量重建(256 key,零错误)+ 板上重验 11 例:8/8 新用例 PASS 且数字与 rebase 前逐位一致(master kernel 重构与 config 6/7 共存无语义冲突),D=128/256 旧路径 PASS,000010 数字 bit 级不变(存量问题定性维持)。已 --force-with-lease 推送(a0a4d3e1f→a7f235c4f);testkit 分支已含最新 master 无需 rebase。注意:git rebase --continue 在本环境需 GIT_EDITOR=true 前缀(默认编辑器会挂起)。
  • 2026-09-04 (v7 二次 rebase master): master 再前移 33d8246e0→60239e6d8(12 提交,关键为 7c1e5fa50 op_api 重构:删除 flash_attn 的 aclnn_flash_attn.h/inner/flash_attn.{h,cpp} 共 5 文件并入 aclnn_flash_attn.cpp,metadata 的 aclnn_flash_attn_metadata.h 删除、声明以 ACLNN_API 宏内联进 .cpp)。与分支冲突仅 metadata 两文件:.h 为 modify/delete(按 master 意图删除,签名改动重放进 .cpp),.cpp 三处冲突按"master 的 ACLNN_API 内联结构 + headDimV 增量"合并;主算子 op_api 分支未改过(v2 零参数设计在此兑现价值)自动通过;csrc 无直接头依赖(ACLNN_CMD 宏)、examples 走 aclnnop/ 安装路径不受影响。重建 256 key 零错误,板上重验 9/9 PASS 且数字与前两次验证逐位一致(master 重构与 config 6/7 语义正交)。已 --force-with-lease 推送(a7f235c4f→2b8bb838b)。
  • 2026-09-04 (v8 百例广义验证): 扩展 testkit 生成器 gen_flash_cases.py 支持 --dv((D,DV) 组合白名单与算子 checker 一致、est_mem 按 DV 估算),生成 100 条 QK=192/V=128 广义用例(seed=20260904,覆盖 7 序列场景 × 12 布局组合 × mask 0/3/4 × 4 数据分布,TND varlen/截断/PA decode/边界长度全含)。板上精度 --precision-only:98/100 PASS;2 例 FAIL(S4_exact_TND_TND_m4_3943、S1_varlen_TND_TND_m4_49542)均为 fp16 + TND + mask4 + 小值域(±0.1/±1.0)模式,特征:LSE 100% 全对、max_abs≤0.002、99.3% 元素通过——D=DV=128 同形状同 seed 对照组同样 FAIL 且数字一致(99.32/99.33 vs 99.28/99.34,max_abs 相同),定性为 fp16 小值相对误差的固有边缘,非 192/128 路径回归。testkit 已 push(343d9d6)。
  • 2026-09-07 (v9 三次 rebase master + testkit rebase): master 60239e6d8→119ea5f1c(hd72 分支落地:checker 加 72、config 2/3 复用、PA_NZ 16 对齐检查;sinks kernel 优化 9 文件;测试框架 cpu_golden/npu/data 改动)。ops 仓 3 处冲突:common_checker.cpp(qk 白名单并集 {64,72,128,192,256},v 白名单 {64,72,128,256},组合表加 (72,72))、metadata_check.h(同并集 + headDimV 白名单 + 组合表)、README(72 复用说明 + config6/7 合并)。testkit 仓 bd02f14→ab1c0df(sinks 合入、生成器 29 行改动)干净落地,106 用例代码生成验证通过。环境事故与排查:第三次 rebase 后板上 9/9 全挂,报 expected at most 16 arguments but received 17——根因是 9月6日 00:23 外部会话向 site-packages 装了 master 版包(无 head_dim_v),重装本 worktree whl 后 9/9 恢复 PASS 且数字与前几轮逐位一致。D=72 兼容性:跑 master 新合入的 functional_hd72 用例 2/3 PASS、HD72_DN_BNSD_LSE(bf16) FAIL 2.0257%——纯 master(119ea5f1c detached + master whl) 对照跑出完全相同数字(5974/294912=2.0257%,MaxAbsErr 0.020071),定性 master 自身 hd72 精度问题(本分支对 config 0-5 路径 diff 为 nil:EnableSoftmaxDn 仅加 config==6、ConfigValue 仅追加索引 6/7、tiling key 对 0-5 编码不变)。ops 已 push(2b8bb838b→bf096cc6c)、testkit 已 push(343d9d6→9e499db)。
  • 2026-09-07 (v10 csv 行尾噪声修正): 两个 UT csv 因 python 重写引入行尾噪声——test_flash_attn_tiling.csv 用 csv.writer 追加时默认 lineterminator='\r\n' 把基线 LF(339 行)整体转 CRLF(diff 342/339,真实改动仅 3 行);test_flash_attn_shape_infershape.csv 基线本为 CRLF,PA_NZ 行修正时 python 文本模式读写整体转 LF(diff 14/10,真实改动 4 行)。均已按各自基线风格恢复(前者转 LF、后者转 CRLF),相对 master 净差异恢复为 3+4 行,已 push(1b65dc318)。教训:改 csv 追加行用行级 append(echo >> 或 binary 模式),避免 csv.writer/文本模式全量重写;提交前用 git diff master --stat 抽查噪声行数。
  • 2026-09-07 (v11 pre-commit 格式补齐): 用户侧 pre-commit clang-format 失败——根因是三次 rebase 手工解冲突时 aclnn_flash_attn_metadata.cpp 的 DFX 宏/ParamsCheck 折行风格与 clang-format v18.1.8 不一致(rebase --continue 不跑 pre-commit 钩子,冲突解法未过格式门)。注意两点:① hook 日志的 "Formatting [N/M]" 是 --verbose 对所有处理文件的打印,实际修改的只有 1 个文件(14 个中 13 个本就干净);② 本地裸 clang-format 是 v21,与 pre-commit 锁定的 v18.1.8 输出可能不同,必须用 pre-commit run clang-format --files ... 复现。修复:提交格式对齐(231c55db8,纯换行调整无语义变化),全量 pre-commit 13 钩子 Passed,已 push。教训:rebase 解冲突后跑一遍 pre-commit run --files $(git diff master --name-only) 再 push。
  • 2026-09-07 (v12 code review 修复): 双轴 review(Standards/Spec 并行 sub-agent)发现并修复 Spec 轴 c-1——graph.py/npu_preprocess.py 对 PA_NZ 的 head_dim_v 仍取 v.shape[3](=block_size),与 csrc/infershape 已按 v4 修复的 D=dim2×dim4 漂移,注释亦矛盾("PA_NZ 的 D 均在 index 3");当前用例因 block_size=DV=128 巧合掩盖(v4 同款掩盖模式,5 处推导副本只修了 C++ 侧 2 处)。修复:两处 python 改 dim2×dim4 + 新增回归用例 000020(PA_NZ + block_size=256≠DV=128,旧代码会被组合白名单 (192,256) 拒绝、新代码 PASS 0.0000%)+ 000018 结果不变。已 push(345abe66e),pre-commit 全绿。review 其余发现(未修,供后续):Standards 轴——组合白名单三处手工同步(算子内 checker/tiling 可收敛)、head_dim 布局推导 5 处副本(graph.py 隔离是刻意的,错误注释已随本次修复);Spec 轴——负例用例缺失((256,128)/(192,192) 拒绝)、metadata -1/显式等价性用例缺失、UT expectTilingKey 用 UINT64_MAX 跳过(v3 已记录原因)、§10-#3 graph_convert 条目前提与代码不符(应回写勘误)。
  • 2026-09-07 (v13 squash): 按 PR 规范将 10 个 commit 合并为单 commit 56b8b72e8(git reset --soft master + 综合提交信息),净 diff 与合并前逐字节一致(25 files, 390+/74-),pre-commit 全绿,已 force-with-lease 推送。
  • 2026-09-07 (v14 review P-2/P-3/P-4/P-5 实施): P-2 tiling 兜底拒绝——UpdateTilingKeyConfig 非法组合分支由只打日志改为返回 GRAPH_FAILED 沿调用链传播(防御纵深,UT faker 直接调 tiling 时原为静默接受 bug;正常路径零影响),UT csv 加 Invalid_D256_DV128/Invalid_D192_DV192 两行 FAIL 断言。P-3 新增 tools/metadata_equiv_check.py:head_dim_v 缺省 vs 显式等价性(64/72/128/256 逐元素一致)+ 显式 (192,128) + 三类负例拒绝,板上 ALL PASS。P-4 UT expectTilingKey 填真实 key(786432/851968/917504,UT 环境确认裸位拼编码与构建产物一致);顺带修正 MaskCausal 行 layout 列错位(BNSD→BSND,曾致 tiling 失败)、S1_1 行补 max_seqlen_q=1/kv=2048(缺省 -1 不触发 32/256 分档、config 7 不可达)。P-5 §10-#3 勘误(graph_convert 无 metadata 函数,零改动自洽)。UT 360 例全绿、板上 sanity 3/3 PASS、pre-commit 全绿,已 push(592edb06c)。v12 review 遗留项至此全部闭环。
  • 2026-09-07 (v15): 按用户决定删除 tools/metadata_equiv_check.py(独立脚本不进自动化、维护价值有限);等价性/负例验证结果保留在本修订记录(v14:板上 ALL PASS),组合白名单拒绝的持续看护由 UT csv 负例行承担。已 push(87baed92a)。
  • 2026-09-07 (v16): 按用户设计裁定撤销 P-2 的 tiling 兜底校验——(qk,v) 组合合法性由 checker 单层保证,tiling 只做映射:UpdateTilingKeyConfig/UpdateTilingKeyInfo/GenTilingKey 返回类型回退 void,else 拒绝分支(日志+return)整体移除;依赖 tiling 兜底的 UT 负例行(Invalid_D256_DV128/Invalid_D192_DV192)同步删除。勘误 v14 的"静默接受 bug"定性——分层视角下这不是 bug,tiling 收到非法输入本身就是 checker 失效的异常场景,防御纵深反而引入了职责重复(review S-1 组合白名单三处同步的根源)。注意:期间一次 UT 跑出 360 假绿(负例行 OK)系旧二进制未重编所致,重编后 358 真绿。板上 sanity 2/2 PASS、256 key 不变,已 push(a9ae4fdff)。
likedislike
Jjiang-lirui成员
26 天前 关联了pull request:docs: flash_attn qk!=v (192/128) 文档更新
游震成员
26 天前 评论:

/assign @jiang-lirui

likedislike
CANN-robotCANN-robot成员
26 天前 将 jiang-lirui 设为负责人
CANN-robotCANN-robot成员
26 天前 关闭了 issue
CANN-robotCANN-robot成员
26 天前 添加了label:resolved