使用tf.data.Dataset遇AttributeError:TensorDataset无output_shapes属性
解决tf.data.Dataset获取输出形状的AttributeError问题
问题场景
运行以下代码时触发属性错误:
import numpy as np import tensorflow as tf x, y = np.array([1, 2, 3, 4]), np.array([5, 6, 7, 8]) d = tf.data.Dataset.from_tensors((x,y)) print(d.output_shapes)
报错信息:
AttributeError: 'TensorDataset' object has no attribute 'output_shapes'
原因说明
TensorFlow 2.x版本中,output_shapes、output_types这类旧属性已被弃用,官方推荐使用element_spec属性来获取数据集元素的类型和形状信息。
解决方案
方法1:通过element_spec提取形状
element_spec会返回数据集元素的完整规格对象,从中可以直接取出每个元素的形状:
import numpy as np import tensorflow as tf x, y = np.array([1, 2, 3, 4]), np.array([5, 6, 7, 8]) d = tf.data.Dataset.from_tensors((x,y)) element_spec = d.element_spec x_shape = element_spec[0].shape y_shape = element_spec[1].shape print(f"x的形状:{x_shape}") print(f"y的形状:{y_shape}")
运行输出:
x的形状:(4,) y的形状:(4,)
方法2:迭代取出元素查看形状
如果数据集规模较小,可直接取出单个元素,通过元素的shape属性获取形状:
import numpy as np import tensorflow as tf x, y = np.array([1, 2, 3, 4]), np.array([5, 6, 7, 8]) d = tf.data.Dataset.from_tensors((x,y)) first_x, first_y = next(iter(d)) print(f"x的形状:{first_x.shape}") print(f"y的形状:{first_y.shape}")
兼容旧版本(不推荐)
若需适配TensorFlow 1.x的代码写法,可使用兼容模块,但不建议长期依赖:
import numpy as np import tensorflow as tf x, y = np.array([1, 2, 3, 4]), np.array([5, 6, 7, 8]) d = tf.data.Dataset.from_tensors((x,y)) print(tf.compat.v1.data.get_output_shapes(d))
运行输出:
(TensorShape([4]), TensorShape([4]))
内容的提问来源于stack exchange,提问作者Hina
相关产品推荐
相关产品推荐

