tf.data.Dataset.from_generator适配多输入网络output_types配置报错问题
问题原因
你的报错核心是output_types的结构和生成器每次yield返回的元素结构不匹配:
x_data_gen每次迭代返回的是长度为num_inputs的张量列表,对应单条样本的所有输入- 你当前定义的
x_types是长度为num_exp的列表,结构完全不对应,TensorFlow无法正确解析类型
修复方案
方案1:修正output_types配置
把x_types的生成长度从num_exp改成num_inputs即可:
# 原错误写法:x_types = [ tf.float32 for _ in range(num_exp) ] # 修正后: x_types = [ tf.float32 for _ in range(num_inputs) ] y_types = tf.float32 inputs = tf.data.Dataset.from_generator(x_data_gen, output_types=x_types) y = tf.data.Dataset.from_generator(y_data_gen, output_types=y_types)
方案2:用output_signature更稳妥
如果想避免后续出现形状不匹配的隐式报错,更推荐用output_signature参数同时指定类型和形状:
x_signature = [ tf.TensorSpec(shape=(1,6000,1), dtype=tf.float32) for _ in range(num_inputs) ] y_signature = tf.TensorSpec(shape=(1,6000,1), dtype=tf.float32) inputs = tf.data.Dataset.from_generator(x_data_gen, output_signature=x_signature) y = tf.data.Dataset.from_generator(y_data_gen, output_signature=y_signature)
Y生成器运行正常是因为它每次yield单个张量,你传的y_types是单个tf.float32类型,结构完全匹配,没有解析问题。
内容的提问来源于stack exchange,提问作者Bill Connelly
相关产品推荐
相关产品推荐

