【问题标题】:TensorFlow 2.1 - TypeError: padded_batch() missing 1 required positional argument: 'padded_shapes'TensorFlow 2.1 - 类型错误:padded_batch() 缺少 1 个必需的位置参数:'padded_shapes'
【发布时间】:2020-12-17 22:54:25
【问题描述】:

我正在尝试在我拥有 CUDA 的本地环境中复制 https://keras.io/examples/vision/retinanet/ 教程。但是,由于我使用的是 Windows,TensorFlow 版本是 2.1,而不是 2.4。在本教程所针对的文档(TensorFlow 2.4 版本)中,padded_shapes 似乎是可选参数,而在 TensorFlow 2.1 中。版本是必需的。如何避免这种情况或如何将其设置为正确的值?

代码如下:

autotune = tf.data.experimental.AUTOTUNE
train_dataset = train_dataset.map(preprocess_data, num_parallel_calls=autotune)
train_dataset = train_dataset.shuffle(8 * batch_size)
train_dataset = train_dataset.padded_batch(
    batch_size=batch_size, padding_values=(0.0, 1e-8, -1), drop_remainder=True
)
train_dataset = train_dataset.map(
    label_encoder.encode_batch, num_parallel_calls=autotune
)
train_dataset = train_dataset.apply(tf.data.experimental.ignore_errors())
train_dataset = train_dataset.prefetch(autotune)

val_dataset = val_dataset.map(preprocess_data, num_parallel_calls=autotune)
val_dataset = val_dataset.padded_batch(
    batch_size=1, padding_values=(0.0, 1e-8, -1), drop_remainder=True
)
val_dataset = val_dataset.map(label_encoder.encode_batch, num_parallel_calls=autotune)
val_dataset = val_dataset.apply(tf.data.experimental.ignore_errors())
val_dataset = val_dataset.prefetch(autotune)

这是错误:

---------------------------------------------------------------------------
TypeError                                 Traceback (most recent call last)
<ipython-input-363-22fc45aaf608> in <module>
      3 train_dataset = train_dataset.shuffle(8 * batch_size)
      4 train_dataset = train_dataset.padded_batch(
----> 5     batch_size=batch_size, padding_values=(0.0, 1e-8, -1), drop_remainder=True
      6 )
      7 train_dataset = train_dataset.map(

TypeError: padded_batch() missing 1 required positional argument: 'padded_shapes'

【问题讨论】:

  • 我回答了你的问题还是有什么不清楚的地方?

标签: python tensorflow tensorflow2.0 tensorflow-datasets


【解决方案1】:

v2.1 methodv2.4 method 相比具有不同的签名
解决方案:

  1. 更新到v2.2+
  2. 尝试:val_dataset = val_dataset.padded_batch( batch_size=1, padded_shapes=[None], padding_values=(0.0, 1e-8, -1), drop_remainder=True)

【讨论】:

    猜你喜欢
    • 2022-01-11
    • 1970-01-01
    • 2018-09-12
    • 2021-08-05
    • 2021-07-06
    • 2021-08-05
    • 2017-07-23
    • 2020-12-11
    相关资源
    最近更新 更多