525 lines
26 KiB
C++
525 lines
26 KiB
C++
/**
|
||
* 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`)向量转换为对应的 MeTensor(MindSpore
|
||
// 引擎中的张量)向量,并返回转换后的结果。 `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`)向量转换为对应的 MeTensor(MindSpore 引擎中的张量)向量,并返回转换后的结果。
|
||
//`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 函数用于将给定的 GeDataType(Ascend AI Core 引擎中的数据类型)转换为对应的 MeDataType(MindSpore引擎中的数据类型)。
|
||
//`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
|