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

如何正确合并两个TensorFlow Dataset用于模型训练?

问题描述

现有两个TensorFlow数据集,需分别独立处理生成特征、目标对应的不同尺寸滑动窗口,初始实现代码如下:

window_size_x = 3
window_size_y = 2
shift_size = 1

x = np.arange(10)
y = x * 10

x = x[:-window_size_y]
y = y[window_size_x:]

ds_x = tf.data.Dataset.from_tensor_slices(x).window(window_size_x, shift=shift_size, drop_remainder=True)
ds_y = tf.data.Dataset.from_tensor_slices(y).window(window_size_y, shift=shift_size, drop_remainder=True)

for i, j in zip(ds_x, ds_y):
  print(list(i.as_numpy_iterator()), list(j.as_numpy_iterator()))

代码运行输出符合窗口预期:

[0, 1, 2] [30, 40]
[1, 2, 3] [40, 50]
[2, 3, 4] [50, 60]
[3, 4, 5] [60, 70]
[4, 5, 6] [70, 80]
[5, 6, 7] [80, 90]

后续训练时遇到两类报错:

  • 直接调用model.fit(ds_x, ds_y)传入两个数据集,触发报错:ValueError: y argument is not supported when using dataset as input.
  • 尝试通过ds_all = tf.data.Dataset.from_tensor_slices((ds_x, ds_y))合并数据集,触发报错:ValueError: Slicing dataset elements is not supported for rank 0.
解决方案

错误原因

  1. 当给model.fit传入Dataset类型输入时,要求Dataset本身输出(特征张量, 标签张量)的元组结构,不支持分开传入两个Dataset分别作为x和y。
  2. from_tensor_slices仅支持对张量、数组等内存数据做切片,不支持直接传入Dataset对象做切片操作。
  3. window方法返回的每个窗口是嵌套的子Dataset对象,不是可直接输入模型的实际张量。

正确实现代码

使用tf.data.Dataset.zip按顺序对齐两个数据集的元素,同时通过flat_map+batch将每个窗口的子Dataset展平为固定长度张量,最终生成符合训练要求的数据集:

import tensorflow as tf
import numpy as np

window_size_x = 3
window_size_y = 2
shift_size = 1

x = np.arange(10)
y = x * 10

x = x[:-window_size_y]
y = y[window_size_x:]

ds_x = tf.data.Dataset.from_tensor_slices(x).window(window_size_x, shift=shift_size, drop_remainder=True)
ds_y = tf.data.Dataset.from_tensor_slices(y).window(window_size_y, shift=shift_size, drop_remainder=True)

# 展平窗口为张量
ds_x = ds_x.flat_map(lambda win: win.batch(window_size_x))
ds_y = ds_y.flat_map(lambda win: win.batch(window_size_y))
# 对齐合并特征、标签数据集
ds_all = tf.data.Dataset.zip((ds_x, ds_y))

# 验证数据集输出
for feat, label in ds_all:
    print(feat.numpy(), label.numpy())

运行后输出和预期窗口完全一致:

[0 1 2] [30 40]
[1 2 3] [40 50]
[2 3 4] [50 60]
[3 4 5] [60 70]
[4 5 6] [70 80]
[5 6 7] [80 90]

训练配置补充

传入模型前可按需添加批次划分、预加载等常规数据流水线配置,例如:

# 按批次大小2划分,开启自动预加载
ds_all = ds_all.batch(2).prefetch(tf.data.AUTOTUNE)
# 直接传入合并后的数据集即可开始训练
model.fit(ds_all, epochs=10)

内容的提问来源于stack exchange,提问作者Mykola Zotko

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.27 23:24:26