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

TensorFlow时间序列教程中WindowGenerator参数设置异常问题求助

TensorFlow时间序列教程中WindowGenerator参数设置异常问题求助

我正在跟着TensorFlow的时间序列教程做天气数据预测温度的项目,但对WindowGenerator的参数逻辑和遇到的问题完全摸不着头脑,想请教大家:

一开始我用下面的参数配置是完全正常的:

wide_window = WindowGenerator(input_width=24, label_width=1,
shift=1, label_columns=['T (degC)'])

这个设置是用24小时的历史数据预测未来1小时的温度值,模型能正常输出结果,WindowGenerator的绘图函数也能正常工作。

但只要我改动其中两个参数,就会出现各种崩溃问题:

问题1:shift设为非1值时,模型无预测结果导致绘图崩溃

当我把shift(控制预测的未来偏移量)改成3(也就是预测3小时后的温度):

wide_window = WindowGenerator(input_width=24, label_width=1,
shift=3, label_columns=['T (degC)'])

模型就完全没有预测结果,直接导致WindowGenerator的plot函数因为数据不匹配而崩溃。

问题2:label_width设为非1/24值时,model.fit抛出维度不匹配错误

如果我把label_width(控制预测的步数)改成3:

wide_window = WindowGenerator(input_width=24, label_width=3,
shift=1, label_columns=['T (degC)'])

运行model.fit的时候就会直接抛出以下错误:

Exception has occurred: ValueError
Dimensions must be equal, but are 3 and 24 for '{{node compile_loss/mean_squared_error/sub}} = Sub[T=DT_FLOAT](data_1, sequential_1/dense_1/Add)' with input shapes: [?,3,1], [?,24,1].

以下是我用到的WindowGenerator类和相关模型代码(省略了数据下载和预处理部分):

# WINDOW GENERATOR CLASS
class WindowGenerator():
  def __init__(self, input_width, label_width, shift,
               train_df=train_df, val_df=val_df, test_df=test_df,
               label_columns=None):
    # Store the raw data.
    self.train_df = train_df
    self.val_df = val_df
    self.test_df = test_df

    # Work out the label column indices.
    self.label_columns = label_columns
    if label_columns is not None:
      self.label_columns_indices = {name: i for i, name in
                                    enumerate(label_columns)}
    self.column_indices = {name: i for i, name in
                           enumerate(train_df.columns)}

    # Work out the window parameters.
    self.input_width = input_width
    self.label_width = label_width
    self.shift = shift

    self.total_window_size = input_width + shift

    self.input_slice = slice(0, input_width)
    self.input_indices = np.arange(self.total_window_size)[self.input_slice]

    self.label_start = self.total_window_size - self.label_width
    self.labels_slice = slice(self.label_start, None)
    self.label_indices = np.arange(self.total_window_size)[self.labels_slice]

  def __repr__(self):
    return '\n'.join([
        f'Total window size: {self.total_window_size}',
        f'Input indices: {self.input_indices}',
        f'Label indices: {self.label_indices}',
        f'Label column name(s): {self.label_columns}'])

  # SPLIT FUNCTION
  def split_window(self, features):
    inputs = features[:, self.input_slice, :]
    labels = features[:, self.labels_slice, :]
    if self.label_columns is not None:
        labels = tf.stack(
            [labels[:, :, self.column_indices[name]] for name in self.label_columns],
            axis=-1)

    # Slicing doesn't preserve static shape information, so set the shapes
    # manually. This way the `tf.data.Datasets` are easier to inspect.
    inputs.set_shape([None, self.input_width, None])
    labels.set_shape([None, self.label_width, None])

    return inputs, labels
  
  # PLOT FUNCTION
  def plot(self, model=None, plot_col='T (degC)', max_subplots=3):
    inputs, labels = self.example
    plt.figure(figsize=(12, 8))
    plot_col_index = self.column_indices[plot_col]
    max_n = min(max_subplots, len(inputs))
    for n in range(max_n):
        plt.subplot(max_n, 1, n+1)
        plt.ylabel(f'{plot_col} [normed]')
        plt.plot(self.input_indices, inputs[n, :, plot_col_index],
                label='Inputs', marker='.', zorder=-10)

        if self.label_columns:
            label_col_index = self.label_columns_indices.get(plot_col, None)
        else:
            label_col_index = plot_col_index

        if label_col_index is None:
            continue

        plt.scatter(self.label_indices, labels[n, :, label_col_index],
                    edgecolors='k', label='Labels', c='#2ca02c', s=64)
        if model is not None:
            predictions = model(inputs)
            plt.scatter(self.label_indices, predictions[n, self.label_indices[0]-1:self.label_indices[-1], label_col_index],
                        marker='X', edgecolors='k', label='Predictions',
                        c='#ff7f0e', s=64)
            # ERROR WHEN SHIFT IS NOT 1 BECAUSE NO PREDICTION:
            # Exception has occurred: ValueError. x and y must be the same size

        if n == 0:
            plt.legend()

    plt.xlabel('Time [h]')
    plt.show()

  # MAKE DATASET FUNCTION
  def make_dataset(self, data):
    data = np.array(data, dtype=np.float32)
    ds = tf.keras.utils.timeseries_dataset_from_array(
        data=data,
        targets=None,
        sequence_length=self.total_window_size,
        sequence_stride=1,
        shuffle=True,
        batch_size=BATCH_SIZE,)

    ds = ds.map(self.split_window)

    return ds

  @property
  def train(self):
    return self.make_dataset(self.train_df)

  @property
  def val(self):
    return self.make_dataset(self.val_df)

  @property
  def test(self):
    return self.make_dataset(self.test_df)

  @property
  def example(self):
    """Get and cache an example batch of `inputs, labels` for plotting."""
    result = getattr(self, '_example', None)
    if result is None:
        # No example batch was found, so get one from the `.train` dataset
        result = next(iter(self.train))
        # And cache it for next time
        self._example = result
    return result


