将二维NumPy数组转换为TensorFlow Dataset时遇类型错误求助
解决TensorFlow Dataset转换时的Decimal类型错误
这个错误的核心原因很明确:你的NumPy特征数组里包含了Python原生的Decimal类型对象,但TensorFlow的tf.data.Dataset.from_tensor_slices()只支持处理原生数值类型(比如NumPy的float32、int64等),无法直接识别Decimal对象,所以抛出了类型错误。
下面是两种简单有效的解决方法:
方法1:批量转换Decimal为NumPy浮点类型
可以通过列表推导或者NumPy的向量化操作,把所有Decimal元素转换成float,再重新构建成标准的浮点型NumPy数组:
import numpy as np import tensorflow as tf from decimal import Decimal # 假设你的features是包含Decimal的(n,12)数组 # 方式一:列表推导逐行转换 features_float = np.array([[float(decimal_val) for decimal_val in row] for row in features], dtype=np.float32) # 方式二:用np.vectorize简化转换(适合大规模数组) decimal_to_float = np.vectorize(lambda x: float(x)) features_float = decimal_to_float(features).astype(np.float32) # 现在用转换后的数组创建Dataset就没问题了 dataset = tf.data.Dataset.from_tensor_slices((features_float, labels))
方法2:从数据源阶段避免生成Decimal
如果你的特征数据是从外部文件(比如CSV)读取的,建议在读取时直接指定解析为浮点类型,而不是Decimal。比如用pandas.read_csv()时,设置dtype参数强制指定列类型为float,从根源上避免混入Decimal:
import pandas as pd # 读取CSV时直接指定浮点类型,跳过Decimal解析 df = pd.read_csv("your_data.csv", dtype={f"col_{i}": float for i in range(12)}) features = df.values # 此时features是标准的float型NumPy数组
额外检查
虽然你提到标签是整型数组,但也可以快速确认下标签里是否混入了Decimal:
print(np.unique([type(x) for x in labels]))
如果有Decimal,同样用类似的方法转换成int即可。
内容的提问来源于stack exchange,提问作者James
相关产品推荐
相关产品推荐

