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

调用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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.06 16:25:33