已合并
feat: 新增AscendQuantV2/AscendAntiQuantV2/DynamicQuant/EmbeddingDenseGrad算子host fallback支持 #8392
杨金翰50065292创建于 8月7日
feat: 新增AscendQuantV2/AscendAntiQuantV2/DynamicQuant/EmbeddingDenseGrad算子host fallback支持 #8392
已合并
共 4 个文件变更+253-0
| @@ -0,0 +1,64 @@ | |||
| 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 | + | ||
| 12 | + | ||
| 13 | + | ||
| 14 | +extern "C" { | ||
| 15 | + | ||
| 16 | +namespace fallback { | ||
| 17 | + | ||
| 18 | +using namespace ge; | ||
| 19 | +using namespace gert; | ||
| 20 | +const static size_t EMBEDDING_DENSE_GRAD_IN_GRAD = 0; | ||
| 21 | +const static size_t EMBEDDING_DENSE_GRAD_IN_INDICES = 1; | ||
| 22 | +const static size_t EMBEDDING_DENSE_GRAD_OUT_Y = 0; | ||
| 23 | +const static size_t EMBEDDING_DENSE_GRAD_ATTR_NUM_WEIGHTS = 0; | ||
| 24 | +const static size_t EMBEDDING_DENSE_GRAD_ATTR_PADDING_IDX = 1; | ||
| 25 | +const static size_t EMBEDDING_DENSE_GRAD_ATTR_SCALE_GRAD_BY_FREQ = 2; | ||
| 26 | + | ||
| 27 | +graphStatus EmbeddingDenseGradHostExecuteFunc(OpExecuteContext* host_api_ctx) | ||
| 28 | +{ | ||
| 29 | + OP_CHECK_IF(host_api_ctx == nullptr, OP_LOGE("aclnnfallback", "host_api_ctx is null"), return GRAPH_FAILED); | ||
| 30 | + | ||
| 31 | + auto inputGrad = host_api_ctx->GetInputTensor(EMBEDDING_DENSE_GRAD_IN_GRAD); | ||
| 32 | + OP_CHECK_IF(inputGrad == nullptr, OP_LOGE("aclnnfallback", "grad is null"), return GRAPH_FAILED); | ||
| 33 | + | ||
| 34 | + auto indices = host_api_ctx->GetInputTensor(EMBEDDING_DENSE_GRAD_IN_INDICES); | ||
| 35 | + OP_CHECK_IF(indices == nullptr, OP_LOGE("aclnnfallback", "indices is null"), return GRAPH_FAILED); | ||
| 36 | + | ||
| 37 | + auto output = host_api_ctx->GetOutputTensor(EMBEDDING_DENSE_GRAD_OUT_Y); | ||
| 38 | + OP_CHECK_IF(output == nullptr, OP_LOGE("aclnnfallback", "output is null"), return GRAPH_FAILED); | ||
| 39 | + | ||
| 40 | + auto attrs = host_api_ctx->GetAttrs(); | ||
| 41 | + OP_CHECK_IF(attrs == nullptr, OP_LOGE("aclnnfallback", "attrs is null"), return GRAPH_FAILED); | ||
| 42 | + | ||
| 43 | + const uint64_t* numWeight = attrs->GetAttrPointer<uint64_t>(EMBEDDING_DENSE_GRAD_ATTR_NUM_WEIGHTS); | ||
| 44 | + OP_CHECK_IF(numWeight == nullptr, OP_LOGE("aclnnfallback", "numWeight is null"), return GRAPH_FAILED); | ||
| 45 | + const uint64_t* paddingIdx = attrs->GetAttrPointer<uint64_t>(EMBEDDING_DENSE_GRAD_ATTR_PADDING_IDX); | ||
| 46 | + OP_CHECK_IF(paddingIdx == nullptr, OP_LOGE("aclnnfallback", "paddingIdx is null"), return GRAPH_FAILED); | ||
| 47 | + const bool* scaleGrad = attrs->GetAttrPointer<bool>(EMBEDDING_DENSE_GRAD_ATTR_SCALE_GRAD_BY_FREQ); | ||
| 48 | + OP_CHECK_IF(scaleGrad == nullptr, OP_LOGE("aclnnfallback", "scaleGrad is null"), return GRAPH_FAILED); | ||
| 49 | + | ||
| 50 | + OP_LOGD("aclnnFallback", "EmbeddingDenseGrad fallback begin"); | ||
| 51 | + auto api_ret = CANN_OPS_OPB_SYN_EXEC_ACLNN(host_api_ctx, aclnnEmbeddingDenseBackward, inputGrad, indices, | ||
| 52 | + *numWeight, *paddingIdx, *scaleGrad, output); | ||
| 53 | + | ||
| 54 | + OP_CHECK_IF(api_ret != GRAPH_SUCCESS, OP_LOGE("aclnnfallback", "api_ret faild:%d", api_ret), return GRAPH_FAILED); | ||
| 55 | + | ||
| 56 | + return GRAPH_SUCCESS; | ||
| 57 | +} | ||
| 58 | + | ||
| 59 | +IMPL_OP(EmbeddingDenseGrad).OpExecuteFunc(EmbeddingDenseGradHostExecuteFunc); | ||
| 60 | +} // namespace fallback | ||
| 61 | + | ||
| 62 | + | ||
| 63 | +} | ||
| 64 | + | ||
| @@ -0,0 +1,64 @@ | |||
| 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 | + | ||
| 12 | + | ||
| 13 | + | ||
| 14 | +extern "C" { | ||
| 15 | + | ||
| 16 | +namespace fallback { | ||
| 17 | + | ||
| 18 | +using namespace ge; | ||
| 19 | +using namespace gert; | ||
| 20 | +static const size_t INPUT_X_INDEX = 0; | ||
| 21 | +static const size_t SCALE_INDEX = 1; | ||
| 22 | +static const size_t OFFSET_INDEX = 2; | ||
| 23 | +static const size_t ATTR_DSTDATATYPE_INDEX = 0; | ||
| 24 | +static const size_t ATTR_SQRT_MODE_INDEX = 1; | ||
| 25 | + | ||
| 26 | +static graphStatus AntiQuantHostExecuteFunc(OpExecuteContext* host_api_ctx) | ||
| 27 | +{ | ||
| 28 | + OP_CHECK_IF(host_api_ctx == nullptr, OP_LOGE("aclnnfallback", "host_api_ctx is null"), return GRAPH_FAILED); | ||
| 29 | + | ||
| 30 | + auto input_x = host_api_ctx->GetInputTensor(INPUT_X_INDEX); | ||
| 31 | + OP_CHECK_IF(input_x == nullptr, OP_LOGE("aclnnfallback", "input_x is null"), return GRAPH_FAILED); | ||
| 32 | + | ||
| 33 | + auto scale = host_api_ctx->GetInputTensor(SCALE_INDEX); | ||
| 34 | + OP_CHECK_IF(scale == nullptr, OP_LOGE("aclnnfallback", "scale is null"), return GRAPH_FAILED); | ||
| 35 | + | ||
| 36 | + auto offset = host_api_ctx->GetOptionalInputTensor(OFFSET_INDEX); | ||
| 37 | + | ||
| 38 | + auto output = host_api_ctx->GetOutputTensor(0); | ||
| 39 | + OP_CHECK_IF(output == nullptr, OP_LOGE("aclnnfallback", "output is null"), return GRAPH_FAILED); | ||
| 40 | + | ||
| 41 | + auto attrs = host_api_ctx->GetAttrs(); | ||
| 42 | + OP_CHECK_IF(attrs == nullptr, OP_LOGE("aclnnfallback", "attrs is null"), return GRAPH_FAILED); | ||
| 43 | + | ||
| 44 | + const int64_t* dstDtype = attrs->GetAttrPointer<int64_t>(ATTR_DSTDATATYPE_INDEX); | ||
| 45 | + OP_CHECK_IF(dstDtype == nullptr, OP_LOGE("aclnnfallback", "dstDtype is null"), return GRAPH_FAILED); | ||
| 46 | + const bool* sqrt_mode = attrs->GetAttrPointer<bool>(ATTR_SQRT_MODE_INDEX); | ||
| 47 | + OP_CHECK_IF(sqrt_mode == nullptr, OP_LOGE("aclnnfallback", "sqrt_mode is null"), return GRAPH_FAILED); | ||
| 48 | + | ||
| 49 | + OP_LOGD("aclnnFallback", "AscendAntiQuant fallback begin"); | ||
| 50 | + | ||
| 51 | + auto api_ret = CANN_OPS_OPB_SYN_EXEC_ACLNN(host_api_ctx, aclnnAscendAntiQuant, input_x, scale, offset, *dstDtype, | ||
| 52 | + *sqrt_mode, output); | ||
| 53 | + | ||
| 54 | + OP_CHECK_IF(api_ret != GRAPH_SUCCESS, OP_LOGE("aclnnfallback", "api_ret faild:%d", api_ret), return GRAPH_FAILED); | ||
| 55 | + | ||
| 56 | + return GRAPH_SUCCESS; | ||
| 57 | +} | ||
| 58 | + | ||
| 59 | +IMPL_OP(AscendAntiQuantV2).OpExecuteFunc(AntiQuantHostExecuteFunc); | ||
| 60 | +} // namespace fallback | ||
| 61 | + | ||
| 62 | + | ||
| 63 | +} | ||
| 64 | + | ||
| @@ -0,0 +1,66 @@ | |||
| 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 | + | ||
| 12 | + | ||
| 13 | + | ||
| 14 | +extern "C" { | ||
| 15 | + | ||
| 16 | +namespace fallback { | ||
| 17 | + | ||
| 18 | +using namespace ge; | ||
| 19 | +using namespace gert; | ||
| 20 | +static const size_t INPUT_X_INDEX = 0; | ||
| 21 | +static const size_t SCALE_INDEX = 1; | ||
| 22 | +static const size_t OFFSET_INDEX = 2; | ||
| 23 | +static const size_t ATTR_SQRT_MODE_INDEX = 0; | ||
| 24 | +static const size_t ATTR_ROUND_MODE_INDEX = 1; | ||
| 25 | +static const size_t ATTR_DST_DATATYPE_INDEX = 2; | ||
| 26 | + | ||
| 27 | +static graphStatus AscendQuantHostExecuteFunc(OpExecuteContext* host_api_ctx) | ||
| 28 | +{ | ||
| 29 | + OP_CHECK_IF(host_api_ctx == nullptr, OP_LOGE("aclnnfallback", "host_api_ctx is null"), return GRAPH_FAILED); | ||
| 30 | + | ||
| 31 | + auto inputX = host_api_ctx->GetInputTensor(INPUT_X_INDEX); | ||
| 32 | + OP_CHECK_IF(inputX == nullptr, OP_LOGE("aclnnfallback", "input_x is null"), return GRAPH_FAILED); | ||
| 33 | + | ||
| 34 | + auto scale = host_api_ctx->GetInputTensor(SCALE_INDEX); | ||
| 35 | + OP_CHECK_IF(scale == nullptr, OP_LOGE("aclnnfallback", "scale is null"), return GRAPH_FAILED); | ||
| 36 | + | ||
| 37 | + auto offset = host_api_ctx->GetOptionalInputTensor(OFFSET_INDEX); | ||
| 38 | + | ||
| 39 | + auto output = host_api_ctx->GetOutputTensor(0); | ||
| 40 | + OP_CHECK_IF(output == nullptr, OP_LOGE("aclnnfallback", "output is null"), return GRAPH_FAILED); | ||
| 41 | + | ||
| 42 | + auto attrs = host_api_ctx->GetAttrs(); | ||
| 43 | + OP_CHECK_IF(attrs == nullptr, OP_LOGE("aclnnfallback", "attrs is null"), return GRAPH_FAILED); | ||
| 44 | + | ||
| 45 | + const bool* sqrtMode = attrs->GetAttrPointer<bool>(ATTR_SQRT_MODE_INDEX); | ||
| 46 | + OP_CHECK_IF(sqrtMode == nullptr, OP_LOGE("aclnnfallback", "sqrtMode is null"), return GRAPH_FAILED); | ||
| 47 | + const char* roundMode = attrs->GetAttrPointer<char>(ATTR_ROUND_MODE_INDEX); | ||
| 48 | + OP_CHECK_IF(roundMode == nullptr, OP_LOGE("aclnnfallback", "roundMode is null"), return GRAPH_FAILED); | ||
| 49 | + const int32_t* dstDtype = attrs->GetAttrPointer<int32_t>(ATTR_DST_DATATYPE_INDEX); | ||
| 50 | + OP_CHECK_IF(dstDtype == nullptr, OP_LOGE("aclnnfallback", "dstDtype is null"), return GRAPH_FAILED); | ||
| 51 | + | ||
| 52 | + OP_LOGD("aclnnFallback", "AscendQuantV2 fallback begin"); | ||
| 53 | + auto api_ret = CANN_OPS_OPB_SYN_EXEC_ACLNN(host_api_ctx, aclnnAscendQuant, inputX, scale, offset, *sqrtMode, | ||
| 54 | + roundMode, *dstDtype, output); | ||
| 55 | + | ||
| 56 | + OP_CHECK_IF(api_ret != GRAPH_SUCCESS, OP_LOGE("aclnnfallback", "api_ret faild:%d", api_ret), return GRAPH_FAILED); | ||
| 57 | + | ||
| 58 | + return GRAPH_SUCCESS; | ||
| 59 | +} | ||
| 60 | + | ||
| 61 | +IMPL_OP(AscendQuantV2).OpExecuteFunc(AscendQuantHostExecuteFunc); | ||
| 62 | +} // namespace fallback | ||
| 63 | + | ||
| 64 | + | ||
| 65 | +} | ||
| 66 | + | ||
| @@ -0,0 +1,59 @@ | |||
| 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 | + | ||
| 12 | + | ||
| 13 | + | ||
| 14 | +extern "C" { | ||
| 15 | + | ||
| 16 | +namespace fallback { | ||
| 17 | + | ||
| 18 | +using namespace ge; | ||
| 19 | +using namespace gert; | ||
| 20 | +constexpr size_t INPUT_X_INDEX = 0; | ||
| 21 | +constexpr size_t INPUT_SMOOTH_INDEX = 1; | ||
| 22 | +constexpr size_t INPUT_GROUP_INDEX = 2; | ||
| 23 | +constexpr size_t OUTPUT_Y_INDEX = 0; | ||
| 24 | +constexpr size_t OUTPUT_SCALE_INDEX = 1; | ||
| 25 | + | ||
| 26 | +static graphStatus DynamicQuantExecuteFunc(OpExecuteContext* host_api_ctx) | ||
| 27 | +{ | ||
| 28 | + OP_CHECK_IF(host_api_ctx == nullptr, OP_LOGE("fallback_dynamic_quant", "host_api_ctx is null"), | ||
| 29 | + return GRAPH_FAILED); | ||
| 30 | + OP_LOGD(host_api_ctx->GetNodeName(), "Enter DynamicQuantExecuteFunc."); | ||
| 31 | + | ||
| 32 | + auto x = host_api_ctx->GetInputTensor(INPUT_X_INDEX); | ||
| 33 | + OP_CHECK_IF(x == nullptr, OP_LOGE(host_api_ctx->GetNodeName(), "x is null"), return GRAPH_FAILED); | ||
| 34 | + | ||
| 35 | + auto y = host_api_ctx->GetOutputTensor(OUTPUT_Y_INDEX); | ||
| 36 | + OP_CHECK_IF(y == nullptr, OP_LOGE(host_api_ctx->GetNodeName(), "y is null"), return GRAPH_FAILED); | ||
| 37 | + | ||
| 38 | + auto scale = host_api_ctx->GetOutputTensor(OUTPUT_SCALE_INDEX); | ||
| 39 | + OP_CHECK_IF(scale == nullptr, OP_LOGE(host_api_ctx->GetNodeName(), "scale is null"), return GRAPH_FAILED); | ||
| 40 | + | ||
| 41 | + auto smooth_scales = host_api_ctx->GetOptionalInputTensor(INPUT_SMOOTH_INDEX); | ||
| 42 | + | ||
| 43 | + auto group_index = host_api_ctx->GetOptionalInputTensor(INPUT_GROUP_INDEX); | ||
| 44 | + | ||
| 45 | + // execute opapi | ||
| 46 | + auto api_ret = CANN_OPS_OPB_SYN_EXEC_ACLNN(host_api_ctx, aclnnDynamicQuantV2, x, smooth_scales, group_index, y, | ||
| 47 | + scale); | ||
| 48 | + OP_CHECK_IF(api_ret != GRAPH_SUCCESS, OP_LOGE(host_api_ctx->GetNodeName(), "api_ret faild:%d", api_ret), | ||
| 49 | + return GRAPH_FAILED); | ||
| 50 | + | ||
| 51 | + return GRAPH_SUCCESS; | ||
| 52 | +} | ||
| 53 | + | ||
| 54 | +IMPL_OP(DynamicQuant).OpExecuteFunc(DynamicQuantExecuteFunc); | ||
| 55 | +} // namespace fallback | ||
| 56 | + | ||
| 57 | + | ||
| 58 | +} | ||
| 59 | + | ||