使用tfds.as_numpy()在map函数处理后崩溃的问题咨询
问题:使用
tfa.image.rotate后tfds.as_numpy()导致程序崩溃? 先还原你的场景:在Google Colab的TensorFlow 2.2环境中,你尝试用map函数给图像旋转20度,但加上tfa.image.rotate代码后程序就崩溃,移除后一切正常,你怀疑是tfds.as_numpy()无法读取map处理后的结果。
先给结论:不是tfds.as_numpy()的问题,核心问题出在tfa.image.rotate的输出和数据集后续处理不兼容,大概率是版本兼容性或图像形状变化导致的。下面是具体分析和解决办法:
可能的原因
- TensorFlow Addons版本不兼容:TensorFlow 2.2对标的TensorFlow Addons版本是0.10.0,如果你的tfa版本过高或过低,
rotate函数的行为可能出现差异,比如输出形状、数据类型的异常变化。 - 旋转后图像形状不一致:默认情况下,
tfa.image.rotate会将旋转后的图像裁剪到最小包围矩形,导致输出图像的形状和输入不一致。而tf.data.Dataset.batch()要求每个batch里的张量形状必须统一,形状不匹配就会触发崩溃,tfds.as_numpy()只是在这个问题暴露时被你注意到而已。
解决办法
1. 确保TensorFlow Addons版本兼容
在Colab中先检查当前tfa版本:
import tensorflow_addons as tfa print(tfa.__version__)
如果不是0.10.0,安装兼容版本:
!pip uninstall -y tensorflow-addons !pip install tensorflow-addons==0.10.0
2. 旋转时保持图像形状不变
修改你的map_func,旋转时指定填充参数,让输出图像和输入形状一致:
import tensorflow_datasets as tfds import tensorflow as tf import tensorflow_addons as tfa import numpy as np MODE_AUTOTUNE = tf.data.experimental.AUTOTUNE batch_size = 128 def map_func(image, label): image = tf.cast(image, tf.float32) / 255. angle = 20. * np.pi / 180.0 # 用np.pi比手动写3.14更准确 # 指定fill_value填充旋转后缺失的区域,保持形状和输入一致 image = tfa.image.rotate( images=image, angles=angle, interpolation='NEAREST', fill_value=0.0 # 用0填充黑色,你也可以换成1.0(白色)等其他值 ) return image, label # 假设你已经加载好了x_train和y_train train_data = tf.data.Dataset.from_tensor_slices((x_train, y_train)) train_data = train_data.shuffle(len(x_train)) train_data = train_data.map(map_func=map_func, num_parallel_calls=MODE_AUTOTUNE) train_data = train_data.batch(batch_size) train_data = train_data.prefetch(buffer_size=MODE_AUTOTUNE) # 现在再尝试读取就不会崩溃了 datas = tfds.as_numpy(train_data) for data in datas: print(data[0].shape) # 打印形状确认是否统一
3. 调试技巧:检查旋转前后的形状
如果你还是不确定问题所在,可以在map_func里添加打印语句,监控图像形状的变化:
def map_func(image, label): image = tf.cast(image, tf.float32) / 255. tf.print("Before rotation shape:", tf.shape(image)) # 打印输入形状 angle = 20. * np.pi / 180.0 image = tfa.image.rotate(images=image, angles=angle, interpolation='NEAREST') tf.print("After rotation shape:", tf.shape(image)) # 打印输出形状 return image, label
如果旋转前后形状不一样,就说明是裁剪导致的问题,用上面的fill_value参数即可解决。
替代方案:用原生TF实现旋转
如果你不想折腾tfa的版本,也可以用TensorFlow原生的方法实现旋转,比如构建旋转矩阵:
def rotate_image(image, angle): # 获取图像尺寸 height, width = image.shape[0], image.shape[1] # 计算旋转矩阵 rotation_matrix = tf.convert_to_tensor([ [tf.cos(angle), -tf.sin(angle), 0], [tf.sin(angle), tf.cos(angle), 0] ], dtype=tf.float32) # 应用仿射变换,保持形状不变 image = tf.keras.preprocessing.image.apply_affine_transform( image, transform_matrix=rotation_matrix, fill_mode='nearest', cval=0.0 ) return image def map_func(image, label): image = tf.cast(image, tf.float32) / 255. angle = 20. * np.pi / 180.0 image = rotate_image(image, angle) return image, label
内容的提问来源于stack exchange,提问作者Paul
相关产品推荐
相关产品推荐

