如何优雅拆分TensorFlow BatchDataset适配多输入LSTM模型?
双输入LSTM模型的数据集拆分优化方案
你可以用tf.split()实现更优雅的数据集拆分,之前报错的核心原因是tf.split()返回的是张量列表,TensorFlow Dataset会试图将列表打包为单个高阶张量,无法匹配模型需要的两个独立输入。只需要将拆分结果转为元组即可解决问题:
import tensorflow as tf # 使用tf.split拆分并转为元组 input_dataset2 = input_dataset.map( lambda x, y: (tuple(tf.split(x, num_or_size_splits=[1, 2], axis=-1)), y) ) # 训练模型 model.fit( input_dataset2, steps_per_epoch=20, epochs=50, verbose=0, shuffle=True )
补充说明
num_or_size_splits=[1,2]对应分类变量1列、数值变量2列的拆分规则,axis=-1表示在特征维度(最后一维)拆分;- 转为元组
tuple(...)是关键,这样Dataset会将其识别为两个独立的输入张量,完美匹配模型的双输入要求。
如果你觉得匿名lambda不够直观,也可以用命名函数增强可读性,逻辑和你原来的切片方案一致,但代码更简洁:
def split_features(x, y): # 提取分类变量(最后一维第0列) cat_input = x[..., :1] # 提取数值变量(最后一维第1到末尾列) num_input = x[..., 1:] return (cat_input, num_input), y input_dataset2 = input_dataset.map(split_features)
内容的提问来源于stack exchange,提问作者Requin
相关产品推荐
相关产品推荐

