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

TensorFlow2生成式LSTM模型训练时维度不匹配问题求助

生成式LSTM模型训练维度不兼容问题解决

问题现象

构建生成式LSTM模型时,已设置LSTM(256, return_sequences=True),但训练时抛出维度不兼容错误:

ValueError: Exception encountered when calling layer 'caption_model' (type Functional).
Input 0 of layer "lstm_5" is incompatible with the layer: expected ndim=3, found ndim=2. Full shape received: (None, 128)

代码与模型信息

核心模型代码

vec = layers.TextVectorization(output_sequence_length=maxlen,
                               max_tokens=379)
vec.adapt(use_for_vocab(train)) 

voc = vec.get_vocabulary()
voc_size = len(voc)

embed = layers.Embedding(input_dim=voc_size,
                         output_dim=256,
                         mask_zero=True)

inp_word = layers.Input(shape=(maxlen+2,), # maxlen为文本中句子的最大长度
                   name="word_input")      # 加2是为了容纳开始和结束标记
x_word = embed(inp_word)
x_word = layers.Dropout(0.5)(x_word)
x_word = layers.LSTM(256, return_sequences=True)(x_word)
ops_word = layers.GlobalAveragePooling1D(name="word_gap")(x_word)

模型摘要

Layer (type)Output ShapeParam #
word_input (InputLayer)[(None, 35)]0
embedding_1 (Embedding)(None, 35, 128)45184
dropout_6 (Dropout)(None, 35, 128)0
lstm_5 (LSTM)(None, 35, 256)394240
word_gap (GlobalAveragePooling1D)(None, 256)0

Total params: 439,424
Trainable params: 439,424
Non-trainable params: 0

数据集处理代码

from tensorflow.data import Dataset, AUTOTUNE

def format_dataset(x, y):
  y = Dataset.from_tensor_slices(y)
  data = Dataset.zip((x, y))
  return data.batch(32).prefetch(AUTOTUNE)

问题根源

模型定义中,Embedding层预期输入为(batch_size, sequence_length)(即(None,35)),输出为(batch_size,35,128),但训练时输入到LSTM层的张量形状是(None,128),说明:

  1. 训练数据的输入维度错误,丢失了序列长度维度,导致Embedding输出变成(batch_size,128)而非预期的三维张量。
  2. 数据集处理环节可能未正确保留序列维度,比如输入x是一维张量(每个样本是单个标量而非序列),批量后无法形成二维的序列批量。

解决方案

  1. 验证输入数据形状
    在训练前打印数据集的元素形状,确认输入维度是否符合预期:

    for batch in train_dataset.take(1):
        print("Input shape:", batch[0].shape)  # 应输出 (32, 35)
        print("Label shape:", batch[1].shape)
    
  2. 修正数据集处理逻辑
    确保输入x是包含序列的数据集,若x是未封装的张量,需先转换为Dataset:

    def format_dataset(x, y):
        x_ds = Dataset.from_tensor_slices(x)  # 确保x是形状为(num_samples, 35)的张量
        y_ds = Dataset.from_tensor_slices(y)
        data = Dataset.zip((x_ds, y_ds))
        return data.batch(32).prefetch(AUTOTUNE)
    
  3. 确认文本向量化结果
    检查vec处理后的样本形状,确保每个输出是长度为maxlen+2的序列:

    sample_text = train[0]
    vec_sample = vec(sample_text)
    print("Vectorized sample shape:", vec_sample.shape)  # 应输出 (35,)
    

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.25 20:45:11