TensorFlow文本分类模型的保存、复用及测试方法咨询
嘿,我来帮你理顺这些问题!你碰到的ValueError: No variables to save其实是因为tf.train.Saver和tf.estimator这套高层API不兼容——estimator有自己的模型保存机制,完全不用手动折腾Saver。下面一步步给你解决:
1. 解决模型保存问题(替换
tf.train.Saver) tf.estimator.DNNClassifier自带模型保存功能,你只需要在创建分类器时指定model_dir参数,训练过程中它会自动把模型 checkpoint 保存到这个目录里,根本不需要调用tf.train.Saver。示例代码如下:
import tensorflow as tf import tensorflow_hub as hub # 定义TF Hub文本嵌入特征列 embedded_text_feature_column = hub.text_embedding_column( key="sentence", module_spec="https://tfhub.dev/google/nnlm-en-dim128/2" ) # 创建分类器时指定模型保存路径 classifier = tf.estimator.DNNClassifier( hidden_units=[64, 32], feature_columns=[embedded_text_feature_column], n_classes=2, model_dir="./sentiment_model" # 模型会自动保存在这个文件夹 ) # 训练代码(和你之前的逻辑一致) train_input_fn = ... # 你的训练输入函数 classifier.train(input_fn=train_input_fn, steps=1000)
训练完成后,./sentiment_model目录下会生成checkpoint文件、变量文件等,下次直接加载就能复用。
2. 加载已保存的模型
要复用模型,只需要用同一个model_dir重新创建DNNClassifier,estimator会自动加载目录里最新的checkpoint,无需重新训练:
# 加载已保存的模型,参数要和训练时完全一致 loaded_classifier = tf.estimator.DNNClassifier( hidden_units=[64, 32], feature_columns=[embedded_text_feature_column], n_classes=2, model_dir="./sentiment_model" )
3. 用自有测试集评估模型
把你的测试数据转换成estimator能识别的输入函数,然后调用evaluate方法就能得到测试集的评估指标(准确率、损失等):
import pandas as pd # 定义测试集输入函数(假设你的测试数据是CSV格式,包含sentence和label列) def test_input_fn(df): return tf.compat.v1.estimator.inputs.pandas_input_fn( x={"sentence": df["sentence"]}, y=df["label"], batch_size=32, shuffle=False # 测试集不需要打乱 ) # 加载你的测试数据(替换成你自己的加载逻辑) test_df = pd.read_csv("your_test_dataset.csv") # 运行评估 eval_results = loaded_classifier.evaluate(input_fn=lambda: test_input_fn(test_df)) print("测试集评估结果:") for key, value in eval_results.items(): print(f"{key}: {value:.4f}")
如果你的测试数据是numpy数组格式,可以用tf.compat.v1.estimator.inputs.numpy_input_fn来创建输入函数。
4. 用单条/样本文本做预测
针对单条或多条样本文本,同样需要先构建输入函数,再调用predict方法获取结果:
import numpy as np # 定义预测用的输入函数 def predict_input_fn(sentences): return tf.compat.v1.estimator.inputs.numpy_input_fn( x={"sentence": np.array(sentences)}, batch_size=1, shuffle=False ) # 测试样本文本 sample_sentences = [ "This is the best meal I've ever had!", "The service was terrible and the food was cold." ] # 获取预测结果(predict返回生成器,需转成列表遍历) predictions = list(loaded_classifier.predict(input_fn=lambda: predict_input_fn(sample_sentences))) # 解析并打印结果 for sentence, pred in zip(sample_sentences, predictions): predicted_label = "正面情感" if pred["class_ids"][0] == 1 else "负面情感" confidence = pred["probabilities"].max() print(f"文本:{sentence}") print(f"预测结果:{predicted_label},置信度:{confidence:.4f}\n")
predict返回的结果里,class_ids是预测的类别索引,probabilities是每个类别的概率值。
内容的提问来源于stack exchange,提问作者John Aristo
相关产品推荐
相关产品推荐

