TensorFlow输入参数数据类型不匹配报错求助:float64不在允许列表中
解决TensorFlow的DataType float64不兼容问题
这个报错本质是你传入模型的输入数据类型是float64,但TensorFlow的目标操作只支持float16、bfloat16或float32。看你代码里已经尝试过类型转换但注释掉了,大概率是转换的时机或方式不对,我给你梳理两种可行的解决方案:
先确认数据类型
首先你可以先打印下数据的类型,明确问题:
print("Images dtype:", train_images.dtype) print("Annotations dtype:", train_annotations.dtype)
如果输出是float64,就按照下面的方法处理。
情况1:next_batch()返回numpy数组
如果你的train_dataset_reader.next_batch()返回的是numpy数组(不是TensorFlow张量),直接用numpy的类型转换最直接,不需要用TensorFlow的函数:
for itr in xrange(MAX_ITERATION): train_images, train_annotations = train_dataset_reader.next_batch(batch_size) # 直接转换numpy数组类型 train_images = train_images.astype(np.float32) train_annotations = train_annotations.astype(np.float32) # 后续模型操作使用转换后的变量 # ... 你的模型训练代码 ...
情况2:next_batch()返回TensorFlow张量
如果返回的是TensorFlow张量,就用tf.cast来转换类型(注意要确保转换后的张量被传入后续操作):
for itr in xrange(MAX_ITERATION): train_images, train_annotations = train_dataset_reader.next_batch(batch_size) # 转换张量类型为float32 train_images = tf.cast(train_images, tf.float32) train_annotations = tf.cast(train_annotations, tf.float32) # 后续模型操作使用转换后的张量 # ... 你的模型训练代码 ...
另外,如果你需要把图像数据归一化到[0,1]区间,可以用tf.image.convert_image_dtype,这个函数会自动完成类型转换+归一化:
train_images = tf.image.convert_image_dtype(train_images, tf.float32)
关键注意点
一定要确保后续模型训练时使用的是转换后的float32数据,而不是原来的float64变量。很多时候注释掉转换代码后忘了恢复,或者转换后没有赋值给原变量,都会导致报错依然存在。
内容的提问来源于stack exchange,提问作者Wey Shi
相关产品推荐
相关产品推荐

