已合并
【质量加固】显式声明部分 lambda 捕获列表 #1756
小白学习者创建于 8月17日
【质量加固】显式声明部分 lambda 捕获列表 #1756
已合并
共 5 个文件变更+32-27
| @@ -249,9 +249,10 @@ bool ArgListReorder::IsReduceInputTensor(const TensorPtr &tensor, const SubAxis | |||
| 249 | if ((tensor == nullptr) || (tensor->data_type_size == 0U)) { | 249 | if ((tensor == nullptr) || (tensor->data_type_size == 0U)) { |
| 250 | return false; | 250 | return false; |
| 251 | } | 251 | } |
| 252 | - return std::any_of(tensor->dim_info.begin(), tensor->dim_info.end(), [&](const SubAxis *dim) { | 252 | + return std::any_of(tensor->dim_info.begin(), tensor->dim_info.end(), |
| 253 | - return IsSameOrRelatedAxis(dim, reduce_axis) || IsReduceOrigAxis(dim, reduce_axis_ori_axes_set); | 253 | + [this, reduce_axis, &reduce_axis_ori_axes_set](const SubAxis *dim) { |
| 254 | - }); | 254 | + return IsSameOrRelatedAxis(dim, reduce_axis) || IsReduceOrigAxis(dim, reduce_axis_ori_axes_set); |
| 255 | + }); | ||
| 255 | } | 256 | } |
| 256 | 257 | ||
| 257 | bool ArgListReorder::CollectReduceTailAxis(const TensorPtr &tensor, const SubAxis *reduce_axis, | 258 | bool ArgListReorder::CollectReduceTailAxis(const TensorPtr &tensor, const SubAxis *reduce_axis, |
| @@ -282,9 +283,10 @@ bool ArgListReorder::IsReduceTailTileAxis(const SubAxis *tail_axis) const { | |||
| 282 | if (tail_axis == nullptr) { | 283 | if (tail_axis == nullptr) { |
| 283 | return false; | 284 | return false; |
| 284 | } | 285 | } |
| 285 | - return std::any_of(tuning_space_->sub_axes.begin(), tuning_space_->sub_axes.end(), [&](const SubAxisPtr &axis) { | 286 | + return std::any_of( |
| 286 | - return (axis->name == tail_axis->name) && (axis->axis_type == AxisPosition::INNER) && !axis->is_bind_multi_core; | 287 | + tuning_space_->sub_axes.begin(), tuning_space_->sub_axes.end(), [&tail_axis](const SubAxisPtr &axis) { |
| 287 | - }); | 288 | + return (axis->name == tail_axis->name) && (axis->axis_type == AxisPosition::INNER) && !axis->is_bind_multi_core; |
| 289 | + }); | ||
| 288 | } | 290 | } |
| 289 | 291 | ||
| 290 | bool ArgListReorder::TryGetReduceTailTileInfo(const NodeInfo &node, | 292 | bool ArgListReorder::TryGetReduceTailTileInfo(const NodeInfo &node, |
| @@ -684,7 +686,7 @@ void ArgListReorder::MakeSureLoadStoreInnerestSameOrder(const std::vector<AttAxi | |||
| 684 | } | 686 | } |
| 685 | 687 | ||
| 686 | const bool overlaps_load_store = std::any_of( | 688 | const bool overlaps_load_store = std::any_of( |
| 687 | - equal_order_reduce_tail_axes_.begin(), equal_order_reduce_tail_axes_.end(), [&](const std::string &name) { | 689 | + equal_order_reduce_tail_axes_.begin(), equal_order_reduce_tail_axes_.end(), [this](const std::string &name) { |
| 688 | return load_store_inner_most_dims_.find(name) != load_store_inner_most_dims_.end(); | 690 | return load_store_inner_most_dims_.find(name) != load_store_inner_most_dims_.end(); |
| 689 | }); | 691 | }); |
| 690 | if (overlaps_load_store) { | 692 | if (overlaps_load_store) { |
| @@ -3847,11 +3847,13 @@ void TilingCodeGenImpl::GenPGOByCoreNumFunctionHead(size_t impl_graph_id) { | |||
| 3847 | 3847 | ||
| 3848 | bool TilingCodeGenImpl::TryGenPGOByCoreNumReuseTiling(size_t asc_graph_id, size_t impl_graph_id, size_t group_id, | 3848 | bool TilingCodeGenImpl::TryGenPGOByCoreNumReuseTiling(size_t asc_graph_id, size_t impl_graph_id, size_t group_id, |
| 3849 | uint32_t group_index) { | 3849 | uint32_t group_index) { |
| 3850 | - const auto iter = std::find_if(tiling_model_info_.cbegin(), tiling_model_info_.cend(), [&](const auto &model_info) { | 3850 | + const auto iter = std::find_if(tiling_model_info_.cbegin(), tiling_model_info_.cend(), |
| 3851 | - const auto &ident = model_info.schedule_group_ident; | 3851 | + [asc_graph_id, impl_graph_id, group_id](const auto &model_info) { |
| 3852 | - return ident.asc_graph_id == asc_graph_id && ident.impl_graph_id == impl_graph_id && ident.group_id == group_id && | 3852 | + const auto &ident = model_info.schedule_group_ident; |
| 3853 | - model_info.reuse_schedule_group != nullptr && model_info.reuse_schedule_group->IsReuseGroup(ident); | 3853 | + return ident.asc_graph_id == asc_graph_id && ident.impl_graph_id == impl_graph_id && |
| 3854 | - }); | 3854 | + ident.group_id == group_id && model_info.reuse_schedule_group != nullptr && |
| 3855 | + model_info.reuse_schedule_group->IsReuseGroup(ident); | ||
| 3856 | + }); | ||
| 3855 | if (iter == tiling_model_info_.cend()) { | 3857 | if (iter == tiling_model_info_.cend()) { |
| 3856 | return false; | 3858 | return false; |
| 3857 | } | 3859 | } |
| @@ -5175,7 +5177,7 @@ std::pair<std::string, bool> TilingCodeGenImpl::GenConflictExprContextCode( | |||
| 5175 | auto input_vars = GetVarsNames(args_manager.GetInputVars()); | 5177 | auto input_vars = GetVarsNames(args_manager.GetInputVars()); |
| 5176 | input_var_names.insert(input_vars.begin(), input_vars.end()); | 5178 | input_var_names.insert(input_vars.begin(), input_vars.end()); |
| 5177 | } | 5179 | } |
| 5178 | - auto emit_decl = [&](const std::string &name, const std::string &src) { | 5180 | + auto emit_decl = [&code, &declared_symbols](const std::string &name, const std::string &src) { |
| 5179 | code += " auto " + name + " = " + src + ".get_" + name + "();\n"; | 5181 | code += " auto " + name + " = " + src + ".get_" + name + "();\n"; |
| 5180 | declared_symbols.insert(name); | 5182 | declared_symbols.insert(name); |
| 5181 | }; | 5183 | }; |
| @@ -396,7 +396,7 @@ Status Codegen::GenerateForInductor(const ascir::FusedScheduledResult &fused_sch | |||
| 396 | fallback.tiling_lib_.DisableInductorPgo(); | 396 | fallback.tiling_lib_.DisableInductorPgo(); |
| 397 | return fallback.GenerateForInductor(fused_schedule_result, result); | 397 | return fallback.GenerateForInductor(fused_schedule_result, result); |
| 398 | } | 398 | } |
| 399 | - const auto generate_tiling_without_pgo = [&]() { | 399 | + const auto generate_tiling_without_pgo = [this, &fused_schedule_result, &result]() { |
| 400 | if (!tiling_lib_.IsInductorPgoEnabled() || ascgen_utils::IsCubeFusedScheduled(fused_schedule_result) || | 400 | if (!tiling_lib_.IsInductorPgoEnabled() || ascgen_utils::IsCubeFusedScheduled(fused_schedule_result) || |
| 401 | !ascgen_utils::IsStaticSchedResult(fused_schedule_result)) { | 401 | !ascgen_utils::IsStaticSchedResult(fused_schedule_result)) { |
| 402 | return af::FAILED; | 402 | return af::FAILED; |
| @@ -304,16 +304,17 @@ bool FilterScheduledResult(FilterState &state, ascir::ScheduledResult &scheduled | |||
| 304 | const size_t before_group_size = schedule_group.impl_graphs.size(); | 304 | const size_t before_group_size = schedule_group.impl_graphs.size(); |
| 305 | size_t impl_idx = 0UL; | 305 | size_t impl_idx = 0UL; |
| 306 | schedule_group.impl_graphs.erase( | 306 | schedule_group.impl_graphs.erase( |
| 307 | - std::remove_if(schedule_group.impl_graphs.begin(), schedule_group.impl_graphs.end(), | 307 | + std::remove_if( |
| 308 | - [&](const af::AscGraph &impl_graph) { | 308 | + schedule_group.impl_graphs.begin(), schedule_group.impl_graphs.end(), |
| 309 | - const size_t current_impl_idx = impl_idx++; | 309 | + [&schedule_group, &impl_idx, &state, node_idx, result_idx, group_idx](const af::AscGraph &impl_graph) { |
| 310 | - const TemplatePosition position = {node_idx, result_idx, group_idx, current_impl_idx}; | 310 | + const size_t current_impl_idx = impl_idx++; |
| 311 | - const bool drop = ShouldDropImplGraph(impl_graph, state.ub_size, position); | 311 | + const TemplatePosition position = {node_idx, result_idx, group_idx, current_impl_idx}; |
| 312 | - if (drop) { | 312 | + const bool drop = ShouldDropImplGraph(impl_graph, state.ub_size, position); |
| 313 | - schedule_group.graph_name_to_score_funcs.erase(impl_graph.GetName()); | 313 | + if (drop) { |
| 314 | - } | 314 | + schedule_group.graph_name_to_score_funcs.erase(impl_graph.GetName()); |
| 315 | - return drop; | 315 | + } |
| 316 | - }), | 316 | + return drop; |
| 317 | + }), | ||
| 317 | schedule_group.impl_graphs.end()); | 318 | schedule_group.impl_graphs.end()); |
| 318 | if (before_group_size > 0UL && schedule_group.impl_graphs.empty()) { | 319 | if (before_group_size > 0UL && schedule_group.impl_graphs.empty()) { |
| 319 | GELOGD( | 320 | GELOGD( |
| @@ -333,7 +334,7 @@ af::Status FilterNodeScheduledResults(ascir::FusedScheduledResult &fused_schedul | |||
| 333 | FilterState state = {fused_scheduled_result, ub_size}; | 334 | FilterState state = {fused_scheduled_result, ub_size}; |
| 334 | size_t result_idx = 0UL; | 335 | size_t result_idx = 0UL; |
| 335 | scheduled_results.erase(std::remove_if(scheduled_results.begin(), scheduled_results.end(), | 336 | scheduled_results.erase(std::remove_if(scheduled_results.begin(), scheduled_results.end(), |
| 336 | - [&](ascir::ScheduledResult &scheduled_result) { | 337 | + [&state, node_idx, &result_idx](ascir::ScheduledResult &scheduled_result) { |
| 337 | const size_t current_result_idx = result_idx++; | 338 | const size_t current_result_idx = result_idx++; |
| 338 | return !FilterScheduledResult(state, scheduled_result, node_idx, | 339 | return !FilterScheduledResult(state, scheduled_result, node_idx, |
| 339 | current_result_idx); | 340 | current_result_idx); |
| @@ -290,10 +290,10 @@ bool RecomputeCaseGenerator::IsRecomputableNode(ascir::HintGraph &hint_graph, af | |||
| 290 | auto output_tensors = node->outputs(); | 290 | auto output_tensors = node->outputs(); |
| 291 | if (is_static_graph_) { | 291 | if (is_static_graph_) { |
| 292 | return std::all_of(output_tensors.begin(), output_tensors.end(), | 292 | return std::all_of(output_tensors.begin(), output_tensors.end(), |
| 293 | - [&](af::AscTensor *&tensor) { return check_static_tensor(tensor->attr); }); | 293 | + [&check_static_tensor](af::AscTensor *&tensor) { return check_static_tensor(tensor->attr); }); |
| 294 | } | 294 | } |
| 295 | return std::all_of(output_tensors.begin(), output_tensors.end(), | 295 | return std::all_of(output_tensors.begin(), output_tensors.end(), |
| 296 | - [&](af::AscTensor *&tensor) { return check_dynamic_tensor(tensor->attr); }); | 296 | + [&check_dynamic_tensor](af::AscTensor *&tensor) { return check_dynamic_tensor(tensor->attr); }); |
| 297 | } | 297 | } |
| 298 | 298 | ||
| 299 | bool RecomputeCaseGenerator::IsRecomputableAlwaysBetter() const { | 299 | bool RecomputeCaseGenerator::IsRecomputableAlwaysBetter() const { |