val_performance = {}
performance = {}

# WIDE WINDOW
wide_window = WindowGenerator(
    input_width=24, label_width=1, shift=1,
    label_columns=['T (degC)'])

# LINEAR MODEL
linear = tf.keras.Sequential([
    tf.keras.layers.Dense(units=1)
])

# COMPILE AND FIT FUNCTION
MAX_EPOCHS = 20

def compile_and_fit(model, window, patience=2):
  early_stopping = tf.keras.callbacks.EarlyStopping(monitor='val_loss',
                                                    patience=patience,
                                                    mode='min')

  model.compile(loss=tf.keras.losses.MeanSquaredError(),
                optimizer=tf.keras.optimizers.Adam(),
                metrics=[tf.keras.metrics.MeanAbsoluteError()])

  history = model.fit(window.train, epochs=MAX_EPOCHS,
                      validation_data=window.val,
                      callbacks=[early_stopping])
  return history

# COMPILE AND FIT THE LINEAR MODEL ONTO THE WIDE WINDOW
history = compile_and_fit(linear, wide_window)

print('Input shape:', wide_window.example[0].shape)
print('Output shape:', linear(wide_window.example[0]).shape)

val_performance['Linear'] = linear.evaluate(wide_window.val, return_dict=True)
performance['Linear'] = linear.evaluate(wide_window.test, verbose=0, return_dict=True)

wide_window.plot(linear)

我核心的疑问是:

  • 为什么shift只能设为1?设为其他值时模型就没有预测结果?
  • 为什么label_width只能是1或24?设为其他值就会出现维度不匹配的错误?

希望有人能帮我理清WindowGenerator这几个参数的实际作用逻辑,以及怎么解决这些问题,让模型支持任意合理的shift和label_width设置。


备注:内容来源于stack exchange,提问作者Erken

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.04.14 15:53:06