如何将Prefetch Dataset的int32转换为tf.int64以匹配输入签名?
问题场景
你遇到的错误是数据集输出的int32类型张量与train_step_signature要求的tf.int64类型不兼容,报错详情如下:
ValueError: Python inputs incompatible with input_signature:
inputs: (
tf.Tensor(
[[ 77 111 110 ... 0 0 0]
[ 83 105 110 ... 0 0 0]
[ 71 97 115 ... 0 0 0]
...
[ 80 114 111 ... 0 0 0]
[ 70 114 97 ... 0 0 0]
[ 65 110 233 ... 0 0 0]], shape=(64, 605), dtype=int32),
tf.Tensor(
[[ 68 101 115 ... 0 0 0]
[ 76 101 32 ... 0 0 0]
[ 76 101 32 ... 0 0 0]
...
[ 68 97 110 ... 0 0 0]
[ 85 110 101 ... 0 0 0]
[ 68 97 110 ... 0 0 0]], shape=(64, 936), dtype=int32))
input_signature: (
TensorSpec(shape=(None, None), dtype=tf.int64, name=None),
TensorSpec(shape=(None, None), dtype=tf.int64, name=None)).
解决方案
可以通过tf.data.Dataset.map()方法,对数据集的每个元素应用类型转换,将int32转为tf.int64,具体实现如下:
1. 针对输入-目标二元组的数据集
如果你的数据集每个元素是(inputs, targets)这样的二元组,直接对两个张量分别转换:
import tensorflow as tf def convert_to_int64(inputs, targets): inputs = tf.cast(inputs, tf.int64) targets = tf.cast(targets, tf.int64) return inputs, targets # 在prefetch前加入类型转换的map操作 dataset = dataset.map(convert_to_int64).prefetch(tf.data.AUTOTUNE)
2. 针对字典结构的数据集
如果数据集元素是字典格式,遍历字典键值对转换类型:
import tensorflow as tf def convert_dict_to_int64(data_dict): for key, tensor in data_dict.items(): if tensor.dtype == tf.int32: data_dict[key] = tf.cast(tensor, tf.int64) return data_dict dataset = dataset.map(convert_dict_to_int64).prefetch(tf.data.AUTOTUNE)
原理说明
tf.cast()是TensorFlow原生的类型转换函数,能安全地将张量从一种数值类型转为另一种;map()操作会遍历数据集的每一个元素,批量应用转换逻辑,之后再执行prefetch()就能得到符合train_step_signature要求的tf.int64类型数据集。
内容的提问来源于stack exchange,提问作者Antoine23

