位置編碼與量化 KV Cache 寫入的實戰(zhàn)指南)
aclnnDequantRopeQuantKvcache 算子深度解析NPU 上融合反量化、RoPE 旋轉(zhuǎn)位置編碼與量化 KV Cache 寫入的實戰(zhàn)指南【免費下載鏈接】ops-transformer本項目是CANN提供的transformer類大模型算子庫實現(xiàn)網(wǎng)絡在NPU上加速計算。項目地址: https://gitcode.com/cann/ops-transformer導讀aclnnDequantRopeQuantKvcache是 CANN ops-transformer 算子庫中面向大模型推理場景的高階融合算子它將反量化Dequant→ QKV 切分 → 旋轉(zhuǎn)位置編碼RoPE→ 量化Quant→ KV Cache 寫入更新五個步驟融合為一次 NPU 算子調(diào)用避免了多個中間張量的 Device?Host 往返與多次 Kernel 啟動開銷。閱讀本文后你將掌握該算子的完整計算流程、兩段式 aclnn 接口的每個參數(shù)語義與取值約束、兩種 KV Cache 更新模式contiguous / page的差異并能夠依據(jù)倉庫中的完整調(diào)用示例與源碼實現(xiàn)在自己的推理管線中正確接入該算子。本算子位于倉庫的 posembedding/dequant_rope_quant_kvcache 目錄完整源碼可參考該目錄下的 op_host 與 op_kernel 實現(xiàn)。產(chǎn)品支持情況根據(jù)算子官方文檔與目錄內(nèi) README.md 說明當前支持情況如下產(chǎn)品是否支持Ascend 950PR / Ascend 950DT√Atlas A3 訓練系列產(chǎn)品 / Atlas A3 推理系列產(chǎn)品√Atlas A2 訓練系列產(chǎn)品 / Atlas A2 推理系列產(chǎn)品√Atlas 200I/500 A2 推理產(chǎn)品×Atlas 推理系列產(chǎn)品×Atlas 訓練系列產(chǎn)品×Kirin X90 處理器系列產(chǎn)品√Kirin 9030 處理器系列產(chǎn)品√需要特別說明兩點Kirin 平臺不支持 BFLOAT16README 明確標注因此在 Kirin X90 / Kirin 9030 上使用本算子時x 僅支持 FLOAT16 / INT32 兩種類型。這一差異在算子注冊源碼中也有體現(xiàn)dequant_rope_quant_kvcache_def.cpp 為 Kirin 平臺單獨定義了XDtypeListKirin/cosDtypeListKirin等數(shù)據(jù)類型列表其中不含DT_BF16且 AICore 配置kirinx90、kirin9030通過GetKirinCoreConfig()統(tǒng)一掛載。在 dequant_rope_quant_kvcache_def.cpp 中可以看到算子通過this-AICore().AddConfig(ascend910b)、AddConfig(ascend910_93)、AddConfig(ascend950)完成 AICore 配置注冊與上表支持的產(chǎn)品一一對應。算子功能與計算流程功能總覽算子對輸入張量x執(zhí)行如下流水線Dequant可選對輸入x進行反量化恢復高精度浮點表示切分Split按屬性sizeSplits給出的長度對尾軸dim-1進行切分得到q、k、vOut三段RoPE旋轉(zhuǎn)位置編碼對q、k應用基于cos、sin的旋轉(zhuǎn)位置編碼生成qOut和kOutQuant量化對kOut與vOut分別使用scaleK/offsetK、scaleV/offsetV進行靜態(tài)量化輸出 INT8 數(shù)據(jù)KV Cache 更新根據(jù)indices指定的 token 位置信息將量化后的 k、v 寫入kCacheRef與vCacheRef。從計算路徑看該算子本質(zhì)上是把 LLM 推理中權(quán)重反量化 QKV 投影結(jié)果切分 位置編碼 KV 量化緩存這段高頻熱點串成一條端到端的單算子流水適合與 PagedAttention 類推理框架配合使用。計算步驟公式算子的計算過程可形式化為以下五步第 1 步反量化可選$$ dequantX Dequant(x, weightScaleOptional, activationScaleOptional, biasOptional) $$第 2 步尾軸切分$$ q, k, vOut SplitTensor(dequantX, dim-1, sizeSplits) $$第 3 步旋轉(zhuǎn)位置編碼$$ qOut, kOut ApplyRotaryPosEmb(q, k, cos, sin) $$第 4 步靜態(tài)量化$$ quantK Quant(kOut, scaleK, offsetKOptional) $$$$ quantV Quant(vOut, scaleV, offsetVOptional) $$第 5 步KV Cache 更新兩種模式見下文兩種 KV Cache 更新模式模式一cacheModeOptional contiguous默認連續(xù)式緩存按 batch 維度逐位置寫入$$ kCacheRef[i][indices[i]] quantK[i] $$$$ vCacheRef[i][indices[i]] quantV[i] $$模式二cacheModeOptional page分頁式緩存Paged KV Cache先將 4 維 cache 張量展平成[總頁數(shù), 頁內(nèi)行數(shù), 列數(shù)]的視圖再按下標寫入$$ kCacheRefView kCacheRef.view(-1, kCacheRef[-2], kCacheRef[-1]) $$$$ vCacheRefView vCacheRef.view(-1, vCacheRef[-2], vCacheRef[-1]) $$$$ kCacheRefView[indices[i]] quantK[i] $$$$ vCacheRefView[indices[i]] quantV[i] $$兩種模式的選擇會直接影響indices的 shape 語義見參數(shù)表與約束說明并且 tiling 階段會通過CheckPaCacheMode()見 dequant_rope_quant_kvcache_tiling.cpp識別是否為 page 模式進而把batch展開為B*S、seqlen置 1 參與任務劃分。兩段式接口與函數(shù)原型該算子采用 CANN 算子庫標準的兩段式接口詳見倉庫文檔 兩段式接口說明必須先調(diào)用aclnnDequantRopeQuantKvcacheGetWorkspaceSize接口獲取入?yún)⑿r灲Y(jié)果與所需 workspace 大小再調(diào)用aclnnDequantRopeQuantKvcache接口執(zhí)行計算。第一段接口原型aclnnStatus aclnnDequantRopeQuantKvcacheGetWorkspaceSize( const aclTensor *x, const aclTensor *cos, const aclTensor *sin, aclTensor *kCacheRef, aclTensor *vCacheRef, const aclTensor *indices, const aclTensor *scaleK, const aclTensor *scaleV, const aclTensor *offsetKOptional, const aclTensor *offsetVOptional, const aclTensor *weightScaleOptional, const aclTensor *activationScaleOptional, const aclTensor *biasOptional, const aclIntArray *sizeSplits, char *quantModeOptional, char *layoutOptional, bool kvOutput, char *cacheModeOptional, const aclTensor *qOut, const aclTensor *kOut, const aclTensor *vOut, uint64_t *workspaceSize, aclOpExecutor **executor)第二段接口原型aclnnStatus aclnnDequantRopeQuantKvcache( void* workspace, uint64_t workspaceSize, aclOpExecutor* executor, aclrtStream stream)第一段接口完成入?yún)⑿r炁c workspace 大小計算第二段接口在指定stream上真正下發(fā)執(zhí)行。兩段接口均返回aclnnStatus狀態(tài)碼具體取值參見倉庫文檔 aclnn 返回碼說明。參數(shù)詳解第一段接口參數(shù)表下表完整列出aclnnDequantRopeQuantKvcacheGetWorkspaceSize的全部參數(shù)語義。其中非連續(xù) Tensor標記為 √ 表示該輸入支持非連續(xù)帶 stride張量。參數(shù)名輸入/輸出描述使用說明數(shù)據(jù)類型數(shù)據(jù)格式維度(shape)非連續(xù)Tensorx輸入公式中用于切分的輸入 xshape 為[B, S, H]或[B, H]H(NqNkvNkv)*D。x 的尾軸小于等于 4096且按 64 對齊FLOAT16、BFLOAT16、INT32ND2-3√cos輸入公式中用于位置編碼的輸入 cosx 為 3 維時 shape 為[B, S, 1, D]x 為 2 維時 shape 為[B, D]FLOAT16、BFLOAT16ND24√sin輸入公式中用于位置編碼的輸入 sinx 為 3 維時 shape 為[B, S, 1, D]x 為 2 維時 shape 為[B, D]和 cos 保持一致ND24√kCacheRef輸入公式中用于緩存 k 的輸入 kCacheRefshape 為[C_1, C_2, Nkv, D]INT8ND4√vCacheRef輸入公式中用于緩存 v 的輸入 vCacheRefshape 為[C_1, C_2, Nkv, D]INT8ND4√indices輸入表示 Kvcache 的 token 位置信息的輸入 indices當 cache_mode 為 page 且 x 為 3 維時 shape 為[B*S]否則 shape 為[B]INT32ND1√scaleK輸入公式中的輸入 scaleK用于量化 k 的 scale 因子元素個數(shù)為Nkv*D推薦 shape 為[Nkv, D]兼容一維展平 shape[Nkv*D]FLOATND≥1√scaleV輸入公式中的輸入 scaleV用于量化 v 的 scale 因子元素個數(shù)為Nkv*D推薦 shape 為[Nkv, D]兼容一維展平 shape[Nkv*D]FLOATND≥1√offsetKOptional輸入公式中的輸入 offsetKOptional用于量化 k 的 offset 因子元素個數(shù)為Nkv*D推薦 shape 為[Nkv, D]兼容一維展平 shape[Nkv*D]FLOATND≥1√offsetVOptional輸入公式中的輸入 offsetVOptional用于量化 v 的 offset 因子元素個數(shù)為Nkv*D推薦 shape 為[Nkv, D]兼容一維展平 shape[Nkv*D]FLOATND≥1√weightScaleOptional輸入公式中的輸入 weightScaleOptional用于反量化的權(quán)重 scale 因子shape 為[H]FLOATND1√activationScaleOptional輸入公式中的輸入 activationScaleOptional用于反量化的激活 scale 因子x 為 3 維時 shape 為[B*S]x 為 2 維時 shape 為[B]FLOATND1√biasOptional輸入公式中的輸入用于反量化的偏置 biasOptionalshape 為[H]FLOAT、FLOAT16、INT32、BFLOAT16ND1√sizeSplits輸入表示輸入的 qkv 進行切分的長度size 大小為 3值為[Nq*D, Nkv*D, Nkv*D]AclIntArray---quantModeOptional輸入表示支持的量化類型目前僅傳入staticCHAR---layoutOptional輸入表示支持的數(shù)據(jù)格式目前僅支持BSNDCHAR---kvOutput輸入Host 側(cè)布爾值表示是否輸出 kOut 和 vOut為 true 時輸出有效 shape 的 kOut 和 vOut為 false 時 kOut 和 vOut 的 shape 為空BOOL---cacheModeOptional輸入表示 kCacheRef 的更新方式目前僅支持page和contiguous默認為contiguousCHAR---qOut輸出公式中經(jīng)旋轉(zhuǎn)位置編碼后的 qx 為 3 維時 shape 為[B, S, Nq, D]x 為 2 維時 shape 為[B, Nq, D]。數(shù)據(jù)類型與 cos、sin 保持一致FLOAT16、BFLOAT16ND3-4×kOut輸出公式中經(jīng)旋轉(zhuǎn)位置編碼后的 kkvOutput 為 true 時x 為 3 維時 shape 為[B, S, Nkv, D]x 為 2 維時 shape 為[B, Nkv, D]kvOutput 為 false 時 shape 為空。數(shù)據(jù)類型與 cos、sin 保持一致FLOAT16、BFLOAT16ND1-4×vOut輸出公式中切分得到的 vkvOutput 為 true 時x 為 3 維時 shape 為[B, S, Nkv, D]x 為 2 維時 shape 為[B, Nkv, D]kvOutput 為 false 時 shape 為空。數(shù)據(jù)類型與 cos、sin 保持一致FLOAT16、BFLOAT16ND1-4×workspaceSize輸出返回需要在 Device 側(cè)申請的 workspace 大小-----executor輸出返回 op 執(zhí)行器包含了算子計算流程-----參數(shù)語義的源碼佐證上述參數(shù)定義可以在算子注冊源碼 dequant_rope_quant_kvcache_def.cpp 中得到印證x、cos、sin、k_cache、v_cache、indices、scale_k、scale_v為REQUIRED必選輸入offset_k、offset_v、weight_scale、activation_scale、bias為OPTIONAL可選輸入所有輸入均聲明了AutoContiguous()與參數(shù)表中非連續(xù) Tensor的支持情況一致屬性側(cè)size_splits為 REQUIRED 的ListIntquant_mode默認static、layout默認BSND、kv_output默認false、cache_mode默認contiguous與文檔描述完全對應輸入/輸出均限定FORMAT_ND格式即文檔參數(shù)表中的數(shù)據(jù)格式 ND。第二段接口參數(shù)表參數(shù)名輸入/輸出描述workspace輸入在 Device 側(cè)申請的 workspace 內(nèi)存地址workspaceSize輸入在 Device 側(cè)申請的 workspace 大小由第一段接口aclnnDequantRopeQuantKvcacheGetWorkspaceSize獲取executor輸入op 執(zhí)行器包含了算子計算流程stream輸入指定執(zhí)行任務的 Stream返回值與錯誤碼兩段接口均返回aclnnStatus狀態(tài)碼。第一段接口完成入?yún)⑿r灣霈F(xiàn)以下場景時報錯返回值錯誤碼描述ACLNN_ERR_PARAM_NULLPTR161001輸入和輸出的 Tensor 是空指針ACLNN_ERR_PARAM_INVALID161002輸入和輸出的數(shù)據(jù)類型不在支持的范圍內(nèi)完整返回碼語義參見倉庫文檔 aclnn 返回碼說明。約束說明使用該算子時必須滿足以下約束確定性計算aclnnDequantRopeQuantKvcache默認確定性實現(xiàn)。contiguous 模式下的 indices 取值kCacheRef的第 0 維大于等于 x 的第 0 維。x 為 3 維時indices 數(shù)據(jù)值大于等于 0 且小于等于kCacheRef的第 1 維減 x 的第 1 維x 為 2 維時indices 數(shù)據(jù)值大于等于 0 且小于等于kCacheRef的第 1 維減 1。page 模式下的 indices 取值indices 數(shù)據(jù)值大于等于 0小于kCacheRef的第 0 維 × 第 1 維且不重復。非 INT32 輸入輸入 x 不為 INT32 時x、cos、sin 與輸出 qOut、kOut、vOut 的數(shù)據(jù)類型保持一致此時activationScaleOptional、weightScaleOptional、biasOptional不生效。INT32 輸入反量化路徑輸入 x 為 INT32 時cos、sin 與輸出 qOut、kOut、vOut 的數(shù)據(jù)類型保持一致此時weightScaleOptional必選activationScaleOptional、biasOptional可選biasOptional不需要與其他輸入類型一致。尾軸限制x 的尾軸小于等于 4096且按 64 對齊。Kirin 平臺Kirin X90 / Kirin 9030 處理器系列產(chǎn)品不支持 BFLOAT16。這些約束在 tiling 階段會被進一步強制校驗。例如 dequant_rope_quant_kvcache_tiling.cpp 中會檢查sizeSplits長度必須為 3、k 與 v 的切分長度必須相等、hiddenSize必須為 16 的倍數(shù)、qHiddenSize/vHiddenSize必須為hiddenSize的整數(shù)倍由此推出 Nq、Nkv 必須為整數(shù)、x 的尾軸必須等于三段切分長度之和等同時 tiling 還會校驗 scale/offset 的元素個數(shù)必須等于Nkv*DquantShapeSize與參數(shù)表要求一致。調(diào)用示例倉庫在 examples/test_aclnn_dequant_rope_quant_kvcache.cpp 中提供了可直接參考的完整示例其調(diào)用流程與本算子文檔中的示例代碼一致編譯和執(zhí)行過程請參考倉庫文檔 編譯與運行樣例。核心代碼如下#include iostream #include vector #include acl/acl.h #include aclnnop/aclnn_dequant_rope_quant_kvcache.h #define CHECK_RET(cond, return_expr) \ do { \ if (!(cond)) { \ return_expr; \ } \ } while (0) #define LOG_PRINT(message, ...) \ do { \ printf(message, ##__VA_ARGS__); \ } while (0) int64_t GetShapeSize(const std::vectorint64_t shape) { int64_t shapeSize 1; for (auto i : shape) { shapeSize * i; } return shapeSize; } void PrintOutResult(std::vectorint64_t shape, void** deviceAddr) { auto size GetShapeSize(shape); std::vectorint8_t resultData(size, 0); auto ret aclrtMemcpy(resultData.data(), resultData.size() * sizeof(resultData[0]), *deviceAddr, size * sizeof(resultData[0]), ACL_MEMCPY_DEVICE_TO_HOST); CHECK_RET(ret ACL_SUCCESS, LOG_PRINT(copy result from device to host failed. ERROR: %d\n, ret); return); for (int64_t i 0; i size; i) { LOG_PRINT(mean result[%ld] is: %d\n, i, resultData[i]); } } int Init(int32_t deviceId, aclrtStream* stream) { // 固定寫法資源初始化 auto ret aclInit(nullptr); CHECK_RET(ret ACL_SUCCESS, LOG_PRINT(aclInit failed. ERROR: %d\n, ret); return ret); ret aclrtSetDevice(deviceId); CHECK_RET(ret ACL_SUCCESS, LOG_PRINT(aclrtSetDevice failed. ERROR: %d\n, ret); return ret); ret aclrtCreateStream(stream); CHECK_RET(ret ACL_SUCCESS, LOG_PRINT(aclrtCreateStream failed. ERROR: %d\n, ret); return ret); return 0; } template typename T int CreateAclTensor(const std::vectorT hostData, const std::vectorint64_t shape, void** deviceAddr, aclDataType dataType, aclTensor** tensor) { auto size GetShapeSize(shape) * sizeof(T); // 調(diào)用aclrtMalloc申請device側(cè)內(nèi)存 auto ret aclrtMalloc(deviceAddr, size, ACL_MEM_MALLOC_HUGE_FIRST); CHECK_RET(ret ACL_SUCCESS, LOG_PRINT(aclrtMalloc failed. ERROR: %d\n, ret); return ret); // 調(diào)用aclrtMemcpy將host側(cè)數(shù)據(jù)拷貝到device側(cè)內(nèi)存上 ret aclrtMemcpy(*deviceAddr, size, hostData.data(), size, ACL_MEMCPY_HOST_TO_DEVICE); CHECK_RET(ret ACL_SUCCESS, LOG_PRINT(aclrtMemcpy failed. ERROR: %d\n, ret); return ret); // 計算連續(xù)tensor的strides std::vectorint64_t strides(shape.size(), 1); for (int64_t i shape.size() - 2; i 0; i--) { strides[i] shape[i 1] * strides[i 1]; } // 調(diào)用aclCreateTensor接口創(chuàng)建aclTensor *tensor aclCreateTensor(shape.data(), shape.size(), dataType, strides.data(), 0, aclFormat::ACL_FORMAT_ND, shape.data(), shape.size(), *deviceAddr); return 0; } int main() { // 1. 固定寫法device/stream初始化參考acl API手冊 // 根據(jù)自己的實際device填寫deviceId int32_t deviceId 0; aclrtStream stream; auto ret Init(deviceId, stream); CHECK_RET(ret ACL_SUCCESS, LOG_PRINT(Init acl failed. ERROR: %d\n, ret); return ret); // 2. 構(gòu)造輸入與輸出需要根據(jù)API的接口定義構(gòu)造 int64_t shapeB 1; // batch int64_t shapeS 1; // seqlen int64_t shapeNq 2; // query 頭數(shù) int64_t shapeNkv 1; // kv 頭數(shù) int64_t shapeD 32; // 每頭維度 int64_t shapeH shapeD * (shapeNq shapeNkv shapeNkv); // 尾軸長度 (NqNkvNkv)*D std::vectorint64_t inputShape {shapeB, shapeS, shapeH}; std::vectorint64_t cosShape {shapeB, shapeS, 1, shapeD}; std::vectorint64_t sinShape {shapeB, shapeS, 1, shapeD}; std::vectorint64_t kcacheShape {shapeB, shapeH, 1, shapeD}; std::vectorint64_t vcacheShape {shapeB, shapeH, 1, shapeD}; std::vectorint64_t indicesShape {shapeB}; std::vectorint64_t kscaleShape {shapeNkv, shapeD}; std::vectorint64_t vscaleShape {shapeNkv, shapeD}; std::vectorint64_t koffsetShape {shapeNkv, shapeD}; std::vectorint64_t voffsetShape {shapeNkv, shapeD}; std::vectorint64_t weightShape {shapeH}; std::vectorint64_t activationShape {shapeB * shapeS}; std::vectorint64_t biasShape {shapeH}; // 以 INT32 輸入觸發(fā)反量化路徑需提供 weightScale/activationScale/bias std::vectorint32_t inputHostData(shapeB * shapeS * shapeH, 1); std::vectorint16_t cosHostData(shapeB * shapeS * shapeD, 1); std::vectorint16_t sinHostData(shapeB * shapeS * shapeD, 1); std::vectorint8_t kcacheHostData(shapeB * shapeH * shapeD, 6); std::vectorint8_t vcacheHostData(shapeB * shapeH * shapeD, 6); std::vectorint32_t indicesHostData(shapeB, 0); std::vectorfloat kscaleHostData(shapeNkv * shapeD, 2); std::vectorfloat vscaleHostData(shapeNkv * shapeD, 2); std::vectorfloat koffsetHostData(shapeNkv * shapeD, 2); std::vectorfloat voffsetHostData(shapeNkv * shapeD, 2); std::vectorfloat weightHostData(shapeH, 2); std::vectorfloat activationHostData(shapeB * shapeS, 2); std::vectorfloat biasHostData(shapeH, 2); // 省略逐個調(diào)用 CreateAclTensor 創(chuàng)建各輸入/輸出的 aclTensor 與 device 內(nèi)存 std::vectorint64_t splitData {shapeNq * shapeD, shapeNkv * shapeD, shapeNkv * shapeD}; aclIntArray *sizeSplits aclCreateIntArray(splitData.data(), splitData.size()); char quantMode[] static; char layout[] BSND; char cacheMode[] contiguous; // 3. 調(diào)用CANN算子庫API uint64_t workspaceSize 0; aclOpExecutor* executor; // 調(diào)用aclnnDequantRopeQuantKvcache第一段接口 ret aclnnDequantRopeQuantKvcacheGetWorkspaceSize(input, cos, sin, kcache, vcache, indices, kscale, vscale, koffset, voffset, weight, activation, bias, sizeSplits, quantMode, layout, true, cacheMode, q, k, v, workspaceSize, executor); CHECK_RET(ret ACL_SUCCESS, LOG_PRINT(aclnnDequantRopeQuantKvcacheGetWorkspaceSize failed. ERROR: %d\n, ret); return ret); // 根據(jù)第一段接口計算出的workspaceSize申請device內(nèi)存 void* workspaceAddr nullptr; if (workspaceSize 0) { ret aclrtMalloc(workspaceAddr, workspaceSize, ACL_MEM_MALLOC_HUGE_FIRST); CHECK_RET(ret ACL_SUCCESS, LOG_PRINT(allocate workspace failed. ERROR: %d\n, ret); return ret); } // 調(diào)用aclnnDequantRopeQuantKvcache第二段接口 ret aclnnDequantRopeQuantKvcache(workspaceAddr, workspaceSize, executor, stream); CHECK_RET(ret ACL_SUCCESS, LOG_PRINT(aclnnDequantRopeQuantKvcache failed. ERROR: %d\n, ret); return ret); // 4. 固定寫法同步等待任務執(zhí)行結(jié)束 ret aclrtSynchronizeStream(stream); CHECK_RET(ret ACL_SUCCESS, LOG_PRINT(aclrtSynchronizeStream failed. ERROR: %d\n, ret); return ret); // 5. 獲取輸出的值將device側(cè)內(nèi)存上的結(jié)果拷貝至host側(cè) PrintOutResult(kcacheShape, kcacheDeviceAddr); PrintOutResult(vcacheShape, vcacheDeviceAddr); // 6. 釋放aclTensor和aclIntArray // 7. 釋放device資源aclrtFree / aclrtDestroyStream / aclrtResetDevice / aclFinalize return 0; }上述示例中sizeSplits {Nq*D, Nkv*D, Nkv*D}與參數(shù)表要求一致。需要注意示例中indices全部為 0 且cacheMode contiguous因此量化后的 k/v 會被寫入每一條 batch 對應 cache 的第 0 個 token 位置若要在真實推理中做增量寫入需按 token 位置更新indices。源碼級原理深挖算子定義與數(shù)據(jù)類型注冊dequant_rope_quant_kvcache_def.cpp 中通過OP_ADD(DequantRopeQuantKvcache)完成算子注冊。該文件用 16 組 dtype 組合FLOAT16 / INT32 / BF16 三種 x 類型分別與 cos/sin、bias、scale、cache 類型的組合精確刻畫了文檔約束中數(shù)據(jù)類型保持一致的規(guī)則例如x為 FLOAT16 時cos/sin必須為 FLOAT16XDtypeList與cosDtypeList前 4 組一一對應x為 INT32 時bias可以是 FLOAT16 / BF16 / INT32 / FLOAT 四種biasDtypeList中對應條目即約束第 5 條biasOptional 不需要與其他輸入類型一致k_cache/v_cache固定為 INT8indices固定為 INT32scale/offset/weight_scale/activation_scale固定為 FLOAT。Shape 推導邏輯dequant_rope_quant_kvcache_infershape.cpp 中InferShapeForDequantRopeQuantKvcache負責輸出 shape 推導要求 x 為 2 維或 3 維cache 為 4 維從size_splits[0]與 cache 的隱藏維度推導qHead size_splits[0] / hiddenSizex 為 3 維時qOut為[B, S, Nq, D]kOut/vOut為[B, S, Nkv, D]x 為 2 維時seqlen 視為 1qOut為[B, Nq, D]kOut/vOut為[B, Nkv, D]當kv_output為 false 時kOut/vOut的第 0 維被置為 0shape 為空這與參數(shù)表中kvOutput 為 false 時 shape 為空完全對應InferDataTypeForDequantRopeQuantKvcache將 q/k/v 的輸出類型設置為與cos一致印證數(shù)據(jù)類型與 cos、sin 保持一致。Tiling 策略與 workspacedequant_rope_quant_kvcache_tiling.cpp 中的TilingDequantRopeQuantKvcache負責在 Host 側(cè)完成任務切分任務量建模taskNum batch * seqlen并按 AIV 核數(shù)GetCoreNumAiv()與 UB 容量GetCoreMemSize計算blockFactor前核塊因子與tailCoreBlockFactor尾核塊因子實現(xiàn)多核負載均衡UB 單次承載量OnceUBMaxS由 UB 剩余空間除以單次處理所需的 buffer 總量q/k/v/cos/sin/indices 的緩沖區(qū)之和全部按 32 字節(jié)BLOCK_SIZE對齊計算得到Kernel 按該值循環(huán)搬數(shù)Workspacetiling 階段統(tǒng)一申請MINIMAL_WORKSPACE 16MB的 workspace見 dequant_rope_quant_kvcache_tiling.cpp作為空 tensor 與 kernel 計算時的中間緩沖TilingKey以 bias 的數(shù)據(jù)類型作為SetTilingKey的取值FLOAT0 / FLOAT161 / INT322 / BF163驅(qū)動 Kernel 側(cè)模板實例化。tiling 數(shù)據(jù)的字段定義位于 dequant_rope_quant_kvcache_tiling.hqHeadNum、kvHeadNum、hiddenSize、OnceUBMaxS、isPA、ifKVout、hasBias、hasAS等這些字段直接決定了 Kernel 內(nèi)部的分支走向。Kernel 實現(xiàn)要點Kernel 側(cè)入口為 dequant_rope_quant_kvcache.cpp按 TilingKey 0/1/2/3 分別實例化RopeQuantKvcacheV2DTYPE_X, bias類型, DTYPE_COS模板bias 類型對應 FLOAT / half / int32_t / bfloat16_t。核心類實現(xiàn)在 dequant_rope_quant_kvcache.h反量化dequantUb對 INT32 輸入先Cast到 float再乘weight_scaleMul可選乘激活 scaleMuls與加 biasAdd切分搬數(shù)通過DataCopyPad與DataCopyExtParams以塊內(nèi) stride方式從inputGm中帶間隔地抽取 q、k、v 三段見dataCopyParamsQ_/K_/V_的構(gòu)造srcStride恰好跳過其他兩段避免三塊獨立訪存RoPE 計算按hiddenSize / 2拆分奇偶半段執(zhí)行k*cos ± 旋轉(zhuǎn)(k)*sin的旋轉(zhuǎn)位置編碼代碼中對sin先乘以 -1Muls再通過兩次Mul與一次Add完成標準 RoPE 公式量化對 kOut/vOut 先Div除以 scale可選Addoffset再經(jīng)Cast(CAST_RINT)到 INT16、轉(zhuǎn) half、最終Cast(CAST_NONE)到 INT8Cache 寫入copyOutcachepage 模式下以index * kvHeadNum * hiddenSize計算頁內(nèi)偏移contiguous 模式下以(bOffset bIndex) * cacheSeqlen index sIndex計算連續(xù)偏移隨后DataCopy寫回kCacheGm/vCacheGm流水并行通過MTE2_S、V_MTE3、MTE3_MTE2、MTE2_MTE3等硬件事件SetFlag/WaitFlag對搬入MTE2、向量計算V、搬出MTE3三段流水做同步編排降低訪存延遲。測試與驗證倉庫為算子提供了多級測試保障ST 測試目錄 tests/st/aclnnDequantRopeQuantKvcache 下的atk_aclnnDequantRopeQuantKvcache.json定義了基于 ATK 的端到端用例用例 0 覆蓋 x 為[1, 2304]的 FLOAT16 輸入Nq*D1536、Nkv*D384、Nkv*D384即 2D 輸入 cacheModepage用例 1 覆蓋 x 為[1, 3584]的 INT32 反量化路徑sizeSplits{1792, 896, 896}bias 為 BF16cacheModepage。兩個用例均攜帶backward: true可用于精度對比基準cv_fused_double_benchmarkUT 測試目錄 tests/ut 下包含 op_host 層的test_dequant_rope_quant_kvcache_infershape.cpp、test_dequant_rope_quant_kvcache_tiling.cpp以及 op_kernel 層的test_dequant_rope_quant_kvcache.cpp分別驗證 shape 推導、tiling 數(shù)據(jù)與 Kernel 計算結(jié)果的正確性。小結(jié)aclnnDequantRopeQuantKvcache是 CANN ops-transformer 中把反量化 → QKV 切分 → RoPE → 量化 → KV Cache 寫入五步融合的單算子實現(xiàn)。通過兩段式 aclnn 接口開發(fā)者可以在 Host 側(cè)一次性完成參數(shù)校驗、workspace 計算與執(zhí)行器構(gòu)建隨后在 Stream 上異步執(zhí)行。本文完整梳理了其計算流程、兩段式接口的全部參數(shù)語義、兩種 cache 更新模式與 7 條使用約束并結(jié)合倉庫中的算子定義、shape 推導、tiling 策略與 Kernel 實現(xiàn)揭示了底層的數(shù)據(jù)類型組合規(guī)則、多核任務劃分與流水并行機制。對于正在自研或接入大模型推理框架、需要在 NPU 上高效維護量化 KV Cache 的開發(fā)者該算子可作為 PagedAttention 場景中 KV 預處理環(huán)節(jié)的落地參考?!久赓M下載鏈接】ops-transformer本項目是CANN提供的transformer類大模型算子庫實現(xiàn)網(wǎng)絡在NPU上加速計算。項目地址: https://gitcode.com/cann/ops-transformer創(chuàng)作聲明:本文部分內(nèi)容由AI輔助生成(AIGC),僅供參考