已合并
feat(hard_shrink): def 驱动 dtype + abs/le 公式对齐竞品 #8756
leaving__创建于 8月17日
feat(hard_shrink): def 驱动 dtype + abs/le 公式对齐竞品 #8756
已合并
共 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 | + | ||
| 22 | + | ||
| 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 / coreNum | 18 | * - 多核切分:totalNum / coreNum |
| 19 | * - UB 切分:ubSize / bufferNum / typeSize,256B 对齐 | 19 | * - UB 切分:ubSize / bufferNum / typeSize,256B 对齐 |
| 20 | * - 双缓冲阈值:totalLength > 1024 → BUFFER_MODE=1 | 20 | * - 双缓冲阈值: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 | 26 | ||
| @@ -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.cpp | 14 | * \file hard_shrink_infershape.cpp |
| 15 | * \brief HardShrink 算子形状推导实现 | 15 | * \brief HardShrink 算子形状推导实现 |
| 16 | * | 16 | * |
| 17 | - * HardShrink 是逐元素算子:输出 shape = 输入 shape,输出 dtype = 输入 dtype | 17 | + * 复用 canndev 仓 runtime2.0 的逐元素算子公共推导(InferShape4Elewise): |
| 18 | + * 输出 shape = 输入 shape。canndev 仓原 InferShape 逻辑保留不删,长尾算子仍使用。 | ||
| 18 | */ | 19 | */ |
| 19 | 20 | ||
| 21 | + | ||
| 20 | 22 | ||
| 21 | -#include "exe_graph/runtime/infer_shape_context.h" | 23 | +#include "log/log.h" |
| 22 | 24 | ||
| 23 | using namespace ge; | 25 | using namespace ge; |
| 24 | - | ||
| 25 | namespace ops { | 26 | namespace ops { |
| 26 | - | ||
| 27 | static ge::graphStatus InferShape4HardShrink(gert::InferShapeContext* context) | 27 | static 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 | ||
| 47 | IMPL_OP_INFERSHAPE(HardShrink).InferShape(InferShape4HardShrink); | 32 | IMPL_OP_INFERSHAPE(HardShrink).InferShape(InferShape4HardShrink); |
| 48 | - | ||
| 49 | } // namespace ops | 33 | } // namespace ops |
| @@ -14,23 +14,26 @@ | |||
| 14 | * \file hard_shrink.h | 14 | * \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 0 | 17 | + * 计算公式: 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 : 0 | 20 | + * 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 : tmp1 | 21 | + * torch: x if |x| > lambd else 0 |
| 22 | - * Step 3: mask_nan = (x != x), out = mask_nan ? x : tmp2 | 22 | + * Step 1: absX = |x| (AscendC::Abs) |
| 23 | - * 说明:IEEE 754 下 NaN 满足 (NaN != NaN)==true,借此识别并透传 NaN | 23 | + * 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 计算,必须设为 1 | 34 | + * · 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 | ||
| 47 | using namespace AscendC; | 50 | using namespace AscendC; |
| 48 | 51 | ||
| 49 | -template <typename T, int BUFFER_MODE, int NEED_UPCAST> | 52 | +template <typename T, int BUFFER_MODE> |
| 50 | class HardShrink { | 53 | class 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 : T | 58 | // 计算类型: 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 | // 获取计算辅助 buffer | 146 | // 获取计算辅助 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 : 0 | 169 | + // 公式(与竞品 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 : tmp | 174 | + 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 output | 179 | // 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 : 0 | 189 | + // 公式(与竞品 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.h | 14 | * \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 | 26 | ||
| @@ -25,35 +28,12 @@ | |||
| 25 | 28 | ||
| 26 | 29 | ||
| 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 | ||
| 33 | ASCENDC_TPL_SEL( | 33 | ASCENDC_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 | 39 | ||
| @@ -14,23 +14,24 @@ | |||
| 14 | * \file hard_shrink.cpp | 14 | * \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 | 26 | ||
| 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 | } |