You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何优雅拆分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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.02 01:55:20