使用predict函数遍历DataFrame出错:所有行预测结果一致
问题分析与解决
你的代码存在两个关键错误,导致所有行的预测结果一致:
错误1:predict函数未使用传入参数,硬编码固定文本
predict函数里使用了外部变量test_text,完全忽略了传入的input参数。不管你传入什么内容,函数都会基于test_text生成预测结果,自然所有行输出相同。
修正后的predict函数:
def predict(input_text): predict_input = loaded_tokenizer.encode(str(input_text), truncation=True, padding=True, return_tensors="tf") output = loaded_model(predict_input)[0] prediction_value = tf.argmax(output, axis=1).numpy()[0] return prediction_value
错误2:apply时传入整个列而非当前行文本
遍历代码中,lambda row: predict(df['text'])是把整个text列传给predict函数,而非当前行的row['text']。结合第一个错误,进一步导致结果全部相同。
修正后的遍历代码:
df['pred'] = df.apply(lambda row: predict(row['text']), axis=1)
内容的提问来源于stack exchange,提问作者Kamil Filipek
相关产品推荐
相关产品推荐

