如何使用形状不同的Numpy数组列表创建tf.data Dataset
问题原因
你使用tf.data.Dataset.from_tensor_slices报错的原因是,该接口会先尝试将输入的整个序列转换为一个统一的矩形张量,而你的列表元素形状各不相同,无法完成转换,因此抛出Can't convert non-rectangular Python sequence to Tensor错误。
解决方案
以下两种实现方式都可以保证拉取的张量和原Numpy数组的形状、数值完全一致:
方案1:使用from_generator构造(推荐,适合大数据集)
该方案不会一次性加载所有数组为张量,内存占用更低,适配任意规模的数组列表:
import tensorflow as tf import numpy as np # 示例:形状不同的Numpy数组列表 list_of_arrays = [ np.array([1, 2]), np.array([[1, 2], [3, 4]]), np.array([1, 2, 3, 4, 5]) ] # 定义生成器,遍历返回每个数组 def array_generator(): for arr in list_of_arrays: yield arr # 构造数据集,注意dtype需和你的数组实际类型匹配,shape=None表示动态形状 dataset = tf.data.Dataset.from_generator( generator=array_generator, output_signature=tf.TensorSpec(shape=None, dtype=tf.int64) ) # 验证输出 for tensor in dataset: print(f"形状:{tensor.shape},数值:\n{tensor.numpy()}")
方案2:单元素数据集拼接(适合小数据集)
如果数组数量较少,可以直接逐个构造单元素数据集再拼接,逻辑更简单:
dataset = tf.data.Dataset.from_tensors(list_of_arrays[0]) for arr in list_of_arrays[1:]: dataset = dataset.concatenate(tf.data.Dataset.from_tensors(arr))
补充说明
如果后续需要对动态形状的元素执行批量操作,可以使用padded_batch接口对不同形状的张量做填充后再批量,不需要批量操作的话可以直接忽略。
内容的提问来源于stack exchange,提问作者MarcoM
相关产品推荐
相关产品推荐

