已合并
feat(hard_shrink): def 驱动 dtype + abs/le 公式对齐竞品 #8756
feat(hard_shrink): def 驱动 dtype + abs/le 公式对齐竞品 #8756
已合并
leaving__创建于 8月17日
共 6 个文件变更+116-126
@@ -0,0 +1,34 @@
1+/**
2+ * Copyright (c) 2026 Huawei Technologies Co., Ltd.
3+ * This program is free software, you can redistribute it and/or modify it under the terms and conditions of
4+ * CANN Open Software License Agreement Version 2.0 (the "License").
5+ * Please refer to the License for details. You may not use this file except in compliance with the License.
6+ * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
7+ * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
8+ * See LICENSE in the root of the software repository for the full text of the License.
9+ */
10+ 
11+/* Generated By CANNBot */
12+ 
13+/*!
14+ * \file hard_shrink_infer.cpp
15+ * \brief HardShrink 算子 InferDataType 实现(图模式)
16+ *
17+ * 复用 canndev 仓 runtime2.0 逐元素算子的 dtype 推导模式(InferDataType4SameAsInput):
18+ * 输出 dtype = 输入 dtype。与 op_host/hard_shrink_infershape.cpp 分文件放置(参考 ops-cv PR#1126)。
19+ */
20+ 
21+#include "register/op_impl_registry.h"
22+#include "log/log.h"
23+ 
24+using namespace ge;
25+namespace ops {
26+static ge::graphStatus InferDataType4HardShrink(gert::InferDataTypeContext* context)
27+{
28+ const ge::DataType inputDataType = context->GetInputDataType(0);
29+ context->SetOutputDataType(0, inputDataType);
30+ return ge::GRAPH_SUCCESS;
31+}
32+ 
33+IMPL_OP(HardShrink).InferDataType(InferDataType4HardShrink);
34+} // namespace ops
@@ -18,7 +18,9 @@
18 * - 多核切分:totalNum / coreNum18 * - 多核切分:totalNum / coreNum
19 * - UB 切分:ubSize / bufferNum / typeSize,256B 对齐19 * - UB 切分:ubSize / bufferNum / typeSize,256B 对齐
20 * - 双缓冲阈值:totalLength > 1024 → BUFFER_MODE=120 * - 双缓冲阈值:totalLength > 1024 → BUFFER_MODE=1
21- * - 升精处理:FP16 / BF16 走 NEED_UPCAST=1(cast 到 fp32 计算);FP32 走 NEED_UPCAST=0(直通)21+ * - 升精处理:FP16 / BF16 走 fp32 升精计算;FP32 直通
22+ * NEED_UPCAST 由 kernel 侧据 DTYPE_SELF 编译期推导,host 仅据 dtype 计算 UB 切分,
23+ * 不再作为 TilingKey 维度编码(def 已通过 DTYPE_SELF 宏覆盖 dtype 维度)
22 */24 */
23 25 
24#include "register/op_def_registry.h"26#include "register/op_def_registry.h"
@@ -113,11 +115,10 @@ static ge::graphStatus HandleEmptyTensor(gert::TilingContext* context, ge::DataT
113 size_t* currentWorkspace = context->GetWorkspaceSizes(1);115 size_t* currentWorkspace = context->GetWorkspaceSizes(1);
114 OP_CHECK_NULL_WITH_CONTEXT(context, currentWorkspace);116 OP_CHECK_NULL_WITH_CONTEXT(context, currentWorkspace);
115 currentWorkspace[0] = WS_SYS_SIZE;117 currentWorkspace[0] = WS_SYS_SIZE;
116- uint32_t dType = static_cast<uint32_t>(dataType);118+ (void)dataType;
119+ // def 驱动 dtype:tiling_key 只编码 bufferMode,dtype 由 def 注入 DTYPE_SELF 宏
117 uint32_t bufferMode = 0;120 uint32_t bufferMode = 0;
118- // FP16/BF16 走升精,FP32 直通121+ ASCENDC_TPL_SEL_PARAM(context, bufferMode);
119- uint32_t needUpcast = (dataType == ge::DT_FLOAT) ? 0 : 1;
120- ASCENDC_TPL_SEL_PARAM(context, dType, bufferMode, needUpcast);
121 return ge::GRAPH_SUCCESS;122 return ge::GRAPH_SUCCESS;
122}123}
123 124 
@@ -198,9 +199,9 @@ static ge::graphStatus HardShrinkTilingFunc(gert::TilingContext* context)
198 199 
199 context->SetBlockDim(usedCoreNum);200 context->SetBlockDim(usedCoreNum);
200 201 
201- uint32_t dType = static_cast<uint32_t>(inputInfo.dataType);202+ // def 驱动 dtype:tiling_key 只编码 bufferMode,dtype 由 def 注入 DTYPE_SELF 宏,
202- uint32_t needUpcastFlag = needUpcast ? 1 : 0;203+ // NEED_UPCAST 由 kernel 侧据 DTYPE_SELF 编译期推导,host 不再编码
203- ASCENDC_TPL_SEL_PARAM(context, dType, bufferMode, needUpcastFlag);204+ ASCENDC_TPL_SEL_PARAM(context, bufferMode);
204 205 
205 return ge::GRAPH_SUCCESS;206 return ge::GRAPH_SUCCESS;
206}207}
@@ -14,36 +14,20 @@
14 * \file hard_shrink_infershape.cpp14 * \file hard_shrink_infershape.cpp
15 * \brief HardShrink 算子形状推导实现15 * \brief HardShrink 算子形状推导实现
16 *16 *
17- * HardShrink 是逐元素算子:输出 shape = 输入 shape,输出 dtype = 输入 dtype17+ * 复用 canndev 仓 runtime2.0 的逐元素算子公共推导(InferShape4Elewise):
18+ * 输出 shape = 输入 shape。canndev 仓原 InferShape 逻辑保留不删,长尾算子仍使用。
18 */19 */
19 20 
21+#include "infershape_elewise_util.h"
20#include "register/op_impl_registry.h"22#include "register/op_impl_registry.h"
21-#include "exe_graph/runtime/infer_shape_context.h"23+#include "log/log.h"
22 24 
23using namespace ge;25using namespace ge;
24- 
25namespace ops {26namespace ops {
26- 
27static ge::graphStatus InferShape4HardShrink(gert::InferShapeContext* context)27static ge::graphStatus InferShape4HardShrink(gert::InferShapeContext* context)
28{28{
29- // 获取输入 self 的形状29+ return Ops::Base::InferShape4Elewise(context);
30- const gert::Shape* inputShape = context->GetInputShape(0);
31- if (inputShape == nullptr) {
32- return ge::GRAPH_FAILED;
33- }
34- 
35- // 获取输出 out 的形状
36- gert::Shape* outputShape = context->GetOutputShape(0);
37- if (outputShape == nullptr) {
38- return ge::GRAPH_FAILED;
39- }
40- 
41- // 输出 shape = 输入 shape
42- *outputShape = *inputShape;
43- 
44- return ge::GRAPH_SUCCESS;
45}30}
46 31 
47IMPL_OP_INFERSHAPE(HardShrink).InferShape(InferShape4HardShrink);32IMPL_OP_INFERSHAPE(HardShrink).InferShape(InferShape4HardShrink);
48- 
49} // namespace ops33} // namespace ops
@@ -14,23 +14,26 @@
14 * \file hard_shrink.h14 * \file hard_shrink.h
15 * \brief HardShrink 算子 Kernel 类定义(arch35 架构,Ascend950)15 * \brief HardShrink 算子 Kernel 类定义(arch35 架构,Ascend950)
16 *16 *
17- * 计算公式: HardShrink(x) = x if (|x| > lambd) or isnan(x), else 017+ * 计算公式: HardShrink(x) = x if |x| > lambd, else 0 (边界 |x|==lambd 归 0)
18 *18 *
19- * 实现方案: 三次 Compare+Select(NaN 透传 + 正侧 + 负侧)19+ * 实现方案: Abs + Compare(LE) + Select(与竞品 canndev/torch 公式顺序一致)
20- * Step 1: mask_gt = (x > lambd), tmp1 = mask_gt ? x : 020+ * canndev tbe: input_x_abs = vabs(x); result = vcmpsel(|x|, lambd, 'le', 0, x)
21- * Step 2: mask_lt = (x < -lambd), tmp2 = mask_lt ? x : tmp121+ * torch: x if |x| > lambd else 0
22- * Step 3: mask_nan = (x != x), out = mask_nan ? x : tmp222+ * Step 1: absX = |x| (AscendC::Abs)
23- * 说明:IEEE 754 下 NaN 满足 (NaN != NaN)==true,借此识别并透传 NaN23+ * Step 2: mask = (|x| <= lambd) (Compare LE)
24- * (避免引入 Or API 兼容性风险,三次 Select 数学等价于24+ * Step 3: out = mask ? 0 : x (Select TENSOR_TENSOR,le→0, else→x)
25- * select(isnan(x) || (x < -lambd) || (x > lambd), x, 0))25+ * NaN/inf 天然透传:IEEE754 下 NaN<=lambd 与 inf<=lambd 均为 false → 走 else 选原值透传
26+ * (与 canndev vcmpsel LE 语义一致,无需额外 NaN 分支)
26 *27 *
27 * 模板参数:28 * 模板参数:
28- * - T: IO 数据类型 (half/float/bfloat16_t)29+ * - T: IO 数据类型 (half/float/bfloat16_t),由 def 驱动的 DTYPE_SELF 宏注入
29- * - BUFFER_MODE: 0=单缓冲, 1=双缓冲30+ * - BUFFER_MODE: 0=单缓冲, 1=双缓冲(唯一由 tiling_key 编码的维度)
30- * - NEED_UPCAST: 是否在内部把 IO 升精到 fp32 计算31+ *
31- * · 0: 直接以 T 计算(fp32 路径)32+ * NEED_UPCAST 不再作为 tiling_key 维度传入:它由 IO 类型 T 唯一决定,
32- * · 1: Cast IO_T → float → 计算 → Cast float → IO_T(fp16 / bf16 路径)33+ * 在此编译期推导为静态常量,避免 tiling_key 重复编码 dtype 衍生信息:
33- * · 当 T == bfloat16_t 时,硬件不直接支持 bf16 计算,必须设为 134+ * · fp16 / bf16 → NEED_UPCAST=1(Cast IO_T → fp32 计算,与 PyTorch CPU 升精路径对齐)
35+ * · fp32 → NEED_UPCAST=0(直通)
36+ * · 当 T == bfloat16_t 时,硬件不直接支持 bf16 计算,必须升精
34 * · 当 T == half 时,遵循 CANN FP16 算子默认走 FP32 中间计算的约定,37 * · 当 T == half 时,遵循 CANN FP16 算子默认走 FP32 中间计算的约定,
35 * 与 PyTorch CPU fp16 vectorized 路径对齐 (cpu/Activation.cpp:642-663)38 * 与 PyTorch CPU fp16 vectorized 路径对齐 (cpu/Activation.cpp:642-663)
36 */39 */
@@ -46,10 +49,12 @@ namespace NsHardShrink {
46 49 
47using namespace AscendC;50using namespace AscendC;
48 51 
49-template <typename T, int BUFFER_MODE, int NEED_UPCAST>52+template <typename T, int BUFFER_MODE>
50class HardShrink {53class HardShrink {
51 // IO 类型: 由模板参数 T 直接决定(half / float / bfloat16_t)54 // IO 类型: 由模板参数 T 直接决定(half / float / bfloat16_t)
52 using IO_T = T;55 using IO_T = T;
56+ // NEED_UPCAST 由 IO 类型编译期推导:fp16/bf16 升精到 fp32,fp32 直通
57+ static constexpr int NEED_UPCAST = (std::is_same<IO_T, float>::value) ? 0 : 1;
53 // 计算类型: NEED_UPCAST ? float : T58 // 计算类型: NEED_UPCAST ? float : T
54 // · fp16 / bf16 → COMPUTE_T = float(与 PyTorch CPU 升精路径对齐)59 // · fp16 / bf16 → COMPUTE_T = float(与 PyTorch CPU 升精路径对齐)
55 // · fp32 → COMPUTE_T = float(保持不变)60 // · fp32 → COMPUTE_T = float(保持不变)
@@ -84,9 +89,9 @@ private:
84 float lambd_ = 0.5f;89 float lambd_ = 0.5f;
85};90};
86 91 
87-template <typename T, int BUFFER_MODE, int NEED_UPCAST>92+template <typename T, int BUFFER_MODE>
88-__aicore__ inline void HardShrink<T, BUFFER_MODE, NEED_UPCAST>::Init(GM_ADDR self, GM_ADDR out,93+__aicore__ inline void HardShrink<T, BUFFER_MODE>::Init(GM_ADDR self, GM_ADDR out,
89- const HardShrinkTilingData* tilingData)94+ const HardShrinkTilingData* tilingData)
90{95{
91 int64_t remainderLength = tilingData->totalNum - tilingData->blockFactor * AscendC::GetBlockIdx();96 int64_t remainderLength = tilingData->totalNum - tilingData->blockFactor * AscendC::GetBlockIdx();
92 blockLength_ = (remainderLength > tilingData->blockFactor) ? tilingData->blockFactor : remainderLength;97 blockLength_ = (remainderLength > tilingData->blockFactor) ? tilingData->blockFactor : remainderLength;
@@ -101,9 +106,10 @@ __aicore__ inline void HardShrink<T, BUFFER_MODE, NEED_UPCAST>::Init(GM_ADDR sel
101 pipe.InitBuffer(outputQueue, BUFFER_NUM, ubLength_ * sizeof(IO_T));106 pipe.InitBuffer(outputQueue, BUFFER_NUM, ubLength_ * sizeof(IO_T));
102 107 
103 // 初始化计算辅助 buffer(使用计算类型 COMPUTE_T)108 // 初始化计算辅助 buffer(使用计算类型 COMPUTE_T)
109+ // 公式 result = |x| <= lambd ? 0 : x 需要:lambd 常量、全 0 常量、|x| 临时、(升精)floatIn 临时
104 pipe.InitBuffer(lambdBuf, ubLength_ * sizeof(COMPUTE_T));110 pipe.InitBuffer(lambdBuf, ubLength_ * sizeof(COMPUTE_T));
105- pipe.InitBuffer(negLambdBuf, ubLength_ * sizeof(COMPUTE_T));111+ pipe.InitBuffer(negLambdBuf, ubLength_ * sizeof(COMPUTE_T)); // 复用为 zeroBuf:填全 0
106- pipe.InitBuffer(tmpBuf, ubLength_ * sizeof(COMPUTE_T));112+ pipe.InitBuffer(tmpBuf, ubLength_ * sizeof(COMPUTE_T)); // |x| 临时 + Select 结果
107 if constexpr (NEED_UPCAST == 1) {113 if constexpr (NEED_UPCAST == 1) {
108 pipe.InitBuffer(floatInBuf, ubLength_ * sizeof(COMPUTE_T));114 pipe.InitBuffer(floatInBuf, ubLength_ * sizeof(COMPUTE_T));
109 }115 }
@@ -111,15 +117,15 @@ __aicore__ inline void HardShrink<T, BUFFER_MODE, NEED_UPCAST>::Init(GM_ADDR sel
111 int64_t maskBytes = ((ubLength_ / 8) + 31) / 32 * 32;117 int64_t maskBytes = ((ubLength_ / 8) + 31) / 32 * 32;
112 pipe.InitBuffer(cmpMaskBuf, maskBytes);118 pipe.InitBuffer(cmpMaskBuf, maskBytes);
113 119 
114- // 一次性填充 lambd / -lambd 常量120+ // 一次性填充 lambd / 0 常量
115 LocalTensor<COMPUTE_T> lambdLocal = lambdBuf.Get<COMPUTE_T>();121 LocalTensor<COMPUTE_T> lambdLocal = lambdBuf.Get<COMPUTE_T>();
116- LocalTensor<COMPUTE_T> negLambdLocal = negLambdBuf.Get<COMPUTE_T>();122+ LocalTensor<COMPUTE_T> zeroLocal = negLambdBuf.Get<COMPUTE_T>();
117 AscendC::Duplicate(lambdLocal, static_cast<COMPUTE_T>(lambd_), ubLength_);123 AscendC::Duplicate(lambdLocal, static_cast<COMPUTE_T>(lambd_), ubLength_);
118- AscendC::Duplicate(negLambdLocal, static_cast<COMPUTE_T>(-lambd_), ubLength_);124+ AscendC::Duplicate(zeroLocal, static_cast<COMPUTE_T>(0), ubLength_);
119}125}
120 126 
121-template <typename T, int BUFFER_MODE, int NEED_UPCAST>127+template <typename T, int BUFFER_MODE>
122-__aicore__ inline void HardShrink<T, BUFFER_MODE, NEED_UPCAST>::CopyIn(int64_t progress, int64_t currentNum)128+__aicore__ inline void HardShrink<T, BUFFER_MODE>::CopyIn(int64_t progress, int64_t currentNum)
123{129{
124 AscendC::LocalTensor<IO_T> inputLocal = inputQueue.template AllocTensor<IO_T>();130 AscendC::LocalTensor<IO_T> inputLocal = inputQueue.template AllocTensor<IO_T>();
125 AscendC::DataCopyParams copyParams;131 AscendC::DataCopyParams copyParams;
@@ -131,16 +137,16 @@ __aicore__ inline void HardShrink<T, BUFFER_MODE, NEED_UPCAST>::CopyIn(int64_t p
131 inputQueue.EnQue(inputLocal);137 inputQueue.EnQue(inputLocal);
132}138}
133 139 
134-template <typename T, int BUFFER_MODE, int NEED_UPCAST>140+template <typename T, int BUFFER_MODE>
135-__aicore__ inline void HardShrink<T, BUFFER_MODE, NEED_UPCAST>::Compute(int64_t currentNum)141+__aicore__ inline void HardShrink<T, BUFFER_MODE>::Compute(int64_t currentNum)
136{142{
137 AscendC::LocalTensor<IO_T> inputLocal = inputQueue.template DeQue<IO_T>();143 AscendC::LocalTensor<IO_T> inputLocal = inputQueue.template DeQue<IO_T>();
138 AscendC::LocalTensor<IO_T> outputLocal = outputQueue.template AllocTensor<IO_T>();144 AscendC::LocalTensor<IO_T> outputLocal = outputQueue.template AllocTensor<IO_T>();
139 145 
140 // 获取计算辅助 buffer146 // 获取计算辅助 buffer
141 AscendC::LocalTensor<COMPUTE_T> lambdLocal = lambdBuf.Get<COMPUTE_T>();147 AscendC::LocalTensor<COMPUTE_T> lambdLocal = lambdBuf.Get<COMPUTE_T>();
142- AscendC::LocalTensor<COMPUTE_T> negLambdLocal = negLambdBuf.Get<COMPUTE_T>();148+ AscendC::LocalTensor<COMPUTE_T> zeroLocal = negLambdBuf.Get<COMPUTE_T>(); // Init 已填全 0
143- AscendC::LocalTensor<COMPUTE_T> tmpLocal = tmpBuf.Get<COMPUTE_T>();149+ AscendC::LocalTensor<COMPUTE_T> absLocal = tmpBuf.Get<COMPUTE_T>(); // 放 |x|(Select 后可覆盖)
144 AscendC::LocalTensor<uint8_t> maskLocal = cmpMaskBuf.Get<uint8_t>();150 AscendC::LocalTensor<uint8_t> maskLocal = cmpMaskBuf.Get<uint8_t>();
145 151 
146 // 对齐 currentNum 到 256B 边界用于 Compare/Select(硬件要求 count 所占空间 256B 对齐)152 // 对齐 currentNum 到 256B 边界用于 Compare/Select(硬件要求 count 所占空间 256B 对齐)
@@ -160,46 +166,30 @@ __aicore__ inline void HardShrink<T, BUFFER_MODE, NEED_UPCAST>::Compute(int64_t
160 // Cast inputLocal(IO_T) → floatInLocal(float)166 // Cast inputLocal(IO_T) → floatInLocal(float)
161 AscendC::Cast(floatInLocal, inputLocal, AscendC::RoundMode::CAST_NONE, currentNum);167 AscendC::Cast(floatInLocal, inputLocal, AscendC::RoundMode::CAST_NONE, currentNum);
162 168 
163- // Step 1: mask = (x > lambd), tmp = mask ? x : 0169+ // 公式(与竞品 canndev/torch 顺序一致):result = |x| <= lambd ? 0 : x
164- AscendC::Compare(maskLocal, floatInLocal, lambdLocal, AscendC::CMPMODE::GT, alignedNum);170+ // canndev tbe: vabs(x) → vcmpsel(|x|, lambd, 'le', 0, x) (le→0, else→x)
165- AscendC::Select(tmpLocal, maskLocal, floatInLocal, static_cast<COMPUTE_T>(0),171+ // torch: x if |x| > lambd else 0 (边界 ==lambd 归 0)
166- AscendC::SELMODE::VSEL_TENSOR_SCALAR_MODE, alignedNum);172+ // NaN 天然透传:|NaN|=NaN,IEEE754 下 NaN<=lambd 为 false → 走 else 选 x(NaN) 透传
167- 173+ // inf 天然透传:|inf|=inf,inf<=lambd 为 false → 走 else 选 inf 透传
168- // Step 2: mask = (x < -lambd), tmp = mask ? x : tmp174+ AscendC::Abs(absLocal, floatInLocal, currentNum);
169- AscendC::Compare(maskLocal, floatInLocal, negLambdLocal, AscendC::CMPMODE::LT, alignedNum);175+ AscendC::Compare(maskLocal, absLocal, lambdLocal, AscendC::CMPMODE::LE, alignedNum);
170- AscendC::Select(tmpLocal, maskLocal, floatInLocal, tmpLocal, AscendC::SELMODE::VSEL_TENSOR_TENSOR_MODE,176+ AscendC::Select(absLocal, maskLocal, zeroLocal, floatInLocal, AscendC::SELMODE::VSEL_TENSOR_TENSOR_MODE,
171- alignedNum);
172- 
173- // Step 3: NaN 透传 —— mask = (x != x),仅 NaN 满足
174- // tmp = mask ? x : tmp
175- AscendC::Compare(maskLocal, floatInLocal, floatInLocal, AscendC::CMPMODE::NE, alignedNum);
176- AscendC::Select(tmpLocal, maskLocal, floatInLocal, tmpLocal, AscendC::SELMODE::VSEL_TENSOR_TENSOR_MODE,
177 alignedNum);177 alignedNum);
178 178 
179 // Cast fp32 result → IO_T output179 // Cast fp32 result → IO_T output
180 // 对 bfloat16_t 使用 CAST_RINT(round-to-nearest-even),对 half 使用 CAST_NONE(默认 RTNE)180 // 对 bfloat16_t 使用 CAST_RINT(round-to-nearest-even),对 half 使用 CAST_NONE(默认 RTNE)
181 if constexpr (std::is_same<IO_T, bfloat16_t>::value) {181 if constexpr (std::is_same<IO_T, bfloat16_t>::value) {
182- AscendC::Cast(outputLocal, tmpLocal, AscendC::RoundMode::CAST_RINT, currentNum);182+ AscendC::Cast(outputLocal, absLocal, AscendC::RoundMode::CAST_RINT, currentNum);
183 } else {183 } else {
184- AscendC::Cast(outputLocal, tmpLocal, AscendC::RoundMode::CAST_NONE, currentNum);184+ AscendC::Cast(outputLocal, absLocal, AscendC::RoundMode::CAST_NONE, currentNum);
185 }185 }
186 } else {186 } else {
187 // 直通路径(fp32):COMPUTE_T = T = IO_T,inputLocal 和 outputLocal 直接参与计算187 // 直通路径(fp32):COMPUTE_T = T = IO_T,inputLocal 和 outputLocal 直接参与计算
188 188 
189- // Step 1: mask = (x > lambd), tmp = mask ? x : 0189+ // 公式(与竞品 canndev/torch 顺序一致):result = |x| <= lambd ? 0 : x
190- AscendC::Compare(maskLocal, inputLocal, lambdLocal, AscendC::CMPMODE::GT, alignedNum);190+ AscendC::Abs(absLocal, inputLocal, currentNum);
191- AscendC::Select(tmpLocal, maskLocal, inputLocal, static_cast<COMPUTE_T>(0),191+ AscendC::Compare(maskLocal, absLocal, lambdLocal, AscendC::CMPMODE::LE, alignedNum);
192- AscendC::SELMODE::VSEL_TENSOR_SCALAR_MODE, alignedNum);192+ AscendC::Select(outputLocal, maskLocal, zeroLocal, inputLocal, AscendC::SELMODE::VSEL_TENSOR_TENSOR_MODE,
193- 
194- // Step 2: mask = (x < -lambd), tmp = mask ? x : tmp
195- AscendC::Compare(maskLocal, inputLocal, negLambdLocal, AscendC::CMPMODE::LT, alignedNum);
196- AscendC::Select(tmpLocal, maskLocal, inputLocal, tmpLocal, AscendC::SELMODE::VSEL_TENSOR_TENSOR_MODE,
197- alignedNum);
198- 
199- // Step 3: NaN 透传 —— mask = (x != x),仅 NaN 满足
200- // out = mask ? x : tmp
201- AscendC::Compare(maskLocal, inputLocal, inputLocal, AscendC::CMPMODE::NE, alignedNum);
202- AscendC::Select(outputLocal, maskLocal, inputLocal, tmpLocal, AscendC::SELMODE::VSEL_TENSOR_TENSOR_MODE,
203 alignedNum);193 alignedNum);
204 }194 }
205 195 
@@ -207,8 +197,8 @@ __aicore__ inline void HardShrink<T, BUFFER_MODE, NEED_UPCAST>::Compute(int64_t
207 inputQueue.FreeTensor(inputLocal);197 inputQueue.FreeTensor(inputLocal);
208}198}
209 199 
210-template <typename T, int BUFFER_MODE, int NEED_UPCAST>200+template <typename T, int BUFFER_MODE>
211-__aicore__ inline void HardShrink<T, BUFFER_MODE, NEED_UPCAST>::CopyOut(int64_t progress, int64_t currentNum)201+__aicore__ inline void HardShrink<T, BUFFER_MODE>::CopyOut(int64_t progress, int64_t currentNum)
212{202{
213 AscendC::LocalTensor<IO_T> outputLocal = outputQueue.template DeQue<IO_T>();203 AscendC::LocalTensor<IO_T> outputLocal = outputQueue.template DeQue<IO_T>();
214 AscendC::DataCopyParams copyParams;204 AscendC::DataCopyParams copyParams;
@@ -220,8 +210,8 @@ __aicore__ inline void HardShrink<T, BUFFER_MODE, NEED_UPCAST>::CopyOut(int64_t
220 outputQueue.FreeTensor(outputLocal);210 outputQueue.FreeTensor(outputLocal);
221}211}
222 212 
223-template <typename T, int BUFFER_MODE, int NEED_UPCAST>213+template <typename T, int BUFFER_MODE>
224-__aicore__ inline void HardShrink<T, BUFFER_MODE, NEED_UPCAST>::Process()214+__aicore__ inline void HardShrink<T, BUFFER_MODE>::Process()
225{215{
226 if (blockLength_ <= 0) {216 if (blockLength_ <= 0) {
227 return; // 空 Tensor 或当前核无任务217 return; // 空 Tensor 或当前核无任务
@@ -14,10 +14,13 @@
14 * \file hard_shrink_tiling_key.h14 * \file hard_shrink_tiling_key.h
15 * \brief HardShrink TilingKey 模板参数定义15 * \brief HardShrink TilingKey 模板参数定义
16 *16 *
17+ * def 驱动 dtype 模式:dtype 由 _def.cpp 的 DataType 列表经构建系统注入
18+ * DTYPE_SELF 编译宏覆盖,tiling_key 只编码 def 未覆盖的维度(调度/缓冲模式)。
19+ * 因此这里不再重复编码 D_T,也不再编码由 dtype 唯一决定的 NEED_UPCAST
20+ *(后者在 kernel 内由 DTYPE_SELF 编译期推导)。
21+ *
17 * 模板参数:22 * 模板参数:
18- * - D_T: IO 数据类型 (half/float/bfloat16_t)
19 * - BUFFER_MODE: 缓冲模式 (0=单缓冲, 1=双缓冲)23 * - BUFFER_MODE: 缓冲模式 (0=单缓冲, 1=双缓冲)
20- * - NEED_UPCAST: 是否升精到 fp32 计算(0=fp32 直通, 1=fp16/bf16 升精到 fp32)
21 */24 */
22 25 
23#ifndef __HARD_SHRINK_TILING_KEY_H__26#ifndef __HARD_SHRINK_TILING_KEY_H__
@@ -25,35 +28,12 @@
25 28 
26#include "ascendc/host_api/tiling/template_argument.h"29#include "ascendc/host_api/tiling/template_argument.h"
27 30 
28-ASCENDC_TPL_ARGS_DECL(HardShrink,31+ASCENDC_TPL_ARGS_DECL(HardShrink, ASCENDC_TPL_UINT_DECL(BUFFER_MODE, 8, ASCENDC_TPL_UI_LIST, 0, 1));
29- ASCENDC_TPL_DATATYPE_DECL(D_T, C_DT_FLOAT16, C_DT_FLOAT, C_DT_BF16, ASCENDC_TPL_INPUT(0)),
30- ASCENDC_TPL_UINT_DECL(BUFFER_MODE, 8, ASCENDC_TPL_UI_LIST, 0, 1),
31- ASCENDC_TPL_UINT_DECL(NEED_UPCAST, 8, ASCENDC_TPL_UI_LIST, 0, 1));
32 32 
33ASCENDC_TPL_SEL(33ASCENDC_TPL_SEL(
34- // fp16 单缓冲(NEED_UPCAST=1,FP16 走 FP32 中间计算)34+ // 单缓冲
35- ASCENDC_TPL_ARGS_SEL(ASCENDC_TPL_DATATYPE_SEL(D_T, C_DT_FLOAT16),35+ ASCENDC_TPL_ARGS_SEL(ASCENDC_TPL_UINT_SEL(BUFFER_MODE, ASCENDC_TPL_UI_LIST, 0)),
36- ASCENDC_TPL_UINT_SEL(BUFFER_MODE, ASCENDC_TPL_UI_LIST, 0),36+ // 双缓冲
37- ASCENDC_TPL_UINT_SEL(NEED_UPCAST, ASCENDC_TPL_UI_LIST, 1)),37+ ASCENDC_TPL_ARGS_SEL(ASCENDC_TPL_UINT_SEL(BUFFER_MODE, ASCENDC_TPL_UI_LIST, 1)), );
38- // fp16 双缓冲
39- ASCENDC_TPL_ARGS_SEL(ASCENDC_TPL_DATATYPE_SEL(D_T, C_DT_FLOAT16),
40- ASCENDC_TPL_UINT_SEL(BUFFER_MODE, ASCENDC_TPL_UI_LIST, 1),
41- ASCENDC_TPL_UINT_SEL(NEED_UPCAST, ASCENDC_TPL_UI_LIST, 1)),
42- // fp32 单缓冲(NEED_UPCAST=0,直通)
43- ASCENDC_TPL_ARGS_SEL(ASCENDC_TPL_DATATYPE_SEL(D_T, C_DT_FLOAT),
44- ASCENDC_TPL_UINT_SEL(BUFFER_MODE, ASCENDC_TPL_UI_LIST, 0),
45- ASCENDC_TPL_UINT_SEL(NEED_UPCAST, ASCENDC_TPL_UI_LIST, 0)),
46- // fp32 双缓冲
47- ASCENDC_TPL_ARGS_SEL(ASCENDC_TPL_DATATYPE_SEL(D_T, C_DT_FLOAT),
48- ASCENDC_TPL_UINT_SEL(BUFFER_MODE, ASCENDC_TPL_UI_LIST, 1),
49- ASCENDC_TPL_UINT_SEL(NEED_UPCAST, ASCENDC_TPL_UI_LIST, 0)),
50- // bf16 单缓冲(NEED_UPCAST=1,bf16 → fp32 计算)
51- ASCENDC_TPL_ARGS_SEL(ASCENDC_TPL_DATATYPE_SEL(D_T, C_DT_BF16),
52- ASCENDC_TPL_UINT_SEL(BUFFER_MODE, ASCENDC_TPL_UI_LIST, 0),
53- ASCENDC_TPL_UINT_SEL(NEED_UPCAST, ASCENDC_TPL_UI_LIST, 1)),
54- // bf16 双缓冲
55- ASCENDC_TPL_ARGS_SEL(ASCENDC_TPL_DATATYPE_SEL(D_T, C_DT_BF16),
56- ASCENDC_TPL_UINT_SEL(BUFFER_MODE, ASCENDC_TPL_UI_LIST, 1),
57- ASCENDC_TPL_UINT_SEL(NEED_UPCAST, ASCENDC_TPL_UI_LIST, 1)), );
58 38 
59#endif39#endif
@@ -14,23 +14,24 @@
14 * \file hard_shrink.cpp14 * \file hard_shrink.cpp
15 * \brief HardShrink 算子 Kernel 入口(arch35 架构,Ascend950)15 * \brief HardShrink 算子 Kernel 入口(arch35 架构,Ascend950)
16 *16 *
17+ * def 驱动 dtype 模式:dtype 由 _def.cpp 的 DataType 列表经构建系统注入
18+ * DTYPE_SELF 编译宏,kernel 入口直接使用该宏作为 IO 类型,无需 tiling_key 编码。
19+ *
17 * 模板参数(与 hard_shrink_tiling_key.h 中 ASCENDC_TPL_ARGS_DECL 定义对应):20 * 模板参数(与 hard_shrink_tiling_key.h 中 ASCENDC_TPL_ARGS_DECL 定义对应):
18- * - D_T: IO 数据类型 (half/float/bfloat16_t)
19 * - BUFFER_MODE: 缓冲模式 (0=单缓冲, 1=双缓冲)21 * - BUFFER_MODE: 缓冲模式 (0=单缓冲, 1=双缓冲)
20- * - NEED_UPCAST: 是否升精到 fp32 计算 (0=fp32 直通, 1=fp16/bf16 升精)
21 *22 *
22 * 注:lambd 作为 Attr 通过 TilingData 传递,不作为 kernel 参数23 * 注:lambd 作为 Attr 通过 TilingData 传递,不作为 kernel 参数
23 */24 */
24 25 
25#include "arch35/hard_shrink.h"26#include "arch35/hard_shrink.h"
26 27 
27-template <typename D_T, int BUFFER_MODE, int NEED_UPCAST>28+template <int BUFFER_MODE>
28__global__ __aicore__ void hard_shrink(GM_ADDR self, GM_ADDR out, GM_ADDR workspace, GM_ADDR tiling)29__global__ __aicore__ void hard_shrink(GM_ADDR self, GM_ADDR out, GM_ADDR workspace, GM_ADDR tiling)
29{30{
30 REGISTER_TILING_DEFAULT(HardShrinkTilingData);31 REGISTER_TILING_DEFAULT(HardShrinkTilingData);
31 GET_TILING_DATA_WITH_STRUCT(HardShrinkTilingData, tilingData, tiling);32 GET_TILING_DATA_WITH_STRUCT(HardShrinkTilingData, tilingData, tiling);
32 33 
33- NsHardShrink::HardShrink<D_T, BUFFER_MODE, NEED_UPCAST> op;34+ NsHardShrink::HardShrink<DTYPE_SELF, BUFFER_MODE> op;
34 op.Init(self, out, &tilingData);35 op.Init(self, out, &tilingData);
35 op.Process();36 op.Process();
36}37}