TF2.0+TF Hub训练时输入维度不匹配问题求助
咱们直接说核心问题:你的输入数据形状和TF Hub层预期的不匹配,这一切都源于你在解析TFRecord时对特征的定义方式。
为什么会出现形状差异?
你看IMDB示例里的数据集形状是((None,), (None,)),其中文本输入是一维的(每个样本是单个字符串);但你的数据集里文本输入形状是(64,1),这是因为你在featdef里把question定义成了tf.io.FixedLenFeature([1], tf.string)——这个[1]表示每个question是一个长度为1的数组,批量处理后就会得到(batch_size, 1)的二维张量。而TF Hub的文本嵌入层(比如你用的NNLM模型)预期的输入是一维张量(每个样本是单个字符串,形状为(batch_size,)),这就是报错的直接原因。
另外你的标签也有类似问题:label定义成FixedLenFeature([1], tf.int64),经过tf.one_hot后会变成(64,1,NUM_CLASSES),而模型的损失函数预期的标签形状是(64,NUM_CLASSES),这也会导致后续不兼容。
分步解决方法
方法1:修改TFRecord特征定义(如果可以重新生成TFRecord的话)
直接把特征定义里的[1]改成空数组[],表示每个特征是单个值而非数组:
def _dataset_parser(value): """Parse a record from value.""" featdef={ 'id': tf.io.FixedLenFeature([], tf.int64), 'question': tf.io.FixedLenFeature([], tf.string), 'label': tf.io.FixedLenFeature([], tf.int64) } example = tf.io.parse_single_example(value, featdef) label = tf.cast(example['label'], tf.int32) question = tf.cast(example['question'], tf.string) return question, label
方法2:在解析时去除多余维度(如果无法重新生成TFRecord)
如果你的TFRecord已经生成好、没法修改特征定义,可以在解析后用tf.squeeze去掉多余的维度:
def _dataset_parser(value): """Parse a record from value.""" featdef={ 'id': tf.io.FixedLenFeature([1], tf.int64), 'question': tf.io.FixedLenFeature([1], tf.string), 'label': tf.io.FixedLenFeature([1], tf.int64) } example = tf.io.parse_single_example(value, featdef) # 去除每个特征的多余维度 label = tf.squeeze(tf.cast(example['label'], tf.int32), axis=1) question = tf.squeeze(tf.cast(example['question'], tf.string), axis=1) return question, label
优化输入管道(适配TF2.0最佳实践)
你当前的_input函数用了TF1.x的make_one_shot_iterator,在TF2.0里没必要这么做,直接返回tf.data.Dataset对象即可,model.fit可以直接接收它。同时调整标签的one-hot处理,确保形状正确:
def _input(epochs, batch_size, filenames): dataset = tf.data.TFRecordDataset(filenames) # 先解析数据 dataset = dataset.map(_dataset_parser) # 打乱、重复、批量处理 dataset = dataset.shuffle(100) dataset = dataset.repeat(epochs) dataset = dataset.batch(batch_size) # 预取优化 dataset = dataset.prefetch(tf.data.AUTOTUNE) # 处理one-hot标签 def add_one_hot(question, label): return question, tf.one_hot(label, NUM_CLASSES) dataset = dataset.map(add_one_hot) return dataset # 生成训练数据集 train_ds = _input(20, 64, ['train_xs.tfrecords'])
验证修改效果
修改后你可以打印train_ds的元素形状,应该会得到类似((None,), (None, NUM_CLASSES))的结果,和IMDB示例的形状一致,此时再运行model.fit(train_ds, steps_per_epoch=...)就不会再出现形状不匹配的错误了。
内容的提问来源于stack exchange,提问作者Daniela Velásquez Garzón

