新版TensorFlow Sequential无predict_classes属性文本生成报错解决
报错修复:AttributeError: 'Sequential' object has no attribute 'predict_classes'
报错原因
该报错出现在TensorFlow 2.6及以上版本中,官方已经彻底移除了Keras Sequential/Functional模型的predict_classes方法,网上多数旧教程基于低版本TensorFlow编写,直接复用代码就会触发该错误。
自行修改代码无效的核心原因
你之前尝试用np.argmax搭配predict()替代predict_classes的思路是正确的,但代码存在两处逻辑错误导致输出异常:
- 维度取值错误:
model.predict()返回的是形状为(batch_size, 词表大小)的概率矩阵,np.argmax按轴取值后返回的是长度为batch_size的一维数组,直接用数组和词表存储的整数索引做相等判断,永远无法匹配到正确的词 - 缩进逻辑错误:原代码中拼接种子文本、追加生成词的逻辑被错误写到了遍历词表的for循环内部,会导致还没匹配到目标词就提前拼接空值,甚至第一次循环就直接返回结果,根本无法完成多轮词生成
修复后的可用代码
import numpy as np from tensorflow.keras.preprocessing.sequence import pad_sequences def generate_text_seq(model, tokenizer, text_seq_length, seed_text, n_words): generated_words = [] current_input = seed_text for _ in range(n_words): # 文本转序列、补长截断,匹配模型输入维度要求 encoded_seq = tokenizer.texts_to_sequences([current_input])[0] encoded_seq = pad_sequences([encoded_seq], maxlen=text_seq_length, truncating='pre') # 预测概率分布,取最大概率对应索引,效果完全等价于旧版predict_classes pred_result = model.predict(encoded_seq, verbose=0) target_index = np.argmax(pred_result, axis=-1)[0] # 匹配索引对应的目标词 output_word = "" for word, idx in tokenizer.word_index.items(): if idx == target_index: output_word = word break # 更新下一轮预测的输入文本,记录当前生成的词 current_input += " " + output_word generated_words.append(output_word) return " ".join(generated_words)
扩展提示
如果想要生成的文本更具多样性、避免内容生硬重复,可以不用固定选取概率最高的词,而是按预测得到的概率分布做随机采样,替换掉上述代码中np.argmax的取值逻辑即可。
内容的提问来源于stack exchange,提问作者Vijay Prasanna Gurubaran
相关产品推荐
相关产品推荐

