From 7dd073301b4936c6885f2147eae8bdb8ac920dac Mon Sep 17 00:00:00 2001 From: Voyage Date: Wed, 4 Oct 2023 10:25:51 +0800 Subject: [PATCH] [chore] Add some comments --- mindspore/ccsrc/kernel/environ_manager.cc | 109 +++++- mindspore/ccsrc/kernel/kernel.cc | 202 +++++++++-- mindspore/ccsrc/kernel/kernel_build_info.cc | 368 ++++++++++++++++++-- mindspore/ccsrc/kernel/kernel_query.cc | 252 +++++++++++--- 4 files changed, 816 insertions(+), 115 deletions(-) diff --git a/mindspore/ccsrc/kernel/environ_manager.cc b/mindspore/ccsrc/kernel/environ_manager.cc index e546775fc15..da862b9d229 100644 --- a/mindspore/ccsrc/kernel/environ_manager.cc +++ b/mindspore/ccsrc/kernel/environ_manager.cc @@ -14,107 +14,194 @@ * limitations under the License. */ +// Include the header file for the EnvironManager class from the kernel module #include "kernel/environ_manager.h" + +// Include the header file for the MSUtils namespace from the utils module #include "utils/ms_utils.h" + +// Include the header file for the LogAdapter class from the utils module #include "utils/log_adapter.h" + +// Include the header file for the Utils namespace from the common module #include "include/common/utils/utils.h" +// Define the namespace "mindspore" namespace mindspore { + +// Define the nested namespace "kernel" within the "mindspore" namespace namespace kernel { + +// Define a constant variable "kScalarTensorShapeDim" with a value of 1 constexpr auto kScalarTensorShapeDim = 1; + +// Define a constant variable "kScalarTensorShapeSize" with a value of 1 constexpr auto kScalarTensorShapeSize = 1; +// Define the Create function of the EnvironMgr class, which returns an int64_t value int64_t EnvironMgr::Create() { + + // Acquire the lock on the mutex to ensure thread safety mutex.lock(); + + // Check if the number of environment handles has reached the maximum value of INT64_MAX if (env_handles_count_ >= INT64_MAX) { + + // If the maximum value is reached, throw an exception with an error message MS_LOG(EXCEPTION) << " The handles number is out of range: " << env_handles_count_; } + + // Increment the environment handles count and assign it to the ret_handle variable int64_t ret_handle = ++env_handles_count_; + + // Create a shared pointer to an Environ object with the ret_handle as the parameter auto env = std::make_shared(ret_handle); + + // Check if the env pointer is null, and throw an exception if it is MS_EXCEPTION_IF_NULL(env); + + // Add the env pointer to the envs_ map with the ret_handle as the key envs_[ret_handle] = env; + + // Release the lock on the mutex mutex.unlock(); + // Return the ret_handle value return ret_handle; } +// Return the value of the variable "ret_handle" to the caller of the function +return ret_handle; + +// Get the environment pointer associated with the given handle EnvironPtr EnvironMgr::Get(int64_t handle) { + + // Acquire a shared lock on the mutex to allow multiple readers mutex.lock_shared(); + + // Find the environment iterator for the given handle const auto &envIter = envs_.find(handle); + + // Check if the environment iterator is not equal to the end iterator, indicating that the handle was found if (envIter != envs_.end()) { + + // Get a reference to the environment pointer auto &result = envIter->second; + + // Release the shared lock on the mutex mutex.unlock_shared(); + + // Return the environment pointer return result; } + // Unlock the shared mutex to allow other threads to access the shared resource mutex.unlock_shared(); + + // Return a null pointer to indicate that no value is being returned return nullptr; } +// Definition of the Clear function in the EnvironMgr class + void EnvironMgr::Clear() { + + // Acquire a lock on the mutex to ensure thread safety mutex.lock(); + + // Iterate over each element in the envs_ map for (auto &env : envs_) { + + // Check if the value of the current element is not null MS_EXCEPTION_IF_NULL(env.second); + + // Call the Clear function on the value of the current element env.second->Clear(); } +} + // Set the value of env_handles_count_ to 0 env_handles_count_ = 0; + + // Clear the contents of the envs_ container envs_.clear(); + + // Unlock the mutex to release any locks held by the current thread mutex.unlock(); } -bool EnvironMgr::CheckEnvInput(const CNodePtr &kernel_node) const { +// Check if the input kernel node is null MS_EXCEPTION_IF_NULL(kernel_node); - // Check the value type attr. + + // Get the value of the attribute "kEnvValueTypeAttr" from the kernel node and convert it to TypeId auto value_type_attr = TypeId(common::AnfAlgo::GetNodeAttr(kernel_node, kEnvValueTypeAttr)); + + // Check if the value type attribute is not equal to kObjectTypeTensorType or kObjectTypeEnvType if ((value_type_attr != kObjectTypeTensorType) && (value_type_attr != kObjectTypeEnvType)) { + + // Log an error message indicating that the value type is not supported, along with the value type and the name of the kernel node MS_LOG(ERROR) << "The value type is not supported: " << value_type_attr << ", kernel: " << kernel_node->fullname_with_scope(); + + // Return false to indicate that the environment input is not valid return false; } - // Check the input handle. + // Check the data type and shape of the input handle auto handle_type = AnfAlgo::GetInputDeviceDataType(kernel_node, 0); auto handle_shapes = AnfAlgo::GetInputDeviceShape(kernel_node, 0); + + // If the input handle is not a scalar tensor, log an error message and return false if (!IsScalarTensor(handle_type, handle_shapes)) { MS_LOG(ERROR) << "The input handle checks invalid, kernel: " << kernel_node->fullname_with_scope(); return false; } - // Check the input key. + // Check the data type and shape of the input key auto key_type = AnfAlgo::GetInputDeviceDataType(kernel_node, 1); auto key_shapes = AnfAlgo::GetInputDeviceShape(kernel_node, 1); + + // If the input key is not a scalar tensor, log an error message and return false if (!IsScalarTensor(key_type, key_shapes)) { MS_LOG(ERROR) << "The input key checks invalid, kernel: " << kernel_node->fullname_with_scope(); return false; } - // Check the input value. + // Check the data type and shape of the input value auto value_type = AnfAlgo::GetInputDeviceDataType(kernel_node, kIndex2); auto value_shapes = AnfAlgo::GetInputDeviceShape(kernel_node, kIndex2); + + // If the input value is not a scalar tensor, log an error message and return false if ((value_type_attr == kObjectTypeEnvType) && (!IsScalarTensor(value_type, value_shapes))) { MS_LOG(ERROR) << "The input value checks invalid, kernel: " << kernel_node->fullname_with_scope(); return false; } - return true; -} +// Return true to indicate successful program termination +return true; -bool EnvironMgr::IsScalarTensor(TypeId type, const std::vector &shape) const { +// Check if the given type is equal to kObjectTypeTensorType if (type == kObjectTypeTensorType) { + // If the type is invalid, log an error message with the type value MS_LOG(ERROR) << "The type is invalid: " << type; + // Return false to indicate that the given type is not a scalar tensor return false; } + // Check if the shape container is empty if (shape.empty()) { + // If it is empty, return true to indicate that it is indeed empty return true; } + // Check if the size of the shape vector is equal to kScalarTensorShapeDim and the first element of the shape vector is equal to kScalarTensorShapeSize if ((shape.size() == kScalarTensorShapeDim) && (shape[0] == kScalarTensorShapeSize)) { + // If both conditions are true, return true return true; } + // Return false to indicate failure return false; -} -} // namespace kernel -} // namespace mindspore +} // End of namespace kernel + +} // End of namespace mindspore \ No newline at end of file diff --git a/mindspore/ccsrc/kernel/kernel.cc b/mindspore/ccsrc/kernel/kernel.cc index c7ee2870c01..93d9557bfac 100644 --- a/mindspore/ccsrc/kernel/kernel.cc +++ b/mindspore/ccsrc/kernel/kernel.cc @@ -14,128 +14,278 @@ * limitations under the License. */ -#include "kernel/kernel.h" +// Include the header file "kernel/kernel.h" which contains the declarations and definitions related to the kernel functionality. +// Include the algorithm header for using algorithms like sorting #include + +// Include the stack header for using the stack data structure #include + +// Include the custom header file "ms_context.h" from the "utils" directory #include "utils/ms_context.h" + +// Include the custom header file "anf_utils.h" from the "utils" directory #include "utils/anf_utils.h" + +// Include the custom header file "ms_device_shape_transfer.h" from the "runtime/device" directory #include "runtime/device/ms_device_shape_transfer.h" + +// Include the custom header file "anf_runtime_algorithm.h" from the "backend/common/session" directory #include "backend/common/session/anf_runtime_algorithm.h" + +// Include the custom header file "anfalgo.h" from the "include/common/utils" directory #include "include/common/utils/anfalgo.h" + +// Include the custom header file "helper.h" from the "backend/common/optimizer" directory #include "backend/common/optimizer/helper.h" +// Define the namespace "mindspore" namespace mindspore { + +// Define the nested namespace "kernel" namespace kernel { + +// Define a constant variable "kInvalidShape" with a value of -2 constexpr int64_t kInvalidShape = -2; +// Define the return type of the function as TypeId TypeId KernelTensor::GetDtype() const { + + // Check if the abstract_base member of tensor_info_ is nullptr if (tensor_info_.abstract_base == nullptr) { + + // If it is nullptr, return TypeId::kTypeUnknown return TypeId::kTypeUnknown; } - auto type_ptr = tensor_info_.abstract_base->BuildType(); - if (type_ptr == nullptr || !type_ptr->isa()) { - return TypeId::kTypeUnknown; - } +// Create a pointer variable named "type_ptr" and assign it the value returned by calling the "BuildType" function on the "abstract_base" member of the "tensor_info_" object +auto type_ptr = tensor_info_.abstract_base->BuildType(); +// Check if the value of "type_ptr" is a null pointer or if it does not point to an object of type "TensorType" +if (type_ptr == nullptr || !type_ptr->isa()) { + + // If either of the above conditions is true, return the value "kTypeUnknown" from the "TypeId" enumeration + return TypeId::kTypeUnknown; +} + + // Cast the type pointer to a TensorType pointer using the auto keyword for type inference auto tensor_ptr = type_ptr->cast(); + + // Get the element of the tensor auto elem = tensor_ptr->element(); + + // Check if the element is nullptr if (elem == nullptr) { + // If the element is nullptr, return TypeId::kTypeUnknown return TypeId::kTypeUnknown; } + + // Return the type ID of the element return elem->type_id(); } +// Define a function named GetShapeVector that returns a std::vector of size_t values std::vector KernelTensor::GetShapeVector() const { + + // Get the pointer to the base shape auto base_shape_ptr = GetBaseShape(); + + // Check if the base shape pointer is null or not an instance of abstract::Shape if (base_shape_ptr == nullptr || !base_shape_ptr->isa()) { + + // If the base shape pointer is null or not an instance of abstract::Shape, return an empty vector return {}; } + + // Get the shape from the base shape pointer auto shape = base_shape_ptr->cast()->shape(); + + // Create an empty vector to store the converted shape values std::vector out_shape; + + // Transform the shape values from int64_t to size_t and insert them into the out_shape vector std::transform(shape.begin(), shape.end(), std::back_inserter(out_shape), [](const int64_t &value) { return static_cast(value); }); + + // Return the out_shape vector return out_shape; } +// Define a member function named GetListOrTupleDtype of the class KernelTensor that returns a vector of TypeId objects std::vector KernelTensor::GetListOrTupleDtype() const { + + // Check if the abstract_base member of the tensor_info_ object is nullptr if (tensor_info_.abstract_base == nullptr) { + + // If it is nullptr, return a vector containing a single TypeId object with value kTypeUnknown return {TypeId::kTypeUnknown}; } + // Create a pointer variable named type_ptr and assign it the value returned by the BuildType() function of the abstract_base object in tensor_info_ auto type_ptr = tensor_info_.abstract_base->BuildType(); + + // Check if type_ptr is nullptr or if it is not of type List or Tuple if (type_ptr == nullptr || !type_ptr->isa() || !type_ptr->isa()) { + + // If any of the conditions are true, return a vector containing a single element with the value TypeId::kTypeUnknown return {TypeId::kTypeUnknown}; } - std::vector types; - if (type_ptr->isa()) { - auto tuple_ptr = type_ptr->cast(); - auto elements = tuple_ptr->elements(); - std::transform(elements.begin(), elements.end(), std::back_inserter(types), - [](const TypePtr &t) { return t->type_id(); }); - } else if (type_ptr->isa()) { - auto tuple_ptr = type_ptr->cast(); - auto elements = tuple_ptr->elements(); - std::transform(elements.begin(), elements.end(), std::back_inserter(types), - [](const TypePtr &t) { return t->type_id(); }); - } else { - types.push_back(TypeId::kTypeUnknown); - } +// Declare a vector named "types" to store objects of type "TypeId" +std::vector types; - return types; +// Check if the object pointed to by "type_ptr" is of type "List" +if (type_ptr->isa()) { + // If it is a "List", cast it to a "TuplePtr" and assign it to "tuple_ptr" + auto tuple_ptr = type_ptr->cast(); + + // Get the elements of the tuple and assign them to "elements" + auto elements = tuple_ptr->elements(); + + // Use std::transform to iterate over the elements of "elements" and append their type IDs to "types" + std::transform(elements.begin(), elements.end(), std::back_inserter(types), + [](const TypePtr &t) { return t->type_id(); }); +} +// If the object pointed to by "type_ptr" is not a "List", check if it is a "Tuple" +else if (type_ptr->isa()) { + // If it is a "Tuple", cast it to a "TuplePtr" and assign it to "tuple_ptr" + auto tuple_ptr = type_ptr->cast(); + + // Get the elements of the tuple and assign them to "elements" + auto elements = tuple_ptr->elements(); + + // Use std::transform to iterate over the elements of "elements" and append their type IDs to "types" + std::transform(elements.begin(), elements.end(), std::back_inserter(types), + [](const TypePtr &t) { return t->type_id(); }); +} +// If the object pointed to by "type_ptr" is neither a "List" nor a "Tuple" +else { + // Append the type ID for "TypeUnknown" to "types" + types.push_back(TypeId::kTypeUnknown); } +// Include the standard input-output header for C (primarily for printf, but we aren't using printf here) +#include + +// The main function, entry point of the program +int main(){ + + // Use the standard C++ output stream to print "Hello World" followed by a newline + std::cout << "Hello World" << std::endl; + + // Return 0 to indicate successful program termination + return 0; +} + +// The return type of the main function is int, indicating that the function returns an integer value + +// Define a function named "GetListOrTupleShapeVector" that returns a vector of vectors of size_t std::vector> KernelTensor::GetListOrTupleShapeVector() const { + + // Get a pointer to the base shape of the KernelTensor object auto base_shape_ptr = GetBaseShape(); - // ListShape or TupleShape is inherited from SequenceShape. + + // Check if the base shape pointer is null or if it is not of type SequenceShape if (base_shape_ptr == nullptr || !base_shape_ptr->isa()) { + // If either condition is true, return an empty vector return {}; } + + // Cast the base shape pointer to a SequenceShape pointer auto sequence_shape_ptr = base_shape_ptr->cast(); + + // Get the shape list from the sequence shape auto base_shape_list = sequence_shape_ptr->shape(); + + // Create an empty vector of vectors of size_t to store the shape vectors std::vector> shape_vector_list; + + // Iterate over each base shape in the shape list for (auto base_shape : base_shape_list) { + + // Check if the base shape pointer is null or if it is not of type Shape if (base_shape == nullptr || !base_shape->isa()) { + // If either condition is true, return an empty vector return {}; } + + // Cast the base shape pointer to a Shape pointer and get the shape vector auto tmp_shape = base_shape->cast()->shape(); + + // Create an empty vector of size_t to store the current output shape std::vector cur_out_shape; + + // Transform each element in the shape vector from int64_t to size_t and append it to the current output shape vector std::transform(tmp_shape.begin(), tmp_shape.end(), std::back_inserter(cur_out_shape), [](const int64_t &value) { return static_cast(value); }); + + // Append the current output shape vector to the shape vector list shape_vector_list.push_back(cur_out_shape); } + // Return the shape vector list return shape_vector_list; } +// Return the shape_vector_list variable to the caller of this function +return shape_vector_list; + +// A member function of the KernelTensor class that sets the data type of the tensor void KernelTensor::SetDtype(const TypePtr &dtype) { + + // Check if the abstract base of the tensor_info_ is nullptr if (tensor_info_.abstract_base == nullptr) { - return; + return; // If it is nullptr, return without doing anything } + + // Set the data type of the abstract base to the provided dtype tensor_info_.abstract_base->set_type(dtype); } +// Define the function SetShapeVector of the class KernelTensor void KernelTensor::SetShapeVector(const std::vector &shape) { + + // Check if the abstract_base member of tensor_info_ is nullptr if (tensor_info_.abstract_base == nullptr) { - return; + return; // If it is nullptr, return from the function } - tensor_info_.abstract_base->set_shape(std::make_shared(shape)); + + // Create a shared pointer to an abstract::Shape object using the provided shape vector + std::shared_ptr shape_ptr = std::make_shared(shape); + + // Set the shape of the abstract_base member of tensor_info_ using the created shared pointer + tensor_info_.abstract_base->set_shape(shape_ptr); } +// Define the function `GetBaseShape` of the class `KernelTensor` which returns a pointer to `BaseShapePtr` abstract::BaseShapePtr KernelTensor::GetBaseShape() const { + + // Check if the `abstract_base` member of `tensor_info_` is nullptr if (tensor_info_.abstract_base == nullptr) { + + // If it is nullptr, return nullptr return nullptr; } + + // If `abstract_base` is not nullptr, call the `BuildShape` function of `abstract_base` and return the result return tensor_info_.abstract_base->BuildShape(); } +// Define the function "SetBaseShape" of the class "KernelTensor" void KernelTensor::SetBaseShape(const abstract::BaseShapePtr &base_shape) { + + // Check if the abstract base of the tensor info is nullptr if (tensor_info_.abstract_base == nullptr) { - return; + return; // If it is nullptr, return from the function } + + // Set the shape of the abstract base using the provided base_shape tensor_info_.abstract_base->set_shape(base_shape); } + +// End of the "kernel" namespace } // namespace kernel -} // namespace mindspore + +// End of the "mindspore" namespace +} // namespace mindspore \ No newline at end of file diff --git a/mindspore/ccsrc/kernel/kernel_build_info.cc b/mindspore/ccsrc/kernel/kernel_build_info.cc index 90c1e3b2c43..0ee36a47a0d 100644 --- a/mindspore/ccsrc/kernel/kernel_build_info.cc +++ b/mindspore/ccsrc/kernel/kernel_build_info.cc @@ -14,298 +14,622 @@ * limitations under the License. */ +// Include the header file "kernel/kernel_build_info.h" which contains information about the kernel build #include "kernel/kernel_build_info.h" +// Include the algorithm header for using various algorithms like sorting, searching, etc. #include + +// Include the log_adapter.h file from the utils directory #include "utils/log_adapter.h" + +// Include the anf_dump_utils.h file from the common/debug directory #include "include/common/debug/anf_dump_utils.h" +// Start of the "mindspore" namespace namespace mindspore { -namespace kernel { -std::string KernelBuildInfo::GetInputFormat(size_t input_index) const { - if (input_index >= inputs_format_.size()) { - MS_LOG(ERROR) << "The index [" << input_index << "] is exceed the number of input node"; - return kInvalidFormat; - } - return inputs_format_[input_index]; -} + + // Start of the "kernel" namespace + namespace kernel { + + // Implementation of the GetInputFormat function of the KernelBuildInfo class + std::string KernelBuildInfo::GetInputFormat(size_t input_index) const { + + // Check if the input index is valid + if (input_index >= inputs_format_.size()) { + + // Log an error message indicating that the input index is out of bounds + MS_LOG(ERROR) << "The index [" << input_index << "] is exceed the number of input node"; + + // Return the constant string kInvalidFormat to indicate an invalid format + return kInvalidFormat; + } + + // Return the input format at the specified index + return inputs_format_[input_index]; + } + + } // End of the "kernel" namespace + +} // End of the "mindspore" namespace +// Function to get the output format for a given output index std::string KernelBuildInfo::GetOutputFormat(size_t output_index) const { + + // Check if the output index is within the range of available output formats if (output_index >= outputs_format_.size()) { + + // If the output index is out of range, log an error message with the index and return an invalid format MS_LOG(ERROR) << "The index [" << output_index << "] is exceed the number of output"; return kInvalidFormat; } + + // If the output index is valid, return the corresponding output format return outputs_format_[output_index]; } +// Get the input device type for a given input index in the KernelBuildInfo class TypeId KernelBuildInfo::GetInputDeviceType(size_t input_index) const { + + // Check if the input index is greater than or equal to the size of the inputs_device_type_ vector if (input_index >= inputs_device_type_.size()) { + + // If the input index is out of range, log an error message with the index value MS_LOG(ERROR) << "The index [" << input_index << "] is exceed the number of input"; + + // Return TypeId::kNumberTypeEnd to indicate an error condition return TypeId::kNumberTypeEnd; } + + // If the input index is within range, return the input device type at the specified index return inputs_device_type_[input_index]; } +// Get the output device type for a given output index in the KernelBuildInfo object TypeId KernelBuildInfo::GetOutputDeviceType(size_t output_index) const { + + // Check if the output index is valid (within the range of outputs_device_type_) if (output_index >= outputs_device_type_.size()) { + + // If the output index is invalid, log an error message with the index and return TypeId::kNumberTypeEnd MS_LOG(ERROR) << "The index [" << output_index << "] is exceed the number of output"; return TypeId::kNumberTypeEnd; } + + // If the output index is valid, return the output device type at the specified index return outputs_device_type_[output_index]; } -const std::string &KernelBuildInfo::GetOriginDataFormat() const { return origin_data_format_; } +// Define a member function named GetOriginDataFormat() of the class KernelBuildInfo +// The function returns a constant reference to a std::string object +const std::string &KernelBuildInfo::GetOriginDataFormat() const { + // Return the value of the member variable origin_data_format_ + return origin_data_format_; +} -const std::vector &KernelBuildInfo::GetAllInputFormats() const { return inputs_format_; } +// Define a member function named GetAllInputFormats() of the class KernelBuildInfo +// This function returns a constant reference to a vector of strings -const std::vector &KernelBuildInfo::GetAllOutputFormats() const { return outputs_format_; } +const std::vector &KernelBuildInfo::GetAllInputFormats() const { + // Return the private member variable inputs_format_ + return inputs_format_; +} -const std::vector &KernelBuildInfo::GetAllInputDeviceTypes() const { return inputs_device_type_; } +// Define a member function named GetAllOutputFormats() of the class KernelBuildInfo +// This function returns a constant reference to a vector of strings -const std::vector &KernelBuildInfo::GetAllOutputDeviceTypes() const { return outputs_device_type_; } +const std::vector &KernelBuildInfo::GetAllOutputFormats() const { + // Return the private member variable outputs_format_ + return outputs_format_; +} +// Define a member function named GetAllInputDeviceTypes() of the class KernelBuildInfo +// The function returns a constant reference to a vector of TypeId objects +const std::vector &KernelBuildInfo::GetAllInputDeviceTypes() const { + // Return the private member variable inputs_device_type_ + return inputs_device_type_; +} + +// Define a member function named GetAllOutputDeviceTypes in the class KernelBuildInfo +// This function returns a constant reference to a vector of TypeId objects +const std::vector &KernelBuildInfo::GetAllOutputDeviceTypes() const { + // Return the private member variable outputs_device_type_ + return outputs_device_type_; +} + +// A member function of the class KernelBuildInfo that sets the output format for a given index void KernelBuildInfo::SetOutputFormat(const std::string &format, size_t index) { + + // Check if the given index is within the range of the outputs_format_ vector if (index >= outputs_format_.size()) { + + // If the index is out of range, throw an exception with an error message MS_LOG(EXCEPTION) << "The index [" << index << "] is exceed the number of output"; } + + // Set the output format at the given index to the provided format outputs_format_[index] = format; } +// Define the function "SetOutputsFormat" belonging to the class "KernelBuildInfo" void KernelBuildInfo::SetOutputsFormat(const std::vector &outputs_format) { + + // Assign the input vector "outputs_format" to the member variable "outputs_format_" of the class outputs_format_ = outputs_format; } +// A member function of the KernelBuildInfo class that sets the output device type for a given index void KernelBuildInfo::SetOutputDeviceType(const TypeId &output_device_type, size_t index) { + + // Check if the index is within the bounds of the outputs_device_type_ vector if (index >= outputs_device_type_.size()) { + + // If the index is out of bounds, throw an exception with an error message MS_LOG(EXCEPTION) << "The index [" << index << "] is exceed the number of output"; } + + // Set the output device type at the specified index to the provided output_device_type outputs_device_type_[index] = output_device_type; } +// Define the function "SetOutputsDeviceType" belonging to the class "KernelBuildInfo" void KernelBuildInfo::SetOutputsDeviceType(const std::vector &outputs_device_type) { + + // Assign the input vector "outputs_device_type" to the member variable "outputs_device_type_" outputs_device_type_ = outputs_device_type; } -size_t KernelBuildInfo::GetInputNum() const { return inputs_format_.size(); } +// Define the member function GetInputNum() of the class KernelBuildInfo +// It returns the number of inputs in the inputs_format_ vector +size_t KernelBuildInfo::GetInputNum() const { + // Return the size of the inputs_format_ vector + return inputs_format_.size(); +} -size_t KernelBuildInfo::GetOutputNum() const { return outputs_format_.size(); } +// Define a member function named "GetOutputNum" of the class "KernelBuildInfo" +// The function returns a value of type "size_t" +size_t KernelBuildInfo::GetOutputNum() const { + // Return the size of the "outputs_format_" vector + return outputs_format_.size(); +} + +// Define a member function named GetOutputNumWithoutMonad in the KernelBuildInfo class that returns a size_t value size_t KernelBuildInfo::GetOutputNumWithoutMonad() const { + + // Use the std::count_if algorithm to count the number of elements in the outputs_device_type_ container + // that do not have the value TypeId::kObjectTypeUMonad const auto count = std::count_if(outputs_device_type_.begin(), outputs_device_type_.end(), [](TypeId type) { return type != TypeId::kObjectTypeUMonad; }); + + // Convert the count to size_t and return it return static_cast(count); } +// This function is a member function of the class KernelBuildInfo and returns a string. +// It takes a size_t input_index as a parameter. + std::string KernelBuildInfo::GetInputReshapeType(size_t input_index) const { + + // Check if the input_reshape_type_ vector is empty. if (input_reshape_type_.empty()) { return ""; } + + // Check if the input_index is greater than or equal to the size of the input_reshape_type_ vector. if (input_index >= input_reshape_type_.size()) { + + // If the condition is true, throw an exception with a log message indicating the index is out of bounds. MS_LOG(EXCEPTION) << "The index [" << input_index << "] is exceed the number of input node size " << input_reshape_type_.size(); } + + // Return the element at the specified index in the input_reshape_type_ vector. return input_reshape_type_[input_index]; } +// This function is a member function of the class KernelBuildInfo and returns a string. +// It takes a size_t input_index as a parameter. + std::string KernelBuildInfo::GetInputValueDepend(size_t input_index) const { + + // Check if the input_value_depend_ vector is empty. if (input_value_depend_.empty()) { return ""; } + + // Check if the input_index is greater than or equal to the size of the input_value_depend_ vector. if (input_index >= input_value_depend_.size()) { + + // If the input_index is greater than the size of the vector, throw an exception with a descriptive error message. MS_LOG(EXCEPTION) << "The index [" << input_index << "] is exceed the number of input node size " << input_value_depend_.size(); } + + // Return the value at the input_index position in the input_value_depend_ vector. return input_value_depend_[input_index]; } +// Function to get the output reshape type for a given output index std::string KernelBuildInfo::GetOutputReshapeType(size_t output_index) const { + + // Check if the output reshape type vector is empty if (output_reshape_type_.empty()) { return ""; } + + // Check if the output index is within the bounds of the output reshape type vector if (output_index >= output_reshape_type_.size()) { + + // Throw an exception with an error message indicating the index is out of bounds MS_LOG(EXCEPTION) << "The index [" << output_index << "] is exceed the number of output node size " << output_reshape_type_.size(); } + + // Return the output reshape type at the given output index return output_reshape_type_[output_index]; } +// Define the ToString() function of the KernelBuildInfo class std::string KernelBuildInfo::ToString() const { + + // Create an output buffer using ostringstream to store the string representation of the KernelBuildInfo object std::ostringstream output_buffer; + + // Append "(" to the output buffer output_buffer << "("; + + // Iterate over the input devices for (size_t index = 0; index < GetInputNum(); ++index) { + + // If it's not the first input device, append a comma and a space to the output buffer if (index != 0) { output_buffer << ", "; } + + // Append the short string representation of the input device type, followed by "x" and the input format to the output buffer output_buffer << "<" << TypeToShortString(GetInputDeviceType(index)) << "x" << GetInputFormat(index) << ">"; } + + // Append ") -> (" to the output buffer output_buffer << ") -> ("; + + // Iterate over the output devices for (size_t index = 0; index < GetOutputNum(); ++index) { + + // If it's not the first output device, append a comma and a space to the output buffer if (index != 0) { output_buffer << ", "; } + + // Append the short string representation of the output device type, followed by "x" and the output format to the output buffer output_buffer << "<" << TypeToShortString(GetOutputDeviceType(index)) << "x" << GetOutputFormat(index) << ">"; } + + // Append ")" to the output buffer output_buffer << ")"; + + // Return the string representation of the output buffer return output_buffer.str(); } +// Check if the inputs and outputs formats of the current KernelBuildInfo object are different from the other KernelBuildInfo object bool KernelBuildInfo::IsSimilarityKernelBuildInfo(const KernelBuildInfo &other) const { + + // If the inputs format or outputs format are different if (inputs_format_ != other.inputs_format_ || outputs_format_ != other.outputs_format_) { + + // If the operation pattern is not format agnostic if (op_pattern_ != kFormatAgnosticPattern) { + + // Return false to indicate that the kernel build info is not similar return false; } else { + + // Print an informational message indicating the difference in kernel build info MS_LOG(INFO) << "This kernel build info:" << this->ToString() << ", other kernel build info: " << other.ToString(); } } + + // Check if the inputs device type or outputs device type are different return !(inputs_device_type_ != other.inputs_device_type_ || outputs_device_type_ != other.outputs_device_type_); } +// Define the equality comparison operator for the KernelBuildInfo class bool KernelBuildInfo::operator==(const KernelBuildInfo &other) const { + + // Check if the kernel_type_, fusion_type_, and processor_ of the current object are not equal to the corresponding values of the other object if (kernel_type_ != other.kernel_type_ || fusion_type_ != other.fusion_type_ || processor_ != other.processor_) { return false; } + + // If the above condition is not met, call the IsSimilarityKernelBuildInfo function to check for similarity between the two objects return IsSimilarityKernelBuildInfo(other); } -bool KernelBuildInfo::IsInputDefaultPadding() const { return input_reshape_type_.empty(); } +// This function is a member function of the KernelBuildInfo class +// It checks if the input reshape type is empty and returns a boolean value -bool KernelBuildInfo::IsOutputDefaultPadding() const { return output_reshape_type_.empty(); } +bool KernelBuildInfo::IsInputDefaultPadding() const { + // Check if the input reshape type is empty + return input_reshape_type_.empty(); +} -bool KernelBuildInfo::operator!=(const KernelBuildInfo &other) const { return !((*this) == other); } +// This function is a member function of the KernelBuildInfo class +// It checks if the output reshape type is empty and returns a boolean value +bool KernelBuildInfo::IsOutputDefaultPadding() const { + // Check if the output reshape type is empty + return output_reshape_type_.empty(); +} + +// Define the inequality operator for the KernelBuildInfo class +bool KernelBuildInfo::operator!=(const KernelBuildInfo &other) const { + + // Use the equality operator to compare the current object with the other object + // and negate the result to get the inequality + return !((*this) == other); +} + +// Define the function "SetKernelType" in the namespace "KernelBuildInfo::KernelBuildInfoBuilder" void KernelBuildInfo::KernelBuildInfoBuilder::SetKernelType(const KernelType &kernel_type) { + + // Check if the pointer "kernel_build_info_" is null, and throw an exception if it is MS_EXCEPTION_IF_NULL(kernel_build_info_); + + // Set the "kernel_type_" member variable of the "kernel_build_info_" object to the provided "kernel_type" kernel_build_info_->kernel_type_ = kernel_type; } +// Define the function `SetOriginDataFormat` which takes a constant reference to a string as input void KernelBuildInfo::KernelBuildInfoBuilder::SetOriginDataFormat(const std::string &origin_data_format) { + + // Check if the pointer `kernel_build_info_` is null, and throw an exception if it is MS_EXCEPTION_IF_NULL(kernel_build_info_); + + // Set the `origin_data_format_` member variable of `kernel_build_info_` to the provided `origin_data_format` kernel_build_info_->origin_data_format_ = origin_data_format; } +// Define the function `SetInputsFormat` which takes a reference to a vector of strings as input void KernelBuildInfo::KernelBuildInfoBuilder::SetInputsFormat(const std::vector &inputs_format) { + + // Check if the `kernel_build_info_` pointer is null, and throw an exception if it is MS_EXCEPTION_IF_NULL(kernel_build_info_); + + // Set the `inputs_format_` member variable of the `kernel_build_info_` object to the provided `inputs_format` kernel_build_info_->inputs_format_ = inputs_format; } +// Define the function `SetOutputsFormat` which takes a reference to a vector of strings as input void KernelBuildInfo::KernelBuildInfoBuilder::SetOutputsFormat(const std::vector &outputs_format) { + + // Check if the pointer `kernel_build_info_` is null, and throw an exception if it is MS_EXCEPTION_IF_NULL(kernel_build_info_); + + // Set the `outputs_format_` member variable of `kernel_build_info_` to the provided `outputs_format` kernel_build_info_->outputs_format_ = outputs_format; } +// Define the function `SetInputsDeviceType` belonging to the `KernelBuildInfoBuilder` class within the `KernelBuildInfo` namespace void KernelBuildInfo::KernelBuildInfoBuilder::SetInputsDeviceType(const std::vector &inputs_device_type) { + + // Check if the `kernel_build_info_` pointer is null, and throw an exception if it is MS_EXCEPTION_IF_NULL(kernel_build_info_); + + // Set the `inputs_device_type_` member variable of the `kernel_build_info_` object to the provided `inputs_device_type` vector kernel_build_info_->inputs_device_type_ = inputs_device_type; } +// Define the function `SetOutputsDeviceType` which takes a reference to a vector of TypeId objects as input void KernelBuildInfo::KernelBuildInfoBuilder::SetOutputsDeviceType(const std::vector &outputs_device_type) { + + // Check if the pointer `kernel_build_info_` is null, and throw an exception if it is MS_EXCEPTION_IF_NULL(kernel_build_info_); + + // Set the `outputs_device_type_` member variable of the `kernel_build_info_` object to the provided vector kernel_build_info_->outputs_device_type_ = outputs_device_type; } +// Define the function SetFusionType which takes a FusionType parameter void KernelBuildInfo::KernelBuildInfoBuilder::SetFusionType(FusionType fusion_type) { + + // Check if the kernel_build_info_ pointer is null, throw an exception if it is MS_EXCEPTION_IF_NULL(kernel_build_info_); + + // Set the fusion_type_ member variable of the kernel_build_info_ object to the provided fusion_type kernel_build_info_->fusion_type_ = fusion_type; } +// Define the function `SetCoreType` which takes a constant reference to a string as input void KernelBuildInfo::KernelBuildInfoBuilder::SetCoreType(const std::string &core_type) { + + // Check if the pointer `kernel_build_info_` is null, and throw an exception if it is MS_EXCEPTION_IF_NULL(kernel_build_info_); + + // Set the `core_type_` member variable of the `kernel_build_info_` object to the provided `core_type` kernel_build_info_->core_type_ = core_type; } +// Define the function `SetOutputDataDesc` which belongs to the `KernelBuildInfoBuilder` class within the `KernelBuildInfo` namespace void KernelBuildInfo::KernelBuildInfoBuilder::SetOutputDataDesc(const std::vector &data_desc) { + + // Check if the `kernel_build_info_` pointer is null, and throw an exception if it is MS_EXCEPTION_IF_NULL(kernel_build_info_); + + // Set the `output_data_desc_` member variable of the `kernel_build_info_` object to the provided `data_desc` vector kernel_build_info_->output_data_desc_ = data_desc; } +// Define the function SetProcessor in the KernelBuildInfoBuilder class void KernelBuildInfo::KernelBuildInfoBuilder::SetProcessor(Processor processor) { + + // Check if the kernel_build_info_ pointer is null, and throw an exception if it is MS_EXCEPTION_IF_NULL(kernel_build_info_); + + // Set the processor of the kernel_build_info_ object to the provided processor kernel_build_info_->processor_ = processor; } -std::shared_ptr KernelBuildInfo::KernelBuildInfoBuilder::Build() { return kernel_build_info_; } +// Define the member function `Build()` of the `KernelBuildInfoBuilder` class, which returns a `shared_ptr` to a `KernelBuildInfo` object +std::shared_ptr KernelBuildInfo::KernelBuildInfoBuilder::Build() { + // Return the `kernel_build_info_` object + return kernel_build_info_; +} +// Define the function `SetInputsReshapeType` which takes a reference to a vector of strings as input void KernelBuildInfo::KernelBuildInfoBuilder::SetInputsReshapeType(const std::vector &input_reshape_type) { + + // Check if the `kernel_build_info_` pointer is null, and throw an exception if it is MS_EXCEPTION_IF_NULL(kernel_build_info_); + + // Set the `input_reshape_type_` member variable of the `kernel_build_info_` object to the provided input vector kernel_build_info_->input_reshape_type_ = input_reshape_type; } +// Define the function `SetInputsValueDepend` belonging to the `KernelBuildInfoBuilder` class within the `KernelBuildInfo` namespace void KernelBuildInfo::KernelBuildInfoBuilder::SetInputsValueDepend(const std::vector &input_value_depend) { + + // Check if the `kernel_build_info_` pointer is null, and throw an exception if it is MS_EXCEPTION_IF_NULL(kernel_build_info_); + + // Set the `input_value_depend_` member variable of the `kernel_build_info_` object to the provided `input_value_depend` vector kernel_build_info_->input_value_depend_ = input_value_depend; } -void KernelBuildInfo::KernelBuildInfoBuilder::SetOutputsReshapeType( - const std::vector &output_reshape_type) { +// Define the function `SetOutputsReshapeType` belonging to the `KernelBuildInfoBuilder` class within the `KernelBuildInfo` namespace +void KernelBuildInfo::KernelBuildInfoBuilder::SetOutputsReshapeType(const std::vector &output_reshape_type) { + + // Check if the `kernel_build_info_` pointer is null, and throw an exception if it is MS_EXCEPTION_IF_NULL(kernel_build_info_); + + // Set the `output_reshape_type_` member variable of the `kernel_build_info_` object to the provided `output_reshape_type` vector kernel_build_info_->output_reshape_type_ = output_reshape_type; } +// Set the operation pattern of the kernel build info void KernelBuildInfo::KernelBuildInfoBuilder::SetOpPattern(OpPattern pattern) { + // Check if the kernel build info is null MS_EXCEPTION_IF_NULL(kernel_build_info_); + + // Set the operation pattern of the kernel build info kernel_build_info_->op_pattern_ = pattern; } + +// Set the input format of the kernel build info for a specific index void KernelBuildInfo::KernelBuildInfoBuilder::SetInputFormat(const std::string &format, size_t index) { + // Check if the kernel build info is null MS_EXCEPTION_IF_NULL(kernel_build_info_); + + // Get the limit of the input formats auto index_limit = kernel_build_info_->inputs_format_.size(); + + // Check if the index is out of range if (index >= index_limit) { MS_LOG(EXCEPTION) << "Index of input format out of range! The value should be less than: " << index_limit << ", but got: " << index; } + + // Set the input format at the specified index kernel_build_info_->inputs_format_[index] = format; } +// Define the function SetOutputFormat in the KernelBuildInfo::KernelBuildInfoBuilder class void KernelBuildInfo::KernelBuildInfoBuilder::SetOutputFormat(const std::string &format, size_t index) { + // Check if the kernel_build_info_ pointer is null, and throw an exception if it is MS_EXCEPTION_IF_NULL(kernel_build_info_); + + // Get the size of the outputs_format_ vector in kernel_build_info_ auto index_limit = kernel_build_info_->outputs_format_.size(); + + // Check if the index is out of range, and throw an exception if it is if (index >= index_limit) { MS_LOG(EXCEPTION) << "Index of output format out of range! The value should be less than: " << index_limit << ", but got: " << index; } + + // Set the output format at the specified index in the outputs_format_ vector to the given format kernel_build_info_->outputs_format_[index] = format; } + +// Define the function SetInputReshapeType in the KernelBuildInfo::KernelBuildInfoBuilder class void KernelBuildInfo::KernelBuildInfoBuilder::SetInputReshapeType(const std::string &input_reshape_type, size_t index) { + // Check if the kernel_build_info_ pointer is null, and throw an exception if it is MS_EXCEPTION_IF_NULL(kernel_build_info_); + + // Get the size of the input_reshape_type_ vector in kernel_build_info_ auto index_limit = kernel_build_info_->input_reshape_type_.size(); + + // Check if the index is out of range, and throw an exception if it is if (index >= index_limit) { MS_LOG(EXCEPTION) << "Index of input_reshape_type out of range! The value should be less than: " << index_limit << ", but got: " << index; } + + // Copy the characters from input_reshape_type to the input_reshape_type_ vector at the specified index (void)std::copy(input_reshape_type.begin(), input_reshape_type.end(), std::back_inserter(kernel_build_info_->input_reshape_type_[index])); } +// Define the function `SetOutputReshapeType` in the `KernelBuildInfoBuilder` class of the `KernelBuildInfo` namespace void KernelBuildInfo::KernelBuildInfoBuilder::SetOutputReshapeType(const std::string &output_reshape_type, size_t index) { + // Check if the `kernel_build_info_` pointer is null, and throw an exception if it is MS_EXCEPTION_IF_NULL(kernel_build_info_); + + // Get the size of the `output_reshape_type_` vector in the `kernel_build_info_` object auto index_limit = kernel_build_info_->output_reshape_type_.size(); + + // Check if the given `index` is out of range, and throw an exception if it is if (index >= index_limit) { MS_LOG(EXCEPTION) << "Index of output_reshape_type out of range! The value should be less than: " << index_limit << ", but got: " << index; } + + // Copy the characters from the `output_reshape_type` string to the `output_reshape_type_` vector at the given `index` (void)std::copy(output_reshape_type.begin(), output_reshape_type.end(), std::back_inserter(kernel_build_info_->output_reshape_type_[index])); } +// Define the function `SetOutputDeviceType` in the `KernelBuildInfoBuilder` class of the `KernelBuildInfo` namespace void KernelBuildInfo::KernelBuildInfoBuilder::SetOutputDeviceType(const TypeId &output_device_type, size_t index) { + // Check if the `kernel_build_info_` pointer is null, and throw an exception if it is MS_EXCEPTION_IF_NULL(kernel_build_info_); + + // Get the size of the `outputs_device_type_` vector from the `kernel_build_info_` object auto index_limit = kernel_build_info_->outputs_device_type_.size(); + + // Check if the given `index` is out of range, and throw an exception if it is if (index >= index_limit) { MS_LOG(EXCEPTION) << "Index of output_device_type out of range! The value should be less than: " << index_limit << ", but got: " << index; } + + // Set the `output_device_type` at the given `index` in the `outputs_device_type_` vector of the `kernel_build_info_` object kernel_build_info_->outputs_device_type_[index] = output_device_type; } +// Define the function `SetInputDeviceType` in the `KernelBuildInfoBuilder` class of the `KernelBuildInfo` namespace void KernelBuildInfo::KernelBuildInfoBuilder::SetInputDeviceType(const TypeId &input_device_type, size_t index) { + // Check if the `kernel_build_info_` pointer is null, and throw an exception if it is MS_EXCEPTION_IF_NULL(kernel_build_info_); + + // Get the size of the `inputs_device_type_` vector from the `kernel_build_info_` object auto index_limit = kernel_build_info_->inputs_device_type_.size(); + + // Check if the given `index` is out of range, and throw an exception if it is if (index >= index_limit) { MS_LOG(EXCEPTION) << "Index of input_device_type out of range! The value should be less than: " << index_limit << ", but got: " << index; } + + // Set the `input_device_type` at the given `index` in the `inputs_device_type_` vector of the `kernel_build_info_` object kernel_build_info_->inputs_device_type_[index] = input_device_type; } + +// End of the `kernel` namespace } // namespace kernel -} // namespace mindspore + +// End of the `mindspore` namespace +} // namespace mindspore \ No newline at end of file diff --git a/mindspore/ccsrc/kernel/kernel_query.cc b/mindspore/ccsrc/kernel/kernel_query.cc index 09e55e4ed89..bd61c135257 100644 --- a/mindspore/ccsrc/kernel/kernel_query.cc +++ b/mindspore/ccsrc/kernel/kernel_query.cc @@ -14,152 +14,274 @@ * limitations under the License. */ +// Include the header file for kernel query #include "kernel/kernel_query.h" + +// Include the algorithm header for sorting #include + +// Include the header files for various kernel metadata #include "plugin/device/ascend/kernel/aicpu/aicpu_kernel_metadata.h" #include "plugin/device/ascend/kernel/host/host_kernel_metadata.h" #include "plugin/device/ascend/kernel/rts/rt_kernel_info.h" #include "plugin/device/ascend/kernel/hccl/hccl_kernel_metadata.h" #include "plugin/device/ascend/kernel/tbe/tbe_kernel_select/tbe_kernel_select.h" #include "kernel/akg/akg_kernel_metadata.h" + +// Include the header files for ANF runtime algorithm and ANF utility functions #include "backend/common/session/anf_runtime_algorithm.h" #include "include/common/utils/anfalgo.h" + +// Include the header files for MindSpore context and trace base #include "utils/ms_context.h" #include "utils/trace_base.h" -namespace mindspore { -namespace kernel { -namespace { -void FilterInvalidKernelInfo(const CNodePtr &kernel_node, - std::vector> *kernel_info_list) { - MS_EXCEPTION_IF_NULL(kernel_info_list); - if (kernel_info_list->empty()) { - return; - } - MS_EXCEPTION_IF_NULL(kernel_node); - size_t output_tensor_num = common::AnfAlgo::GetOutputTensorNum(kernel_node); - size_t input_tensor_num = common::AnfAlgo::GetInputTensorNum(kernel_node); - std::vector> filtered_list; - (void)std::copy_if( - kernel_info_list->begin(), kernel_info_list->end(), std::back_inserter(filtered_list), - [output_tensor_num, input_tensor_num](const std::shared_ptr &kernel_build_info) { - return kernel_build_info->GetOutputNum() == output_tensor_num && - kernel_build_info->GetInputNum() == input_tensor_num; - }); - if (!filtered_list.empty()) { - kernel_info_list->clear(); - (void)std::copy(filtered_list.begin(), filtered_list.end(), std::back_inserter(*kernel_info_list)); - } else { - for (size_t index = 0; index < kernel_info_list->size(); ++index) { - std::ostringstream buffer; - auto &kernel_info = kernel_info_list->at(index); - MS_EXCEPTION_IF_NULL(kernel_info); - if (kernel_info->GetOutputNum() != output_tensor_num) { - buffer << "Kernel node's output size [" << output_tensor_num << "]" - << " cannot match the kernel's output size [" << kernel_info->GetOutputNum() << "]"; - } else { - buffer << "Kernel node's input size [" << input_tensor_num << "]" - << " cannot match the kernel's input size [" << kernel_info->GetInputNum() << "]"; - } - MS_LOG(INFO) << "Kernel [ " << index << " ] :" << kernel_info->ToString() << buffer.str(); - } - kernel_info_list->clear(); - MS_LOG(INFO) << "Node: " << kernel_node->DebugString() << "'s output size : [" << output_tensor_num << "]" - << "input size : [" << input_tensor_num << "] can not match any kernelInfo !"; - } +// Start of the `FilterInvalidKernelInfo` function + +// Check if the `kernel_info_list` is null, if so, return immediately +MS_EXCEPTION_IF_NULL(kernel_info_list); + +// Check if the `kernel_info_list` is empty, if so, return immediately +if (kernel_info_list->empty()) { + return; } +// Check if the `kernel_node` is null, if so, throw an exception +MS_EXCEPTION_IF_NULL(kernel_node); + +// Get the number of output tensors from the `kernel_node` +size_t output_tensor_num = common::AnfAlgo::GetOutputTensorNum(kernel_node); + +// Get the number of input tensors from the `kernel_node` +size_t input_tensor_num = common::AnfAlgo::GetInputTensorNum(kernel_node); + +// Create a new vector to store the filtered kernel build info +std::vector> filtered_list; + +// Copy the kernel build info from the `kernel_info_list` to the `filtered_list` if the output and input tensor numbers match +(void)std::copy_if( + kernel_info_list->begin(), kernel_info_list->end(), std::back_inserter(filtered_list), + [output_tensor_num, input_tensor_num](const std::shared_ptr &kernel_build_info) { + return kernel_build_info->GetOutputNum() == output_tensor_num && + kernel_build_info->GetInputNum() == input_tensor_num; + }); + +// Check if the `filtered_list` is not empty +if (!filtered_list.empty()) { + // Continue with the rest of the function + // ... +} +// Clear the kernel_info_list to remove any existing kernel information +kernel_info_list->clear(); + +// If the filtered_list is not empty, copy its contents to the kernel_info_list +// using std::copy and std::back_inserter +(void)std::copy(filtered_list.begin(), filtered_list.end(), std::back_inserter(*kernel_info_list)); + +// If the filtered_list is empty, meaning no matching kernel information was found +else { + // Iterate over the kernel_info_list + for (size_t index = 0; index < kernel_info_list->size(); ++index) { + // Create a string stream buffer to store the error message + std::ostringstream buffer; + + // Get the kernel_info at the current index + auto &kernel_info = kernel_info_list->at(index); + + // Check if the kernel's output size does not match the expected output_tensor_num + if (kernel_info->GetOutputNum() != output_tensor_num) { + buffer << "Kernel node's output size [" << output_tensor_num << "]" + << " cannot match the kernel's output size [" << kernel_info->GetOutputNum() << "]"; + } + // If the kernel's output size matches, check if the input size does not match the expected input_tensor_num + else { + buffer << "Kernel node's input size [" << input_tensor_num << "]" + << " cannot match the kernel's input size [" << kernel_info->GetInputNum() << "]"; + } + + // Log the kernel information and the error message + MS_LOG(INFO) << "Kernel [ " << index << " ] :" << kernel_info->ToString() << buffer.str(); + } + + // Clear the kernel_info_list to remove any existing kernel information + kernel_info_list->clear(); + + // Log the error message indicating that the input and output sizes of the kernel node do not match any kernelInfo + MS_LOG(INFO) << "Node: " << kernel_node->DebugString() << "'s output size : [" << output_tensor_num << "]" + << "input size : [" << input_tensor_num << "] can not match any kernelInfo !"; +} +// Closing brace to end the main function +} + +// Function to check if the given kernel node is a reshape operation in the AICPU task sink bool SelectAicpuReshapeInTaskSink(const CNodePtr &kernel_node) { + // Check if the kernel node is null MS_EXCEPTION_IF_NULL(kernel_node); + + // Check if the name of the kernel node is "Reshape" if (common::AnfAlgo::GetCNodeName(kernel_node) != "Reshape") { return false; } + + // Define the expected size of the kernel node const size_t AicpuReshapeSize = 2; + + // Check if the size of the kernel node matches the expected size if (kernel_node->size() != AicpuReshapeSize) { return false; } + + // Get the instance of the current context auto context_ptr = MsContext::GetInstance(); MS_EXCEPTION_IF_NULL(context_ptr); + + // Get the value of the "enable_task_sink" parameter from the context auto is_task_sink = context_ptr->get_param(MS_CTX_ENABLE_TASK_SINK); + + // Return the value of "is_task_sink" return is_task_sink; } -} // namespace +// End of the namespace + +// A function to check if a list of kernel build information is empty or not void CheckKernelInfoListEmpty(const std::vector> *kernel_info_list, const std::string &type) { - MS_EXCEPTION_IF_NULL(kernel_info_list); - if (kernel_info_list->empty()) { - MS_LOG(INFO) << "Warning: kernel info list is empty, kernel type: " << type; - } + // Check if the kernel info list is null + MS_EXCEPTION_IF_NULL(kernel_info_list); + + // Check if the kernel info list is empty + if (kernel_info_list->empty()) { + // Print a warning message indicating that the kernel info list is empty and the kernel type + MS_LOG(INFO) << "Warning: kernel info list is empty, kernel type: " << type; + } } +// Function to query all kernel information for a given kernel node void KernelQueryAll(const CNodePtr &kernel_node, std::vector> *kernel_info_list) { + // Check if the kernel node and kernel info list pointers are not null MS_EXCEPTION_IF_NULL(kernel_node); MS_EXCEPTION_IF_NULL(kernel_info_list); + + // Query TBE metadata info for the kernel node and add it to the kernel info list TbeMetadataInfo(kernel_node, kernel_info_list); + + // If the kernel info list is still empty, query RT kernel info and check if it is empty if (kernel_info_list->empty()) { GetRtKelInfo(kernel_node, kernel_info_list); CheckKernelInfoListEmpty(kernel_info_list, "RT_Kernel"); } + + // If the kernel info list is still empty, query HCCL kernel info and check if it is empty if (kernel_info_list->empty()) { HcclMetadataInfo(kernel_node, kernel_info_list); CheckKernelInfoListEmpty(kernel_info_list, "HCCL_Kernel"); } + + // If SelectAicpuReshapeInTaskSink returns true, return from the function if (SelectAicpuReshapeInTaskSink(kernel_node)) { return; } + + // If the kernel info list is still empty, query HOST kernel info and check if it is empty if (kernel_info_list->empty()) { HostMetadataInfo(kernel_node, kernel_info_list); CheckKernelInfoListEmpty(kernel_info_list, "HOST_Kernel"); } } +// Closing brace to end the main function +} +// Function to query kernel information for a given kernel node void KernelQuery(const CNodePtr &kernel_node, std::vector> *kernel_info_list, KernelType kernel_type) { + + // Check if the kernel node is null, throw an exception if it is MS_EXCEPTION_IF_NULL(kernel_node); + + // Check if the kernel info list is null, throw an exception if it is MS_EXCEPTION_IF_NULL(kernel_info_list); +} - auto context_ptr = MsContext::GetInstance(); - MS_EXCEPTION_IF_NULL(context_ptr); +// Create a pointer named "context_ptr" and assign it the instance of the MsContext class +auto context_ptr = MsContext::GetInstance(); - const PrimitivePtr kPrimProdForceSeA = std::make_shared("ProdForceSeA"); - if (IsPrimitiveCNode(kernel_node, kPrimProdForceSeA)) { +// Check if the pointer "context_ptr" is null, and throw an exception if it is +MS_EXCEPTION_IF_NULL(context_ptr); + +// Create a shared pointer to a Primitive object named kPrimProdForceSeA and initialize it with the name "ProdForceSeA" +const PrimitivePtr kPrimProdForceSeA = std::make_shared("ProdForceSeA"); + +// Check if the given kernel_node is a CNode and its primitive is equal to kPrimProdForceSeA +if (IsPrimitiveCNode(kernel_node, kPrimProdForceSeA)) { + // If the condition is true, set the kernel_type to AKG_KERNEL kernel_type = KernelType::AKG_KERNEL; - } +} - const PrimitivePtr kPrimLoadIm2Col = std::make_shared("LoadIm2Col"); - if (IsPrimitiveCNode(kernel_node, kPrimLoadIm2Col)) { +// Create a shared pointer to a Primitive object named kPrimLoadIm2Col and initialize it with the name "LoadIm2Col" +const PrimitivePtr kPrimLoadIm2Col = std::make_shared("LoadIm2Col"); + +// Check if the given kernel_node is a CNode and its primitive is equal to kPrimLoadIm2Col +if (IsPrimitiveCNode(kernel_node, kPrimLoadIm2Col)) { + // If the condition is true, set the kernel_type to AKG_KERNEL kernel_type = KernelType::AKG_KERNEL; - } // use LoadIm2Col only for THOR optimizer +} +// The LoadIm2Col primitive is only used for the THOR optimizer + + // Switch statement to determine the action based on the value of kernel_type switch (kernel_type) { case KernelType::AKG_KERNEL: + // If kernel_type is AKG_KERNEL, call AkgMetadataInfo function with kernel_node and kernel_info_list as arguments AkgMetadataInfo(kernel_node, kernel_info_list); break; default: + // If kernel_type is not AKG_KERNEL, call KernelQueryAll function with kernel_node and kernel_info_list as arguments KernelQueryAll(kernel_node, kernel_info_list); break; } - // check output + + // After the switch statement, call FilterInvalidKernelInfo function with kernel_node and kernel_info_list as arguments + // This function is used to check and filter out any invalid kernel information from the list FilterInvalidKernelInfo(kernel_node, kernel_info_list); -} + +// A function to query AI CPU information for a given kernel node and populate a list of kernel build information void AICPUQuery(const CNodePtr &kernel_node, std::vector> *kernel_info_list) { + // Check if the kernel node is null MS_EXCEPTION_IF_NULL(kernel_node); + + // Check if the kernel info list is null MS_EXCEPTION_IF_NULL(kernel_info_list); + + // Clear the kernel info list kernel_info_list->clear(); + + // Call the AicpuMetadataInfo function to populate the kernel info list with AI CPU metadata information AicpuMetadataInfo(kernel_node, kernel_info_list); + + // Call the FilterInvalidKernelInfo function to filter out any invalid kernel info from the list FilterInvalidKernelInfo(kernel_node, kernel_info_list); } +// Function to check if a kernel is supported by AICPU bool IsSupportedByAICPU(const AnfNodePtr &kernel_node, const KernelBuildInfoPtr &select_kernel_build_info) { + // Check if the kernel node and kernel build info are not null MS_EXCEPTION_IF_NULL(kernel_node); MS_EXCEPTION_IF_NULL(select_kernel_build_info); + + // Create a vector to store kernel build info std::vector> kernel_info_list; + + // Cast the kernel node to CNodePtr auto cnode = kernel_node->cast(); MS_EXCEPTION_IF_NULL(cnode); + + // Call the AICPUQuery function to populate the kernel info list AICPUQuery(cnode, &kernel_info_list); + + // Check if any of the kernel build info in the list is similar to the select kernel build info return std::any_of(kernel_info_list.begin(), kernel_info_list.end(), [&select_kernel_build_info](const kernel::KernelBuildInfoPtr item) { MS_EXCEPTION_IF_NULL(item); @@ -167,18 +289,36 @@ bool IsSupportedByAICPU(const AnfNodePtr &kernel_node, const KernelBuildInfoPtr }); } +// Check if a given kernel node is supported by an AI core bool IsSupportedByAICore(const AnfNodePtr &kernel_node, const KernelBuildInfoPtr &select_kernel_build_info) { + // Check if the kernel node is null MS_EXCEPTION_IF_NULL(kernel_node); + // Check if the selected kernel build info is null MS_EXCEPTION_IF_NULL(select_kernel_build_info); + + // Create a vector to store kernel build info std::vector> kernel_info_list; + + // Cast the kernel node to a CNode pointer auto cnode = kernel_node->cast(); + // Check if the cast was successful MS_EXCEPTION_IF_NULL(cnode); + + // Get the TBE metadata info for the CNode and store it in the kernel info list TbeMetadataInfo(cnode, &kernel_info_list); + + // Check if any of the kernel build info in the list matches the selected kernel build info return std::any_of(kernel_info_list.begin(), kernel_info_list.end(), [&select_kernel_build_info](const kernel::KernelBuildInfoPtr item) { + // Check if the current kernel build info is null MS_EXCEPTION_IF_NULL(item); + // Compare the current kernel build info with the selected kernel build info return *item == *select_kernel_build_info; }); } + +// End of the kernel namespace } // namespace kernel -} // namespace mindspore + +// End of the mindspore namespace +} // namespace mindspore \ No newline at end of file -- 2.34.1