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

使用自定义生成器喂入Keras model.fit时的形状不兼容问题求助

问题描述

使用自定义生成器函数向Keras LSTM模型喂入数据时,出现如下错误:

WARNING:tensorflow:Model was constructed with shape (None, 3177, 2) for input
KerasTensor(type_spec=TensorSpec(shape=(None, 3177, 2), dtype=tf.float32, name='masking_9_input'),
name='masking_9_input', description="created by layer 'masking_9_input'"), but it was called on an
input with incompatible shape (None, None).

生成器函数

当前使用的生成器代码:

def padded_generator(trajectories=trajectories, max_length=3177):
    X = []
    Y = []
    for trajectory in trajectories.values:
        curr_X = np.hstack([trajectory[0][0]])  
        curr_Y = np.hstack([trajectory[0][2]])
        temp = (np.hstack([trajectory[0][1:]]))
    
        for i, point in enumerate(temp):
            if i >= temp.shape[0] - 1: # Should break at second to last sample. 
                break
            curr_X = np.vstack((curr_X, point)) # Stack next point on existing X
            padded_X = np.squeeze(tf.keras.utils.pad_sequences([curr_X], 
                                                               maxlen= 3177, 
                                                               padding='post',
                                                               dtype=float, 
                                                               value=-10))
            curr_Y = temp[i+1] # Point added to X in next iter. is current target. 
            yield (padded_X, curr_Y)
        
data_gen = padded_generator()    

完整轨迹是如下形式的点数组:

[[-0.1843775   0.6867699 ]
 [-1.0841161  -3.0429556 ]
 [ 1.3582058  -0.6040352 ]
 [ 1.8754534  -1.7010269 ]
 ...
 [-2.4015598   0.3573116 ]
 [-1.3986164  -0.95052546]
 [-0.705326   -1.3387672 ]
 [-1.455082   -0.57572746]
 [-3.1130497  -2.7871382 ]]

每次调用生成器时,返回的填充后轨迹X和标签Y的形状:

Shape of X: (3177, 2)
Shape of Y: (2,)
模型定义

当前的模型结构:

model = Sequential()
model.add(tf.keras.layers.Masking(mask_value=-10, input_shape=(3177, 2)))
model.add(LSTM(100, activation='relu', return_sequences=True))
model.add(LSTM(50, activation='relu', return_sequences=True))
model.add(LSTM(25, activation='relu'))
model.add(Dense(10, activation='relu'))
model.add(Dense(2))
model.compile(optimizer='adam', loss='mse')

执行训练代码后触发错误:

model.fit(data_gen, verbose=1)
问题原因与解决方法

错误原因

Keras模型的输入要求带批量维度:模型输入形状(None, 3177, 2)里的None代表批量大小,而你的生成器返回的X是(3177,2),缺少了批量维度(也就是没有(1, 3177, 2)这个形状),导致模型无法识别输入的维度匹配。

解决方法

有两种简单的修改方式:

方法1:修改生成器,给返回的X增加批量维度

在生成器的yield语句前,给padded_X和curr_Y各增加一个维度,让输出形状匹配模型期望的批量输入:

def padded_generator(trajectories=trajectories, max_length=3177):
    for trajectory in trajectories.values:
        curr_X = np.hstack([trajectory[0][0]])  
        temp = (np.hstack([trajectory[0][1:]]))
    
        for i, point in enumerate(temp):
            if i >= temp.shape[0] - 1: 
                break
            curr_X = np.vstack((curr_X, point)) 
            padded_X = np.squeeze(tf.keras.utils.pad_sequences([curr_X], 
                                                               maxlen=3177, 
                                                               padding='post',
                                                               dtype=float, 
                                                               value=-10))
            # 给X增加批量维度,从(3177,2)变为(1,3177,2)
            padded_X = np.expand_dims(padded_X, axis=0)
            curr_Y = temp[i+1] 
            # 给Y增加批量维度,从(2,)变为(1,2)
            curr_Y = np.expand_dims(curr_Y, axis=0)
            yield (padded_X, curr_Y)

方法2:用tf.data.Dataset包装生成器,自动处理批量

如果你不想修改生成器,可以用TensorFlow的Dataset来包装生成器,显式指定输出形状并添加批量:

import tensorflow as tf

data_gen = padded_generator()
# 将生成器转换为Dataset,定义输出的形状和类型
dataset = tf.data.Dataset.from_generator(
    lambda: data_gen,
    output_signature=(
        tf.TensorSpec(shape=(3177, 2), dtype=tf.float32),
        tf.TensorSpec(shape=(2,), dtype=tf.float32)
    )
)
# 设置批量大小,可根据显存调整
dataset = dataset.batch(1)

# 使用Dataset训练模型
model.fit(dataset, verbose=1)

额外提示

  • 生成器里的curr_Y = np.hstack([trajectory[0][2]])看起来有问题:你的轨迹是2维点数组,没有第三列,建议检查这部分逻辑是否是笔误。
  • LSTM层通常更适合用tanh作为激活函数,而非relu,可以尝试替换看看效果。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.13 08:50:23