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

使用Keras与Transformers训练多模型时遇InvalidArgumentError求助

Keras与Transformers多模型训练故障排查求助

模型1:LSTM实现代码

(X_train, y_train), (X_test, y_test) = keras.datasets.imdb.load_data()
X_train = X_train[:2500]
y_train = y_train[:2500]
X_test = X_test[:500]
y_test = y_test[:500]

def dekodeeri(tekstijada):
    # Abifunktsioon numbritest tagasi tähtede saamiseks
    word_index = keras.datasets.imdb.get_word_index()
    index_word = {0: "<PAD>", 1: "<START>", 2: "<UNK>", 3: "<UNUSED>"}
    index_word[1] = "[START]"
    index_word[2] = "[OOV]"
    for (word, i) in word_index.items():
        index_word[i + 3] = word
    return " ".join(index_word[i] for i in tekstijada)

print(X_train.shape,y_train.shape)
print(X_test.shape,y_test.shape)
print()
print(X_train[0])
print(dekodeeri(X_train[0]))
print(y_train[0])

max_features = 100000
maxlen = 500

# Padding sequences
print('Pad sequences (samples x time)')
X_train = sequence.pad_sequences(X_train, maxlen=maxlen)
X_test = sequence.pad_sequences(X_test, maxlen=maxlen)
print('X_train shape:', X_train.shape)
print('X_test shape:', X_test.shape)

model = Sequential()
model.add(Embedding(max_features, 256))
model.add(SpatialDropout1D(0.4))
model.add(LSTM(100, dropout=0.2, recurrent_dropout=0.2))
model.add(Dense(1, activation='sigmoid'))
model.compile(loss='binary_crossentropy', optimizer='adam', metrics=['accuracy'])

# Train the model
model.fit(X_train, y_train, batch_size=64, epochs=5, validation_data=(X_test, y_test))

模型2:DistilBERT实现代码

import tensorflow as tf
from transformers import TFDistilBertModel, DistilBertConfig
from tensorflow.keras.models import Model
from tensorflow.keras.layers import Input, Dense, GlobalAveragePooling1D

# Define input shape
input_layer = Input(shape=(maxlen,), dtype=tf.int32)

# Load DistilBERT model
config = DistilBertConfig(dropout=0.2, attention_dropout=0.2)
distil_bert_model = TFDistilBertModel.from_pretrained('distilbert-base-uncased', config=config)

# Freeze DistilBERT layers
for layer in distil_bert_model.layers:
    layer.trainable = False

# Get DistilBERT output
distil_bert_output = distil_bert_model(input_layer)[0]

# Add pooling layer
pooled_output = GlobalAveragePooling1D()(distil_bert_output)

# Add dense layer for classification
output_layer = Dense(1, activation='sigmoid')(pooled_output)

# Create model
model_2_1 = Model(inputs=input_layer, outputs=output_layer)

# Compile the model
model_2_1.compile(loss='binary_crossentropy', optimizer='adam', metrics=['accuracy'])

# Train the model
model_2_1.fit(X_train_padded, y_train, batch_size=64, epochs=5, validation_data=(X_test_padded, y_test))

报错信息

InvalidArgumentError: Graph execution error. InvalidArgumentError Traceback (most recent call last)  in <cell line: 36>() 34 model_2_1.summary() 35 # Hindamistulemuste saamine ---> 36 model_2_1.fit(X_train_padded, y_train, batch_size=64, epochs=5, validation_data=(X_test_padded, y_test)) 1 frames /usr/local/lib/python3.10/dist-packages/tensorflow/python/eager/execute.py in quick_execute(op_name, num_outputs, inputs, attrs, ctx, name) 52 try:

排查建议

  1. 输入数据适配性检查
    • DistilBERT要求输入的token ID必须在自身预训练词表范围内(distilbert-base-uncased词表大小约30522),不能直接使用Keras IMDB数据集的原始索引。需用DistilBertTokenizer.from_pretrained('distilbert-base-uncased')重新处理文本,生成符合要求的input_ids和attention_mask后再传入模型。
  2. 变量名一致性验证
    • 模型1预处理后的变量是X_train/X_test,但模型2使用X_train_padded/X_test_padded,需确认这些变量已正确定义,且数据格式与模型输入要求匹配。
  3. 版本兼容性排查
    • 确保transformers与TensorFlow版本匹配,建议使用transformers>=4.20.0搭配tensorflow>=2.8.0,避免版本不兼容引发图执行错误。
  4. 完整报错信息获取
    • 当前报错信息截断,需打印完整错误堆栈,明确是维度不匹配、数据类型错误还是其他操作引发的异常,这是定位问题的核心依据。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.27 02:34:56