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