Keras IMDB文本二分类训练报错:无法将NumPy数组转为Tensor
解决Keras IMDB数据集训练时的ValueError(无法将NumPy数组转换为Tensor)
Hey there! Let's fix this error you're hitting. The problem here is straightforward: the raw IMDB sequences you're loading are variable-length integer lists, and Keras can't directly convert those into a Tensor for training—your Embedding layer expects uniformly shaped input data.
解决方案:统一序列长度
The fix is to pad (or truncate) all your sequences to the same length using keras.preprocessing.sequence.pad_sequences. Here's your adjusted code with this critical step added:
import tensorflow as tf from tensorflow import keras import numpy as np data = keras.datasets.imdb (x_train,y_train),(x_test,y_test) = data.load_data() dictionary = data.get_word_index() dictionary = {k:(v+3) for k,v in dictionary.items()} dictionary['<PAD>'] = 0 dictionary['<START>'] = 1 dictionary['<UNKNOWN>'] = 2 dictionary['<UNUSED>'] = 3 dictionary = dict([(v,k) for (k,v) in dictionary.items()]) # ---------------------- 新增的序列统一处理步骤 ---------------------- # 将所有序列调整为256长度,超长截断,不足则用<PAD>(对应值0)补在末尾 max_sequence_length = 256 x_train = keras.preprocessing.sequence.pad_sequences( x_train, value=dictionary['<PAD>'], padding='post', maxlen=max_sequence_length ) x_test = keras.preprocessing.sequence.pad_sequences( x_test, value=dictionary['<PAD>'], padding='post', maxlen=max_sequence_length ) # ------------------------------------------------------------------- model = keras.Sequential([ keras.layers.Embedding(10000,16), keras.layers.GlobalAveragePooling1D(), keras.layers.Dense(16,activation='relu'), keras.layers.Dense(1,activation='sigmoid') ]) model.compile( optimizer='adam', loss='binary_crossentropy', metrics=['accuracy'] ) print(model.summary()) history = model.fit(x_train,y_train,epochs=50,batch_size=32,verbose=1) prediction = model.predict(x_test) print(prediction)
关键步骤解释
pad_sequences把原本长度不一的列表转换成了二维NumPy数组(形状为(样本数量, 统一序列长度)),这样TensorFlow就能正常将其转换为模型可接受的Tensor。padding='post'把补全的<PAD>token放在序列末尾,这种方式更适配你用的GlobalAveragePooling1D层,不会影响有效文本的特征提取。- 你可以根据需求调整
max_sequence_length,常用值有100、256或512;如果设置的长度短于部分序列,默认会截断序列末尾的多余token。
内容的提问来源于stack exchange,提问作者Philip Purwoko
相关产品推荐
相关产品推荐

