已合并
gmm support ascend350 #10843
gmm support ascend350 #10843
已合并
马琦钧创建于 8 天前
9 个文件变更+170-166
@@ -24,6 +24,12 @@ if (BUILD_OPEN_PROJECT)
24 COMPUTE_UNIT Ascend950PR_959924 COMPUTE_UNIT Ascend950PR_9599
25 OPTIONS -mllvm -cce-aicore-dcci-before-kernel-end=false25 OPTIONS -mllvm -cce-aicore-dcci-before-kernel-end=false
26 )26 )
27+ elseif("${ASCEND_COMPUTE_UNIT}" STREQUAL "ascend350")
28+ add_ops_compile_options(
29+ OP_NAME GroupedMatmul
30+ COMPUTE_UNIT Ascend350_355e
31+ OPTIONS -mllvm -cce-aicore-dcci-before-kernel-end=false
32+ )
27 endif()33 endif()
28endif()34endif()
29 35 
@@ -35,6 +41,13 @@ if("${CONDITION_UNIT}" STREQUAL "ascend950")
35 COMPUTE_UNIT Ascend950PR_959941 COMPUTE_UNIT Ascend950PR_9599
36 OPTIONS -DENABLE_CV_COMM_VIA_SSBUF=true42 OPTIONS -DENABLE_CV_COMM_VIA_SSBUF=true
37 )43 )
44+elseif("${CONDITION_UNIT}" STREQUAL "ascend350")
45+ set(grouped_matmul_depends gmm/common CACHE INTERNAL "Dependencies for grouped_matmul")
46+ add_ops_compile_options(
47+ OP_NAME GroupedMatmul
48+ COMPUTE_UNIT Ascend350_355e
49+ OPTIONS -DENABLE_CV_COMM_VIA_SSBUF=true
50+ )
38endif()51endif()
39 52 
40if(NOT BUILD_OPS_RTY_KERNEL)53if(NOT BUILD_OPS_RTY_KERNEL)
@@ -992,6 +992,7 @@ public:
992 .ExtendCfgInfo("aclnnSupport.value", "support_aclnn")992 .ExtendCfgInfo("aclnnSupport.value", "support_aclnn")
993 .ExtendCfgInfo("opFile.value", "grouped_matmul_apt");993 .ExtendCfgInfo("opFile.value", "grouped_matmul_apt");
994 this->AICore().AddConfig("ascend950", config950);994 this->AICore().AddConfig("ascend950", config950);
995+ this->AICore().AddConfig("ascend350", config950);
995 996 
996 OpAICoreConfig config_kirin = GetKirinCoreConfig();997 OpAICoreConfig config_kirin = GetKirinCoreConfig();
997 this->AICore().AddConfig("kirinx90", config_kirin);998 this->AICore().AddConfig("kirinx90", config_kirin);
@@ -138,6 +138,7 @@ public:
138 .ExtendCfgInfo("coreType.value", "AiCore")138 .ExtendCfgInfo("coreType.value", "AiCore")
139 .ExtendCfgInfo("opFile.value", "grouped_matmul_activation_quant_apt");139 .ExtendCfgInfo("opFile.value", "grouped_matmul_activation_quant_apt");
140 this->AICore().AddConfig("ascend950", config950);140 this->AICore().AddConfig("ascend950", config950);
141+ this->AICore().AddConfig("ascend350", config950);
141 }142 }
142};143};
143 144 
@@ -11,6 +11,8 @@
11if (BUILD_OPEN_PROJECT)11if (BUILD_OPEN_PROJECT)
12 if("${ASCEND_COMPUTE_UNIT}" STREQUAL "ascend950")12 if("${ASCEND_COMPUTE_UNIT}" STREQUAL "ascend950")
13 set(grouped_matmul_add_depends gmm/common CACHE INTERNAL "Dependencies for grouped_matmul_add")13 set(grouped_matmul_add_depends gmm/common CACHE INTERNAL "Dependencies for grouped_matmul_add")
14+ elseif("${ASCEND_COMPUTE_UNIT}" STREQUAL "ascend350")
15+ set(grouped_matmul_add_depends gmm/common CACHE INTERNAL "Dependencies for grouped_matmul_add")
14 endif()16 endif()
15 target_sources(op_host_aclnnExc PRIVATE17 target_sources(op_host_aclnnExc PRIVATE
16 grouped_matmul_add_def.cpp18 grouped_matmul_add_def.cpp
@@ -31,6 +33,13 @@ if (BUILD_OPS_RTY_KERNEL) # 回黄kernel
31 COMPUTE_UNIT Ascend950PR_959933 COMPUTE_UNIT Ascend950PR_9599
32 OPTIONS -mllvm -cce-aicore-dcci-before-kernel-end=false34 OPTIONS -mllvm -cce-aicore-dcci-before-kernel-end=false
33 )35 )
36+ elseif("${ASCEND_COMPUTE_UNIT}" STREQUAL "ascend350")
37+ set(grouped_matmul_add_depends gmm/common CACHE INTERNAL "Dependencies for grouped_matmul_add")
38+ add_ops_compile_options(
39+ OP_NAME GroupedMatmulAdd
40+ COMPUTE_UNIT Ascend350_355e
41+ OPTIONS -mllvm -cce-aicore-dcci-before-kernel-end=false
42+ )
34 endif()43 endif()
35else() # custom44else() # custom
36 add_modules_sources(45 add_modules_sources(
@@ -8,7 +8,6 @@
8 * See LICENSE in the root of the software repository for the full text of the License.8 * See LICENSE in the root of the software repository for the full text of the License.
9 */9 */
10 10 
11- 
12/*!11/*!
13 * \file grouped_matmul_add_def.cpp12 * \file grouped_matmul_add_def.cpp
14 * \brief13 * \brief
@@ -18,10 +17,10 @@
18 17 
19namespace ops {18namespace ops {
20constexpr int64_t INDEX_GROUP_LIST = 2;19constexpr int64_t INDEX_GROUP_LIST = 2;
21-class GroupedMatmulAdd : public OpDef20+class GroupedMatmulAdd : public OpDef {
22-{
23public:21public:
24- explicit GroupedMatmulAdd(const char* name) : OpDef(name)22+ explicit GroupedMatmulAdd(const char *name)
23+ : OpDef(name)
25 {24 {
26 this->Input("x")25 this->Input("x")
27 .ParamType(REQUIRED)26 .ParamType(REQUIRED)
@@ -67,7 +66,8 @@ public:
67 .ExtendCfgInfo("prebuildPattern.value", "Opaque")66 .ExtendCfgInfo("prebuildPattern.value", "Opaque")
68 .ExtendCfgInfo("coreType.value", "AiCore");67 .ExtendCfgInfo("coreType.value", "AiCore");
69 this->AICore().AddConfig("ascend950", config91095);68 this->AICore().AddConfig("ascend950", config91095);
69+ this->AICore().AddConfig("ascend350", config91095);
70 }70 }
71};71};
72OP_ADD(GroupedMatmulAdd);72OP_ADD(GroupedMatmulAdd);
73-} // namespace ops73+} // namespace ops
@@ -17,63 +17,76 @@
17namespace ops {17namespace ops {
18class GroupedMatmulFinalizeRouting : public OpDef {18class GroupedMatmulFinalizeRouting : public OpDef {
19public:19public:
20- explicit GroupedMatmulFinalizeRouting(const char* name) : OpDef(name)20+ explicit GroupedMatmulFinalizeRouting(const char *name)
21+ : OpDef(name)
21 {22 {
22 this->Input("x")23 this->Input("x")
23 .ParamType(REQUIRED)24 .ParamType(REQUIRED)
24 .DataType({ge::DT_INT8, ge::DT_INT8, ge::DT_INT8, ge::DT_INT8, ge::DT_INT8, ge::DT_INT8})25 .DataType({ge::DT_INT8, ge::DT_INT8, ge::DT_INT8, ge::DT_INT8, ge::DT_INT8, ge::DT_INT8})
25 .Format({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND})26 .Format({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND})
26- .UnknownShapeFormat({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND});27+ .UnknownShapeFormat(
28+ {ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND});
27 this->Input("w")29 this->Input("w")
28 .ParamType(REQUIRED)30 .ParamType(REQUIRED)
29 .DataType({ge::DT_INT8, ge::DT_INT8, ge::DT_INT4, ge::DT_INT4, ge::DT_INT8, ge::DT_INT8})31 .DataType({ge::DT_INT8, ge::DT_INT8, ge::DT_INT4, ge::DT_INT4, ge::DT_INT8, ge::DT_INT8})
30- .Format({ge::FORMAT_FRACTAL_NZ, ge::FORMAT_FRACTAL_NZ, ge::FORMAT_ND, ge::FORMAT_FRACTAL_NZ, ge::FORMAT_FRACTAL_NZ, ge::FORMAT_FRACTAL_NZ})32+ .Format({ge::FORMAT_FRACTAL_NZ, ge::FORMAT_FRACTAL_NZ, ge::FORMAT_ND, ge::FORMAT_FRACTAL_NZ,
31- .UnknownShapeFormat({ge::FORMAT_FRACTAL_NZ, ge::FORMAT_FRACTAL_NZ, ge::FORMAT_ND, ge::FORMAT_FRACTAL_NZ, ge::FORMAT_FRACTAL_NZ, ge::FORMAT_FRACTAL_NZ});33+ ge::FORMAT_FRACTAL_NZ, ge::FORMAT_FRACTAL_NZ})
34+ .UnknownShapeFormat({ge::FORMAT_FRACTAL_NZ, ge::FORMAT_FRACTAL_NZ, ge::FORMAT_ND, ge::FORMAT_FRACTAL_NZ,
35+ ge::FORMAT_FRACTAL_NZ, ge::FORMAT_FRACTAL_NZ});
32 this->Input("scale")36 this->Input("scale")
33 .ParamType(OPTIONAL)37 .ParamType(OPTIONAL)
34 .DataType({ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_INT64, ge::DT_INT64, ge::DT_FLOAT, ge::DT_BF16})38 .DataType({ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_INT64, ge::DT_INT64, ge::DT_FLOAT, ge::DT_BF16})
35 .Format({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND})39 .Format({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND})
36- .UnknownShapeFormat({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND});40+ .UnknownShapeFormat(
41+ {ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND});
37 this->Input("bias")42 this->Input("bias")
38 .ParamType(OPTIONAL)43 .ParamType(OPTIONAL)
39 .DataType({ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_BF16, ge::DT_BF16})44 .DataType({ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_BF16, ge::DT_BF16})
40 .Format({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND})45 .Format({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND})
41- .UnknownShapeFormat({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND});46+ .UnknownShapeFormat(
47+ {ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND});
42 this->Input("pertoken_scale")48 this->Input("pertoken_scale")
43 .ParamType(OPTIONAL)49 .ParamType(OPTIONAL)
44 .DataType({ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT})50 .DataType({ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT})
45 .Format({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND})51 .Format({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND})
46- .UnknownShapeFormat({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND});52+ .UnknownShapeFormat(
53+ {ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND});
47 this->Input("group_list")54 this->Input("group_list")
48 .ParamType(OPTIONAL)55 .ParamType(OPTIONAL)
49 .DataType({ge::DT_INT64, ge::DT_INT64, ge::DT_INT64, ge::DT_INT64, ge::DT_INT64, ge::DT_INT64})56 .DataType({ge::DT_INT64, ge::DT_INT64, ge::DT_INT64, ge::DT_INT64, ge::DT_INT64, ge::DT_INT64})
50 .Format({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND})57 .Format({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND})
51- .UnknownShapeFormat({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND});58+ .UnknownShapeFormat(
59+ {ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND});
52 this->Input("shared_input")60 this->Input("shared_input")
53 .ParamType(OPTIONAL)61 .ParamType(OPTIONAL)
54 .DataType({ge::DT_BF16, ge::DT_BF16, ge::DT_BF16, ge::DT_BF16, ge::DT_BF16, ge::DT_BF16})62 .DataType({ge::DT_BF16, ge::DT_BF16, ge::DT_BF16, ge::DT_BF16, ge::DT_BF16, ge::DT_BF16})
55 .Format({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND})63 .Format({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND})
56- .UnknownShapeFormat({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,ge::FORMAT_ND, ge::FORMAT_ND});64+ .UnknownShapeFormat(
65+ {ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND});
57 this->Input("logit")66 this->Input("logit")
58 .ParamType(OPTIONAL)67 .ParamType(OPTIONAL)
59 .DataType({ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT})68 .DataType({ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT})
60 .Format({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND})69 .Format({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND})
61- .UnknownShapeFormat({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND});70+ .UnknownShapeFormat(
71+ {ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND});
62 this->Input("row_index")72 this->Input("row_index")
63 .ParamType(OPTIONAL)73 .ParamType(OPTIONAL)
64 .DataType({ge::DT_INT64, ge::DT_INT32, ge::DT_INT64, ge::DT_INT64, ge::DT_INT64, ge::DT_INT64})74 .DataType({ge::DT_INT64, ge::DT_INT32, ge::DT_INT64, ge::DT_INT64, ge::DT_INT64, ge::DT_INT64})
65 .Format({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND})75 .Format({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND})
66- .UnknownShapeFormat({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND});76+ .UnknownShapeFormat(
77+ {ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND});
67 this->Input("offset")78 this->Input("offset")
68 .ParamType(OPTIONAL)79 .ParamType(OPTIONAL)
69 .DataType({ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT})80 .DataType({ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT})
70 .Format({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND})81 .Format({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND})
71- .UnknownShapeFormat({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND});82+ .UnknownShapeFormat(
83+ {ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND});
72 this->Output("y")84 this->Output("y")
73 .ParamType(REQUIRED)85 .ParamType(REQUIRED)
74 .DataType({ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT})86 .DataType({ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT})
75 .Format({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND})87 .Format({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND})
76- .UnknownShapeFormat({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND});88+ .UnknownShapeFormat(
89+ {ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND});
77 this->Attr("dtype").AttrType(OPTIONAL).Int(0);90 this->Attr("dtype").AttrType(OPTIONAL).Int(0);
78 this->Attr("shared_input_weight").AttrType(OPTIONAL).Float(1.0);91 this->Attr("shared_input_weight").AttrType(OPTIONAL).Float(1.0);
79 this->Attr("shared_input_offset").AttrType(OPTIONAL).Int(0);92 this->Attr("shared_input_offset").AttrType(OPTIONAL).Int(0);
@@ -84,11 +97,11 @@ public:
84 this->Attr("tuning_config").AttrType(OPTIONAL).ListInt({0});97 this->Attr("tuning_config").AttrType(OPTIONAL).ListInt({0});
85 OpAICoreConfig aicConfig;98 OpAICoreConfig aicConfig;
86 aicConfig.DynamicCompileStaticFlag(true)99 aicConfig.DynamicCompileStaticFlag(true)
87- .DynamicFormatFlag(true)100+ .DynamicFormatFlag(true)
88- .DynamicRankSupportFlag(true)101+ .DynamicRankSupportFlag(true)
89- .DynamicShapeSupportFlag(true)102+ .DynamicShapeSupportFlag(true)
90- .NeedCheckSupportFlag(false)103+ .NeedCheckSupportFlag(false)
91- .ExtendCfgInfo("softsync.flag", "true");104+ .ExtendCfgInfo("softsync.flag", "true");
92 this->AICore().AddConfig("ascend910b", aicConfig);105 this->AICore().AddConfig("ascend910b", aicConfig);
93 this->AICore().AddConfig("ascend910_93", aicConfig);106 this->AICore().AddConfig("ascend910_93", aicConfig);
94 107 
@@ -96,173 +109,128 @@ public:
96 config91095.Input("x")109 config91095.Input("x")
97 .ParamType(REQUIRED)110 .ParamType(REQUIRED)
98 .DataType({ge::DT_FLOAT8_E5M2, ge::DT_FLOAT8_E5M2, ge::DT_FLOAT8_E4M3FN, ge::DT_FLOAT8_E4M3FN,111 .DataType({ge::DT_FLOAT8_E5M2, ge::DT_FLOAT8_E5M2, ge::DT_FLOAT8_E4M3FN, ge::DT_FLOAT8_E4M3FN,
99- ge::DT_FLOAT4_E2M1,112+ ge::DT_FLOAT4_E2M1, ge::DT_INT8, ge::DT_INT8, ge::DT_INT8, ge::DT_INT8, ge::DT_FLOAT8_E4M3FN,
100- ge::DT_INT8,ge::DT_INT8,ge::DT_INT8,ge::DT_INT8,113+ ge::DT_HIFLOAT8, ge::DT_FLOAT8_E4M3FN, ge::DT_HIFLOAT8, ge::DT_FLOAT8_E4M3FN,
101- ge::DT_FLOAT8_E4M3FN,ge::DT_HIFLOAT8,114+ ge::DT_FLOAT8_E4M3FN})
102- ge::DT_FLOAT8_E4M3FN,ge::DT_HIFLOAT8,115+ .Format({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
103- ge::DT_FLOAT8_E4M3FN,ge::DT_FLOAT8_E4M3FN116+ ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
104- })117+ ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND})
105- .Format({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,ge::FORMAT_ND,118+ .UnknownShapeFormat({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
106- ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,ge::FORMAT_ND,ge::FORMAT_ND,119+ ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
107- ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND120+ ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND});
108- })
109- .UnknownShapeFormat({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,ge::FORMAT_ND,
110- ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,ge::FORMAT_ND,ge::FORMAT_ND,
111- ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND
112- });
113 config91095.Input("w")121 config91095.Input("w")
114 .ParamType(REQUIRED)122 .ParamType(REQUIRED)
115 .DataType({ge::DT_FLOAT8_E5M2, ge::DT_FLOAT8_E4M3FN, ge::DT_FLOAT8_E5M2, ge::DT_FLOAT8_E4M3FN,123 .DataType({ge::DT_FLOAT8_E5M2, ge::DT_FLOAT8_E4M3FN, ge::DT_FLOAT8_E5M2, ge::DT_FLOAT8_E4M3FN,
116- ge::DT_FLOAT4_E2M1,124+ ge::DT_FLOAT4_E2M1, ge::DT_INT8, ge::DT_INT8, ge::DT_INT8, ge::DT_INT8, ge::DT_FLOAT8_E4M3FN,
117- ge::DT_INT8,ge::DT_INT8,ge::DT_INT8,ge::DT_INT8,125+ ge::DT_HIFLOAT8, ge::DT_FLOAT8_E4M3FN, ge::DT_HIFLOAT8, ge::DT_FLOAT4_E2M1, ge::DT_FLOAT4_E2M1})
118- ge::DT_FLOAT8_E4M3FN,ge::DT_HIFLOAT8,126+ .Format({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_FRACTAL_NZ,
119- ge::DT_FLOAT8_E4M3FN,ge::DT_HIFLOAT8,127+ ge::FORMAT_FRACTAL_NZ, ge::FORMAT_FRACTAL_NZ, ge::FORMAT_FRACTAL_NZ, ge::FORMAT_FRACTAL_NZ,
120- ge::DT_FLOAT4_E2M1, ge::DT_FLOAT4_E2M1128+ ge::FORMAT_FRACTAL_NZ, ge::FORMAT_FRACTAL_NZ, ge::FORMAT_FRACTAL_NZ, ge::FORMAT_FRACTAL_NZ_C0_32,
121- })129+ ge::FORMAT_FRACTAL_NZ})
122- .Format({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
123- ge::FORMAT_FRACTAL_NZ, ge::FORMAT_FRACTAL_NZ, ge::FORMAT_FRACTAL_NZ,ge::FORMAT_FRACTAL_NZ, ge::FORMAT_FRACTAL_NZ, ge::FORMAT_FRACTAL_NZ,
124- ge::FORMAT_FRACTAL_NZ, ge::FORMAT_FRACTAL_NZ,
125- ge::FORMAT_FRACTAL_NZ_C0_32, ge::FORMAT_FRACTAL_NZ
126- })
127 .UnknownShapeFormat({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,130 .UnknownShapeFormat({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
128- ge::FORMAT_FRACTAL_NZ, ge::FORMAT_FRACTAL_NZ, ge::FORMAT_FRACTAL_NZ,ge::FORMAT_FRACTAL_NZ, ge::FORMAT_FRACTAL_NZ, ge::FORMAT_FRACTAL_NZ,131+ ge::FORMAT_FRACTAL_NZ, ge::FORMAT_FRACTAL_NZ, ge::FORMAT_FRACTAL_NZ,
129- ge::FORMAT_FRACTAL_NZ, ge::FORMAT_FRACTAL_NZ,132+ ge::FORMAT_FRACTAL_NZ, ge::FORMAT_FRACTAL_NZ, ge::FORMAT_FRACTAL_NZ,
130- ge::FORMAT_FRACTAL_NZ_C0_32, ge::FORMAT_FRACTAL_NZ133+ ge::FORMAT_FRACTAL_NZ, ge::FORMAT_FRACTAL_NZ, ge::FORMAT_FRACTAL_NZ_C0_32,
131- });134+ ge::FORMAT_FRACTAL_NZ});
132 config91095.Input("scale")135 config91095.Input("scale")
133 .ParamType(OPTIONAL)136 .ParamType(OPTIONAL)
134 .DataType({ge::DT_FLOAT8_E8M0, ge::DT_FLOAT8_E8M0, ge::DT_FLOAT8_E8M0, ge::DT_FLOAT8_E8M0,137 .DataType({ge::DT_FLOAT8_E8M0, ge::DT_FLOAT8_E8M0, ge::DT_FLOAT8_E8M0, ge::DT_FLOAT8_E8M0,
135- ge::DT_FLOAT8_E8M0, 138+ ge::DT_FLOAT8_E8M0, ge::DT_BF16, ge::DT_FLOAT, ge::DT_BF16, ge::DT_FLOAT, ge::DT_BF16,
136- ge::DT_BF16, ge::DT_FLOAT,139+ ge::DT_BF16, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT8_E8M0, ge::DT_FLOAT8_E8M0})
137- ge::DT_BF16, ge::DT_FLOAT,140+ .Format({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
138- ge::DT_BF16, ge::DT_BF16,141+ ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
139- ge::DT_FLOAT, ge::DT_FLOAT,142+ ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND})
140- ge::DT_FLOAT8_E8M0, ge::DT_FLOAT8_E8M0
141- })
142- .Format({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
143- ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,ge::FORMAT_ND,ge::FORMAT_ND,
144- ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND
145- })
146 .UnknownShapeFormat({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,143 .UnknownShapeFormat({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
147- ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,ge::FORMAT_ND,ge::FORMAT_ND,144+ ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
148- ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND145+ ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND});
149- });
150 config91095.Input("bias")146 config91095.Input("bias")
151 .ParamType(OPTIONAL)147 .ParamType(OPTIONAL)
152- .DataType({ge::DT_BF16, ge::DT_BF16, ge::DT_BF16, ge::DT_BF16, ge::DT_BF16,148+ .DataType({ge::DT_BF16, ge::DT_BF16, ge::DT_BF16, ge::DT_BF16, ge::DT_BF16, ge::DT_BF16, ge::DT_BF16,
153- ge::DT_BF16, ge::DT_BF16, ge::DT_BF16, ge::DT_BF16, ge::DT_BF16,ge::DT_BF16,149+ ge::DT_BF16, ge::DT_BF16, ge::DT_BF16, ge::DT_BF16, ge::DT_BF16, ge::DT_BF16, ge::DT_BF16,
154- ge::DT_BF16, ge::DT_BF16, ge::DT_BF16, ge::DT_BF16150+ ge::DT_BF16})
155- })151+ .Format({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
156- .Format({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,152+ ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
157- ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,ge::FORMAT_ND,ge::FORMAT_ND,153+ ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND})
158- ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND
159- })
160 .UnknownShapeFormat({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,154 .UnknownShapeFormat({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
161- ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,ge::FORMAT_ND,ge::FORMAT_ND,155+ ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
162- ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND156+ ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND});
163- });
164 config91095.Input("pertoken_scale")157 config91095.Input("pertoken_scale")
165 .ParamType(OPTIONAL)158 .ParamType(OPTIONAL)
166- .DataType({ge::DT_FLOAT8_E8M0, ge::DT_FLOAT8_E8M0, ge::DT_FLOAT8_E8M0, ge::DT_FLOAT8_E8M0,ge::DT_FLOAT8_E8M0,159+ .DataType({ge::DT_FLOAT8_E8M0, ge::DT_FLOAT8_E8M0, ge::DT_FLOAT8_E8M0, ge::DT_FLOAT8_E8M0,
167- ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,ge::DT_FLOAT, ge::DT_FLOAT,160+ ge::DT_FLOAT8_E8M0, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
168- ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT8_E8M0, ge::DT_FLOAT8_E8M0161+ ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT8_E8M0, ge::DT_FLOAT8_E8M0})
169- })162+ .Format({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
170- .Format({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,163+ ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
171- ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,ge::FORMAT_ND,ge::FORMAT_ND,164+ ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND})
172- ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND
173- })
174 .UnknownShapeFormat({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,165 .UnknownShapeFormat({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
175- ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,ge::FORMAT_ND,ge::FORMAT_ND,166+ ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
176- ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND167+ ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND});
177- });
178 config91095.Input("group_list")168 config91095.Input("group_list")
179 .ParamType(OPTIONAL)169 .ParamType(OPTIONAL)
180- .DataType({ge::DT_INT64, ge::DT_INT64, ge::DT_INT64, ge::DT_INT64, ge::DT_INT64,170+ .DataType({ge::DT_INT64, ge::DT_INT64, ge::DT_INT64, ge::DT_INT64, ge::DT_INT64, ge::DT_INT64, ge::DT_INT64,
181- ge::DT_INT64, ge::DT_INT64, ge::DT_INT64, ge::DT_INT64, ge::DT_INT64,ge::DT_INT64,171+ ge::DT_INT64, ge::DT_INT64, ge::DT_INT64, ge::DT_INT64, ge::DT_INT64, ge::DT_INT64, ge::DT_INT64,
182- ge::DT_INT64, ge::DT_INT64, ge::DT_INT64, ge::DT_INT64172+ ge::DT_INT64})
183- })173+ .Format({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
184- .Format({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,174+ ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
185- ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,ge::FORMAT_ND,ge::FORMAT_ND,175+ ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND})
186- ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND
187- })
188 .UnknownShapeFormat({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,176 .UnknownShapeFormat({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
189- ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,ge::FORMAT_ND,ge::FORMAT_ND,177+ ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
190- ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND178+ ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND});
191- });
192 config91095.Input("shared_input")179 config91095.Input("shared_input")
193 .ParamType(OPTIONAL)180 .ParamType(OPTIONAL)
194- .DataType({ge::DT_BF16, ge::DT_BF16, ge::DT_BF16, ge::DT_BF16, ge::DT_BF16,181+ .DataType({ge::DT_BF16, ge::DT_BF16, ge::DT_BF16, ge::DT_BF16, ge::DT_BF16, ge::DT_BF16, ge::DT_BF16,
195- ge::DT_BF16, ge::DT_BF16, ge::DT_BF16, ge::DT_BF16, ge::DT_BF16,ge::DT_BF16,182+ ge::DT_BF16, ge::DT_BF16, ge::DT_BF16, ge::DT_BF16, ge::DT_BF16, ge::DT_BF16, ge::DT_BF16,
196- ge::DT_BF16, ge::DT_BF16, ge::DT_BF16, ge::DT_BF16183+ ge::DT_BF16})
197- })184+ .Format({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
198- .Format({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,185+ ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
199- ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,ge::FORMAT_ND,ge::FORMAT_ND,186+ ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND})
200- ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND
201- })
202 .UnknownShapeFormat({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,187 .UnknownShapeFormat({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
203- ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,ge::FORMAT_ND,ge::FORMAT_ND,188+ ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
204- ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND189+ ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND});
205- });
206 config91095.Input("logit")190 config91095.Input("logit")
207 .ParamType(OPTIONAL)191 .ParamType(OPTIONAL)
208- .DataType({ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,192+ .DataType({ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
209- ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,193+ ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
210- ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT194+ ge::DT_FLOAT})
211- })195+ .Format({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
212- .Format({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,196+ ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
213- ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,ge::FORMAT_ND,ge::FORMAT_ND,197+ ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND})
214- ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND
215- })
216 .UnknownShapeFormat({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,198 .UnknownShapeFormat({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
217- ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,ge::FORMAT_ND,ge::FORMAT_ND,199+ ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
218- ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND200+ ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND});
219- });
220 config91095.Input("row_index")201 config91095.Input("row_index")
221 .ParamType(OPTIONAL)202 .ParamType(OPTIONAL)
222- .DataType({ge::DT_INT64, ge::DT_INT64, ge::DT_INT64, ge::DT_INT64, ge::DT_INT64,203+ .DataType({ge::DT_INT64, ge::DT_INT64, ge::DT_INT64, ge::DT_INT64, ge::DT_INT64, ge::DT_INT64, ge::DT_INT64,
223- ge::DT_INT64, ge::DT_INT64,204+ ge::DT_INT32, ge::DT_INT32, ge::DT_INT64, ge::DT_INT64, ge::DT_INT64, ge::DT_INT64, ge::DT_INT64,
224- ge::DT_INT32, ge::DT_INT32,205+ ge::DT_INT64})
225- ge::DT_INT64, ge::DT_INT64,206+ .Format({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
226- ge::DT_INT64, ge::DT_INT64,207+ ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
227- ge::DT_INT64, ge::DT_INT64208+ ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND})
228- })
229- .Format({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
230- ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,ge::FORMAT_ND,ge::FORMAT_ND,
231- ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND
232- })
233 .UnknownShapeFormat({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,209 .UnknownShapeFormat({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
234- ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,ge::FORMAT_ND,ge::FORMAT_ND,210+ ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
235- ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND211+ ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND});
236- });
237 config91095.Input("offset")212 config91095.Input("offset")
238 .ParamType(OPTIONAL)213 .ParamType(OPTIONAL)
239- .DataType({ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,214+ .DataType({ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
240- ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,ge::DT_FLOAT,215+ ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
241- ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT216+ ge::DT_FLOAT})
242- })217+ .Format({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
243- .Format({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,218+ ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
244- ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,ge::FORMAT_ND,ge::FORMAT_ND,219+ ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND})
245- ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND
246- })
247 .UnknownShapeFormat({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,220 .UnknownShapeFormat({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
248- ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,ge::FORMAT_ND,ge::FORMAT_ND,221+ ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
249- ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND222+ ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND});
250- });
251 config91095.Output("y")223 config91095.Output("y")
252 .ParamType(REQUIRED)224 .ParamType(REQUIRED)
253- .DataType({ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,225+ .DataType({ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
254- ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,ge::DT_FLOAT,226+ ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
255- ge::DT_FLOAT, ge::DT_FLOAT,227+ ge::DT_FLOAT})
256- ge::DT_FLOAT, ge::DT_FLOAT228+ .Format({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
257- })229+ ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
258- .Format({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,230+ ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND})
259- ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,ge::FORMAT_ND,ge::FORMAT_ND,
260- ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND
261- })
262 .UnknownShapeFormat({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,231 .UnknownShapeFormat({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
263- ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,ge::FORMAT_ND,ge::FORMAT_ND,232+ ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
264- ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND233+ ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND});
265- });
266 config91095.DynamicCompileStaticFlag(true)234 config91095.DynamicCompileStaticFlag(true)
267 .DynamicFormatFlag(true)235 .DynamicFormatFlag(true)
268 .DynamicRankSupportFlag(true)236 .DynamicRankSupportFlag(true)
@@ -272,11 +240,12 @@ public:
272 .ExtendCfgInfo("prebuildPattern.value", "Opaque")240 .ExtendCfgInfo("prebuildPattern.value", "Opaque")
273 .ExtendCfgInfo("coreType.value", "AiCore")241 .ExtendCfgInfo("coreType.value", "AiCore")
274 .ExtendCfgInfo("aclnnSupport.value", "support_aclnn")242 .ExtendCfgInfo("aclnnSupport.value", "support_aclnn")
275- .ExtendCfgInfo("opFile.value","grouped_matmul_finalize_routing_apt");243+ .ExtendCfgInfo("opFile.value", "grouped_matmul_finalize_routing_apt");
276 this->AICore().AddConfig("ascend950", config91095);244 this->AICore().AddConfig("ascend950", config91095);
245+ this->AICore().AddConfig("ascend350", config91095);
277 }246 }
278};247};
279 248 
280OP_ADD(GroupedMatmulFinalizeRouting);249OP_ADD(GroupedMatmulFinalizeRouting);
281 250 
282-}251+} // namespace ops
@@ -353,6 +353,7 @@ public:
353 .ExtendCfgInfo("aclnnSupport.value", "support_aclnn")353 .ExtendCfgInfo("aclnnSupport.value", "support_aclnn")
354 .ExtendCfgInfo("opFile.value", "grouped_matmul_swiglu_quant_v2_apt");354 .ExtendCfgInfo("opFile.value", "grouped_matmul_swiglu_quant_v2_apt");
355 this->AICore().AddConfig("ascend950", config91095);355 this->AICore().AddConfig("ascend950", config91095);
356+ this->AICore().AddConfig("ascend350", config91095);
356 }357 }
357};358};
358 359 
@@ -11,6 +11,8 @@
11if (BUILD_OPEN_PROJECT)11if (BUILD_OPEN_PROJECT)
12 if("${ASCEND_COMPUTE_UNIT}" STREQUAL "ascend950")12 if("${ASCEND_COMPUTE_UNIT}" STREQUAL "ascend950")
13 set(quant_grouped_matmul_inplace_add_apt_depends gmm/common CACHE INTERNAL "Dependencies for quant_grouped_matmul_inplace_add_apt")13 set(quant_grouped_matmul_inplace_add_apt_depends gmm/common CACHE INTERNAL "Dependencies for quant_grouped_matmul_inplace_add_apt")
14+ elseif("${ASCEND_COMPUTE_UNIT}" STREQUAL "ascend350")
15+ set(quant_grouped_matmul_inplace_add_apt_depends gmm/common CACHE INTERNAL "Dependencies for quant_grouped_matmul_inplace_add_apt")
14 endif()16 endif()
15 target_sources(op_host_aclnnExc PRIVATE17 target_sources(op_host_aclnnExc PRIVATE
16 quant_grouped_matmul_inplace_add_def.cpp18 quant_grouped_matmul_inplace_add_def.cpp
@@ -32,6 +34,13 @@ if (BUILD_OPS_RTY_KERNEL) # 回黄kernel
32 COMPUTE_UNIT Ascend950PR_959934 COMPUTE_UNIT Ascend950PR_9599
33 OPTIONS -mllvm -cce-aicore-dcci-before-kernel-end=false35 OPTIONS -mllvm -cce-aicore-dcci-before-kernel-end=false
34 )36 )
37+ elseif("${ASCEND_COMPUTE_UNIT}" STREQUAL "ascend350")
38+ set(quant_grouped_matmul_inplace_add_apt_depends gmm/common CACHE INTERNAL "Dependencies for quant_grouped_matmul_inplace_add_apt")
39+ add_ops_compile_options(
40+ OP_NAME QuantGroupedMatmulInplaceAdd
41+ COMPUTE_UNIT Ascend350_355e
42+ OPTIONS -mllvm -cce-aicore-dcci-before-kernel-end=false
43+ )
35 endif()44 endif()
36 45 
37else()46else()
@@ -17,7 +17,8 @@
17namespace ops {17namespace ops {
18class QuantGroupedMatmulInplaceAdd : public OpDef {18class QuantGroupedMatmulInplaceAdd : public OpDef {
19public:19public:
20- explicit QuantGroupedMatmulInplaceAdd(const char *name) : OpDef(name)20+ explicit QuantGroupedMatmulInplaceAdd(const char *name)
21+ : OpDef(name)
21 {22 {
22 this->Input("x1")23 this->Input("x1")
23 .ParamType(REQUIRED)24 .ParamType(REQUIRED)
@@ -43,8 +44,7 @@ public:
43 .Format({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND});44 .Format({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND});
44 this->Input("scale1")45 this->Input("scale1")
45 .ParamType(OPTIONAL)46 .ParamType(OPTIONAL)
46- .DataType(47+ .DataType({ge::DT_FLOAT8_E8M0, ge::DT_FLOAT8_E8M0, ge::DT_FLOAT8_E8M0, ge::DT_FLOAT8_E8M0, ge::DT_FLOAT})
47- {ge::DT_FLOAT8_E8M0, ge::DT_FLOAT8_E8M0, ge::DT_FLOAT8_E8M0, ge::DT_FLOAT8_E8M0, ge::DT_FLOAT})
48 .Format({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND});48 .Format({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND});
49 this->Output("y")49 this->Output("y")
50 .ParamType(REQUIRED)50 .ParamType(REQUIRED)
@@ -52,7 +52,7 @@ public:
52 .Format({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND});52 .Format({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND});
53 this->Attr("group_list_type")53 this->Attr("group_list_type")
54 .AttrType(OPTIONAL)54 .AttrType(OPTIONAL)
55- .Int(0); // Indicates whether the value in group_dist is cumsum or count.55+ .Int(0); // Indicates whether the value in group_dist is cumsum or count.
56 this->Attr("group_size").AttrType(OPTIONAL).Int(0);56 this->Attr("group_size").AttrType(OPTIONAL).Int(0);
57 OpAICoreConfig config91095;57 OpAICoreConfig config91095;
58 config91095.DynamicCompileStaticFlag(true)58 config91095.DynamicCompileStaticFlag(true)
@@ -64,10 +64,11 @@ public:
64 .ExtendCfgInfo("prebuildPattern.value", "Opaque")64 .ExtendCfgInfo("prebuildPattern.value", "Opaque")
65 .ExtendCfgInfo("coreType.value", "AiCore")65 .ExtendCfgInfo("coreType.value", "AiCore")
66 .ExtendCfgInfo("aclnnSupport.value", "support_aclnn")66 .ExtendCfgInfo("aclnnSupport.value", "support_aclnn")
67- .ExtendCfgInfo("opFile.value","quant_grouped_matmul_inplace_add_apt");67+ .ExtendCfgInfo("opFile.value", "quant_grouped_matmul_inplace_add_apt");
68 this->AICore().AddConfig("ascend950", config91095);68 this->AICore().AddConfig("ascend950", config91095);
69+ this->AICore().AddConfig("ascend350", config91095);
69 }70 }
70};71};
71 72 
72OP_ADD(QuantGroupedMatmulInplaceAdd);73OP_ADD(QuantGroupedMatmulInplaceAdd);
73-} // namespace ops74+} // namespace ops