diff --git a/mindspore/ccsrc/transform-update/datatypes.py b/mindspore/ccsrc/transform-update/datatypes.py new file mode 100644 index 00000000000..aae15d42f59 --- /dev/null +++ b/mindspore/ccsrc/transform-update/datatypes.py @@ -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 数据类型映射到 CDE(MindSpore 数据增强库)的数据类型 +# 这个字典用于将 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 数据类型映射到 CDE(MindSpore 数据增强库)的数据类型 +# 这个字典用于将 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 +