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

使用tf.data.Dataset时如何实现自定义窗口分批喂入训练数据?

TensorFlow自定义token序列喂入逻辑实现方案

你可以完全通过tf.data.Dataset的内置转换算子自定义喂入规则,所有逻辑都可以在数据流水线层面实现,不需要修改模型代码。

核心实现思路

通过flat_map算子将单条长token序列展开为多个自定义范围的子序列,再做批次处理即可匹配你的需求。

代码示例

1. 基础实现(固定切片范围)

直接按照你的示例需求写死切片规则:

import tensorflow as tf

# 模拟原始数据集,每个元素为1维token序列
raw_ds = tf.data.Dataset.from_tensor_slices([
    [0,1,2,3,4,5],
    [6,7,8,9,10,11]
])

def custom_slice(seq):
    # 自定义切片逻辑,示例返回两个长度为3的子序列
    # 对应你需求的[1,2,3]、[3,4,5]
    return [seq[1:4], seq[3:6]]

# 把单条样本生成的多个子序列展开为独立的数据集元素
processed_ds = raw_ds.flat_map(
    lambda x: tf.data.Dataset.from_tensor_slices(custom_slice(x))
)

# 按批次加载,批次大小可按需调整
batch_size = 2
final_ds = processed_ds.batch(batch_size)

测试输出:

for batch in final_ds:
    print(batch.numpy())

输出结果:

[[1 2 3]
 [7 8 9]]
[[3 4 5]
 [9 10 11]]

2. 通用滑动窗口实现

如果需要适配不同长度的序列、动态调整窗口参数,可以用更通用的滑动窗口逻辑:

window_size = 3  # 子序列长度
step = 2         # 窗口滑动步长
start_offset = 1 # 第一个窗口的起始索引偏移

def sliding_window(seq):
    max_start = tf.shape(seq)[0] - window_size
    # 生成所有合法的窗口起始索引
    starts = tf.range(start_offset, max_start + 1, delta=step)
    # 批量生成所有子序列
    return tf.map_fn(lambda s: seq[s:s+window_size], starts, dtype=seq.dtype)

processed_ds = raw_ds.flat_map(
    lambda x: tf.data.Dataset.from_tensor_slices(sliding_window(x))
)
final_ds = processed_ds.batch(2)

注意事项

  • 你可以任意修改自定义切片函数的逻辑,实现任意范围的token截取、不同样本差异化处理等需求
  • 如果数据集包含标签,只要在切片函数中同步处理特征和标签即可,比如原始数据为(feature_seq, label)元组,切片后返回(sub_seq, label)即可
  • 整个流水线属于TensorFlow计算图的一部分,支持prefetch、shuffle、多进程并行加载等常规优化,不会影响训练性能

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.30 23:06:00