ADD file via upload

This commit is contained in:
saltyfish 2023-10-03 09:22:15 +08:00
parent 5b1b14448a
commit fef5b1745f
1 changed files with 116 additions and 0 deletions

View File

@ -0,0 +1,116 @@
# Copyright 2019-2022 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.
# ==============================================================================
"""
Define the data types.
定义数据类型
"""
import numpy as np
import mindspore._c_dataengine as cde
from mindspore._c_expression import typing
import mindspore.common.dtype as mstype
def nptype_to_detype(type_):
"""
Get de data type corresponding to numpy dtype.
获取numpy dtype对应的de数据类型
Args:
type_ (numpy.dtype): Numpy's dtype.
Returns:
The data type of de.
"""
# 如果传入的 'type_' 不是 NumPy 数据类型对象np.dtype则将其转换为 np.dtype 对象
if not isinstance(type_, np.dtype):
type_ = np.dtype(type_)
# 创建一个字典,将 NumPy 数据类型映射到 CDEMindSpore 数据增强库)的数据类型
# 这个字典用于将 NumPy 数据类型转换为 CDE 数据类型
return {
np.dtype("bool"): cde.DataType("bool"),
np.dtype("int8"): cde.DataType("int8"),
np.dtype("int16"): cde.DataType("int16"),
np.dtype("int32"): cde.DataType("int32"),
np.dtype("int64"): cde.DataType("int64"),
np.dtype("uint8"): cde.DataType("uint8"),
np.dtype("uint16"): cde.DataType("uint16"),
np.dtype("uint32"): cde.DataType("uint32"),
np.dtype("uint64"): cde.DataType("uint64"),
np.dtype("float16"): cde.DataType("float16"),
np.dtype("float32"): cde.DataType("float32"),
np.dtype("float64"): cde.DataType("float64"),
np.dtype("str"): cde.DataType("string"),
}.get(type_)
def mstype_to_detype(type_):
"""
Get de data type corresponding to mindspore dtype.
获取mindspore数据类型对应的de数据类型
Args:
type_ (mindspore.dtype): MindSpore's dtype.
Returns:
The data type of de.
"""
# 如果传入的 'type_' 不是 NumPy 数据类型对象np.dtype则将其转换为 np.dtype 对象
if not isinstance(type_, np.dtype):
type_ = np.dtype(type_)
# 创建一个字典,将 NumPy 数据类型映射到 CDEMindSpore 数据增强库)的数据类型
# 这个字典用于将 NumPy 数据类型转换为 CDE 数据类型
return {
np.dtype("bool"): cde.DataType("bool"),
np.dtype("int8"): cde.DataType("int8"),
np.dtype("int16"): cde.DataType("int16"),
np.dtype("int32"): cde.DataType("int32"),
np.dtype("int64"): cde.DataType("int64"),
np.dtype("uint8"): cde.DataType("uint8"),
np.dtype("uint16"): cde.DataType("uint16"),
np.dtype("uint32"): cde.DataType("uint32"),
np.dtype("uint64"): cde.DataType("uint64"),
np.dtype("float16"): cde.DataType("float16"),
np.dtype("float32"): cde.DataType("float32"),
np.dtype("float64"): cde.DataType("float64"),
np.dtype("str"): cde.DataType("string"),
}.get(type_)
def mstypelist_to_detypelist(type_list):
"""
Get list[de type] corresponding to list[mindspore.dtype].
获取列表[mindspore.dtype]对应的列表[detype]
Args:
type_list (list[mindspore.dtype]): a list of MindSpore's dtype.
Returns:
The list of de data type.
"""
# 遍历传入的 type_list 列表
for index, _ in enumerate(type_list):
# 如果列表中的元素不为 None
if type_list[index] is not None:
# 调用 mstype_to_detype 函数将 MindSpore 数据类型转换为 CDE 数据类型
type_list[index] = mstype_to_detype(type_list[index])
else:
# 如果列表中的元素为 None则将其设置为空字符串的 CDE 数据类型
type_list[index] = cde.DataType("")
# 返回转换后的 type_list 列表
return type_list