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

如何修改WindowGenerator让LSTM训练时加入预测日特征做电价预测

解决方案

核心修改思路

你当前的窗口逻辑只把前14天的全量特征作为输入,要加入第15天的已知协变量(风力、温度等非电价特征),只需要调整split_window的特征提取逻辑,同时给WindowGenerator新增已知协变量的配置项即可,具体修改如下:

第一步:修改WindowGenerator类的初始化方法

新增known_cov_columns参数,用来指定未来可提前获取的非标签特征列名:

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, known_cov_columns=None): # 新增已知协变量参数
    # 原有逻辑保留
    self.train_df = train_df
    self.val_df = val_df
    self.test_df = test_df

    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)}
    
    # 新增:保存已知协变量的列索引
    self.known_cov_columns = known_cov_columns
    if known_cov_columns is not None:
      self.known_cov_indices = [self.column_indices[name] for name in known_cov_columns]

    # 原有窗口参数逻辑保留
    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}',
        f'Known covariate column name(s): {self.known_cov_columns}']) # 新增打印项

第二步:重写split_window方法

同时提取历史全量特征和未来段的已知协变量,拼接为最终输入:

def split_window(self, features):
  # 1. 提取前14天的全量特征(历史输入)
  history_inputs = features[:, self.input_slice, :]
  # 2. 提取第15天的已知协变量(未来已知输入)
  future_known_inputs = tf.gather(features[:, self.labels_slice, :], self.known_cov_indices, axis=-1)
  
  # 适配普通LSTM的处理:把未来已知协变量广播到历史每个时间步,拼接特征维度
  # 如果你用的是编码解码LSTM,可以直接把history_inputs传给编码器,future_known_inputs传给解码器
  future_known_inputs_tiled = tf.tile(tf.expand_dims(future_known_inputs, 1), [1, self.input_width, 1, 1])
  future_known_inputs_flat = tf.reshape(future_known_inputs_tiled, [-1, self.input_width, self.label_width*len(self.known_cov_indices)])
  inputs = tf.concat([history_inputs, future_known_inputs_flat], axis=-1)

  # 原有标签提取逻辑保留
  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)

  # 固定形状
  inputs.set_shape([None, self.input_width, None])
  labels.set_shape([None, self.label_width, None])

  return inputs, labels

WindowGenerator.split_window = split_window

第三步:实例化WindowGenerator

提前定义好已知协变量的列名(除了Price之外的所有列即可):

OUT_STEPS = 24
INPUT_WIDTH = 336
# 构造已知协变量列表:所有不是Price的列
known_cov_cols = [col for col in train_df.columns if col != 'Price']
w1 = WindowGenerator(input_width=INPUT_WIDTH, label_width=OUT_STEPS, shift=OUT_STEPS, 
                     label_columns=['Price'], known_cov_columns=known_cov_cols)

形状验证

你可以运行原有示例代码验证形状是否符合预期:

example_window = tf.stack([np.array(test_df[:w1.total_window_size])])
example_inputs, example_labels = w1.split_window(example_window)

print('All shapes are: (batch, time, features)')
print(f'Window shape: {example_window.shape}')
print(f'Inputs shape: {example_inputs.shape}') # 特征维度会变成15 + 24*14=351,符合预期
print(f'labels shape: {example_labels.shape}')

可选优化

如果你用的是序列到序列(Seq2Seq)结构的LSTM,不需要广播未来协变量,直接把history_inputs作为编码器输入,future_known_inputs作为解码器输入即可,输出和原来的标签对齐。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.10.05 01:51:01