ValueError: NumPy数组转Tensor失败(不支持list类型) IMDB分类问题排查
问题修复方案
报错根因
报错ValueError: Failed to convert a NumPy array to a Tensor (Unsupported object type list)的直接原因是测试集未做长度统一的填充处理,原生IMDB数据集中的影评是长度不等的整数列表,无法直接转换为Tensor输入模型。
代码存在的所有问题
- 序列填充步骤赋值错误
原代码中第二行填充语句错误将处理后的测试集覆盖到了train_data变量上,既丢失了训练集的填充结果,也没有对test_data完成填充,导致evaluate阶段传入的test_data仍是嵌套列表格式触发报错。 - 训练数据与标签不匹配
被覆盖后的train_data实际为测试集内容,和train_labels的标签完全错位,即使解决报错,模型也无法学习到有效特征。 - Embedding层参数配置不合理
加载数据集时指定了num_words=10000,说明输入序列的最大索引不会超过9999,Embedding层的输入维度设为10000即可,88000的参数会造成不必要的资源浪费。 - 变量名拼写不规范(非报错原因,建议修正)
标签变量名拼写为train_lables/test_lables,正确拼写为train_labels/test_labels,前后统一的情况下不影响运行,建议修正避免后续维护出错。
修正后的完整代码
import tensorflow as tf from tensorflow import keras import numpy imdb = keras.datasets.imdb (train_data, train_labels), (test_data, test_labels) = imdb.load_data(num_words=10000) _word_index = imdb.get_word_index() word_index = {k:(v+3) for k,v in _word_index.items()} word_index["<PAD>"] = 0 word_index["<START>"] = 1 word_index["<UNK>"] = 2 word_index["<UNUSED>"] = 3 reverse_word_index = dict([(value, key) for (key, value) in word_index.items()]) def decode_review(text): return " ".join([reverse_word_index.get(i, "?") for i in text]) # 修正填充逻辑,分别处理训练集和测试集 train_data = keras.preprocessing.sequence.pad_sequences(train_data, value=word_index["<PAD>"], padding="post", maxlen=250) test_data = keras.preprocessing.sequence.pad_sequences(test_data, value=word_index["<PAD>"], padding="post", maxlen=250) model = keras.Sequential() # 修正Embedding层输入维度 model.add(keras.layers.Embedding(10000, 16)) model.add(keras.layers.GlobalAveragePooling1D()) model.add(keras.layers.Dense(16, activation="relu")) model.add(keras.layers.Dense(1, activation="sigmoid")) model.summary() model.compile(optimizer="adam", loss="binary_crossentropy", metrics=["accuracy"]) x_val = train_data[:10000] x_train = train_data[10000:] y_val = train_labels[:10000] y_train = train_labels[10000:] fitModel = model.fit(x_train, y_train, epochs=40, batch_size=512, validation_data=(x_val, y_val), verbose=1) results = model.evaluate(test_data, test_labels) print(results)
内容的提问来源于stack exchange,提问作者Codi Fredericks
相关产品推荐
相关产品推荐

