位置編碼反向計算的融合實現(xiàn)與調(diào)用指南)
CANN ops-transformer 中 ApplyRotaryPosEmbGrad 算子詳解雙路旋轉(zhuǎn)位置編碼反向計算的融合實現(xiàn)與調(diào)用指南【免費下載鏈接】ops-transformer本項目是CANN提供的transformer類大模型算子庫實現(xiàn)網(wǎng)絡(luò)在NPU上加速計算。項目地址: https://gitcode.com/cann/ops-transformerApplyRotaryPosEmbGrad 是 CANN ops-transformer 算子庫中旋轉(zhuǎn)位置編碼RoPERotary Position Embedding系列的反向算子它將 query 與 key 兩路的 RoPE 梯度計算融合進一次 kernel 調(diào)用同時可選輸出 cos/sin 的梯度。本文以 README 為主體結(jié)合倉庫內(nèi) aclnn 接口文檔、PyTorch 封裝、tiling 與 kernel 源碼完整講解該算子的數(shù)學(xué)原理、參數(shù)約束、兩種調(diào)用方式以及底層多模板調(diào)度實現(xiàn)幫助你直接上手訓(xùn)練場景下的 RoPE 反向計算。一、算子功能與應(yīng)用場景該算子是雙路旋轉(zhuǎn)位置編碼算子 ApplyRotaryPosEmb 的反向算子核心功能如下執(zhí)行雙路反向計算同時計算query和key的 rope 反向梯度融合為一次 kernel 調(diào)用節(jié)省開銷相比分別對 query、key 各執(zhí)行一次反向 kernel融合實現(xiàn)節(jié)省了 cos/sin 的重復(fù)加載和 kernel launch 開銷可選計算 cos/sin 梯度當(dāng)正向輸入query、key同時傳入時額外計算grad_cos與grad_sin供需要更新 cos/sin 參與反傳的場景使用。從產(chǎn)品支持情況看該算子目前僅面向Ascend 950PR / Ascend 950DT產(chǎn)品即源碼配置目錄 config/ascend950 對應(yīng)的 SoCAtlas A2/A3、Atlas 200I/500 A2、Atlas 推理/訓(xùn)練系列等產(chǎn)品均不支持。表格如下產(chǎn)品是否支持Ascend 950PR/Ascend 950DT√Atlas A3 訓(xùn)練系列產(chǎn)品/Atlas A3 推理系列產(chǎn)品×Atlas A2 訓(xùn)練系列產(chǎn)品/Atlas A2 推理系列產(chǎn)品×Atlas 200I/500 A2 推理產(chǎn)品×Atlas 推理系列產(chǎn)品×Atlas 訓(xùn)練系列產(chǎn)品×二、數(shù)學(xué)原理與計算公式旋轉(zhuǎn)位置編碼RoPE的核心思想是在 Head-Dim 維度上將每個頭的向量按 D/2 拆成前后兩半通過 cos/sin 旋轉(zhuǎn)矩陣完成位置信息注入。反向計算即對該旋轉(zhuǎn)過程求導(dǎo)。取正向計算中cos、sin發(fā)生 broadcast 的軸列表為dims即 cos/sin 中取值為 1、而grad_query_embed/grad_key_embed中對應(yīng)維度大于 1 的軸包含 N 軸以及 BSND、SBND 布局下可選的 B 軸rotary_mode為half時的計算公式如下。1對輸入梯度做末維對半切分$$ grad_q_1, grad_q_2 chunk(grad_query_embed, chunks2, dim-1) $$$$ grad_k_1, grad_k_2 chunk(grad_key_embed, chunks2, dim-1) $$$$ cos_1, cos_2 chunk(cos, chunks2, dim-1) $$$$ sin_1, sin_2 chunk(sin, chunks2, dim-1) $$2構(gòu)造旋轉(zhuǎn)后的正向向量用于 cos/sin 梯度$$ query_rotate cat((-query_2, query_1), dim-1) $$$$ key_rotate cat((-key_2, key_1), dim-1) $$3計算 query/key 的梯度$$ grad_query cat(cos_1 * grad_q_1 sin_2 * grad_q_2, cos_2 * grad_q_2 - sin_1 * grad_q_1, dim-1) $$$$ grad_key cat(cos_1 * grad_k_1 sin_2 * grad_k_2, cos_2 * grad_k_2 - sin_1 * grad_k_1, dim-1) $$4當(dāng)同時傳入 query 和 key 時沿廣播軸 dims 歸約得到 cos/sin 梯度$$ grad_cos sum(grad_query_embed * query grad_key_embed * key, dims) $$$$ grad_sin sum(grad_query_embed * query_rotate grad_key_embed * key_rotate, dims) $$倉庫中的測試 golden 腳本 tests/assets/golden.py 以注釋形式完整復(fù)現(xiàn)了上述公式并注明“所有路徑統(tǒng)一升 FP32 計算結(jié)果轉(zhuǎn)回輸入 dtype”可供理解參考。三、參數(shù)說明各參數(shù)的完整說明如下表參數(shù)名輸入/輸出/屬性描述數(shù)據(jù)類型數(shù)據(jù)格式grad_query_embed輸入正向輸出 query 的導(dǎo)數(shù)對應(yīng)公式中 $grad_q_{embed}$。BFLOAT16、FLOAT16、FLOAT32NDgrad_key_embed輸入正向輸出 key 的導(dǎo)數(shù)對應(yīng)公式中 $grad_k_{embed}$。BFLOAT16、FLOAT16、FLOAT32NDcos輸入正向計算輸入 cos需與 grad_query_embed 數(shù)據(jù)類型一致。BFLOAT16、FLOAT16、FLOAT32NDsin輸入正向計算輸入 sin需與 grad_query_embed 數(shù)據(jù)類型一致。BFLOAT16、FLOAT16、FLOAT32NDquery可選輸入正向計算輸入 query。如果為空指針則不計算 grad_cos 和 grad_sin必須與 key 同時傳入或同時不傳入。BFLOAT16、FLOAT16、FLOAT32NDkey可選輸入正向計算輸入 key。如果為空指針則不計算 grad_cos 和 grad_sin必須與 query 同時傳入或同時不傳入。BFLOAT16、FLOAT16、FLOAT32NDrotary_mode屬性旋轉(zhuǎn)模式僅支持 half。STRING-layout屬性輸入 Tensor 的布局格式。1BSND2SBND4TND。默認值為 1。INT64-grad_query輸出正向計算輸入 query 的導(dǎo)數(shù)shape 與 grad_query_embed 相同。BFLOAT16、FLOAT16、FLOAT32NDgrad_key輸出正向計算輸入 key 的導(dǎo)數(shù)shape 與 grad_key_embed 相同。BFLOAT16、FLOAT16、FLOAT32NDgrad_cos輸出正向計算輸入 cos 的導(dǎo)數(shù)僅當(dāng) query 和 key 非空時有效。BFLOAT16、FLOAT16、FLOAT32NDgrad_sin輸出正向計算輸入 sin 的導(dǎo)數(shù)僅當(dāng) query 和 key 非空時有效。BFLOAT16、FLOAT16、FLOAT32ND關(guān)于 layout 的補充說明BBatch批量大小SSeq-Length序列長度NHead-Num多頭數(shù)DHead-Dim每個頭的隱藏維度大小TB 和 S 的合軸layout4時輸入為 3 維 Tensor其他 layout 下為 4 維。從 算子定義源碼 可以看到host 側(cè)注冊的輸入輸出 dtype 均為DT_FLOAT16 / DT_FLOAT / DT_BF16格式為FORMAT_ND屬性默認值分別為rotary_modehalf、layout1與上述參數(shù)表完全對應(yīng)同時注冊了DynamicCompileStaticFlag / DynamicRankSupportFlag / DynamicShapeSupportFlag表明算子支持動態(tài) shape。四、約束說明輸入輸出 Tensor 只支持 3 維或 4 維layout 為 1 或 2 時為 4 維layout 為 4 時為 3 維。輸入輸出 Tensor 的 dtype 必須相同。輸入輸出 Tensor 不支持空 Tensor各維度必須大于 0。輸入輸出 Tensor 的 layout 必須相同。輸入輸出 Tensor 的 D 軸必須相同在 half 模式下必須 ≤ 1024 且能被 2 整除。grad_query_embed、grad_query的 shape 必須相同grad_key_embed、grad_key的 shape 必須相同。對于任意 layoutgrad_query_embed和grad_key_embed除 N 維度外其它維度必須相同。cos、sin的 N 維度必須等于 1layout 為 1BSND或 2SBND時cos、sin的 B 維度可以等于 1也可以和grad_query_embed的 B 維度一致layout 為 4TND時cos、sin的 T 維度必須和grad_query_embed的 T 維度一致除 N及 BSND、SBND 布局下可選廣播的 B維度外其余維度需與grad_query_embed一致。cos、sin、grad_cos、grad_sin的 shape 必須相同。query維度需與grad_query_embed一致key維度需與grad_key_embed一致且query和key必須同時傳入或同時不傳入。rotary_mode僅支持 half。layout僅支持 {1, 2, 4}對應(yīng) {BSND, SBND, TND}。3BNSD 為預(yù)留暫不支持。這些約束在 tiling 源碼中均有對應(yīng)的顯式校驗。例如 apply_rotary_pos_emb_grad_tiling.cpp 中CheckRotaryModeShapeRelation()校驗 D 軸≤ 1024D_LIMIT且% 2 0HALF_MODE_COEFValidateBroadcastByLayout()按 BSND/SBND/TND 分別校驗 cos/sin 的 B/S/T 與 N 軸廣播關(guān)系TND 下 cos 的 T 軸必須等于 grad_query_embed 的 TN 軸必須為 1BSND/SBND 下 cos 的 B 軸必須為 1 或等于輸入 BS 軸必須一致CheckShape()校驗grad_query_embed與grad_key_embed除 N 軸4D 布局下的 dim 2外各維度相同CheckOptionalInput()校驗 query 與 grad_query_embed、key 與 grad_key_embed、cos 與 grad_cos、sin 與 grad_sin 的 shape 全等。五、aclnn 調(diào)用方式兩段式接口aclnn 調(diào)用遵循 CANN 單算子調(diào)用的兩段式接口規(guī)范必須先調(diào)用第一段aclnnApplyRotaryPosEmbGradGetWorkspaceSize完成入?yún)⑿r灢⒂嬎?workspace 大小再調(diào)用第二段aclnnApplyRotaryPosEmbGrad執(zhí)行計算。函數(shù)原型如下aclnnStatus aclnnApplyRotaryPosEmbGradGetWorkspaceSize( const aclTensor *gradQueryEmbed, const aclTensor *gradKeyEmbed, const aclTensor *cos, const aclTensor *sin, const aclTensor *queryOptional, const aclTensor *keyOptional, char *rotaryModeOptional, int64_t layout, const aclTensor *gradQueryOut, const aclTensor *gradKeyOut, const aclTensor *gradCosOut, const aclTensor *gradSinOut, uint64_t *workspaceSize, aclOpExecutor **executor)aclnnStatus aclnnApplyRotaryPosEmbGrad( void *workspace, uint64_t workspaceSize, aclOpExecutor *executor, aclrtStream stream)5.1 第一段接口參數(shù)與返回值第一段接口的參數(shù)細節(jié)完整見 aclnnApplyRotaryPosEmbGrad 接口文檔參數(shù)名輸入/輸出描述使用說明數(shù)據(jù)類型維度(shape)非連續(xù)TensorgradQueryEmbed輸入正向輸出 query 的導(dǎo)數(shù)對應(yīng) grad_q_embed不支持空 TensorBFLOAT16/FLOAT16/FLOAT324(layout 1/2)或 3(layout 4)√gradKeyEmbed輸入正向輸出 key 的導(dǎo)數(shù)對應(yīng) grad_k_embed與 gradQueryEmbed 類型和維度一致不支持空 Tensor同上同上√cos輸入正向計算輸入 cos與 gradQueryEmbed 類型和維度一致N 維必須為 1同上同上√sin輸入正向計算輸入 sin與 gradQueryEmbed 類型和維度一致N 維必須為 1同上同上√queryOptional可選輸入正向輸入 query空指針時不計算 gradCos/gradSin與 keyOptional 必須同時傳入或同時不傳同上同上√keyOptional可選輸入正向輸入 key空指針時不計算 gradCos/gradSin與 queryOptional 必須同時傳入或同時不傳同上同上√rotaryModeOptional輸入旋轉(zhuǎn)模式僅支持 halfSTRING--layout輸入輸入 Tensor 布局1-BSND2-SBND4-TND3-BNSD(預(yù)留)INT64--gradQueryOut輸出query 的導(dǎo)數(shù)與 gradQueryEmbed 類型和維度一致同上同上×gradKeyOut輸出key 的導(dǎo)數(shù)與 gradQueryEmbed 類型和維度一致同上同上×gradCosOut輸出cos 的導(dǎo)數(shù)query/key 非空時有效與 gradQueryEmbed 類型和維度一致同上同上×gradSinOut輸出sin 的導(dǎo)數(shù)query/key 非空時有效與 gradQueryEmbed 類型和維度一致同上同上×workspaceSize輸出Device 側(cè)需申請的 workspace 大小----executor輸出op 執(zhí)行器包含算子計算流程----第一段接口完成入?yún)⑿r灣霈F(xiàn)以下場景時報錯具體錯誤碼定義見 aclnn 返回碼返回值錯誤碼描述ACLNN_ERR_PARAM_NULLPTR161001必選輸入 gradQueryEmbed、gradKeyEmbed、cos、sin 和必選輸出 gradQueryOut、gradKeyOut 是空指針ACLNN_ERR_PARAM_INVALID161002輸入輸出數(shù)據(jù)類型/格式不在支持范圍、shape 不滿足校驗、維度不在支持范圍、queryOptional 與 keyOptional 未成對傳入、或 rotaryMode/layout 不符合支持值第二段接口接收workspace、workspaceSize、executor與stream四個參數(shù)workspace為 Device 側(cè)申請的臨時內(nèi)存地址workspaceSize由第一段接口計算得出executor為第一段接口返回的 op 執(zhí)行器stream指定任務(wù)執(zhí)行的 Stream 流。注意第二段接口不可重復(fù)調(diào)用每次執(zhí)行需重新走兩段式流程。5.2 完整調(diào)用示例倉庫提供了可參考的完整示例 examples/test_aclnn_apply_rotary_pos_emb_grad.cpp核心流程如下編譯與運行方法參考編譯與運行樣例#include acl/acl.h #include aclnnop/aclnn_apply_rotary_pos_emb_grad.h #include iostream #include vector // 1. 資源初始化aclInit / aclrtSetDevice / aclrtCreateStream固定寫法 // 2. 構(gòu)造輸入輸出以 BSND layout、D128 為例 std::vectorint64_t gradQEmbedShape {1, 1, 1, 128}; std::vectorint64_t gradKEmbedShape {1, 1, 1, 128}; std::vectorint64_t cosShape {1, 1, 1, 128}; // N 維必須為 1 std::vectorint64_t sinShape {1, 1, 1, 128}; std::vectorint64_t queryShape {1, 1, 1, 128}; std::vectorint64_t keyShape {1, 1, 1, 128}; std::vectorint64_t gradQueryOutShape {1, 1, 1, 128}; std::vectorint64_t gradKeyOutShape {1, 1, 1, 128}; std::vectorint64_t gradCosOutShape {1, 1, 1, 128}; std::vectorint64_t gradSinOutShape {1, 1, 1, 128}; int64_t layout 1; // BSND const char *rotaryModeOptional half; // 通過 aclrtMalloc aclrtMemcpy 將 host 數(shù)據(jù)搬入 device // 再以 aclCreateTensor(..., ACL_FORMAT_ND, ...) 構(gòu)造各 aclTensor // 3. 第一段接口計算 workspace 大小并創(chuàng)建 executor uint64_t workspaceSize 0; aclOpExecutor *executor; ret aclnnApplyRotaryPosEmbGradGetWorkspaceSize( gradQueryEmbed, gradKeyEmbed, cos, sin, query, key, const_castchar *(rotaryModeOptional), layout, gradQueryOut, gradKeyOut, gradCosOut, gradSinOut, workspaceSize, executor); // 4. 按 workspaceSize 申請 device 內(nèi)存 void *workspaceAddr nullptr; if (workspaceSize 0) { ret aclrtMalloc(workspaceAddr, workspaceSize, ACL_MEM_MALLOC_HUGE_FIRST); } // 5. 第二段接口執(zhí)行計算 ret aclnnApplyRotaryPosEmbGrad(workspaceAddr, workspaceSize, executor, stream); // 6. 同步并取回結(jié)果 ret aclrtSynchronizeStream(stream); ret aclrtMemcpy(resultData.data(), ..., gradQueryOutDeviceAddr, ..., ACL_MEMCPY_DEVICE_TO_HOST); // 7. 釋放 aclTensor、device 內(nèi)存與 stream示例中 shape 均取{1, 1, 1, 128}BSN1、D128滿足 D ≤ 1024 且可被 2 整除query/key 同時傳入以觸發(fā) grad_cos/grad_sin 計算。實際使用時可根據(jù)模型配置替換為 BSND、SBND 或 TND 布局并注意 cos/sin 的 N 維必須為 1。六、PyTorch API 調(diào)用方式PyTorch 側(cè)封裝位于 torch_extension/apply_rotary_pos_emb_grad.py通過torch.library注冊自定義算子并以PrivateUse1后端分發(fā)到 NPU。函數(shù)原型如下完整文檔見 torchapi_apply_rotary_pos_emb_gradcann_ops_transformer.apply_rotary_pos_emb_grad( grad_query_embed, grad_key_embed, cos, sin, *, queryNone, keyNone, rotary_modehalf, layout1, ) - (Tensor, Tensor, Optional[Tensor], Optional[Tensor])6.1 參數(shù)說明參數(shù)名參數(shù)類型可選/必選描述數(shù)據(jù)類型維度(shape)grad_query_embedTensor必選正向輸出query_out的梯度bfloat16、float16、float32layout4 時為 (T, Nq, D)其他 layout 下為 4 維 Tensorgrad_key_embedTensor必選正向輸出key_out的梯度除 N 維度外 shape 需與 grad_query_embed 一致同 grad_query_embedlayout4 時為 (T, Nk, D)其他 layout 下為 4 維 TensorcosTensor必選正向計算輸入的余弦值張量N 維度必須等于 1同 grad_query_embed與輸入布局對應(yīng)的 3 維或 4 維 TensorsinTensor必選正向計算輸入的正弦值張量shape 需與 cos 一致同 grad_query_embed同 cosqueryTensor可選正向計算輸入 query傳入時計算 grad_cos 和 grad_sin必須與 key 同時傳入或同時不傳入默認 None同 grad_query_embed與 grad_query_embed 一致keyTensor可選正向計算輸入 key傳入時計算 grad_cos 和 grad_sin必須與 query 同時傳入或同時不傳入默認 None同 grad_query_embed與 grad_key_embed 一致rotary_modestr可選旋轉(zhuǎn)編碼模式僅支持 half默認 half--layoutint可選1 表示 BSND2 表示 SBND4 表示 TND默認 1--返回值說明grad_queryTensor正向輸入 query 的梯度shape 和數(shù)據(jù)類型與 grad_query_embed 一致grad_keyTensor正向輸入 key 的梯度shape 和數(shù)據(jù)類型與 grad_key_embed 一致grad_cosOptional[Tensor]正向輸入 cos 的梯度query 和 key 均傳入時 shape 與 cos 一致否則為Nonegrad_sinOptional[Tensor]正向輸入 sin 的梯度query 和 key 均傳入時 shape 與 sin 一致否則為None。封裝源碼中的_check_inputs函數(shù)在 Python 側(cè)提前完成了與上節(jié)約束一致的校驗dtype 僅支持 float16/float32/bfloat16 且必須一致、維度必須為 3 或 4、TND(4) 布局要求 3 維輸入、query/key 必須成對出現(xiàn)、query 與 grad_query_embed shape 相等、key 與 grad_key_embed shape 相等、rotary_mode 僅支持 half 等。Meta 實現(xiàn)apply_rotary_pos_emb_grad_meta則負責(zé) shape/dtype 推導(dǎo)支撐 Autograd 與 FakeTensor 場景。6.2 單算子模式調(diào)用示例import torch import torch_npu from cann_ops_transformer.ops import apply_rotary_pos_emb_grad torch_npu.npu.set_device(0) B 1 S 64 N 8 D 128 grad_query_embed torch.randn(B, S, N, D, devicenpu, dtypetorch.float16) grad_key_embed torch.randn(B, S, N, D, devicenpu, dtypetorch.float16) cos torch.randn(B, S, 1, D, devicenpu, dtypetorch.float16) # N 維為 1可廣播 sin torch.randn(B, S, 1, D, devicenpu, dtypetorch.float16) query torch.randn(B, S, N, D, devicenpu, dtypetorch.float16) key torch.randn(B, S, N, D, devicenpu, dtypetorch.float16) grad_query, grad_key, grad_cos, grad_sin apply_rotary_pos_emb_grad( grad_query_embed, grad_key_embed, cos, sin, queryquery, keykey, rotary_modehalf, layout1, # BSND ) print(fgrad_query shape: {grad_query.shape}) print(fgrad_key shape: {grad_key.shape}) print(fgrad_cos shape: {grad_cos.shape}) print(fgrad_sin shape: {grad_sin.shape})該示例展示了 BSND 布局下帶 cos/sin 梯度的完整調(diào)用cos/sin取(B, S, 1, D)N 維為 1 從而沿 N 軸廣播若不需要 grad_cos/grad_sin將query、key均置為None即可返回的 grad_cos/grad_sin 為None。6.3 與正向算子的配套使用該算子為 apply_rotary_pos_emb 的反向算子。正向接口使用rotary_modehalf時對 loss 執(zhí)行.backward()會自動觸發(fā)本算子僅在需要顯式控制梯度時才需要手動調(diào)用本接口。該接口支持訓(xùn)練場景下單算子模式調(diào)用且默認支持確定性計算aclnn 與 torch API 文檔均明確標(biāo)注“默認確定性實現(xiàn)”。七、底層實現(xiàn)從算子定義到多模板調(diào)度7.1 算子定義與 shape 推導(dǎo)apply_rotary_pos_emb_grad_def.cpp注冊 6 個輸入grad_query_embed、grad_key_embed、cos、sin 為 REQUIREDquery、key 為 OPTIONAL、4 個輸出grad_query、grad_key 為 REQUIREDgrad_cos、grad_sin 為 OPTIONALdtype 支持 FLOAT16/FLOAT/BF16格式統(tǒng)一 ND并聲明DynamicCompileStaticFlag(true)、DynamicRankSupportFlag(true)、DynamicShapeSupportFlag(true)AICore 配置僅注冊 ascend950。apply_rotary_pos_emb_grad_infershape.cppgrad_query 繼承 grad_query_embed 的 shape、grad_key 繼承 grad_key_embed 的 shape、grad_cos 繼承 cos 的 shape、grad_sin 繼承 sin 的 shapeSetGradOutputShapedtype 同理逐輸出透傳動態(tài) shape 下以 -2unknown rank/-1unknown dim占位。7.2 tiling 階段的廣播判定與三模板調(diào)度tiling 核心邏輯位于 apply_rotary_pos_emb_grad_tiling.cpp 及三個模板實現(xiàn)文件_bab/_ab/_a。host 側(cè)先執(zhí)行一整套參數(shù)校驗dtype、維度、D 軸上限與奇偶、cos/sin 廣播關(guān)系、query/key 成對性等隨后按 shape 關(guān)系判定內(nèi)部ApplyRopeGradLayoutTND 布局3 維輸入退化為 B1 的 BSND 處理若 N1 則走 NO_BROADCASTA 模板BSND 布局cos 的 B 軸等于 1 時進入 BSND 廣播BAB 模板cos 的 B 軸與輸入 B 一致時進入 SBNDAB 模板若 cos 與輸入各維度完全一致含 gk 的 N 軸則判定為無廣播A 模板SBND 布局shape 完全一致時回落 NO_BROADCASTA 模板否則按 AB 模板處理。最終通過三種 kernel tiling key 調(diào)度見 apply_rotary_pos_emb_grad_apt.cpp 中的枚舉Tiling Key枚舉值適用場景TILING_KEY_BAB203BSND 布局 cos B 軸1B/S/N 三層廣播迭代TILING_KEY_AB204SBND 布局或 BSND 下 cos B 軸與輸入一致TILING_KEY_A205無廣播shape 完全一致最簡路徑7.3 kernel 執(zhí)行流程與 workspacekernel 側(cè)以__global__ __aicore__模板函數(shù)實現(xiàn)AIC 核直接返回、僅 AIV 核執(zhí)行。計算分為兩個 PhasePhase 1計算 grad_query / grad_key。BAB/AB 模板下若同時需要 grad_cos/grad_sinDcosFlag1會在 kernel 內(nèi)同步累加 grad_cos/grad_sin 的部分積并寫入 workspaceA 模板則分三個階段先算 dx再預(yù)計算rotate(query)/rotate(key)寫入 workspace最后通過ApplyDcosDsin高層 Mul 與 Q/K 累加直接寫回 GM。Phase 2Reduce 歸約。僅廣播模板BAB/AB需要——將 Phase 1 產(chǎn)生的 dcos/dsin 部分積沿廣播軸N 軸等跨核歸約得到最終的 grad_cos/grad_sinA 模板因無廣播無需 Reduce。workspace 大小由 tiling 計算并寫入 tiling datausrWorkSpaceSize b * s * max(nQ, nK) * d * partialTypeSize * INPUT_OUTPUT_NUM兩份部分積廣播模板下 partialTypeSize 為 float 大小A 模板還有 16MB 的系統(tǒng) workspace 預(yù)留這與第一段接口返回的workspaceSize直接對應(yīng)。7.4 配置與測試驗證算子二進制按 dtype 拆分為三個 binApplyRotaryPosEmbGrad_1/2/3對應(yīng) float16/bfloat16/float32見 apply_rotary_pos_emb_grad_binary.json單測覆蓋 infershape 與 tiling 校驗邏輯tests/ut/op_hostkernel 級 UT 直接包含_apt.cpp源碼完成模板實例化覆蓋 BAB/AB/A 三條路徑tests/ut/op_kernel/arch35/test_apply_rotary_pos_emb_grad.cpp多路徑 golden 腳本 tests/assets/golden.py 同時為 kernel spec、aclnn spec 與 E2E spec 提供參照實現(xiàn)q/k 兩路分別歸約以兼容 Nq ≠ Nk。八、總結(jié)ApplyRotaryPosEmbGrad 是 CANN ops-transformer 中面向 Ascend 950PR/950DT 的雙路 RoPE 反向算子它一次 kernel 調(diào)用同時產(chǎn)出 grad_query、grad_key并在傳入 query/key 時額外產(chǎn)出 grad_cos、grad_sin內(nèi)部依據(jù)布局與廣播關(guān)系在 BAB/AB/A 三種模板間自動選擇配合 Reduce 歸約與部分積 workspace 完成帶廣播的反向計算。實際使用時重點把握三類約束D 軸 ≤ 1024 且為偶數(shù)、cos/sin 的 N 維必須為 1、query 與 key 必須成對出現(xiàn)。相關(guān)接口文檔與示例代碼可直接參考 aclnnApplyRotaryPosEmbGrad 文檔、torchapi 文檔 與 調(diào)用示例?!久赓M下載鏈接】ops-transformer本項目是CANN提供的transformer類大模型算子庫實現(xiàn)網(wǎng)絡(luò)在NPU上加速計算。項目地址: https://gitcode.com/cann/ops-transformer創(chuàng)作聲明:本文部分內(nèi)容由AI輔助生成(AIGC),僅供參考