TensorFlow调用padded_batch报序列长度不匹配错误如何解决
报错原因
触发该错误的核心原因是padded_batch传入的padded_shapes嵌套结构,与数据集实际输出的结构不匹配:
- 经过
tf.py_function包装的映射函数返回2个float32类型张量,因此d2的单元素为长度2的元组结构,分别对应计算得到的y1、y2 - 调用
padded_batch时传入的padded_shapes=(None,)是长度为1的结构,无法和数据集长度为2的输出结构对齐,因此抛出序列长度不匹配的错误。
解决方案
将padded_shapes修改为与数据集输出结构完全匹配的形式即可,修正后可正常运行的代码如下:
import tensorflow as tf d1 = tf.data.Dataset.range(10) d1 = d1.map(lambda x:tf.cast([x], tf.float32)) def func1(x): y1 = 2.0 * x y2 = -3.0 * x return tuple([y1, y2]) d2 = d1.map(lambda x: tf.py_function(func1, [x], [tf.float32, tf.float32])) # 修复:padded_shapes对应两个输出的形状,每个输出第一维度自动按batch内最大长度填充 d3 = d2.padded_batch(3, padded_shapes=((None,), (None,))) # 验证批次迭代 for x_batch, y_batch in d3.as_numpy_iterator(): print(f"x批次形状:{x_batch.shape}, y批次形状:{y_batch.shape}")
注意事项
- 使用
padded_batch时,padded_shapes、padding_values参数的嵌套结构、元素数量必须和数据集单输出的结构完全一一对应:如果数据集单步输出n个张量组成的元组,上述两个参数也必须传入长度为n的元组,逐位对应每个张量的填充规则 - 若不同输出需要指定不同填充值,按相同结构传参即可,例如两个输出分别填充0和-99时,传入
padding_values=(tf.constant(0, dtype=tf.float32), tf.constant(-99, dtype=tf.float32)) - 原代码中迭代
d2的逻辑无问题,报错仅由padded_batch参数不匹配导致。
内容的提问来源于stack exchange,提问作者Jitendra Dhiman
相关产品推荐
相关产品推荐

