指南)
CANN ops-transformer QuantLightningIndexer 算子深度解析稀疏 Attention 前處理與量化索引技術實戰(zhàn)指南【免費下載鏈接】ops-transformer本項目是CANN提供的transformer類大模型算子庫實現(xiàn)網(wǎng)絡在NPU上加速計算。項目地址: https://gitcode.com/cann/ops-transformer本文以 CANN ops-transformer 開源倉庫中 experimental/attention/quant_lightning_indexer/README.md 為骨架結合算子 Host 側(cè)定義、Tiling、InferShape 與 Kernel 側(cè)源碼系統(tǒng)講解 QuantLightningIndexer下文簡稱 QLI算子的功能原理、計算公式、參數(shù)語義、硬件適配約束以及單算子與 aclgraph 兩種調(diào)用方式。讀者讀完后將掌握如何在 Atlas A3 與 Ascend 950 系列產(chǎn)品上正確配置并調(diào)用該算子理解其稀疏選擇 存8算8的量化索引機制并能結合倉庫源碼定位實現(xiàn)細節(jié)。一、算子定位與功能概述QuantLightningIndexer 是 CANN ops-transformer 倉庫中面向推理場景的稀疏 Attention 前處理算子。它的核心職責有兩項稀疏 Token 選擇在長序列推理中從大量前序 Key 中高效選出與當前 Query 相關性最強的 Top-k 個稀疏 Token作為后續(xù)稀疏 Attention 計算的輸入索引量化索引存8算8輸入 Query 與 Key 以 INT8 或 FLOAT8 等低精度格式存儲存8在計算相關性分數(shù)時通過反量化系數(shù)還原精度算8在保證精度的前提下降低存儲與帶寬開銷獲取最大收益。從代碼定位看該算子同時存在于 experimental/attention/quant_lightning_indexer實驗分支本指南關聯(lián)文檔所在位置與 attention/quant_lightning_indexer正式目錄含 aclnn 接口與 examples 示例兩個目錄中二者共享相同的算子原型、Tiling 與 Kernel 設計后者額外提供了 aclnnQuantLightningIndexer.md 接口文檔與 test_npu_quant_lightning_indexer.py 示例腳本。在 算子原型定義 中該算子注冊的 AICore 配置覆蓋ascend910b、ascend910_93與ascend950三類 SoC并聲明了動態(tài)編譯、動態(tài) Format、動態(tài) Rank、動態(tài) Shape 支持以及PrecisionReduceFlag(true)精度降低標志說明其本身即是為低精度量化計算而設計。二、計算公式與算法流程QLI 算子的計算公式如下$$out \text{Top-}k\left{[1]{1\times g}\left[(W[1]{1\times S_{k}})\odot\text{ReLU}\left(\left(Scale_QScale_K^T\right)\odot\left(Q_{index}^{Quant}{\left(K_{index}^{Quant}\right)}^T\right)\right)\right]\right}$$其中各符號含義為符號對應入?yún)⒕S度含義$Q_{index}^{Quant}$query$\mathbb{R}^{g\times d}$當前 Token 的量化 Index Query$K_{index}^{Quant}$key$\mathbb{R}^{S_{k}\times d}$壓縮后的量化 Index Key$Scale_Q$query_dequant_scale-Index Query 的反量化系數(shù)$Scale_K^T$key_dequant_scale-Index Key 的反量化系數(shù)$W$weights-權重系數(shù)$g$--Group 數(shù)N1/N2算法分為三個主要計算步驟相關性計算將當前 Token 對應的輸入 Query 乘以給定上下文 Key得到二者間的相關性即量化域內(nèi)的 $QK^T$ 矩陣乘反量化與激活過濾相關性結果與 Query、Key 對應的反量化系數(shù) $Scale_Q$、$Scale_K^T$ 相乘完成反量化再經(jīng)過ReLU激活函數(shù)過濾無效的負相關信號得到當前 Token 與所有前序 Token 的相關性分數(shù)向量加權選 Top-k將分數(shù)向量與權重系數(shù) $W$ 相乘后沿 g 方向即 Query 的 Head 方向選取前 Top-k 個索引值得到輸出sparse_indices作為后續(xù) Attention 的輸入。從 Kernel 實現(xiàn)看該流程在 quant_lightning_indexer.cpp 中通過QLIPreload模板類完成內(nèi)部由QLIMatmulCube 側(cè)矩陣乘與QLIVectorVector 側(cè)激活、加權與 Topk兩個服務模塊協(xié)同并按KERNEL_TYPE_MIX_AIC_1_2采用 AICCube AIVVector混合核架構執(zhí)行。三、參數(shù)說明QLI 算子的完整參數(shù)如下表所示。其中輸入/輸出/屬性列區(qū)分了該參數(shù)在算子圖中的角色數(shù)據(jù)格式列中 ND 表示按維度順序連續(xù)排布。參數(shù)名輸入/輸出/屬性描述數(shù)據(jù)類型數(shù)據(jù)格式query輸入公式中的 $Q_{index}^{Quant}\in\mathbb{R}^{g\times d}$表示輸入 Index Query不支持非連續(xù)。INT8、FLOAT8_e4m3fn、HIFLOAT8NDkey輸入公式中的 $K_{index}^{Quant}\in\mathbb{R}^{S_{k}\times d}$表示壓縮后的輸入 Index Key支持 0 軸非連續(xù)。INT8、FLOAT8_e4m3fn、HIFLOAT8NDweights輸入公式中的 $W$表示權重系數(shù)不支持非連續(xù)。FLOAT16、FLOAT32NDquery_dequant_scale輸入公式中的 $Scale_Q$表示 Index Query 的反量化系數(shù)不支持非連續(xù)。FLOAT16、FLOAT32NDkey_dequant_scale輸入公式中的 $Scale_K$表示 Index Key 的反量化系數(shù)不支持非連續(xù)。FLOAT16、FLOAT32NDactual_seq_lengths_query可選輸入表示不同 Batch 中query的有效 Token 數(shù)。INT32NDactual_seq_lengths_key可選輸入表示不同 Batch 中key的有效 Token 數(shù)。INT32NDblock_table可選輸入表示 PageAttention 中 KV 存儲使用的 block 映射表。INT32NDmetadata可選輸入QuantLightningIndexerMetadata 算子傳入的分核信息包含使用核數(shù)、分塊大小以及每個核處理數(shù)據(jù)的起始點等內(nèi)容shape 大小為 [1024]當前不支持傳空。INT32NDquery_quant_mode屬性用于標識輸入query的量化模式當前支持 Per-Token-Head 量化模式當前僅支持傳入 0。INT32-key_quant_mode屬性用于標識輸入key的量化模式當前支持 Per-Token-Head 量化模式當前僅支持傳入 0。INT32-layout_query可選屬性用于標識輸入query的數(shù)據(jù)排布格式當前支持 BSND、TND默認值 BSND。STRING-layout_key可選屬性用于標識輸入key的數(shù)據(jù)排布格式當前僅支持傳入 PA_BSND默認值 PA_BSND。STRING-sparse_count可選屬性代表 topK 階段需要保留的 block 數(shù)量支持 [1, 2048]默認值 2048。INT32-sparse_mode可選屬性表示 sparse 的模式支持 0/3。sparse_mode 為 0 時代表 defaultMask 模式為 3 時代表 rightDownCausal 模式的 mask對應以右頂點為劃分的下三角場景。INT32-pre_tokens可選屬性預留參數(shù)表示 attention 需要和前幾個 Token 計算關聯(lián)僅支持默認值 2^63-1。INT64-next_tokens可選屬性預留參數(shù)表示 attention 需要和后面幾個 Token 計算關聯(lián)僅支持默認值 2^63-1。INT64-cmp_ratio可選屬性用于稀疏計算表示 key 的壓縮倍數(shù)。Atlas A3 推理系列產(chǎn)品支持 1/2/4/8/16/32/64/128Ascend 950PR/Ascend 950DT 支持 1/4/128默認值 1。INT32-return_value可選屬性表示是否輸出sparse_values。True 表示輸出False 表示不輸出僅支持默認值 False。BOOL-sparse_indices輸出公式中的輸出 Out參與稀疏 Attention 計算的 Token 索引值。INT32NDsparse_values輸出公式中 Indices 輸出對應的 value 值目前暫不支持返回 sparse_values。FLOAT32ND不同硬件平臺的數(shù)據(jù)類型差異Ascend 950PR/Ascend 950DTquery、key不支持 INT8weights、query_dequant_scale和key_dequant_scale不支持 FLOAT16。對應源碼中ascend950配置僅注冊DT_FLOAT8_E4M3FN、DT_HIFLOAT8輸入與DT_FLOAT權重/反量化系數(shù)。Atlas A3 訓練系列產(chǎn)品/Atlas A3 推理系列產(chǎn)品、Atlas A2 訓練系列產(chǎn)品/Atlas A2 推理系列產(chǎn)品query、key不支持 FLOAT8_e4m3fn 和 HIFLOAT8weights、query_dequant_scale和key_dequant_scale不支持 FLOAT32。該平臺差異在 Tiling 源碼 的GetAndCheckInOutDataType()中有嚴格校驗ASCEND910B/910_93 分支強制 query/key 為 INT8、scale 為 FLOAT16、weights 為 FLOAT16ASCEND950 分支強制 query/key 為 FLOAT8_E4M3FN 或 HIFLOAT8、scale 與 weights 為 FLOAT。由源碼確認的參數(shù)校驗規(guī)則從 quant_lightning_indexer_tiling.cpp 的CheckAttrParaInfo()與 shape 校驗邏輯可以進一步確認以下規(guī)則sparse_count必須滿足0 sparse_count 2048常量SPARSE_LIMITASCEND910B/910_93 上cmp_ratio必須為 2 的冪且0 cmp_ratio 128ASCEND950 上僅允許 1、4、128layout_key僅支持PA_BSNDlayout_query僅支持BSND或TNDPA 場景外layout_query與layout_key必須一致sparse_mode僅支持 0 或 3pre_tokens/next_tokens僅支持9223372036854775807即 int64 最大值2^63-1query_quant_mode/key_quant_mode僅支持 0return_values僅支持 Falsequery與key數(shù)據(jù)類型必須一致兩個 dequant_scale 數(shù)據(jù)類型必須一致query的 N 軸N1與key的 N 軸N2之比 g 必須等于 64即 N164、N21常量G_SIZE_LIMITD 軸head_dim必須等于 128常量HEAD_DIM_LIMITblock_size必須是 16 的整數(shù)倍且屬于 (0, 1024]常量BLOCK_SIZE_FACTOR/BLOCK_SIZE_LIMITmetadata的 shape 必須為 [1024]常量METADATA_LIMIT。四、約束說明該接口支持圖模式。該接口要求 $W \odot Scale_Q$ 的結果在float16Atlas A3/float32Ascend 950PR/Ascend 950DT的表示范圍內(nèi)否則計算精度無法保證。該接口的 TopK 過程對 NAN 排序是未定義行為輸入數(shù)據(jù)中不應包含 NAN。參數(shù)query中的 D 軸和參數(shù)key中的 D 軸值相等為 128。參數(shù)query和key中的 N 軸分別僅支持 64 和 1。當layout_query為 TND 時actual_seq_lengths_query必須傳入且以該入?yún)⒃氐臄?shù)量作為 B 值該入?yún)⒅忻總€元素的值表示當前 batch 與之前所有 batch 的 Token 數(shù)總和即前綴和因此后一個元素的值必須大于等于前一個元素的值不能出現(xiàn)負值。當layout_key為 PA_BSND 時actual_seq_lengths_key該入?yún)⒈仨殏魅?。注意其語義與actual_seq_lengths_query不同在 PageAttention 場景下它表示每個 batch 的實際 Token 數(shù)非前綴和這在測試參數(shù)文件 test_quant_lightning_indexer_paramset.py 的注釋PA場景非前綴和表示每個batch_size的實際token數(shù)中也有明確標注。PageAttention 場景下block_table必須為二維第一維長度需要等于 B第二維長度不能小于maxBlockNumPerSeq每個 batch 中最大actual_seq_lengths_key對應的 block 數(shù)量支持 block_size 取值為 16 的整數(shù)倍最大支持到 1024。query、key、weights、query_dequant_scale、key_dequant_scale數(shù)據(jù)排布格式支持從多種維度解讀其中BBatch Size表示輸入樣本批量大小、SSequence Length表示輸入樣本序列長度、HHead Size表示 hidden 層的大小、NHead Num表示多頭數(shù)、DHead Dim表示 hidden 層最小的單元尺寸且滿足 DH/NT 表示所有 Batch 輸入樣本序列長度的累加和。Layout 與 Shape 的對應關系從 InferShape 實現(xiàn) 與 Tiling 的 shape 校驗ValidateInputShapesMatch()中的注釋可以得出各 Layout 下輸入輸出的 shape 約定場景querykeyweightsblock_tableoutBSND[B, S1, N1, D][BlockNum, BlockSize, N2, D][B, S1, N1][B, MaxBlockNumPerBatch][B, S1, N2, TopK]TND[T, N1, D][BlockNum, BlockSize, N2, D][T, N1][B, MaxBlockNumPerBatch][T, N2, TopK]其中S2 block_table.dim1 * block_sizePageAttention 場景即參與比較的 Key 序列長度由 block 表換算而來。sparse_values輸出在return_valueFalse時 shape 為 [0]空若開啟返回則與sparse_indices形狀一致。五、Atlas A3 推理系列產(chǎn)品調(diào)用說明5.1 單算子模式調(diào)用單算子直調(diào)模式通過torch.ops.custom.npu_quant_lightning_indexer直接執(zhí)行算子。完整可運行示例import torch import torch_npu import numpy as np import torch.nn as nn import math import custom_ops n1 64 n2 1 d 128 block_size 128 layout_key PA_BSND layout_query BSND query_quant_mode 0 key_quant_mode 0 np.random.seed(0) # ------------- b 24 t None s1 4 s2 512 act_seq_q None act_seq_k None sparse_mode 0 sparse_count 512 cmp_ratio 1 max_block_table_num (s2 block_size - 1) // block_size block_table torch.tensor([range(b * max_block_table_num)], dtype torch.int32).reshape(b, -1) key torch.tensor(np.random.uniform(-128, 127, (b * max_block_table_num, block_size, n2, d))).to(torch.int8) key_dequant_scale torch.tensor(np.random.uniform(0, 10, (b * max_block_table_num, block_size, n2))) key_dequant_scale key_dequant_scale.to(torch.float16) query torch.tensor(np.random.uniform(-128, 127, (b, s1, n1, d))).to(torch.int8) query_dequant_scale torch.tensor(np.random.uniform(0, 10, (b, s1, n1))).to(torch.float16) weights torch.tensor(np.random.uniform(0, 0.01, (b, s1, n1))).to(torch.float16) actual_seq_lengths_query torch.tensor(np.random.uniform(s1, s1, (b))).to(torch.int32) \ if act_seq_q is None else torch.tensor(act_seq_q).to(torch.int32) actual_seq_lengths_key torch.tensor(np.random.uniform(s2, s2, (b))).to(torch.int32) \ if act_seq_k is None else torch.tensor(act_seq_k).to(torch.int32) max_seqlen_q actual_seq_lengths_query.max().item() max_seqlen_k actual_seq_lengths_key.max().item() metadata torch.ops.custom.npu_quant_lightning_indexer_metadata ( actual_seq_lengths_query actual_seq_lengths_query.npu(), actual_seq_lengths_key actual_seq_lengths_key.npu(), num_heads_q n1, num_heads_k n2, head_dim d, query_quant_mode query_quant_mode, key_quant_mode key_quant_mode, batch_size b, max_seqlen_q max_seqlen_q, max_seqlen_k max_seqlen_k, layout_query layout_query, layout_key layout_key, sparse_count sparse_count, sparse_mode sparse_mode, pre_tokens (163)-1, next_tokens (163)-1, cmp_ratio cmp_ratio, device npu:0) sparse_indices, sparse_values torch.ops.custom.npu_quant_lightning_indexer(query.npu(), key.npu(), weights.npu(), query_dequant_scale.npu(), key_dequant_scale.npu(), actual_seq_lengths_queryactual_seq_lengths_query.npu(), actual_seq_lengths_keyactual_seq_lengths_key.npu(), block_tableblock_table.npu(), metadata metadata, query_quant_modequery_quant_mode, key_quant_modekey_quant_mode, layout_querylayout_query, layout_keylayout_key, sparse_countsparse_count, sparse_modesparse_mode, pre_tokens(163)-1, next_tokens(163)-1, cmp_ratiocmp_ratio)關鍵點說明必須先調(diào)用 Metadata 算子npu_quant_lightning_indexer_metadata生成 shape 為 [1024] 的分核信息使用核數(shù)、分塊大小、每核處理數(shù)據(jù)的起始點等該結果作為metadata入?yún)鹘o主算子且當前不支持傳空Tensor 需先.npu()上板所有輸入張量在傳入前均需移到 NPU 設備上量化模式與 Layout示例中query_quant_mode0、key_quant_mode0Per-Token-Headlayout_queryBSND、layout_keyPA_BSND返回值為二元組sparse_indices與sparse_values其中sparse_values目前固定不返回有效內(nèi)容。5.2 aclgraph 調(diào)用aclgraphtorchair 圖編譯模式通過torch.compile配合torchair.get_npu_backend將網(wǎng)絡整體編譯后執(zhí)行適合將 QLI 算子嵌入完整推理圖。完整可運行示例import torch import torch_npu import numpy as np import torch.nn as nn import math import torchair import custom_ops from torchair.configs.compiler_config import CompilerConfig n1 64 n2 1 d 128 block_size 128 layout_key PA_BSND layout_query BSND query_quant_mode 0 key_quant_mode 0 np.random.seed(0) # ------------- b 24 t None s1 4 s2 512 act_seq_q None act_seq_k None sparse_mode 3 sparse_count 512 pre_tokens(163)-1 next_tokens(163)-1 cmp_ratio 4 max_block_table_num (s2 block_size - 1) // block_size block_table torch.tensor([range(b * max_block_table_num)], dtype torch.int32).reshape(b, -1).npu() key torch.tensor(np.random.uniform(-128, 127, (b * max_block_table_num, block_size, n2, d))).to(torch.int8).npu() key_dequant_scale torch.tensor(np.random.uniform(0, 10, (b * max_block_table_num, block_size, n2))).npu() key_dequant_scale key_dequant_scale.to(torch.float16).npu() query torch.tensor(np.random.uniform(-128, 127, (b, s1, n1, d))).to(torch.int8).npu() query_dequant_scale torch.tensor(np.random.uniform(0, 10, (b, s1, n1))).to(torch.float16).npu() weights torch.tensor(np.random.uniform(0, 0.01, (b, s1, n1))).to(torch.float16).npu() actual_seq_lengths_query torch.tensor(np.random.uniform(s1, s1, (b))).to(torch.int32).npu() \ if act_seq_q is None else torch.tensor(act_seq_q).to(torch.int32).npu() actual_seq_lengths_key torch.tensor(np.random.uniform(s2, s2, (b))).to(torch.int32).npu() \ if act_seq_k is None else torch.tensor(act_seq_k).to(torch.int32).npu() max_seqlen_q actual_seq_lengths_query.max().item() max_seqlen_k actual_seq_lengths_key.max().item() class QLINetwork(nn.Module): def __init__(self): super(QLINetwork, self).__init__() def forward(self, query, key, weights, q_scale, k_scale, query_quant_mode, key_quant_mode, batch_size, num_heads_q, num_heads_k, head_dim, actual_seq_lengths_queryNone, actual_seq_lengths_keyNone, block_tableNone, layout_queryBSND, layout_keyBSND, sparse_count512, sparse_mode3, pre_tokens(163)-1, next_tokens(163)-1, cmp_ratiocmp_ratio, return_valueFalse): metadata torch.ops.custom.npu_quant_lightning_indexer_metadata( actual_seq_lengths_query actual_seq_lengths_query, actual_seq_lengths_key actual_seq_lengths_key, num_heads_q num_heads_q, num_heads_k num_heads_k, head_dim head_dim, query_quant_mode query_quant_mode, key_quant_mode key_quant_mode, batch_size batch_size, max_seqlen_q max_seqlen_q, max_seqlen_k max_seqlen_k, layout_query layout_query, layout_key layout_key, sparse_count sparse_count, sparse_mode sparse_mode, pre_tokens (163)-1, next_tokens (163)-1, cmp_ratio cmp_ratio, device npu:0) sparse_indices, sparse_values torch.ops.custom.npu_quant_lightning_indexer(query, key, weights, q_scale, k_scale, actual_seq_lengths_queryactual_seq_lengths_query, actual_seq_lengths_keyactual_seq_lengths_key, block_tableblock_table, metadatametadata, query_quant_modequery_quant_mode, key_quant_modekey_quant_mode, layout_querylayout_query, layout_keylayout_key, sparse_countsparse_count, sparse_modesparse_mode,pre_tokenspre_tokens, next_tokensnext_tokens, cmp_ratiocmp_ratio, return_valuereturn_value) return sparse_indices config CompilerConfig() config.mode reduce-overhead npu_backend torchair.get_npu_backend(compiler_configconfig) torch._dynamo.reset() npu_mode torch.compile(QLINetwork().npu(), fullgraphTrue, backendnpu_backend, dynamicFalse) sparse_indices npu_mode( query, key, weights, query_dequant_scale, key_dequant_scale, query_quant_mode, key_quant_mode, b, n1, n2, d, actual_seq_lengths_queryactual_seq_lengths_query, actual_seq_lengths_keyactual_seq_lengths_key, block_tableblock_table, layout_querylayout_query, layout_keylayout_key, sparse_countsparse_count, sparse_modesparse_mode, pre_tokenspre_tokens, next_tokensnext_tokens, cmp_ratiocmp_ratio, return_valueFalse)與單算子模式相比aclgraph 模式的關鍵差異需要額外導入torchair與CompilerConfig通過定義nn.Module網(wǎng)絡并配合torch.compile(..., fullgraphTrue, backendnpu_backend, dynamicFalse)將算子整體編譯為圖執(zhí)行CompilerConfig的mode設置為reduce-overhead以減少圖執(zhí)行開銷示例中使用了sparse_mode3rightDownCausal 下三角 mask與cmp_ratio4Key 壓縮 4 倍可作為稀疏 Attention 長序列場景的參考配置。六、Metadata 算子與分核機制metadata入?yún)⒂膳涮姿阕觧pu_quant_lightning_indexer_metadata生成是 QLI 算子正確執(zhí)行的前提。其底層布局定義位于 quant_lightning_indexer_metadata.hMetadata 總大小固定為1024 個 int32 元素QLI_META_SIZE 1024對應入?yún)?shape [1024]內(nèi)部包含兩類分核信息LI Metadata大小為 8共AIC_CORE_NUM 36組記錄每個 Cube 核的使能標志、BN2/M/S2 的起始與結束索引、首個 LD 數(shù)據(jù)的 workspace 索引等見常量LI_CORE_ENABLE_INDEX至LI_FIRST_LD_DATA_WORKSPACE_IDX_INDEXLD Metadata大小為 8共AIV_CORE_NUM 72組記錄每個 Vector 核的使能標志、BN2/M 索引、workspace 索引與數(shù)量、M 起始位置與數(shù)量等見常量LD_CORE_ENABLE_INDEX至LD_M_NUM_INDEX通過GetAttrAbsIndex(coreIdx, metaIdx, isAIV)計算屬性在 1024 元素中的絕對偏移非 AIV 數(shù)據(jù)位于LI_METADATA_SIZE * coreIdx metaIdxAIV 數(shù)據(jù)位于LI_METADATA_SIZE * AIC_CORE_NUM LD_METADATA_SIZE * coreIdx metaIdx。從 Tiling 實現(xiàn) 看Tiling 階段會通過CalcTschBlockDim(aivNum, aicNum, aivNum)計算 blockDim并按M_BASE_SIZE 512、S2_BASE_SIZE 512arch22或s1BaseSize 4、s2BaseSize 128arch35/DAV_3510等基本塊規(guī)格為各核分配 S2 循環(huán)與 DecodeLD中間結果的 workspace臨時存儲索引/值與參數(shù)信息最終以tilingData與tilingKey由數(shù)據(jù)類型、PA 標志、Q/K Layout 組合而成下發(fā)給 Kernel。七、Kernel 實現(xiàn)要點Kernel 入口 quant_lightning_indexer.cpp 通過編譯期宏區(qū)分兩套架構實現(xiàn)arch22ASCEND910B / ASCEND910_93即 Atlas A3模板參數(shù)為int8_t, int8_t, int32_t即 Query/Key 走 INT8、反量化輸出 INT32對應KERNEL_TYPE_MIX_AIC_1_2混合核arch35ASCEND950根據(jù)原始數(shù)據(jù)類型選擇hifloat8或fp8_e4m3fn_t模板輸出類型為float反量化系數(shù)為uint16_t編碼的 FLOAT8 相關格式對應 FLOAT8 存算方案。在 arch22 Kernel 頭文件 中可以看到關鍵設計常量M_BASE_SIZE 256、S2_BASE_SIZE 2048S2 循環(huán)的基本分塊大小HEAD_DIM 128、K_HEAD_NUM 1與約束Query N64、Key N1、D128對應MM1_OUT_T floatQK^T 矩陣乘的中間累加精度為 FP32WS_DOUBLE 2workspace 采用雙緩沖LD_PREFETCH_LEN 2Decode 預取深度為 2TempLoopInfo結構體記錄了 B/N/S1/S2 方向的循環(huán)邊界、實際有效長度actS1Size/actS2Size、是否需要 LDisNeedLD以及尾塊大小等運行時信息用于驅(qū)動主循環(huán)。Kernel 的 Process 流程按照Cube 計算相關性 → Vector 反量化/ReLU/加權 → TopK 選索引的流水線組織Cube 與 Vector 之間通過同步標志SYNC_C1_V1_FLAG 4、SYNC_V1_C1_FLAG 5實現(xiàn)跨核同步。八、測試框架與驗證方法關聯(lián)目錄提供了基于 pytest 的完整測試框架說明文檔見 tests/pytest/README.md。其驗證思路為CPU 側(cè)通過 quant_lightning_indexer_golden.py 復現(xiàn)算子功能生成 golden 數(shù)據(jù)NPU 側(cè)通過 TorchNPU 直調(diào)算子獲取實際輸出精度對比由 result_compare_method.py 完成 CPU 與 NPU 結果比對。使用方法單用例調(diào)測在experimental/attention/quant_lightning_indexer/tests/pytest目錄下手動配置 test_quant_lightning_indexer_paramset.py 中的參數(shù)然后執(zhí)行bash test_run.sh single批量生成與測試在 excel 路徑下放置用例表格在test_run.sh中設置 excel 路徑與 pt 文件存放路徑然后執(zhí)行bash test_run.sh batch在 test_quant_lightning_indexer_paramset.py 中可以看到覆蓋三類硬件場景的典型用例quant_li_default_a5qk_dtype torch.float8_e4m3fn、dequant_dtype torch.float32、layout_query BSND、layout_key PA_BSND對應 Ascend 950 場景quant_li_default_hifp8_a5qk_dtype torch.uint8HIFLOAT8 載體、dequant_dtype torch.float32同樣對應 Ascend 950 場景quant_li_default_a3qk_dtype torch.int8、dequant_dtype torch.float16、layout_query TND、layout_key PA_BSND對應 Atlas A3 場景。其中act_seq_k的取值如[28,24,80,96,47,76,0,111]與PA 場景非前綴和表示每個 batch_size 的實際 token 數(shù)的注釋相印證同時block_size512、block_num8也對應了 block 表維度的設計約定。正式目錄 attention/quant_lightning_indexer/tests 下還提供了對應的 test_quant_lightning_indexer_acl_graph.py 圖模式用例與 Host 側(cè) tiling/infershape 的 UT見 tests/ut可作為回歸驗證的參考。九、總結與適用場景QuantLightningIndexer 是 CANN ops-transformer 中面向推理場景的稀疏 Attention 前處理算子通過存8算8的量化索引與 Top-k 稀疏選擇將長序列 Attention 的計算范圍收斂到與當前 Query 最相關的 Key 子集從而降低存儲帶寬與計算開銷。在使用時需注意先元數(shù)據(jù)后主算子metadata必須由配套的npu_quant_lightning_indexer_metadata算子生成并傳入shape 固定為 [1024]按硬件選擇數(shù)據(jù)類型Atlas A3 使用 INT8 FLOAT16Ascend 950PR/950DT 使用 FLOAT8_e4m3fn/HIFLOAT8 FLOAT32嚴格遵守 shape 約束N164、N21、D128、block_size 為 16 的整數(shù)倍且不超過 1024區(qū)分兩種 seq 語義TND 下actual_seq_lengths_query為前綴和PA_BSND 下actual_seq_lengths_key為各 batch 實際 Token 數(shù)兩種調(diào)用方式單算子直調(diào)適合調(diào)試驗證aclgraph 圖編譯適合嵌入完整推理鏈路。如需深入源碼可繼續(xù)閱讀 算子原型定義、Tiling 實現(xiàn)、InferShape 以及 Kernel 入口 與兩套架構的 Kernel 頭文件arch22 / arch35?!久赓M下載鏈接】ops-transformer本項目是CANN提供的transformer類大模型算子庫實現(xiàn)網(wǎng)絡在NPU上加速計算。項目地址: https://gitcode.com/cann/ops-transformer創(chuàng)作聲明:本文部分內(nèi)容由AI輔助生成(AIGC),僅供參考