基于MLPClassifier的文本分类:训练后模型预测及代码实现咨询
基于MLPClassifier的文本分类预测实现
我来帮你梳理并完善这个基于MLP的文本分类预测流程,结合你给出的代码片段,咱们一步步来:
1. 模型Pipeline初始化(NeuralNetworkModel类)
你已经实现了NeuralNetworkModel类的静态方法_init_pipeline,用来构建从文本预处理到模型的完整流程,这里我把代码整理并做下说明:
from sklearn.pipeline import Pipeline from sklearn.feature_extraction.text import CountVectorizer, TfidfTransformer from sklearn.neural_network import MLPClassifier class NeuralNetworkModel: @staticmethod def _init_pipeline(stopwords): pipe_line = Pipeline([ # 第一步:将文本转换为词袋特征,去除停用词并统一小写 ("vect", CountVectorizer(stop_words=stopwords, lowercase=True)), # 第二步:将词袋特征转换为TF-IDF权重特征,提升重要词汇的权重占比 ("tfidf", TfidfTransformer()), # 第三步:初始化MLP分类器,设置两层各15个神经元的隐藏层,用SGD优化器训练 ("clf", MLPClassifier( hidden_layer_sizes=(15, 15), activation="relu", solver="sgd", max_iter=3000, tol=1e-6, verbose=True )) ]) return pipe_line
这个Pipeline把文本预处理和模型整合到了一起,训练时直接调用fit方法就能完成从文本到模型训练的全流程,避免了预处理和模型环节的特征不一致问题。
2. 预测类TextClassificationPredict的完善
你提到已经创建了TextClassificationPredict类,目前只给出了__init__的开头,这里我帮你补全核心的预测逻辑,让它能加载训练好的模型并完成文本分类预测:
class TextClassificationPredict(object): def __init__(self, trained_model): # 初始化时传入训练好的完整Pipeline模型 self.trained_model = trained_model def predict_single_text(self, text): """ 对单条文本进行分类预测 :param text: 待预测的原始文本字符串 :return: 预测的类别结果 """ # 注意传入的文本要放在列表中,符合sklearn模型的输入格式 return self.trained_model.predict([text])[0] def predict_batch_texts(self, texts): """ 对多条文本进行批量分类预测 :param texts: 待预测的原始文本列表 :return: 每个文本对应的预测类别列表 """ return self.trained_model.predict(texts)
使用示例
假设你已经用NeuralNetworkModel完成了模型训练,就可以这样使用预测类:
# 假设已经完成模型训练,得到训练好的Pipeline # stopwords = 你的停用词列表 # trained_pipeline = NeuralNetworkModel._init_pipeline(stopwords).fit(train_texts, train_labels) # 初始化预测类实例 predictor = TextClassificationPredict(trained_pipeline) # 单文本预测 sample_text = "这里是你需要分类的测试文本内容" predicted_label = predictor.predict_single_text(sample_text) print(f"该文本的预测类别:{predicted_label}") # 批量文本预测 sample_texts = ["测试文本1", "测试文本2", "测试文本3"] predicted_labels = predictor.predict_batch_texts(sample_texts) for text, label in zip(sample_texts, predicted_labels): print(f"文本:{text},预测类别:{label}")
关键注意点
- 训练和预测时要使用完全一致的停用词表,否则会导致预处理后的特征维度不匹配
- 如果需要将模型保存到本地后再加载使用,建议用
joblib库完整保存整个Pipeline对象,而不是单独保存MLP模型,这样能保证预处理流程和模型的一致性
内容的提问来源于stack exchange,提问作者December
相关产品推荐
相关产品推荐

