如何在sklearn 20NewsGroups数据集上剪枝词汇优化朴素贝叶斯分类器
解决方案:20NewsGroups文本剪枝与朴素贝叶斯训练流程
嘿,我来帮你搞定这个文本剪枝+分类器训练的事儿!你已经拿到了20NewsGroups的训练/测试集,还有邮箱匹配正则,接下来只需要把「移除邮箱、清理特定字符串、过滤停用词」这几个逻辑整合到sklearn的文本预处理流程里就行,全程用代码就能串起来,步骤很清晰:
1. 先搞定依赖与基础准备
首先导入需要的库,同时定义好你的停用词集合(你提到的an/the/is这些,也可以扩展成更全面的停用词库):
import re from sklearn.datasets import fetch_20newsgroups from sklearn.feature_extraction.text import CountVectorizer from sklearn.naive_bayes import MultinomialNB from sklearn.pipeline import make_pipeline # 自定义停用词(你可以根据需求扩展) custom_stop_words = {"an", "the", "is", "are", "and", "or", "a", "in", "on"} # 替换成你已有的邮箱正则表达式,这里给个通用示例 EMAIL_PATTERN = r'\b[A-Za-z0-9._%+-]+@[A-Za-z0-9.-]+\.[A-Z|a-z]{2,}\b'
如果你想用更全面的英文停用词库,可以用nltk的stopwords(需要先下载一次),替换上面的custom_stop_words就行:
import nltk nltk.download('stopwords') from nltk.corpus import stopwords custom_stop_words = set(stopwords.words('english'))
2. 写一个文本清理函数
把「移除邮箱、清理特定字符串」的逻辑都塞到这个函数里,作为预处理步骤:
# 要移除的特定字符串(比如你说的w32w,可以添加更多) TARGET_STRINGS = {"w32w"} def clean_text(text): # 第一步:用正则移除所有邮箱地址 text = re.sub(EMAIL_PATTERN, '', text) # 第二步:移除指定的特定字符串 for s in TARGET_STRINGS: text = text.replace(s, '') # 第三步:统一转小写(可选,但能让相同单词的大小写形式合并) text = text.lower() return text
可选:批量匹配类似w32w的字符串
如果你的目标字符串是「字母+数字+字母」这类模式(不止w32w),可以用正则批量替换,不用一个个列出来:
# 替换TARGET_STRINGS的处理逻辑,用正则匹配所有字母数字混合的短字符串 TARGET_PATTERN = r'\b[a-zA-Z]+[0-9]+[a-zA-Z]+\b' text = re.sub(TARGET_PATTERN, '', text)
3. 构建完整的训练Pipeline
用sklearn的make_pipeline把「文本预处理(词袋模型)+ 朴素贝叶斯分类器」串起来,这样数据会自动按流程处理:
# 加载数据集(如果你已经加载过,可以跳过这步) train_data = fetch_20newsgroups(subset='train', remove=('headers', 'footers', 'quotes')) test_data = fetch_20newsgroups(subset='test', remove=('headers', 'footers', 'quotes')) # 构建管道:先做文本清理→生成词袋→训练朴素贝叶斯 model = make_pipeline( CountVectorizer( preprocessor=clean_text, # 用我们写的清理函数预处理文本 stop_words=custom_stop_words, # 自动过滤停用词 # 可选优化:过滤低频词,比如只保留出现过至少2次的词 min_df=2, # 可选优化:加入n元组特征,比如同时考虑单个词和词组 # ngram_range=(1,2) ), MultinomialNB() ) # 训练模型 model.fit(train_data.data, train_data.target) # 测试准确率 accuracy = model.score(test_data.data, test_data.target) print(f"模型测试准确率: {accuracy:.2f}")
4. 为什么这么做?
- 用
preprocessor参数:它会在分词之前对整个文本做处理,先把邮箱、垃圾字符串删掉,避免这些内容被拆分成无意义的token。 - 停用词过滤:直接用CountVectorizer的
stop_words参数,比自己手动过滤更高效,sklearn会自动在分词后移除这些词。 - Pipeline的好处:把预处理和模型打包成一个整体,避免训练/测试数据的预处理不一致(比如数据泄露问题)。
内容的提问来源于stack exchange,提问作者tushariyer
相关产品推荐
相关产品推荐

