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

tf.keras.utils.timeseries_dataset_from_array采样率>1时如何获取完整序列?

问题分析与解决:TensorFlow timeseries_dataset_from_array 序列缺失问题

问题描述

我有一个包含数据列和类别列的Pandas DataFrame,想要创建用于神经网络的Dataset。使用timeseries_dataset_from_array函数时,发现生成的序列提前终止,本应能生成更多序列,请问为什么没有生成另外两个序列?该如何修复(最好使用该函数)?

数据与代码

构造DataFrame

import pandas as pd
import numpy as np
import tensorflow as tf

arr100_cat = pd.DataFrame(np.arange(100)*10, columns=['range_100'])
arr100_cat['category'] = 0
arr100_cat.loc[arr100_cat.index % 3 == 1, 'category'] = 1
arr100_cat.loc[arr100_cat.index % 3 == 2, 'category'] = 2

DataFrame示例:

indexrange_100category
000
1101
2202
3300
4401

生成时间序列Dataset

input_width = 5
label_width = 3
timeseries_dataset = tf.keras.utils.timeseries_dataset_from_array(
    arr100_cat.to_numpy(),
    targets=None,
    sequence_length=input_width+label_width,
    sequence_stride=1,
    sampling_rate=3,
    batch_size=1,
    shuffle=False,
    seed=None,
    start_index=None,
    end_index=None
)

遍历Dataset的代码

timeseries_dataset_iter = iter(timeseries_dataset)
for i in range(100):
    timeseries_dataset_next = next(timeseries_dataset_iter)
    print(timeseries_dataset_next)

终止现象

tf.Tensor(
[[[750   0]
  [780   0]
  [810   0]
  [840   0]
  [870   0]
  [900   0]
  [930   0]
  [960   0]]], shape=(1, 8, 2), dtype=int64)
tf.Tensor(
[[[760   1]
  [790   1]
  [820   1]
  [850   1]
  [880   1]
  [910   1]
  [940   1]
  [970   1]]], shape=(1, 8, 2), dtype=int64)
---------------------------------------------------------------------------
OutOfRangeError  

原因分析

  1. 序列生成规则限制
    当设置sampling_rate=3和sequence_length=8时,每个序列的元素是从起始索引s开始,依次取s, s+3, s+6, ..., s+3*(8-1),共8个元素。要求最后一个元素的索引s+21必须≤DataFrame最大索引(99),因此起始索引s的取值范围是0 ≤ s ≤78,总共会生成79个序列。

  2. 遍历报错的真实原因
    你的遍历代码强制循环100次,但实际只有79个序列,因此第80次调用next()时触发OutOfRangeError。你看到的最后两个序列并非全部序列的末尾,只是截取的部分输出,s=77和s=78对应的序列实际已生成,只是未遍历到就报错了。

如果你的预期是生成更多序列,大概率是混淆了sampling_rate的作用:

  • 若想要从连续原始数据窗口生成序列(而非间隔采样),应将sampling_rate=1,此时序列数为100-8+1=93。
  • 若想要每个序列内元素属于同一类别,当前sampling_rate=3的设置是正确的,只需确认实际生成的序列数量是否符合预期。

修复方案

方案1:正确遍历所有序列

调整遍历逻辑,避免超出Dataset实际长度:

# 遍历所有生成的序列
for seq in timeseries_dataset:
    print(seq)

方案2:拆分输入与标签序列

若你需要的是5个输入点+3个标签点的拆分结构(而非将8个点作为单一序列),可调整参数如下:

input_width =5
label_width=3
sampling_rate=3

# 生成输入-标签配对的序列
dataset = tf.keras.utils.timeseries_dataset_from_array(
    data=arr100_cat.to_numpy(),
    # 标签序列对应输入序列的后续3个同类别点
    targets=arr100_cat.to_numpy()[input_width*sampling_rate:],
    sequence_length=input_width,
    sequence_stride=1,
    sampling_rate=sampling_rate,
    batch_size=1,
    shuffle=False
)

# 查看输入与标签
for input_seq, label_seq in dataset:
    print("输入序列:", input_seq)
    print("标签序列:", label_seq)
    print("---")

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.20 15:56:59