调用model.fit训练RNN模型时触发InvalidArgumentError: Graph execution error
问题排查与修复:RNN训练时的InvalidArgumentError
错误根源分析
你的代码触发InvalidArgumentError主要有两个核心问题:
- 损失函数不匹配:
sparse_categorical_crossentropy要求目标张量是单维度的类别索引(对应每个样本预测一个类别),但你的output_sequences是二维张量(形状[batch_size, max_seq_len]),每个样本是一个词序列,每个位置都是词索引,两者不兼容。 - 模型结构不匹配:当前LSTM层默认只返回最后一个时间步的输出(形状
[batch_size, 64]),但你需要预测整个输出序列的每个时间步,所以LSTM需要返回所有时间步的输出。
修复步骤与完整代码
1. 调整模型结构,让LSTM返回序列
修改LSTM层,添加return_sequences=True参数,使LSTM输出每个时间步的特征(形状[batch_size, max_seq_len, 64]),确保后续能对每个时间步的词进行预测。
2. 适配序列预测的输出层
使用TimeDistributed包裹Dense层,将分类逻辑应用到每个时间步的输出上,保证每个位置都能生成对应词的概率分布。
修复后的完整代码
import nltk nltk.download('punkt') nltk.download('wordnet') nltk.download('stopwords') nltk.download('omw-1.4') def preprocess_data(data): # Tokenize as sentences data = [nltk.word_tokenize(sent) for sent in data] # Remove stopwords stopwords = nltk.corpus.stopwords.words('portuguese') data = [[word for word in sent if word not in stopwords] for sent in data] # Perform lemmatization lemmatizer = nltk.stem.WordNetLemmatizer() data = [[lemmatizer.lemmatize(word) for word in sent] for sent in data] return data # Example usage data = ["oi tudo bem?", "sim, eu estou com fome e voce?", "eu estou, quero uma pizza", "vou pedir duas pizzas", "Gosto de livros", "quais livros voce gosta de ler?", "oi tudo bem?", "sim, eu estou com fome e voce?", "eu estou, quero uma pizza", "vou pedir duas pizzas", "Gosto de livros", "quais livros voce gosta de ler?", "eu gosto de ler thrillers e ficção científica", "e você, o que mais gosta de ler?", "eu também gosto de ler romances e livros de autoajuda", "qual é o seu livro favorito?", "um dos meus livros favoritos é o 'O Alquimista' de Paulo Coelho", "eu também gosto muito desse livro! Qual é o seu gênero literário favorito?", "eu gosto de todos os gêneros, mas talvez meu favorito seja a ficção científica", "eu também adoro ficção científica. Qual é o seu livro de ficção científica favorito?", "meu livro de ficção científica favorito é 'Dune' de Frank Herbert", "que legal, eu também gosto muito de 'Dune'! Já leu algum outro livro do Frank Herbert?", "sim, eu também gostei muito de 'O Imperador-Deus de Dune' e 'Herejia de Dune'"] processed_data = preprocess_data(data) print(processed_data) import tensorflow as tf # Define the input and output sequences input_sequences = processed_data[:-1] output_sequences = processed_data[1:] def flatten(l): return [item for sublist in l for item in sublist] # Create a vocabulary of unique words vocab = sorted(set(flatten(input_sequences + output_sequences))) # Create word-to-index and index-to-word mappings word_to_index = {word: i for i, word in enumerate(vocab)} index_to_word = {i: word for i, word in enumerate(vocab)} # Convert the input and output sequences to integers input_sequences = [[word_to_index[word] for word in seq] for seq in input_sequences] output_sequences = [[word_to_index[word] for word in seq] for seq in output_sequences] # Find the maximum sequence length max_seq_len = max(len(seq) for seq in input_sequences + output_sequences) # Pad the sequences with zeros to the maximum sequence length input_sequences = tf.keras.preprocessing.sequence.pad_sequences( input_sequences, maxlen=max_seq_len, padding='post') output_sequences = tf.keras.preprocessing.sequence.pad_sequences( output_sequences, maxlen=max_seq_len, padding='post') # Create a dataset from the padded sequences batch_size = 32 # Replace 32 with the desired batch size dataset = tf.data.Dataset.from_tensor_slices( (input_sequences, output_sequences)).batch(batch_size) # Define the RNN model model = tf.keras.Sequential([ tf.keras.layers.Embedding(len(vocab), 64, input_length=max_seq_len), tf.keras.layers.LSTM(64, return_sequences=True), # 返回所有时间步输出 tf.keras.layers.TimeDistributed(tf.keras.layers.Dense(len(vocab), activation='softmax')) # 每个时间步独立预测 ]) # Compile the model model.compile(optimizer='adam', loss='sparse_categorical_crossentropy', metrics=['accuracy']) # Train the model model.fit(dataset, epochs=5)
额外说明
sparse_categorical_crossentropy本身支持二维目标序列,无需修改损失函数,直接适配序列预测场景。- 你的数据集样本量较小(仅22个),训练时容易过拟合,可考虑增加数据量,或添加
tf.keras.layers.Dropout(0.2)层做正则化处理。
内容的提问来源于stack exchange,提问作者Lucas Torres
相关产品推荐
相关产品推荐

