已合并
feat: 新增AscendQuantV2/AscendAntiQuantV2/DynamicQuant/EmbeddingDenseGrad算子host fallback支持 #8392
杨金翰50065292创建于 8月7日
feat: 新增AscendQuantV2/AscendAntiQuantV2/DynamicQuant/EmbeddingDenseGrad算子host fallback支持 #8392
已合并
杨金翰50065292创建于 8月7日
共 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+#include "op_fallback.h"
12+ 
13+#ifdef __cplusplus
14+extern "C" {
15+#endif
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+#ifdef __cplusplus
63+}
64+#endif
@@ -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+#include "op_fallback.h"
12+ 
13+#ifdef __cplusplus
14+extern "C" {
15+#endif
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+#ifdef __cplusplus
63+}
64+#endif
@@ -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+#include "op_fallback.h"
12+ 
13+#ifdef __cplusplus
14+extern "C" {
15+#endif
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+#ifdef __cplusplus
65+}
66+#endif
@@ -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+#include "op_fallback.h"
12+ 
13+#ifdef __cplusplus
14+extern "C" {
15+#endif
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+#ifdef __cplusplus
58+}
59+#endif