ADD file via upload
This commit is contained in:
parent
5b1b14448a
commit
fef5b1745f
|
|
@ -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
|
||||
|
||||
Loading…
Reference in New Issue