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 Shape | Param # |
|---|---|---|
| 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),说明:
- 训练数据的输入维度错误,丢失了序列长度维度,导致Embedding输出变成
(batch_size,128)而非预期的三维张量。 - 数据集处理环节可能未正确保留序列维度,比如输入x是一维张量(每个样本是单个标量而非序列),批量后无法形成二维的序列批量。
解决方案
验证输入数据形状
在训练前打印数据集的元素形状,确认输入维度是否符合预期:for batch in train_dataset.take(1): print("Input shape:", batch[0].shape) # 应输出 (32, 35) print("Label shape:", batch[1].shape)修正数据集处理逻辑
确保输入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)确认文本向量化结果
检查vec处理后的样本形状,确保每个输出是长度为maxlen+2的序列:sample_text = train[0] vec_sample = vec(sample_text) print("Vectorized sample shape:", vec_sample.shape) # 应输出 (35,)
内容的提问来源于stack exchange,提问作者j raynukem
相关产品推荐
相关产品推荐

