TensorFlow 2.0:将TFRecord读取的MapDataset转为模型输入格式
把字典结构的MapDataset转换成模型可用的输入格式
你完全不用依赖生成器,用TensorFlow Dataset的map方法就能完美解决这个问题——既保留Dataset的所有便利特性(比如批处理、洗牌、预取),又能直接适配model.fit的要求,包括validation_data参数。
核心解决方案:自定义转换函数
只需要写一个简单的函数,从字典格式的样本里提取你需要的特征和标签,返回成模型需要的元组格式就行:
def format_for_model(sample_dict): # 从字典中取出Signal A和标签 input_feature = sample_dict['Signal A'] label = sample_dict['label'] # 返回(输入特征,标签)的元组 return (input_feature, label) # 对训练集和验证集应用转换 training_dataset = training_data.map(format_for_model) val_dataset = val_data.map(format_for_model)
转换后你再查看数据集信息,就会得到你想要的格式:
<MapDataset shapes: ((150,), ()), types: (tf.float32, tf.int64)>
适配未来多输入的需求
如果之后要用到Signal B,只需要修改转换函数,返回多输入的元组即可(适配多输入模型):
def format_for_multi_input(sample_dict): input_a = sample_dict['Signal A'] input_b = sample_dict['Signal B'] label = sample_dict['label'] # 多输入模型需要把所有输入打包成一个元组,再和标签配对 return ((input_a, input_b), label) # 更新数据集 training_dataset = training_data.map(format_for_multi_input) val_dataset = val_data.map(format_for_multi_input)
为什么这比生成器好?
用Dataset的map是TensorFlow原生的操作,它能:
- 保留Dataset的所有优化特性(比如
batch()、shuffle()、prefetch(tf.data.AUTOTUNE)这些都能直接链式调用) - 支持并行处理(可以在
map里加num_parallel_calls参数提升效率) - 完美兼容
model.fit的所有参数,包括validation_data,不需要额外处理
转换后的数据集可以直接用来训练:
model.fit( training_dataset.batch(32).shuffle(1000).prefetch(tf.data.AUTOTUNE), validation_data=val_dataset.batch(32), epochs=10 )
内容的提问来源于stack exchange,提问作者Alb
相关产品推荐
相关产品推荐

