天天向上队——pipeline文件夹注释 #26

Open
WEI_4614 wants to merge 26 commits from WEI_4614/mindspore2022:comp into master
21 changed files with 1661 additions and 354 deletions

View File

@ -94,6 +94,9 @@ void UpdateFuncGraphParameter(const FuncGraphPtr &func_graph) {
} }
func_graph->set_parameters(new_paras); func_graph->set_parameters(new_paras);
} }
//The function of this code is to update the parameters in the passed function graph
// leaving only the parameters that have no default value and meet certain conditions,
//and update the new parameter list to the function graph
bool IsDynamicShapeGraph(const FuncGraphPtr &func_graph) { bool IsDynamicShapeGraph(const FuncGraphPtr &func_graph) {
MS_EXCEPTION_IF_NULL(func_graph); MS_EXCEPTION_IF_NULL(func_graph);
@ -141,6 +144,7 @@ void TaskEmitActionForMindRT(const ResourcePtr &res) {
res->SetResult(kOutput, actor_info); res->SetResult(kOutput, actor_info);
} }
// Get the graph information, construct the pointer of the execution function, execute the graph and return the result
void ExecuteActionForMindRT(const ResourcePtr &res) { void ExecuteActionForMindRT(const ResourcePtr &res) {
MS_EXCEPTION_IF_NULL(res); MS_EXCEPTION_IF_NULL(res);
const auto actor_info = res->GetResult(kOutput).cast<compile::ActorInfo>(); const auto actor_info = res->GetResult(kOutput).cast<compile::ActorInfo>();
@ -206,6 +210,7 @@ void ModifyOutputNode(const FuncGraphPtr &func_graph) {
func_graph->set_output(merge_node); func_graph->set_output(merge_node);
// Clear // Clear
func_graph->set_modify_output(true); func_graph->set_modify_output(true);
func_graph->ClearUsedForwardNodes(); func_graph->ClearUsedForwardNodes();
} }
@ -228,6 +233,7 @@ abstract::AnalysisResult AbstractAnalyze(const ResourcePtr &resource, const Func
MS_EXCEPTION_IF_NULL(node); MS_EXCEPTION_IF_NULL(node);
// Handle previous inferred value for CNode if is loaded from MindIR // Handle previous inferred value for CNode if is loaded from MindIR
if (resource->is_load()) { if (resource->is_load()) {
// If the primitive is not defined in front end, keep the inferred value loaded from MindIR. // If the primitive is not defined in front end, keep the inferred value loaded from MindIR.
auto primitive = GetCNodePrimitive(node); auto primitive = GetCNodePrimitive(node);
@ -241,9 +247,9 @@ abstract::AnalysisResult AbstractAnalyze(const ResourcePtr &resource, const Func
} }
const AbstractBasePtr &prev_inferred = node->abstract(); const AbstractBasePtr &prev_inferred = node->abstract();
// Keep previous inferred value for ValueNode if the inferred value is not AbstractFunction. // Keep previous inferred value for ValueNode if the inferred value is not AbstractFunction
if (!node->isa<ValueNode>() || (prev_inferred != nullptr && prev_inferred->isa<abstract::AbstractFunction>())) { if (!node->isa<ValueNode>() || (prev_inferred != nullptr && prev_inferred->isa<abstract::AbstractFunction>())) {
// Reset tuple/list abstract use flags. // Reset tuple/list abstract use flags
if (enable_eliminate_unused_element && prev_inferred != nullptr && if (enable_eliminate_unused_element && prev_inferred != nullptr &&
prev_inferred->isa<abstract::AbstractSequence>()) { prev_inferred->isa<abstract::AbstractSequence>()) {
SetSequenceNodeElementsUseFlags(node, nullptr); SetSequenceNodeElementsUseFlags(node, nullptr);
@ -322,6 +328,7 @@ const FuncGraphPtr GetLoadedGraph(const ResourcePtr &res) {
MS_LOG(EXCEPTION) << "The loaded sub graph currently should be less than 2, but got " << loaded_graph_num; MS_LOG(EXCEPTION) << "The loaded sub graph currently should be less than 2, but got " << loaded_graph_num;
} }
// Check that the root diagram input shape and type are consistent with the loaded diagram
void CheckRootInputShapeAndType(const ResourcePtr &res, const FuncGraphPtr &loaded_graph) { void CheckRootInputShapeAndType(const ResourcePtr &res, const FuncGraphPtr &loaded_graph) {
MS_EXCEPTION_IF_NULL(res); MS_EXCEPTION_IF_NULL(res);
auto manager = res->manager(); auto manager = res->manager();
@ -374,6 +381,9 @@ void CheckRootInputShapeAndType(const ResourcePtr &res, const FuncGraphPtr &load
} }
} }
// Parsing a Python object into a graph includes the process of obtaining a source input object, initializing the parser environment,
// setting up Python paths, converting an input object into a graph, creating a top-level graph, updating the parser and manager, and returning true values
bool ParseAction(const ResourcePtr &res) { bool ParseAction(const ResourcePtr &res) {
MS_EXCEPTION_IF_NULL(res); MS_EXCEPTION_IF_NULL(res);
TraceManager::OpenRecordDebugInfoFlag(); TraceManager::OpenRecordDebugInfoFlag();
@ -419,6 +429,7 @@ bool ParseAction(const ResourcePtr &res) {
// This step do this optimize: graph1(x){xx(fv1),xxx(fv2)}, graph2(x){xxx(fv3),xxx(fv4)}-> // This step do this optimize: graph1(x){xx(fv1),xxx(fv2)}, graph2(x){xxx(fv3),xxx(fv4)}->
// graph1(x){base_graph(x, fv1, fv2)}, graph1(x){base_graph(x, fv3, fv4)}, base_graph(x, fv...){xxx,xxx} // graph1(x){base_graph(x, fv1, fv2)}, graph1(x){base_graph(x, fv3, fv4)}, base_graph(x, fv...){xxx,xxx}
// all obj_map's graph shared base_graph // all obj_map's graph shared base_graph
bool CombineLikeGraphs(const ResourcePtr &res) { bool CombineLikeGraphs(const ResourcePtr &res) {
MS_EXCEPTION_IF_NULL(res); MS_EXCEPTION_IF_NULL(res);
auto &obj_map = parse::data_converter::GetObjGraphs(); auto &obj_map = parse::data_converter::GetObjGraphs();
@ -555,6 +566,10 @@ bool EliminateUnusedParameterAction(const ResourcePtr &res) {
return true; return true;
} }
// Perform graph abstraction and specialization operations, including obtaining graph objects, parameter specification lists, parallel context objects,
// initializing shape information, obtaining originally loaded graph, processing default parameters, performing abstract analysis,
// updating top-level graph, specializing graph, removing unused nodes, checking input shapes and types, updating graph parameters, and so on.
bool AbstractSpecializeAction(const ResourcePtr &res) { bool AbstractSpecializeAction(const ResourcePtr &res) {
MS_EXCEPTION_IF_NULL(res); MS_EXCEPTION_IF_NULL(res);
if (res->func_graph() == nullptr) { if (res->func_graph() == nullptr) {
@ -587,14 +602,11 @@ bool AbstractSpecializeAction(const ResourcePtr &res) {
} }
// Analyze // Analyze
AnalysisResult result = AbstractAnalyze(res, func_graph, args_spec); AnalysisResult result = AbstractAnalyze(res, func_graph, args_spec);
// The top graph may be replaced by infer, update the top graph when the infer is done // The top graph may be replaced by infer, update the top graph when the infer is done
parse::Parser::UpdateTopFuncGraph(result.context->func_graph()); parse::Parser::UpdateTopFuncGraph(result.context->func_graph());
// Specialize // Specialize
FuncGraphPtr new_fg = ProgramSpecialize(res, result.context->func_graph(), result.context); FuncGraphPtr new_fg = ProgramSpecialize(res, result.context->func_graph(), result.context);
res->set_func_graph(new_fg); res->set_func_graph(new_fg);
// Remove unused nodes in cnode order list, this is prepared for auto-monad. // Remove unused nodes in cnode order list, this is prepared for auto-monad.
if (new_fg) { if (new_fg) {
new_fg->EraseUnusedNodeInOrder(); new_fg->EraseUnusedNodeInOrder();
@ -710,6 +722,8 @@ bool CheckGraphOutputConstOrParameter(const FuncGraphPtr &func_graph) {
return false; return false;
} }
// Eliminate forward CNode nodes in Pynative mode, including obtaining graph actuator and Pynative actuator instance, checking execution mode, obtaining process phase,
// processing derived graph and forward process, running gradient calculation and replacing forward node, setting forward eliminating flag, setting gradient graph, modifying output node, etc.
bool EliminateForwardCNode(const ResourcePtr &res) { bool EliminateForwardCNode(const ResourcePtr &res) {
// This function only works in Pynative mode. The func_graph is decorated by ms_function. // This function only works in Pynative mode. The func_graph is decorated by ms_function.
if (MsContext::GetInstance()->get_param<int>(MS_CTX_EXECUTION_MODE) == kGraphMode) { if (MsContext::GetInstance()->get_param<int>(MS_CTX_EXECUTION_MODE) == kGraphMode) {
@ -721,6 +735,7 @@ bool EliminateForwardCNode(const ResourcePtr &res) {
auto phase = graph_executor->phase(); auto phase = graph_executor->phase();
MS_LOG(DEBUG) << "The phase of current pipeline graph is: " << phase; MS_LOG(DEBUG) << "The phase of current pipeline graph is: " << phase;
// Exporting graph in PyNative mode or only running forward process no need to do this action. // Exporting graph in PyNative mode or only running forward process no need to do this action.
auto pynative_exec = pynative::PynativeExecutor::GetInstance(); auto pynative_exec = pynative::PynativeExecutor::GetInstance();
if (phase.find("export") == 0 || !pynative_exec->grad_flag()) { if (phase.find("export") == 0 || !pynative_exec->grad_flag()) {
MS_LOG(DEBUG) << "When exporting graph or only running forward process, no need to eliminate forward cnode."; MS_LOG(DEBUG) << "When exporting graph or only running forward process, no need to eliminate forward cnode.";
@ -762,6 +777,7 @@ bool EliminateAdRelatedSpecialOpNode(const ResourcePtr &res) {
return EliminateAdRelatedSpecialOpOptPass(res); return EliminateAdRelatedSpecialOpOptPass(res);
} }
// The process of determining whether indirect calls exist includes traversing all nodes, determining Partial, Switch, SwitchLayer, and Call nodes, printing log information, and returning true or false values.
bool HasIncorporateCall(const std::vector<AnfNodePtr> &all_nodes) { bool HasIncorporateCall(const std::vector<AnfNodePtr> &all_nodes) {
for (const auto &node : all_nodes) { for (const auto &node : all_nodes) {
if (!node->isa<CNode>()) { if (!node->isa<CNode>()) {
@ -872,6 +888,7 @@ void SetRunMode(const FuncGraphPtr &func_graph, compile::Backend *backend_ptr) {
const auto &all_nodes = TopoSort(func_graph->return_node(), SuccDeeperSimple, AlwaysInclude); const auto &all_nodes = TopoSort(func_graph->return_node(), SuccDeeperSimple, AlwaysInclude);
// GPU/CPU no need set any context. // GPU/CPU no need set any context.
if (!ExistTarget(all_nodes, kAscendDevice)) { if (!ExistTarget(all_nodes, kAscendDevice)) {
return; return;
} }
@ -924,12 +941,15 @@ void SetRunMode(const FuncGraphPtr &func_graph, compile::Backend *backend_ptr) {
return; return;
} }
// GRAPH | normal network and if/for/switch scenario etc : MultiGraph path in MindRT. // GRAPH | normal network and if/for/switch scenario etc : MultiGraph path in Mind
MS_LOG(INFO) << "Run graph mode with multigraph sink."; MS_LOG(INFO) << "Run graph mode with multigraph sink.";
set_ctx(true, true, true); set_ctx(true, true, true);
return; return;
} }
// Set the running mode according to the function graph properties, execution mode, device target, and back-end policy,
//and set the corresponding flag bit and print log information according to the conditions
void OriginSetRunMode(const ResourcePtr &res) { void OriginSetRunMode(const ResourcePtr &res) {
FuncGraphPtr func_graph = res->func_graph(); FuncGraphPtr func_graph = res->func_graph();
MS_EXCEPTION_IF_NULL(func_graph); MS_EXCEPTION_IF_NULL(func_graph);
@ -964,6 +984,7 @@ void OriginSetRunMode(const ResourcePtr &res) {
} }
} }
// Perform task launch operations and set the run mode and corresponding graph compilation operations by calling different functions.
bool TaskEmitAction(const ResourcePtr &res) { bool TaskEmitAction(const ResourcePtr &res) {
MS_EXCEPTION_IF_NULL(res); MS_EXCEPTION_IF_NULL(res);
FuncGraphPtr func_graph = res->func_graph(); FuncGraphPtr func_graph = res->func_graph();
@ -989,6 +1010,7 @@ bool TaskEmitAction(const ResourcePtr &res) {
!is_parallel; !is_parallel;
if (context_ptr->get_param<bool>(MS_CTX_ENABLE_MINDRT) && common::GetEnv("DISABLE_ASCEND_MINDRT") != "1") { if (context_ptr->get_param<bool>(MS_CTX_ENABLE_MINDRT) && common::GetEnv("DISABLE_ASCEND_MINDRT") != "1") {
// Run in GRAPH_MODE if the func_graph is ms_function or the func_graph contain multi-subgraph. // Run in GRAPH_MODE if the func_graph is ms_function or the func_graph contain multi-subgraph.
if (pynative_switch_to_graph_mode) { if (pynative_switch_to_graph_mode) {
context_ptr->set_param<int>(MS_CTX_EXECUTION_MODE, kGraphMode); context_ptr->set_param<int>(MS_CTX_EXECUTION_MODE, kGraphMode);
MS_LOG(INFO) << "PyNative graph Compile and Run in GRAPH_MODE"; MS_LOG(INFO) << "PyNative graph Compile and Run in GRAPH_MODE";
@ -1002,6 +1024,7 @@ bool TaskEmitAction(const ResourcePtr &res) {
MS_EXCEPTION_IF_NULL(bc_ptr); MS_EXCEPTION_IF_NULL(bc_ptr);
std::string backend = context_ptr->backend_policy(); std::string backend = context_ptr->backend_policy();
// The graph compiling of mindRT. // The graph compiling of mindRT.
if ((backend == kMsConvert) && context_ptr->get_param<bool>(MS_CTX_ENABLE_MINDRT)) { if ((backend == kMsConvert) && context_ptr->get_param<bool>(MS_CTX_ENABLE_MINDRT)) {
TaskEmitActionForMindRT(res); TaskEmitActionForMindRT(res);
if (pynative_switch_to_graph_mode) { if (pynative_switch_to_graph_mode) {
@ -1010,7 +1033,7 @@ bool TaskEmitAction(const ResourcePtr &res) {
return true; return true;
} }
// The graph compiling of control sink. // The graph compiling of control sink
if (IsCtrlSink() && backend == kMsConvert) { if (IsCtrlSink() && backend == kMsConvert) {
auto graph_id = bc_ptr->CompileGraph(NOT_NULL(func_graph)); auto graph_id = bc_ptr->CompileGraph(NOT_NULL(func_graph));
res->SetResult(kOutput, graph_id); res->SetResult(kOutput, graph_id);
@ -1037,12 +1060,14 @@ bool ExecuteAction(const ResourcePtr &res) {
} }
std::string backend = MsContext::GetInstance()->backend_policy(); std::string backend = MsContext::GetInstance()->backend_policy();
// The graph running of mindRT. // The graph running of mindRT.
if ((backend == kMsConvert) && MsContext::GetInstance()->get_param<bool>(MS_CTX_ENABLE_MINDRT)) { if ((backend == kMsConvert) && MsContext::GetInstance()->get_param<bool>(MS_CTX_ENABLE_MINDRT)) {
ExecuteActionForMindRT(res); ExecuteActionForMindRT(res);
return true; return true;
} }
// The graph running of control sink. // The graph running of control sink.
if (IsCtrlSink() && backend == kMsConvert) { if (IsCtrlSink() && backend == kMsConvert) {
auto graph_id = res->GetResult(kOutput).cast<GraphId>(); auto graph_id = res->GetResult(kOutput).cast<GraphId>();
std::shared_ptr<compile::Backend> bc_ptr = res->GetResult(kBackend).cast<std::shared_ptr<compile::Backend>>(); std::shared_ptr<compile::Backend> bc_ptr = res->GetResult(kBackend).cast<std::shared_ptr<compile::Backend>>();
@ -1097,6 +1122,8 @@ bool StartPSServerAction(const ResourcePtr &res) {
return true; return true;
} }
// Initialize the server according to the configuration parameters and run the server.
bool StartServerAction(const ResourcePtr &res) { bool StartServerAction(const ResourcePtr &res) {
MS_EXCEPTION_IF_NULL(res); MS_EXCEPTION_IF_NULL(res);
FuncGraphPtr func_graph = res->func_graph(); FuncGraphPtr func_graph = res->func_graph();
@ -1107,6 +1134,7 @@ bool StartServerAction(const ResourcePtr &res) {
// Update model threshold is a certain ratio of start_fl_job threshold. // Update model threshold is a certain ratio of start_fl_job threshold.
// update_model_threshold = start_fl_job_threshold * update_model_ratio. // update_model_threshold = start_fl_job_threshold * update_model_ratio.
size_t start_fl_job_threshold = ps::PSContext::instance()->start_fl_job_threshold(); size_t start_fl_job_threshold = ps::PSContext::instance()->start_fl_job_threshold();
float update_model_ratio = ps::PSContext::instance()->update_model_ratio(); float update_model_ratio = ps::PSContext::instance()->update_model_ratio();
size_t update_model_threshold = static_cast<size_t>(std::ceil(start_fl_job_threshold * update_model_ratio)); size_t update_model_threshold = static_cast<size_t>(std::ceil(start_fl_job_threshold * update_model_ratio));
@ -1216,6 +1244,7 @@ bool DistributedSplitAction(const ResourcePtr &res) {
// that will result in a synchronization error due to different executing order. // that will result in a synchronization error due to different executing order.
// Here we temporarily avoid the problem by skipping valuenode merging used by parallel related primitive, // Here we temporarily avoid the problem by skipping valuenode merging used by parallel related primitive,
// the final solution will be proposed later as a parallel feature. // the final solution will be proposed later as a parallel feature.
bool KeepValueNodeDuplication(const AnfNodePtr &value_node, const ResourcePtr &res) { bool KeepValueNodeDuplication(const AnfNodePtr &value_node, const ResourcePtr &res) {
MS_EXCEPTION_IF_NULL(res); MS_EXCEPTION_IF_NULL(res);
MS_EXCEPTION_IF_NULL(res->manager()); MS_EXCEPTION_IF_NULL(res->manager());
@ -1248,6 +1277,7 @@ bool RemoveValueNodeDuplicationsAction(const ResourcePtr &res) {
} }
auto manager = res->manager(); auto manager = res->manager();
// Remove duplicated value nodes, due to replace operation, can't use reference. // Remove duplicated value nodes, due to replace operation, can't use reference.
auto value_nodes = func_graph->value_nodes(); auto value_nodes = func_graph->value_nodes();
HashCache hash_cache; HashCache hash_cache;
HashValue hashes; HashValue hashes;
@ -1266,6 +1296,8 @@ bool ValidateAction(const ResourcePtr &res) { return ValidatePass(res); }
bool GeSpecializedAction(const ResourcePtr &res) { return GeSpecializedPass(res); } bool GeSpecializedAction(const ResourcePtr &res) { return GeSpecializedPass(res); }
// Based on the MindIR model information in the resource pointer, convert it to FuncGraphPtr and set it in the resource pointer.
bool SetMindIRGraphAction(const ResourcePtr &res) { bool SetMindIRGraphAction(const ResourcePtr &res) {
MS_EXCEPTION_IF_NULL(res); MS_EXCEPTION_IF_NULL(res);
res->set_is_load(true); res->set_is_load(true);
@ -1322,6 +1354,7 @@ bool SetMindIRGraphAction(const ResourcePtr &res) {
if (!is_equal_input_args) { if (!is_equal_input_args) {
// Use InferMindir which will find c++ infer in eval_map and backend_eval_map; // Use InferMindir which will find c++ infer in eval_map and backend_eval_map;
(void)InferMindir(res->func_graph(), args_spec_list, true); (void)InferMindir(res->func_graph(), args_spec_list, true);
} }
return true; return true;
@ -1342,6 +1375,7 @@ bool PreAdActionPyStub(const ResourcePtr &res) {
return true; return true;
} }
// Run the Python optimization procedure on the computation graph associated with the resource pointer
bool OptActionVmPyStub(const ResourcePtr &res) { bool OptActionVmPyStub(const ResourcePtr &res) {
if (ActionPyStub(res, opt::python_pass::Phase::OPT)) { if (ActionPyStub(res, opt::python_pass::Phase::OPT)) {
if (opt::python_pass::PyPassManager::GetInstance()->ShouldRenorm()) { if (opt::python_pass::PyPassManager::GetInstance()->ShouldRenorm()) {
@ -1384,12 +1418,12 @@ bool OptActionGePyStub(const ResourcePtr &res) {
return true; return true;
} }
// Returns a vector containing multiple Actionitems
static std::vector<ActionItem> CommonPipeline() { static std::vector<ActionItem> CommonPipeline() {
std::vector<ActionItem> actions; std::vector<ActionItem> actions;
// Parse the python ast to ANF graph // Parse the python ast to ANF graph
(void)actions.emplace_back(std::make_pair("parse", ParseAction)); (void)actions.emplace_back(std::make_pair("parse", ParseAction));
// Resolve the python func // Resolve the python func
(void)actions.emplace_back(std::make_pair("symbol_resolve", SymbolResolveAction)); (void)actions.emplace_back(std::make_pair("symbol_resolve", SymbolResolveAction));
@ -1434,18 +1468,13 @@ std::vector<ActionItem> VmPipeline(const ResourcePtr &resource) {
// If enable compilation cache and the cache is read successfully, only do the backend actions. // If enable compilation cache and the cache is read successfully, only do the backend actions.
if (!resource->EnableCompileCache() || resource->func_graph() == nullptr) { if (!resource->EnableCompileCache() || resource->func_graph() == nullptr) {
actions = CommonPipeline(); actions = CommonPipeline();
// Optimize // Optimize
(void)actions.emplace_back(std::make_pair("optimize", VmOptimizeAction)); (void)actions.emplace_back(std::make_pair("optimize", VmOptimizeAction));
// Add opt-stage python pass stub // Add opt-stage python pass stub
(void)actions.emplace_back(std::make_pair("py_opt", OptActionVmPyStub)); (void)actions.emplace_back(std::make_pair("py_opt", OptActionVmPyStub));
(void)actions.emplace_back(std::make_pair("auto_monad_reorder", OrderEnforceAction)); (void)actions.emplace_back(std::make_pair("auto_monad_reorder", OrderEnforceAction));
// Eliminate forward cnode for grad graph // Eliminate forward cnode for grad graph
(void)actions.emplace_back(std::make_pair("eliminate_forward_cnode", EliminateForwardCNode)); (void)actions.emplace_back(std::make_pair("eliminate_forward_cnode", EliminateForwardCNode));
// Eliminate the virtual mirror node // Eliminate the virtual mirror node
(void)actions.emplace_back(std::make_pair("eliminate_ad_related_special_op_node", EliminateAdRelatedSpecialOpNode)); (void)actions.emplace_back(std::make_pair("eliminate_ad_related_special_op_node", EliminateAdRelatedSpecialOpNode));
@ -1468,7 +1497,6 @@ std::vector<ActionItem> VmPipeline(const ResourcePtr &resource) {
#endif #endif
// Compile the ANF graph // Compile the ANF graph
(void)actions.emplace_back(std::make_pair("task_emit", TaskEmitAction)); (void)actions.emplace_back(std::make_pair("task_emit", TaskEmitAction));
// Execute the graph // Execute the graph
(void)actions.emplace_back(std::make_pair("execute", ExecuteAction)); (void)actions.emplace_back(std::make_pair("execute", ExecuteAction));

View File

@ -47,6 +47,8 @@ constexpr char kRolePServer[] = "pserver_";
constexpr char kRolePScheduler[] = "pscheduler_"; constexpr char kRolePScheduler[] = "pscheduler_";
constexpr char kGroupCkptFileName[] = "group.ckpt"; constexpr char kGroupCkptFileName[] = "group.ckpt";
// Get cache path defined by user.
// The cache path is in MsContext.
std::string GetUserDefinedCachePath() { std::string GetUserDefinedCachePath() {
auto user_defined_path = MsContext::GetInstance()->get_param<std::string>(MS_CTX_COMPILE_CACHE_PATH); auto user_defined_path = MsContext::GetInstance()->get_param<std::string>(MS_CTX_COMPILE_CACHE_PATH);
if (!user_defined_path.empty()) { if (!user_defined_path.empty()) {
@ -68,6 +70,9 @@ std::string GetCompileCacheDir() {
return compile_cache_dir; return compile_cache_dir;
} }
// Get current role.
// Not support for windows.
// Roles: kRoleServer, kRolePServer, kRolePScheduler
std::string GetRole() { std::string GetRole() {
#if ((defined ENABLE_CPU) && (!defined _WIN32)) #if ((defined ENABLE_CPU) && (!defined _WIN32))
const std::string &server_mode = ps::PSContext::instance()->server_mode(); const std::string &server_mode = ps::PSContext::instance()->server_mode();
@ -110,6 +115,7 @@ std::string GetDepFilesHashPath() {
std::string GetGroupCkptSavePath() { return GetCompileCacheDir() + "/" + kGroupCkptFileName; } std::string GetGroupCkptSavePath() { return GetCompileCacheDir() + "/" + kGroupCkptFileName; }
// Get hash code of compiled dependency files.
std::string GetCompileDepFilesHash(const py::list &dep_files) { std::string GetCompileDepFilesHash(const py::list &dep_files) {
MS_LOG(DEBUG) << "Dependency files size: " << dep_files.size(); MS_LOG(DEBUG) << "Dependency files size: " << dep_files.size();
std::vector<std::string> dep_files_path; std::vector<std::string> dep_files_path;
@ -251,6 +257,8 @@ bool CompileCacheManager::CheckDepFilesHashConsistency() {
return true; return true;
} }
// Load and return the cached function graph based on parallel mode and compilation cache information.
// If loading fails, perform all compilation operations and return a null pointer
FuncGraphPtr CompileCacheManager::GetCachedFuncGraph(const FuncGraphManagerPtr &manager, const py::dict &weights, FuncGraphPtr CompileCacheManager::GetCachedFuncGraph(const FuncGraphManagerPtr &manager, const py::dict &weights,
const std::string &queue_name) { const std::string &queue_name) {
// Determine whether to load parallel information. // Determine whether to load parallel information.

View File

@ -541,3 +541,6 @@ PYBIND11_MODULE(_c_expression, m) {
#endif #endif
(void)m.def("_ms_memory_recycle", &mindspore::pipeline::MemoryRecycle, "Recycle memory used by mindspore."); (void)m.def("_ms_memory_recycle", &mindspore::pipeline::MemoryRecycle, "Recycle memory used by mindspore.");
} }
//This file is mainly used to initialize functions
//and facilitate direct calls to subsequent files.

View File

@ -42,6 +42,7 @@ FunctionBlock::FunctionBlock(const Parser &parser) : parser_(parser) {
void FunctionBlock::AddPrevBlock(const FunctionBlockPtr &block) { prev_blocks_.push_back(block.get()); } void FunctionBlock::AddPrevBlock(const FunctionBlockPtr &block) { prev_blocks_.push_back(block.get()); }
// Determine whether a node can be isolated based on its type, name information, and whether it has side effects.
static bool CanBeIsolatedNode(const std::string &var_name, const AnfNodePtr &node) { static bool CanBeIsolatedNode(const std::string &var_name, const AnfNodePtr &node) {
auto cnode = dyn_cast<CNode>(node); auto cnode = dyn_cast<CNode>(node);
if (cnode == nullptr || cnode->inputs().empty()) { if (cnode == nullptr || cnode->inputs().empty()) {
@ -107,6 +108,8 @@ void FunctionBlock::WriteVariable(const std::string &var_name, const AnfNodePtr
} }
} }
// Based on the given variable name, in the assigned_ Vars_ Find the corresponding node in and
// return the value of that node as a local variable. At the same time, mark that the variable has been used.
AnfNodePtr FunctionBlock::ReadLocalVariable(const std::string &var_name) { AnfNodePtr FunctionBlock::ReadLocalVariable(const std::string &var_name) {
auto found = assigned_vars_.find(var_name); auto found = assigned_vars_.find(var_name);
if (found != assigned_vars_.end()) { if (found != assigned_vars_.end()) {
@ -268,6 +271,7 @@ AnfNodePtr FunctionBlock::HandleNamespaceInfo(const py::tuple &info) {
return GetResolveNode(info); return GetResolveNode(info);
} }
// Process built-in namespace information and add it to global variables.
AnfNodePtr FunctionBlock::HandleBuiltinNamespaceInfo(const py::tuple &info) { AnfNodePtr FunctionBlock::HandleBuiltinNamespaceInfo(const py::tuple &info) {
constexpr size_t closure_info_size = 2; constexpr size_t closure_info_size = 2;
constexpr size_t namespace_info_size = 4; constexpr size_t namespace_info_size = 4;
@ -336,6 +340,7 @@ AnfNodePtr FunctionBlock::MakeResolveSymbol(const std::string &value) {
} }
} }
// Create a parsing operation and return the parsing node for subsequent processing and use
AnfNodePtr FunctionBlock::MakeResolveOperation(const std::string &value) { AnfNodePtr FunctionBlock::MakeResolveOperation(const std::string &value) {
auto ast = parser_.ast(); auto ast = parser_.ast();
MS_EXCEPTION_IF_NULL(ast); MS_EXCEPTION_IF_NULL(ast);
@ -395,6 +400,8 @@ void FunctionBlock::SetPhiArgument(const ParameterPtr &phi) {
} }
} }
// Search for and replace nodes in the preceding block, find a replacement node that meets the condition,
// and return it. Otherwise, return a null pointer
AnfNodePtr FunctionBlock::SearchReplaceNode(const std::string &var, const ParameterPtr &phi) { AnfNodePtr FunctionBlock::SearchReplaceNode(const std::string &var, const ParameterPtr &phi) {
AnfNodePtr arg_node = nullptr; AnfNodePtr arg_node = nullptr;
MS_LOG(DEBUG) << "Prev_blocks size: " << prev_blocks_.size(); MS_LOG(DEBUG) << "Prev_blocks size: " << prev_blocks_.size();
@ -622,6 +629,8 @@ void FunctionBlock::FindIsolatedNodes() {
void FunctionBlock::AddIsolatedNode(const AnfNodePtr &target) { isolated_nodes_.add(target); } void FunctionBlock::AddIsolatedNode(const AnfNodePtr &target) { isolated_nodes_.add(target); }
// Before returning the function block, add isolated nodes to the dependency and create a new depend_ node.
// The node replaces the original output node, achieving the effect of adding dependency on the new state node on the original output node.
void FunctionBlock::AttachIsolatedNodesBeforeReturn() { void FunctionBlock::AttachIsolatedNodesBeforeReturn() {
if (isolated_nodes_.empty()) { if (isolated_nodes_.empty()) {
return; return;

View File

@ -130,6 +130,8 @@ void Parser::CleanParserResource() {
ScopeManager::GetInstance().ClearScope(); ScopeManager::GetInstance().ClearScope();
} }
// This function is used to check for missing return statements in the function graph
// and throw an exception when missing statements are found.
void CheckFuncReturn(const FuncGraphPtr &fn, const std::shared_ptr<ParseFunctionAst> &ast) { void CheckFuncReturn(const FuncGraphPtr &fn, const std::shared_ptr<ParseFunctionAst> &ast) {
// Check whether the functions referred by this function and itself are missing 'return' statement // Check whether the functions referred by this function and itself are missing 'return' statement
auto manager = Manage(fn, false); auto manager = Manage(fn, false);
@ -137,6 +139,8 @@ void CheckFuncReturn(const FuncGraphPtr &fn, const std::shared_ptr<ParseFunction
for (const auto &func_graph : manager->func_graphs()) { for (const auto &func_graph : manager->func_graphs()) {
MS_EXCEPTION_IF_NULL(func_graph); MS_EXCEPTION_IF_NULL(func_graph);
if (func_graph->get_return() != nullptr) { if (func_graph->get_return() != nullptr) {
// If not null, it indicates that the function graph already has a return statement,
// skipping the processing of the current function graph.
continue; continue;
} }
py::object node = ast->GetAstNode(); py::object node = ast->GetAstNode();
@ -152,6 +156,8 @@ void CheckFuncReturn(const FuncGraphPtr &fn, const std::shared_ptr<ParseFunction
} }
} }
// Find free variables in the function graph and return information about these free variables.
// Free variables refer to input objects that reference other function graphs in the current function graph.
std::vector<std::pair<CNodePtr, size_t>> GetFreeVariable(const FuncGraphPtr &func_graph) { std::vector<std::pair<CNodePtr, size_t>> GetFreeVariable(const FuncGraphPtr &func_graph) {
// Considering the performance, we didn't use Manager here. // Considering the performance, we didn't use Manager here.
std::vector<std::pair<CNodePtr, size_t>> free_variables; std::vector<std::pair<CNodePtr, size_t>> free_variables;
@ -183,6 +189,11 @@ std::vector<std::pair<CNodePtr, size_t>> GetFreeVariable(const FuncGraphPtr &fun
return free_variables; return free_variables;
} }
// This function is used to elevate the free variable of the scrolling function to
// its caller's parameter list, and modify the input parameter nodes in the referenced
// function graph and the input nodes in the calling node. This process,
// also known as free variable lifting, can eliminate dependencies between function graphs
// and achieve better reusability between them.
void Parser::LiftRolledBodyGraphFV() { void Parser::LiftRolledBodyGraphFV() {
for (auto &rolled_call_pair : rolled_body_calls_) { for (auto &rolled_call_pair : rolled_body_calls_) {
auto rolled_call_cnode = rolled_call_pair.first; auto rolled_call_cnode = rolled_call_pair.first;
@ -204,6 +215,11 @@ void Parser::LiftRolledBodyGraphFV() {
} }
} }
// This function is used to elevate a free variable in a conditional statement branch
// to its caller's parameter list, and modify the input parameter nodes and call nodes
// in the true and false branch function graphs. This process is similar to the free variable
// enhancement of rolling functions, which can eliminate dependencies between function
// graphs and improve code reusability.
void Parser::LiftIfBranchGraphFV() { void Parser::LiftIfBranchGraphFV() {
for (auto &branch_call_tuple : if_branch_calls_) { for (auto &branch_call_tuple : if_branch_calls_) {
auto call_cnode = std::get<0>(branch_call_tuple); auto call_cnode = std::get<0>(branch_call_tuple);
@ -245,6 +261,11 @@ void Parser::LiftIfBranchGraphFV() {
} }
namespace { namespace {
// This function converts the first half of the parallel call into a call to the
// intermediate graph, and achieves the conversion process by creating a
// new input node list and updating the output nodes of the first half of the
// call graph. This process may involve graph optimization or automatic parallelization
// in deep learning frameworks.
void TransformParallelCallFormerToMiddle(const FuncGraphPtr &former_call_graph, const FuncGraphPtr &latter_call_graph, void TransformParallelCallFormerToMiddle(const FuncGraphPtr &former_call_graph, const FuncGraphPtr &latter_call_graph,
size_t middle_graph_output_cnode_size, bool use_arguments_pack) { size_t middle_graph_output_cnode_size, bool use_arguments_pack) {
// The 'former_graph_output' is middle graph call. // The 'former_graph_output' is middle graph call.
@ -264,6 +285,12 @@ void TransformParallelCallFormerToMiddle(const FuncGraphPtr &former_call_graph,
former_call_graph->set_output(new_output); former_call_graph->set_output(new_output);
} }
// This function converts the call to the middle graph into a call to the second half,
// determines whether parameter packaging (tuples) is needed based on the number
// of input parameters, adjusts the output of the middle graph based on the existence
// of dependent nodes, and updates the output of the second half call graph.
// This process may be used in areas such as graph optimization or automatic parallelization
// in deep learning frameworks.
bool TransformParallelCallMiddleToLatter(const FuncGraphPtr &middle_call_graph, bool TransformParallelCallMiddleToLatter(const FuncGraphPtr &middle_call_graph,
const CNodePtr &middle_graph_output_cnode, const CNodePtr &middle_graph_output_cnode,
const AnfNodePtr &middle_graph_dependency_node, const AnfNodePtr &middle_graph_dependency_node,
@ -301,6 +328,11 @@ bool IsDependOfIsolatedNodes(const AnfNodePtr &node) {
return sort_rhs_first; return sort_rhs_first;
} }
// This function is used to obtain the actual output nodes of the intermediate graph.
// If the output node of the middle graph is a null pointer, an exception is thrown.
// If the output node is a Dependent node, obtain the actual output node and dependent
// node, and return them. This function may be used in scenarios such as graph
// optimization or graph transformation in deep learning frameworks.
std::pair<CNodePtr, AnfNodePtr> GetRealMiddleOutputNodes(const FuncGraphPtr &middle_call_graph) { std::pair<CNodePtr, AnfNodePtr> GetRealMiddleOutputNodes(const FuncGraphPtr &middle_call_graph) {
auto middle_graph_output = middle_call_graph->output(); auto middle_graph_output = middle_call_graph->output();
if (middle_graph_output == nullptr) { if (middle_graph_output == nullptr) {
@ -374,6 +406,12 @@ void Parser::TransformParallelCall() {
LiftRolledBodyGraphFV(); LiftRolledBodyGraphFV();
} }
// Parse functions in Python code and convert them into FuncGraph objects.
// In the parsing process, first determine whether the function type is FunctionDef
// or Lambda based on the node type in AST, and then call the corresponding parsing
// functions for parsing. After parsing is completed, a series of post-processing is required,
// including removing irrelevant Phi functions, checking function return values, and concurrently
// calling and replacing nodes. Finally, the parsed FuncGraph object is returned.
FuncGraphPtr Parser::ParseFuncGraph() { FuncGraphPtr Parser::ParseFuncGraph() {
// Get ast FunctionDef node // Get ast FunctionDef node
py::object node = ast_->GetAstNode(); py::object node = ast_->GetAstNode();
@ -418,6 +456,10 @@ AnfNodePtr GetMixedPrecisionCastHelp(const FuncGraphPtr &func_graph, const AnfNo
return cast; return cast;
} }
// This function generates function parameter nodes based on the attribute information of function nodes,
// and stores parameter names and node objects in the function block object block.
// When generating parameter nodes, it will determine whether there are variable length parameters and
// keyword parameters, as well as handle other parameter related information
void Parser::GenerateArgsNodeForFunction(const FunctionBlockPtr &block, const py::object &fn_node) { void Parser::GenerateArgsNodeForFunction(const FunctionBlockPtr &block, const py::object &fn_node) {
py::object func_args = python_adapter::GetPyObjAttr(fn_node, "args"); py::object func_args = python_adapter::GetPyObjAttr(fn_node, "args");
py::object var_arg_node = python_adapter::GetPyObjAttr(func_args, "vararg"); py::object var_arg_node = python_adapter::GetPyObjAttr(func_args, "vararg");
@ -453,6 +495,9 @@ void Parser::GenerateArgsNodeForFunction(const FunctionBlockPtr &block, const py
} }
} }
// Generate default values for function parameters and store them in the function
// graph object. When generating default values, it will determine whether the
// default values for the parsing target type and parameters are None and process them separately.
void Parser::GenerateArgsDefaultValueForFunction(const FunctionBlockPtr &block, const py::object &fn_node) { void Parser::GenerateArgsDefaultValueForFunction(const FunctionBlockPtr &block, const py::object &fn_node) {
MS_EXCEPTION_IF_NULL(block); MS_EXCEPTION_IF_NULL(block);
py::list defaults = ast_->GetArgsDefaultValues(fn_node); py::list defaults = ast_->GetArgsDefaultValues(fn_node);
@ -493,6 +538,7 @@ ScopePtr Parser::GetScopeForParseFunction() {
return scope; return scope;
} }
// Parse and generate function nodes.
FunctionBlockPtr Parser::ParseDefFunction(const py::object &node, const FunctionBlockPtr &block) { FunctionBlockPtr Parser::ParseDefFunction(const py::object &node, const FunctionBlockPtr &block) {
ScopePtr scope = GetScopeForParseFunction(); ScopePtr scope = GetScopeForParseFunction();
// The node created in the parsefunction context, will inherit the scope created using scope_guard // The node created in the parsefunction context, will inherit the scope created using scope_guard
@ -559,6 +605,7 @@ FunctionBlockPtr Parser::ParseDefFunction(const py::object &node, const Function
return func_block; return func_block;
} }
// Parse and generate lambda function nodes.
FunctionBlockPtr Parser::ParseLambdaFunction(const py::object &node, const FunctionBlockPtr &block) { FunctionBlockPtr Parser::ParseLambdaFunction(const py::object &node, const FunctionBlockPtr &block) {
MS_EXCEPTION_IF_NULL(ast_); MS_EXCEPTION_IF_NULL(ast_);
ScopePtr scope = GetScopeForParseFunction(); ScopePtr scope = GetScopeForParseFunction();
@ -731,6 +778,7 @@ void Parser::UpdateBlockPyParams(const FunctionBlockPtr &block, const FunctionBl
block->UpdateLocalPyParam(keys, values); block->UpdateLocalPyParam(keys, values);
} }
// Used to generate conditional blocks, including true_block and false_block.
void Parser::MakeConditionBlocks(const FunctionBlockPtr &pre_block, const FunctionBlockPtr &true_block, void Parser::MakeConditionBlocks(const FunctionBlockPtr &pre_block, const FunctionBlockPtr &true_block,
const FunctionBlockPtr &false_block) { const FunctionBlockPtr &false_block) {
MS_EXCEPTION_IF_NULL(true_block); MS_EXCEPTION_IF_NULL(true_block);
@ -1002,6 +1050,7 @@ std::vector<AnfNodePtr> Parser::ParseException(const FunctionBlockPtr &block, co
return node_inputs; return node_inputs;
} }
// Parse the function call in the raise statement and return the parsing result
std::vector<AnfNodePtr> Parser::ParseRaiseCall(const FunctionBlockPtr &block, const py::object &node) { std::vector<AnfNodePtr> Parser::ParseRaiseCall(const FunctionBlockPtr &block, const py::object &node) {
MS_LOG(DEBUG) << "Process ast Call, the current node is raise."; MS_LOG(DEBUG) << "Process ast Call, the current node is raise.";
// Process function call // Process function call
@ -1095,6 +1144,9 @@ AnfNodePtr Parser::GenerateAnfNodeForCall(const FunctionBlockPtr &block, const A
return call_anf_node; return call_anf_node;
} }
// Parse the parameters of the function call and store the parsing results in packed_ Arguments
// and groups_ Arguments. At the same time, the function also returns a Boolean value of need_ Unpack,
// indicating whether the parameter needs to be unpacked.
bool Parser::ParseArgsInCall(const FunctionBlockPtr &block, const py::list &args, bool *need_fallback, bool Parser::ParseArgsInCall(const FunctionBlockPtr &block, const py::list &args, bool *need_fallback,
std::vector<AnfNodePtr> *packed_arguments, std::vector<AnfNodePtr> *group_arguments) { std::vector<AnfNodePtr> *packed_arguments, std::vector<AnfNodePtr> *group_arguments) {
MS_LOG(DEBUG) << "Process ast args in call"; MS_LOG(DEBUG) << "Process ast args in call";
@ -1124,6 +1176,9 @@ bool Parser::ParseArgsInCall(const FunctionBlockPtr &block, const py::list &args
return need_unpack; return need_unpack;
} }
// Parse keyword parameters in function calls and store the parsing results in
// packed_ Arguments. Meanwhile, the function returns a Boolean value of need_ Unpack,
// indicating whether the parameter needs to be unpacked.
bool Parser::ParseKeywordsInCall(const FunctionBlockPtr &block, const py::object &node, bool Parser::ParseKeywordsInCall(const FunctionBlockPtr &block, const py::object &node,
std::vector<AnfNodePtr> *packed_arguments) { std::vector<AnfNodePtr> *packed_arguments) {
MS_LOG(DEBUG) << "Process ast key words in call"; MS_LOG(DEBUG) << "Process ast key words in call";
@ -1252,6 +1307,23 @@ AnfNodePtr Parser::ParseCompare(const FunctionBlockPtr &block, const py::object
return new_node; return new_node;
} }
// This is a function that parses Boolean operations in Python syntax.
// The function parameters include block representing the current function
// block, value_ List represents the list of Boolean operation nodes processed,
// and mode represents the type of Boolean operation (and/or).
// When there is only one node in the node list, directly call the ParseExprNode()
// function to parse the node and return it.
// When there are multiple nodes in the node list, the first node is removed, and
// the remaining nodes rest to form a new list. Then, create two new function blocks true_ Block
// and false_ Block and hijack the tracker TraceGuard in two separate blocks to record its call stack information.
//Next, call the MakeConditionBlocks() function to create a condition block and set it,
// and then determine the Boolean operation type to select sub blocks b1 and b2. For
// the and operation, reset_ Node wrapped in b1, test_ Node wrapped in b2; For the
// or operation, test_ Node wrapped in b1, rest_ Node is wrapped in b2. Rest_ The node
// is obtained by recursively calling the ProcessBoolOpValueList() function.
//Finally, use the conditional node prim:: kPrimSwitch to convert cond_ Node as the
// branching condition, set true_ Block and false_ Run two function blocks as branches and
// switch them_ Add app to block_ In fg. Finally, switch_ The app is returned as the output of the function block.
AnfNodePtr Parser::ProcessBoolOpValueList(const FunctionBlockPtr &block, const py::list &value_list, AstSubType mode) { AnfNodePtr Parser::ProcessBoolOpValueList(const FunctionBlockPtr &block, const py::list &value_list, AstSubType mode) {
// If there is only one bool op now // If there is only one bool op now
MS_EXCEPTION_IF_NULL(block); MS_EXCEPTION_IF_NULL(block);
@ -2317,6 +2389,7 @@ void Parser::HandleAssignSubscript(const FunctionBlockPtr &block, const py::obje
block->WriteVariable(var_name, setitem_app); block->WriteVariable(var_name, setitem_app);
} }
// Choose appropriate processing methods to handle assignment statements based on the different types of target objects
void Parser::WriteAssignVars(const FunctionBlockPtr &block, const py::object &target_object, void Parser::WriteAssignVars(const FunctionBlockPtr &block, const py::object &target_object,
const AnfNodePtr &value_node) { const AnfNodePtr &value_node) {
MS_EXCEPTION_IF_NULL(value_node); MS_EXCEPTION_IF_NULL(value_node);
@ -2409,6 +2482,7 @@ bool Parser::IsTensorType(const AnfNodePtr &node, const std::string &script_text
return false; return false;
} }
// Create an interpretation node and handle global and local parameters.
AnfNodePtr Parser::MakeInterpretNode(const FunctionBlockPtr &block, const AnfNodePtr &value_node, AnfNodePtr Parser::MakeInterpretNode(const FunctionBlockPtr &block, const AnfNodePtr &value_node,
const string &script_text) { const string &script_text) {
MS_EXCEPTION_IF_NULL(block); MS_EXCEPTION_IF_NULL(block);

View File

@ -33,6 +33,9 @@ static const std::set<std::string> unchanged_named_primitive = {
parse::NAMED_PRIMITIVE_ATTRIBUTE, parse::NAMED_PRIMITIVE_NAMECONSTANT, parse::NAMED_PRIMITIVE_CONSTANT, parse::NAMED_PRIMITIVE_ATTRIBUTE, parse::NAMED_PRIMITIVE_NAMECONSTANT, parse::NAMED_PRIMITIVE_CONSTANT,
parse::NAMED_PRIMITIVE_NUM, parse::NAMED_PRIMITIVE_STR}; parse::NAMED_PRIMITIVE_NUM, parse::NAMED_PRIMITIVE_STR};
// Resolve the name of the node based on the incoming node object and
// type, and return the name. It can be used to process and distinguish
// different types of nodes during the parsing process.
std::string DynamicParser::ParseNodeName(const std::shared_ptr<parse::ParseFunctionAst> &ast, const py::object &node, std::string DynamicParser::ParseNodeName(const std::shared_ptr<parse::ParseFunctionAst> &ast, const py::object &node,
parse::AstMainType type) { parse::AstMainType type) {
MS_EXCEPTION_IF_NULL(ast); MS_EXCEPTION_IF_NULL(ast);
@ -63,6 +66,9 @@ void DynamicParser::ParseInputArgs(const std::shared_ptr<parse::ParseFunctionAst
} }
} }
// Parse the expression in the if/while statement and obtain the name of the variable
// that needs further processing. It can be used to further parse the conditional judgment
// of if/while statements and extract variable information to adapt to different contextual requirements
bool DynamicParser::ParseIfWhileExprNode(const std::shared_ptr<parse::ParseFunctionAst> &ast, const py::object &node) { bool DynamicParser::ParseIfWhileExprNode(const std::shared_ptr<parse::ParseFunctionAst> &ast, const py::object &node) {
MS_LOG(DEBUG) << "Parse if/while expr"; MS_LOG(DEBUG) << "Parse if/while expr";
py::object test_node = python_adapter::GetPyObjAttr(node, parse::NAMED_PRIMITIVE_TEST); py::object test_node = python_adapter::GetPyObjAttr(node, parse::NAMED_PRIMITIVE_TEST);
@ -112,6 +118,9 @@ bool DynamicParser::ParseIfWhileExprNode(const std::shared_ptr<parse::ParseFunct
return false; return false;
} }
// Analyze the expression in the assignment statement and determine if further
// processing is necessary. It can be used to identify variable information that needs
// to be retained and process it accordingly according to different contextual requirements.
bool DynamicParser::ParseAssignExprNode(const std::shared_ptr<parse::ParseFunctionAst> &ast, const py::object &node) { bool DynamicParser::ParseAssignExprNode(const std::shared_ptr<parse::ParseFunctionAst> &ast, const py::object &node) {
MS_LOG(DEBUG) << "Parse assign expr"; MS_LOG(DEBUG) << "Parse assign expr";
py::object value_node = python_adapter::GetPyObjAttr(node, parse::NAMED_PRIMITIVE_VALUE); py::object value_node = python_adapter::GetPyObjAttr(node, parse::NAMED_PRIMITIVE_VALUE);
@ -140,6 +149,9 @@ bool DynamicParser::ParseAssignExprNode(const std::shared_ptr<parse::ParseFuncti
return false; return false;
} }
// Parse the expression in an augmented assignment statement and determine if further processing
// is required. It compares the value of assign_prim with the values in the given string vector
// compare_prim to determine if a specific operation needs to be performed.
bool DynamicParser::ParseAugAssignExprNode(const std::shared_ptr<parse::ParseFunctionAst> &, const py::object &node, bool DynamicParser::ParseAugAssignExprNode(const std::shared_ptr<parse::ParseFunctionAst> &, const py::object &node,
const std::vector<std::string> &compare_prim) { const std::vector<std::string> &compare_prim) {
MS_LOG(DEBUG) << "Parse augassign expr"; MS_LOG(DEBUG) << "Parse augassign expr";
@ -168,6 +180,9 @@ bool DynamicParser::ParseAugAssignExprNode(const std::shared_ptr<parse::ParseFun
return ret; return ret;
} }
// this function parses the body of a for expression, iterates over the nodes in the body,
// and checks if any of them are assignment expressions. It returns true if an assignment
// expression is found, otherwise false.
bool DynamicParser::ParseForExprNode(const std::shared_ptr<parse::ParseFunctionAst> &ast, const py::object &node) { bool DynamicParser::ParseForExprNode(const std::shared_ptr<parse::ParseFunctionAst> &ast, const py::object &node) {
MS_LOG(DEBUG) << "Parse for expr"; MS_LOG(DEBUG) << "Parse for expr";
py::object body_node = python_adapter::GetPyObjAttr(node, parse::NAMED_PRIMITIVE_BODY); py::object body_node = python_adapter::GetPyObjAttr(node, parse::NAMED_PRIMITIVE_BODY);
@ -188,6 +203,9 @@ bool DynamicParser::ParseForExprNode(const std::shared_ptr<parse::ParseFunctionA
return false; return false;
} }
// this function parses the body context of a function or cell by iterating over the nodes
// in the body and calling specific parsing functions based on the node type. It returns
// true if any dynamic expressions are found, otherwise false.
bool DynamicParser::ParseBodyContext(const std::shared_ptr<parse::ParseFunctionAst> &ast, const py::object &fn_node, bool DynamicParser::ParseBodyContext(const std::shared_ptr<parse::ParseFunctionAst> &ast, const py::object &fn_node,
const std::vector<std::string> &compare_prim) { const std::vector<std::string> &compare_prim) {
MS_EXCEPTION_IF_NULL(ast); MS_EXCEPTION_IF_NULL(ast);
@ -230,6 +248,8 @@ std::string DynamicParser::GetCellInfo(const py::object &cell) {
return ""; return "";
} }
// IsDynamicCell checks whether a cell contains any dynamic expressions by creating
// an AST for the cell's code and parsing its input arguments and body context
bool DynamicParser::IsDynamicCell(const py::object &cell) { bool DynamicParser::IsDynamicCell(const py::object &cell) {
std::string cell_info = GetCellInfo(cell); std::string cell_info = GetCellInfo(cell);
if (ignore_judge_dynamic_cell.find(cell_info) != ignore_judge_dynamic_cell.end()) { if (ignore_judge_dynamic_cell.find(cell_info) != ignore_judge_dynamic_cell.end()) {

View File

@ -48,6 +48,7 @@ std::string ReplaceSpecialChar(const std::string &str) {
} }
return oss.str(); return oss.str();
} }
//What this code does is replace "<" in the string with """ and ">" with """ to get a new string
struct AnfDumpHandlerRegister { struct AnfDumpHandlerRegister {
AnfDumpHandlerRegister() { AnfDumpHandlerRegister() {
@ -77,6 +78,10 @@ abstract::AbstractBasePtr ClassObject::ToAbstract() {
auto func_ptr = std::make_shared<abstract::PrimitiveAbstractClosure>(prim::kPrimMakeRecord); auto func_ptr = std::make_shared<abstract::PrimitiveAbstractClosure>(prim::kPrimMakeRecord);
return std::make_shared<abstract::PartialAbstractClosure>(func_ptr, args_spec_list); return std::make_shared<abstract::PartialAbstractClosure>(func_ptr, args_spec_list);
} }
//What this code does is create an abstract object and return its pointer.
//An abstract object consists of a PartialAbstractClosure
//that contains a PrimitiveAbstractClosure object
//and a parameter list containing an AbstractScalar object as part of the application
static inline bool IsSupportedCreateInstanceType(const py::object &obj) { static inline bool IsSupportedCreateInstanceType(const py::object &obj) {
py::module mod = python_adapter::GetPyModule(PYTHON_MOD_PARSE_MODULE); py::module mod = python_adapter::GetPyModule(PYTHON_MOD_PARSE_MODULE);
@ -87,6 +92,9 @@ static inline bool IsSupportedCreateInstanceType(const py::object &obj) {
} }
return res.cast<bool>(); return res.cast<bool>();
} }
//The purpose of this code is to call a function in Python to determine
// whether a given Python object is a type that supports creating instances,
//and return the result
abstract::AbstractBasePtr ClassType::ToAbstract() { abstract::AbstractBasePtr ClassType::ToAbstract() {
auto abs_scalar = auto abs_scalar =
@ -194,6 +202,9 @@ void BroadenCNodeAbstract(const FuncGraphPtr &func_graph) {
} }
} }
} }
//The function of this code is to expand the CNode node in the given function graph,
// implement the extension by calling the Broaden function,
//and update the abstract properties of the node
void ConvertLoadedGraph(const FuncGraphPtr &func_graph, const ValuePtr &value) { void ConvertLoadedGraph(const FuncGraphPtr &func_graph, const ValuePtr &value) {
if (!value->isa<FuncGraph>()) { if (!value->isa<FuncGraph>()) {
@ -224,6 +235,8 @@ void ConvertLoadedGraph(const FuncGraphPtr &func_graph, const ValuePtr &value) {
resolved_graph->set_parameters(input_params); resolved_graph->set_parameters(input_params);
BroadenCNodeAbstract(resolved_graph); BroadenCNodeAbstract(resolved_graph);
} }
//The purpose of this code is to convert the loaded subgraph object into the corresponding FuncGraph object,
//and update the parameter information and expand the abstract information
bool ResolveObjectToNode(const FuncGraphPtr &func_graph, const py::object &obj, AnfNodePtr *const node) { bool ResolveObjectToNode(const FuncGraphPtr &func_graph, const py::object &obj, AnfNodePtr *const node) {
AnfNodePtr output = nullptr; AnfNodePtr output = nullptr;
@ -266,6 +279,10 @@ bool ResolveObjectToNode(const FuncGraphPtr &func_graph, const py::object &obj,
*node = output; *node = output;
return true; return true;
} }
//The function of this code is to parse the incoming Python object
//into the corresponding AnfNodePtr node,
// and perform parameter parsing, creating CNode or value nodes
//according to different situations, and converting data types
bool IsAllFuncInValueSequence(const std::vector<ValuePtr> &value_vec) { bool IsAllFuncInValueSequence(const std::vector<ValuePtr> &value_vec) {
if (value_vec.empty()) { if (value_vec.empty()) {
@ -285,6 +302,9 @@ bool IsAllFuncInValueSequence(const std::vector<ValuePtr> &value_vec) {
} }
return true; return true;
} }
//The purpose of this code is to determine
//whether the element types in the incoming value_vec are all FuncGraph or Primitive objects.
//Returns false if there are non-eligible elements.
AnfNodePtr TransformToMakeTupleNodes(const FuncGraphManagerPtr &manager, const FuncGraphPtr &func_graph, AnfNodePtr TransformToMakeTupleNodes(const FuncGraphManagerPtr &manager, const FuncGraphPtr &func_graph,
const std::vector<ValuePtr> &value_vec) { const std::vector<ValuePtr> &value_vec) {
@ -310,8 +330,10 @@ AnfNodePtr TransformToMakeTupleNodes(const FuncGraphManagerPtr &manager, const F
auto cnode = func_graph->NewCNode(std::move(nodes)); auto cnode = func_graph->NewCNode(std::move(nodes));
return cnode; return cnode;
} }
//The function of this code is to convert the incoming value_vec into a MakeTuple node,
//and create and add nodes to the node vector nodes according to different situations
// Transform the ValueTuple or ValueList of graph/primitive node to make tuple of const graph/primitive node // Transform the ValueTuple or ValueList of graph/primitive node to make tuple of const graph/primitive node
bool TransformVectorFuncValueNode(const FuncGraphManagerPtr &manager, const FuncGraphPtr &func_graph, bool TransformVectorFuncValueNode(const FuncGraphManagerPtr &manager, const FuncGraphPtr &func_graph,
const ValueNodePtr &value_node, AnfNodePtr *const transformed) { const ValueNodePtr &value_node, AnfNodePtr *const transformed) {
MS_EXCEPTION_IF_NULL(value_node); MS_EXCEPTION_IF_NULL(value_node);
@ -320,14 +342,14 @@ bool TransformVectorFuncValueNode(const FuncGraphManagerPtr &manager, const Func
return false; return false;
} }
// (1) The celllist or ordered_cell will be parsed as valuetuple of const graph in it, // (1) The celllist or ordered_cell will be parsed as valuetuple of const graph in it,
// So if has graph in list, try to replace the node with make tuple of graph value node. // So if has graph in list, try to replace the node with make tuple of graph value node.
// We do this because the graph manager won't investigate the graph inside valuetuple, // We do this because the graph manager won't investigate the graph inside valuetuple,
// change the vector of graph to be make_tuple of graph value node. // change the vector of graph to be make_tuple of graph value node.
// (2) the primitive valuetuple or valuelist may encounter to abstract error, make it all // (2) the primitive valuetuple or valuelist may encounter to abstract error, make it all
// independent nodes. // independent nodes.
auto node_tuple_graphs = TransformToMakeTupleNodes(manager, func_graph, value_vec); auto node_tuple_graphs = TransformToMakeTupleNodes(manager, func_graph, value_vec);
// Replace the ret ptr to be make tuple of graph value node // Replace the ret ptr to be make tuple of graph value node
*transformed = node_tuple_graphs; *transformed = node_tuple_graphs;
return true; return true;
@ -348,16 +370,16 @@ AnfNodePtr ResolveObjectAndAddToManager(const FuncGraphManagerPtr &manager, cons
manager->AddFuncGraph(new_fg); manager->AddFuncGraph(new_fg);
} }
// If the constant node is constant of vector of graph, add graph to manager. // If the constant node is constant of vector of graph, add graph to manager.
if (IsValueNode<ValueTuple>(resolved_node) || IsValueNode<ValueList>(resolved_node)) { if (IsValueNode<ValueTuple>(resolved_node) || IsValueNode<ValueList>(resolved_node)) {
(void)TransformVectorFuncValueNode(manager, node->func_graph(), resolved_node->cast<ValueNodePtr>(), (void)TransformVectorFuncValueNode(manager, node->func_graph(), resolved_node->cast<ValueNodePtr>(),
&resolved_node); &resolved_node);
} }
return resolved_node; return resolved_node;
} }
} // namespace } // namespace
// Get python object with index from a list or the whole list if the index is not fixed. // Get python object with index from a list or the whole list if the index is not fixed.
py::object GetObjectFromSequence(const NameSpacePtr &name_space, const SymbolPtr &symbol, const AnfNodePtr &node, py::object GetObjectFromSequence(const NameSpacePtr &name_space, const SymbolPtr &symbol, const AnfNodePtr &node,
const AnfNodePtr &index_node) { const AnfNodePtr &index_node) {
MS_EXCEPTION_IF_NULL(node); MS_EXCEPTION_IF_NULL(node);
@ -375,7 +397,7 @@ py::object GetObjectFromSequence(const NameSpacePtr &name_space, const SymbolPtr
// Index is not fixed, return the whole list. // Index is not fixed, return the whole list.
return obj; return obj;
} }
// It index is a value node, get the item of index directly. // It index is a value node, get the item of index directly.
const std::string fn = PYTHON_MOD_GET_ITEM_FROM_SEQUENCE; const std::string fn = PYTHON_MOD_GET_ITEM_FROM_SEQUENCE;
const std::string module = "mindspore._extends.parse.parser"; const std::string module = "mindspore._extends.parse.parser";
auto index = imm_value->value(); auto index = imm_value->value();
@ -511,6 +533,10 @@ bool IsGetItemCNode(const AnfNodePtr &node) {
constexpr auto prim_index = 0; constexpr auto prim_index = 0;
return IsResolveNodeWithGetItem(cnode->input(prim_index)); return IsResolveNodeWithGetItem(cnode->input(prim_index));
} }
//The purpose of this code is to determine
//whether the incoming node is a GetItem node,
//by checking the node type, the number of inputs,
//and the node that parses the GetItem. Returns true if all conditions are met. Otherwise, false is returned.
AnfNodePtr ResolveMsClassWithAttr(const FuncGraphManagerPtr &manager, const MsClassObjectPtr &ms_class, AnfNodePtr ResolveMsClassWithAttr(const FuncGraphManagerPtr &manager, const MsClassObjectPtr &ms_class,
const std::string &attr, const AnfNodePtr &node) { const std::string &attr, const AnfNodePtr &node) {
@ -551,6 +577,8 @@ bool ResolveFuncGraph(const FuncGraphPtr &func_graph, const pipeline::ResourceBa
MS_LOG(ERROR) << "func_graph or resource is null"; MS_LOG(ERROR) << "func_graph or resource is null";
return false; return false;
} }
//What this code does is parse the incoming func_graph
//and print an error message and return false if the parameter is invalid
opt::irpass::ResolveIRPassLib irpass; opt::irpass::ResolveIRPassLib irpass;
opt::OptimizerPtr opt_resolve = opt::OptimizerPtr opt_resolve =
opt::Optimizer::MakeOptimizer("opt_resolve", res, GetOptResolvePasses(irpass), false, false, false); opt::Optimizer::MakeOptimizer("opt_resolve", res, GetOptResolvePasses(irpass), false, false, false);

View File

@ -90,6 +90,9 @@ bool SimplifyDataStructuresPass(const ResourcePtr &res) {
UpdateArgsSpec(func_graph, res); UpdateArgsSpec(func_graph, res);
return true; return true;
} }
//Optimization of Simplified Data Structures PASS aims to
//optimize and simplify data structures in function graphs
//to improve compute performance and reduce memory footprint
bool TransformTopGraphPass(const ResourcePtr &res) { bool TransformTopGraphPass(const ResourcePtr &res) {
MS_EXCEPTION_IF_NULL(res); MS_EXCEPTION_IF_NULL(res);
@ -109,6 +112,10 @@ bool TransformTopGraphPass(const ResourcePtr &res) {
} }
return true; return true;
} }
//The transformation pass of the top-level function graph
//mainly performs conversion operations on the tuple inputs
//that may exist in the function graph,
//splitting the tuple parameters into a single parameter for subsequent optimization and processing
bool CleanAfterOptAPass(const ResourcePtr &res) { bool CleanAfterOptAPass(const ResourcePtr &res) {
MS_EXCEPTION_IF_NULL(res); MS_EXCEPTION_IF_NULL(res);
@ -118,6 +125,11 @@ bool CleanAfterOptAPass(const ResourcePtr &res) {
UpdateArgsSpec(func_graph, res); UpdateArgsSpec(func_graph, res);
return true; return true;
} }
//The purpose of cleaning up the optimized pass is to
//do some additional processing and cleaning on the optimized function graph
//to improve the readability and execution efficiency of the code.
//This pass may perform some cleaning operations according to the specific optimization technique
//to ensure that the structure and logic of the function graph are correct and optimal.
FuncGraphPtr PrimBpOptPassStep1(const opt::irpass::OptimizeIRPassLib &irpass, const ResourcePtr &res) { FuncGraphPtr PrimBpOptPassStep1(const opt::irpass::OptimizeIRPassLib &irpass, const ResourcePtr &res) {
MS_EXCEPTION_IF_NULL(res); MS_EXCEPTION_IF_NULL(res);
@ -125,6 +137,8 @@ FuncGraphPtr PrimBpOptPassStep1(const opt::irpass::OptimizeIRPassLib &irpass, co
opt::OptPassConfig pynative_eliminate = opt::OptPassConfig({ opt::OptPassConfig pynative_eliminate = opt::OptPassConfig({
irpass.pynative_eliminate_, irpass.pynative_eliminate_,
}); });
//Provides a flexible way to configure and organize the execution of optimized passes,
//selecting the required passes and setting their parameters according to your needs
opt::OptPassConfig switch_simplify = opt::OptPassConfig({ opt::OptPassConfig switch_simplify = opt::OptPassConfig({
irpass.switch_simplify_, irpass.switch_simplify_,
@ -144,6 +158,9 @@ FuncGraphPtr PrimBpOptPassStep1(const opt::irpass::OptimizeIRPassLib &irpass, co
}; };
return func_graph; return func_graph;
} }
//By defining different pass groupings and organizing them in the desired order,
//complex optimization logic can be implemented
//and multiple rounds of optimization of the function graph can be performed to achieve the final result.
FuncGraphPtr PrimBpOptPassStep2(const opt::irpass::OptimizeIRPassLib &irpass, const ResourcePtr &res) { FuncGraphPtr PrimBpOptPassStep2(const opt::irpass::OptimizeIRPassLib &irpass, const ResourcePtr &res) {
MS_EXCEPTION_IF_NULL(res); MS_EXCEPTION_IF_NULL(res);
@ -162,10 +179,10 @@ FuncGraphPtr PrimBpOptPassStep2(const opt::irpass::OptimizeIRPassLib &irpass, co
auto re_auto_monadwrapper = [](const FuncGraphPtr &root, const opt::OptimizerPtr &) -> bool { auto re_auto_monadwrapper = [](const FuncGraphPtr &root, const opt::OptimizerPtr &) -> bool {
return ReAutoMonad(root); return ReAutoMonad(root);
}; };
OptPassGroupMap map({{"ad_renormalize", opt::OptPassConfig::Renormalize()}, OptPassGroupMap map({{"ad_renormalize", opt::OptPassConfig::Renormalize()},//An optimized pass configuration object for renormalization
{"ad_inline", inline_opt}, {"ad_inline", inline_opt},//Used for inline optimization
{"ad_special_op_simplify", special_op_simplify}, {"ad_special_op_simplify", special_op_simplify},//Simplified optimization for performing special operations
{"auto_monad_grad", opt::OptPassConfig(re_auto_monadwrapper)}}); {"auto_monad_grad", opt::OptPassConfig(re_auto_monadwrapper)}});//Used to automatically process Monad in automatic differentiation and generate gradient functions.
auto prim_bprop_opt_step_2 = opt::Optimizer::MakeOptimizer("prim_bprop_opt_step_2", res, map); auto prim_bprop_opt_step_2 = opt::Optimizer::MakeOptimizer("prim_bprop_opt_step_2", res, map);
FuncGraphPtr func_graph = res->func_graph(); FuncGraphPtr func_graph = res->func_graph();
@ -209,6 +226,10 @@ FuncGraphPtr BpropGraphFinalOptPass(const ResourcePtr &res) {
}); });
(void)map.emplace_back(std::make_pair("environ_eliminate", environ_eliminate)); (void)map.emplace_back(std::make_pair("environ_eliminate", environ_eliminate));
} }
//Three additional pass groupings are dynamically added to the map
//as needed for appropriate optimization steps in subsequent optimization processes.
//This gives you the flexibility to configure and adjust optimization processes
// to your specific needs for better performance and results.
auto bprop_graph_final_opt = opt::Optimizer::MakeOptimizer("bprop_graph_final_opt", res, map); auto bprop_graph_final_opt = opt::Optimizer::MakeOptimizer("bprop_graph_final_opt", res, map);
FuncGraphPtr func_graph = res->func_graph(); FuncGraphPtr func_graph = res->func_graph();
@ -232,6 +253,7 @@ bool parallel_mode() {
std::string parallel_mode = parallel::ParallelContext::GetInstance()->parallel_mode(); std::string parallel_mode = parallel::ParallelContext::GetInstance()->parallel_mode();
return (parallel_mode == parallel::kAutoParallel) || (parallel_mode == parallel::kSemiAutoParallel); return (parallel_mode == parallel::kAutoParallel) || (parallel_mode == parallel::kSemiAutoParallel);
} }
//Determine whether you are currently in parallel mode
void AddParallelRenormalize(OptPassGroupMap *map_a) { void AddParallelRenormalize(OptPassGroupMap *map_a) {
if (parallel_mode()) { if (parallel_mode()) {
@ -242,6 +264,9 @@ void AddParallelRenormalize(OptPassGroupMap *map_a) {
} }
} }
} }
//The purpose of this code is to find the optimization step group "meta_fg_expand"
// based on whether it is currently in parallel mode,
//and insert a parallel optimization step group named "parallel_renormalize" before the combination.
opt::OptPassConfig GetOptPassA1(const opt::irpass::OptimizeIRPassLib &irpass) { opt::OptPassConfig GetOptPassA1(const opt::irpass::OptimizeIRPassLib &irpass) {
return opt::OptPassConfig({ return opt::OptPassConfig({
@ -304,6 +329,9 @@ opt::OptPassConfig GetGeTensorArrayPass(const opt::irpass::OptimizeIRPassLib &ir
irpass.ge_tensor_array_cast_index_, irpass.ge_tensor_array_cast_index_,
}); });
} }
//What this code does is create a function called GetGeTensorArrayPass
//that encapsulates two pass functions
//and returns an optimized pass configuration object containing both passes
OptPassGroupMap GetOptPassesA(const opt::irpass::OptimizeIRPassLib &irpass) { OptPassGroupMap GetOptPassesA(const opt::irpass::OptimizeIRPassLib &irpass) {
opt::OptPassConfig a_1 = GetOptPassA1(irpass); opt::OptPassConfig a_1 = GetOptPassA1(irpass);
@ -391,6 +419,9 @@ OptPassGroupMap GetA1A2(const opt::irpass::OptimizeIRPassLib &irpass) {
OptPassGroupMap a1_a2(opt_a.begin(), opt_a.begin() + a1_a2_len); OptPassGroupMap a1_a2(opt_a.begin(), opt_a.begin() + a1_a2_len);
return a1_a2; return a1_a2;
} }
//What this code does is create a function named GetA1A2
//that extracts a pass configuration combination named a1_a2
//that contains the first 9 pass configurations from the pass function obtained from irpass
OptPassGroupMap GetOptPassesAfterCconv(const opt::irpass::OptimizeIRPassLib &irpass) { OptPassGroupMap GetOptPassesAfterCconv(const opt::irpass::OptimizeIRPassLib &irpass) {
opt::OptPassConfig c_1 = opt::OptPassConfig({ opt::OptPassConfig c_1 = opt::OptPassConfig({
@ -495,6 +526,7 @@ OptPassGroupMap GetOptPassesPynativeElim(const opt::irpass::OptimizeIRPassLib &i
}); });
return map; return map;
} }
//Create an optimized pass configuration combination map that contains a pass
OptPassGroupMap GetOptPassesC(const opt::irpass::OptimizeIRPassLib &) { OptPassGroupMap GetOptPassesC(const opt::irpass::OptimizeIRPassLib &) {
return OptPassGroupMap({{"renormalize", opt::OptPassConfig::Renormalize()}}); return OptPassGroupMap({{"renormalize", opt::OptPassConfig::Renormalize()}});
@ -508,6 +540,12 @@ OptPassGroupMap GetControlPhases(const opt::irpass::OptimizeIRPassLib &) {
}); });
return map; return map;
} }
//What this code does is create an optimized pass configuration combination map with two passes.
//One of the passes has the name "control_group"
//and the corresponding configuration is a control flow optimization pass;
//The other pass has the name "renormalize",
//which corresponds to a renormalized pass.
// The function returns this configuration combination as a result.
OptPassGroupMap GetGeSpecializedPhases() { OptPassGroupMap GetGeSpecializedPhases() {
opt::OptPassConfig ge_ta_size_group = opt::OptPassConfig(opt::irpass::GeTensorArrayPrepare()); opt::OptPassConfig ge_ta_size_group = opt::OptPassConfig(opt::irpass::GeTensorArrayPrepare());
@ -519,6 +557,12 @@ OptPassGroupMap GetGeSpecializedPhases() {
}); });
return map; return map;
} }
//What this code does is create an optimized pass configuration combination map with two passes.
//One of the passes is named "ge_ta_size_group",
//and the corresponding configuration is to handle the pass of GeTensorArrayPrepare;
//The other pass, named "ge_ta_passes",
//corresponds to a set of passes used to optimize GeTensorArray.
//The function returns this configuration combination as a result
OptPassGroupMap GetOptPynativeGradEpiloguePhases(const opt::irpass::OptimizeIRPassLib &irpass) { OptPassGroupMap GetOptPynativeGradEpiloguePhases(const opt::irpass::OptimizeIRPassLib &irpass) {
auto opt_a = GetOptPassesA(irpass); auto opt_a = GetOptPassesA(irpass);
@ -578,6 +622,7 @@ void ReclaimOptimizer() {
} }
g_pass_opts.clear(); g_pass_opts.clear();
} }
//Free up optimizer-related assets
bool OptPassGroup(const ResourcePtr &res, const std::string &name) { bool OptPassGroup(const ResourcePtr &res, const std::string &name) {
MS_EXCEPTION_IF_NULL(res); MS_EXCEPTION_IF_NULL(res);
@ -585,6 +630,11 @@ bool OptPassGroup(const ResourcePtr &res, const std::string &name) {
MS_LOG(ERROR) << "Opt passes int64_t error"; MS_LOG(ERROR) << "Opt passes int64_t error";
return false; return false;
} }
//This code completes the null determination of resources
// and determines whether the function graph is empty,
//if the function graph is empty,
//it prints an error message and returns false,
//otherwise continue the subsequent optimization pass operation
FuncGraphPtr func_graph = res->func_graph(); FuncGraphPtr func_graph = res->func_graph();
MS_LOG(DEBUG) << "Start " << name << " func graph:" << func_graph->ToString() << ", " MS_LOG(DEBUG) << "Start " << name << " func graph:" << func_graph->ToString() << ", "
@ -622,6 +672,8 @@ bool SliceRecomputeActivationPass(const ResourcePtr &res) {
opt::SliceRecomputedActivationNodes(res->func_graph()); opt::SliceRecomputedActivationNodes(res->func_graph());
return true; return true;
} }
//The function of this code is to perform the SliceRecomputeActivation optimization operation
//on the function graph in the passed resource and return the optimization execution result.
bool CommOpAddAttrs(const ResourcePtr &res) { bool CommOpAddAttrs(const ResourcePtr &res) {
MS_EXCEPTION_IF_NULL(res); MS_EXCEPTION_IF_NULL(res);
@ -693,6 +745,8 @@ bool CconvPass(const ResourcePtr &res) {
res->set_func_graph(new_fg); res->set_func_graph(new_fg);
return true; return true;
} }
//The function of this code is to clone the function graph in the incoming resource
//and update the cloned function graph to the resource
bool PipelineSplitPass(const ResourcePtr &res) { return PipelineSplit(res); } bool PipelineSplitPass(const ResourcePtr &res) { return PipelineSplit(res); }

View File

@ -135,6 +135,9 @@ std::string GetBaseNameForIR(int64_t stage_idx, const std::string &action_name)
oss << std::setfill('0') << std::setw(spaces) << stage_idx << "_" << action_name; oss << std::setfill('0') << std::setw(spaces) << stage_idx << "_" << action_name;
return oss.str(); return oss.str();
} }
//The definition of an anonymous namespace
//The namespace contains an implementation of the function GetBaseNameForIR
//Based on the given stage index and action name, a baseline name is generated for IR
#endif #endif
bool CheckAllTensor(const ValueTuplePtr &value_tuple) { bool CheckAllTensor(const ValueTuplePtr &value_tuple) {
@ -147,6 +150,12 @@ bool CheckAllTensor(const ValueTuplePtr &value_tuple) {
} }
return true; return true;
} }
//A function called CheckAllTensor is defined
//Determines whether a value tuple object contains all tensors
//If there are non-tensor elements in the value tuple,
//or if the element itself is not a value tuple or MetaTensor type
//The function returns false
//The function returns true only if all elements are tensors.
AbstractBasePtr ArgsToAbstract(const ValuePtr &value, bool enable_tuple_broaden = false) { AbstractBasePtr ArgsToAbstract(const ValuePtr &value, bool enable_tuple_broaden = false) {
MS_EXCEPTION_IF_NULL(value); MS_EXCEPTION_IF_NULL(value);
@ -156,13 +165,22 @@ AbstractBasePtr ArgsToAbstract(const ValuePtr &value, bool enable_tuple_broaden
return abstract::FromValue(value, broaden); return abstract::FromValue(value, broaden);
} }
//A function called ArgsToAbstract is defined
//The purpose of this function is to convert the given ValuePtr object into an AbstractBasePtr object
//The judgment logic inside the function determines whether a type extension operation
//is required based on different types and conditions
//The final return is the converted AbstractBasePtr object
bool CheckArgValid(const py::handle &arg) { bool CheckArgValid(const py::handle &arg) {
if (py::isinstance<py::list>(arg) || py::isinstance<py::tuple>(arg)) { if (py::isinstance<py::list>(arg) || py::isinstance<py::tuple>(arg)) {
auto vector_arg = py::cast<py::list>(arg); auto vector_arg = py::cast<py::list>(arg);
return std::all_of(vector_arg.begin(), vector_arg.end(), CheckArgValid); return std::all_of(vector_arg.begin(), vector_arg.end(), CheckArgValid);
} }
//A function called CheckArgValid is defined
//What this function does is check if a given Python object is legitimate
//If the object is a list or tuple, each element in it is checked recursively
//If the object is not a list or tuple, the judgment false is returned directly
if (py::isinstance<py::dict>(arg)) { if (py::isinstance<py::dict>(arg)) {
auto dict_arg = py::cast<py::dict>(arg); auto dict_arg = py::cast<py::dict>(arg);
return std::all_of(dict_arg.begin(), dict_arg.end(), [](const auto &pair) { return CheckArgValid(pair.second); }); return std::all_of(dict_arg.begin(), dict_arg.end(), [](const auto &pair) { return CheckArgValid(pair.second); });
@ -184,7 +202,8 @@ bool CheckArgValid(const py::handle &arg) {
"For more details, please refer to the FAQ at https://www.mindspore.cn."; "For more details, please refer to the FAQ at https://www.mindspore.cn.";
} }
} }
//What this code does is check if a given object is of type Tensor or not
//Special handling of boolean tensors prints a warning message
return py::isinstance<py::int_>(arg) || py::isinstance<py::float_>(arg) || py::isinstance<py::none>(arg) || return py::isinstance<py::int_>(arg) || py::isinstance<py::float_>(arg) || py::isinstance<py::none>(arg) ||
py::isinstance<Number>(arg) || py::isinstance<Number>(arg) ||
((py::isinstance<Tensor>(arg) || py::isinstance<CSRTensor>(arg) || py::isinstance<COOTensor>(arg)) && ((py::isinstance<Tensor>(arg) || py::isinstance<CSRTensor>(arg) || py::isinstance<COOTensor>(arg)) &&
@ -215,6 +234,9 @@ void SetLoopCount(const ResourcePtr &resource) {
MS_LOG(INFO) << "Change vm_loop_flag to " << resource->vm_loop_flag() << ", set loop_size to " << loop_size; MS_LOG(INFO) << "Change vm_loop_flag to " << resource->vm_loop_flag() << ", set loop_size to " << loop_size;
} }
} }
//Set the number of cycles and the vm_loop_flag flag bit based on the current operating environment
//These flag bits are passed to the virtual machine engine
//You can use these flag bits in the virtual machine engine to control how the graph performs
std::map<string, string> GenerateJitConfigMap(const py::dict &jit_config) { std::map<string, string> GenerateJitConfigMap(const py::dict &jit_config) {
std::map<string, string> ret{}; std::map<string, string> ret{};
@ -225,6 +247,8 @@ std::map<string, string> GenerateJitConfigMap(const py::dict &jit_config) {
} }
return ret; return ret;
} }
//What this code does is convert the given dictionary jit_config of type py::dict to a key-value pair mapping of type std::map
//and returns the mapping result
void RecordInitStatus() { void RecordInitStatus() {
static bool printed = false; static bool printed = false;
@ -233,6 +257,9 @@ void RecordInitStatus() {
printed = true; printed = true;
} }
} }
//A function is defined, RecordInitStatus
//Record the system status during the system initialization phase
//Ensure that the status is logged only once
void RecordExitStatus() { MS_LOG(INFO) << "Status record: system exit."; } void RecordExitStatus() { MS_LOG(INFO) << "Status record: system exit."; }
} // namespace } // namespace
@ -258,6 +285,10 @@ void CheckArgsValid(const py::object &source_obj, const py::tuple &args) {
} }
} }
} }
//A function is defined, CheckArgsValid
//Check the validity of the input parameters
//and throw a type error exception when there are invalid parameters
py::object GraphExecutorPy::GenerateArgumentsKey(const py::tuple &args, bool enable_tuple_broaden) { py::object GraphExecutorPy::GenerateArgumentsKey(const py::tuple &args, bool enable_tuple_broaden) {
MS_LOG(DEBUG) << "GenerateArgumentsKey args size:" << args.size(); MS_LOG(DEBUG) << "GenerateArgumentsKey args size:" << args.size();
@ -277,6 +308,7 @@ py::object GraphExecutorPy::GenerateArgumentsKey(const py::tuple &args, bool ena
} }
// If cache matched no need CheckArgsValid // If cache matched no need CheckArgsValid
auto iter = g_args_cache.find(args_spec); auto iter = g_args_cache.find(args_spec);
if (iter != g_args_cache.end()) { if (iter != g_args_cache.end()) {
return py::int_(iter->second); return py::int_(iter->second);
@ -287,14 +319,18 @@ py::object GraphExecutorPy::GenerateArgumentsKey(const py::tuple &args, bool ena
MS_LOG(INFO) << "Generate a new compile key for new args, key: " << key_counter; MS_LOG(INFO) << "Generate a new compile key for new args, key: " << key_counter;
return py::int_(key_counter++); return py::int_(key_counter++);
} }
//A simple caching mechanism is implemented
py::bool_ VerifyInputSignature(const py::list &input_signature, const py::tuple &inputs) { py::bool_ VerifyInputSignature(const py::list &input_signature, const py::tuple &inputs) {
MS_LOG(DEBUG) << "Verify args size:" << inputs.size(); MS_LOG(DEBUG) << "Verify args size:" << inputs.size();
if (inputs.size() != input_signature.size()) { if (inputs.size() != input_signature.size()) {
MS_LOG(ERROR) << "Signature size not equal to args size"; MS_LOG(ERROR) << "Signature size not equal to args size";
return false; return false;
} }
//A function VerifyInputSignature is implemented
//The log message outputs the number of input parameters
//in order to debug and troubleshoot errors
size_t count = 0; size_t count = 0;
for (auto arg_obj : inputs) { for (auto arg_obj : inputs) {
if (py::isinstance<Tensor>(arg_obj)) { if (py::isinstance<Tensor>(arg_obj)) {
@ -304,6 +340,10 @@ py::bool_ VerifyInputSignature(const py::list &input_signature, const py::tuple
MS_LOG(ERROR) << "Verify Tensor error, get ptr is null"; MS_LOG(ERROR) << "Verify Tensor error, get ptr is null";
return false; return false;
} }
//This code implements a one-by-one traversal of the input parameters
//and verify that each parameter is of type Tensor
//Any parameter that is not of a Tensor type will cause validation to fail
auto sig = input_signature[count].cast<std::shared_ptr<MetaTensor>>(); auto sig = input_signature[count].cast<std::shared_ptr<MetaTensor>>();
ShapeVector sig_shape = sig->shape(); ShapeVector sig_shape = sig->shape();
TypePtr sig_type = sig->Dtype(); TypePtr sig_type = sig->Dtype();
@ -325,7 +365,11 @@ py::bool_ VerifyInputSignature(const py::list &input_signature, const py::tuple
return true; return true;
} }
//Validation of the data type of each input Tensor is implemented one by one
//The data type is compared to the data type specified in the signature
//If the data type of any input Tensor does not match the signature
//the signature verification fails
ResourcePtr GraphExecutorPy::GetResource(const std::string &phase) { ResourcePtr GraphExecutorPy::GetResource(const std::string &phase) {
MS_LOG(DEBUG) << "Phase size:" << info_.size(); MS_LOG(DEBUG) << "Phase size:" << info_.size();
if (info_.count(phase) == 0) { if (info_.count(phase) == 0) {
@ -341,6 +385,9 @@ FuncGraphPtr GraphExecutorPy::GetFuncGraph(const std::string &phase) {
} }
return info_[phase]->func_graph; return info_[phase]->func_graph;
} }
//The corresponding resources are obtained according to the given phase
//If a given stage exists in a info_ map container
//the corresponding resource is returned
FuncGraphPtr GraphExecutorPy::GetGradGraph(const std::string &phase) { FuncGraphPtr GraphExecutorPy::GetGradGraph(const std::string &phase) {
if (phase.empty()) { if (phase.empty()) {
@ -349,13 +396,15 @@ FuncGraphPtr GraphExecutorPy::GetGradGraph(const std::string &phase) {
if (info_.count(phase) == 0) { if (info_.count(phase) == 0) {
MS_LOG(EXCEPTION) << "No phase in executor:" << phase; MS_LOG(EXCEPTION) << "No phase in executor:" << phase;
} }
auto execute_info = info_[phase]; auto execute_info = info_[phase];
MS_EXCEPTION_IF_NULL(execute_info); MS_EXCEPTION_IF_NULL(execute_info);
auto grad_graph = execute_info->grad_graph; auto grad_graph = execute_info->grad_graph;
MS_EXCEPTION_IF_NULL(grad_graph); MS_EXCEPTION_IF_NULL(grad_graph);
return grad_graph; return grad_graph;
} }
//The corresponding gradient map is obtained according to the given phase
//The corresponding gradient map can only be obtained
//if a given stage is present in the actuator
void GraphExecutorPy::SetGradGraph(const FuncGraphPtr &grad_graph, const std::string &phase) { void GraphExecutorPy::SetGradGraph(const FuncGraphPtr &grad_graph, const std::string &phase) {
if (phase.empty()) { if (phase.empty()) {
@ -383,6 +432,11 @@ compile::VmEvalFuncPtr GraphExecutorPy::GetVmEvalFunc(const std::string &phase)
MS_LOG(ERROR) << "GetVmEvalFunc vm model can't find kOutput:" << kOutput; MS_LOG(ERROR) << "GetVmEvalFunc vm model can't find kOutput:" << kOutput;
return nullptr; return nullptr;
} }
//Implements obtaining the corresponding VmEvalFunc according to a given phase
//Start by getting a resource pointer for a given stage
//Then check if the resource contains a result with the name kOutput
//And the type of the result is compile::VmEvalFuncPtr
//If the condition is met, the result is returned
bool GraphExecutorPy::HasCompiled(const std::string &phase) const { bool GraphExecutorPy::HasCompiled(const std::string &phase) const {
if (info_.count(phase) == 0) { if (info_.count(phase) == 0) {
@ -426,6 +480,11 @@ py::bytes GraphExecutorPy::GetFuncGraphProto(const std::string &phase, const std
MS_LOG(EXCEPTION) << "Unknown ir type: " << ir_type; MS_LOG(EXCEPTION) << "Unknown ir type: " << ir_type;
} }
//Implements serialization strings that obtain the corresponding function graph
//according to the given phase and IR type (ir_type).
//First get the function graph pointer corresponding to the given stage
//It is then processed differently depending on the IR type
//Finally, the corresponding function graph serialized string is returned
py::bytes GraphExecutorPy::GetOptimizeGraphProto(const std::string &phase) { py::bytes GraphExecutorPy::GetOptimizeGraphProto(const std::string &phase) {
if (info_.count(phase) == 0) { if (info_.count(phase) == 0) {
@ -442,6 +501,12 @@ py::bytes GraphExecutorPy::GetOptimizeGraphProto(const std::string &phase) {
} }
return proto_str; return proto_str;
} }
//A serialized string for the function graph that obtains the optimized graph
//according to a given phase is implemented
//First check if a given stage is present in the actuator
//Then get the function graph pointer of the optimized graph
//by calling the optimize_graph method of the resource
//Finally, the function graph is serialized to a string and returned
void GraphExecutorPy::SetJitConfig(const py::dict &jit_config) { jit_config_ = GenerateJitConfigMap(jit_config); } void GraphExecutorPy::SetJitConfig(const py::dict &jit_config) { jit_config_ = GenerateJitConfigMap(jit_config); }
@ -455,6 +520,10 @@ py::dict GraphExecutorPy::GetParallelGraphInfo(const std::string &phase) {
return mindspore::parallel::GetParallelCNodeInfoFromGraph(graph); return mindspore::parallel::GetParallelCNodeInfoFromGraph(graph);
} }
//This code implements the functions of setting the JIT configuration
//and obtaining information about the parallel graph
//The SetJitConfig method converts the incoming Python dictionary into an internal configuration map
//The GetParallelGraphInfo method gets the function graph based on the given stage name
py::dict GraphExecutorPy::GetParameterLayout(const std::string &phase) { py::dict GraphExecutorPy::GetParameterLayout(const std::string &phase) {
MS_LOG(DEBUG) << "GetParameterLayout!"; MS_LOG(DEBUG) << "GetParameterLayout!";
@ -466,6 +535,8 @@ py::dict GraphExecutorPy::GetParameterLayout(const std::string &phase) {
} }
return mindspore::parallel::GetParameterLayoutFromGraph(graph); return mindspore::parallel::GetParameterLayoutFromGraph(graph);
} }
//The function is to obtain parameter layout information based on the given phase phase
//and return it as a Python dictionary
py::dict GraphExecutorPy::GetCNodeStrategy(const std::string &phase) { py::dict GraphExecutorPy::GetCNodeStrategy(const std::string &phase) {
MS_LOG(DEBUG) << "GetCNodeStrategy!"; MS_LOG(DEBUG) << "GetCNodeStrategy!";
@ -481,11 +552,14 @@ py::list GraphExecutorPy::GetParallelParameterNameList(const std::string &phase)
} }
return mindspore::parallel::GetParallelParameterNameListFromGraph(graph); return mindspore::parallel::GetParallelParameterNameListFromGraph(graph);
} }
//The function is to get a list of parallel parameter names based on the given phase phase
void GraphExecutorPy::SetCNodeStrategy(const std::string &name, const parallel::Strategys &strategy) { void GraphExecutorPy::SetCNodeStrategy(const std::string &name, const parallel::Strategys &strategy) {
MS_LOG(DEBUG) << "SetCNodeStrategy!"; MS_LOG(DEBUG) << "SetCNodeStrategy!";
stra_dict_[phase_][py::str(name)] = strategy; stra_dict_[phase_][py::str(name)] = strategy;
} }
//The function is to store the parallel policy strategy corresponding to the given node name
//in the stra_dict_ member variable of the executor object
size_t GraphExecutorPy::GetNumOpsInfo(const std::string &phase) { size_t GraphExecutorPy::GetNumOpsInfo(const std::string &phase) {
MS_LOG(DEBUG) << "GetNumOpsInfo!"; MS_LOG(DEBUG) << "GetNumOpsInfo!";
@ -530,6 +604,9 @@ void GraphExecutorPy::DelNetRes(const py::set &id) {
} }
#endif #endif
} }
//This method is mainly used to delete specified network resources
//and reset the number of iterations after deletion
void GraphExecutorPy::DelOneNetRes(const py::handle &py_phase) { void GraphExecutorPy::DelOneNetRes(const py::handle &py_phase) {
if (!pybind11::isinstance<py::str>(py_phase)) { if (!pybind11::isinstance<py::str>(py_phase)) {
MS_LOG(ERROR) << "Expect string phase, but got " << py::str(py_phase); MS_LOG(ERROR) << "Expect string phase, but got " << py::str(py_phase);
@ -547,6 +624,8 @@ void GraphExecutorPy::DelOneNetRes(const py::handle &py_phase) {
MS_LOG(DEBUG) << "Delete phase: " << phase << ", info size: " << info_.size(); MS_LOG(DEBUG) << "Delete phase: " << phase << ", info size: " << info_.size();
} }
} }
//The DeleteSource method deletes network resources
//and related information at a specified stage and outputs some related log information.
void GraphExecutorPy::ClearRes() { void GraphExecutorPy::ClearRes() {
MS_LOG(INFO) << "Clean executor resource!"; MS_LOG(INFO) << "Clean executor resource!";
@ -638,6 +717,10 @@ std::map<std::string, std::pair<PrimitivePyAdapterPtr, std::string>> GraphExecut
return !(IsPrimitiveCNode(node, prim::kPrimConv2D) || IsPrimitiveCNode(node, prim::kPrimMatMul) || return !(IsPrimitiveCNode(node, prim::kPrimConv2D) || IsPrimitiveCNode(node, prim::kPrimMatMul) ||
IsPrimitiveCNode(node, prim::kPrimDepthwiseConv2dNative)); IsPrimitiveCNode(node, prim::kPrimDepthwiseConv2dNative));
}; };
//:kPrimConv2D:Determine whether it is a convolution node.
//kPrimMatMul:Determines whether it is a matrix multiplication node
//kPrimDepthwiseConv2dNative:Determines whether it is a deeply separable convolution node
//You can filter out nodes of the specified type
std::vector<AnfNodePtr> nodes = DeepScopedGraphSearchWithFilter(func_graph->get_return(), AlwaysInclude, filter); std::vector<AnfNodePtr> nodes = DeepScopedGraphSearchWithFilter(func_graph->get_return(), AlwaysInclude, filter);
auto is_quant_cnode = [](const AnfNodePtr &node) { auto is_quant_cnode = [](const AnfNodePtr &node) {
return IsPrimitiveCNode(node, prim::kPrimFakeQuantPerLayer) || return IsPrimitiveCNode(node, prim::kPrimFakeQuantPerLayer) ||
@ -645,6 +728,12 @@ std::map<std::string, std::pair<PrimitivePyAdapterPtr, std::string>> GraphExecut
IsPrimitiveCNode(node, prim::kPrimFakeLearnedScaleQuantPerLayer) || IsPrimitiveCNode(node, prim::kPrimFakeLearnedScaleQuantPerLayer) ||
IsPrimitiveCNode(node, prim::kPrimFakeLearnedScaleQuantPerChannel); IsPrimitiveCNode(node, prim::kPrimFakeLearnedScaleQuantPerChannel);
}; };
//Determine whether the node is a specific quantization operation type
//kPrimFakeQuantPerLayer:Quantization operation nodes for each layer
//kPrimFakeQuantPerChannel:Quantization operation nodes per channel
//kPrimFakeLearnedScaleQuantPerLayer:Learning scaling quantization operation nodes for each layer
//kPrimFakeLearnedScaleQuantPerChannel:Learning scaling quantization operation nodes per channel
const size_t root_node_size = 3; const size_t root_node_size = 3;
const size_t weight_index = 2; const size_t weight_index = 2;
for (const auto &node : nodes) { for (const auto &node : nodes) {
@ -701,6 +790,9 @@ void GraphExecutorPy::SaveCompiledGraph(const std::string &phase) {
} else { } else {
MS_LOG(DEBUG) << "Save model parallel parameter layout graph null!"; MS_LOG(DEBUG) << "Save model parallel parameter layout graph null!";
} }
//If there is no result with a key value of kStepParallelGraph in the res object
//a DEBUG level log is output,
//indicating that the model parallel parameter layout graph is empty
MS_LOG(INFO) << "End save compiled func graph!"; MS_LOG(INFO) << "End save compiled func graph!";
} }
@ -712,6 +804,8 @@ void GraphExecutorPy::GetGeBackendPolicy() const {
MS_LOG(EXCEPTION) << backend << " backend policy is not supported under ge backend!"; MS_LOG(EXCEPTION) << backend << " backend policy is not supported under ge backend!";
} }
} }
//This code is used to check if the current backend policy is GE
//and if not, throw an exception
bool IsPhaseExportAir(const std::string &phase) { bool IsPhaseExportAir(const std::string &phase) {
auto phase_to_export = "export.air"; auto phase_to_export = "export.air";
@ -740,17 +834,22 @@ std::vector<ActionItem> GetPipeline(const ResourcePtr &resource, const std::stri
ps::PSContext::instance()->is_server()) { ps::PSContext::instance()->is_server()) {
return ServerPipeline(resource); return ServerPipeline(resource);
} }
//It determines whether the server mode is federated learning mode or mixed mode
//And whether the current process is a server process
if (ps::PSContext::instance()->is_server()) { if (ps::PSContext::instance()->is_server()) {
resource->SetResult(kBackend, compile::CreateBackend()); resource->SetResult(kBackend, compile::CreateBackend());
return PServerPipeline(resource); return PServerPipeline(resource);
} }
//Determines whether the current process is a server process
if (ps::PSContext::instance()->is_scheduler()) { if (ps::PSContext::instance()->is_scheduler()) {
return PSchedulerPipeline(resource); return PSchedulerPipeline(resource);
} }
//Determines whether the current process is a scheduler process
if (distributed::cluster::ClusterContext::instance()->initialized()) { if (distributed::cluster::ClusterContext::instance()->initialized()) {
auto node = distributed::cluster::ClusterContext::instance()->node(); auto node = distributed::cluster::ClusterContext::instance()->node();
MS_EXCEPTION_IF_NULL(node); MS_EXCEPTION_IF_NULL(node);
MS_LOG(INFO) << "Cluster is initialized. This node role is " << node->role(); MS_LOG(INFO) << "Cluster is initialized. This node role is " << node->role();
//Determines whether the cluster environment is initialized
switch (node->role()) { switch (node->role()) {
case ps::core::NodeRole::SERVER: case ps::core::NodeRole::SERVER:
return PServerPipeline(resource); return PServerPipeline(resource);
@ -771,6 +870,8 @@ std::vector<ActionItem> GetPipeline(const ResourcePtr &resource, const std::stri
} }
return GePipeline(); return GePipeline();
} }
//Select different pipeline functions according to different conditions and parameters,
//and return the corresponding results
void GraphExecutorPy::InitCompileCacheInfo(const ResourcePtr &resource, const std::string &phase) { void GraphExecutorPy::InitCompileCacheInfo(const ResourcePtr &resource, const std::string &phase) {
// The compilation cache only support for training cell or ms_function currently. // The compilation cache only support for training cell or ms_function currently.
@ -965,6 +1066,9 @@ void CacheValidateFuncGraph(const ResourcePtr &resource) {
MsProfile::StatTime("SaveCacheFuncGraph", t2 - t1); MsProfile::StatTime("SaveCacheFuncGraph", t2 - t1);
#endif #endif
} }
//When the compilation cache function is enabled,
//the computational graph is cached and verified,
//and performance statistics can optionally be recorded
void CheckInterpretNodeLineInfos() { void CheckInterpretNodeLineInfos() {
auto &line_infos = InterpretNodeRecorder::GetInstance().LineInfos(); auto &line_infos = InterpretNodeRecorder::GetInstance().LineInfos();
@ -983,6 +1087,9 @@ void CheckInterpretNodeLineInfos() {
MS_LOG(INFO) << ss.str(); MS_LOG(INFO) << ss.str();
InterpretNodeRecorder::GetInstance().Clear(); InterpretNodeRecorder::GetInstance().Clear();
} }
//Check the line information for the interpretation node,
//and if there is line information,
//print out the code that runs in the JIT fallback and empty the line information for the interpreter node logger
#ifdef ENABLE_DUMP_IR #ifdef ENABLE_DUMP_IR
void RDRRecordGraph(const size_t action_index, const size_t action_size, const std::string &filename, void RDRRecordGraph(const size_t action_index, const size_t action_size, const std::string &filename,
@ -1007,6 +1114,8 @@ void RDRRecordGraph(const size_t action_index, const size_t action_size, const s
} }
} }
#endif #endif
//The function is to record the function graph in the pipeline
//using the RDR recorder with macro definition turned on
#ifdef ENABLE_DUMP_IR #ifdef ENABLE_DUMP_IR
void RecordIR(const size_t action_index, const size_t action_size, const std::string &action_name, void RecordIR(const size_t action_index, const size_t action_size, const std::string &action_name,
@ -1015,19 +1124,22 @@ void RecordIR(const size_t action_index, const size_t action_size, const std::st
*user_graph = graph; *user_graph = graph;
std::string base_name = GetBaseNameForIR(SizeToLong(action_index), action_name); std::string base_name = GetBaseNameForIR(SizeToLong(action_index), action_name);
// Generate IR file in human readable format // Generate IR file in human readable format
if (action_index == action_size - 1) { if (action_index == action_size - 1) {
DumpIR(base_name + ".ir", graph, false, kWholeStack); DumpIR(base_name + ".ir", graph, false, kWholeStack);
} else { } else {
DumpIR(base_name + ".ir", graph, false, kTopStack); DumpIR(base_name + ".ir", graph, false, kTopStack);
} }
// Generate IR file in a heavily commented format, which can also be reloaded // Generate IR file in a heavily commented format, which can also be reloaded
ExportIR(base_name + ".dat", graph); ExportIR(base_name + ".dat", graph);
// Generate IR file in dot format, which can be converted to svg file using graphviz dot command // Generate IR file in dot format, which can be converted to svg file using graphviz dot command
draw::Draw(base_name + ".dot", graph); draw::Draw(base_name + ".dot", graph);
} }
} }
#endif #endif
//This function generates IR files in different formats
//according to the current action index and action name
//to record the IR representation of the function graph if the conditions are met
#ifndef ENABLE_SECURITY #ifndef ENABLE_SECURITY
void SaveGraphForReadability(const std::string &action_name, const FuncGraphPtr graph, const ResourcePtr resource) { void SaveGraphForReadability(const std::string &action_name, const FuncGraphPtr graph, const ResourcePtr resource) {
@ -1108,7 +1220,7 @@ bool Pipeline::NeedCreateBackend() {
return std::any_of(actions_.begin(), actions_.end(), return std::any_of(actions_.begin(), actions_.end(),
[](const ActionItem &action) { return action.first == "task_emit" || action.first == "execute"; }); [](const ActionItem &action) { return action.first == "task_emit" || action.first == "execute"; });
} }
//If there is an action for some condition, the function returns a true value
void ProcessVmArgInner(const py::tuple &args, const ResourcePtr &res, VectorRef *const arg_list) { void ProcessVmArgInner(const py::tuple &args, const ResourcePtr &res, VectorRef *const arg_list) {
MS_EXCEPTION_IF_NULL(arg_list); MS_EXCEPTION_IF_NULL(arg_list);
std::size_t size = args.size(); std::size_t size = args.size();
@ -1156,7 +1268,8 @@ void ProcessVmArgInner(const py::tuple &args, const ResourcePtr &res, VectorRef
void GraphExecutorPy::ProcessVmArg(const py::tuple &args, const std::string &phase, VectorRef *const arg_list) { void GraphExecutorPy::ProcessVmArg(const py::tuple &args, const std::string &phase, VectorRef *const arg_list) {
ProcessVmArgInner(args, GetResource(phase), arg_list); ProcessVmArgInner(args, GetResource(phase), arg_list);
} }
//The ProcessVmArg function implements the processing of args
//and stores the processed results in arg_list.
#ifdef ENABLE_DEBUGGER #ifdef ENABLE_DEBUGGER
void GraphExecutorPy::TerminateDebugger() { void GraphExecutorPy::TerminateDebugger() {
if (Common::GetDebugTerminate()) { if (Common::GetDebugTerminate()) {
@ -1166,6 +1279,8 @@ void GraphExecutorPy::TerminateDebugger() {
} }
} }
#endif #endif
//With the debugger enabled, the appropriate debugger termination logic
//can be executed when the program terminates
py::object GraphExecutorPy::Run(const py::tuple &args, const py::object &phase_obj) { py::object GraphExecutorPy::Run(const py::tuple &args, const py::object &phase_obj) {
// Mindspore debugger notify main thread to exit after one step, and will not run next step // Mindspore debugger notify main thread to exit after one step, and will not run next step
@ -1218,8 +1333,8 @@ py::object GraphExecutorPy::Run(const py::tuple &args, const py::object &phase_o
if (vm_loop_flag) { if (vm_loop_flag) {
vm_loop = loop_size; vm_loop = loop_size;
} else { } else {
// Set the loop size in config if graphs nums is 1(is_loop_sin=True), then there will be a loop embrace // Set the loop size in config if graphs nums is 1(is_loop_sin=True), then there will be a loop embrace
// 'Execute(graph)' in GPUSession. // 'Execute(graph)' in GPUSession.
ConfigManager::GetInstance().set_gpu_loopsink_size(loop_size); ConfigManager::GetInstance().set_gpu_loopsink_size(loop_size);
} }
MS_LOG(INFO) << "VM loop size " << vm_loop << ", loopsink size " << vm_loop; MS_LOG(INFO) << "VM loop size " << vm_loop << ", loopsink size " << vm_loop;
@ -1233,7 +1348,7 @@ py::object GraphExecutorPy::Run(const py::tuple &args, const py::object &phase_o
} }
MS_LOG(DEBUG) << "Run end"; MS_LOG(DEBUG) << "Run end";
return ret; return ret;
} // namespace pipeline } // namespace pipeline
FuncGraphPtr GraphExecutorPy::BuildGraph(const py::dict &init_params, const std::string &phase, FuncGraphPtr GraphExecutorPy::BuildGraph(const py::dict &init_params, const std::string &phase,
const py::object &broadcast_params) const { const py::object &broadcast_params) const {
@ -1243,6 +1358,8 @@ FuncGraphPtr GraphExecutorPy::BuildGraph(const py::dict &init_params, const std:
return nullptr; return nullptr;
#endif #endif
} }
//Selectively builds different types of graphs based on the macro definition
//and returns the corresponding graph objects.
void GraphExecutorPy::UpdataParamNodeDefaultInput( void GraphExecutorPy::UpdataParamNodeDefaultInput(
const std::string &phase, const std::unordered_map<std::string, tensor::TensorPtr> &params_value) { const std::string &phase, const std::unordered_map<std::string, tensor::TensorPtr> &params_value) {
@ -1261,6 +1378,8 @@ void GraphExecutorPy::UpdataParamNodeDefaultInput(
} }
} }
} }
//You can update the default input for parameter nodes in a dynamic graph,
// providing default input when you run the graph
void GraphExecutorPy::RunInitGraph(const py::dict &init_params, const std::string &phase) const { void GraphExecutorPy::RunInitGraph(const py::dict &init_params, const std::string &phase) const {
#ifdef ENABLE_D #ifdef ENABLE_D
@ -1272,6 +1391,8 @@ void GraphExecutorPy::RunInitGraph(const py::dict &init_params, const std::strin
} }
#endif #endif
} }
//According to the situation defined by the macro and the setting of the back-end policy,
//the corresponding initialization diagram operation is executed
void GraphExecutorPy::PyExePath(const py::object &py_exe_path) { void GraphExecutorPy::PyExePath(const py::object &py_exe_path) {
if (!py::isinstance<py::str>(py_exe_path)) { if (!py::isinstance<py::str>(py_exe_path)) {
@ -1281,6 +1402,7 @@ void GraphExecutorPy::PyExePath(const py::object &py_exe_path) {
auto ms_context = MsContext::GetInstance(); auto ms_context = MsContext::GetInstance();
ms_context->set_param<std::string>(MS_CTX_PYTHON_EXE_PATH, py_exe_path_s); ms_context->set_param<std::string>(MS_CTX_PYTHON_EXE_PATH, py_exe_path_s);
} }
//You can set the global Python executable path parameter
void GraphExecutorPy::KernelBuildServerDir(const py::object &kernel_build_server_dir) { void GraphExecutorPy::KernelBuildServerDir(const py::object &kernel_build_server_dir) {
if (!py::isinstance<py::str>(kernel_build_server_dir)) { if (!py::isinstance<py::str>(kernel_build_server_dir)) {
@ -1290,6 +1412,7 @@ void GraphExecutorPy::KernelBuildServerDir(const py::object &kernel_build_server
auto ms_context = MsContext::GetInstance(); auto ms_context = MsContext::GetInstance();
ms_context->set_param<std::string>(MS_CTX_KERNEL_BUILD_SERVER_DIR, kernel_build_server_dir_s); ms_context->set_param<std::string>(MS_CTX_KERNEL_BUILD_SERVER_DIR, kernel_build_server_dir_s);
} }
//You can set the global kernel build server directory parameter
bool InitExecDataset(const std::string &queue_name, int64_t iter_num, int64_t batch_size, bool InitExecDataset(const std::string &queue_name, int64_t iter_num, int64_t batch_size,
const std::vector<TypePtr> &types, const std::vector<std::vector<int64_t>> &shapes, const std::vector<TypePtr> &types, const std::vector<std::vector<int64_t>> &shapes,
@ -1432,6 +1555,8 @@ void InitHccl() {
return; return;
} }
#endif #endif
//According to the situation defined by the macro and the setting of the backend policy,
//perform the corresponding HCCL initialization operation
mindspore::python_adapter::set_python_env_flag(true); mindspore::python_adapter::set_python_env_flag(true);
uint32_t device_id = ms_context->get_param<uint32_t>(MS_CTX_DEVICE_ID); uint32_t device_id = ms_context->get_param<uint32_t>(MS_CTX_DEVICE_ID);
@ -1486,6 +1611,8 @@ void FinalizeHccl() {
device::DeviceContextManager::GetInstance().ClearDeviceContexts(); device::DeviceContextManager::GetInstance().ClearDeviceContexts();
device::KernelRuntimeManager::Instance().ClearRuntimeResource(); device::KernelRuntimeManager::Instance().ClearRuntimeResource();
} }
//Depending on the macro definition and the settings of the backend policy,
//perform the corresponding HCCL termination operation or resource release operation
uint32_t GetHcclRankId() { uint32_t GetHcclRankId() {
uint32_t rank_id = 0; uint32_t rank_id = 0;
@ -1495,6 +1622,8 @@ uint32_t GetHcclRankId() {
} }
return rank_id; return rank_id;
} }
//You can get the rank id of the current HCCL
// and return the default value if the acquisition fails
uint32_t GetHcclRankSize() { uint32_t GetHcclRankSize() {
uint32_t rank_size = 0; uint32_t rank_size = 0;
@ -1504,6 +1633,8 @@ uint32_t GetHcclRankSize() {
} }
return rank_size; return rank_size;
} }
//You can get the rank number of the current HCCL
//and return the default value if the acquisition fails
void ExportGraph(const std::string &file_name, const std::string &, const std::string &phase) { void ExportGraph(const std::string &file_name, const std::string &, const std::string &phase) {
#ifdef ENABLE_D #ifdef ENABLE_D
@ -1512,6 +1643,7 @@ void ExportGraph(const std::string &file_name, const std::string &, const std::s
MS_EXCEPTION(ValueError) << "Only support export file in 'AIR' format with Ascend backend."; MS_EXCEPTION(ValueError) << "Only support export file in 'AIR' format with Ascend backend.";
#endif #endif
} }
//Perform the appropriate diagram export operation according to the situation defined by the macro
FuncGraphPtr LoadMindIR(const std::string &file_name, char *dec_key, const size_t key_len, FuncGraphPtr LoadMindIR(const std::string &file_name, char *dec_key, const size_t key_len,
const std::string &dec_mode) { const std::string &dec_mode) {
@ -1535,12 +1667,13 @@ void ReleaseGeTsd() {
(void)context::CloseTsd(context_ptr, true); (void)context::CloseTsd(context_ptr, true);
} }
} }
//Free up resources related to GE and TSD
void InitPipeline() { void InitPipeline() {
// set python env flag // set python env flag
RecordInitStatus(); RecordInitStatus();
mindspore::python_adapter::set_python_env_flag(true); mindspore::python_adapter::set_python_env_flag(true);
// open tsd before ge initialize // open tsd before ge initialize
auto ms_context = MsContext::GetInstance(); auto ms_context = MsContext::GetInstance();
MS_EXCEPTION_IF_NULL(ms_context); MS_EXCEPTION_IF_NULL(ms_context);
if (!context::OpenTsd(ms_context)) { if (!context::OpenTsd(ms_context)) {
@ -1555,6 +1688,8 @@ void FinalizeBackend() {
(void)context::FinalizeGe(context_ptr); (void)context::FinalizeGe(context_ptr);
(void)context::CloseTsd(context_ptr); (void)context::CloseTsd(context_ptr);
} }
//Complete the termination of the backend,
//including the termination of GE and the shutdown of TSD
void MemoryRecycle() { void MemoryRecycle() {
#ifdef ENABLE_DUMP_IR #ifdef ENABLE_DUMP_IR
@ -1568,8 +1703,8 @@ void MemoryRecycle() {
abstract::AnalysisResultCacheMgr::GetInstance().Clear(); abstract::AnalysisResultCacheMgr::GetInstance().Clear();
abstract::AnalysisContext::ClearContext(); abstract::AnalysisContext::ClearContext();
g_args_cache.clear(); g_args_cache.clear();
// clean static variable to prevent from crash. As static variable is released after // clean static variable to prevent from crash. As static variable is released after
// Python threads is released. // Python threads is released.
parse::data_converter::ClearObjectCache(); parse::data_converter::ClearObjectCache();
parse::Parser::CleanParserResource(); parse::Parser::CleanParserResource();
parse::CleanDataClassToClassMap(); parse::CleanDataClassToClassMap();
@ -1580,7 +1715,7 @@ void MemoryRecycle() {
void ClearResAtexit() { void ClearResAtexit() {
MS_LOG(INFO) << "Pipeline clear all resource"; MS_LOG(INFO) << "Pipeline clear all resource";
runtime::OpExecutor::GetInstance().WorkerJoin(); runtime::OpExecutor::GetInstance().WorkerJoin();
// When the python process exits, the kernels on the device may not have finished executing. // When the python process exits, the kernels on the device may not have finished executing.
device::KernelRuntimeManager::Instance().WaitTaskFinishOnDevice(); device::KernelRuntimeManager::Instance().WaitTaskFinishOnDevice();
device::DeviceContextManager::GetInstance().WaitTaskFinishOnDevice(); device::DeviceContextManager::GetInstance().WaitTaskFinishOnDevice();
@ -1720,6 +1855,8 @@ py::bytes PyEncrypt(char *plain_data, size_t plain_len, char *key, size_t key_le
auto py_encrypt_data = py::bytes(reinterpret_cast<char *>(encrypt_data.get()), encrypt_len); auto py_encrypt_data = py::bytes(reinterpret_cast<char *>(encrypt_data.get()), encrypt_len);
return py_encrypt_data; return py_encrypt_data;
} }
//You can implement that the specified plaintext data is encrypted with the given key
//and encryption mode, and then the encrypted data is returned as a Python byte object
py::bytes PyDecrypt(const std::string &encrypt_data_path, char *key, size_t key_len, const std::string &dec_mode) { py::bytes PyDecrypt(const std::string &encrypt_data_path, char *key, size_t key_len, const std::string &dec_mode) {
size_t decrypt_len; size_t decrypt_len;
@ -1732,6 +1869,8 @@ py::bytes PyDecrypt(const std::string &encrypt_data_path, char *key, size_t key_
auto py_decrypt_data = py::bytes(reinterpret_cast<char *>(decrypt_data.get()), decrypt_len); auto py_decrypt_data = py::bytes(reinterpret_cast<char *>(decrypt_data.get()), decrypt_len);
return py_decrypt_data; return py_decrypt_data;
} }
//You can implement the specified encrypted data to be decrypted with the given key and decryption mode,
//and then return the decrypted data as a Python byte object
bool PyIsCipherFile(const std::string &file_path) { return mindspore::IsCipherFile(file_path); } bool PyIsCipherFile(const std::string &file_path) { return mindspore::IsCipherFile(file_path); }
} // namespace pipeline } // namespace pipeline

View File

@ -58,9 +58,11 @@ void DoExecNonInputGraph(const std::string &phase) {
MS_LOG(ERROR) << "Can not found GraphRunner"; MS_LOG(ERROR) << "Can not found GraphRunner";
return; return;
} }
//It provides a basic framework for performing non-input graph calculations,
//and the specific calculation logic is implemented in subsequent code
{ {
// Release GIL before calling into (potentially long-running) C++ code // Release GIL before calling into (potentially long-running) C++ code
py::gil_scoped_release release; py::gil_scoped_release release;
Status ret = graph_runner->RunGraph(run_options, ge_tensors, &ge_outputs); Status ret = graph_runner->RunGraph(run_options, ge_tensors, &ge_outputs);
if (ret != Status::SUCCESS) { if (ret != Status::SUCCESS) {
@ -73,6 +75,7 @@ void DoExecNonInputGraph(const std::string &phase) {
void SetGeOption(const std::map<std::string, std::string> &options) { void SetGeOption(const std::map<std::string, std::string> &options) {
ConfigManager::GetInstance().set_ge_initialize_options(options); ConfigManager::GetInstance().set_ge_initialize_options(options);
} }
//We can flexibly set GE's initialization parameters to meet different needs
Status CreateSessionAndGraphRunner(bool is_training = true) { Status CreateSessionAndGraphRunner(bool is_training = true) {
std::shared_ptr<ge::Session> sess = DfGraphManager::GetInstance().GetGeSession(); std::shared_ptr<ge::Session> sess = DfGraphManager::GetInstance().GetGeSession();
@ -98,6 +101,9 @@ Status CreateSessionAndGraphRunner(bool is_training = true) {
DfGraphManager::GetInstance().SetGraphRunner(graph_runner); DfGraphManager::GetInstance().SetGraphRunner(graph_runner);
return Status::SUCCESS; return Status::SUCCESS;
} }
//The role of this code is to create session and graph runner objects,
//and configure the corresponding options,
//which provides the infrastructure for the calculation process of the model
bool InitExecDatasetGe(const std::string &queue_name, int64_t size, int64_t batch_size, bool InitExecDatasetGe(const std::string &queue_name, int64_t size, int64_t batch_size,
const std::vector<TypePtr> &types, const std::vector<std::vector<int64_t>> &shapes, const std::vector<TypePtr> &types, const std::vector<std::vector<int64_t>> &shapes,
@ -108,11 +114,15 @@ bool InitExecDatasetGe(const std::string &queue_name, int64_t size, int64_t batc
}); });
ConfigManager::GetInstance().set_dataset_mode(DatasetMode::DS_SINK_MODE); ConfigManager::GetInstance().set_dataset_mode(DatasetMode::DS_SINK_MODE);
//Set the dataset mode to Data Drop Mode
ConfigManager::GetInstance().set_iter_num(queue_name, size); ConfigManager::GetInstance().set_iter_num(queue_name, size);
//Set the number of iterations
ConfigManager::GetInstance().set_dataset_phase(phase); ConfigManager::GetInstance().set_dataset_phase(phase);
//Set up the dataset stage
DatasetGraphParam param(queue_name, size, batch_size, ge_types, shapes, input_indexes); DatasetGraphParam param(queue_name, size, batch_size, ge_types, shapes, input_indexes);
ConfigManager::GetInstance().set_dataset_param(param); ConfigManager::GetInstance().set_dataset_param(param);
//Set some specific configurations for the dataset
if (transform::BuildDatasetGraph(param, phase) != transform::SUCCESS) { if (transform::BuildDatasetGraph(param, phase) != transform::SUCCESS) {
MS_LOG(ERROR) << "Build dateset graph failed."; MS_LOG(ERROR) << "Build dateset graph failed.";
@ -169,6 +179,8 @@ void ConvertObjectToTensors(const py::dict &dict, TensorOrderMap *const tensors)
(void)tensors->emplace(name, tensor); (void)tensors->emplace(name, tensor);
} }
} }
//By processing key-value pairs in the Python dictionary one by one,
//they are converted into tensors and stored in TensorOrderMap for later use
bool AddDFGraph(const std::map<std::string, ExecutorInfoPtr> &info, const py::dict &init_params, bool AddDFGraph(const std::map<std::string, ExecutorInfoPtr> &info, const py::dict &init_params,
const std::string &phase, const py::object &broadcast_params) { const std::string &phase, const py::object &broadcast_params) {
@ -211,15 +223,15 @@ bool AddDFGraph(const std::map<std::string, ExecutorInfoPtr> &info, const py::di
} }
#ifdef ENABLE_DUMP_IR #ifdef ENABLE_DUMP_IR
if (MsContext::GetInstance()->get_param<bool>(MS_CTX_SAVE_GRAPHS_FLAG)) { if (MsContext::GetInstance()->get_param<bool>(MS_CTX_SAVE_GRAPHS_FLAG)) {
converter.DrawComputeGraph(GetSaveGraphsPathName("ge_graph.dot")); // for debug converter.DrawComputeGraph(GetSaveGraphsPathName("ge_graph.dot")); // for debug
converter.DrawInitGraph(GetSaveGraphsPathName("init_graph.dot")); // for debug converter.DrawInitGraph(GetSaveGraphsPathName("init_graph.dot")); // for debug
converter.DrawSaveCheckpointGraph(GetSaveGraphsPathName("save_checkpoint_graph.dot")); // for debug converter.DrawSaveCheckpointGraph(GetSaveGraphsPathName("save_checkpoint_graph.dot")); // for debug
} }
#endif #endif
std::string init_graph = "init_subgraph." + net_id; std::string init_graph = "init_subgraph." + net_id;
std::string checkpoint_name = "save." + net_id; std::string checkpoint_name = "save." + net_id;
if (phase.find("train") != std::string::npos) { if (phase.find("train") != std::string::npos) {
(void)DfGraphManager::GetInstance().AddGraph(phase, converter.GetComputeGraph(), {{"ge.exec.variable_acc", "1"}}); (void)DfGraphManager::GetInstance().AddGraph(phase, converter.GetComputeGraph(), {{"ge.exec.variable_acc", "1"}});//Add additional properties to the graph
} else { } else {
(void)DfGraphManager::GetInstance().AddGraph(phase, converter.GetComputeGraph()); (void)DfGraphManager::GetInstance().AddGraph(phase, converter.GetComputeGraph());
} }
@ -246,6 +258,8 @@ FuncGraphPtr BuildDFGraph(const std::map<std::string, ExecutorInfoPtr> &info, co
DumpIR("anf_graph.ir", anf_graph, true); DumpIR("anf_graph.ir", anf_graph, true);
} }
#endif #endif
//Computational graphs can be saved in the form of images
//and texts for subsequent visualization, analysis, and debugging
if (!AddDFGraph(info, init_params, phase, broadcast_params)) { if (!AddDFGraph(info, init_params, phase, broadcast_params)) {
MS_LOG(ERROR) << "GenConvertor failed"; MS_LOG(ERROR) << "GenConvertor failed";
@ -300,7 +314,7 @@ void RunGEInitGraph(const py::dict &init_params, const std::string &phase) {
MS_LOG(EXCEPTION) << "Can not found GraphRunner."; MS_LOG(EXCEPTION) << "Can not found GraphRunner.";
} }
{ {
// Release GIL before calling into (potentially long-running) C++ code // Release GIL before calling into (potentially long-running) C++ code
py::gil_scoped_release release; py::gil_scoped_release release;
Status ret = graph_runner->RunGraph(run_options, ge_tensors, &ge_outputs); Status ret = graph_runner->RunGraph(run_options, ge_tensors, &ge_outputs);
if (ret != Status::SUCCESS) { if (ret != Status::SUCCESS) {
@ -329,6 +343,9 @@ py::object ExtractGeneralCnodeRet(const AbstractBasePtr &cnode_data, const py::t
MS_LOG(EXCEPTION) << "The number of elements in the outputs : " << data.size() MS_LOG(EXCEPTION) << "The number of elements in the outputs : " << data.size()
<< " less than the number of elements required. "; << " less than the number of elements required. ";
} }
//This code is used to check whether the abstract tensor data output
//by the compute node is available and determine
//whether the amount of output data is consistent with the required quantity.
BaseShapePtr shape = cnode_data->BuildShape(); BaseShapePtr shape = cnode_data->BuildShape();
if (!shape->isa<abstract::Shape>()) { if (!shape->isa<abstract::Shape>()) {
@ -337,7 +354,7 @@ py::object ExtractGeneralCnodeRet(const AbstractBasePtr &cnode_data, const py::t
auto shape_me = shape->cast<abstract::ShapePtr>()->shape(); auto shape_me = shape->cast<abstract::ShapePtr>()->shape();
auto shape_ge = py::cast<Tensor &>(data[*count]).shape(); auto shape_ge = py::cast<Tensor &>(data[*count]).shape();
if (shape_ge != shape_me) { // dynamic shape if (shape_ge != shape_me) { // dynamic shape
MS_LOG(WARNING) << "The shape of the " << *count << "th tensor returned: " << shape_ge MS_LOG(WARNING) << "The shape of the " << *count << "th tensor returned: " << shape_ge
<< " is not the same as the shape of the tensor derived: " << shape_me; << " is not the same as the shape of the tensor derived: " << shape_me;
} }
@ -350,6 +367,7 @@ py::object ExtractGeneralCnodeRet(const AbstractBasePtr &cnode_data, const py::t
<< "only be a tensor or a tuple of tensor, but got " << cnode_data->BuildValue()->ToString() << "only be a tensor or a tuple of tensor, but got " << cnode_data->BuildValue()->ToString()
<< "."; << ".";
} }
//Used to check whether the data type of the compute node output is an abstract tuple
auto data_tp = cnode_data->cast<AbstractTuplePtr>(); auto data_tp = cnode_data->cast<AbstractTuplePtr>();
auto elements = data_tp->elements(); auto elements = data_tp->elements();
size_t size = data_tp->size(); size_t size = data_tp->size();
@ -380,6 +398,7 @@ py::object StructureOutput(const AnfNodePtr &output_node, const py::tuple &data,
MS_LOG(EXCEPTION) << "The final anf graph could only have constant, parameter, and operator, but got " MS_LOG(EXCEPTION) << "The final anf graph could only have constant, parameter, and operator, but got "
<< output_node->ToString(); << output_node->ToString();
} }
//Used to check whether the final output of the graph is constant, parameter, or operator
if (output_c->IsApply(prim::kPrimMakeTuple)) { if (output_c->IsApply(prim::kPrimMakeTuple)) {
auto input_list = output_c->inputs(); auto input_list = output_c->inputs();
@ -413,7 +432,7 @@ std::shared_ptr<py::object> DoExecGraph(const FuncGraphPtr &graph, const std::ve
} }
{ {
// Release GIL before calling into (potentially long-running) C++ code // Release GIL before calling into (potentially long-running) C++ code
py::gil_scoped_release release; py::gil_scoped_release release;
MS_LOG(DEBUG) << "Run graph begin, inputs size is: " << inputs.size(); MS_LOG(DEBUG) << "Run graph begin, inputs size is: " << inputs.size();
Status ret = graph_runner->RunGraph(run_options, ge_tensors, &ge_outputs); Status ret = graph_runner->RunGraph(run_options, ge_tensors, &ge_outputs);
@ -447,7 +466,7 @@ std::shared_ptr<py::object> DoExecGraph(const FuncGraphPtr &graph, const std::ve
void ProcessGeArg(const std::map<std::string, ExecutorInfoPtr> &info, const py::tuple &args, const std::string &phase, void ProcessGeArg(const std::map<std::string, ExecutorInfoPtr> &info, const py::tuple &args, const std::string &phase,
std::vector<tensor::TensorPtr> *inputs) { std::vector<tensor::TensorPtr> *inputs) {
// check the arg and use the GraphExecutorPy args // check the arg and use the GraphExecutorPy args
std::size_t size = args.size(); std::size_t size = args.size();
if (info.count(phase) == 0) { if (info.count(phase) == 0) {
@ -459,8 +478,8 @@ void ProcessGeArg(const std::map<std::string, ExecutorInfoPtr> &info, const py::
MS_LOG(EXCEPTION) << "The real arg num : size = " << size << ". graph_arg_size = " << arg_size; MS_LOG(EXCEPTION) << "The real arg num : size = " << size << ". graph_arg_size = " << arg_size;
} }
// process the first args of tensor // process the first args of tensor
// only in dataset normal(non-sink) mode, fp_bp graph need input tensors // only in dataset normal(non-sink) mode, fp_bp graph need input tensors
if (ConfigManager::GetInstance().dataset_mode() == DS_NORMAL_MODE) { if (ConfigManager::GetInstance().dataset_mode() == DS_NORMAL_MODE) {
for (std::size_t i = 0; i < size; i++) { for (std::size_t i = 0; i < size; i++) {
ValuePtr converted = nullptr; ValuePtr converted = nullptr;
@ -492,7 +511,7 @@ py::object ExecDFGraph(const std::map<std::string, ExecutorInfoPtr> &info, const
FuncGraphPtr anf_graph = info.at(phase)->func_graph; FuncGraphPtr anf_graph = info.at(phase)->func_graph;
std::shared_ptr<py::object> ret_val = std::make_shared<py::object>(); std::shared_ptr<py::object> ret_val = std::make_shared<py::object>();
// We will not execute graph when output is constant or just input itself. // We will not execute graph when output is constant or just input itself.
if (IsGraphOutputValueNodeOrParameter(info.at(phase)->func_graph->output(), args, ret_val)) { if (IsGraphOutputValueNodeOrParameter(info.at(phase)->func_graph->output(), args, ret_val)) {
ConfigManager::GetInstance().ResetConfig(); ConfigManager::GetInstance().ResetConfig();
return *ret_val; return *ret_val;
@ -517,6 +536,10 @@ void ExportDFGraph(const std::string &file_name, const std::string &phase) {
MS_LOG(ERROR) << "Get graph form DfGraphManager failed!"; MS_LOG(ERROR) << "Get graph form DfGraphManager failed!";
return; return;
} }
//It is mainly used to export deep learning framework diagrams to disk files
//for use by other modules or tools.
//You need to obtain the corresponding DfGraphWrapperPtr object
//through DfGraphManager and then export it through this object
transform::DfGraphPtr ge_graph = wrap_ptr->graph_ptr_; transform::DfGraphPtr ge_graph = wrap_ptr->graph_ptr_;
if (ge_graph == nullptr) { if (ge_graph == nullptr) {
@ -529,5 +552,5 @@ void ExportDFGraph(const std::string &file_name, const std::string &phase) {
} }
MS_LOG(INFO) << "Export air model finish."; MS_LOG(INFO) << "Export air model finish.";
} }
} // namespace pipeline } // namespace pipeline
} // namespace mindspore } // namespace mindspore

View File

@ -49,6 +49,9 @@ std::string GetWorldGroup() {
} }
return world_group; return world_group;
} }
//It is mainly used to obtain the communication groups used in the current running environment
// for parallel operations in scenarios such as distributed training.
//Depending on the back-end device, the corresponding communication group name is returned
static int64_t GetRank() { static int64_t GetRank() {
auto ms_context = MsContext::GetInstance(); auto ms_context = MsContext::GetInstance();
@ -64,6 +67,10 @@ static int64_t GetRank() {
} }
return global_rank; return global_rank;
} }
//It is mainly used to obtain the global ranking of the current process in distributed training.
//First, check whether the global ranking has been set, and if not,
//get the ranking of the current process in the distribution group through the CommManager.
//The rank is then converted to the correct data type and the global rank is returned
static int64_t InferStage(int64_t rank_id, int64_t stage_num, int64_t device_num) { static int64_t InferStage(int64_t rank_id, int64_t stage_num, int64_t device_num) {
if (stage_num == 0) { if (stage_num == 0) {
@ -76,6 +83,12 @@ static int64_t InferStage(int64_t rank_id, int64_t stage_num, int64_t device_num
auto per_stage_rank_num = device_num / stage_num; auto per_stage_rank_num = device_num / stage_num;
return rank_id / per_stage_rank_num; return rank_id / per_stage_rank_num;
} }
//This function is mainly used to infer the stage of the current process
// based on the given number of stages, the number of devices,
//and the ranking of the current process.
//It calculates the stage it is in by dividing the ranking of the current process
//by the number of rankings for each stage,
//thus determining the stage to which the current process belongs.
static bool HasVirtualDataset(const std::vector<AnfNodePtr> &all_nodes) { static bool HasVirtualDataset(const std::vector<AnfNodePtr> &all_nodes) {
for (auto &node : all_nodes) { for (auto &node : all_nodes) {
@ -86,7 +99,11 @@ static bool HasVirtualDataset(const std::vector<AnfNodePtr> &all_nodes) {
} }
return false; return false;
} }
//This function is primarily used to check for the presence of a virtual dataset operation in a given node list.
//It determines whether a node is a virtual dataset operation
//by iterating through each node in the list.
//Returns true if a dummy dataset operation is found; Otherwise, false is returned.
static CNodePtr CreateTupleGetItem(const AnfNodePtr &node, size_t index, const FuncGraphPtr &func_graph) { static CNodePtr CreateTupleGetItem(const AnfNodePtr &node, size_t index, const FuncGraphPtr &func_graph) {
MS_EXCEPTION_IF_NULL(node); MS_EXCEPTION_IF_NULL(node);
MS_EXCEPTION_IF_NULL(func_graph); MS_EXCEPTION_IF_NULL(func_graph);
@ -105,6 +122,10 @@ static CNodePtr CreateTupleGetItem(const AnfNodePtr &node, size_t index, const F
tuple_get_item->set_abstract(tuple_get_item_abstract); tuple_get_item->set_abstract(tuple_get_item_abstract);
return tuple_get_item; return tuple_get_item;
} }
//This function is primarily used to create an tuple_get_item operation
//that gets the element of the specified subscript from a tuple type node.
// It implements the function of fetching the specified subscript element in the tuple
//by creating a new tuple_get_item operation and setting its input node and Abstract.
static CNodePtr CreateVirtualDataset(const FuncGraphPtr &func_graph) { static CNodePtr CreateVirtualDataset(const FuncGraphPtr &func_graph) {
mindspore::parallel::OperatorAttrs attrs; mindspore::parallel::OperatorAttrs attrs;
@ -128,6 +149,11 @@ static CNodePtr CreateVirtualDataset(const FuncGraphPtr &func_graph) {
virtual_dataset_node->set_abstract(std::make_shared<abstract::AbstractTuple>(abstract_list)); virtual_dataset_node->set_abstract(std::make_shared<abstract::AbstractTuple>(abstract_list));
return virtual_dataset_node; return virtual_dataset_node;
} }
//This function is mainly used to create a virtual dataset (VirtualDataset) operation.
// It implements the function of creating a virtual dataset operation
//by creating a ValueNode node and corresponding parameter list,
//and then creating a new virtual dataset operation CNode with these parameters,
//and setting its in_forward_flag and Abstract.
static std::set<FuncGraphPtr> FindForwardGraph(const FuncGraphPtr &root, const std::vector<AnfNodePtr> &all_nodes) { static std::set<FuncGraphPtr> FindForwardGraph(const FuncGraphPtr &root, const std::vector<AnfNodePtr> &all_nodes) {
std::set<FuncGraphPtr> graph_sets; std::set<FuncGraphPtr> graph_sets;
@ -176,6 +202,7 @@ static std::set<FuncGraphPtr> FindForwardGraph(const FuncGraphPtr &root, const s
} }
return graph_sets; return graph_sets;
} }
//Used to find the forward graph associated with a given root node
static void InsertVirtualDataset(const FuncGraphPtr &root, const std::vector<AnfNodePtr> &all_nodes) { static void InsertVirtualDataset(const FuncGraphPtr &root, const std::vector<AnfNodePtr> &all_nodes) {
MS_EXCEPTION_IF_NULL(root); MS_EXCEPTION_IF_NULL(root);
@ -219,6 +246,7 @@ static void InsertVirtualDataset(const FuncGraphPtr &root, const std::vector<Anf
} }
} }
} }
//You can automatically insert a VirtualDataset node to control data parallelism
void GenerateDefaultStrategy(const ValueNodePtr &axes, const std::vector<AnfNodePtr> &nodes, const int64_t device_num, void GenerateDefaultStrategy(const ValueNodePtr &axes, const std::vector<AnfNodePtr> &nodes, const int64_t device_num,
std::vector<std::vector<int64_t>> *default_strategy) { std::vector<std::vector<int64_t>> *default_strategy) {
@ -240,6 +268,8 @@ void GenerateDefaultStrategy(const ValueNodePtr &axes, const std::vector<AnfNode
i += 1; i += 1;
} }
} }
//You can quickly generate default policies to control data parallelism
//based on device_num specified default_strategy and number of devices
bool CheckLayout(const ValueNodePtr &axes, bool *need_default_strategy, size_t *axes_size) { bool CheckLayout(const ValueNodePtr &axes, bool *need_default_strategy, size_t *axes_size) {
auto strategies = axes->value()->cast<ValueTuplePtr>()->value(); auto strategies = axes->value()->cast<ValueTuplePtr>()->value();
@ -262,22 +292,25 @@ bool CheckLayout(const ValueNodePtr &axes, bool *need_default_strategy, size_t *
} }
return true; return true;
} }
//You can check whether the layout meets your requirements
//and determine if you need a default policy
bool IsElementWiseNode(const CNodePtr &cnode) { bool IsElementWiseNode(const CNodePtr &cnode) {
auto prim = GetCNodePrimitive(cnode); auto prim = GetCNodePrimitive(cnode);
MS_EXCEPTION_IF_NULL(prim); MS_EXCEPTION_IF_NULL(prim);
return ELEMENT_WISE_NODE_.find(prim->name()) != ELEMENT_WISE_NODE_.end(); return ELEMENT_WISE_NODE_.find(prim->name()) != ELEMENT_WISE_NODE_.end();
} }
//You can determine whether a node is an element-by-element operation node
void HandleStrategyForOneHot(std::vector<ValuePtr> *strategy) { void HandleStrategyForOneHot(std::vector<ValuePtr> *strategy) {
// onehot needs to set layout for output, modify the strategy with an additional dimension // onehot needs to set layout for output, modify the strategy with an additional dimension
auto input_strategy = GetValue<std::vector<int64_t>>(strategy->at(0)); auto input_strategy = GetValue<std::vector<int64_t>>(strategy->at(0));
input_strategy.push_back(1); input_strategy.push_back(1);
strategy->at(0) = MakeValue(input_strategy); strategy->at(0) = MakeValue(input_strategy);
} }
void HandleStrategyForMatMul(std::vector<ValuePtr> *strategy, const CNodePtr &cnode) { void HandleStrategyForMatMul(std::vector<ValuePtr> *strategy, const CNodePtr &cnode) {
// handle strategy for matmul to deal with corresponding dimension // handle strategy for matmul to deal with corresponding dimension
auto left_matrix_strategy = GetValue<std::vector<int64_t>>(strategy->at(0)); auto left_matrix_strategy = GetValue<std::vector<int64_t>>(strategy->at(0));
auto right_matrix_strategy = GetValue<std::vector<int64_t>>(strategy->at(1)); auto right_matrix_strategy = GetValue<std::vector<int64_t>>(strategy->at(1));
auto index_a = left_matrix_strategy.size() - 1; auto index_a = left_matrix_strategy.size() - 1;
@ -342,6 +375,8 @@ void HandleSpecialStrategy(std::vector<ValuePtr> *strategy, const CNodePtr &cnod
HandleStrategyForElementWiseNode(strategy, cnode); HandleStrategyForElementWiseNode(strategy, cnode);
} }
} }
//You can use the corresponding data parallelism strategy
//according to the special node type
void GetInputNodes(const FuncGraphPtr &func_graph, std::vector<AnfNodePtr> *input_nodes) { void GetInputNodes(const FuncGraphPtr &func_graph, std::vector<AnfNodePtr> *input_nodes) {
auto parameters = func_graph->parameters(); auto parameters = func_graph->parameters();
@ -352,6 +387,8 @@ void GetInputNodes(const FuncGraphPtr &func_graph, std::vector<AnfNodePtr> *inpu
input_nodes->push_back(parameter); input_nodes->push_back(parameter);
} }
} }
//You can get the input nodes in the function graph
//except for the parameter nodes named "u" and "io"
void GetOutputNodes(const FuncGraphPtr &func_graph, std::vector<AnfNodePtr> *output_nodes) { void GetOutputNodes(const FuncGraphPtr &func_graph, std::vector<AnfNodePtr> *output_nodes) {
auto return_node = func_graph->get_return(); auto return_node = func_graph->get_return();
@ -368,6 +405,9 @@ void GetOutputNodes(const FuncGraphPtr &func_graph, std::vector<AnfNodePtr> *out
} }
} }
} }
//You can get the output node in the function graph,
//that is, the input node of the non-Depend child node
//or the MakeTuple node of the return node
bool CheckDeviceNum(const std::vector<std::vector<int64_t>> &strategies, const int64_t &device_num) { bool CheckDeviceNum(const std::vector<std::vector<int64_t>> &strategies, const int64_t &device_num) {
for (size_t i = 0; i < strategies.size(); ++i) { for (size_t i = 0; i < strategies.size(); ++i) {
@ -387,6 +427,8 @@ bool CheckDeviceNum(const std::vector<std::vector<int64_t>> &strategies, const i
} }
return true; return true;
} }
//It is mainly used in distributed training to check
// whether the number of devices meets the requirements of each policy
void SetOutputLayout(const FuncGraphPtr &func_graph, const AnfNodePtr &out_strategy, const int64_t &device_num) { void SetOutputLayout(const FuncGraphPtr &func_graph, const AnfNodePtr &out_strategy, const int64_t &device_num) {
auto out_strategy_tuple = out_strategy->cast<ValueNodePtr>(); auto out_strategy_tuple = out_strategy->cast<ValueNodePtr>();
@ -419,6 +461,9 @@ void SetOutputLayout(const FuncGraphPtr &func_graph, const AnfNodePtr &out_strat
<< " is not equal to out_strategy dimension: " << output_strategy[i].size() << " at index " << " is not equal to out_strategy dimension: " << output_strategy[i].size() << " at index "
<< i; << i;
} }
//It is mainly used in distributed training to check
//whether the shape dimension of the output node matches the expected policy dimension.
//If it doesn't match, it can result in incorrect or incomplete data distribution for distributed training
std::vector<ValuePtr> elements; std::vector<ValuePtr> elements;
elements.push_back(MakeValue(output_strategy[i])); elements.push_back(MakeValue(output_strategy[i]));
auto prim = GetCNodePrimitive(node); auto prim = GetCNodePrimitive(node);
@ -451,6 +496,8 @@ std::vector<ValuePtr> GetStrategyElements(const CNodePtr &cnode, const std::vect
} }
} }
return elements; return elements;
//It is mainly used to process the input information of the model
//and generate corresponding policy information according to the situation
} }
void SetInputLayout(const FuncGraphPtr &func_graph, const AnfNodePtr &in_strategy, const int64_t &device_num) { void SetInputLayout(const FuncGraphPtr &func_graph, const AnfNodePtr &in_strategy, const int64_t &device_num) {
@ -506,6 +553,7 @@ void SetInputLayout(const FuncGraphPtr &func_graph, const AnfNodePtr &in_strateg
attrs_temp[parallel::IN_STRATEGY] = strategy; attrs_temp[parallel::IN_STRATEGY] = strategy;
(void)prim->SetAttrs(attrs_temp); (void)prim->SetAttrs(attrs_temp);
} }
//It is mainly used to input policy information for some special computing node settings
} }
void SetStrategyForShard(const FuncGraphPtr &root, const std::vector<AnfNodePtr> &all_nodes, void SetStrategyForShard(const FuncGraphPtr &root, const std::vector<AnfNodePtr> &all_nodes,
@ -526,6 +574,9 @@ void SetStrategyForShard(const FuncGraphPtr &root, const std::vector<AnfNodePtr>
} }
} }
} }
//This code snippet finds the nodes of the shard operation
//and sets the layout strategy of the input and output
//by iterating through the nodes in the function graph
// Only auto_parallel and semi_auto_parallel support PipelineSplit // Only auto_parallel and semi_auto_parallel support PipelineSplit
bool PipelineSplit(const ResourcePtr &res) { bool PipelineSplit(const ResourcePtr &res) {

View File

@ -40,6 +40,12 @@ void TryToDoReplace(FuncGraphManager *const manager, const AnfNodePtr &node, Has
// Calculate hash value. // Calculate hash value.
size_t h; size_t h;
//What this code does is try to replace a node.
//It first excludes the case where the value node is a function graph,
//then gets the value object of the node and attempts to replace the node.
//Specifically, the code calculates the hash value of the node
//and makes node replacement or caching based on the hash value
auto hash_iter = hash_value->find(node); auto hash_iter = hash_value->find(node);
if (hash_iter == hash_value->end()) { if (hash_iter == hash_value->end()) {
h = hash_combine(to_check_value->hash(), (opt::AbsOf(node)->hash())); h = hash_combine(to_check_value->hash(), (opt::AbsOf(node)->hash()));
@ -47,13 +53,21 @@ void TryToDoReplace(FuncGraphManager *const manager, const AnfNodePtr &node, Has
} else { } else {
h = hash_iter->second; h = hash_iter->second;
} }
//What this code does is perform different operations
//depending on whether the node has a hash value or not.
//If the node does not have a hash,
//a new hash value is calculated and stored in a hash table;
//If the node already has a hash, it is taken out directly and stored in the variable h.
auto bucket_iter = hash_cache->find(h); auto bucket_iter = hash_cache->find(h);
if (bucket_iter == hash_cache->end()) { if (bucket_iter == hash_cache->end()) {
// Meet for the first time, add bucket. // Meet for the first time, add bucket.
(*hash_cache)[h] = {node}; (*hash_cache)[h] = {node};
return; return;
} }
//The function of this code is to find the corresponding cache bucket
//according to the hash value of the node,
//and if it is not found, create a new cache bucket and add the node to it
auto &bucket = bucket_iter->second; auto &bucket = bucket_iter->second;
// Check if need to replace node with value node already met. // Check if need to replace node with value node already met.
@ -75,6 +89,12 @@ void TryToDoReplace(FuncGraphManager *const manager, const AnfNodePtr &node, Has
return; return;
} }
} }
//The function of this code is to compare
//whether the values of two nodes are equal,
//and if they are, replace the nodes; Otherwise, do nothing.
//The specific value comparison method depends on the type of node,
// for nodes of type tensor::Tensor, call the ValueEqual() function to compare whether their values are equal,
//for other types of nodes, use the operator "==" to compare
// Meet for the first time, append node to bucket. // Meet for the first time, append node to bucket.
bucket.emplace_back(node); bucket.emplace_back(node);

View File

@ -285,7 +285,8 @@ Resource::Resource(const py::object &obj)
: engine_(std::make_shared<abstract::AnalysisEngine>(abstract::GetPrimEvaluatorConstructors(), manager_)), : engine_(std::make_shared<abstract::AnalysisEngine>(abstract::GetPrimEvaluatorConstructors(), manager_)),
source_input_(obj), source_input_(obj),
is_cleaned_(false) {} is_cleaned_(false) {}
// The constructor initializes several member variables of the class
//assign initial values to those member variables
Resource::~Resource() { Resource::~Resource() {
MS_LOG(DEBUG) << "Resource clear"; MS_LOG(DEBUG) << "Resource clear";
@ -325,6 +326,9 @@ Any GetMethodOrAttr(const string &name, const TypeId &type_id, const BuiltInType
} }
return method->second; return method->second;
} }
// This function is used to get a method or property
//of the specified name and type from the method_map
//If no matching method or property is found, an empty Any object is returned.
bool Resource::IsTypeInBuiltInMap(const TypeId &type) { bool Resource::IsTypeInBuiltInMap(const TypeId &type) {
TypeId type_id = NormalizeTypeId(type); TypeId type_id = NormalizeTypeId(type);
@ -339,12 +343,18 @@ bool Resource::IsTypeInBuiltInMap(const TypeId &type) {
} }
return true; return true;
} }
// This function is used to determine whether a given type exists
//in the mapping table of built-in methods and properties
//Returns true if present, false otherwise
Any Resource::GetMethodPtr(const TypeId &type, const std::string &name) { Any Resource::GetMethodPtr(const TypeId &type, const std::string &name) {
TypeId type_id = NormalizeTypeId(type); TypeId type_id = NormalizeTypeId(type);
const BuiltInTypeMap &method_map = GetMethodMap(); const BuiltInTypeMap &method_map = GetMethodMap();
return GetMethodOrAttr(name, type_id, method_map); return GetMethodOrAttr(name, type_id, method_map);
} }
//This function is used to get the method pointer of the specified name and type
//from the built-in method mapping table and return it with the Any type.
//If no matching method is found, an empty Any object is returned
Any Resource::GetAttrPtr(const TypeId &type, const std::string &name) { Any Resource::GetAttrPtr(const TypeId &type, const std::string &name) {
TypeId type_id = NormalizeTypeId(type); TypeId type_id = NormalizeTypeId(type);
@ -371,6 +381,11 @@ void Resource::GetCompileCacheResource(const py::list &compile_cache_dep_files,
func_graph_ = compile_cache_manager_->GetCachedFuncGraph(manager_, weights, queue_name); func_graph_ = compile_cache_manager_->GetCachedFuncGraph(manager_, weights, queue_name);
layout_map_ = compile_cache_manager_->layout_map(); layout_map_ = compile_cache_manager_->layout_map();
} }
//This function is used to initialize and fetch compiled cache resources
//It creating a CompileCacheManager object
//Initialize and save the parallel checkpoint file
//Check the hash consistency of dependent files
//Get cached function graphs and layout maps
void Resource::CacheFuncGraph() const { void Resource::CacheFuncGraph() const {
FuncGraphPtr layout_fg = nullptr; FuncGraphPtr layout_fg = nullptr;
@ -381,7 +396,10 @@ void Resource::CacheFuncGraph() const {
} }
compile_cache_manager_->CacheFuncGraph(func_graph_, layout_fg); compile_cache_manager_->CacheFuncGraph(func_graph_, layout_fg);
} }
//This function is used to cache the compiled function graph
//It determines whether the current function graph has automatic parallelism enabled
//If so, further obtain a step-by-step parallel function graph
//Cache function diagrams and layout diagrams
void Resource::Clean() { void Resource::Clean() {
// AbstractTensor->elements() will be saved in AbstractBasePtrList // AbstractTensor->elements() will be saved in AbstractBasePtrList
args_spec_.clear(); args_spec_.clear();

View File

@ -25,6 +25,8 @@ namespace mindspore {
namespace abstract { namespace abstract {
thread_local std::string AnalysisSchedule::thread_id_ = "m"; thread_local std::string AnalysisSchedule::thread_id_ = "m";
// The role of this code in the MindSpore project is to control the scheduling and execution of threads,
// and realize the dynamic management and control of threads by constantly checking the conditions and executing the corresponding actions through the loop.
void AnalysisSchedule::Schedule() { void AnalysisSchedule::Schedule() {
const auto checkPeriod = std::chrono::seconds(3); const auto checkPeriod = std::chrono::seconds(3);
while (run_ || infer_thread_count_.load() > 0) { while (run_ || infer_thread_count_.load() > 0) {
@ -38,6 +40,7 @@ void AnalysisSchedule::Schedule() {
MS_LOG(DEBUG) << "Success to exit."; MS_LOG(DEBUG) << "Success to exit.";
} }
// The thread that performs an asynchronous task frees up CPU resources so that other threads can continue executing.
void AnalysisSchedule::Yield(const AsyncInferTask *async_infer_task) { void AnalysisSchedule::Yield(const AsyncInferTask *async_infer_task) {
MS_EXCEPTION_IF_NULL(async_infer_task); MS_EXCEPTION_IF_NULL(async_infer_task);
{ {
@ -51,6 +54,12 @@ void AnalysisSchedule::Yield(const AsyncInferTask *async_infer_task) {
activate_thread_cv_.notify_one(); activate_thread_cv_.notify_one();
} }
// Analyze a member function of the scheduling class.
// Its role is to handle anomalies that occur during analysis.
// Specifically, it logs the first exception and, if the incoming exception is a Python exception, gets the exception stack and logs it;
// Then release all locks so that other threads can continue running;
// Clear the list of ongoing tasks. Finally, the global raw evaluation cache is cleared to avoid the cache containing invalid results.
void AnalysisSchedule::HandleException(const std::exception &ex) { void AnalysisSchedule::HandleException(const std::exception &ex) {
// Just record the first exception information. // Just record the first exception information.
if (!StaticAnalysisException::Instance().HasException()) { if (!StaticAnalysisException::Instance().HasException()) {
@ -84,12 +93,16 @@ void AnalysisSchedule::HandleException(const std::exception &ex) {
} }
} }
// Stop the analysis task in progress. It stops a task by creating an asynchronous inference task in a stopped state and adding it to the scheduler.
void AnalysisSchedule::Stop() { void AnalysisSchedule::Stop() {
AsyncInferTaskPtr stop_task = AsyncInferTask::MakeShared(std::make_shared<AsyncAbstract>(), kStateStop); AsyncInferTaskPtr stop_task = AsyncInferTask::MakeShared(std::make_shared<AsyncAbstract>(), kStateStop);
Add2Schedule(stop_task); Add2Schedule(stop_task);
MS_LOG(DEBUG) << "Set analysis schedule to stop"; MS_LOG(DEBUG) << "Set analysis schedule to stop";
} }
// Wait for the analysis task to complete.
// It waits for the task by waiting for the condition variable and checking the number of threads,
// and outputs the relevant information and checks the exception after the task is completed.
void AnalysisSchedule::Wait() { void AnalysisSchedule::Wait() {
EnterWaiting(); EnterWaiting();
if (infer_thread_count_.load() > 0) { if (infer_thread_count_.load() > 0) {
@ -104,6 +117,7 @@ void AnalysisSchedule::Wait() {
StaticAnalysisException::Instance().CheckException(); StaticAnalysisException::Instance().CheckException();
} }
// Adds asynchronous inference tasks to the scheduling list and updates related statistics.
void AnalysisSchedule::Add2Schedule(const AsyncInferTaskPtr &async_infer_task_ptr) { void AnalysisSchedule::Add2Schedule(const AsyncInferTaskPtr &async_infer_task_ptr) {
std::lock_guard<std::mutex> lock(activate_thread_lock_); std::lock_guard<std::mutex> lock(activate_thread_lock_);
MS_EXCEPTION_IF_NULL(async_infer_task_ptr); MS_EXCEPTION_IF_NULL(async_infer_task_ptr);
@ -115,6 +129,9 @@ void AnalysisSchedule::Add2Schedule(const AsyncInferTaskPtr &async_infer_task_pt
<< " schedule list size: " << schedule_list_.size(); << " schedule list size: " << schedule_list_.size();
} }
// Set up the next executable analysis task.
// It determines whether to continue waiting or trigger an infinite loop exception by judging the status of the task and the number of threads in the thread pool,
// and marks the task as ready when it finds a result.
void AnalysisSchedule::SetNextReady() { void AnalysisSchedule::SetNextReady() {
if (schedule_list_.empty()) { if (schedule_list_.empty()) {
return; return;
@ -154,6 +171,8 @@ void AnalysisSchedule::SetNextReady() {
<< " address: " << async_task.get(); << " address: " << async_task.get();
} }
// Gets the result of an asynchronous task.
// It determines whether to wait and schedule by judging whether the result is a null pointer, and outputs relevant information after obtaining the result.
AbstractBasePtr AsyncAbstract::GetResult() { AbstractBasePtr AsyncAbstract::GetResult() {
auto ret = TryGetResult(); auto ret = TryGetResult();
if (ret != nullptr) { if (ret != nullptr) {
@ -195,6 +214,8 @@ AbstractFunctionPtr GetAbstractFuncRecursively(const AbstractBasePtr &abs, const
} }
} // namespace } // namespace
// Gets a unique asynchronous abstract function pointer,
// returned directly if it has already been parsed, otherwise retrieved and parsed by a recursive call.
AbstractFunctionPtr AsyncAbstractFuncAtom::GetUnique() { AbstractFunctionPtr AsyncAbstractFuncAtom::GetUnique() {
if (resolved_ != nullptr) { if (resolved_ != nullptr) {
return resolved_; return resolved_;
@ -208,6 +229,9 @@ AbstractFunctionPtr AsyncAbstractFuncAtom::GetUnique() {
return resolved_; return resolved_;
} }
// Converts the AsyncAbstractFuncAtom object to a string representation.
// It determines the content of the returned string by determining whether the member variable resolved_ is a null pointer,
// and calls resolved_'s ToString() method to get more information if needed.
std::string AsyncAbstractFuncAtom::ToString() const { std::string AsyncAbstractFuncAtom::ToString() const {
if (resolved_ == nullptr) { if (resolved_ == nullptr) {
return "AsyncAbstractFuncAtom(Not Resolved)"; return "AsyncAbstractFuncAtom(Not Resolved)";
@ -221,6 +245,7 @@ std::string AsyncAbstractFuncAtom::ToString() const {
return buffer.str(); return buffer.str();
} }
// Clear the cache of analysis results, including the original evaluation cache and three different types of cache objects.
void AnalysisResultCacheMgr::Clear() { void AnalysisResultCacheMgr::Clear() {
prim_eval_cache_->Clear(); prim_eval_cache_->Clear();
std::lock_guard<std::mutex> lock(lock_); std::lock_guard<std::mutex> lock(lock_);
@ -229,6 +254,7 @@ void AnalysisResultCacheMgr::Clear() {
switch_cache_for_check_.clear(); switch_cache_for_check_.clear();
} }
// Initializes the switch value by fetching or creating a new asynchronous abstract result object from the cache.
void AnalysisResultCacheMgr::InitSwitchValue(const AnfNodeConfigPtr &conf) { void AnalysisResultCacheMgr::InitSwitchValue(const AnfNodeConfigPtr &conf) {
std::lock_guard<std::mutex> lock(lock_); std::lock_guard<std::mutex> lock(lock_);
AsyncAbstractPtr async_eval_result = switch_cache_.get(conf); AsyncAbstractPtr async_eval_result = switch_cache_.get(conf);
@ -238,6 +264,7 @@ void AnalysisResultCacheMgr::InitSwitchValue(const AnfNodeConfigPtr &conf) {
} }
} }
// According to the given configuration information, the corresponding switch value is obtained from the analysis result cache.
AbstractBasePtr AnalysisResultCacheMgr::GetSwitchValue(const AnfNodeConfigPtr &conf) { AbstractBasePtr AnalysisResultCacheMgr::GetSwitchValue(const AnfNodeConfigPtr &conf) {
// don't call lock_.lock(). switch_cache is protected. and it waits for result. // don't call lock_.lock(). switch_cache is protected. and it waits for result.
AsyncAbstractPtr async_eval_result = switch_cache_.get(conf); AsyncAbstractPtr async_eval_result = switch_cache_.get(conf);
@ -247,6 +274,7 @@ AbstractBasePtr AnalysisResultCacheMgr::GetSwitchValue(const AnfNodeConfigPtr &c
return async_eval_result->GetResult(); return async_eval_result->GetResult();
} }
// Cache the analysis results and update the asynchronous abstract result objects in the cache by merging the current abstract result with the previous abstract result.
void AnalysisResultCacheMgr::SetCacheValue(const AnfNodeConfigPtr &conf, const AbstractBasePtr &current_abs, void AnalysisResultCacheMgr::SetCacheValue(const AnfNodeConfigPtr &conf, const AbstractBasePtr &current_abs,
AnalysisConfigAsyncResultCache *cache) { AnalysisConfigAsyncResultCache *cache) {
MS_EXCEPTION_IF_NULL(conf); MS_EXCEPTION_IF_NULL(conf);
@ -277,6 +305,7 @@ void AnalysisResultCacheMgr::SetCacheValue(const AnfNodeConfigPtr &conf, const A
} }
} }
// Set and check the cache of switch values in the analysis results cache manager.
void AnalysisResultCacheMgr::CheckSwitchValueJoinable(const AnfNodeConfigPtr &conf, const AbstractBasePtr &arg) { void AnalysisResultCacheMgr::CheckSwitchValueJoinable(const AnfNodeConfigPtr &conf, const AbstractBasePtr &arg) {
SetCacheValue(conf, arg, &switch_cache_for_check_); SetCacheValue(conf, arg, &switch_cache_for_check_);
} }

View File

@ -30,6 +30,7 @@
namespace mindspore { namespace mindspore {
namespace abstract { namespace abstract {
namespace { namespace {
// Record the run logs of the evaluator, including the evaluator name, scope name, and information about the abstract base pointer.
string EvalEntryLogging(const EvaluatorPtr &evaluator, const AbstractBasePtrList &arg_spec_list, string EvalEntryLogging(const EvaluatorPtr &evaluator, const AbstractBasePtrList &arg_spec_list,
const AnfNodeConfigPtr &out_conf) { const AnfNodeConfigPtr &out_conf) {
MS_EXCEPTION_IF_NULL(evaluator); MS_EXCEPTION_IF_NULL(evaluator);
@ -44,6 +45,9 @@ string EvalEntryLogging(const EvaluatorPtr &evaluator, const AbstractBasePtrList
return ss.str(); return ss.str();
} }
// Check whether the evaluator and output configuration are empty,
// get the node and determine the node type, and then output the appropriate error log based on the node type,
// including the evaluator name, node full name, or debugging information.
void EvalFailLogging(const EvaluatorPtr &evaluator, const AbstractBasePtrList &, const AnfNodeConfigPtr &out_conf) { void EvalFailLogging(const EvaluatorPtr &evaluator, const AbstractBasePtrList &, const AnfNodeConfigPtr &out_conf) {
MS_EXCEPTION_IF_NULL(evaluator); MS_EXCEPTION_IF_NULL(evaluator);
if (out_conf != nullptr) { if (out_conf != nullptr) {
@ -59,6 +63,7 @@ void EvalFailLogging(const EvaluatorPtr &evaluator, const AbstractBasePtrList &,
} }
} // namespace } // namespace
// Check whether a given parameter is always evaluated, based on the results of previous analysis and the value of the current parameter.
bool CheckIfAlwaysEval(const AnfNodeConfigPtr &conf, const AbstractBasePtr &arg) { bool CheckIfAlwaysEval(const AnfNodeConfigPtr &conf, const AbstractBasePtr &arg) {
auto new_sequence = dyn_cast<AbstractSequence>(arg); auto new_sequence = dyn_cast<AbstractSequence>(arg);
if (new_sequence != nullptr && new_sequence->sequence_nodes() != nullptr && new_sequence->size() != 0) { if (new_sequence != nullptr && new_sequence->sequence_nodes() != nullptr && new_sequence->size() != 0) {
@ -78,6 +83,12 @@ bool CheckIfAlwaysEval(const AnfNodeConfigPtr &conf, const AbstractBasePtr &arg)
return false; return false;
} }
// Checks if the argument passed in is empty and throws an exception if it is.
// Enter the new func graph. Gets the current node and the current context, and creates the call configuration.
// Create a new evaluator and get a new context.
// Log new context and call configuration entry events.Increase and check function call depth and stack frame depth.
// If the depth of a function call exceeds the maximum depth limit, output an exception log with methods for adjusting the maximum depth of calls and suggestions on how to avoid stack overflows.
// Output a debug log, showing the evaluator type, name, and depth of incoming function calls and stack frame depth information.
void BaseFuncGraphEvaluator::EnterStackFrame(const AnalysisEnginePtr &engine, const StackFramePtr &current_stack_frame, void BaseFuncGraphEvaluator::EnterStackFrame(const AnalysisEnginePtr &engine, const StackFramePtr &current_stack_frame,
const StackFramePtr &new_stack_frame) { const StackFramePtr &new_stack_frame) {
MS_EXCEPTION_IF_NULL(current_stack_frame); MS_EXCEPTION_IF_NULL(current_stack_frame);
@ -111,6 +122,7 @@ void BaseFuncGraphEvaluator::EnterStackFrame(const AnalysisEnginePtr &engine, co
<< "), enter, function call depth: " << FunctionCallDepth() << " - " << StackFrameDepth(); << "), enter, function call depth: " << FunctionCallDepth() << " - " << StackFrameDepth();
} }
// Leave the current function call stack frame and perform the associated operations and records.
void BaseFuncGraphEvaluator::LeaveStackFrame(const AnalysisEnginePtr &, const StackFramePtr &current_stack_frame) { void BaseFuncGraphEvaluator::LeaveStackFrame(const AnalysisEnginePtr &, const StackFramePtr &current_stack_frame) {
MS_EXCEPTION_IF_NULL(current_stack_frame); MS_EXCEPTION_IF_NULL(current_stack_frame);
// Leave current func graph. // Leave current func graph.
@ -174,6 +186,7 @@ AbstractBasePtr BaseFuncGraphEvaluator::LaunchStackFrame(const AnalysisEnginePtr
return res_base; return res_base;
} }
// Recursively executes the function graph and returns the result
AbstractBasePtr BaseFuncGraphEvaluator::LaunchRecursiveEval(const AnalysisEnginePtr &engine, const FuncGraphPtr &fg, AbstractBasePtr BaseFuncGraphEvaluator::LaunchRecursiveEval(const AnalysisEnginePtr &engine, const FuncGraphPtr &fg,
const AnalysisContextPtr &context) { const AnalysisContextPtr &context) {
MS_EXCEPTION_IF_NULL(fg); MS_EXCEPTION_IF_NULL(fg);
@ -207,6 +220,12 @@ AbstractBasePtr BaseFuncGraphEvaluator::LaunchRecursiveEval(const AnalysisEngine
return res_base; return res_base;
} }
// Checks if the argument passed in is empty and throws an exception if it is.
// Enter the new func graph. Gets the current node and the current context, and creates the call configuration.
// Create a new evaluator and get a new context.
// Log new context and call configuration entry events.Increase and check function call depth and stack frame depth.
// If the depth of a function call exceeds the maximum depth limit, output an exception log with methods for adjusting the maximum depth of calls and suggestions on how to avoid stack overflows.
// Output a debug log, showing the evaluator type, name, and depth of incoming function calls and stack frame depth information.
EvalResultPtr BaseFuncGraphEvaluator::Eval(AnalysisEnginePtr engine, const AbstractBasePtrList &args_abs_list, EvalResultPtr BaseFuncGraphEvaluator::Eval(AnalysisEnginePtr engine, const AbstractBasePtrList &args_abs_list,
const AnfNodeConfigPtr &out_conf) { const AnfNodeConfigPtr &out_conf) {
auto eval_result = evaluator_cache_mgr_->GetValue(args_abs_list); auto eval_result = evaluator_cache_mgr_->GetValue(args_abs_list);
@ -301,6 +320,7 @@ EvalResultPtr BaseFuncGraphEvaluator::Eval(AnalysisEnginePtr engine, const Abstr
return res; return res;
} }
// Each parameter in the input parameter list is extended and the extended parameter list is stored at the location pointed by broaded_args.
void BroadenArgs(const AbstractBasePtrList &args_spec_list, AbstractBasePtrList *broaded_args) { void BroadenArgs(const AbstractBasePtrList &args_spec_list, AbstractBasePtrList *broaded_args) {
MS_EXCEPTION_IF_NULL(broaded_args); MS_EXCEPTION_IF_NULL(broaded_args);
(void)std::transform(args_spec_list.begin(), args_spec_list.end(), std::back_inserter(*broaded_args), (void)std::transform(args_spec_list.begin(), args_spec_list.end(), std::back_inserter(*broaded_args),
@ -313,6 +333,7 @@ void BroadenArgs(const AbstractBasePtrList &args_spec_list, AbstractBasePtrList
}); });
} }
// The input parameter list is extended or not extended depending on whether the function graph has a flag for ignoring values.
AbstractBasePtrList FuncGraphEvaluator::NormalizeArgs(const AbstractBasePtrList &args_spec_list) const { AbstractBasePtrList FuncGraphEvaluator::NormalizeArgs(const AbstractBasePtrList &args_spec_list) const {
MS_EXCEPTION_IF_NULL(func_graph_); MS_EXCEPTION_IF_NULL(func_graph_);
if (func_graph_->has_flag(FUNC_GRAPH_FLAG_IGNORE_VALUE)) { if (func_graph_->has_flag(FUNC_GRAPH_FLAG_IGNORE_VALUE)) {
@ -325,6 +346,11 @@ AbstractBasePtrList FuncGraphEvaluator::NormalizeArgs(const AbstractBasePtrList
return args_spec_list; return args_spec_list;
} }
// Checks if the argument passed in is empty and throws an exception if it is.
// If the func graph has an ignore value flag, the parameter specification list (args_spec_list) is returned directly.
// If the function graph has an undetermined flag, set the ignore value flag to true, normalize the list of parameter specifications,
// and output a debug log. Finally, the normalized parameter specification list is returned.
// If the function graph has neither ignored value flags nor undefined flags, the list of parameter specifications is directly returned.
AbstractBasePtrList FuncGraphEvaluator::BroadenUndeterminedArgs(const AbstractBasePtrList &args_spec_list) { AbstractBasePtrList FuncGraphEvaluator::BroadenUndeterminedArgs(const AbstractBasePtrList &args_spec_list) {
MS_EXCEPTION_IF_NULL(func_graph_); MS_EXCEPTION_IF_NULL(func_graph_);
if (func_graph_->has_flag(FUNC_GRAPH_FLAG_IGNORE_VALUE)) { if (func_graph_->has_flag(FUNC_GRAPH_FLAG_IGNORE_VALUE)) {
@ -341,6 +367,8 @@ AbstractBasePtrList FuncGraphEvaluator::BroadenUndeterminedArgs(const AbstractBa
return args_spec_list; return args_spec_list;
} }
// The corresponding function graph object is obtained from the input parameter list.
// If it does not exist in the cache, a new function graph object is generated and added to the cache.
FuncGraphPtr FuncGraphEvaluator::GetFuncGraph(AnalysisEnginePtr engine, const AbstractBasePtrList &args_spec_list) { FuncGraphPtr FuncGraphEvaluator::GetFuncGraph(AnalysisEnginePtr engine, const AbstractBasePtrList &args_spec_list) {
auto iter = func_graph_cache_.find(args_spec_list); auto iter = func_graph_cache_.find(args_spec_list);
FuncGraphPtr res; FuncGraphPtr res;
@ -369,6 +397,13 @@ FuncGraphPtr FuncGraphEvaluator::GetFuncGraph(AnalysisEnginePtr engine, const Ab
return res; return res;
} }
// First, check the cache (func_graph_cache_) to see if a function graph object corresponding to the parameter specification list already exists, and return it directly if it does.
// If no corresponding function graph object exists in the cache, a new function graph object is generated based on whether the bound_node() pointer of the current object is empty.
// If bound_node() is not empty, then meta_func_graph_ and bound_node()->debug_info() are used to generate a new function graph object;
// Otherwise, a new function graph object is also generated using meta_func_graph_ and bound_node()->debug_info().
// Create a new clone function graph object (cloned_func_graph) and add it to the cache (func_graph_cache_).
// Add the newly generated function graph object to the engine's function graph manager.
// Finally, the newly generated function graph object is returned.
FuncGraphPtr MetaFuncGraphEvaluator::GetFuncGraph(AnalysisEnginePtr engine, const AbstractBasePtrList &args_spec_list) { FuncGraphPtr MetaFuncGraphEvaluator::GetFuncGraph(AnalysisEnginePtr engine, const AbstractBasePtrList &args_spec_list) {
auto iter = func_graph_cache_.find(args_spec_list); auto iter = func_graph_cache_.find(args_spec_list);
if (iter != func_graph_cache_.end()) { if (iter != func_graph_cache_.end()) {
@ -396,6 +431,17 @@ FuncGraphPtr MetaFuncGraphEvaluator::GetFuncGraph(AnalysisEnginePtr engine, cons
return cloned_func_graph; return cloned_func_graph;
} }
// The function takes three arguments:
// engine for the analysis engine object,
// args_conf_list for the function parameter configuration list to run,
// out_conf for the output node configuration object.
// converts each configuration object in args_conf_list into a corresponding evaluation result object
// and stores them in args_spec_list. It then normalizes args_spec_list and extends the undefined parameters.
// Next, the function attempts to retrieve the evaluation result object corresponding to args_spec_list from the cache.
// If it does not exist in the cache, the corresponding evaluation function is called to evaluate and the result is stored in the cache.
// If it exists in the cache, the evaluation result object in the cache is returned directly.
// The function also determines whether to update the information of the input sequence node based on the value of the environment variable MS_DEV_ENABLE_DDE before returning the result object.
// If this option is enabled, usage flags for the nodes of the old sequence and the new sequence are recursively synchronized.
EvalResultPtr Evaluator::Run(AnalysisEnginePtr engine, const ConfigPtrList &args_conf_list, EvalResultPtr Evaluator::Run(AnalysisEnginePtr engine, const ConfigPtrList &args_conf_list,
const AnfNodeConfigPtr &out_conf) { const AnfNodeConfigPtr &out_conf) {
AbstractBasePtrList args_spec_list; AbstractBasePtrList args_spec_list;
@ -450,6 +496,10 @@ EvalResultPtr Evaluator::Run(AnalysisEnginePtr engine, const ConfigPtrList &args
return eval_result; return eval_result;
} }
// determine whether the current Evaluator is a Python Prim Evaluator based on the identifier passed in (identifier_), and if so set is_py_eval to true
// Convert the parameter configuration list (args_conf_list) into an abstract base pointer list (args_spec_list) and process each element in it.
// If the current Evaluator is a Python Prim Evaluator and the parameter configuration object is an AbstractRef type, convert it to an AbstractRefPtr type and extend its ref_key.
// Call EvalPrim function, pass the engine, abstract base pointer list (args_spec_list) and other parameters, and return the evaluation result (EvalResultPtr).
EvalResultPtr TrivialPrimEvaluator::Run(AnalysisEnginePtr engine, const ConfigPtrList &args_conf_list, EvalResultPtr TrivialPrimEvaluator::Run(AnalysisEnginePtr engine, const ConfigPtrList &args_conf_list,
const AnfNodeConfigPtr &) { const AnfNodeConfigPtr &) {
AbstractBasePtrList args_spec_list; AbstractBasePtrList args_spec_list;
@ -469,6 +519,9 @@ EvalResultPtr TrivialPrimEvaluator::Run(AnalysisEnginePtr engine, const ConfigPt
return EvalPrim(engine, args_spec_list); return EvalPrim(engine, args_spec_list);
} }
// Checks if args_conf_list is empty, and throws an exception if it is empty and the identifiers are not "MakeTupleEvaluator", "MakeListEvaluator", or" RaiseEvaluator".
// Convert each configuration object in args_conf_list to the corresponding evaluation result object and store them in args_spec_list.
// The EvalPrim() function is called for in-place conversion evaluation and the result is stored in res. Finally, it returns res as the result. Note that because caching is not required, the cache manager is not used.
EvalResultPtr TransitionPrimEvaluator::Run(AnalysisEnginePtr engine, const ConfigPtrList &args_conf_list, EvalResultPtr TransitionPrimEvaluator::Run(AnalysisEnginePtr engine, const ConfigPtrList &args_conf_list,
const AnfNodeConfigPtr &out_conf) { const AnfNodeConfigPtr &out_conf) {
if (args_conf_list.empty() && identifier_ != "MakeTupleEvaluator" && identifier_ != "MakeListEvaluator" && if (args_conf_list.empty() && identifier_ != "MakeTupleEvaluator" && identifier_ != "MakeListEvaluator" &&
@ -486,6 +539,8 @@ EvalResultPtr TransitionPrimEvaluator::Run(AnalysisEnginePtr engine, const Confi
return res; return res;
} }
// Their main function is to run a Prim algorithm by configuring the list (args_conf_list) and identifier_ (identifier_) based on the parameters passed in,
// and return the evaluation result (EvalResultPtr).
EvalResultPtr SymbolicPrimEvaluator::Run(AnalysisEnginePtr, const ConfigPtrList &args_conf_list, EvalResultPtr SymbolicPrimEvaluator::Run(AnalysisEnginePtr, const ConfigPtrList &args_conf_list,
const AnfNodeConfigPtr &) { const AnfNodeConfigPtr &) {
return EvalPrim(args_conf_list); return EvalPrim(args_conf_list);
@ -506,6 +561,12 @@ EvalResultPtr TrackedEvaluator::Run(AnalysisEnginePtr engine, const ConfigPtrLis
return res; return res;
} }
// engine represents the analysis engine object, args_conf_list represents the function parameter configuration list to run, and out_conf represents the output node configuration object
// Convert each configuration object in args_conf_list to the corresponding evaluation result object and store them in args_spec_list.
// Checks if the cache manager contains the evaluation result in evaluator_cache_mgr_ and returns it directly if it does.
// Otherwise, it merges some of the application arguments and the remaining arguments into a new parameter configuration list, partial_args_conf_list,
// and calls the evaluator evaluator_ to evaluate.
// The result of the evaluation is stored in the cache manager and the result is returned.
EvalResultPtr PartialAppEvaluator::Run(AnalysisEnginePtr engine, const ConfigPtrList &args_conf_list, EvalResultPtr PartialAppEvaluator::Run(AnalysisEnginePtr engine, const ConfigPtrList &args_conf_list,
const AnfNodeConfigPtr &out_conf) { const AnfNodeConfigPtr &out_conf) {
AbstractBasePtrList args_spec_list; AbstractBasePtrList args_spec_list;
@ -532,6 +593,7 @@ EvalResultPtr PartialAppEvaluator::Run(AnalysisEnginePtr engine, const ConfigPtr
return res; return res;
} }
// Run a Prim algorithm and return an EvalResultPtr by configuring the args_conf_list and engine based on the parameters passed in.
EvalResultPtr JEvaluator::Run(AnalysisEnginePtr engine, const ConfigPtrList &args_conf_list, const AnfNodeConfigPtr &) { EvalResultPtr JEvaluator::Run(AnalysisEnginePtr engine, const ConfigPtrList &args_conf_list, const AnfNodeConfigPtr &) {
AbstractBasePtrList args_spec_list; AbstractBasePtrList args_spec_list;
(void)std::transform(args_conf_list.begin(), args_conf_list.end(), std::back_inserter(args_spec_list), (void)std::transform(args_conf_list.begin(), args_conf_list.end(), std::back_inserter(args_spec_list),
@ -577,6 +639,7 @@ EvalResultPtr JEvaluator::Run(AnalysisEnginePtr engine, const ConfigPtrList &arg
return res; return res;
} }
// Run the Taylor evaluator.
EvalResultPtr TaylorEvaluator::Run(AnalysisEnginePtr engine, const ConfigPtrList &args_conf_list, EvalResultPtr TaylorEvaluator::Run(AnalysisEnginePtr engine, const ConfigPtrList &args_conf_list,
const AnfNodeConfigPtr &) { const AnfNodeConfigPtr &) {
AbstractBasePtrList args_spec_list; AbstractBasePtrList args_spec_list;
@ -598,6 +661,7 @@ EvalResultPtr TaylorEvaluator::Run(AnalysisEnginePtr engine, const ConfigPtrList
return result; return result;
} }
// Configure a list (args_conf_list) and an engine based on the parameters passed in to run a Prim algorithm and return an EvalResultPtr. To be specific:
EvalResultPtr ShardEvaluator::Run(AnalysisEnginePtr engine, const ConfigPtrList &args_conf_list, EvalResultPtr ShardEvaluator::Run(AnalysisEnginePtr engine, const ConfigPtrList &args_conf_list,
const AnfNodeConfigPtr &) { const AnfNodeConfigPtr &) {
AbstractBasePtrList args_spec_list; AbstractBasePtrList args_spec_list;
@ -621,6 +685,12 @@ EvalResultPtr ShardEvaluator::Run(AnalysisEnginePtr engine, const ConfigPtrList
} }
namespace { namespace {
// Reduce the dimension of the tensor.
// axis represents the dimension index to be reduced, orig_abs represents the original tensor, and axis_size represents the size of each dimension.
// Checks if orig_abs is of type AbstractTensor, and throws an exception if it is not. It then takes the shape of the original tensor and calculates the length of the shape.
// Check that axis is in a valid range and throw an exception if it is not
// Check that axis is in a valid range and throw an exception if it is not
// Removes the dimensions specified in the original tensor and returns a new tensor object whose dimensions have been reduced by the specified dimensions.
AbstractBasePtr ReduceDim(int *axis, const AbstractBasePtr &orig_abs, int *axis_size) { AbstractBasePtr ReduceDim(int *axis, const AbstractBasePtr &orig_abs, int *axis_size) {
if (!orig_abs->isa<abstract::AbstractTensor>()) { if (!orig_abs->isa<abstract::AbstractTensor>()) {
MS_LOG(EXCEPTION) << "ValueError: orig_abs should be AbstractTensor, but got a " << orig_abs->ToString() << "."; MS_LOG(EXCEPTION) << "ValueError: orig_abs should be AbstractTensor, but got a " << orig_abs->ToString() << ".";
@ -646,10 +716,13 @@ AbstractBasePtr ReduceDim(int *axis, const AbstractBasePtr &orig_abs, int *axis_
return abs_clone; return abs_clone;
} }
// Accept the physical view (physical_view_abs), the input axis (in_axes), and the axis size (axis_size) as parameters.
AbstractBasePtr GetLogicalViewAbs(const AbstractBasePtr &physical_view_abs, const ValuePtr &in_axes, int *axis_size) { AbstractBasePtr GetLogicalViewAbs(const AbstractBasePtr &physical_view_abs, const ValuePtr &in_axes, int *axis_size) {
MS_EXCEPTION_IF_NULL(physical_view_abs); MS_EXCEPTION_IF_NULL(physical_view_abs);
MS_EXCEPTION_IF_NULL(in_axes); MS_EXCEPTION_IF_NULL(in_axes);
auto physical_view_abs_sequence = dyn_cast<abstract::AbstractSequence>(physical_view_abs); auto physical_view_abs_sequence = dyn_cast<abstract::AbstractSequence>(physical_view_abs);
// Determines whether the physical view is of a sequence type, and if so,
// calls the GetLogicalViewAbs function recursively to combine the abstract base pointer list of the subviews into a new logical view abstract base pointer list.
if (physical_view_abs_sequence != nullptr) { if (physical_view_abs_sequence != nullptr) {
AbstractBasePtrList abs_list = physical_view_abs_sequence->elements(); AbstractBasePtrList abs_list = physical_view_abs_sequence->elements();
AbstractBasePtrList logical_view_abs_list; AbstractBasePtrList logical_view_abs_list;
@ -670,7 +743,9 @@ AbstractBasePtr GetLogicalViewAbs(const AbstractBasePtr &physical_view_abs, cons
} }
return std::make_shared<AbstractTuple>(logical_view_abs_list); return std::make_shared<AbstractTuple>(logical_view_abs_list);
} }
// If the physical view is not of a sequence type, it is processed according to the type of the input axis.
ValuePtr in_axis = in_axes; ValuePtr in_axis = in_axes;
// If the input axis is Int64Imm, the ReduceDim function is called to reduce the dimension of the physical view and the result is returned.
if (in_axis->isa<Int64Imm>()) { if (in_axis->isa<Int64Imm>()) {
int axis = dyn_cast<Int64Imm>(in_axis)->value(); int axis = dyn_cast<Int64Imm>(in_axis)->value();
auto logical_view_abs = ReduceDim(&axis, physical_view_abs, axis_size); auto logical_view_abs = ReduceDim(&axis, physical_view_abs, axis_size);
@ -684,6 +759,7 @@ AbstractBasePtr GetLogicalViewAbs(const AbstractBasePtr &physical_view_abs, cons
return physical_view_abs; return physical_view_abs;
} }
// Extend the dimensions of the tensor.
AbstractBasePtr ExtendDim(int *axis, const AbstractBasePtr &orig_abs, int axis_size) { AbstractBasePtr ExtendDim(int *axis, const AbstractBasePtr &orig_abs, int axis_size) {
MS_EXCEPTION_IF_NULL(orig_abs); MS_EXCEPTION_IF_NULL(orig_abs);
AbstractBasePtr out_abs = nullptr; AbstractBasePtr out_abs = nullptr;
@ -711,65 +787,91 @@ AbstractBasePtr ExtendDim(int *axis, const AbstractBasePtr &orig_abs, int axis_s
return out_abs; return out_abs;
} }
// Process physical view
AbstractBasePtr GetPhysicalViewAbs(const AbstractBasePtr &logical_view_abs, const ValuePtr &out_axes, int axis_size) { AbstractBasePtr GetPhysicalViewAbs(const AbstractBasePtr &logical_view_abs, const ValuePtr &out_axes, int axis_size) {
// Check whether the logical view abstraction is empty, if it is empty, raise the exception
MS_EXCEPTION_IF_NULL(logical_view_abs); MS_EXCEPTION_IF_NULL(logical_view_abs);
// Attempts to convert the abstraction of a logical view to an abstract sequence type
auto logical_view_abs_sequence = dyn_cast<abstract::AbstractSequence>(logical_view_abs); auto logical_view_abs_sequence = dyn_cast<abstract::AbstractSequence>(logical_view_abs);
// if the conversion is successful, the logical view is a sequence.
if (logical_view_abs_sequence != nullptr) { if (logical_view_abs_sequence != nullptr) {
// Gets the element list of a logical view sequence
AbstractBasePtrList logical_view_abs_list = logical_view_abs_sequence->elements(); AbstractBasePtrList logical_view_abs_list = logical_view_abs_sequence->elements();
AbstractBasePtrList physical_view_abs_list; AbstractBasePtrList physical_view_abs_list;
// Try to convert the value of the output axis to the value sequence type
auto out_axes_seq = dyn_cast<ValueSequeue>(out_axes); auto out_axes_seq = dyn_cast<ValueSequeue>(out_axes);
// if the conversion is successful, the output axis is a sequence
if (out_axes_seq != nullptr) { if (out_axes_seq != nullptr) {
// Check whether the size of the output axis sequence is equal to the size of the logical view sequence. if not, throw an exception
if (logical_view_abs_list.size() != out_axes_seq->size()) { if (logical_view_abs_list.size() != out_axes_seq->size()) {
MS_LOG(EXCEPTION) << "The size of vmap's 'out_axes' should be equal to the number of results of 'fn': " MS_LOG(EXCEPTION) << "The size of vmap's 'out_axes' should be equal to the number of results of 'fn': "
<< logical_view_abs_list.size() << ", but got size: " << out_axes_seq->size() << "."; << logical_view_abs_list.size() << ", but got size: " << out_axes_seq->size() << ".";
} }
} }
// Defines an index variable that traverses the output axis sequence
int index = 0; int index = 0;
// For each element in the logical view sequence, convert according to the corresponding output axis value. And add the result to the physical view sequence
(void)std::transform( (void)std::transform(
logical_view_abs_list.begin(), logical_view_abs_list.end(), std::back_inserter(physical_view_abs_list), logical_view_abs_list.begin(), logical_view_abs_list.end(), std::back_inserter(physical_view_abs_list),
[&axis_size, &index, &out_axes_seq, out_axes](const AbstractBasePtr &arg_spec) -> AbstractBasePtr { [&axis_size, &index, &out_axes_seq, out_axes](const AbstractBasePtr &arg_spec) -> AbstractBasePtr {
// Defines a child output axis value that holds the output axis value corresponding to the current element
ValuePtr sub_out_axes = out_axes; ValuePtr sub_out_axes = out_axes;
// if the output axis isa sequence, take the value corresponding to the current index from it and update the index
if (out_axes->isa<ValueSequeue>()) { if (out_axes->isa<ValueSequeue>()) {
sub_out_axes = (*out_axes_seq)[index]; sub_out_axes = (*out_axes_seq)[index];
index++; index++;
} }
// If the current element is an abstract sequence type, this function is called recursively.
if (arg_spec->isa<AbstractSequence>()) { if (arg_spec->isa<AbstractSequence>()) {
return GetPhysicalViewAbs(arg_spec, sub_out_axes, axis_size); return GetPhysicalViewAbs(arg_spec, sub_out_axes, axis_size);
} }
// If the sub-output axis value is an integer type, then the ExtendDim function is called to extend the dimension of the current element based on the axis value and axis size.
if (sub_out_axes->isa<Int64Imm>()) { if (sub_out_axes->isa<Int64Imm>()) {
int axis = dyn_cast<Int64Imm>(sub_out_axes)->value(); int axis = dyn_cast<Int64Imm>(sub_out_axes)->value();
return ExtendDim(&axis, arg_spec, axis_size); return ExtendDim(&axis, arg_spec, axis_size);
} else if (sub_out_axes->isa<None>()) { } else if (sub_out_axes->isa<None>()) {
// If the suboutput axis value is an empty type, return the current element without any conversion.
return arg_spec; return arg_spec;
} }
// If the suboutput axis value is neither an integer nor an empty type,
MS_LOG(EXCEPTION) << "The axis in vmap's 'out_axes' should be a None or a scalar of type Int64Imm, but got a " MS_LOG(EXCEPTION) << "The axis in vmap's 'out_axes' should be a None or a scalar of type Int64Imm, but got a "
<< sub_out_axes->ToString() << "."; << sub_out_axes->ToString() << ".";
}); });
// If the logical view is an abstract list type, Returns an abstract list type
if (logical_view_abs->isa<AbstractList>()) { if (logical_view_abs->isa<AbstractList>()) {
// Otherwise an abstract tuple consisting of a sequence of physical views is returned.
return std::make_shared<AbstractList>(physical_view_abs_list); return std::make_shared<AbstractList>(physical_view_abs_list);
} }
return std::make_shared<AbstractTuple>(physical_view_abs_list); return std::make_shared<AbstractTuple>(physical_view_abs_list);
} }
// for the single output case, outputs: A, and out_axes: 1 or (1,). // for the single output case, outputs: A, and out_axes: 1 or (1,).
// If the logical view is not a sequence but a single output, then the output axis should also be a single value
// Define a suboutput axis value to hold the value of the output axis
ValuePtr sub_out_axes = out_axes; ValuePtr sub_out_axes = out_axes;
// Try to convert the value of the output axis to the value sequence type
ValueSequeuePtr out_axes_seq = dyn_cast<ValueSequeue>(out_axes); ValueSequeuePtr out_axes_seq = dyn_cast<ValueSequeue>(out_axes);
// if the conversion is successful, the output axis is a sequence
if (out_axes_seq != nullptr) { if (out_axes_seq != nullptr) {
// Check whether the output axis sequence size is 1, if not, throw an exception
if (out_axes_seq->size() != 1) { if (out_axes_seq->size() != 1) {
MS_LOG(EXCEPTION) << "The size of vmap's 'out_axes' should be equal to the result size: 1, but got size: " MS_LOG(EXCEPTION) << "The size of vmap's 'out_axes' should be equal to the result size: 1, but got size: "
<< out_axes_seq->size() << "."; << out_axes_seq->size() << ".";
} }
sub_out_axes = (*out_axes_seq)[0]; sub_out_axes = (*out_axes_seq)[0];
} }
// Define an axis variable that holds the sub-output axis value
int axis = 0; int axis = 0;
// Try to convert the sub-output axis value to an integer type
auto axis_int_ptr = dyn_cast<Int64Imm>(sub_out_axes); auto axis_int_ptr = dyn_cast<Int64Imm>(sub_out_axes);
// if the conversion succeeds, the integer value is assigned to the axis variable
if (axis_int_ptr != nullptr) { if (axis_int_ptr != nullptr) {
axis = LongToInt(axis_int_ptr->value()); axis = LongToInt(axis_int_ptr->value());
} else { } else {
MS_LOG(EXCEPTION) << "The axis in vmap's 'out_axes' should be a None or a scalar of type Int64Imm, but got a " MS_LOG(EXCEPTION) << "The axis in vmap's 'out_axes' should be a None or a scalar of type Int64Imm, but got a "
<< sub_out_axes->ToString() << "."; << sub_out_axes->ToString() << ".";
} }
// Call ExtendDim function, extending the dimension of the logical view based on axis variable and axis size, and return the result
return ExtendDim(&axis, logical_view_abs, axis_size); return ExtendDim(&axis, logical_view_abs, axis_size);
} }
} // namespace } // namespace
@ -829,15 +931,19 @@ EvalResultPtr VmapEvaluator::Run(AnalysisEnginePtr engine, const ConfigPtrList &
return res; return res;
} }
// VirtualEvaluator::Eval method to evaluate the output of VirtualEvaluator
EvalResultPtr VirtualEvaluator::Eval(AnalysisEnginePtr, const AbstractBasePtrList &args_spec_list, EvalResultPtr VirtualEvaluator::Eval(AnalysisEnginePtr, const AbstractBasePtrList &args_spec_list,
const AnfNodeConfigPtr &out_conf) { const AnfNodeConfigPtr &out_conf) {
// Check whether the size of the parameter list is as expected, and throw an exception if it is not
if (args_spec_list.size() != args_spec_list_.size()) { if (args_spec_list.size() != args_spec_list_.size()) {
MS_LOG(EXCEPTION) << "Arguments mismatch, parameters no: " << args_spec_list_.size() MS_LOG(EXCEPTION) << "Arguments mismatch, parameters no: " << args_spec_list_.size()
<< ", arguments no: " << args_spec_list.size(); << ", arguments no: " << args_spec_list.size();
} }
// Gets the value of the environment variable MS_DEV_ENABLE_DDE. If it is not 0, the function to eliminate unused elements is enabled
static const auto enable_eliminate_unused_element = (common::GetEnv("MS_DEV_ENABLE_DDE") != "0"); static const auto enable_eliminate_unused_element = (common::GetEnv("MS_DEV_ENABLE_DDE") != "0");
// Check each parameter and argument match; // Check each parameter and argument match;
for (std::size_t i = 0; i < args_spec_list.size(); i++) { for (std::size_t i = 0; i < args_spec_list.size(); i++) {
// If the argument is null, an exception is thrown
MS_EXCEPTION_IF_NULL(args_spec_list[i]); MS_EXCEPTION_IF_NULL(args_spec_list[i]);
// For VirtualAbstractClosure, likely J's bprop, we just set its tuple arguments as used before really grad. // For VirtualAbstractClosure, likely J's bprop, we just set its tuple arguments as used before really grad.
if (enable_eliminate_unused_element && args_spec_list[i]->isa<abstract::AbstractSequence>()) { if (enable_eliminate_unused_element && args_spec_list[i]->isa<abstract::AbstractSequence>()) {
@ -845,14 +951,18 @@ EvalResultPtr VirtualEvaluator::Eval(AnalysisEnginePtr, const AbstractBasePtrLis
<< "]: " << args_spec_list[i]->ToString(); << "]: " << args_spec_list[i]->ToString();
SetSequenceElementsUseFlagsRecursively(args_spec_list[i], true); SetSequenceElementsUseFlagsRecursively(args_spec_list[i], true);
} }
// Join the parameters with the expected ones, throwing an exception if they are incompatible
(void)args_spec_list[i]->Join(args_spec_list_[i]); (void)args_spec_list[i]->Join(args_spec_list_[i]);
} }
// Returns evaluation results, including output and attribute value mapping
return std::make_shared<EvalResult>(output_, std::make_shared<AttrValueMap>()); return std::make_shared<EvalResult>(output_, std::make_shared<AttrValueMap>());
} }
// Evaluator::SingleRun method for performing a single evaluation
EvalResultPtr Evaluator::SingleRun(AnalysisEnginePtr engine, const ConfigPtrList &args_conf_list, EvalResultPtr Evaluator::SingleRun(AnalysisEnginePtr engine, const ConfigPtrList &args_conf_list,
const AnfNodeConfigPtr &out_conf) { const AnfNodeConfigPtr &out_conf) {
EvalResultPtr result; EvalResultPtr result;
try { try {
// Call the Run method, which implements different logic depending on the type of evaluator
result = this->Run(engine, args_conf_list, out_conf); result = this->Run(engine, args_conf_list, out_conf);
} catch (const std::exception &ex) { } catch (const std::exception &ex) {
MS_LOG(INFO) << "Eval " << ToString() << " throw exception."; MS_LOG(INFO) << "Eval " << ToString() << " throw exception.";

File diff suppressed because it is too large Load Diff

View File

@ -40,17 +40,25 @@ inline AbstractBasePtr GetEvaluatedValue(const AnfNodeConfigPtr &conf) {
} }
AnfNodePtr BuildValueNode(const ValuePtr &v, const AbstractBasePtr &abs_base) { AnfNodePtr BuildValueNode(const ValuePtr &v, const AbstractBasePtr &abs_base) {
// Ensure that the abstract base is not null
MS_EXCEPTION_IF_NULL(abs_base); MS_EXCEPTION_IF_NULL(abs_base);
// Create a new value node with the given value
AnfNodePtr value_node = NewValueNode(v); AnfNodePtr value_node = NewValueNode(v);
// Set the abstract base of the value node to the given abstract base
value_node->set_abstract(abs_base); value_node->set_abstract(abs_base);
// Log a debug message indicating the creation of a new value node with its corresponding abstract base
MS_LOG(DEBUG) << "Create ValueNode: " << value_node->ToString() << ", with abstract: " << abs_base->ToString(); MS_LOG(DEBUG) << "Create ValueNode: " << value_node->ToString() << ", with abstract: " << abs_base->ToString();
// Return the new value node
return value_node; return value_node;
} }
bool IsVisible(FuncGraphPtr fg, const FuncGraphPtr &parent) { bool IsVisible(FuncGraphPtr fg, const FuncGraphPtr &parent) {
// Iterate until the current function graph is nullptr or matches the parent function graph
while (fg != nullptr && fg != parent) { while (fg != nullptr && fg != parent) {
// Move up to the parent function graph
fg = fg->parent(); fg = fg->parent();
} }
// Check if the current function graph matches the parent function graph
return fg == parent; return fg == parent;
} }
@ -77,12 +85,16 @@ bool CanSpecializeValueNode(const AnfNodePtr &node) {
void PurifyAbstractOfSequence(ProgramSpecializer *const specializer) { void PurifyAbstractOfSequence(ProgramSpecializer *const specializer) {
constexpr int recursive_level = 2; constexpr int recursive_level = 2;
// Iterate over the sequence abstract list in the specializer
for (auto &abstract_and_node : specializer->sequence_abstract_list()) { for (auto &abstract_and_node : specializer->sequence_abstract_list()) {
auto &sequence_abs = abstract_and_node.first; auto &sequence_abs = abstract_and_node.first;
// Purify the elements of the abstract value
if (!sequence_abs->PurifyElements()) { if (!sequence_abs->PurifyElements()) {
// If purification fails, log an error message with the abstract value and corresponding node information
MS_LOG(ERROR) << "Purify elements failed, abstract: " << sequence_abs->ToString() MS_LOG(ERROR) << "Purify elements failed, abstract: " << sequence_abs->ToString()
<< ", node: " << abstract_and_node.second->DebugString(recursive_level); << ", node: " << abstract_and_node.second->DebugString(recursive_level);
} else { } else {
// If purification is successful, log a debug message with the abstract value and corresponding node information
MS_LOG(DEBUG) << "Purify elements, abstract: " << sequence_abs->ToString() MS_LOG(DEBUG) << "Purify elements, abstract: " << sequence_abs->ToString()
<< ", node: " << abstract_and_node.second->DebugString(recursive_level); << ", node: " << abstract_and_node.second->DebugString(recursive_level);
} }
@ -169,60 +181,96 @@ void EliminateCollectedSequenceNodes(ProgramSpecializer *const specializer) {
} // namespace } // namespace
FuncGraphPtr ProgramSpecializer::Run(const FuncGraphPtr &fg, const AnalysisContextPtr &context) { FuncGraphPtr ProgramSpecializer::Run(const FuncGraphPtr &fg, const AnalysisContextPtr &context) {
// Check if the function graph and context are not null
MS_EXCEPTION_IF_NULL(fg); MS_EXCEPTION_IF_NULL(fg);
MS_EXCEPTION_IF_NULL(context); MS_EXCEPTION_IF_NULL(context);
// Log a debug message indicating the specialization of the topmost function graph
MS_LOG(DEBUG) << "Specialize topmost function graph: " MS_LOG(DEBUG) << "Specialize topmost function graph: "
<< (context->func_graph() ? context->func_graph()->ToString() : "FG(Null)"); << (context->func_graph() ? context->func_graph()->ToString() : "FG(Null)");
// If top_context_ is null, set it to the given context and log an info message
if (top_context_ == nullptr) { if (top_context_ == nullptr) {
top_context_ = context; top_context_ = context;
MS_LOG(INFO) << "Specialize set top func graph context: " << context->ToString(); MS_LOG(INFO) << "Specialize set top func graph context: " << context->ToString();
} }
// Specialize the function graph using the given context and store the result in 'res'
auto res = SpecializeFuncGraph(fg, context); auto res = SpecializeFuncGraph(fg, context);
// Eliminate collected sequence nodes
EliminateCollectedSequenceNodes(this); EliminateCollectedSequenceNodes(this);
// Return the specialized function graph 'res'
return res; return res;
} }
FuncGraphPtr ProgramSpecializer::SpecializeFuncGraph(const FuncGraphPtr &fg, const AnalysisContextPtr &context) { FuncGraphPtr ProgramSpecializer::SpecializeFuncGraph(const FuncGraphPtr &fg, const AnalysisContextPtr &context) {
// Check if the function graph and context are not null
MS_EXCEPTION_IF_NULL(fg); MS_EXCEPTION_IF_NULL(fg);
MS_EXCEPTION_IF_NULL(context); MS_EXCEPTION_IF_NULL(context);
// Check if a specialization for the given context already exists
auto iter = specializations_.find(context->SpecializeKey()); auto iter = specializations_.find(context->SpecializeKey());
if (iter != specializations_.end()) { if (iter != specializations_.end()) {
// If a specialization exists, return the corresponding specialized function graph
MS_EXCEPTION_IF_NULL(iter->second); MS_EXCEPTION_IF_NULL(iter->second);
return iter->second->specialized_func_graph(); return iter->second->specialized_func_graph();
} }
// Create a new FuncGraphSpecializer instance for the function graph and context
std::shared_ptr<FuncGraphSpecializer> fg_spec = std::make_shared<FuncGraphSpecializer>(this, fg, context); std::shared_ptr<FuncGraphSpecializer> fg_spec = std::make_shared<FuncGraphSpecializer>(this, fg, context);
// Get the specialized function graph from the FuncGraphSpecializer
FuncGraphPtr specialized_func_graph = fg_spec->specialized_func_graph(); FuncGraphPtr specialized_func_graph = fg_spec->specialized_func_graph();
// Store the FuncGraphSpecializer instance in the specializations map
specializations_[context->SpecializeKey()] = fg_spec; specializations_[context->SpecializeKey()] = fg_spec;
// Run the specialization process
fg_spec->Run(); fg_spec->Run();
// Return the specialized function graph
return specialized_func_graph; return specialized_func_graph;
} }
std::shared_ptr<FuncGraphSpecializer> ProgramSpecializer::GetFuncGraphSpecializer(const AnalysisContextPtr &context) { std::shared_ptr<FuncGraphSpecializer> ProgramSpecializer::GetFuncGraphSpecializer(const AnalysisContextPtr &context) {
// Check if the context is not null
MS_EXCEPTION_IF_NULL(context); MS_EXCEPTION_IF_NULL(context);
// Check if a specialization for the given context exists
auto iter = specializations_.find(context->SpecializeKey()); auto iter = specializations_.find(context->SpecializeKey());
if (iter != specializations_.end()) { if (iter != specializations_.end()) {
// Return the corresponding FuncGraphSpecializer instance
return iter->second; return iter->second;
} }
// If no specialization exists, return nullptr
return nullptr; return nullptr;
} }
void ProgramSpecializer::PutSpecializedAbstract(const CNodePtr &cnode, const AnfNodePtr &func, void ProgramSpecializer::PutSpecializedAbstract(const CNodePtr &cnode, const AnfNodePtr &func,
const AbstractFunctionPtr &old_abs_func, const AbstractFunctionPtr &old_abs_func,
const AbstractFunctionPtr &new_abs_func) { const AbstractFunctionPtr &new_abs_func) {
// Check if a specialization for the old abstract function already exists in the specialized abstract map
auto iter = specialized_abs_map_.find(old_abs_func); auto iter = specialized_abs_map_.find(old_abs_func);
if (iter == specialized_abs_map_.end()) { if (iter == specialized_abs_map_.end()) {
// If no specialization exists for the old abstract function, add a new entry to the map
MS_LOG(DEBUG) << "Emplace cnode: " << cnode->DebugString() << ", func: " << func->ToString() MS_LOG(DEBUG) << "Emplace cnode: " << cnode->DebugString() << ", func: " << func->ToString()
<< ", old_abstract: " << old_abs_func->ToString() << ", new_abs_func: " << new_abs_func->ToString(); << ", old_abstract: " << old_abs_func->ToString() << ", new_abs_func: " << new_abs_func->ToString();
(void)specialized_abs_map_.emplace(old_abs_func, new_abs_func); (void)specialized_abs_map_.emplace(old_abs_func, new_abs_func);
} else { } else {
// If a specialization already exists, compare the new and existing specialized abstract functions
MS_LOG(DEBUG) << "Duplicate abstract from cnode: " << cnode->DebugString() << ", func: " << func->ToString() MS_LOG(DEBUG) << "Duplicate abstract from cnode: " << cnode->DebugString() << ", func: " << func->ToString()
<< ", old_abstract: " << old_abs_func->ToString() << ", new_abs_func: " << new_abs_func->ToString(); << ", old_abstract: " << old_abs_func->ToString() << ", new_abs_func: " << new_abs_func->ToString();
if (!(*iter->second == *new_abs_func)) { if (!(*iter->second == *new_abs_func)) {
// If the specialized abstract functions do not match, log an error and replace the existing specialization
MS_LOG(DEBUG) << "Duplicate abstract from cnode: " << cnode->DebugString() << ", func: " << func->ToString() MS_LOG(DEBUG) << "Duplicate abstract from cnode: " << cnode->DebugString() << ", func: " << func->ToString()
<< ", old_abstract: " << old_abs_func->ToString() << ", first: " << iter->second->ToString() << ", old_abstract: " << old_abs_func->ToString() << ", first: " << iter->second->ToString()
<< ", new_abs_func: " << new_abs_func->ToString(); << ", new_abs_func: " << new_abs_func->ToString();
// Cannot determined which one to use.
// Replace the existing specialization with an AbstractError indicating a poly node
const auto poly_abstract = std::make_shared<AbstractError>(kPolyNode, func); const auto poly_abstract = std::make_shared<AbstractError>(kPolyNode, func);
iter->second = poly_abstract; iter->second = poly_abstract;
} }
@ -230,19 +278,25 @@ void ProgramSpecializer::PutSpecializedAbstract(const CNodePtr &cnode, const Anf
} }
AbstractBasePtr ProgramSpecializer::GetSpecializedAbstract(const AbstractFunctionPtr &old_abs_func) { AbstractBasePtr ProgramSpecializer::GetSpecializedAbstract(const AbstractFunctionPtr &old_abs_func) {
// Check if a specialization for the old abstract function exists in the specialized abstract map
auto iter = specialized_abs_map_.find(old_abs_func); auto iter = specialized_abs_map_.find(old_abs_func);
if (iter != specialized_abs_map_.end()) { if (iter != specialized_abs_map_.end()) {
// If a specialization is found, log the details and return the specialized abstract function
MS_LOG(DEBUG) << "Find abstract for old_abstract: " << old_abs_func->ToString() MS_LOG(DEBUG) << "Find abstract for old_abstract: " << old_abs_func->ToString()
<< ", new_abs_func: " << iter->second->ToString(); << ", new_abs_func: " << iter->second->ToString();
// Check if the specialized abstract function is of type AbstractFunction
if (iter->second->isa<AbstractFunction>()) { if (iter->second->isa<AbstractFunction>()) {
return iter->second; return iter->second;
} }
// Return nullptr if the specialized abstract function is not of type AbstractFunction
return nullptr; return nullptr;
} }
// If no specialization is found, log an error and return nullptr
MS_LOG(DEBUG) << "Cannot find abstract for old_abstract: " << old_abs_func->ToString(); MS_LOG(DEBUG) << "Cannot find abstract for old_abstract: " << old_abs_func->ToString();
return nullptr; return nullptr;
} }
AbstractBasePtr ProgramSpecializer::SpecializeAbstractFuncRecursively(const AbstractFunctionPtr &old_abs_func) { AbstractBasePtr ProgramSpecializer::SpecializeAbstractFuncRecursively(const AbstractFunctionPtr &old_abs_func) {
AbstractBasePtr new_abs = nullptr; AbstractBasePtr new_abs = nullptr;
if (old_abs_func->isa<AbstractFuncUnion>()) { if (old_abs_func->isa<AbstractFuncUnion>()) {
@ -301,23 +355,32 @@ AbstractBasePtr ProgramSpecializer::SpecializeAbstractFuncRecursively(const Abst
} }
void ProgramSpecializer::SpecializeCNodeInput0FuncGraph() { void ProgramSpecializer::SpecializeCNodeInput0FuncGraph() {
// Retrieve all nodes in the manager.
const auto &all_nodes = mng_->all_nodes(); const auto &all_nodes = mng_->all_nodes();
// Iterate over each node.
for (auto node : all_nodes) { for (auto node : all_nodes) {
// Skip nodes that are not CNodes.
if (!node->isa<CNode>()) { if (!node->isa<CNode>()) {
continue; continue;
} }
// Get the input 0 of the CNode.
auto &input0 = node->cast<CNodePtr>()->input(0); auto &input0 = node->cast<CNodePtr>()->input(0);
MS_EXCEPTION_IF_NULL(input0); MS_EXCEPTION_IF_NULL(input0);
// Skip if the input is a ValueNode of type FuncGraph.
if (IsValueNode<FuncGraph>(input0)) { if (IsValueNode<FuncGraph>(input0)) {
continue; continue;
} }
// Check the abstract value of input0 and skip if it does not match any specific types.
const auto &old_abs = input0->abstract(); const auto &old_abs = input0->abstract();
if (!(old_abs->isa<FuncGraphAbstractClosure>() || old_abs->isa<MetaFuncGraphAbstractClosure>() || if (!(old_abs->isa<FuncGraphAbstractClosure>() || old_abs->isa<MetaFuncGraphAbstractClosure>() ||
old_abs->isa<AbstractFuncUnion>() || old_abs->isa<PartialAbstractClosure>())) { old_abs->isa<AbstractFuncUnion>() || old_abs->isa<PartialAbstractClosure>())) {
continue; continue;
} }
// Cast the abstract value to AbstractFunctionPtr.
auto old_abs_func = old_abs->cast<AbstractFunctionPtr>(); auto old_abs_func = old_abs->cast<AbstractFunctionPtr>();
// Specialize the abstract function recursively.
auto new_abs_func = SpecializeAbstractFuncRecursively(old_abs_func); auto new_abs_func = SpecializeAbstractFuncRecursively(old_abs_func);
// Update the abstract value of input0 if specialization is successful.
if (new_abs_func != nullptr) { if (new_abs_func != nullptr) {
input0->set_abstract(new_abs_func); input0->set_abstract(new_abs_func);
MS_LOG(DEBUG) << "Find specialized abstract for node: " << input0->DebugString() MS_LOG(DEBUG) << "Find specialized abstract for node: " << input0->DebugString()
@ -330,6 +393,7 @@ void ProgramSpecializer::SpecializeCNodeInput0FuncGraph() {
} }
} }
static int64_t GetNextCounter() { static int64_t GetNextCounter() {
static int64_t g_CloneCounter = 1; static int64_t g_CloneCounter = 1;
return g_CloneCounter++; return g_CloneCounter++;
@ -338,38 +402,49 @@ static int64_t GetNextCounter() {
FuncGraphSpecializer::FuncGraphSpecializer(ProgramSpecializer *const s, const FuncGraphPtr &fg, FuncGraphSpecializer::FuncGraphSpecializer(ProgramSpecializer *const s, const FuncGraphPtr &fg,
const AnalysisContextPtr &context) const AnalysisContextPtr &context)
: specializer_(s), func_graph_(fg), context_(context) { : specializer_(s), func_graph_(fg), context_(context) {
// Retrieve the parent function graph specializer from the program specializer.
parent_ = s->GetFuncGraphSpecializer(context->parent()); parent_ = s->GetFuncGraphSpecializer(context->parent());
if (parent_ == nullptr && context->parent()->func_graph() != nullptr) { // If context's not dummy context. // If the parent is not null and the parent's context has a function graph (not a dummy context),
// then throw an exception.
if (parent_ == nullptr && context->parent()->func_graph() != nullptr) {
MS_LOG(EXCEPTION) << "Parent func graph should be handled in advance, fg: " << fg->ToString() MS_LOG(EXCEPTION) << "Parent func graph should be handled in advance, fg: " << fg->ToString()
<< ", context: " << context->ToString() << ", parent context: " << context->parent()->ToString(); << ", context: " << context->ToString() << ", parent context: " << context->parent()->ToString();
} }
engine_ = s->engine(); // Retrieve the engine from the program specializer.
engine_ = s->engine()
// Clone the original function graph using the TraceSpecialize clone method.
cloner_ = SpecializerClone(fg, std::make_shared<TraceSpecialize>(GetNextCounter())); cloner_ = SpecializerClone(fg, std::make_shared<TraceSpecialize>(GetNextCounter()));
// Get the specialized function graph from the cloned function graphs.
specialized_func_graph_ = cloner_->cloned_func_graphs().find(fg)->second; specialized_func_graph_ = cloner_->cloned_func_graphs().find(fg)->second;
// Add the return node and the parameter nodes of the function graph as todo items.
AddTodoItem(fg->get_return()); AddTodoItem(fg->get_return());
AddTodoItem(fg->parameters()); AddTodoItem(fg->parameters());
} }
AnfNodePtr FuncGraphSpecializer::ReplicateDisconnectedNode(const AnfNodePtr &node) { AnfNodePtr FuncGraphSpecializer::ReplicateDisconnectedNode(const AnfNodePtr &node) {
MS_EXCEPTION_IF_NULL(node); MS_EXCEPTION_IF_NULL(node);
// If the node is a ValueNode, simply return it as it doesn't need to be replicated.
if (node->isa<ValueNode>()) { if (node->isa<ValueNode>()) {
return node; return node;
} }
// Get the top specializer for the node.
std::shared_ptr<FuncGraphSpecializer> specializer = GetTopSpecializer(node); std::shared_ptr<FuncGraphSpecializer> specializer = GetTopSpecializer(node);
// Check if the node has already been replicated, and if so, return the replicated node.
// If had replicated, just return that.
auto iter = specializer->cloned_nodes().find(node); auto iter = specializer->cloned_nodes().find(node);
if (iter != specializer->cloned_nodes().end()) { if (iter != specializer->cloned_nodes().end()) {
return iter->second; return iter->second;
} }
// Clone the disconnected node using the specializer's cloner.
auto new_node = specializer->cloner_->CloneDisconnected(node); auto new_node = specializer->cloner_->CloneDisconnected(node);
// If the original node is a CNode, ensure that the cloned node is also a CNode and update its inputs.
if (node->isa<CNode>()) { if (node->isa<CNode>()) {
if (!new_node->isa<CNode>()) { if (!new_node->isa<CNode>()) {
MS_LOG(EXCEPTION) << "new_node must be a CNode, but is " << new_node->DebugString() << "."; MS_LOG(EXCEPTION) << "new_node must be a CNode, but is " << new_node->DebugString() << ".";
} }
UpdateNewCNodeInputs(node, new_node); UpdateNewCNodeInputs(node, new_node);
} }
// Check if the node has been replicated and ensure it is not the same as the original node.
iter = specializer->cloned_nodes().find(node); iter = specializer->cloned_nodes().find(node);
if (iter != specializer->cloned_nodes().end()) { if (iter != specializer->cloned_nodes().end()) {
if (iter->second == node) { if (iter->second == node) {
@ -381,33 +456,42 @@ AnfNodePtr FuncGraphSpecializer::ReplicateDisconnectedNode(const AnfNodePtr &nod
return new_node; return new_node;
} }
void FuncGraphSpecializer::UpdateNewCNodeInputs(const AnfNodePtr &node, const AnfNodePtr &new_node) { void FuncGraphSpecializer::UpdateNewCNodeInputs(const AnfNodePtr &node, const AnfNodePtr &new_node) {
// Check if node and c_node are not null.
MS_EXCEPTION_IF_NULL(node); MS_EXCEPTION_IF_NULL(node);
auto c_node = node->cast<CNodePtr>(); auto c_node = node->cast<CNodePtr>();
MS_EXCEPTION_IF_NULL(c_node); MS_EXCEPTION_IF_NULL(c_node);
// Get the inputs of the c_node.
auto inputs = c_node->inputs(); auto inputs = c_node->inputs();
// Create a vector to store the new inputs.
std::vector<AnfNodePtr> new_inputs; std::vector<AnfNodePtr> new_inputs;
// Iterate over each input and transform them.
(void)std::transform( (void)std::transform(
inputs.begin(), inputs.end(), std::back_inserter(new_inputs), [this](const AnfNodePtr &inp) -> AnfNodePtr { inputs.begin(), inputs.end(), std::back_inserter(new_inputs), [this](const AnfNodePtr &inp) -> AnfNodePtr {
// Replicate the disconnected node.
auto new_inp = ReplicateDisconnectedNode(inp); auto new_inp = ReplicateDisconnectedNode(inp);
// Refer the comments in BuildReplacedNode. // Check if the input is a CNode.
if (inp->isa<CNode>()) { if (inp->isa<CNode>()) {
auto c_inp = inp->cast<CNodePtr>(); auto c_inp = inp->cast<CNodePtr>();
MS_EXCEPTION_IF_NULL(c_inp); MS_EXCEPTION_IF_NULL(c_inp);
auto c_new_inp = new_inp->cast<CNodePtr>(); auto c_new_inp = new_inp->cast<CNodePtr>();
MS_EXCEPTION_IF_NULL(c_new_inp); MS_EXCEPTION_IF_NULL(c_new_inp);
MS_EXCEPTION_IF_NULL(c_new_inp->func_graph()); MS_EXCEPTION_IF_NULL(c_new_inp->func_graph());
// Replace the original CNode with the replicated CNode in the function graph.
MS_LOG(DEBUG) << "Replace in order, inp node: " << inp->DebugString() << " -> " << new_inp->DebugString(); MS_LOG(DEBUG) << "Replace in order, inp node: " << inp->DebugString() << " -> " << new_inp->DebugString();
c_new_inp->func_graph()->ReplaceInOrder(c_inp, c_new_inp); c_new_inp->func_graph()->ReplaceInOrder(c_inp, c_new_inp);
} }
return new_inp; return new_inp;
}); });
// Set the new inputs for the new_node.
auto c_new_node = new_node->cast<CNodePtr>(); auto c_new_node = new_node->cast<CNodePtr>();
MS_EXCEPTION_IF_NULL(c_new_node); MS_EXCEPTION_IF_NULL(c_new_node);
c_new_node->set_inputs(new_inputs); c_new_node->set_inputs(new_inputs);
} }
AnfNodePtr FuncGraphSpecializer::GetReplicatedNode(const AnfNodePtr &node) { AnfNodePtr FuncGraphSpecializer::GetReplicatedNode(const AnfNodePtr &node) {
std::shared_ptr<FuncGraphSpecializer> specializer = GetTopSpecializer(node); std::shared_ptr<FuncGraphSpecializer> specializer = GetTopSpecializer(node);
auto iter = specializer->cloned_nodes().find(node); auto iter = specializer->cloned_nodes().find(node);
@ -464,13 +548,17 @@ std::shared_ptr<FuncGraphSpecializer> FuncGraphSpecializer::GetTopSpecializer(co
} }
void FuncGraphSpecializer::Run() { void FuncGraphSpecializer::Run() {
// Print debug information about the original and cloned function graphs.
MS_LOG(DEBUG) << "Before run, origin func graph name: " << (func_graph_ ? func_graph_->ToString() : "FG(Null)") MS_LOG(DEBUG) << "Before run, origin func graph name: " << (func_graph_ ? func_graph_->ToString() : "FG(Null)")
<< ", cloned func graph name: " << ", cloned func graph name: "
<< (specialized_func_graph_ ? specialized_func_graph_->ToString() : "FG(Null)") << ", func graph: " << (specialized_func_graph_ ? specialized_func_graph_->ToString() : "FG(Null)") << ", func graph: "
<< (func_graph_ ? func_graph_->get_return() ? func_graph_->get_return()->DebugString() : "return null" << (func_graph_ ? func_graph_->get_return() ? func_graph_->get_return()->DebugString() : "return null"
: "FG(null)"); : "FG(null)");
// Perform the first pass of the specialization process.
FirstPass(); FirstPass();
// Perform the second pass of the specialization process.
SecondPass(); SecondPass();
// Print debug information after the specialization process is completed.
MS_LOG(DEBUG) << "After run, origin func graph name: " << (func_graph_ ? func_graph_->ToString() : "FG(Null)") MS_LOG(DEBUG) << "After run, origin func graph name: " << (func_graph_ ? func_graph_->ToString() : "FG(Null)")
<< ", cloned func graph name: " << ", cloned func graph name: "
<< (specialized_func_graph_ ? specialized_func_graph_->ToString() : "FG(Null)") << ", new func graph: " << (specialized_func_graph_ ? specialized_func_graph_->ToString() : "FG(Null)") << ", new func graph: "
@ -480,6 +568,7 @@ void FuncGraphSpecializer::Run() {
: "FG(null)"); : "FG(null)");
} }
void FuncGraphSpecializer::FirstPass() { void FuncGraphSpecializer::FirstPass() {
while (todo_.size()) { while (todo_.size()) {
AnfNodePtr node = todo_.back(); AnfNodePtr node = todo_.back();
@ -607,34 +696,49 @@ void UpdateSequenceNode(const AnfNodePtr &new_node, const AnfNodePtr &old_node,
// Purify specific input of a CNode. // Purify specific input of a CNode.
template <typename T> template <typename T>
void PurifySequenceValueNode(const CNodePtr &cnode, size_t index, ProgramSpecializer *const specializer) { void PurifySequenceValueNode(const CNodePtr &cnode, size_t index, ProgramSpecializer *const specializer) {
// Get the original input value at the specified index.
const auto &old_input = cnode->input(index); const auto &old_input = cnode->input(index);
// Attempt to cast the input value to a shared pointer of type T.
auto sequence_value = GetValueNode<std::shared_ptr<T>>(old_input); auto sequence_value = GetValueNode<std::shared_ptr<T>>(old_input);
// If the cast fails or the sequence value is null, return without further processing.
if (sequence_value == nullptr) { if (sequence_value == nullptr) {
return; return;
} }
// Retrieve the use flags for the elements in the sequence node.
auto flags = GetSequenceNodeElementsUseFlags(old_input); auto flags = GetSequenceNodeElementsUseFlags(old_input);
// If the flags are null, return without further processing.
if (flags == nullptr) { if (flags == nullptr) {
return; return;
} }
// Initialize variables for collecting dead node positions and updated elements.
std::vector<size_t> dead_node_positions; std::vector<size_t> dead_node_positions;
ValuePtrList elements; ValuePtrList elements;
// Iterate over each element in the sequence node.
for (size_t i = 0; i < (*flags).size(); ++i) { for (size_t i = 0; i < (*flags).size(); ++i) {
// Get the old sequence value at position i.
ValuePtr old_sequence_value = sequence_value->value()[i]; ValuePtr old_sequence_value = sequence_value->value()[i];
auto old_sequence_str_value = old_sequence_value->cast<StringImmPtr>(); // Check if the flag for this element is false. If so, replace the element with zero and log the information.
if (!(*flags)[i]) { if (!(*flags)[i]) {
auto zero = MakeValue(0); auto zero = MakeValue(0);
(void)elements.emplace_back(zero); (void)elements.emplace_back(zero);
MS_LOG(DEBUG) << "Erase elements[" << i << "] as zero for " << old_input->DebugString() << ", which is inputs[" MS_LOG(DEBUG) << "Erase elements[" << i << "] as zero for " << old_input->DebugString() << ", which is inputs["
<< index << "] of " << cnode->DebugString(); << index << "] of " << cnode->DebugString();
} else if (old_sequence_str_value != nullptr && old_sequence_str_value->value() == kDeadNodeName) { }
// Check if the old sequence value is a StringImmPtr and its value is equal to kDeadNodeName.
// If so, collect the position for erasing later and add the old sequence value to the updated elements.
else if (old_sequence_str_value != nullptr && old_sequence_str_value->value() == kDeadNodeName) {
MS_LOG(DEBUG) << "Collect for erasing elements[" << i << "] DeadNode as zero for " << old_input->DebugString() MS_LOG(DEBUG) << "Collect for erasing elements[" << i << "] DeadNode as zero for " << old_input->DebugString()
<< ", which is inputs[" << index << "] of " << cnode->DebugString(); << ", which is inputs[" << index << "] of " << cnode->DebugString();
(void)dead_node_positions.emplace_back(i); (void)dead_node_positions.emplace_back(i);
(void)elements.emplace_back(old_sequence_value); (void)elements.emplace_back(old_sequence_value);
} else { }
// Otherwise, add the old sequence value to the updated elements.
else {
(void)elements.emplace_back(old_sequence_value); (void)elements.emplace_back(old_sequence_value);
} }
} }
}
auto new_sequence_value = std::make_shared<T>(elements); auto new_sequence_value = std::make_shared<T>(elements);
auto new_input = NewValueNode(new_sequence_value); auto new_input = NewValueNode(new_sequence_value);
auto new_input_abs = new_sequence_value->ToAbstract(); auto new_input_abs = new_sequence_value->ToAbstract();

View File

@ -81,17 +81,19 @@ size_t StackFrameDepth() { return stack_frame_depth; }
size_t StackFrameMaxDepth() { return stack_frame_max_depth; } size_t StackFrameMaxDepth() { return stack_frame_max_depth; }
EvalResultPtr PrimitiveEvalCache::Get(const PrimitivePtr &prim, const AbstractBasePtrList &args) const { EvalResultPtr PrimitiveEvalCache::Get(const PrimitivePtr &prim, const AbstractBasePtrList &args) const {
std::lock_guard<std::mutex> guard(mutex_); std::lock_guard<std::mutex> guard(mutex_); // Locks the mutex to ensure atomic execution
auto cache_iter = prim_cache_.find(prim->name());
if (cache_iter == prim_cache_.end()) { auto cache_iter = prim_cache_.find(prim->name()); // Looks up the prim_cache_ map using the name() method of the Primitive object prim
return nullptr; if (cache_iter == prim_cache_.end()) { // If the key is not found in the map
return nullptr; // Returns a null pointer
} }
auto &cache = cache_iter->second; auto &cache = cache_iter->second; // Obtains a reference to the value (a map) corresponding to the key in cache_iter
auto iter = cache.find(PrimitiveEvalCacheKey{prim->attrs(), args}); auto iter = cache.find(PrimitiveEvalCacheKey{prim->attrs(), args}); // Searches the cache map using a PrimitiveEvalCacheKey object created from prim's attributes and args
if (iter == cache.end()) { if (iter == cache.end()) { // If the key is not found in the map
return nullptr; return nullptr; // Returns a null pointer
} }
return iter->second;
return iter->second; // Returns the value (a shared pointer to an EvalResult object) corresponding to the key in iter
} }
void PrimitiveEvalCache::Put(const PrimitivePtr &prim, AttrValueMap &&attrs, const AbstractBasePtrList &args, void PrimitiveEvalCache::Put(const PrimitivePtr &prim, AttrValueMap &&attrs, const AbstractBasePtrList &args,
@ -106,41 +108,54 @@ void PrimitiveEvalCache::Clear() {
} }
AnalysisResult AnalysisEngine::Run(const FuncGraphPtr &func_graph, const AbstractBasePtrList &args_spec_list) { AnalysisResult AnalysisEngine::Run(const FuncGraphPtr &func_graph, const AbstractBasePtrList &args_spec_list) {
StaticAnalysisException::Instance().ClearException(); StaticAnalysisException::Instance().ClearException(); // Clears any previous exceptions in StaticAnalysisException
AnalysisResult result;
AnalysisResult result; // Creates an empty AnalysisResult object
try { try {
MS_EXCEPTION_IF_NULL(func_graph); MS_EXCEPTION_IF_NULL(func_graph); // Checks if func_graph is null and throws an exception if it is
ConfigPtrList args_conf_list;
ConfigPtrList args_conf_list; // Creates an empty list of ConfigPtr objects
// Transforms each element in args_spec_list into a ConfigPtr object using a lambda function and appends it to args_conf_list
(void)std::transform(args_spec_list.begin(), args_spec_list.end(), std::back_inserter(args_conf_list), (void)std::transform(args_spec_list.begin(), args_spec_list.end(), std::back_inserter(args_conf_list),
[](const AbstractBasePtr &arg) -> ConfigPtr { return std::make_shared<VirtualConfig>(arg); }); [](const AbstractBasePtr &arg) -> ConfigPtr { return std::make_shared<VirtualConfig>(arg); });
MS_EXCEPTION_IF_NULL(func_graph_manager_);
func_graph_manager_->AddFuncGraph(func_graph); MS_EXCEPTION_IF_NULL(func_graph_manager_); // Checks if func_graph_manager_ is null and throws an exception if it is
root_func_graph_ = func_graph; func_graph_manager_->AddFuncGraph(func_graph); // Adds func_graph to func_graph_manager_
root_func_graph_ = func_graph; // Sets root_func_graph_ to func_graph
// Running the analyzer. // Running the analyzer.
ResetFunctionCallDepth(); ResetFunctionCallDepth(); // Resets the function call depth counter
ResetStackFrameDepth(); ResetStackFrameDepth(); // Resets the stack frame depth counter
AnalysisContextPtr dummy_context = AnalysisContext::DummyContext();
AnalysisContextPtr root_context = Run(func_graph, dummy_context, args_conf_list); AnalysisContextPtr dummy_context = AnalysisContext::DummyContext(); // Creates a dummy AnalysisContext object
MS_EXCEPTION_IF_NULL(root_context); AnalysisContextPtr root_context = Run(func_graph, dummy_context, args_conf_list); // Runs the analysis with func_graph, dummy_context, and args_conf_list
auto root_context_fg = root_context->func_graph(); MS_EXCEPTION_IF_NULL(root_context); // Checks if root_context is null and throws an exception if it is
MS_EXCEPTION_IF_NULL(root_context_fg);
AnfNodeConfigPtr output_conf = MakeConfig(root_context_fg->get_return(), root_context, root_context_fg); auto root_context_fg = root_context->func_graph(); // Gets the function graph associated with root_context
MS_EXCEPTION_IF_NULL(func_graph); MS_EXCEPTION_IF_NULL(root_context_fg); // Checks if root_context_fg is null and throws an exception if it is
MS_LOG(INFO) << func_graph->ToString() << ": Run finished.";
AnfNodeConfigPtr output_conf = MakeConfig(root_context_fg->get_return(), root_context, root_context_fg); // Creates a config object for the return node of the function graph
MS_EXCEPTION_IF_NULL(func_graph); // Checks if func_graph is null and throws an exception if it is
MS_LOG(INFO) << func_graph->ToString() << ": Run finished."; // Logs an informational message
MS_EXCEPTION_IF_NULL(output_conf); // Checks if output_conf is null and throws an exception if it is
auto eval_result = output_conf->ObtainEvalResult(); // Obtains the evaluation result from output_conf
MS_EXCEPTION_IF_NULL(output_conf);
auto eval_result = output_conf->ObtainEvalResult();
// Set the sequence nodes' elements use flags all true. // Set the sequence nodes' elements use flags all true.
SetSequenceElementsUseFlagsRecursively(eval_result->abstract(), true); SetSequenceElementsUseFlagsRecursively(eval_result->abstract(), true); // Sets the use flags of sequence elements to true recursively
result.eval_result = eval_result;
result.context = root_context; result.eval_result = eval_result; // Sets the eval_result field of the AnalysisResult object
result.context = root_context; // Sets the context field of the AnalysisResult object
} catch (const std::exception &ex) { } catch (const std::exception &ex) {
MS_LOG(INFO) << "Eval " << func_graph->ToString() << " threw exception."; MS_LOG(INFO) << "Eval " << func_graph->ToString() << " threw exception."; // Logs an informational message
AnalysisSchedule::GetInstance().HandleException(ex); AnalysisSchedule::GetInstance().HandleException(ex); // Handles the exception in AnalysisSchedule
} }
AnalysisSchedule::GetInstance().Wait();
return result; AnalysisSchedule::GetInstance().Wait(); // Waits for analysis tasks to complete
return result; // Returns the AnalysisResult object
} }
AnalysisContextPtr AnalysisEngine::Run(const FuncGraphPtr &func_graph, const AnalysisContextPtr &context, AnalysisContextPtr AnalysisEngine::Run(const FuncGraphPtr &func_graph, const AnalysisContextPtr &context,
@ -151,14 +166,22 @@ AnalysisContextPtr AnalysisEngine::Run(const FuncGraphPtr &func_graph, const Ana
} }
void AnalysisEngine::SaveEvalResultInCache(const AnfNodeConfigPtr &conf, const EvalResultPtr &result) { void AnalysisEngine::SaveEvalResultInCache(const AnfNodeConfigPtr &conf, const EvalResultPtr &result) {
// Check that the pointers to AnfNodeConfig and EvalResult objects are not null
MS_EXCEPTION_IF_NULL(conf); MS_EXCEPTION_IF_NULL(conf);
MS_EXCEPTION_IF_NULL(result); MS_EXCEPTION_IF_NULL(result);
// Get an instance of AnalysisResultCacheMgr from the AnalysisResultCacheMgr singleton object
static AnalysisResultCacheMgr &cache_mgr = AnalysisResultCacheMgr::GetInstance(); static AnalysisResultCacheMgr &cache_mgr = AnalysisResultCacheMgr::GetInstance();
// Search for the given AnfNodeConfigPtr object in the cache
auto iter = cache_mgr.GetCache().find(conf); auto iter = cache_mgr.GetCache().find(conf);
// If the object is found in the cache, update the use flags of sequence elements in the cached evaluation result with the use flags in the new evaluation result, if enabled by the MS_DEV_ENABLE_DDE environment variable.
if (iter != cache_mgr.GetCache().end()) { if (iter != cache_mgr.GetCache().end()) {
MS_LOG(DEBUG) << "Found previous result for NodeConfig: " << conf->ToString() MS_LOG(DEBUG) << "Found previous result for NodeConfig: " << conf->ToString()
<< ", result: " << iter->second->abstract().get() << "/" << iter->second->abstract()->ToString(); << ", result: " << iter->second->abstract().get() << "/" << iter->second->abstract()->ToString();
// Update sequence nodes info, if matched in cache.
// If MS_DEV_ENABLE_DDE environment variable is enabled, update sequence nodes info
static const auto enable_eliminate_unused_element = (common::GetEnv("MS_DEV_ENABLE_DDE") != "0"); static const auto enable_eliminate_unused_element = (common::GetEnv("MS_DEV_ENABLE_DDE") != "0");
if (enable_eliminate_unused_element) { if (enable_eliminate_unused_element) {
auto new_sequence = dyn_cast<AbstractSequence>(result->abstract()); auto new_sequence = dyn_cast<AbstractSequence>(result->abstract());
@ -174,20 +197,34 @@ void AnalysisEngine::SaveEvalResultInCache(const AnfNodeConfigPtr &conf, const E
} }
} }
} }
// Log debug message indicating that the new evaluation result is being saved in the cache
MS_LOG(DEBUG) << "Save result for NodeConfig: " << conf->ToString() << ", result: " << result->abstract().get() << "/" MS_LOG(DEBUG) << "Save result for NodeConfig: " << conf->ToString() << ", result: " << result->abstract().get() << "/"
<< result->abstract()->ToString(); << result->abstract()->ToString();
// Save the new evaluation result in the cache using the SetValue() method of AnalysisResultCacheMgr
cache_mgr.SetValue(conf, result); cache_mgr.SetValue(conf, result);
} }
EvalResultPtr AnalysisEngine::ObtainEvalResultWithCache(const AnfNodeConfigPtr &conf) { EvalResultPtr AnalysisEngine::ObtainEvalResultWithCache(const AnfNodeConfigPtr &conf) {
// Check that the pointer to AnfNodeConfig object is not null
MS_EXCEPTION_IF_NULL(conf); MS_EXCEPTION_IF_NULL(conf);
// Get an instance of AnalysisResultCacheMgr from the AnalysisResultCacheMgr singleton object
static AnalysisResultCacheMgr &cache_mgr = AnalysisResultCacheMgr::GetInstance(); static AnalysisResultCacheMgr &cache_mgr = AnalysisResultCacheMgr::GetInstance();
// Search for the given AnfNodeConfigPtr object in the cache
auto result = cache_mgr.GetValue(conf); auto result = cache_mgr.GetValue(conf);
// If the object is found in the cache, return the cached evaluation result
if (result != nullptr) { if (result != nullptr) {
MS_LOG(DEBUG) << "Evaluate cache found for NodeConfig: " << conf->ToString() MS_LOG(DEBUG) << "Evaluate cache found for NodeConfig: " << conf->ToString()
<< ", result: " << result->abstract().get() << "/" << result->abstract()->ToString(); << ", result: " << result->abstract().get() << "/" << result->abstract()->ToString();
return result; return result;
} }
// If the object is not found in the cache, perform evaluation and save the result in the cache before returning it
MS_LOG(DEBUG) << "Evaluate cache miss for NodeConfig: " << conf->ToString(); MS_LOG(DEBUG) << "Evaluate cache miss for NodeConfig: " << conf->ToString();
result = Eval(conf); result = Eval(conf);
if (result == nullptr) { if (result == nullptr) {
@ -199,6 +236,7 @@ EvalResultPtr AnalysisEngine::ObtainEvalResultWithCache(const AnfNodeConfigPtr &
return result; return result;
} }
EvalResultPtr AnalysisEngine::ObtainEvalResultWithoutCache(const AnfNodeConfigPtr &conf) { EvalResultPtr AnalysisEngine::ObtainEvalResultWithoutCache(const AnfNodeConfigPtr &conf) {
MS_EXCEPTION_IF_NULL(conf); MS_EXCEPTION_IF_NULL(conf);
EvalResultPtr result = nullptr; EvalResultPtr result = nullptr;
@ -213,11 +251,20 @@ EvalResultPtr AnalysisEngine::ObtainEvalResultWithoutCache(const AnfNodeConfigPt
} }
EvalResultPtr AnalysisEngine::Eval(const AnfNodeConfigPtr &conf) { EvalResultPtr AnalysisEngine::Eval(const AnfNodeConfigPtr &conf) {
// Check that the pointer to AnfNodeConfig object is not null
MS_EXCEPTION_IF_NULL(conf); MS_EXCEPTION_IF_NULL(conf);
// Get the AnfNodePtr from the AnfNodeConfigPtr object
AnfNodePtr node = conf->node(); AnfNodePtr node = conf->node();
// Initialize the EvalResultPtr object as nullptr
EvalResultPtr eval_result = nullptr; EvalResultPtr eval_result = nullptr;
#ifdef DEBUG #ifdef DEBUG
// Push the current node onto the compute_conf_stack_ vector for debugging purposes
compute_conf_stack_.push_back(node); compute_conf_stack_.push_back(node);
// Build a string representation of the compute_conf_stack_ for debugging purposes
std::ostringstream buffer; std::ostringstream buffer;
buffer << "Compute Config Begin:"; buffer << "Compute Config Begin:";
for (auto iter : compute_conf_stack_) { for (auto iter : compute_conf_stack_) {
@ -225,21 +272,29 @@ EvalResultPtr AnalysisEngine::Eval(const AnfNodeConfigPtr &conf) {
} }
MS_LOG(DEBUG) << buffer.str(); MS_LOG(DEBUG) << buffer.str();
#endif #endif
MS_LOG(DEBUG) << "Begin Eval NodeConfig " << conf->ToString(); MS_LOG(DEBUG) << "Begin Eval NodeConfig " << conf->ToString();
MS_EXCEPTION_IF_NULL(node);
// If the node already has an abstract value, return it as the evaluation result
if (node->abstract() != nullptr) { if (node->abstract() != nullptr) {
MS_LOG(DEBUG) << "Return old abstract: " << node->DebugString(); MS_LOG(DEBUG) << "Return old abstract: " << node->DebugString();
eval_result = std::make_shared<EvalResult>(node->abstract(), std::make_shared<AttrValueMap>()); eval_result = std::make_shared<EvalResult>(node->abstract(), std::make_shared<AttrValueMap>());
} else if (node->isa<ValueNode>()) { }
// If the node is a ValueNode, evaluate its abstract value
else if (node->isa<ValueNode>()) {
auto value_node = node->cast<ValueNodePtr>(); auto value_node = node->cast<ValueNodePtr>();
auto abstract = EvalValueNode(value_node, conf); auto abstract = EvalValueNode(value_node, conf);
eval_result = std::make_shared<EvalResult>(abstract, std::make_shared<AttrValueMap>()); eval_result = std::make_shared<EvalResult>(abstract, std::make_shared<AttrValueMap>());
} else if (node->isa<CNode>()) { }
// If the node is a CNode, evaluate its abstract value
else if (node->isa<CNode>()) {
auto cnode = node->cast<CNodePtr>(); auto cnode = node->cast<CNodePtr>();
trace::TraceEvalCNodeEnter(conf); trace::TraceEvalCNodeEnter(conf);
eval_result = EvalCNode(cnode, conf); eval_result = EvalCNode(cnode, conf);
trace::TraceEvalCNodeLeave(); trace::TraceEvalCNodeLeave();
} else { }
// If the node type is not supported for evaluation, throw an exception
else {
MS_LOG(EXCEPTION) << "Illegal AnfNode for evaluating, node: " << node->DebugString() MS_LOG(EXCEPTION) << "Illegal AnfNode for evaluating, node: " << node->DebugString()
<< "(type:" << node->type_name() << "(type:" << node->type_name()
<< "), fg: " << (node->func_graph() != nullptr ? node->func_graph()->ToString() : "nullgraph") << "), fg: " << (node->func_graph() != nullptr ? node->func_graph()->ToString() : "nullgraph")
@ -247,13 +302,19 @@ EvalResultPtr AnalysisEngine::Eval(const AnfNodeConfigPtr &conf) {
} }
#ifdef DEBUG #ifdef DEBUG
// Pop the current node from the compute_conf_stack_ vector for debugging purposes
compute_conf_stack_.pop_back(); compute_conf_stack_.pop_back();
// If the evaluation result is still nullptr, throw an exception
if (eval_result == nullptr) { if (eval_result == nullptr) {
MS_LOG(EXCEPTION) << "Compute Config failed, node: " << node->DebugString() MS_LOG(EXCEPTION) << "Compute Config failed, node: " << node->DebugString()
<< " NodeInfo: " << trace::GetDebugInfo(node->debug_info()); << " NodeInfo: " << trace::GetDebugInfo(node->debug_info());
} }
#endif #endif
MS_LOG(DEBUG) << "End Eval NodeConfig " << conf->ToString() << ", res: " << eval_result->abstract()->ToString(); MS_LOG(DEBUG) << "End Eval NodeConfig " << conf->ToString() << ", res: " << eval_result->abstract()->ToString();
// Return the evaluation result
return eval_result; return eval_result;
} }
@ -269,25 +330,38 @@ AbstractBasePtr AnalysisEngine::EvalValueNode(const ValueNodePtr &value_node, co
AbstractBasePtr AnalysisEngine::GetCNodeOperatorAbstract(const CNodePtr &cnode, const AnalysisContextPtr &context, AbstractBasePtr AnalysisEngine::GetCNodeOperatorAbstract(const CNodePtr &cnode, const AnalysisContextPtr &context,
const FuncGraphPtr &func_graph) { const FuncGraphPtr &func_graph) {
// Check that the pointer to CNode object is not null
MS_EXCEPTION_IF_NULL(cnode); MS_EXCEPTION_IF_NULL(cnode);
// Get the inputs of the CNode
auto &inputs = cnode->inputs(); auto &inputs = cnode->inputs();
// Check that the inputs are not empty
if (inputs.empty()) { if (inputs.empty()) {
MS_LOG(EXCEPTION) << "CNode->inputs() is empty, CNode: " << cnode->DebugString(); MS_LOG(EXCEPTION) << "CNode->inputs() is empty, CNode: " << cnode->DebugString();
} }
// Get the function node from the inputs
AnfNodePtr func_node = inputs[0]; AnfNodePtr func_node = inputs[0];
// Check that the function node is not null
MS_EXCEPTION_IF_NULL(func_node); MS_EXCEPTION_IF_NULL(func_node);
MS_LOG(DEBUG) << "Current CNode function: " << func_node->DebugString(); MS_LOG(DEBUG) << "Current CNode function: " << func_node->DebugString();
// Create a AnfNodeConfigPtr object for the function node
AnfNodeConfigPtr func_conf = MakeConfig(func_node, context, func_graph); AnfNodeConfigPtr func_conf = MakeConfig(func_node, context, func_graph);
// Check that the pointer to AnfNodeConfig object is not null
MS_EXCEPTION_IF_NULL(func_conf); MS_EXCEPTION_IF_NULL(func_conf);
// Keep it in a local variable, otherwise smart pointer will free it. // Obtain the evaluation result for the function node
auto possible_func_eval_result = func_conf->ObtainEvalResult(); auto possible_func_eval_result = func_conf->ObtainEvalResult();
// Get the abstract value from the evaluation result
AbstractBasePtr possible_func = possible_func_eval_result->abstract(); AbstractBasePtr possible_func = possible_func_eval_result->abstract();
// Check that the abstract value is not null
if (possible_func == nullptr) { if (possible_func == nullptr) {
MS_LOG(EXCEPTION) << "No abstract, func_conf: " << func_conf->ToString(); MS_LOG(EXCEPTION) << "No abstract, func_conf: " << func_conf->ToString();
} }
// Return the abstract value of the function node
return possible_func; return possible_func;
} }
void CheckInterpretedObject(const AbstractBasePtr &abs) { void CheckInterpretedObject(const AbstractBasePtr &abs) {
static const auto support_fallback = common::GetEnv("MS_DEV_ENABLE_FALLBACK"); static const auto support_fallback = common::GetEnv("MS_DEV_ENABLE_FALLBACK");
static const auto use_fallback = (support_fallback != "0"); static const auto use_fallback = (support_fallback != "0");
@ -303,33 +377,49 @@ void CheckInterpretedObject(const AbstractBasePtr &abs) {
} }
EvalResultPtr AnalysisEngine::EvalCNode(const CNodePtr &cnode, const AnfNodeConfigPtr &conf) { EvalResultPtr AnalysisEngine::EvalCNode(const CNodePtr &cnode, const AnfNodeConfigPtr &conf) {
// Check that the pointers to CNode and AnfNodeConfig objects are not null
MS_EXCEPTION_IF_NULL(conf); MS_EXCEPTION_IF_NULL(conf);
MS_EXCEPTION_IF_NULL(cnode); MS_EXCEPTION_IF_NULL(cnode);
// Get the abstract value of the CNode's operator
AbstractBasePtr possible_func = GetCNodeOperatorAbstract(cnode, conf->context(), conf->func_graph()); AbstractBasePtr possible_func = GetCNodeOperatorAbstract(cnode, conf->context(), conf->func_graph());
// Check if the abstract value has undetermined type
if (possible_func->BuildType()->type_id() == kObjectTypeUndeterminedType) { if (possible_func->BuildType()->type_id() == kObjectTypeUndeterminedType) {
MS_LOG(DEBUG) << "EvalCNode eval Undetermined"; MS_LOG(DEBUG) << "EvalCNode eval Undetermined";
return std::make_shared<EvalResult>(possible_func->Clone(), std::make_shared<AttrValueMap>()); return std::make_shared<EvalResult>(possible_func->Clone(), std::make_shared<AttrValueMap>());
} }
// Check if the abstract value can be casted to AbstractFunction
AbstractFunctionPtr func = dyn_cast<AbstractFunction>(possible_func); AbstractFunctionPtr func = dyn_cast<AbstractFunction>(possible_func);
if (func == nullptr) { if (func == nullptr) {
// If not, log an error and throw an exception
CheckInterpretedObject(possible_func); CheckInterpretedObject(possible_func);
MS_LOG(ERROR) << "Can not cast to a AbstractFunction from " << possible_func->ToString() << "."; MS_LOG(ERROR) << "Can not cast to a AbstractFunction from " << possible_func->ToString() << ".";
MS_LOG(ERROR) << "It's called at: " << cnode->DebugString(); MS_LOG(ERROR) << "It's called at: " << cnode->DebugString();
MS_EXCEPTION(ValueError) << "This may be not defined, or it can't be a operator. Please check code."; MS_EXCEPTION(ValueError) << "This may be not defined, or it can't be a operator. Please check code.";
} }
// Create a ConfigPtrList to store the configurations of the arguments
ConfigPtrList args_conf_list; ConfigPtrList args_conf_list;
// Ignore the first node which is function name
// Iterate through the inputs of the CNode, ignoring the first node (function name)
auto &inputs = cnode->inputs(); auto &inputs = cnode->inputs();
for (std::size_t i = 1; i < inputs.size(); i++) { for (std::size_t i = 1; i < inputs.size(); i++) {
const AnfNodePtr &node = inputs[i]; const AnfNodePtr &node = inputs[i];
// Create an AnfNodeConfigPtr object for the argument node
args_conf_list.push_back(MakeConfig(node, conf->context(), conf->func_graph())); args_conf_list.push_back(MakeConfig(node, conf->context(), conf->func_graph()));
} }
// Create a vector to store the evaluators
std::vector<EvaluatorPtr> evaluators; std::vector<EvaluatorPtr> evaluators;
// Define a lambda function to build evaluators for each resolved AtomicAbstractFunc
auto build_evaluator = [this, &evaluators, &cnode](const AbstractFuncAtomPtr &poss) { auto build_evaluator = [this, &evaluators, &cnode](const AbstractFuncAtomPtr &poss) {
auto resolved_atom = poss; auto resolved_atom = poss;
// If the resolved AtomicAbstractFunc is an AsyncAbstractFunc, resolve it to get the actual function
if (poss->isa<AsyncAbstractFuncAtom>()) { if (poss->isa<AsyncAbstractFuncAtom>()) {
const auto &async_abs_func = poss->cast<AsyncAbstractFuncAtomPtr>(); const auto &async_abs_func = poss->cast<AsyncAbstractFuncAtomPtr>();
const auto &resolved_func = async_abs_func->GetUnique(); const auto &resolved_func = async_abs_func->GetUnique();
@ -337,51 +427,73 @@ EvalResultPtr AnalysisEngine::EvalCNode(const CNodePtr &cnode, const AnfNodeConf
MS_EXCEPTION_IF_NULL(resolved_atom); MS_EXCEPTION_IF_NULL(resolved_atom);
MS_LOG(DEBUG) << "Resolved AsyncAbstractFuncAtom is: " << resolved_atom->ToString(); MS_LOG(DEBUG) << "Resolved AsyncAbstractFuncAtom is: " << resolved_atom->ToString();
} }
// Get an evaluator for the resolved AtomicAbstractFunc
auto evaluator = this->GetEvaluatorFor(resolved_atom); auto evaluator = this->GetEvaluatorFor(resolved_atom);
// Set the bound node of the evaluator to the current CNode
evaluator->set_bound_node(cnode); evaluator->set_bound_node(cnode);
// Add the evaluator to the vector
evaluators.push_back(evaluator); evaluators.push_back(evaluator);
}; };
// Visit the AbstractFunction to build evaluators
func->Visit(build_evaluator); func->Visit(build_evaluator);
// Execute the evaluators with the given configurations and return the evaluation result
auto eval_result = ExecuteEvaluators(evaluators, conf, args_conf_list); auto eval_result = ExecuteEvaluators(evaluators, conf, args_conf_list);
return eval_result; return eval_result;
} }
EvalResultPtr AnalysisEngine::Execute(const AbstractFunctionPtr &func, const AbstractBasePtrList &args_spec_list) { EvalResultPtr AnalysisEngine::Execute(const AbstractFunctionPtr &func, const AbstractBasePtrList &args_spec_list) {
// Check that the AbstractFunction pointer is not null
MS_EXCEPTION_IF_NULL(func); MS_EXCEPTION_IF_NULL(func);
// Create a ConfigPtrList to store the configurations of the arguments
ConfigPtrList args_conf_list; ConfigPtrList args_conf_list;
// Transform the input argument list into a list of VirtualConfigs
(void)std::transform(args_spec_list.begin(), args_spec_list.end(), std::back_inserter(args_conf_list), (void)std::transform(args_spec_list.begin(), args_spec_list.end(), std::back_inserter(args_conf_list),
[](const AbstractBasePtr &arg) -> ConfigPtr { return std::make_shared<VirtualConfig>(arg); }); [](const AbstractBasePtr &arg) -> ConfigPtr { return std::make_shared<VirtualConfig>(arg); });
// Create a vector to store the evaluators
std::vector<EvaluatorPtr> infs; std::vector<EvaluatorPtr> infs;
MS_EXCEPTION_IF_NULL(func); // Define a lambda function to build evaluators for each resolved AtomicAbstractFunc
auto build_evaluator = [this, &infs](const AbstractFuncAtomPtr &poss) { auto build_evaluator = [this, &infs](const AbstractFuncAtomPtr &poss) {
auto evaluator = this->GetEvaluatorFor(poss); auto evaluator = this->GetEvaluatorFor(poss);
infs.push_back(evaluator); infs.push_back(evaluator);
}; };
// Visit the AbstractFunction to build evaluators
func->Visit(build_evaluator); func->Visit(build_evaluator);
// Execute the evaluators with the given configurations and return the evaluation result
return ExecuteEvaluators(infs, nullptr, args_conf_list); return ExecuteEvaluators(infs, nullptr, args_conf_list);
} }
void AnalysisEngine::ClearEvaluatorCache() { void AnalysisEngine::ClearEvaluatorCache() {
// Clear cache for evaluators in evaluators_ map
for (auto &element : evaluators_) { for (auto &element : evaluators_) {
EvaluatorPtr evaluator = element.second; EvaluatorPtr evaluator = element.second;
MS_EXCEPTION_IF_NULL(evaluator); MS_EXCEPTION_IF_NULL(evaluator);
MS_EXCEPTION_IF_NULL(evaluator->evaluator_cache_mgr()); MS_EXCEPTION_IF_NULL(evaluator->evaluator_cache_mgr());
evaluator->evaluator_cache_mgr()->Clear(); evaluator->evaluator_cache_mgr()->Clear();
} }
// Clear cache for evaluators in prim_constructors_ map
for (auto &element : prim_constructors_) { for (auto &element : prim_constructors_) {
EvaluatorPtr evaluator = element.second; EvaluatorPtr evaluator = element.second;
MS_EXCEPTION_IF_NULL(evaluator); MS_EXCEPTION_IF_NULL(evaluator);
MS_EXCEPTION_IF_NULL(evaluator->evaluator_cache_mgr()); MS_EXCEPTION_IF_NULL(evaluator->evaluator_cache_mgr());
evaluator->evaluator_cache_mgr()->Clear(); evaluator->evaluator_cache_mgr()->Clear();
} }
// Clear cache for evaluators in prim_py_evaluators_ map
for (auto &element : prim_py_evaluators_) { for (auto &element : prim_py_evaluators_) {
EvaluatorPtr evaluator = element.second; EvaluatorPtr evaluator = element.second;
MS_EXCEPTION_IF_NULL(evaluator); MS_EXCEPTION_IF_NULL(evaluator);
MS_EXCEPTION_IF_NULL(evaluator->evaluator_cache_mgr()); MS_EXCEPTION_IF_NULL(evaluator->evaluator_cache_mgr());
evaluator->evaluator_cache_mgr()->Clear(); evaluator->evaluator_cache_mgr()->Clear();
} }
// Release Exception to avoid hup at exit.
// Clear exceptions in the StaticAnalysisException singleton
StaticAnalysisException::Instance().ClearException(); StaticAnalysisException::Instance().ClearException();
} }
@ -399,38 +511,39 @@ void AnalysisEngine::Clear() {
EvaluatorPtr GetPrimEvaluator(const PrimitivePtr &prim, const AnalysisEnginePtr &engine) { EvaluatorPtr GetPrimEvaluator(const PrimitivePtr &prim, const AnalysisEnginePtr &engine) {
// Custom Primitive with python infer_shape, infer_type // Custom Primitive with python infer_shape, infer_type
MS_EXCEPTION_IF_NULL(prim); MS_EXCEPTION_IF_NULL(prim);
if (prim->isa<prim::DoSignaturePrimitive>()) { if (prim->isa<prim::DoSignaturePrimitive>()) { // Check if it is a custom DoSignaturePrimitive
return std::make_shared<DoSignatureEvaluator>(prim); return std::make_shared<DoSignatureEvaluator>(prim); // Return an instance of DoSignatureEvaluator
} }
if (prim->isa<prim::UnpackGraphPrimitive>()) { if (prim->isa<prim::UnpackGraphPrimitive>()) { // Check if it is a custom UnpackGraphPrimitive
return std::make_shared<UnpackGraphEvaluator>(prim); return std::make_shared<UnpackGraphEvaluator>(prim); // Return an instance of UnpackGraphEvaluator
} }
if (prim->Hash() == prim::kPrimMixedPrecisionCast->Hash() && prim->name() == prim::kPrimMixedPrecisionCast->name()) { if (prim->Hash() == prim::kPrimMixedPrecisionCast->Hash() && prim->name() == prim::kPrimMixedPrecisionCast->name()) {
return std::make_shared<MixedPrecisionCastEvaluator>(prim); // Check if it is a mixed precision cast operation
return std::make_shared<MixedPrecisionCastEvaluator>(prim); // Return an instance of MixedPrecisionCastEvaluator
} }
// Find prim infer function in the prim function map return a standard evaluator // Find prim infer function in the prim function map and return a standard evaluator
auto eval_impl = GetPrimitiveInferImpl(prim); auto eval_impl = GetPrimitiveInferImpl(prim); // Get the infer implementation of the prim from the prim function map
if (eval_impl.infer_shape_impl_ != nullptr && prim->name() != prim::kPrimMakeTuple->name() && if (eval_impl.infer_shape_impl_ != nullptr && prim->name() != prim::kPrimMakeTuple->name() &&
prim->name() != prim::kPrimMakeList->name()) { // Refactoring infer routine soon. prim->name() != prim::kPrimMakeList->name()) { // Check if the infer implementation exists and it is not MakeTuple or MakeList
return std::make_shared<StandardPrimEvaluator>(prim, eval_impl); return std::make_shared<StandardPrimEvaluator>(prim, eval_impl); // Return an instance of StandardPrimEvaluator
} }
// Use python infer function if the infer function not founded in the map return a python evaluator // Use python infer function if the infer function is not found in the map, and return a python evaluator
EvaluatorPtr evaluator = nullptr; EvaluatorPtr evaluator = nullptr;
if (prim->HasPyEvaluator()) { if (prim->HasPyEvaluator()) { // Check if it has a Python infer function
auto prim_py = dyn_cast<PrimitivePy>(prim); auto prim_py = dyn_cast<PrimitivePy>(prim);
if (prim_py != nullptr) { if (prim_py != nullptr) {
if (engine == nullptr) { if (engine == nullptr) { // Check if the analysis engine is provided
return std::make_shared<PythonPrimEvaluator>(prim_py); return std::make_shared<PythonPrimEvaluator>(prim_py); // Return an instance of PythonPrimEvaluator
} }
const auto &iter = engine->prim_py_evaluators_.find(prim_py); const auto &iter = engine->prim_py_evaluators_.find(prim_py); // Find the cached PythonPrimEvaluator in the engine
if (iter != engine->prim_py_evaluators_.end()) { if (iter != engine->prim_py_evaluators_.end()) {
return iter->second; return iter->second; // If already cached, return the cached PythonPrimEvaluator
} }
evaluator = std::make_shared<PythonPrimEvaluator>(prim_py); evaluator = std::make_shared<PythonPrimEvaluator>(prim_py); // Create a new PythonPrimEvaluator
engine->prim_py_evaluators_[prim_py] = evaluator; engine->prim_py_evaluators_[prim_py] = evaluator; // Cache the new PythonPrimEvaluator in the engine
return evaluator; return evaluator;
} }
MS_LOG(ERROR) << "The primitive with python evaluator should be a python primitive."; MS_LOG(ERROR) << "The primitive with python evaluator should be a python primitive.";
@ -438,25 +551,25 @@ EvaluatorPtr GetPrimEvaluator(const PrimitivePtr &prim, const AnalysisEnginePtr
} }
// Return a default evaluator // Return a default evaluator
if (engine == nullptr) { if (engine == nullptr) { // Check if the analysis engine is provided
// If engine is nullptr, get constructor from default. // If engine is nullptr, get constructor from default.
const PrimEvaluatorMap &prim_evaluator_map = GetPrimEvaluatorConstructors(); const PrimEvaluatorMap &prim_evaluator_map = GetPrimEvaluatorConstructors(); // Get the constructor from the default PrimEvaluatorMap
auto iter = prim_evaluator_map.find(prim); auto iter = prim_evaluator_map.find(prim); // Find the constructor that matches the prim
if (iter != prim_evaluator_map.end()) { if (iter != prim_evaluator_map.end()) {
evaluator = iter->second; evaluator = iter->second; // Get the evaluator instance from the constructor
} }
} else { } else {
// If engine is given, get constructor from engine resource. // If engine is given, get constructor from engine resource.
const PrimEvaluatorMap &prim_evaluator_map = engine->PrimConstructors(); const PrimEvaluatorMap &prim_evaluator_map = engine->PrimConstructors(); // Get the PrimEvaluatorMap from the engine
auto iter = prim_evaluator_map.find(prim); auto iter = prim_evaluator_map.find(prim); // Find the constructor that matches the prim
if (iter != prim_evaluator_map.end()) { if (iter != prim_evaluator_map.end()) {
evaluator = iter->second; evaluator = iter->second; // Get the evaluator instance from the constructor
} }
} }
if (evaluator == nullptr) { if (evaluator == nullptr) {
MS_LOG(DEBUG) << "The evaluator of the primitive is not defined (" << prim->name() << ")."; MS_LOG(DEBUG) << "The evaluator of the primitive is not defined (" << prim->name() << ").";
} }
return evaluator; return evaluator; // Return the obtained evaluator instance, which can be nullptr
} }
EvaluatorPtr AnalysisEngine::_GetEvaluatorFor(const std::shared_ptr<PrimitiveAbstractClosure> &func) { EvaluatorPtr AnalysisEngine::_GetEvaluatorFor(const std::shared_ptr<PrimitiveAbstractClosure> &func) {
@ -541,20 +654,32 @@ EvaluatorPtr AnalysisEngine::_GetEvaluatorFor(const std::shared_ptr<VirtualAbstr
} }
EvaluatorPtr AnalysisEngine::_GetEvaluatorFor(const std::shared_ptr<PartialAbstractClosure> &func) { EvaluatorPtr AnalysisEngine::_GetEvaluatorFor(const std::shared_ptr<PartialAbstractClosure> &func) {
// Check if the input argument is not null
MS_EXCEPTION_IF_NULL(func); MS_EXCEPTION_IF_NULL(func);
// Get the original function from the partial closure
AbstractFunctionPtr func_orig = func->fn(); AbstractFunctionPtr func_orig = func->fn();
// Get the evaluator for the original function
EvaluatorPtr evaluator_orig = GetEvaluatorFor(func_orig); EvaluatorPtr evaluator_orig = GetEvaluatorFor(func_orig);
// Create a pair of the original function and its arguments
auto part_pair = std::make_pair(func_orig, func->args()); auto part_pair = std::make_pair(func_orig, func->args());
// Check if an evaluator for the partial closure already exists in the map
auto itr = constructors_app_.find(part_pair); auto itr = constructors_app_.find(part_pair);
if (itr != constructors_app_.end()) { if (itr != constructors_app_.end()) {
// Return the existing evaluator
return itr->second; return itr->second;
} }
// If an evaluator doesn't exist, create a new PartialAppEvaluator
// Pass the original evaluator and the arguments of the partial closure to the constructor
std::shared_ptr<PartialAppEvaluator> partial_evaluator = std::shared_ptr<PartialAppEvaluator> partial_evaluator =
std::make_shared<PartialAppEvaluator>(evaluator_orig, func->args()); std::make_shared<PartialAppEvaluator>(evaluator_orig, func->args());
// Cache the newly created PartialAppEvaluator in the map
constructors_app_[part_pair] = partial_evaluator; constructors_app_[part_pair] = partial_evaluator;
// Return the created PartialAppEvaluator
return partial_evaluator; return partial_evaluator;
} }
EvaluatorPtr AnalysisEngine::_GetEvaluatorFor(const std::shared_ptr<TypedPrimitiveAbstractClosure> &) { EvaluatorPtr AnalysisEngine::_GetEvaluatorFor(const std::shared_ptr<TypedPrimitiveAbstractClosure> &) {
MS_LOG(EXCEPTION) << "Should not be called "; MS_LOG(EXCEPTION) << "Should not be called ";
} }
@ -659,24 +784,29 @@ EvalResultPtr AnalysisEngine::ExecuteEvaluators(const std::vector<EvaluatorPtr>
void AnalysisEngine::SetUndeterminedFlag(const EvaluatorPtr &evaluator, const FuncGraphPtr &possible_parent_fg) { void AnalysisEngine::SetUndeterminedFlag(const EvaluatorPtr &evaluator, const FuncGraphPtr &possible_parent_fg) {
MS_EXCEPTION_IF_NULL(evaluator); MS_EXCEPTION_IF_NULL(evaluator);
static std::mutex fg_lock;
std::lock_guard<std::mutex> infer_lock(fg_lock); static std::mutex fg_lock; // A static mutex to ensure thread safety for modifying the func graph
if (possible_parent_fg != nullptr) { std::lock_guard<std::mutex> infer_lock(fg_lock); // Acquire the lock
possible_parent_fg->set_flag(kFuncGraphFlagUndetermined, true);
if (possible_parent_fg != nullptr) { // If a parent func graph is provided...
possible_parent_fg->set_flag(kFuncGraphFlagUndetermined, true); // Set the undetermined flag of the parent func graph
MS_LOG(DEBUG) << "Set graph undetermined: " << possible_parent_fg->ToString(); MS_LOG(DEBUG) << "Set graph undetermined: " << possible_parent_fg->ToString();
} }
auto fg_eval = evaluator->cast<FuncGraphEvaluatorPtr>();
if (fg_eval == nullptr) { auto fg_eval = evaluator->cast<FuncGraphEvaluatorPtr>(); // Cast the evaluator to FuncGraphEvaluatorPtr
if (fg_eval == nullptr) { // If it doesn't cast to FuncGraphEvaluatorPtr, simply return
return; return;
} }
auto fg = fg_eval->func_graph();
auto fg = fg_eval->func_graph(); // Get the func graph from the FuncGraphEvaluatorPtr
MS_EXCEPTION_IF_NULL(fg); MS_EXCEPTION_IF_NULL(fg);
auto fg_parent = fg->parent();
if (fg_parent != nullptr) { auto fg_parent = fg->parent(); // Get the parent func graph of the current func graph
fg_parent->set_flag(kFuncGraphFlagUndetermined, true); if (fg_parent != nullptr) { // If the parent func graph exists...
fg_parent->set_flag(kFuncGraphFlagUndetermined, true); // Set the undetermined flag of the parent func graph
MS_LOG(DEBUG) << "Set graph undetermined: " << fg_parent->ToString() << " for fg: " << fg->ToString(); MS_LOG(DEBUG) << "Set graph undetermined: " << fg_parent->ToString() << " for fg: " << fg->ToString();
return; return;
} else { } else { // If the parent func graph doesn't exist...
MS_LOG(DEBUG) << "cannot find parent for fg: " << fg->ToString(); MS_LOG(DEBUG) << "cannot find parent for fg: " << fg->ToString();
} }
} }
@ -738,6 +868,7 @@ EvaluatorPtr AnalysisEngine::HandleNestedRecursion(const std::vector<EvaluatorPt
std::string JoinBranchesFailedInfo(const AbstractBasePtr &spec, const AbstractBasePtr &last_spec, std::string JoinBranchesFailedInfo(const AbstractBasePtr &spec, const AbstractBasePtr &last_spec,
const AnfNodePtr &node, const std::string &error_info) { const AnfNodePtr &node, const std::string &error_info) {
constexpr int recursive_level = 2; constexpr int recursive_level = 2;
// Use a stringstream to build the error message string.
std::ostringstream buffer; std::ostringstream buffer;
buffer << "Cannot join the return values of different branches, perhaps you need to make them equal.\n" buffer << "Cannot join the return values of different branches, perhaps you need to make them equal.\n"
<< error_info << "\nThe abstract type of the return value of the current branch is " << spec->ToString() << error_info << "\nThe abstract type of the return value of the current branch is " << spec->ToString()
@ -745,6 +876,7 @@ std::string JoinBranchesFailedInfo(const AbstractBasePtr &spec, const AbstractBa
<< "The node is " << node->DebugString(recursive_level); << "The node is " << node->DebugString(recursive_level);
if (node->isa<CNode>()) { if (node->isa<CNode>()) {
auto cnode = node->cast<CNodePtr>()->input(0); auto cnode = node->cast<CNodePtr>()->input(0);
// If the current node is a Switch node, output the information of the True branch and False branch.
if (IsPrimitiveCNode(cnode, prim::kPrimSwitch)) { if (IsPrimitiveCNode(cnode, prim::kPrimSwitch)) {
// {prim::kPrimSwitch, cond, true_branch, false_branch} // {prim::kPrimSwitch, cond, true_branch, false_branch}
constexpr int true_index = 2; constexpr int true_index = 2;
@ -752,7 +884,9 @@ std::string JoinBranchesFailedInfo(const AbstractBasePtr &spec, const AbstractBa
auto inputs = cnode->cast<CNodePtr>()->inputs(); auto inputs = cnode->cast<CNodePtr>()->inputs();
buffer << ", true branch: " << inputs.at(true_index)->ToString() buffer << ", true branch: " << inputs.at(true_index)->ToString()
<< ", false branch: " << inputs.at(false_index)->ToString(); << ", false branch: " << inputs.at(false_index)->ToString();
} else if (IsPrimitiveCNode(cnode, prim::kPrimSwitchLayer)) { }
// If the current node is a SwitchLayer node, output the information of each branch.
else if (IsPrimitiveCNode(cnode, prim::kPrimSwitchLayer)) {
// {prim::kPrimSwitchLayer, X, {prim::kPrimMakeTuple, branch1, branch2, ...}} // {prim::kPrimSwitchLayer, X, {prim::kPrimMakeTuple, branch1, branch2, ...}}
constexpr int branch_index = 2; constexpr int branch_index = 2;
auto tuple_node = cnode->cast<CNodePtr>()->input(branch_index); auto tuple_node = cnode->cast<CNodePtr>()->input(branch_index);
@ -764,10 +898,12 @@ std::string JoinBranchesFailedInfo(const AbstractBasePtr &spec, const AbstractBa
} }
} }
} }
// Output the source code location of the current node.
buffer << trace::DumpSourceLines(node); buffer << trace::DumpSourceLines(node);
return buffer.str(); return buffer.str();
} }
EvalResultPtr AnalysisEngine::ProcessEvalResults(const AbstractBasePtrList &out_specs, const AnfNodePtr &node) { EvalResultPtr AnalysisEngine::ProcessEvalResults(const AbstractBasePtrList &out_specs, const AnfNodePtr &node) {
if (out_specs.empty()) { if (out_specs.empty()) {
MS_LOG(EXCEPTION) << "There is an endless loop for evaluator."; MS_LOG(EXCEPTION) << "There is an endless loop for evaluator.";

View File

@ -75,6 +75,10 @@ void ValidateOperation(const AnfNodePtr &node) {
MS_LOG(EXCEPTION) << "Illegal primitive: " << prim->name(); MS_LOG(EXCEPTION) << "Illegal primitive: " << prim->name();
} }
//The function of this code is to verify whether a node operation is legal,
//mainly by judging whether the Primitive corresponding to the node is in the whitelist,
//whether there are specific attributes or methods to judge its legitimacy,
//if the node operation is illegal, an exception will be thrown.
bool CheckAbstractScalar(const AnfNodePtr &node) { bool CheckAbstractScalar(const AnfNodePtr &node) {
MS_EXCEPTION_IF_NULL(node); MS_EXCEPTION_IF_NULL(node);
@ -96,6 +100,12 @@ bool CheckAbstractScalar(const AnfNodePtr &node) {
} }
return false; return false;
} }
//What this code does is check whether the abstract value of a node is a scalar type.
//Returns false if the abstract value is not of type AbstractScalar;
// If it is an AbstractScalar type,
//it further checks whether the type of the abstract value is legal,
//and if not, an exception is thrown;
//Returning true if legal indicates that the abstract value is a scalar type.
void ValidateAbstract(const AnfNodePtr &node) { void ValidateAbstract(const AnfNodePtr &node) {
if (node == nullptr) { if (node == nullptr) {
@ -132,6 +142,15 @@ void ValidateAbstract(const AnfNodePtr &node) {
// Other types show exception // Other types show exception
MS_LOG(EXCEPTION) << "Illegal type in the graph: " << abstract->ToString(); MS_LOG(EXCEPTION) << "Illegal type in the graph: " << abstract->ToString();
} }
//The purpose of this code is to verify that the abstract value of a node is valid.
// First check whether the node and abstract value are empty,
//and then verify whether the type of the abstract value is AbstractClass type or AbstractJTagged type, respectively,
//if so, throw an exception;
//Then call the CheckAbstractScalar function to verify whether the abstract value is a scalar type,
//and if so, return;
//Then determine whether the abstract value is of type AbstractError,
//and if so, print the debug log;
//Finally, determine whether the abstract value type is legal, and if so, return.
void ValidateValueNode(const AnfNodePtr &node) { void ValidateValueNode(const AnfNodePtr &node) {
if (node == nullptr) { if (node == nullptr) {
@ -147,6 +166,11 @@ void ValidateValueNode(const AnfNodePtr &node) {
<< "https://www.mindspore.cn/search?inputValue=JIT%20Fallback"; << "https://www.mindspore.cn/search?inputValue=JIT%20Fallback";
} }
} }
//The purpose of this code is to verify the validity of a value node.
//First check whether the node is empty,
//then determine whether the node is a value node of type parse,
//and if so, throw an exception.
//The purpose of this validation function is to ensure that Python objects are not used at runtime
void CheckValueTuple(const AnfNodePtr &node) { void CheckValueTuple(const AnfNodePtr &node) {
const auto &value_node = node->cast<ValueNodePtr>(); const auto &value_node = node->cast<ValueNodePtr>();
@ -162,6 +186,11 @@ void CheckValueTuple(const AnfNodePtr &node) {
ValidateValueNode(input_node); ValidateValueNode(input_node);
} }
} }
//The purpose of this code is to check
// whether a value node is a tuple type and validate each value node in the tuple.
//It first gets the value object of the value node and converts it to a tuple type.
//Each value node in the tuple is then looped through and operational and value node validation is performed
void Validate(const FuncGraphPtr &fg) { void Validate(const FuncGraphPtr &fg) {
FuncGraphManagerPtr mgr = Manage(fg, false); FuncGraphManagerPtr mgr = Manage(fg, false);

View File

@ -84,6 +84,7 @@ std::map<std::string, std::shared_ptr<session::SessionBasic>> kSessionBackends;
std::map<std::string, std::shared_ptr<compile::MindRTBackend>> kMindRtBackends; std::map<std::string, std::shared_ptr<compile::MindRTBackend>> kMindRtBackends;
PyObjectIdCache g_pyobj_id_cache; PyObjectIdCache g_pyobj_id_cache;
// General exception handling function that executes a method and handles exceptions.
template <typename T, typename... Args> template <typename T, typename... Args>
void PynativeExecutorTry(const std::function<void(T *ret, const Args &...)> &method, T *ret, const Args &... args) { void PynativeExecutorTry(const std::function<void(T *ret, const Args &...)> &method, T *ret, const Args &... args) {
const auto inst = PynativeExecutor::GetInstance(); const auto inst = PynativeExecutor::GetInstance();
@ -128,6 +129,7 @@ void PynativeExecutorTry(const std::function<void(T *ret, const Args &...)> &met
} }
} }
// Convert a py::object to a pointer to Value.
inline ValuePtr PyObjToValue(const py::object &obj) { inline ValuePtr PyObjToValue(const py::object &obj) {
ValuePtr converted_ret = parse::data_converter::PyDataToValue(obj); ValuePtr converted_ret = parse::data_converter::PyDataToValue(obj);
if (!converted_ret) { if (!converted_ret) {
@ -144,6 +146,14 @@ std::string GetPyObjId(const py::handle &obj) {
return out.cast<std::string>(); return out.cast<std::string>();
} }
// Get the identifier of the given Python object.
// If obj is Tensor and is not Parameter, then return id.
// If obj is Parameter, then return name.
// If obj is mindspore::Type, then return "type" + ToString(obj).
// If obj is str or int_ or float_, then return string(obj).
// If obj is None, then return "none".
// If obj is tuple or list, then return "tuple"/"list" + "empty"/str(obj[0]):str(obj[1]):...
// If obj is Cell of function, then return GetPyObjId(obj).
std::string GetId(const py::handle &obj) { std::string GetId(const py::handle &obj) {
if (py::isinstance<tensor::Tensor>(obj)) { if (py::isinstance<tensor::Tensor>(obj)) {
auto tensor_ptr = py::cast<tensor::TensorPtr>(obj); auto tensor_ptr = py::cast<tensor::TensorPtr>(obj);
@ -197,6 +207,7 @@ bool IsFunctionType(const py::object &cell) {
return false; return false;
} }
// Find all indexs of types from type_indexes.
void GetTypeIndex(const std::vector<SignatureEnumDType> &dtypes, void GetTypeIndex(const std::vector<SignatureEnumDType> &dtypes,
mindspore::HashMap<SignatureEnumDType, std::vector<size_t>> *type_indexes) { mindspore::HashMap<SignatureEnumDType, std::vector<size_t>> *type_indexes) {
MS_EXCEPTION_IF_NULL(type_indexes); MS_EXCEPTION_IF_NULL(type_indexes);
@ -763,6 +774,7 @@ void RunReplace(const CNodePtr &added_make_tuple, const std::vector<tensor::Tens
} }
} }
// Replace the new output tensor in the gradient graph so that the tensors of the forward and backward nodes can correspond.
void ReplaceNewTensorsInGradGraph(const TopCellInfoPtr &top_cell, const OpExecInfoPtr &op_exec_info, void ReplaceNewTensorsInGradGraph(const TopCellInfoPtr &top_cell, const OpExecInfoPtr &op_exec_info,
const ValuePtr &added_out, const FuncGraphPtr &ms_func_graph, const ValuePtr &added_out, const FuncGraphPtr &ms_func_graph,
const FuncGraphPtr &grad_graph) { const FuncGraphPtr &grad_graph) {
@ -814,6 +826,7 @@ void SaveOpInfo(const TopCellInfoPtr &top_cell, const std::string &op_info,
}); });
} }
// Update the new tensor information to the pre tensor and make inter device memory data copies as needed.
void UpdateTensorInfo(const tensor::TensorPtr &new_tensor, const std::vector<tensor::TensorPtr> &pre_tensors) { void UpdateTensorInfo(const tensor::TensorPtr &new_tensor, const std::vector<tensor::TensorPtr> &pre_tensors) {
MS_EXCEPTION_IF_NULL(new_tensor); MS_EXCEPTION_IF_NULL(new_tensor);
if (pre_tensors.empty() || new_tensor->device_address() == nullptr) { if (pre_tensors.empty() || new_tensor->device_address() == nullptr) {
@ -938,6 +951,7 @@ ValuePtr ShallowCopyValue(const OpExecInfoPtr &op_exec_info, const ValuePtr &val
} }
} // namespace } // namespace
// The true operation of the operator is in this function.
py::object RealRunOp(const py::args &args) { py::object RealRunOp(const py::args &args) {
CheckPyNativeContext(); CheckPyNativeContext();
const auto &executor = PynativeExecutor::GetInstance(); const auto &executor = PynativeExecutor::GetInstance();
@ -1303,6 +1317,7 @@ void ForwardExecutor::DoNopOutput(const OpExecInfoPtr &op_exec_info, ValuePtr *o
MS_LOG(DEBUG) << "New copy value is " << (*out_real_value)->ToString(); MS_LOG(DEBUG) << "New copy value is " << (*out_real_value)->ToString();
} }
// Get output of operator in forward graph.
void ForwardExecutor::GetOpOutput(const OpExecInfoPtr &op_exec_info, void ForwardExecutor::GetOpOutput(const OpExecInfoPtr &op_exec_info,
const abstract::AbstractBasePtrList &args_spec_list, const CNodePtr &cnode, const abstract::AbstractBasePtrList &args_spec_list, const CNodePtr &cnode,
bool prim_cache_hit, py::object *ret) { bool prim_cache_hit, py::object *ret) {
@ -1893,6 +1908,7 @@ void GradExecutor::DoOpGrad(const OpExecInfoPtr &op_exec_info, const CNodePtr &c
} }
} }
// Update tensors in forward graph created by ms_function.
void GradExecutor::UpdateMsFunctionForwardTensors(const OpExecInfoPtr &op_exec_info, void GradExecutor::UpdateMsFunctionForwardTensors(const OpExecInfoPtr &op_exec_info,
const ValuePtr &new_forward_value) { const ValuePtr &new_forward_value) {
MS_LOG(DEBUG) << "Ms func graph has already ran before. The graph phase is: " << graph_phase(); MS_LOG(DEBUG) << "Ms func graph has already ran before. The graph phase is: " << graph_phase();
@ -2000,6 +2016,10 @@ void GradExecutor::MakeAdjointForMsFunction(const FuncGraphPtr &ms_func_graph, c
top_cell()->set_ms_function_flag(true); top_cell()->set_ms_function_flag(true);
} }
// Update forward tensor info in backprop graph.
// If you need to construct a graph, use the SaveOpInfo function to save all tensor information for the current operation.
// Its implementation is determined by the need construct graph function, which returns whether the graph has already been constructed before, and if it has,
// returns false and does not need to save the tensor information again.
void GradExecutor::UpdateForwardTensorInfoInBpropGraph(const OpExecInfoPtr &op_exec_info, const ValuePtr &op_out) { void GradExecutor::UpdateForwardTensorInfoInBpropGraph(const OpExecInfoPtr &op_exec_info, const ValuePtr &op_out) {
if (!grad_flag_) { if (!grad_flag_) {
MS_LOG(DEBUG) << "The grad flag is false, no need to update forward op info in bprop graph"; MS_LOG(DEBUG) << "The grad flag is false, no need to update forward op info in bprop graph";
@ -2126,6 +2146,7 @@ MsBackendPolicy ForwardExecutor::GetBackendPolicy(const OpExecInfoPtr &op_exec_i
return backend_policy; return backend_policy;
} }
// Diffrent backend policy has dirrent handling func.
py::object ForwardExecutor::RunOpWithBackendPolicy(MsBackendPolicy backend_policy, const OpExecInfoPtr &op_exec_info) { py::object ForwardExecutor::RunOpWithBackendPolicy(MsBackendPolicy backend_policy, const OpExecInfoPtr &op_exec_info) {
py::object result; py::object result;
if (backend_policy == kMsBackendVmOnly) { if (backend_policy == kMsBackendVmOnly) {
@ -2551,6 +2572,7 @@ void GradExecutor::NewGraphInner(py::object *ret, const py::object &cell, const
} }
} }
// Create new top-level maps and manage the number and resources of top-level maps.
void GradExecutor::MakeNewTopGraph(const string &cell_id, const py::object &cell, const py::args &args, void GradExecutor::MakeNewTopGraph(const string &cell_id, const py::object &cell, const py::args &args,
bool is_topest) { bool is_topest) {
pipeline::CheckArgsValid(cell, args); pipeline::CheckArgsValid(cell, args);
@ -2631,6 +2653,8 @@ void GradExecutor::SetTupleItemArgsToGraphInfoMap(const FuncGraphPtr &g, const p
} }
} }
// Cleaning and processing logic at the end of the calculation graph execution, including updating the gradient flag,
// popping the stack, setting the output node, dumping the IR graph, and checking the compiled graph
void GradExecutor::EndGraphInner(py::object *ret, const py::object &cell, const py::object &out, const py::args &args) { void GradExecutor::EndGraphInner(py::object *ret, const py::object &cell, const py::object &out, const py::args &args) {
MS_EXCEPTION_IF_NULL(ret); MS_EXCEPTION_IF_NULL(ret);
const auto &cell_id = GetCellId(cell, args); const auto &cell_id = GetCellId(cell, args);
@ -2800,6 +2824,8 @@ void GradExecutor::MarkMsFunctionNodes(const pipeline::ResourcePtr &resource) {
} }
} }
// Execute the backpropagation graph and manage related resources, including creating, configuring and preparing the graph,
// launching and executing the action, and finally performing the necessary cleaning and releasing operations
void GradExecutor::GradNetInner(py::object *ret, const prim::GradOperationPtr &grad, const py::object &cell, void GradExecutor::GradNetInner(py::object *ret, const prim::GradOperationPtr &grad, const py::object &cell,
const py::object &weights, const py::object &grad_position, const py::args &args) { const py::object &weights, const py::object &grad_position, const py::args &args) {
MS_EXCEPTION_IF_NULL(ret); MS_EXCEPTION_IF_NULL(ret);
@ -2925,6 +2951,7 @@ std::vector<size_t> GradExecutor::GetGradPositionArgs(const py::object &grad_pos
MS_LOG(EXCEPTION) << "Grad position only support tuple."; MS_LOG(EXCEPTION) << "Grad position only support tuple.";
} }
// Shallow copy of sens parameters. That is, create a new sens parameter and replace the original sens parameter to share and transfer data.
void GradExecutor::ShallowCopySensValue(const py::tuple &input_args, bool has_sens, VectorRef *run_args) { void GradExecutor::ShallowCopySensValue(const py::tuple &input_args, bool has_sens, VectorRef *run_args) {
if (!has_sens) { if (!has_sens) {
return; return;
@ -3157,6 +3184,7 @@ void GradExecutor::CheckNeedCompileGraph() {
} }
} }
// The execution process of gradient graph calculation is realized, including the processing of input parameters, shallow copy of sensitive parameters, calculation execution, result conversion and so on.
void GradExecutor::RunGradGraph(py::object *ret, const py::object &cell, const py::tuple &args) { void GradExecutor::RunGradGraph(py::object *ret, const py::object &cell, const py::tuple &args) {
MS_EXCEPTION_IF_NULL(ret); MS_EXCEPTION_IF_NULL(ret);
const auto &cell_id = GetCellId(cell, args); const auto &cell_id = GetCellId(cell, args);
@ -3356,6 +3384,8 @@ void GradExecutor::EraseTopCellFromTopCellList(const TopCellInfoPtr &top_cell) {
} }
} }
// The process of gradient graph calculation for ms function type graph is realized, including creating operation execution information,
// updating tensor information, replacing new tensor, cloning calculation graph and generating backpropagation function.
void GradExecutor::GradMsFunctionInner(const std::string &phase, const py::object &out, const py::args &args, void GradExecutor::GradMsFunctionInner(const std::string &phase, const py::object &out, const py::args &args,
const FuncGraphPtr &ms_func_graph, const FuncGraphPtr &grad_graph) { const FuncGraphPtr &ms_func_graph, const FuncGraphPtr &grad_graph) {
// Get actual output value and added output value. // Get actual output value and added output value.
@ -3405,6 +3435,8 @@ void GradExecutor::GradMsFunctionInner(const std::string &phase, const py::objec
MakeAdjointForMsFunction(new_ms_func_graph, new_grad_graph, actual_out, args, actual_out_v); MakeAdjointForMsFunction(new_ms_func_graph, new_grad_graph, actual_out, args, actual_out_v);
} }
// The process of gradient calculation for ms function diagram is realized, including obtaining the phase of the calculation diagram,
// obtaining the original calculation diagram and the gradient calculation diagram, modifying the output, performing the gradient calculation and so on.
py::object GradExecutor::GradMsFunction(const py::object &out, const py::args &args) { py::object GradExecutor::GradMsFunction(const py::object &out, const py::args &args) {
// Get actual forward output object. // Get actual forward output object.
if (graph_phase().empty()) { if (graph_phase().empty()) {

View File

@ -755,5 +755,5 @@ def constexpr(fn=None, get_instance=True, name=None, reuse_result=True):
@_wrap_func @_wrap_func
def _run_op(obj, op_name, args): def _run_op(obj, op_name, args):
"""Single op execution function supported by ge in PyNative mode.""" """Single op execution function supported by ge in PyNative mode."""
output = real_run_op(obj, op_name, args) output = real_run_op(obj, op_name, args) # jump into C++ function: RealRunOp
return output return output