transform/util.cc

525 lines
26 KiB
C++
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

/**
* Copyright 2019 Huawei Technologies Co., Ltd
*
* Licensed under the Apache License, Version 2.0 (the "License");
* you may not use this file except in compliance with the License.
* You may obtain a copy of the License at
*
* http://www.apache.org/licenses/LICENSE-2.0
*
* Unless required by applicable law or agreed to in writing, software
* distributed under the License is distributed on an "AS IS" BASIS,
* WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
* See the License for the specific language governing permissions and
* limitations under the License.
*/
#include "include/transform/graph_ir/util.h"
#include <utility>
#include <map>
#include "securec/include/securec.h"
#include "include/common/utils/convert_utils.h"
#include "include/common/utils/utils.h"
namespace mindspore {
namespace transform {
using std::make_shared;
using std::shared_ptr;
using std::string;
using std::vector;
const size_t kErrorSize = 0;
//该函数主要用于生成一个包含相同 数据的 int64_tvector方便在一些情况下使用。
vector<int64_t> TransformUtil::ConvertIntToList(int64_t data, int size) {
vector<int64_t> list{}; //创建一个空的 类型的 vector<int64_t> list。
if (size <= 0) { //检查 size 是否小于等于 0
MS_LOG(WARNING) << "size <= 0"; //如果是,输出警告日志并返回空的 list。
return list;
}
for (int i = 0; i < size; ++i) { // 使用 for 循环将 data 复制 size 次并将每次复制的结果添加到list中。
list.push_back(data);
}
return list;//循环结束后返回存储有复制结果的list
}
static std::map<MeDataType, GeDataType> datatype_trans_map = {
{MeDataType::kNumberTypeFloat16, GeDataType::DT_FLOAT16}, {MeDataType::kNumberTypeFloat32, GeDataType::DT_FLOAT},
{MeDataType::kNumberTypeFloat64, GeDataType::DT_DOUBLE}, {MeDataType::kNumberTypeInt8, GeDataType::DT_INT8},
{MeDataType::kNumberTypeInt16, GeDataType::DT_INT16}, {MeDataType::kNumberTypeInt32, GeDataType::DT_INT32},
{MeDataType::kNumberTypeInt64, GeDataType::DT_INT64}, {MeDataType::kNumberTypeUInt8, GeDataType::DT_UINT8},
{MeDataType::kNumberTypeUInt16, GeDataType::DT_UINT16}, {MeDataType::kNumberTypeUInt32, GeDataType::DT_UINT32},
{MeDataType::kNumberTypeUInt64, GeDataType::DT_UINT64}, {MeDataType::kNumberTypeBool, GeDataType::DT_BOOL}};
// ConvertDataType 函数用于将 MeDataType 类型的数据转换为 GeDataType 类型。
GeDataType TransformUtil::ConvertDataType(const MeDataType &type) {
MS_LOG(DEBUG) << "Convert me data type: " << TypeIdLabel(type) << " to ge data type";
// 在 datatype_trans_map 中查找 MeDataType 对应的 GeDataType。
if (datatype_trans_map.find(type) != datatype_trans_map.end()) {
return datatype_trans_map[type];
} else { // 如果找不到对应的映射关系,则返回 GeDataType::DT_UNDEFINED。
return GeDataType::DT_UNDEFINED;
}
}
static std::map<MeDataType, size_t> datatype_size_map = {
{MeDataType::kNumberTypeFloat16, sizeof(float) / 2}, {MeDataType::kNumberTypeFloat32, sizeof(float)}, // 1/2 of float
{MeDataType::kNumberTypeFloat64, sizeof(double)}, {MeDataType::kNumberTypeInt8, sizeof(int8_t)},
{MeDataType::kNumberTypeInt16, sizeof(int16_t)}, {MeDataType::kNumberTypeInt32, sizeof(int32_t)},
{MeDataType::kNumberTypeInt64, sizeof(int64_t)}, {MeDataType::kNumberTypeUInt8, sizeof(uint8_t)},
{MeDataType::kNumberTypeUInt16, sizeof(uint16_t)}, {MeDataType::kNumberTypeUInt32, sizeof(uint32_t)},
{MeDataType::kNumberTypeUInt64, sizeof(uint64_t)}, {MeDataType::kNumberTypeBool, sizeof(bool)}};
// GetDataTypeSize 函数用于获取给定 MeDataType 类型数据在内存中占据的字节大小。
size_t TransformUtil::GetDataTypeSize(const MeDataType &type) {
if (datatype_size_map.find(type) != datatype_size_map.end()) { // 在 datatype_size_map 中查找 MeDataType 对应的字节大小。
return datatype_size_map[type];
} else { // 如果找不到对应的大小输出错误日志并返回一个特定的错误大小kErrorSize
MS_LOG(ERROR) << "Illegal tensor data type!";
return kErrorSize;
}
}
// ConvertFormat 函数用于将给定的字符串格式 `format` 转换为对应的 GeFormat 枚举类型。
GeFormat TransformUtil::ConvertFormat(const string &format) {
// 通过比较字符串格式 `format`,判断对应的 GeFormat 枚举类型,并进行转换并返回。
if (format == kOpFormat_NCHW) {
return GeFormat::FORMAT_NCHW;
} else if (format == kOpFormat_NDHWC) {
return GeFormat::FORMAT_NDHWC;
} else if (format == kOpFormat_NCDHW) {
return GeFormat::FORMAT_NCDHW;
} else if (format == kOpFormat_DHWNC) {
return GeFormat::FORMAT_DHWNC;
} else if (format == kOpFormat_DHWCN) {
return GeFormat::FORMAT_DHWCN;
} else if (format == kOpFormat_NC1HWC0) {
return GeFormat::FORMAT_NC1HWC0;
} else if (format == kOpFormat_NHWC) {
return GeFormat::FORMAT_NHWC;
} else if (format == kOpFormat_HWCN) {
return GeFormat::FORMAT_HWCN;
} else if (format == kOpFormat_ND) {
return GeFormat::FORMAT_ND;
} else { // 如果 `format` 不是支持的数据格式,输出错误日志并返回默认的 GeFormat::FORMAT_ND 格式。
MS_LOG(ERROR) << "Illegal tensor data format: (" << format << "). Use ND format instead.";
return GeFormat::FORMAT_ND;
}
}
static int64_t IntegerCastFunc(size_t temp) { return static_cast<int64_t>(temp); }
// GetGeTensorDesc 函数用于根据给定的 MeTensor 的形状ShapeVector、数据类型MeDataType和格式format
// 创建对应的 GeTensorDesc 对象,并返回一个指向该对象的 shared_ptr。
// GeTensorDesc 是 Ascend AI Core 引擎中定义的张量描述类,用于描述张量的形状、数据类型和数据格式。
std::shared_ptr<GeTensorDesc> TransformUtil::GetGeTensorDesc(const ShapeVector &me_shape, const MeDataType &me_type,
const std::string &format) {
// convert me shape to ge shape
// 将 MeTensor 的形状ShapeVector转换为 GeShape 对象std::vector<int64_t>
std::vector<int64_t> ge_shape;
if (me_shape.size() == 1) {
ge_shape.push_back(static_cast<int64_t>(me_shape[0]));
} else {
ge_shape.resize(me_shape.size());
(void)std::transform(me_shape.begin(), me_shape.end(), ge_shape.begin(), IntegerCastFunc);
}
GeShape shape(ge_shape);
if (shape.GetDimNum() == 0) { // 如果 GeShape 对象的维度数为 0则输出提示信息日志。
MS_LOG(INFO) << "The dims size of Ge tensor is zero";
}
// convert me format to ge format
// 将 MeTensor 的格式format转换为 GeFormat 枚举类型。
GeFormat ge_format = ConvertFormat(format);
if (ge_format == GeFormat::FORMAT_ND) {
MS_LOG(INFO) << "Set ND data format";
}
// convert me datatype to ge datatype
// 将 MeTensor 的数据类型me_type转换为 GeDataType 枚举类型。
GeDataType data_type = ConvertDataType(me_type);
if (data_type == GeDataType::DT_UNDEFINED) { // 如果数据类型转换失败,则输出错误日志,并返回空指针。
MS_LOG(ERROR) << "undefined data type :" << me_type;
return nullptr;
}
// 创建 GeTensorDesc 对象,并设置相应的形状、数据类型和数据格式信息。
auto desc = std::make_shared<GeTensorDesc>(shape, ge_format, data_type);
if (desc == nullptr) {
MS_LOG(ERROR) << "Create GeTensorDesc failed!";
return nullptr;
}
// 设置实际维度数量,即 GeTensorDesc 对象的形状维度数。
MS_LOG(INFO) << "SetRealDimCnt is :" << me_shape.size();
desc->SetRealDimCnt(SizeToInt(me_shape.size()));
return desc; // 返回指向创建的 GeTensorDesc 对象的 shared_ptr。
}
// if failed, return empty vector.
// ConvertInputTensors 函数用于将给定的 MeTensor 列表me_tensors转换为对应的 GeTensor 列表,并返回转换后的结果。
// MeTensor 是 MindSpore 引擎中定义的张量类,用于存储张量的数据和相关信息。
// GeTensor 是 Ascend AI Core 引擎中定义的张量类,用于存储张量的数据和相关信息。
// 函数遍历输入的 MeTensor 列表,对每个 MeTensor 进行转换,并将转换后的 GeTensor 添加到结果列表 ge_tensors 中。
// 如果在转换过程中遇到错误,则输出相应的错误日志,并返回空列表。
std::vector<GeTensorPtr> TransformUtil::ConvertInputTensors(const std::vector<MeTensorPtr> &me_tensors,
const std::string &format) {
std::vector<GeTensorPtr> ge_tensors;
// 遍历输入的 MeTensor 列表,对每个 MeTensor 进行转换,并将转换后的 GeTensor 添加到结果列表 ge_tensors 中。
for (size_t index = 0; index < me_tensors.size(); index++) {
MS_EXCEPTION_IF_NULL(me_tensors[index]);
// 输出当前 MeTensor 的数据大小、形状和数据类型信息。
MS_LOG(INFO) << "me_tensor " << index << " 's data size is: " << me_tensors[index]->DataSize();
auto shape = me_tensors[index]->shape();
std::string shape_str;
for (size_t i = 0; i < shape.size(); i++) {
shape_str += std::to_string(shape[i]);
shape_str += " ";
}
MS_LOG(INFO) << "me_tensor " << index << " 's shape is: { " << shape_str << "}";
MS_LOG(INFO) << "me_tensor " << index << " 's type is: " << me_tensors[index]->data_type();
// 调用 ConvertTensor 函数将当前的 MeTensor 转换为对应的 GeTensor。
auto ge_tensor_ptr = TransformUtil::ConvertTensor(me_tensors[index], format);
if (ge_tensor_ptr != nullptr) { // 如果转换成功,则将转换后的 GeTensor 添加到结果列表 ge_tensors 中。
ge_tensors.emplace_back(ge_tensor_ptr);
} else { // 如果转换过程中遇到错误,则输出相应的错误日志,并清空结果列表 ge_tensors并返回空列表。
MS_LOG(ERROR) << "Convert me_tensor " << index << " to Ge Tensor failed!";
ge_tensors.clear();
return ge_tensors;
}
}
return ge_tensors; // 返回转换后的 GeTensor 列表。
}
// ConvertTensor 函数用于将给定的 MeTensor`tensor`)转换为对应的 GeTensor并返回转换后的结果。
// MeTensor 是 MindSpore 引擎中定义的张量类,用于存储张量的数据和相关信息。
// GeTensor 是 Ascend AI Core 引擎中定义的张量类,用于存储张量的数据和相关信息。
GeTensorPtr TransformUtil::ConvertTensor(const MeTensorPtr &tensor, const std::string &format) {
// get tensor data type size
// 获取 MeTensor 的数据类型大小type_size即数据类型占用的字节数。
MS_EXCEPTION_IF_NULL(tensor);
size_t type_size = GetDataTypeSize(tensor->data_type());
if (type_size == kErrorSize) { // 如果数据类型大小获取失败,则输出错误日志,并返回空指针。
MS_LOG(ERROR) << "The Me Tensor data type size is wrong, type size is: " << type_size;
return nullptr;
}
// 获取 MeTensor 的元素数量和数据缓冲区大小data_buff_size
size_t elements_num = IntToSize(tensor->ElementsNum());
// get tensor buff size
size_t data_buff_size = elements_num * type_size;
if (data_buff_size == 0) { // 如果数据缓冲区大小为 0则输出提示信息日志。
MS_LOG(INFO) << "The Me Tensor data buff size is 0.";
}
// create ge tensor
// 创建 GeTensorDesc 对象,并根据 MeTensor 的形状、数据类型和数据格式信息创建对应的 GeTensor 对象。
auto desc = GetGeTensorDesc(tensor->shape_c(), tensor->data_type(), format);
if (desc == nullptr) { // 如果创建 GeTensorDesc 对象失败,则输出错误日志,并返回空指针。
MS_LOG(ERROR) << "Failed to get Tensor Desc";
return nullptr;
}
// 创建 GeTensor 对象,并设置相应的数据缓冲区、形状、数据类型等信息。
GeTensorPtr tensor_ptr = make_shared<GeTensor>(*desc, static_cast<uint8_t *>(tensor->data_c()), data_buff_size);
if (tensor_ptr != nullptr) { // 如果创建 GeTensor 对象成功,则输出转换成功的提示信息,并返回指向创建的 GeTensor 对象的 shared_ptr。
MS_LOG(INFO) << "Convert Me Tensor to Ge Tensor success!";
}
return tensor_ptr;
}
// ConvertGeTensors 函数用于将给定的 GeTensor`ge_tensors`)向量转换为对应的 MeTensorMindSpore
// 引擎中的张量)向量,并返回转换后的结果。 `ge_tensors` 是 Ascend AI Core 引擎中的张量向量,用于存储计算结果的张量。
// `request_dims` 是一个与 `ge_tensors` 同样大小的向量,用于指定转换后的 MeTensor 的形状ShapeVector
std::vector<MeTensorPtr> TransformUtil::ConvertGeTensors(const std::vector<GeTensorPtr> &ge_tensors,
const std::vector<ShapeVector> &request_dims) {
std::vector<MeTensorPtr> outputs;
// 遍历 `ge_tensors` 向量中的每个 GeTensor将其转换为对应的 MeTensor 对象
for (size_t index = 0; index < ge_tensors.size(); index++) {
MeTensorPtr me_tensor_ptr = nullptr;
// 根据索引值 `index` 来获取对应的请求形状(`request_dims`)。
if (index < request_dims.size()) {
me_tensor_ptr = ConvertGeTensor(ge_tensors[index], request_dims[index]);
} else { // 如果请求形状的向量长度小于当前索引 `index`,则使用空的形状向量来进行转换。
ShapeVector empty_shape;
me_tensor_ptr = ConvertGeTensor(ge_tensors[index], empty_shape);
}
if (me_tensor_ptr != nullptr) { // 如果转换成功,则将转换后的 MeTensor 存储在输出向量 `outputs` 中。
outputs.emplace_back(me_tensor_ptr);
} else { // 如果转换失败,则输出相应的错误日志,并返回已经成功转换的 MeTensor 向量 `outputs`。
MS_LOG(ERROR) << "Convert Ge Tensor " << index << " to Me Tensor failed!";
return outputs;
}
}
return outputs; // 返回转换后的 MeTensor 向量 `outputs`。
}
// ConvertGeTensors 函数用于将给定的 GeTensor`ge_tensors`)向量转换为对应的 MeTensorMindSpore 引擎中的张量)向量,并返回转换后的结果。
//`ge_tensors` 是 Ascend AI Core 引擎中的张量向量,用于存储计算结果的张量。
std::vector<MeTensorPtr> TransformUtil::ConvertGeTensors(const std::vector<GeTensorPtr> &ge_tensors) {
std::vector<MeTensorPtr> outputs;
// 遍历 `ge_tensors` 向量中的每个 GeTensor将其转换为对应的 MeTensor 对象。
for (size_t index = 0; index < ge_tensors.size(); index++) {
MeTensorPtr me_tensor_ptr = ConvertGeTensor(ge_tensors[index]);
if (me_tensor_ptr != nullptr) { // 如果转换成功,则将转换后的 MeTensor 存储在输出向量 `outputs` 中。
outputs.emplace_back(me_tensor_ptr);
} else { // 如果转换失败,则输出相应的错误日志,并返回已经成功转换的 MeTensor 向量 `outputs`。
MS_LOG(ERROR) << "Convert Ge Tensor " << index << " to Me Tensor failed!";
return outputs;
}
}
return outputs; // 返回转换后的 MeTensor 向量 `outputs`。
}
// ConvertGeDataType 函数用于将给定的 GeDataTypeAscend AI Core 引擎中的数据类型)转换为对应的 MeDataTypeMindSpore引擎中的数据类型
//`type` 是 Ascend AI Core 引擎中的数据类型,需要被转换为对应的 MeDataType。
MeDataType TransformUtil::ConvertGeDataType(const GeDataType &type) {
switch (type) { // 对于不同的 GeDataType 值,根据其对应的数据类型进行转换。
case GeDataType::DT_FLOAT16:
return MeDataType::kNumberTypeFloat16;
case GeDataType::DT_FLOAT:
return MeDataType::kNumberTypeFloat32;
case GeDataType::DT_DOUBLE:
return MeDataType::kNumberTypeFloat64;
case GeDataType::DT_INT64:
return MeDataType::kNumberTypeInt64;
case GeDataType::DT_INT32:
return MeDataType::kNumberTypeInt32;
case GeDataType::DT_INT16:
return MeDataType::kNumberTypeInt16;
case GeDataType::DT_INT8:
return MeDataType::kNumberTypeInt8;
case GeDataType::DT_BOOL:
return MeDataType::kNumberTypeBool;
case GeDataType::DT_UINT8:
return MeDataType::kNumberTypeUInt8;
case GeDataType::DT_UINT16:
return MeDataType::kNumberTypeUInt16;
case GeDataType::DT_UINT32:
return MeDataType::kNumberTypeUInt32;
case GeDataType::DT_UINT64:
return MeDataType::kNumberTypeUInt64;
// 对于其他未列出的 GeDataType 值,或者无法转换的值,返回 MeDataType::kTypeUnknown表示未知数据类型。
case GeDataType::DT_UNDEFINED:
case GeDataType::DT_DUAL_SUB_UINT8:
case GeDataType::DT_DUAL_SUB_INT8:
case GeDataType::DT_DUAL:
return MeDataType::kTypeUnknown;
default:
return MeDataType::kTypeUnknown;
}
}
namespace {
// IsGeShapeCompatible 函数用于检查给定的 GeTensor 的形状 `ge_shape` 是否与请求的形状 `request_dims` 兼容。
// `ge_shape` 是 GeTensor 的形状,`request_dims` 是请求的形状。
bool IsGeShapeCompatible(const GeShape &ge_shape, const ShapeVector &request_dims) {
MS_LOG(INFO) << "GeTensor's shape is " << TransformUtil::PrintVector(ge_shape.GetDims());
MS_LOG(INFO) << "Me request shape is " << TransformUtil::PrintVector(request_dims);
const int GE_DIMS = 4;
std::vector<int64_t> ge_dims = ge_shape.GetDims();
if (request_dims.size() > ge_dims.size()) { // 如果请求的维度数量大于 GeTensor 的维度数量,说明形状不兼容,返回 false。
MS_LOG(ERROR) << "Request shape's dims count greater than ge shape's";
return false;
}
// convert NHWC to NCHW
// 如果请求的维度数量等于 1且 GeTensor 的维度数量等于 4且对应维度值满足 NHWC 到 NCHW 的转换条件,返回兼容。
if ((request_dims.size() == 1) && (ge_dims.size() == GE_DIMS) && (request_dims[0] == ge_dims[1]) &&
(ge_dims[0] == 1) && (ge_dims[2] == 1) && (ge_dims[3] == 1)) {
MS_LOG(INFO) << "Ge tensor shape and request shape is compatible";
return true;
}
std::string::size_type i = 0;
// 逐一比较 `request_dims` 和 `ge_shape` 的维度值,如果任何维度值不相等,说明形状不兼容,返回 false。
for (; i < request_dims.size(); i++) {
if (ge_dims[i] != request_dims[i]) {
MS_LOG(ERROR) << "Request shape's dims value not equal to ge shape's";
return false;
}
}
// 对于 `request_dims` 中已比较的维度后的未提供的维度,要求其在 `ge_shape` 中对应的维度值为1否则说明形状不兼容返回 false。
for (; i < ge_dims.size(); i++) {
if (ge_dims[i] != 1) {
MS_LOG(ERROR) << "GeShape's extend dims is not equal to 1";
return false;
}
}
// 形状兼容,返回 true。
MS_LOG(INFO) << "Ge tensor shape and request shape is compatible";
return true;
}
} // namespace
// ConvertMeShape 函数用于将 MeTensor 的形状表示转换为 GeTensor 的形状表示。
GeShape TransformUtil::ConvertMeShape(const ShapeVector &me_dims) {
std::vector<int64_t> ge_dims;
//将 `me_dims` 中的维度值拷贝到一个新的 vector `ge_dims` 中,并使用这个新的 vector 创建一个 GeShape 对象
(void)std::copy(me_dims.begin(), me_dims.end(), std::back_inserter(ge_dims));
return GeShape(ge_dims); //返回 GeShape 对象,表示 MeTensor 形状转换为 GeTensor 形状的结果。
}
// ConvertGeShape 函数用于将 GeTensor 的形状表示转换为 MeTensor 的形状表示。
ShapeVector TransformUtil::ConvertGeShape(const GeShape &ge_shape) {
//将 GeShape 对象中的维度值拷贝到一个新的 ShapeVector `me_dims` 中,并使用这个新的 ShapeVector 表示 MeTensor
ShapeVector me_dims;
std::vector<int64_t> ge_dims = ge_shape.GetDims();
(void)std::copy(ge_dims.begin(), ge_dims.end(), std::back_inserter(me_dims));
return me_dims; //返回 `me_dims`,表示 GeTensor 形状转换为 MeTensor 形状的结果
}
// ConvertGeShape 函数用于将 GeTensor 的形状表示转换为 MeTensor 的形状表示。
// 参数 `ge_shape` 是 GeTensor 的形状,以 GeShape 对象表示。
// 参数 `request_dims` 是 MeTensor 请求的形状,以 ShapeVector 表示。
ShapeVector TransformUtil::ConvertGeShape(const GeShape &ge_shape, const ShapeVector &request_dims) {
vector<int64_t> ret;
if (ge_shape.GetDimNum() == 0) {
MS_LOG(DEBUG) << "GeTensor's shape is scalar";
return ret;
}
if (IsGeShapeCompatible(ge_shape, request_dims) == true) { //比较 GeShape 和 MeTensor 请求形状是否兼容
ret = request_dims; //如果兼容则返回 MeTensor 请求形状 `request_dims`
} else { //否则返回 GeTensor 形状转换后的 MeTensor 形状 `ret`
MS_LOG(ERROR) << "GeShape and Me request shape are incompatible, return GeShape";
ret = ConvertGeShape(ge_shape);
}
return ret;
}
// GenerateMeTensor 函数用于根据给定的 GeTensor 对象 `ge_tensor`,生成相应的 MeTensor 对象并将其形状和数据复制到MeTensor 中。
// 参数 `ge_tensor` 是给定的 GeTensor 对象,表示待转换的 GeTensor。
// 参数 `me_dims` 是 MeTensor 的形状以ShapeVector 表示。
// 参数 `me_type` 是 MeTensor 的数据类型,以 TypeId 表示。
MeTensorPtr TransformUtil::GenerateMeTensor(const GeTensorPtr &ge_tensor, const ShapeVector &me_dims,
const TypeId &me_type) {
MeTensor me_tensor(me_type, me_dims);
// Get the writable data pointer of the tensor and cast it to its data type
// 获取 MeTensor 的可写数据指针,并将其转换为指定的数据类型
auto me_data_ptr = reinterpret_cast<uint8_t *>(me_tensor.data_c());
size_t me_data_size = static_cast<size_t>(me_tensor.data().nbytes());
MS_EXCEPTION_IF_NULL(me_data_ptr);
MS_EXCEPTION_IF_NULL(ge_tensor);
if (me_data_size < ge_tensor->GetSize()) {
MS_LOG(ERROR) << "ME tensor data size[" << me_data_size << " bytes] is less than GE tensor ["
<< ge_tensor->GetSize() << " bytes]";
return nullptr;
}
// Copy or use the writable data pointer of the ME tensor
// 复制或使用 MeTensor 的可写数据指针
MS_EXCEPTION_IF_NULL(ge_tensor->GetData());
if (ge_tensor->GetSize() == 0) {
MS_LOG(ERROR) << "GE tensor data size is zero!";
return nullptr;
}
// Use memcpy here, not memcpy_s, just because the size of ge_tensor may be bigger than 2GB
// which is the size limit of memcpy_s
// 使用 memcpy 进行数据拷贝,而不使用 memcpy_s这是因为 ge_tensor 的大小可能大于 2GB
// 而 memcpy_s 有 2GB 的大小限制。
(void)memcpy(me_data_ptr, ge_tensor->GetData(), ge_tensor->GetSize());
return make_shared<MeTensor>(me_tensor);
}
// ConvertGeTensor 函数用于将给定的 GeTensor 对象 `ge_tensor` 转换为相应的 MeTensor 对象。
// 参数 `ge_tensor` 是给定的 GeTensor 对象,表示待转换的 GeTensor。
MeTensorPtr TransformUtil::ConvertGeTensor(const GeTensorPtr &ge_tensor) {
MS_EXCEPTION_IF_NULL(ge_tensor);
// 获取 GeTensor 的形状,并将其转换为 MeTensor 的形状
GeShape ge_shape = ge_tensor->GetTensorDesc().GetShape();
vector<int64_t> me_dims = ConvertGeShape(ge_shape);
// 获取 GeTensor 的数据类型,并将其转换为 MeTensor 的数据类型
TypeId type_id = ConvertGeDataType(ge_tensor->GetTensorDesc().GetDataType());
if (type_id == MeDataType::kTypeUnknown) {
MS_LOG(ERROR) << "Could not convert Ge Tensor because of unsupported data type: "
<< static_cast<int>(ge_tensor->GetTensorDesc().GetDataType());
return nullptr;
}
// 调用 GenerateMeTensor 函数,根据转换后的 MeTensor 形状和数据类型,生成相应的 MeTensor 对象
return GenerateMeTensor(ge_tensor, me_dims, type_id);
}
// if request_dims is empty, use ge tensor's shape,otherwise convert to request shape
// ConvertGeTensor 函数用于将给定的 GeTensor 对象 `ge_tensor` 转换为相应的 MeTensor 对象,并根据给定的 MeTensor 形状 `request_dims` 进行转换。
// 参数 `ge_tensor` 是给定的 GeTensor 对象,表示待转换的 GeTensor。
// 参数 `request_dims` 是要求的 MeTensor 形状,用于与 GeTensor 的形状进行兼容性检查和转换。
MeTensorPtr TransformUtil::ConvertGeTensor(const GeTensorPtr ge_tensor, const ShapeVector &request_dims) {
MS_EXCEPTION_IF_NULL(ge_tensor);
// 获取 GeTensor 的形状,并将其与给定的 MeTensor 形状 `request_dims` 进行兼容性检查和转换
GeShape ge_shape = ge_tensor->GetTensorDesc().GetShape();
vector<int64_t> me_dims = ConvertGeShape(ge_shape, request_dims);
// 输出 GeTensor 的数据类型
MS_LOG(INFO) << "GE tensor type is " << static_cast<int>(ge_tensor->GetTensorDesc().GetDataType());
// Create a tensor with wanted data type and shape
// 创建具有指定数据类型和形状的 MeTensor 对象
TypeId type_id = ConvertGeDataType(ge_tensor->GetTensorDesc().GetDataType());
if (type_id == MeDataType::kTypeUnknown) {
MS_LOG(ERROR) << "Could not convert Ge Tensor because of unsupported data type: "
<< static_cast<int>(ge_tensor->GetTensorDesc().GetDataType());
return nullptr; //如果转换失败,则返回 nullptr
}
return GenerateMeTensor(ge_tensor, me_dims, type_id); //返回其指针,表示转换成功
}
//PrintGeTensor 函数用于打印给定的 GeTensor 对象的数据内容
std::string TransformUtil::PrintGeTensor(const GeTensorPtr ge_tensor) {
std::string ret;
if (ge_tensor == nullptr) { //检查输入的 ge_tensor 是否为空
MS_LOG(ERROR) << "Input ge tensor is nullptr"; //如果为空,则输出错误日志并返回空字符串
return ret;
}
//获取了ge_tensor 的数据类型,并根据数据类型使用 MakeVector 函数将 ge_tensor 的数据内容转换成相应的向量
//根据数据类型的不同,函数将数据内容转换成不同类型的向量,包括 uint32_t、float_t、int32_t、double_t、int64_t、uint64_t、int16_t、uint16_t、int8_t 和 uint8_t。
//使用 PrintVector 函数将转换后的向量打印成字符串,并返回该字符串
MS_LOG(INFO) << "Ge Tensor data type is : " << static_cast<int>(ge_tensor->GetTensorDesc().GetDataType());
switch (static_cast<int>(ge_tensor->GetTensorDesc().GetDataType())) {
case GeDataType::DT_UINT32:
ret = PrintVector(MakeVector<uint32_t>(ge_tensor->GetData(), ge_tensor->GetSize()));
break;
case GeDataType::DT_FLOAT:
ret = PrintVector(MakeVector<float_t>(ge_tensor->GetData(), ge_tensor->GetSize()));
break;
case GeDataType::DT_INT32:
ret = PrintVector(MakeVector<int32_t>(ge_tensor->GetData(), ge_tensor->GetSize()));
break;
case GeDataType::DT_DOUBLE:
ret = PrintVector(MakeVector<double_t>(ge_tensor->GetData(), ge_tensor->GetSize()));
break;
case GeDataType::DT_INT64:
ret = PrintVector(MakeVector<int64_t>(ge_tensor->GetData(), ge_tensor->GetSize()));
break;
case GeDataType::DT_UINT64:
ret = PrintVector(MakeVector<uint64_t>(ge_tensor->GetData(), ge_tensor->GetSize()));
break;
case GeDataType::DT_INT16:
ret = PrintVector(MakeVector<int16_t>(ge_tensor->GetData(), ge_tensor->GetSize()));
break;
case GeDataType::DT_UINT16:
ret = PrintVector(MakeVector<uint16_t>(ge_tensor->GetData(), ge_tensor->GetSize()));
break;
case GeDataType::DT_DUAL_SUB_INT8:
case GeDataType::DT_INT8:
ret = PrintVector(MakeVector<int8_t>(ge_tensor->GetData(), ge_tensor->GetSize()));
break;
case GeDataType::DT_UINT8:
case GeDataType::DT_DUAL_SUB_UINT8:
ret = PrintVector(MakeVector<uint8_t>(ge_tensor->GetData(), ge_tensor->GetSize()));
break;
case GeDataType::DT_FLOAT16:
case GeDataType::DT_BOOL:
case GeDataType::DT_UNDEFINED:
case GeDataType::DT_DUAL:
//如果给定的 ge_tensor 的数据类型不在上述支持的类型列表中(如 DT_FLOAT16、DT_BOOL、DT_UNDEFINED 和 DT_DUAL则输出错误日志并返回空字符串。
default:
MS_LOG(ERROR) << "Unsupported to print type:" << static_cast<int>(ge_tensor->GetTensorDesc().GetDataType())
<< " ge tensor";
break;
}
return ret;
}
} // namespace transform
} // namespace mindspore