向生成器传参构建tf.data.Dataset时遇类型转换错误问题
问题原因分析与解决方案
问题原因
tf.data.Dataset.from_generator的args参数要求传入的每个参数都能被转换为Tensor对象。你的some_list同时包含元组和列表两种不同类型的元素,属于混合类型序列,TensorFlow无法为这种序列创建类型统一的Tensor,因此抛出ValueError: Can't convert Python sequence with mixed types to Tensor。
- 当你不传入参数、直接使用全局变量
some_list时,生成器每次返回的是单一类型的整数和字符串,TensorFlow只需处理这两个输出类型,无需转换混合类型的输入参数,因此正常运行。 - 当传入纯整数列表时,所有元素类型一致,可被转换为统一类型的Tensor,因此也能正常运行。
解决方案
方案1:统一列表元素类型
将some_list中的所有元素统一为同一种类型(全元组或全列表),这样TensorFlow可以将整个列表转换为结构一致的Tensor:
import tensorflow as tf import random # 统一所有元素为元组 some_list = [(1,'One'), (2,'Two'), (3,'Three'), (4,'Four'), (5,'Five'), (6,'Six'), (7,'Seven'), (8,'Eight')] def text_gen2(file_list): random.shuffle(file_list) size = len(file_list) i = 0 while True: yield file_list[i][0], file_list[i][1] i += 1 if i >= size: # 修正原代码越界bug:i>size会访问不存在的索引 i = 0 random.shuffle(file_list) tf_dataset2 = tf.data.Dataset.from_generator( text_gen2, args=[some_list], output_types=(tf.int32, tf.string), output_shapes=((), ()) ) for count_batch in tf_dataset2.repeat().batch(3).take(2): print(count_batch)
方案2:用Lambda闭包传递参数
通过lambda包装生成器,避免直接将混合类型列表传入args,而是利用闭包传递参数,这样TensorFlow不会尝试转换some_list:
import tensorflow as tf import random some_list = [(1,'One'),[2,'Two'],[3,'Three'],[4,'Four'], (5,'Five'),[6,'Six'],[7,'Seven'],[8,'Eight']] def text_gen2(file_list): random.shuffle(file_list) size = len(file_list) i = 0 while True: yield file_list[i][0], file_list[i][1] i += 1 if i >= size: i = 0 random.shuffle(file_list) # 使用lambda闭包传递参数,无需args tf_dataset2 = tf.data.Dataset.from_generator( lambda: text_gen2(some_list), output_types=(tf.int32, tf.string), output_shapes=((), ()) ) for count_batch in tf_dataset2.repeat().batch(3).take(2): print(count_batch)
内容的提问来源于stack exchange,提问作者Arindam
相关产品推荐
相关产品推荐

