已合并
【fix】: 修复FusedBackend场景的一些问题 #1834
【fix】: 修复FusedBackend场景的一些问题 #1834
已合并
liyuewei创建于 8月26日
共 15 个文件变更+609-39
@@ -67,6 +67,21 @@ bool HasGatherNode(const std::vector<std::vector<af::AscGraph>> &schedule_groups
67 return false;67 return false;
68}68}
69 69 
70+void FilterOutputAliasWorkspace(ModelInfo &model_info) {
71+ for (const auto &output_node : model_info.output_nodes) {
72+ if (output_node == nullptr || !af::ops::IsOps<af::ascir_op::Output>(output_node)) {
73+ continue;
74+ }
75+ const int64_t tensor_id = output_node->inputs[0].attr.mem.tensor_id;
76+ const auto workspace_iter = model_info.workspace_size_map.find(tensor_id);
77+ if (workspace_iter == model_info.workspace_size_map.end()) {
78+ continue;
79+ }
80+ GELOGD("Filter output alias tensor id [%ld] from workspace size map.", tensor_id);
81+ model_info.workspace_size_map.erase(workspace_iter);
82+ }
83+}
84+ 
70bool ShouldEnableGatherReducePenalty(const std::vector<std::vector<af::AscGraph>> &schedule_groups,85bool ShouldEnableGatherReducePenalty(const std::vector<std::vector<af::AscGraph>> &schedule_groups,
71 const size_t group_id, const bool enable_group_parallel) {86 const size_t group_id, const bool enable_group_parallel) {
72 if (enable_group_parallel) {87 if (enable_group_parallel) {
@@ -539,6 +554,7 @@ af::Status ProcessAndSetScheduleGroupInfo(const std::vector<std::vector<af::AscG
539 model_info.schedule_group_ident.group_id = schedule_group_id;554 model_info.schedule_group_ident.group_id = schedule_group_id;
540 model_info.input_nodes = schedule_results.input_nodes;555 model_info.input_nodes = schedule_results.input_nodes;
541 model_info.output_nodes = schedule_results.output_nodes;556 model_info.output_nodes = schedule_results.output_nodes;
557+ FilterOutputAliasWorkspace(model_info);
542 auto it = all_graph_score_funcs.find(model_info.graph_name);558 auto it = all_graph_score_funcs.find(model_info.graph_name);
543 if (it != all_graph_score_funcs.end()) {559 if (it != all_graph_score_funcs.end()) {
544 model_info.score_func = it->second;560 model_info.score_func = it->second;
@@ -58,6 +58,19 @@ inline const std::string &AddSlogExtend() {
58 }58 }
59 return kGeLogUtils;59 return kGeLogUtils;
60}60}
61+ 
62+std::set<int64_t> GetWorkspaceTensorIds(const TensorIdSet &workspace_tensor_id_set) {
63+ std::set<int64_t> workspace_ids;
64+ for (const auto &[asc_graph_id, impl_graph_ids] : workspace_tensor_id_set) {
65+ (void)asc_graph_id;
66+ for (const auto &[impl_graph_id, tensor_ids] : impl_graph_ids) {
67+ (void)impl_graph_id;
68+ workspace_ids.insert(tensor_ids.begin(), tensor_ids.end());
69+ }
70+ }
71+ return workspace_ids;
72+}
73+ 
61template <typename T>74template <typename T>
62af::Status IsUpperBoundValid(const Expr &min_expr, const Expr &max_expr) {75af::Status IsUpperBoundValid(const Expr &min_expr, const Expr &max_expr) {
63 T min_value{};76 T min_value{};
@@ -3018,6 +3031,23 @@ af::Status TilingCodeGenImpl::GenPGOByCoreNumSearchTilingKeyCollectTilingData(Fu
3018 }3031 }
3019 }3032 }
3020 3033 
3034+ std::set<int64_t> workspace_ids;
3035+ for (const auto &asc_graph_workspace_ids : workspace_tensor_id_set_) {
3036+ for (const auto &impl_graph_workspace_ids : asc_graph_workspace_ids.second) {
3037+ workspace_ids.insert(impl_graph_workspace_ids.second.begin(), impl_graph_workspace_ids.second.end());
3038+ }
3039+ }
3040+ for (const auto &tensor_id : workspace_ids) {
3041+ const auto tensor_id_str = std::to_string(tensor_id);
3042+ for (const auto &asc_graph_map_iter : namespace_map) {
3043+ const auto asc_graph_id = asc_graph_map_iter.first;
3044+ tiling_func_.AddLine(" tiling_data_tmp.set_workspace" + tensor_id_str +
3045+ "(std::max(tiling_data_tmp.get_workspace" + tensor_id_str + "(), ascgraph_tiling_data_" +
3046+ std::to_string(asc_graph_id) + ".get_workspace" + tensor_id_str + "()));");
3047+ }
3048+ }
3049+ GenWorkspaceOffsetFinalize("tiling_data_tmp");
3050+ 
3021 // The root search starts from the current core-count probe. After composing3051 // The root search starts from the current core-count probe. After composing
3022 // graph candidates, normalize the outer block_dim to the shared block range3052 // graph candidates, normalize the outer block_dim to the shared block range
3023 // instead of carrying that probe value into the candidate identity.3053 // instead of carrying that probe value into the candidate identity.
@@ -3069,6 +3099,7 @@ af::Status TilingCodeGenImpl::GenPGOByCoreNumSearchTilingKey() {
3069 "tiling_data, uint32_t max_block_dim) {");3099 "tiling_data, uint32_t max_block_dim) {");
3070 tiling_func_.AddLine(" (void)tiling_data_list; (void)tiling_data; (void)max_block_dim;");3100 tiling_func_.AddLine(" (void)tiling_data_list; (void)tiling_data; (void)max_block_dim;");
3071 tiling_func_.AddLine(" bool ret = true;");3101 tiling_func_.AddLine(" bool ret = true;");
3102+ GenWorkspaceOffsetReset("*tiling_data");
3072 tiling_func_.AddLine(" for (uint32_t block_dim_i=1; block_dim_i <= max_block_dim; block_dim_i++) {");3103 tiling_func_.AddLine(" for (uint32_t block_dim_i=1; block_dim_i <= max_block_dim; block_dim_i++) {");
3073 tiling_func_.AddLine(" int32_t tiling_case;");3104 tiling_func_.AddLine(" int32_t tiling_case;");
3074 tiling_func_.AddLine(" AutofuseTilingData tiling_data_tmp;");3105 tiling_func_.AddLine(" AutofuseTilingData tiling_data_tmp;");
@@ -3498,12 +3529,53 @@ void TilingCodeGenImpl::GenGetScheduleResultTail(
3498 tiling_func_.AddLine("}");3529 tiling_func_.AddLine("}");
3499}3530}
3500 3531 
3532+void TilingCodeGenImpl::GenWorkspaceOffsetHelpers() {
3533+ const auto workspace_ids = GetWorkspaceTensorIds(workspace_tensor_id_set_);
3534+ if (workspace_ids.empty()) {
3535+ return;
3536+ }
3537+ 
3538+ const auto &tiling_data_type = config_.tiling_data_type_name;
3539+ // ATT 内部仍使用 workspace 字段暂存全局大小,所有子图完成后再统一转换为 offset。
3540+ tiling_func_.AddLine("inline void ResetWorkspaceSizes(" + tiling_data_type + " &tiling_data) {");
3541+ for (const auto &tensor_id : workspace_ids) {
3542+ tiling_func_.AddLine(" tiling_data.set_workspace" + std::to_string(tensor_id) + "(0);");
3543+ }
3544+ tiling_func_.AddLine("}");
3545+ 
3546+ tiling_func_.AddLine("inline void FinalizeWorkspaceOffsets(" + tiling_data_type + " &tiling_data) {");
3547+ tiling_func_.AddLine(" uint32_t workspace_offset = 0U;");
3548+ for (const auto &tensor_id : workspace_ids) {
3549+ const auto tensor_id_str = std::to_string(tensor_id);
3550+ tiling_func_.AddLine(" const uint32_t workspace_size_" + tensor_id_str + " = tiling_data.get_workspace" +
3551+ tensor_id_str + "();");
3552+ tiling_func_.AddLine(" tiling_data.set_workspace" + tensor_id_str + "(workspace_offset);");
3553+ tiling_func_.AddLine(" workspace_offset += workspace_size_" + tensor_id_str + ";");
3554+ }
3555+ tiling_func_.AddLine("}");
3556+}
3557+ 
3558+void TilingCodeGenImpl::GenWorkspaceOffsetReset(const std::string &tiling_data_name) {
3559+ if (GetWorkspaceTensorIds(workspace_tensor_id_set_).empty()) {
3560+ return;
3561+ }
3562+ tiling_func_.AddLine(" ResetWorkspaceSizes(" + tiling_data_name + ");");
3563+}
3564+ 
3565+void TilingCodeGenImpl::GenWorkspaceOffsetFinalize(const std::string &tiling_data_name) {
3566+ if (GetWorkspaceTensorIds(workspace_tensor_id_set_).empty()) {
3567+ return;
3568+ }
3569+ tiling_func_.AddLine(" FinalizeWorkspaceOffsets(" + tiling_data_name + ");");
3570+}
3571+ 
3501void TilingCodeGenImpl::GenUpdateWorkspace(const size_t asc_graph_id, const size_t impl_graph_id) {3572void TilingCodeGenImpl::GenUpdateWorkspace(const size_t asc_graph_id, const size_t impl_graph_id) {
3502 for (const auto &tensor_id : workspace_tensor_id_set_[asc_graph_id][impl_graph_id]) {3573 for (const auto &tensor_id : workspace_tensor_id_set_[asc_graph_id][impl_graph_id]) {
3503 auto tensor_id_str = to_string(tensor_id);3574 auto tensor_id_str = to_string(tensor_id);
3504 tiling_func_.AddLine(" auto it" + tensor_id_str + " = workspace_map.find(" + tensor_id_str + ");");3575 tiling_func_.AddLine(" auto it" + tensor_id_str + " = workspace_map.find(" + tensor_id_str + ");");
3505 tiling_func_.AddLine(" if (it" + tensor_id_str + " != workspace_map.end()) {");3576 tiling_func_.AddLine(" if (it" + tensor_id_str + " != workspace_map.end()) {");
3506- tiling_func_.AddLine(" tiling_data.set_workspace" + tensor_id_str + "(it" + tensor_id_str + "->second);");3577+ tiling_func_.AddLine(" tiling_data.set_workspace" + tensor_id_str + "(std::max(tiling_data.get_workspace" +
3578+ tensor_id_str + "(), static_cast<uint32_t>(it" + tensor_id_str + "->second)));");
3507 tiling_func_.AddLine(" }");3579 tiling_func_.AddLine(" }");
3508 }3580 }
3509}3581}
@@ -4334,6 +4406,7 @@ af::Status TilingCodeGenImpl::GenFusedScheduleResultsGetTilingDefine(const Fused
4334 "Generate init and query cache code failed.");4406 "Generate init and query cache code failed.");
4335 }4407 }
4336 tiling_func_.AddLine(" bool ret = true;"); // 声明ret变量用于缓存保存操作4408 tiling_func_.AddLine(" bool ret = true;"); // 声明ret变量用于缓存保存操作
4409+ GenWorkspaceOffsetReset("tiling_data");
4337 4410 
4338 size_t asc_graph_id = 0UL;4411 size_t asc_graph_id = 0UL;
4339 const std::string failed_log_level =4412 const std::string failed_log_level =
@@ -4355,6 +4428,7 @@ af::Status TilingCodeGenImpl::GenFusedScheduleResultsGetTilingDefine(const Fused
4355 "max_block_dim;");4428 "max_block_dim;");
4356 asc_graph_id++;4429 asc_graph_id++;
4357 }4430 }
4431+ GenWorkspaceOffsetFinalize("tiling_data");
4358 tiling_func_.AddLine(" tiling_data.set_block_dim(max_block_dim);");4432 tiling_func_.AddLine(" tiling_data.set_block_dim(max_block_dim);");
4359 4433 
4360 // Save only automatic tilings; explicit case/PGO requests bypass operator cache.4434 // Save only automatic tilings; explicit case/PGO requests bypass operator cache.
@@ -4387,6 +4461,7 @@ af::Status TilingCodeGenImpl::GenPGOByCoreNumFusedScheduleResultsGetTilingDefine
4387 auto core_num = BaseTypeUtils::DumpHardware(HardwareDef::CORENUM);4461 auto core_num = BaseTypeUtils::DumpHardware(HardwareDef::CORENUM);
4388 tiling_func_.AddLine(" tiling_data->set_block_dim(block_dim_i);");4462 tiling_func_.AddLine(" tiling_data->set_block_dim(block_dim_i);");
4389 tiling_func_.AddLine(" tiling_data->set_" + core_num + "(block_dim_i);");4463 tiling_func_.AddLine(" tiling_data->set_" + core_num + "(block_dim_i);");
4464+ GenWorkspaceOffsetReset("*tiling_data");
4390 for (const auto &asc_graph_namespace_map : namespace_map) {4465 for (const auto &asc_graph_namespace_map : namespace_map) {
4391 const std::string &asc_graph_namespace = "AscGraph" + std::to_string(asc_graph_namespace_map.first);4466 const std::string &asc_graph_namespace = "AscGraph" + std::to_string(asc_graph_namespace_map.first);
4392 tiling_func_.AddLine(" if (!" + asc_graph_namespace + "::PGOByCoreNumSearchTilingKey(vec" +4467 tiling_func_.AddLine(" if (!" + asc_graph_namespace + "::PGOByCoreNumSearchTilingKey(vec" +
@@ -4416,11 +4491,26 @@ af::Status TilingCodeGenImpl::GenPGOFusedScheduleResultsGetTilingDefine(const Fu
4416 tiling_func_.AddLine(" double cur_perf = DBL_MAX;");4491 tiling_func_.AddLine(" double cur_perf = DBL_MAX;");
4417 tiling_func_.AddLine(" uint32_t cur_block_dim = 1;");4492 tiling_func_.AddLine(" uint32_t cur_block_dim = 1;");
4418 tiling_func_.AddLine(" uint32_t ori_block_dim = tiling_data.get_block_dim();");4493 tiling_func_.AddLine(" uint32_t ori_block_dim = tiling_data.get_block_dim();");
4494+ GenWorkspaceOffsetReset("tiling_data");
4419 tiling_func_.AddLine(" AutofuseTilingData tilingTmp;");4495 tiling_func_.AddLine(" AutofuseTilingData tilingTmp;");
4420 tiling_func_.AddLine(" tilingTmp = tiling_data;");4496 tiling_func_.AddLine(" tilingTmp = tiling_data;");
4497+ 
4498+ // Rebuild size aggregation base by calling each AscGraph::GetTiling to populate workspace values
4499+ // This is the critical fix to ensure each subgraph's workspace values are included in tilingTmp
4500+ for (const auto &asc_graph_namespace_map : namespace_map) {
4501+ const std::string &asc_graph_namespace = "AscGraph" + std::to_string(asc_graph_namespace_map.first);
4502+ tiling_func_.AddLine(" if (!" + asc_graph_namespace + "::GetTiling(tilingTmp, -1, nullptr)) {");
4503+ tiling_func_.AddLine(" OP_LOGW(OP_NAME, \"Failed to rebuild size base for " + asc_graph_namespace + ".\");");
4504+ tiling_func_.AddLine(" return false;");
4505+ tiling_func_.AddLine(" }");
4506+ }
4507+ 
4421 std::string block_dim_list_arg = "multi_group_block_dim_list";4508 std::string block_dim_list_arg = "multi_group_block_dim_list";
4422 GenPGOMultiGroupBlockDimList(namespace_map, block_dim_list_arg);4509 GenPGOMultiGroupBlockDimList(namespace_map, block_dim_list_arg);
4423 4510 
4511+ // Record baseline before search loop to identify new candidates
4512+ tiling_func_.AddLine(" const size_t ws_baseline = tiling_data_list.size();");
4513+ 
4424 for (const auto &asc_graph_namespace_map : namespace_map) {4514 for (const auto &asc_graph_namespace_map : namespace_map) {
4425 const std::string &asc_graph_namespace = "AscGraph" + std::to_string(asc_graph_namespace_map.first);4515 const std::string &asc_graph_namespace = "AscGraph" + std::to_string(asc_graph_namespace_map.first);
4426 tiling_func_.AddLine(" if (!" + asc_graph_namespace +4516 tiling_func_.AddLine(" if (!" + asc_graph_namespace +
@@ -4435,6 +4525,17 @@ af::Status TilingCodeGenImpl::GenPGOFusedScheduleResultsGetTilingDefine(const Fu
4435 tiling_func_.AddLine(" }");4525 tiling_func_.AddLine(" }");
4436 }4526 }
4437 4527 
4528+ // Finalize only the new candidates from this search iteration (skip baseline candidates)
4529+ // This prevents double-finalization of existing offset-semantic candidates
4530+ // 仅在存在 workspace tensor 时外抛:helper 定义由 GenWorkspaceOffsetHelpers 条件生成,
4531+ // 无 workspace 时调用点必须同步跳过,否则生成的 tiling func 编译失败。
4532+ if (!GetWorkspaceTensorIds(workspace_tensor_id_set_).empty()) {
4533+ tiling_func_.AddLine(" for (size_t i = ws_baseline; i < tiling_data_list.size(); ++i) {");
4534+ tiling_func_.AddLine(" FinalizeWorkspaceOffsets(tiling_data_list[i].tiling_data);");
4535+ tiling_func_.AddLine(" }");
4536+ }
4537+ 
4538+ GenWorkspaceOffsetFinalize("tiling_data");
4438 tiling_func_.AddLine(" OP_LOGI(OP_NAME, \"End PGOSearchTilingKey root.\");");4539 tiling_func_.AddLine(" OP_LOGI(OP_NAME, \"End PGOSearchTilingKey root.\");");
4439 tiling_func_.AddLine(" return true;");4540 tiling_func_.AddLine(" return true;");
4440 tiling_func_.AddLine("}");4541 tiling_func_.AddLine("}");
@@ -4535,6 +4636,7 @@ af::Status TilingCodeGenImpl::GenGetTilingForScheduleResult() {
4535 GE_ASSERT_SUCCESS(GenGetTilingForAllSchedulesResults(asc_graph_id, asc_graph_map),4636 GE_ASSERT_SUCCESS(GenGetTilingForAllSchedulesResults(asc_graph_id, asc_graph_map),
4536 "Generate GetTiling for all schedules results failed, asc_graph_id = %zu.", asc_graph_id);4637 "Generate GetTiling for all schedules results failed, asc_graph_id = %zu.", asc_graph_id);
4537 }4638 }
4639+ GenWorkspaceOffsetHelpers();
4538 GE_ASSERT_SUCCESS(GenEnableGroupParallelFunctions(namespace_map));4640 GE_ASSERT_SUCCESS(GenEnableGroupParallelFunctions(namespace_map));
4539 GE_ASSERT_SUCCESS(GenFusedScheduleResultsGetTilingDefine(namespace_map));4641 GE_ASSERT_SUCCESS(GenFusedScheduleResultsGetTilingDefine(namespace_map));
4540 GE_ASSERT_SUCCESS(GenIsStaticShape());4642 GE_ASSERT_SUCCESS(GenIsStaticShape());
@@ -127,6 +127,9 @@ class TilingCodeGenImpl {
127 const std::string &schedule_result_prefix);127 const std::string &schedule_result_prefix);
128 void GenGetScheduleResultTail(const std::map<size_t, std::pair<std::string, std::string>> &graph_info);128 void GenGetScheduleResultTail(const std::map<size_t, std::pair<std::string, std::string>> &graph_info);
129 void GenUpdateWorkspace(const size_t asc_graph_id, const size_t impl_graph_id);129 void GenUpdateWorkspace(const size_t asc_graph_id, const size_t impl_graph_id);
130+ void GenWorkspaceOffsetHelpers();
131+ void GenWorkspaceOffsetReset(const std::string &tiling_data_name);
132+ void GenWorkspaceOffsetFinalize(const std::string &tiling_data_name);
130 // 生成DoGroupTiling公共函数,支持首次Tiling和二次Tiling133 // 生成DoGroupTiling公共函数,支持首次Tiling和二次Tiling
131 af::Status GenDoGroupTilingFunction(const size_t asc_graph_id, const size_t impl_graph_id,134 af::Status GenDoGroupTilingFunction(const size_t asc_graph_id, const size_t impl_graph_id,
132 const std::map<size_t, std::pair<std::string, std::string>> &graph_info);135 const std::map<size_t, std::pair<std::string, std::string>> &graph_info);
@@ -1063,7 +1063,7 @@ std::string codegen::Tiler::BlockOutterAxisDefine() {
1063 << " - block_offset : " << this->block_dim.name << " + GetBlockNum() - block_offset;" << std::endl;1063 << " - block_offset : " << this->block_dim.name << " + GetBlockNum() - block_offset;" << std::endl;
1064 // block_dim范围在调用前校验了,此处不需要重复校验1064 // block_dim范围在调用前校验了,此处不需要重复校验
1065 } else {1065 } else {
1066- code << "if (" << this->block_dim.name << " >= " << tiling_data.name << "->block_dim) { " << std::endl1066+ code << "if (" << this->block_dim.name << " >= " << tiling_data.name << "->block_dim) {" << std::endl
1067 << " return;" << std::endl1067 << " return;" << std::endl
1068 << "}" << std::endl;1068 << "}" << std::endl;
1069 }1069 }
@@ -1075,8 +1075,7 @@ std::string codegen::Tiler::BlockOutterAxisDefine() {
1075 1075 
1076 stringstream axis_value;1076 stringstream axis_value;
1077 axis_value << this->block_dim.name << " % " << axis.loop_size;1077 axis_value << this->block_dim.name << " % " << axis.loop_size;
1078- code << axis.Define(axis_value.str(), true);1078+ code << axis.Define(axis_value.str(), true) << std::endl;
1079- code << " " << std::endl;
1080 if (axis.from.size() > 1) {1079 if (axis.from.size() > 1) {
1081 BlockOutterAxisDefine(id, code);1080 BlockOutterAxisDefine(id, code);
1082 }1081 }
@@ -2093,8 +2092,6 @@ Status Kernel::AppendWorkspaceTensorInit(std::stringstream &ss,
2093 ss << workspace_buffer_arg.AsArg() << " = " << workspace_buffer_arg_override.c_str() << ";" << std::endl;2092 ss << workspace_buffer_arg.AsArg() << " = " << workspace_buffer_arg_override.c_str() << ";" << std::endl;
2094 }2093 }
2095 2094 
2096- std::stringstream offset_ss;
2097- offset_ss << "0";
2098 auto it_ws_tensors = this->workspace_tensors.begin();2095 auto it_ws_tensors = this->workspace_tensors.begin();
2099 for (size_t i = 0UL; i < this->workspaces.size(); i++) {2096 for (size_t i = 0UL; i < this->workspaces.size(); i++) {
2100 GELOGI("Define workspace tensor id: %ld", it_ws_tensors->first);2097 GELOGI("Define workspace tensor id: %ld", it_ws_tensors->first);
@@ -2105,11 +2102,12 @@ Status Kernel::AppendWorkspaceTensorInit(std::stringstream &ss,
2105 }2102 }
2106 2103 
2107 ss << tensor->second.Define() << std::endl;2104 ss << tensor->second.Define() << std::endl;
2105+ std::stringstream offset_ss;
2106+ offset_ss << this->workspaces[i];
2108 std::string local_result;2107 std::string local_result;
2109 GE_CHK_STATUS_RET(tensor->second.SetGlobalBuffer(workspace_buffer_arg, offset_ss.str(), local_result),2108 GE_CHK_STATUS_RET(tensor->second.SetGlobalBuffer(workspace_buffer_arg, offset_ss.str(), local_result),
2110 "Codegen set global buffer failed");2109 "Codegen set global buffer failed");
2111 ss << local_result << std::endl;2110 ss << local_result << std::endl;
2112- offset_ss << " + " << "(" << this->workspaces[i] << ")";
2113 it_ws_tensors++;2111 it_ws_tensors++;
2114 }2112 }
2115 return af::SUCCESS;2113 return af::SUCCESS;
@@ -2237,7 +2235,7 @@ Status Kernel::ParseGraph(const ascir::ImplGraph &graph, const ascir::FusedSched
2237 kernel.output_tensors.emplace_back(pair.second.second);2235 kernel.output_tensors.emplace_back(pair.second.second);
2238 }2236 }
2239 2237 
2240- std::vector<ascir::TensorId> workspace_tensor_id = GetWorkspaceTensorIdListInOneScheduleResult(fused_schedule_result);2238+ std::vector<ascir::TensorId> workspace_tensor_id = GetWorkspaceTensorIdListInOneGraph(fused_schedule_result, graph);
2241 for (auto tId : workspace_tensor_id) {2239 for (auto tId : workspace_tensor_id) {
2242 std::string workspaceStr = "workspace";2240 std::string workspaceStr = "workspace";
2243 workspaceStr = workspaceStr + std::to_string(tId);2241 workspaceStr = workspaceStr + std::to_string(tId);
@@ -3075,6 +3073,10 @@ Status Kernel::GenKernelFuncWithParseTilingData(const ascir::FusedScheduledResul
3075 } else {3073 } else {
3076 GE_ASSERT_SUCCESS(GenMulGroupKernelWithParseTilingData(fused_schedule_result, graph_id, config, ss, ss1,3074 GE_ASSERT_SUCCESS(GenMulGroupKernelWithParseTilingData(fused_schedule_result, graph_id, config, ss, ss1,
3077 use_list_tensor, kernel_file_ptr));3075 use_list_tensor, kernel_file_ptr));
3076+ ss1 << std::endl;
3077+ if ((graph_id + 1) < fused_schedule_result.node_idx_to_scheduled_results.size()) {
3078+ ss1 << " SyncAll();" << std::endl;
3079+ }
3078 }3080 }
3079 }3081 }
3080 return af::SUCCESS;3082 return af::SUCCESS;
@@ -1084,8 +1084,9 @@ void TilingLib::TilingSetShapeDim(std::stringstream &tiling_set_shape_dim, const
1084 const std::string &tiling_expr) const {1084 const std::string &tiling_expr) const {
1085 for (size_t i = 0; i < fused_schedule_result.node_idx_to_scheduled_results.size(); i++) {1085 for (size_t i = 0; i < fused_schedule_result.node_idx_to_scheduled_results.size(); i++) {
1086 auto scheduled_results = fused_schedule_result.node_idx_to_scheduled_results[i];1086 auto scheduled_results = fused_schedule_result.node_idx_to_scheduled_results[i];
1087- if ((scheduled_results.empty()) ||1087+ if ((fused_schedule_result.node_idx_to_scheduled_results.size() == 1) &&
1088- ((scheduled_results.size() == 1) && (scheduled_results[0].schedule_groups.size() == 1))) {1088+ (scheduled_results.empty() ||
1089+ ((scheduled_results.size() == 1) && (scheduled_results[0].schedule_groups.size() == 1)))) {
1089 // 检查变量是否被此 schedule_group 使用1090 // 检查变量是否被此 schedule_group 使用
1090 if (!IsVarUsedInScheduleGroup(var_define, scheduled_results[0].schedule_groups[0])) {1091 if (!IsVarUsedInScheduleGroup(var_define, scheduled_results[0].schedule_groups[0])) {
1091 continue;1092 continue;
@@ -681,6 +681,32 @@ std::vector<ascir::TensorId> GetWorkspaceTensorIdListInOneScheduleResult(
681 return tensorId;681 return tensorId;
682}682}
683 683 
684+std::vector<ascir::TensorId> GetWorkspaceTensorIdListInOneGraph(
685+ const ascir::FusedScheduledResult &fused_schedule_result, const ascir::ImplGraph &graph) {
686+ std::vector<ascir::TensorId> tensorId;
687+ for (auto workspace : fused_schedule_result.workspace_nodes) {
688+ GE_ASSERT_NOTNULL(workspace, "fused schedule result workspace node is null");
689+ ascir::TensorId tId = workspace->outputs[0].attr.mem.tensor_id;
690+ GELOGI("Get workspace tensor id: %ld", tId);
691+ // 检查是否在 graph 中存在
692+ bool is_exist = false;
693+ for (auto node : graph.GetAllNodes()) {
694+ if (IsOps<Workspace>(node)) {
695+ ascir::TensorId tensor_id = node->outputs[0].attr.mem.tensor_id;
696+ if (tId == tensor_id) {
697+ is_exist = true;
698+ break;
699+ }
700+ }
701+ }
702+ auto index = std::find(tensorId.begin(), tensorId.end(), tId);
703+ if (index == tensorId.end() && is_exist) {
704+ tensorId.emplace_back(tId);
705+ }
706+ }
707+ return tensorId;
708+}
709+ 
684bool IsScalarNextNodeSupportBlkTensor(const af::AscNodePtr &node) {710bool IsScalarNextNodeSupportBlkTensor(const af::AscNodePtr &node) {
685 for (auto &out : node->outputs()) {711 for (auto &out : node->outputs()) {
686 for (auto &peer_input : out->anchor.GetPeerInDataAnchors()) {712 for (auto &peer_input : out->anchor.GetPeerInDataAnchors()) {
@@ -735,6 +761,9 @@ bool IsSingleGroup(const ascir::FusedScheduledResult &fused_schedule_result) {
735}761}
736 762 
737bool CanUseTilingKey(const ascir::FusedScheduledResult &fused_schedule_result) {763bool CanUseTilingKey(const ascir::FusedScheduledResult &fused_schedule_result) {
764+ if (fused_schedule_result.node_idx_to_scheduled_results.size() > 1U) {
765+ return false;
766+ }
738 for (const auto &schedule_result_list : fused_schedule_result.node_idx_to_scheduled_results) {767 for (const auto &schedule_result_list : fused_schedule_result.node_idx_to_scheduled_results) {
739 for (const auto &schedule_result : schedule_result_list) {768 for (const auto &schedule_result : schedule_result_list) {
740 if (schedule_result.enable_group_parallel) {769 if (schedule_result.enable_group_parallel) {
@@ -156,6 +156,8 @@ af::Expression CalculateWorkspaceSize(const std::vector<af::AscNodePtr> &workspa
156af::Expression CalcExtraTmpBufForAscGraph(const ascir::ImplGraph &graph);156af::Expression CalcExtraTmpBufForAscGraph(const ascir::ImplGraph &graph);
157std::vector<ascir::TensorId> GetWorkspaceTensorIdListInOneScheduleResult(157std::vector<ascir::TensorId> GetWorkspaceTensorIdListInOneScheduleResult(
158 const ascir::FusedScheduledResult &fused_schedule_result);158 const ascir::FusedScheduledResult &fused_schedule_result);
159+std::vector<ascir::TensorId> GetWorkspaceTensorIdListInOneGraph(
160+ const ascir::FusedScheduledResult &fused_schedule_result, const ascir::ImplGraph &graph);
159 161 
160af::Status GetApiTilingTypeName(const ascir::NodeView &node, std::string &type_name);162af::Status GetApiTilingTypeName(const ascir::NodeView &node, std::string &type_name);
161af::Status GetApiTilingFieldName(const ascir::NodeView &node, std::string &field_name);163af::Status GetApiTilingFieldName(const ascir::NodeView &node, std::string &field_name);
@@ -241,6 +241,7 @@ struct ShareGraph {
241 static af::ComputeGraphPtr VfScalarFusionComprehensiveFusedGraph();241 static af::ComputeGraphPtr VfScalarFusionComprehensiveFusedGraph();
242 static af::ComputeGraphPtr RemainderBf16FusedGraph(size_t dims_size);242 static af::ComputeGraphPtr RemainderBf16FusedGraph(size_t dims_size);
243 static af::ComputeGraphPtr ArgMaxFusedGraph(size_t dims_size);243 static af::ComputeGraphPtr ArgMaxFusedGraph(size_t dims_size);
244+ static af::ComputeGraphPtr FusedBackendElewiseGraph(size_t dims_size);
244};245};
245} // namespace ascir246} // namespace ascir
246#endif247#endif
@@ -13440,4 +13440,228 @@ af::ComputeGraphPtr ShareGraph::RandnStoreFusedGraph(size_t dims_size) {
13440 return compute_graph;13440 return compute_graph;
13441}13441}
13442 13442 
13443+static void CreateAscBackendGraphTwoInTwoOut(std::shared_ptr<af::AscGraph> &graph, const std::string &prefix,
13444+ int64_t axis_num = 2) {
13445+ auto ONE = af::Symbol(1);
13446+ std::vector<int64_t> axis_ids;
13447+ std::vector<af::Expression> repeats;
13448+ for (int64_t i = 0; i < axis_num; ++i) {
13449+ const af::Expression exp = graph->CreateSizeVar("s" + std::to_string(i));
13450+ auto axis = graph->CreateAxis("z" + std::to_string(i), exp);
13451+ axis_ids.push_back(i);
13452+ repeats.push_back(exp);
13453+ }
13454+ 
13455+ std::vector<af::Expression> strides(repeats.size(), af::ops::One);
13456+ if (axis_num > 1) {
13457+ for (int64_t i = axis_num - 2; i >= 0; --i) {
13458+ strides[i] = repeats[i + 1] * strides[i + 1];
13459+ }
13460+ }
13461+ 
13462+ af::ascir_op::Data data0(std::string(prefix + "_data0").c_str(), *graph);
13463+ data0.attr.sched.axis = axis_ids;
13464+ *data0.y.axis = axis_ids;
13465+ *data0.y.repeats = repeats;
13466+ *data0.y.strides = strides;
13467+ data0.ir_attr.SetIndex(0);
13468+ data0.y.dtype = ge::DT_FLOAT;
13469+ 
13470+ af::ascir_op::Load load0(std::string(prefix + "_load0").c_str());
13471+ load0.x = data0.y;
13472+ load0.attr.sched.axis = axis_ids;
13473+ *load0.y.axis = axis_ids;
13474+ *load0.y.repeats = repeats;
13475+ *load0.y.strides = strides;
13476+ 
13477+ af::ascir_op::Data data1(std::string(prefix + "_data1").c_str(), *graph);
13478+ data1.attr.sched.axis = axis_ids;
13479+ *data1.y.axis = axis_ids;
13480+ *data1.y.repeats = repeats;
13481+ *data1.y.strides = strides;
13482+ data1.ir_attr.SetIndex(1);
13483+ data1.y.dtype = ge::DT_FLOAT;
13484+ 
13485+ af::ascir_op::Load load1(std::string(prefix + "_load1").c_str());
13486+ load1.x = data1.y;
13487+ load1.attr.sched.axis = axis_ids;
13488+ *load1.y.axis = axis_ids;
13489+ *load1.y.repeats = repeats;
13490+ *load1.y.strides = strides;
13491+ 
13492+ af::ascir_op::Add add(std::string(prefix + "_add").c_str());
13493+ add.x1 = load0.y;
13494+ add.x2 = load1.y;
13495+ add.attr.sched.axis = axis_ids;
13496+ *add.y.axis = axis_ids;
13497+ *add.y.repeats = repeats;
13498+ *add.y.strides = strides;
13499+ 
13500+ af::ascir_op::Store store0(std::string(prefix + "_store0").c_str());
13501+ store0.x = add.y;
13502+ store0.attr.sched.axis = axis_ids;
13503+ *store0.y.axis = axis_ids;
13504+ *store0.y.repeats = repeats;
13505+ *store0.y.strides = strides;
13506+ 
13507+ af::ascir_op::Output y0(std::string(prefix + "_out0").c_str());
13508+ y0.x = store0.y;
13509+ y0.ir_attr.SetIndex(0);
13510+ y0.y.dtype = ge::DT_FLOAT;
13511+ 
13512+ af::ascir_op::Store store1(std::string(prefix + "_store1").c_str());
13513+ store1.x = add.y;
13514+ store1.attr.sched.axis = axis_ids;
13515+ *store1.y.axis = axis_ids;
13516+ *store1.y.repeats = repeats;
13517+ *store1.y.strides = strides;
13518+ 
13519+ af::ascir_op::Output y1(std::string(prefix + "_out1").c_str());
13520+ y1.x = store1.y;
13521+ y1.ir_attr.SetIndex(1);
13522+ y1.y.dtype = ge::DT_FLOAT;
13523+}
13524+ 
13525+static void CreateAscBackendGraphTwoInOneOut(std::shared_ptr<af::AscGraph> &graph, const std::string &prefix,
13526+ int64_t axis_num = 2) {
13527+ auto ONE = af::Symbol(1);
13528+ std::vector<int64_t> axis_ids;
13529+ std::vector<af::Expression> repeats;
13530+ for (int64_t i = 0; i < axis_num; ++i) {
13531+ const af::Expression exp = graph->CreateSizeVar("s" + std::to_string(i));
13532+ auto axis = graph->CreateAxis("z" + std::to_string(i), exp);
13533+ axis_ids.push_back(i);
13534+ repeats.push_back(exp);
13535+ }
13536+ 
13537+ std::vector<af::Expression> strides(repeats.size(), af::ops::One);
13538+ if (axis_num > 1) {
13539+ for (int64_t i = axis_num - 2; i >= 0; --i) {
13540+ strides[i] = repeats[i + 1] * strides[i + 1];
13541+ }
13542+ }
13543+ 
13544+ af::ascir_op::Data data0(std::string(prefix + "_data0").c_str(), *graph);
13545+ data0.attr.sched.axis = axis_ids;
13546+ *data0.y.axis = axis_ids;
13547+ *data0.y.repeats = repeats;
13548+ *data0.y.strides = strides;
13549+ data0.ir_attr.SetIndex(0);
13550+ data0.y.dtype = ge::DT_FLOAT;
13551+ 
13552+ af::ascir_op::Load load0(std::string(prefix + "_load0").c_str());
13553+ load0.x = data0.y;
13554+ load0.attr.sched.axis = axis_ids;
13555+ *load0.y.axis = axis_ids;
13556+ *load0.y.repeats = repeats;
13557+ *load0.y.strides = strides;
13558+ 
13559+ af::ascir_op::Data data1(std::string(prefix + "_data1").c_str(), *graph);
13560+ data1.attr.sched.axis = axis_ids;
13561+ *data1.y.axis = axis_ids;
13562+ *data1.y.repeats = repeats;
13563+ *data1.y.strides = strides;
13564+ data1.ir_attr.SetIndex(1);
13565+ data1.y.dtype = ge::DT_FLOAT;
13566+ 
13567+ af::ascir_op::Load load1(std::string(prefix + "_load1").c_str());
13568+ load1.x = data1.y;
13569+ load1.attr.sched.axis = axis_ids;
13570+ *load1.y.axis = axis_ids;
13571+ *load1.y.repeats = repeats;
13572+ *load1.y.strides = strides;
13573+ 
13574+ af::ascir_op::Add add(std::string(prefix + "_add").c_str());
13575+ add.x1 = load0.y;
13576+ add.x2 = load1.y;
13577+ add.attr.sched.axis = axis_ids;
13578+ *add.y.axis = axis_ids;
13579+ *add.y.repeats = repeats;
13580+ *add.y.strides = strides;
13581+ 
13582+ af::ascir_op::Store store0(std::string(prefix + "_store0").c_str());
13583+ store0.x = add.y;
13584+ store0.attr.sched.axis = axis_ids;
13585+ *store0.y.axis = axis_ids;
13586+ *store0.y.repeats = repeats;
13587+ *store0.y.strides = strides;
13588+ 
13589+ af::ascir_op::Output y0(std::string(prefix + "_out0").c_str());
13590+ y0.x = store0.y;
13591+ y0.ir_attr.SetIndex(0);
13592+ y0.y.dtype = ge::DT_FLOAT;
13593+}
13594+ 
13595+static NodePtr CreateAscbcToAscGraph(const std::string &name, ComputeGraphPtr &compute_graph, int64_t in_num = 1,
13596+ int64_t out_num = 1) {
13597+ OpDescBuilder op_desc_builder(name, "AscBackend");
13598+ op_desc_builder.AddDynamicInput("x", in_num);
13599+ op_desc_builder.AddDynamicOutput("y", out_num);
13600+ const auto &op_desc = op_desc_builder.Build();
13601+ auto node = compute_graph->AddNode(op_desc);
13602+ node->SetOwnerComputeGraph(compute_graph);
13603+ return node;
13604+}
13605+ 
13606+af::ComputeGraphPtr ShareGraph::FusedBackendElewiseGraph(size_t dims_size) {
13607+ std::shared_ptr<af::AscGraph> g0 = std::make_shared<af::AscGraph>("g0");
13608+ CreateAscBackendGraphTwoInTwoOut(g0, "g0", dims_size);
13609+ std::shared_ptr<af::AscGraph> g1 = std::make_shared<af::AscGraph>("g1");
13610+ CreateAscBackendGraphTwoInOneOut(g1, "g1", dims_size);
13611+ std::shared_ptr<af::AscGraph> g2 = std::make_shared<af::AscGraph>("g2");
13612+ CreateAscBackendGraphTwoInOneOut(g2, "g2", dims_size);
13613+ 
13614+ af::AscGraph fused_asc_graph("fused_backend_elewise_test");
13615+ af::ascir_op::Data data0("data0", fused_asc_graph);
13616+ auto ir_attr0 = data0.attr.ir_attr->DownCastTo<af::AscDataIrAttrDef>();
13617+ ir_attr0->SetIndex(0);
13618+ 
13619+ af::ascir_op::Data data1("data1", fused_asc_graph);
13620+ auto ir_attr1 = data1.attr.ir_attr->DownCastTo<af::AscDataIrAttrDef>();
13621+ ir_attr1->SetIndex(1);
13622+ 
13623+ af::ascir_op::Data data2("data2", fused_asc_graph);
13624+ auto ir_attr2 = data2.attr.ir_attr->DownCastTo<af::AscDataIrAttrDef>();
13625+ ir_attr2->SetIndex(2);
13626+ 
13627+ auto fused_graph = af::AscGraphUtils::GetComputeGraph(fused_asc_graph);
13628+ auto data0_node = fused_asc_graph.FindNode("data0");
13629+ auto data1_node = fused_asc_graph.FindNode("data1");
13630+ auto data2_node = fused_asc_graph.FindNode("data2");
13631+ 
13632+ auto ascbc1 = CreateAscbcToAscGraph("ascbc1", fused_graph, 2, 2);
13633+ auto ascbc2 = CreateAscbcToAscGraph("ascbc2", fused_graph, 2, 1);
13634+ auto ascbc3 = CreateAscbcToAscGraph("ascbc3", fused_graph, 2, 1);
13635+ 
13636+ af::GraphUtils::AddEdge(data0_node->GetOutDataAnchor(0), ascbc1->GetInDataAnchor(0));
13637+ af::GraphUtils::AddEdge(data1_node->GetOutDataAnchor(0), ascbc1->GetInDataAnchor(1));
13638+ af::GraphUtils::AddEdge(data2_node->GetOutDataAnchor(0), ascbc2->GetInDataAnchor(0));
13639+ af::GraphUtils::AddEdge(ascbc1->GetOutDataAnchor(0), ascbc2->GetInDataAnchor(1));
13640+ af::GraphUtils::AddEdge(ascbc2->GetOutDataAnchor(0), ascbc3->GetInDataAnchor(0));
13641+ af::GraphUtils::AddEdge(ascbc1->GetOutDataAnchor(1), ascbc3->GetInDataAnchor(1));
13642+ 
13643+ af::ascir_op::Output output0("output0");
13644+ auto out0_ir_attr = output0.attr.ir_attr->DownCastTo<af::AscDataIrAttrDef>();
13645+ out0_ir_attr->SetIndex(0);
13646+ auto out0_desc = OpDescUtils::GetOpDescFromOperator(output0);
13647+ auto output0_node = fused_graph->AddNode(out0_desc);
13648+ 
13649+ af::ascir_op::Output output1("output1");
13650+ auto out1_ir_attr = output1.attr.ir_attr->DownCastTo<af::AscDataIrAttrDef>();
13651+ out1_ir_attr->SetIndex(1);
13652+ auto out1_desc = OpDescUtils::GetOpDescFromOperator(output1);
13653+ auto output1_node = fused_graph->AddNode(out1_desc);
13654+ af::GraphUtils::AddEdge(ascbc3->GetOutDataAnchor(0), output0_node->GetInDataAnchor(0));
13655+ af::GraphUtils::AddEdge(ascbc1->GetOutDataAnchor(1), output1_node->GetInDataAnchor(0));
13656+ 
13657+ auto fuse1_attrs = ascbc1->GetOpDesc()->GetOrCreateAttrsGroup<ge::AutoFuseAttrs>();
13658+ fuse1_attrs->SetAscGraph(g0);
13659+ auto fuse2_attrs = ascbc2->GetOpDesc()->GetOrCreateAttrsGroup<ge::AutoFuseAttrs>();
13660+ fuse2_attrs->SetAscGraph(g1);
13661+ auto fuse3_attrs = ascbc3->GetOpDesc()->GetOrCreateAttrsGroup<ge::AutoFuseAttrs>();
13662+ fuse3_attrs->SetAscGraph(g2);
13663+ fused_graph->TopologicalSorting();
13664+ return fused_graph;
13665+}
13666+ 
13443} // namespace ascir13667} // namespace ascir
@@ -1023,7 +1023,9 @@ TEST_F(TestGenModelInfo, gen_workspace_with_tensor_id) {
1023 EXPECT_EQ(GenTilingImplAutoFuseV3("FlashSoftmax", fused_schedule_result, options, tiling_funcs, true), true);1023 EXPECT_EQ(GenTilingImplAutoFuseV3("FlashSoftmax", fused_schedule_result, options, tiling_funcs, true), true);
1024 std::string tiling_func;1024 std::string tiling_func;
1025 CombineTilings(tiling_funcs, tiling_func);1025 CombineTilings(tiling_funcs, tiling_func);
1026- EXPECT_NE(tiling_func.find("tiling_data.set_workspace0(it0->second);"), std::string::npos);1026+ EXPECT_NE(tiling_func.find("tiling_data.set_workspace0(std::max(tiling_data.get_workspace0(), "
1027+ "static_cast<uint32_t>(it0->second)));"),
1028+ std::string::npos);
1027}1029}
1028 1030 
1029TEST_F(TestGenModelInfo, gen_schedule_group_reduce_tile_r) {1031TEST_F(TestGenModelInfo, gen_schedule_group_reduce_tile_r) {
@@ -261,12 +261,12 @@ TEST(CodegenKernel, Tiler_BlockOutterAxisDefine) {
261 261 
262 auto result_code = tiler.BlockOutterAxisDefine();262 auto result_code = tiler.BlockOutterAxisDefine();
263 EXPECT_EQ(result_code, std::string{"int block_dim = GetBlockIdx();\n"263 EXPECT_EQ(result_code, std::string{"int block_dim = GetBlockIdx();\n"
264- "if (block_dim >= t->block_dim) { \n"264+ "if (block_dim >= t->block_dim) {\n"
265 " return;\n"265 " return;\n"
266 "}\n"266 "}\n"
267- "const int z0 = block_dim % z0_loop_size; \n"267+ "const int z0 = block_dim % z0_loop_size;\n"
268- "const int z1 = block_dim % z1_loop_size; \n"268+ "const int z1 = block_dim % z1_loop_size;\n"
269- "const int z2 = block_dim % z2_loop_size; \n"});269+ "const int z2 = block_dim % z2_loop_size;\n"});
270}270}
271 271 
272TEST(CodegenKernel, Tiler_GetAxisVar) {272TEST(CodegenKernel, Tiler_GetAxisVar) {
@@ -3817,18 +3817,17 @@ TEST(CodegenKernel, TwoWorkspaceCodegen) {
3817 codegen::Kernel::ParseGraph(graph, fused_schedule_result, kernel);3817 codegen::Kernel::ParseGraph(graph, fused_schedule_result, kernel);
3818 std::string result;3818 std::string result;
3819 kernel.GlobalTensorInit(result);3819 kernel.GlobalTensorInit(result);
3820- EXPECT_EQ(3820+ EXPECT_EQ(result,
3821- result,3821+ std::string{"GlobalTensor<half> global_0;\n"
3822- std::string{"GlobalTensor<half> global_0;\n"3822+ "global_0.SetGlobalBuffer((__gm__ half*)x);\n"
3823- "global_0.SetGlobalBuffer((__gm__ half*)x);\n"3823+ "GlobalTensor<half> global_3;\n"
3824- "GlobalTensor<half> global_3;\n"3824+ "global_3.SetGlobalBuffer((__gm__ half*)y1);\n"
3825- "global_3.SetGlobalBuffer((__gm__ half*)y1);\n"3825+ "GlobalTensor<half> global_4;\n"
3826- "GlobalTensor<half> global_4;\n"3826+ "global_4.SetGlobalBuffer((__gm__ half*)y2);\n"
3827- "global_4.SetGlobalBuffer((__gm__ half*)y2);\n"3827+ "GlobalTensor<half> global_1;\n"
3828- "GlobalTensor<half> global_1;\n"3828+ "global_1.SetGlobalBuffer((__gm__ half*)((__gm__ uint8_t*)(workspace) + (workspace1)));\n"
3829- "global_1.SetGlobalBuffer((__gm__ half*)workspace);\n"3829+ "GlobalTensor<half> global_2;\n"
3830- "GlobalTensor<half> global_2;\n"3830+ "global_2.SetGlobalBuffer((__gm__ half*)((__gm__ uint8_t*)(workspace) + (workspace2)));\n"});
3831- "global_2.SetGlobalBuffer((__gm__ half*)((__gm__ uint8_t*)(workspace) + (0 + (workspace1))));\n"});
3832}3831}
3833 3832 
3834TEST(CodegenKernel, TwoWorkspaceReuseAsInputCodegen) {3833TEST(CodegenKernel, TwoWorkspaceReuseAsInputCodegen) {
@@ -3923,7 +3922,8 @@ TEST(CodegenKernel, TwoWorkspaceReuseAsInputCodegen) {
3923 "GlobalTensor<half> global_4;\n"3922 "GlobalTensor<half> global_4;\n"
3924 "global_4.SetGlobalBuffer((__gm__ half*)y2);\n"3923 "global_4.SetGlobalBuffer((__gm__ half*)y2);\n"
3925 "GlobalTensor<half> global_1;\n"3924 "GlobalTensor<half> global_1;\n"
3926- "global_1.SetGlobalBuffer((__gm__ half*)workspace);\n"});3925+ "global_1.SetGlobalBuffer((__gm__ half*)((__gm__ uint8_t*)(workspace) + "
3926+ "(workspace1)));\n"});
3927}3927}
3928 3928 
3929TEST(CodegenKernel, GlobalTensorInitShouldUseWorkspaceReuseOverride) {3929TEST(CodegenKernel, GlobalTensorInitShouldUseWorkspaceReuseOverride) {
@@ -3974,7 +3974,8 @@ TEST(CodegenKernel, GlobalTensorInitShouldUseWorkspaceReuseOverride) {
3974 "global_0.SetGlobalBuffer((__gm__ half*)x);\n"3974 "global_0.SetGlobalBuffer((__gm__ half*)x);\n"
3975 "GM_ADDR workspace_reuse = y0;\n"3975 "GM_ADDR workspace_reuse = y0;\n"
3976 "GlobalTensor<half> global_1;\n"3976 "GlobalTensor<half> global_1;\n"
3977- "global_1.SetGlobalBuffer((__gm__ half*)workspace_reuse);\n"});3977+ "global_1.SetGlobalBuffer((__gm__ half*)((__gm__ uint8_t*)(workspace_reuse) + "
3978+ "(workspace1)));\n"});
3978}3979}
3979 3980 
3980TEST(CodegenKernel, GenCubeCommonTilingSingleFuncCallShouldUseOutputOverride) {3981TEST(CodegenKernel, GenCubeCommonTilingSingleFuncCallShouldUseOutputOverride) {
@@ -4638,7 +4639,7 @@ TEST(CodegenKernel, Kernel_GenerateKernel_Multi_ScheduleGroup) {
4638 "inline __aicore__ void test_kernel_general_0_nil_0_nil(GM_ADDR x, GM_ADDR y, GM_ADDR workspace, const "4639 "inline __aicore__ void test_kernel_general_0_nil_0_nil(GM_ADDR x, GM_ADDR y, GM_ADDR workspace, const "
4639 "AscGraph0ScheduleResult0G0TilingData *t) {\n"4640 "AscGraph0ScheduleResult0G0TilingData *t) {\n"
4640 "int block_dim = GetBlockIdx();\n"4641 "int block_dim = GetBlockIdx();\n"
4641- "if (block_dim >= t->block_dim) { \n"4642+ "if (block_dim >= t->block_dim) {\n"
4642 " return;\n"4643 " return;\n"
4643 "}\n\n"4644 "}\n\n"
4644 "GlobalTensor<half> global_0;\n"4645 "GlobalTensor<half> global_0;\n"
@@ -4659,7 +4660,7 @@ TEST(CodegenKernel, Kernel_GenerateKernel_Multi_ScheduleGroup) {
4659 "inline __aicore__ void test_kernel_general_1_nil_1_nil(GM_ADDR x, GM_ADDR y, GM_ADDR workspace, const "4660 "inline __aicore__ void test_kernel_general_1_nil_1_nil(GM_ADDR x, GM_ADDR y, GM_ADDR workspace, const "
4660 "AscGraph0ScheduleResult0G0TilingData *t) {\n"4661 "AscGraph0ScheduleResult0G0TilingData *t) {\n"
4661 "int block_dim = GetBlockIdx();\n"4662 "int block_dim = GetBlockIdx();\n"
4662- "if (block_dim >= t->block_dim) { \n"4663+ "if (block_dim >= t->block_dim) {\n"
4663 " return;\n"4664 " return;\n"
4664 "}\n\n"4665 "}\n\n"
4665 "GlobalTensor<half> global_0;\n"4666 "GlobalTensor<half> global_0;\n"
@@ -4680,7 +4681,7 @@ TEST(CodegenKernel, Kernel_GenerateKernel_Multi_ScheduleGroup) {
4680 "inline __aicore__ void test_kernel_general_2_nil_2_nil(GM_ADDR x, GM_ADDR y, GM_ADDR workspace, const "4681 "inline __aicore__ void test_kernel_general_2_nil_2_nil(GM_ADDR x, GM_ADDR y, GM_ADDR workspace, const "
4681 "AscGraph0ScheduleResult0G0TilingData *t) {\n"4682 "AscGraph0ScheduleResult0G0TilingData *t) {\n"
4682 "int block_dim = GetBlockIdx();\n"4683 "int block_dim = GetBlockIdx();\n"
4683- "if (block_dim >= t->block_dim) { \n"4684+ "if (block_dim >= t->block_dim) {\n"
4684 " return;\n"4685 " return;\n"
4685 "}\n\n"4686 "}\n\n"
4686 "GlobalTensor<half> global_0;\n"4687 "GlobalTensor<half> global_0;\n"
@@ -4701,7 +4702,7 @@ TEST(CodegenKernel, Kernel_GenerateKernel_Multi_ScheduleGroup) {
4701 "inline __aicore__ void test_kernel_general_3_nil_3_nil(GM_ADDR x, GM_ADDR y, GM_ADDR workspace, const "4702 "inline __aicore__ void test_kernel_general_3_nil_3_nil(GM_ADDR x, GM_ADDR y, GM_ADDR workspace, const "
4702 "AscGraph0ScheduleResult0G1TilingData *t) {\n"4703 "AscGraph0ScheduleResult0G1TilingData *t) {\n"
4703 "int block_dim = GetBlockIdx();\n"4704 "int block_dim = GetBlockIdx();\n"
4704- "if (block_dim >= t->block_dim) { \n"4705+ "if (block_dim >= t->block_dim) {\n"
4705 " return;\n"4706 " return;\n"
4706 "}\n\n"4707 "}\n\n"
4707 "GlobalTensor<half> global_0;\n"4708 "GlobalTensor<half> global_0;\n"
@@ -4722,7 +4723,7 @@ TEST(CodegenKernel, Kernel_GenerateKernel_Multi_ScheduleGroup) {
4722 "inline __aicore__ void test_kernel_general_4_nil_4_nil(GM_ADDR x, GM_ADDR y, GM_ADDR workspace, const "4723 "inline __aicore__ void test_kernel_general_4_nil_4_nil(GM_ADDR x, GM_ADDR y, GM_ADDR workspace, const "
4723 "AscGraph0ScheduleResult0G1TilingData *t) {\n"4724 "AscGraph0ScheduleResult0G1TilingData *t) {\n"
4724 "int block_dim = GetBlockIdx();\n"4725 "int block_dim = GetBlockIdx();\n"
4725- "if (block_dim >= t->block_dim) { \n"4726+ "if (block_dim >= t->block_dim) {\n"
4726 " return;\n"4727 " return;\n"
4727 "}\n\n"4728 "}\n\n"
4728 "GlobalTensor<half> global_0;\n"4729 "GlobalTensor<half> global_0;\n"
@@ -4743,7 +4744,7 @@ TEST(CodegenKernel, Kernel_GenerateKernel_Multi_ScheduleGroup) {
4743 "inline __aicore__ void test_kernel_general_5_nil_5_nil(GM_ADDR x, GM_ADDR y, GM_ADDR workspace, const "4744 "inline __aicore__ void test_kernel_general_5_nil_5_nil(GM_ADDR x, GM_ADDR y, GM_ADDR workspace, const "
4744 "AscGraph0ScheduleResult1G0TilingData *t) {\n"4745 "AscGraph0ScheduleResult1G0TilingData *t) {\n"
4745 "int block_dim = GetBlockIdx();\n"4746 "int block_dim = GetBlockIdx();\n"
4746- "if (block_dim >= t->block_dim) { \n"4747+ "if (block_dim >= t->block_dim) {\n"
4747 " return;\n"4748 " return;\n"
4748 "}\n\n"4749 "}\n\n"
4749 "GlobalTensor<half> global_0;\n"4750 "GlobalTensor<half> global_0;\n"
@@ -4764,7 +4765,7 @@ TEST(CodegenKernel, Kernel_GenerateKernel_Multi_ScheduleGroup) {
4764 "inline __aicore__ void test_kernel_general_6_nil_6_nil(GM_ADDR x, GM_ADDR y, GM_ADDR workspace, const "4765 "inline __aicore__ void test_kernel_general_6_nil_6_nil(GM_ADDR x, GM_ADDR y, GM_ADDR workspace, const "
4765 "AscGraph0ScheduleResult1G0TilingData *t) {\n"4766 "AscGraph0ScheduleResult1G0TilingData *t) {\n"
4766 "int block_dim = GetBlockIdx();\n"4767 "int block_dim = GetBlockIdx();\n"
4767- "if (block_dim >= t->block_dim) { \n"4768+ "if (block_dim >= t->block_dim) {\n"
4768 " return;\n"4769 " return;\n"
4769 "}\n\n"4770 "}\n\n"
4770 "GlobalTensor<half> global_0;\n"4771 "GlobalTensor<half> global_0;\n"
@@ -4953,7 +4954,7 @@ TEST(CodegenKernel, Kernel_GenerateKernel_Single_ScheduleGroup) {
4953 "inline __aicore__ void test_kernel_general_0_nil_0_nil(GM_ADDR x, GM_ADDR y, GM_ADDR workspace, const "4954 "inline __aicore__ void test_kernel_general_0_nil_0_nil(GM_ADDR x, GM_ADDR y, GM_ADDR workspace, const "
4954 "AutofuseTilingData *t) {\n"4955 "AutofuseTilingData *t) {\n"
4955 "int block_dim = GetBlockIdx();\n"4956 "int block_dim = GetBlockIdx();\n"
4956- "if (block_dim >= t->block_dim) { \n"4957+ "if (block_dim >= t->block_dim) {\n"
4957 " return;\n"4958 " return;\n"
4958 "}\n\n"4959 "}\n\n"
4959 "GlobalTensor<half> global_0;\n"4960 "GlobalTensor<half> global_0;\n"
@@ -4974,7 +4975,7 @@ TEST(CodegenKernel, Kernel_GenerateKernel_Single_ScheduleGroup) {
4974 "inline __aicore__ void test_kernel_general_1_nil_1_nil(GM_ADDR x, GM_ADDR y, GM_ADDR workspace, const "4975 "inline __aicore__ void test_kernel_general_1_nil_1_nil(GM_ADDR x, GM_ADDR y, GM_ADDR workspace, const "
4975 "AutofuseTilingData *t) {\n"4976 "AutofuseTilingData *t) {\n"
4976 "int block_dim = GetBlockIdx();\n"4977 "int block_dim = GetBlockIdx();\n"
4977- "if (block_dim >= t->block_dim) { \n"4978+ "if (block_dim >= t->block_dim) {\n"
4978 " return;\n"4979 " return;\n"
4979 "}\n\n"4980 "}\n\n"
4980 "GlobalTensor<half> global_0;\n"4981 "GlobalTensor<half> global_0;\n"
@@ -4995,7 +4996,7 @@ TEST(CodegenKernel, Kernel_GenerateKernel_Single_ScheduleGroup) {
4995 "inline __aicore__ void test_kernel_general_2_nil_2_nil(GM_ADDR x, GM_ADDR y, GM_ADDR workspace, const "4996 "inline __aicore__ void test_kernel_general_2_nil_2_nil(GM_ADDR x, GM_ADDR y, GM_ADDR workspace, const "
4996 "AutofuseTilingData *t) {\n"4997 "AutofuseTilingData *t) {\n"
4997 "int block_dim = GetBlockIdx();\n"4998 "int block_dim = GetBlockIdx();\n"
4998- "if (block_dim >= t->block_dim) { \n"4999+ "if (block_dim >= t->block_dim) {\n"
4999 " return;\n"5000 " return;\n"
5000 "}\n\n"5001 "}\n\n"
5001 "GlobalTensor<half> global_0;\n"5002 "GlobalTensor<half> global_0;\n"
@@ -217,3 +217,4 @@ add_subdirectory(trunc_to_int_bf16_to_int32_test)
217# add_subdirectory(remainder_bf16_test)217# add_subdirectory(remainder_bf16_test)
218add_subdirectory(rand_store_test)218add_subdirectory(rand_store_test)
219add_subdirectory(randn_store_test)219add_subdirectory(randn_store_test)
220+add_subdirectory(fused_backend_elewise_test)
@@ -0,0 +1,7 @@
1+backend_e2e_st_test(fused_backend_elewise_test
2+ CODEGEN fused_backend_elewise_generate.cpp
3+ KERNEL_SRC
4+ fused_backend_elewise_test_kernel.cpp
5+ fused_backend_elewise_test_tiling.cpp
6+ autofuse_tiling_data.h
7+ TEST_SRC test_e2e_fused_backend_elewise_expect_kernel.cpp)
@@ -0,0 +1,78 @@
1+/**
2+ * Copyright (c) 2026 Huawei Technologies Co., Ltd.
3+ * This program is free software, you can redistribute it and/or modify it under
4+ * the terms and conditions of CANN Open Software License Agreement Version 2.0
5+ * (the "License"). Please refer to the License for details. You may not use
6+ * this file except in compliance with the License. THIS SOFTWARE IS PROVIDED ON
7+ * AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
8+ * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS
9+ * FOR A PARTICULAR PURPOSE. See LICENSE in the root of the software repository
10+ * for the full text of the License.
11+ */
12+ 
13+#include <fstream>
14+#include <gtest/gtest.h>
15+#include <exception>
16+#include <filesystem>
17+#include "codegen.h"
18+#include "optimize.h"
19+#include "share_graph.h"
20+#include "backend_common.h"
21+ 
22+#include <iostream>
23+#include <vector>
24+#include <string>
25+#include "runtime_stub.h"
26+#include "common/platform_context.h"
27+ 
28+class TestFusedBackendElewiseE2e : public testing::Test {
29+ protected:
30+ void SetUp() override {
31+ dlog_setlevel(ASCGEN_MODULE_NAME, DLOG_ERROR, 0);
32+ ge::PlatformContext::GetInstance().Reset();
33+ auto stub_v2 = std::make_shared<af::RuntimeStubV2>();
34+ ge::RuntimeStub::SetInstance(stub_v2);
35+ }
36+ void TearDown() override {
37+ dlog_setlevel(ASCGEN_MODULE_NAME, DLOG_ERROR, 0);
38+ ge::RuntimeStub::Reset();
39+ }
40+};
41+ 
42+TEST_F(TestFusedBackendElewiseE2e, FusedBackendElewiseE2eCodegen) {
43+ bool gen_success = true;
44+ std::string tilig_stub = R"(
45+#define REGISTER_TILING_DEFAULT(tiling)
46+#define GET_TILING_DATA(t, tiling) AutofuseTilingData t = *(AutofuseTilingData*)tiling;
47+)";
48+ 
49+ // shape_info 和 FusedBackendElewiseGraph入参dims_size匹配(个数相同,命名规则为s开头、编号从0开始)
50+ std::map<std::string, std::string> shape_info({{"s0", "stub_s0"}, {"s1", "stub_s1"}});
51+ auto graph = ascir::ShareGraph::FusedBackendElewiseGraph(2);
52+ std::cout << "KERNEL_SRC_LIST=" << KERNEL_SRC_LIST << std::endl;
53+ std::vector<std::string> parts = splitString(KERNEL_SRC_LIST, ':');
54+ std::string kernel_src_file_name = parts[0]; // fused_backend_elewise_test_tiling.cpp
55+ std::string tiling_src_file_name = parts[1]; // fused_backend_elewise_test_kernel.cpp
56+ std::string tiling_data_src_file_name = parts[2]; // autofuse_tiling_data.h
57+ 
58+ try {
59+ optimize::Optimizer optimizer(optimize::OptimizerOptions{.graph_type = optimize::GraphType::kFusedAscBackend});
60+ codegen::Codegen codegen(codegen::CodegenOptions{});
61+ 
62+ std::fstream kernel_file(kernel_src_file_name, std::ios::out);
63+ std::fstream tiling_file(tiling_src_file_name, std::ios::out);
64+ std::fstream tiling_data_file(tiling_data_src_file_name, std::ios::out);
65+ 
66+ ascir::FusedScheduledResult fused_schedule_result;
67+ EXPECT_EQ(optimizer.Optimize(graph, fused_schedule_result), 0);
68+ codegen::CodegenResult result;
69+ EXPECT_EQ(codegen.Generate(shape_info, fused_schedule_result, result), 0);
70+ kernel_file << tilig_stub << RemoveSubDirInclude(result.kernel);
71+ tiling_file << result.tiling;
72+ tiling_data_file << result.tiling_data;
73+ } catch (...) {
74+ gen_success = false;
75+ }
76+ 
77+ EXPECT_EQ(gen_success, true);
78+}
@@ -0,0 +1,101 @@
1+/**
2+ * Copyright (c) Huawei Technologies Co., Ltd. 2026 All rights reserved.
3+ *
4+ * Licensed under the Apache License, Version 2.0 (the "License");
5+ * you may not use this file except in compliance with the License.
6+ * You may obtain a copy of the License at
7+ *
8+ * http://www.apache.org/licenses/LICENSE-2.0
9+ *
10+ * Unless required by applicable law or agreed to in writing, software
11+ * distributed under the License is distributed on an "AS IS" BASIS,
12+ * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
13+ * See the License for the specific language governing permissions and
14+ * limitations under the License.
15+ */
16+#include <gtest/gtest.h>
17+#include <cmath>
18+#include "tikicpulib.h"
19+#include "autofuse_tiling_data.h"
20+ 
21+extern "C" __global__ __aicore__ void fused_backend_elewise_test(GM_ADDR x0, GM_ADDR x1, GM_ADDR x2, GM_ADDR y0,
22+ GM_ADDR y1, GM_ADDR workspace, GM_ADDR gm_tiling_data);
23+extern "C" int64_t AutofuseTiling(uint32_t s0, uint32_t s1, AutofuseTilingData *tiling, uint32_t *workspaceSize,
24+ uint64_t *blockDim, uint32_t aiv_num, uint32_t ub_size);
25+ 
26+namespace {
27+class E2E_FusedBackendElewise_Code : public testing::Test, public testing::WithParamInterface<std::vector<int>> {};
28+ 
29+TEST_P(E2E_FusedBackendElewise_Code, CalculateCorrect) {
30+ auto test_shape = GetParam();
31+ 
32+ uint64_t block_dim = 48;
33+ 
34+ int test_size = test_shape[0] * test_shape[1];
35+ 
36+ AutofuseTilingData tiling_data;
37+ float *x0 = (float *)AscendC::GmAlloc(test_size * sizeof(float) + 32);
38+ float *x1 = (float *)AscendC::GmAlloc(test_size * sizeof(float) + 32);
39+ float *x2 = (float *)AscendC::GmAlloc(test_size * sizeof(float) + 32);
40+ 
41+ float *y0 = (float *)AscendC::GmAlloc(test_size * sizeof(float) + 32);
42+ float *y1 = (float *)AscendC::GmAlloc(test_size * sizeof(float) + 32);
43+ 
44+ float *expect0 = (float *)AscendC::GmAlloc(test_size * sizeof(float) + 32);
45+ float *expect1 = (float *)AscendC::GmAlloc(test_size * sizeof(float) + 32);
46+ 
47+ // Prepare test and expect data
48+ for (int i = 0; i < test_size; i++) {
49+ x0[i] = static_cast<float>(i);
50+ x1[i] = static_cast<float>(i);
51+ x2[i] = static_cast<float>(i);
52+ expect1[i] = x0[i] + x1[i];
53+ expect0[i] = (x0[i] + x1[i]) * 2 + x2[i];
54+ }
55+ 
56+ // Launch
57+ uint32_t ws_size = 0;
58+ AutofuseTiling(test_shape[0], test_shape[1], &tiling_data, &ws_size, &block_dim, 48, 192 * 1024);
59+ printf("g0_tiling_key: %d\n", tiling_data.graph0_tiling_key);
60+ printf("g1_tiling_key: %d\n", tiling_data.graph1_tiling_key);
61+ printf("g2_tiling_key: %d\n", tiling_data.graph2_tiling_key);
62+ printf("ws_size: %d, ws3_size: %d, ws5_size: %d\n", ws_size, tiling_data.workspace3, tiling_data.workspace5);
63+ float *workspace = (float *)AscendC::GmAlloc(ws_size + 32);
64+ 
65+ AscendC::SetKernelMode(KernelMode::AIV_MODE);
66+ ICPU_RUN_KF(fused_backend_elewise_test, tiling_data.block_dim, (uint8_t *)x0, (uint8_t *)x1, (uint8_t *)x2,
67+ (uint8_t *)y0, (uint8_t *)y1, (uint8_t *)workspace, (uint8_t *)&tiling_data);
68+ 
69+ // Count difference
70+ uint32_t diff_count = 0;
71+ for (int i = 0; i < test_size; i++) {
72+ auto diff0 = (double)(y0[i] - expect0[i]);
73+ if (diff0 < -1e-5 || diff0 > 1e-5) {
74+ printf("i: %d, y0: %f, expect0: %f\n", i, y0[i], expect0[i]);
75+ diff_count++;
76+ }
77+ }
78+ for (int i = 0; i < test_size; i++) {
79+ auto diff1 = (double)(y1[i] - expect1[i]);
80+ if (diff1 < -1e-5 || diff1 > 1e-5) {
81+ printf("i: %d, y1: %f, expect1: %f\n", i, y1[i], expect1[i]);
82+ diff_count++;
83+ }
84+ }
85+ 
86+ EXPECT_EQ(diff_count, 0) << " of " << test_size;
87+ 
88+ AscendC::GmFree(x0);
89+ AscendC::GmFree(x1);
90+ AscendC::GmFree(x2);
91+ AscendC::GmFree(y0);
92+ AscendC::GmFree(y1);
93+ AscendC::GmFree(expect0);
94+ AscendC::GmFree(expect1);
95+ AscendC::GmFree(workspace);
96+}
97+ 
98+INSTANTIATE_TEST_SUITE_P(CalcWithDifferentShape, E2E_FusedBackendElewise_Code,
99+ ::testing::Values(std::vector<int>{2, 64}));
100+ 
101+} // namespace