基于TensorFlow Dataset API训练模型后,新数据推理的更优方案问询
更优雅的TensorFlow 1.7 Dataset API多模式(训练/验证/推理)处理方案
嘿,这个问题我之前在做TF1.x项目的时候也踩过类似的坑!你现在加第三个推理数据集的方法是可行的,但确实有更简洁的方式,既能避免数据覆盖的问题,又不用维护三个独立的迭代器。下面给你分享两种更优的思路:
方法一:用可重新初始化的迭代器统一管理多数据集
这种方法的核心是先定义好数据集的统一结构,再为训练、验证、推理分别创建Dataset,最后用一个可重新初始化的迭代器来切换不同模式,完美解决数据覆盖的问题:
# 先定义所有数据集的输出类型和形状,和你的输入输出匹配 output_types = (tf.int32, tf.int32) output_shapes = ( tf.TensorShape([None, None, max_sent_len]), # features的形状 tf.TensorShape([None, num_classes]) # labels的形状 ) # 创建可重新初始化的迭代器 iterator = tf.data.Iterator.from_structure(output_types, output_shapes) input_tensor, labels_tensor = iterator.get_next() # ---------------------- 分别定义训练/验证/推理数据集 ---------------------- # 训练集:带shuffle和batch train_dataset = tf.data.Dataset.from_tensor_slices((features_data_ph, labels_data_ph)) train_dataset = train_dataset.shuffle(buffer_size=100000).batch(batch_size) train_init_op = iterator.make_initializer(train_dataset) # 验证集:仅batch,不shuffle val_dataset = tf.data.Dataset.from_tensor_slices((features_data_ph, labels_data_ph)) val_dataset = val_dataset.batch(batch_size) val_init_op = iterator.make_initializer(val_dataset) # 推理集:单独用一个placeholder(因为推理不需要labels) infer_features_ph = tf.placeholder(tf.int32, [None, None, max_sent_len], name='infer_features_ph') infer_dataset = tf.data.Dataset.from_tensor_slices(infer_features_ph).batch(batch_size) infer_init_op = iterator.make_initializer(infer_dataset) # ---------------------- 模型和损失定义不变 ---------------------- logits = model(input_tensor) loss = get_loss(logits, labels_tensor)
使用的时候就很清晰了,切换模式只需要初始化对应的操作:
# 训练阶段 session.run(train_init_op, feed_dict={ features_data_ph: train_features, labels_data_ph: train_labels }) # 验证阶段 session.run(val_init_op, feed_dict={ features_data_ph: val_features, labels_data_ph: val_labels }) # 推理阶段(不用传labels,用单独的推理placeholder) session.run(infer_init_op, feed_dict={ infer_features_ph: your_new_infer_data }) # 获取推理结果 infer_results = session.run(logits)
这种方式的好处是结构清晰,所有模式共享同一个迭代器和模型输入张量,同时推理用了独立的placeholder,完全不会和训练/验证的数据互相覆盖。
方法二:直接用Placeholder喂数据(适合小批量推理)
如果你的推理场景只是偶尔进行,且数据量不大,可以绕开Dataset API,直接给模型喂numpy数组:
# 修改模型定义,让它可以接受外部传入的输入张量 def model(inputs): # 你的原有模型逻辑 ... return logits # 除了迭代器的input_tensor,额外定义一个推理用的placeholder infer_input_ph = tf.placeholder(tf.int32, [None, None, max_sent_len], name='infer_input_ph') infer_logits = model(infer_input_ph) # 推理时直接喂数据,不用碰Dataset和迭代器 infer_results = session.run(infer_logits, feed_dict={ infer_input_ph: your_new_infer_data })
这种方法的优势是简单快捷,不用维护额外的Dataset和迭代器,但如果是大规模推理,Dataset API的并行加载等优化就享受不到了,所以更适合小批量的快速推理场景。
为什么你之前的方法会出现数据覆盖?
你之前用同一个features_data_ph和labels_data_ph来初始化训练和验证数据集,每次调用iterator.initializer时,placeholder里的旧数据会被新传入的数据覆盖。而上面的方法一通过给推理集单独分配placeholder,彻底避免了这个问题;方法二则完全绕开了共享placeholder的问题。
内容的提问来源于stack exchange,提问作者ted
相关产品推荐
相关产品推荐

