用BERT微调实现垃圾短信过滤:从数据预处理到Gradio部署
2026/9/16 5:51:57 网站建设 项目流程

简介:这是一份面向人工智能、深度学习的毕业设计项目资源,聚焦垃圾短信过滤场景,基于BERT模型搭建文本分类核心,并使用Gradio构建可交互的Web界面。资源压缩包共15个文件,包括Python源代码、pyc编译文件、JSON配置、PyTorch模型权重、词典文本、PDF介绍、二进制模型、CSV数据集及ipynb示例等,整体大小约为397MB。其中data_process.py与data_create.py负责原始短信数据的清洗和BERT输入格式转换,train.py完成模型微调训练,model.py定义网络结构,gradio_web.py提供前端可视化交互;另有训练好的pth权重、bert-base-chinese预训练参数、tokenizer配置等,可直接加载运行,实现输入短信内容后实时输出过滤结果,完整覆盖数据预处理、模型训练、界面部署全流程。目前已有137人学习与下载,适合正在准备毕业设计、课程设计,或希望从头实践BERT文本分类与Gradio应用开发的开发者参考学习。

1. 垃圾短信过滤为什么值得用BERT做一次完整的毕业设计

短信这种场景非常尴尬:广告、诈骗、验证码混在一起,长度通常不超过几十个汉字,却包含各种变形词、谐音和异体字。用传统规则库要不停维护关键词黑白名单,换一套话术就失效;用朴素贝叶斯这类词袋模型,又会把“恭喜您”和“您中奖了”拆散,丢了语序信息。而BERT这类预训练语言模型,能通过双向注意力捕捉短文本里的上下文关系,在几十毫秒内判断是否可信。这个项目把BERT微调、数据预处理、Web交互串成一条完整链路,正好适合人工智能或深度学习方向的毕业设计,也适合第一次想完整跑通NLP分类任务的人。

2. 数据预处理:从原始短信到BERT可读的token序列

2.1 data_create.py 在做什么:先把原始语料变成结构化数据

垃圾短信过滤的第一步不是训练,而是把原始短信整理成message.csv这样的结构化文件。项目里的data_create.py负责收集、合并和初筛原始文本,常见做法是从设备导出、公开数据集或人工标注的Excel中读取,统一字段为labeltext。label 用 0 表示正常短信,1 表示垃圾短信。字段设计可以参考下方表格:

字段名类型示例说明
labelint10为正常,1为垃圾短信
textstr恭喜您获得xx万元大奖原始短信内容,未清洗
sourcestrmanual来源标记,方便追溯数据质量问题

data_create.py时不需要复杂逻辑,重点是保证数据能重入、可追踪。我会在脚本里固定随机种子,避免每次运行shuffle结果不一致。比如:

import pandas as pd import re df = pd.read_excel("raw_sms.xlsx", engine="openpyxl") df = df[["label", "text"]].copy() df["text"] = df["text"].astype(str) df.to_csv("message.csv", index=False, encoding="utf-8-sig") print(df["label"].value_counts())

这段代码把Excel中的原始数据转成csv,encoding="utf-8-sig"是为了让Excel直接打开不乱码。处理完一定要打印类别分布,如果垃圾短信占比远低于正常短信,后面要做类别加权或过采样。

2.2 分词、编码与截断:data_process.py 的完整流程

data_process.py的主要任务是把中文短信转换成BERT的输入格式。BERT不能直接吃字符串,需要借助bert-base-chinese配套的vocab.txttokenizer.json做分词,再转成input_idsattention_masktoken_type_ids。中文BERT默认按字切分,所以不需要引入jieba。

我一般会限制max_length=128,因为垃圾短信一般不会超过这个长度,过长反而会把一些广告的尾部干扰带进来。代码结构大致如下:

import pandas as pd from transformers import BertTokenizer tokenizer = BertTokenizer.from_pretrained("bert-base-chinese") max_len = 128 df = pd.read_csv("message.csv") input_ids_list = [] attention_mask_list = [] labels = [] for text in df["text"].tolist(): encoded = tokenizer( text, max_length=max_len, padding="max_length", truncation=True, return_tensors="pt" ) input_ids_list.append(encoded["input_ids"]) attention_mask_list.append(encoded["attention_mask"]) labels.append(1 if df["label"] else 0) torch.save({"input_ids": input_ids_list, "labels": labels}, "dataset/processed.pt")

注意padding="max_length"会把不足128的token全部补成[PAD],这会让推理时计算量增大。如果追求效率,可以改为padding="longest",等训练完再统一处理。attention_mask会把真实token标记为1,[PAD]标记为0,模型就不会attention到pad位置。

2.3 数据拆分的两个细节:防治泄漏和类别不平衡

很多初学者在预处理时直接对整个数据集做编码、然后随机切分,这在短文本分类里问题不大,但如果原始数据里有重复短信,同一条字符串既出现在训练集又出现在验证集,评估结果就会虚高。另一种更隐蔽的泄漏是处理URL或号码时没有归一化,“12306”和“10086”这类数字串容易被模型当作强特征,换一批数据就失效。

我建议在data_process.py里提前做一次去重和分层采样:

from sklearn.model_selection import train_test_split X_train, X_val, y_train, y_val = train_test_split( df["text"], df["label"], test_size=0.2, stratify=df["label"], random_state=42 )

stratify会按label比例拆分,避免某一折全是正常短信。当垃圾短信占比不足10%时,这个参数非常关键,否则训练出来的模型只要预测“正常”就能拿到90%准确率,实战中没有意义。

3. 模型定义与训练:微调bert-base-chinese的分类头

3.1 model.py 中如何构建BERT分类模型

model.py在项目中承担的是模型结构定义。如果你直接用transformers库的BertForSequenceClassification,其实不需要自己写模型,但为了毕业设计答辩时能讲清结构,还是建议自己包一层。典型的写法是加载bert-base-chinese的权重,取出最后一层的pooler_output,再接一个Dropout和全连接层:

import torch.nn as nn from transformers import BertModel class BertSpamClassifier(nn.Module): def __init__(self, num_labels=2, dropout=0.3): super().__init__() self.bert = BertModel.from_pretrained("bert-base-chinese") self.dropout = nn.Dropout(dropout) self.classifier = nn.Linear(768, num_labels) def forward(self, input_ids, attention_mask): outputs = self.bert(input_ids=input_ids, attention_mask=attention_mask) pooled = outputs.pooler_output return self.classifier(self.dropout(pooled))

pooler_output是BERT对[CLS]这个token做线性变换和tanh激活后的结果,可以理解为整个句子的语义向量。这里的dropout=0.3是防过拟合用的,对短文本分类来说0.2到0.4之间都是合理选择。如果你的训练数据很少,不加载bert-base-chinese的预训练权重,而是随机初始化模型,效果会差很多,这也是这个项目必须带pytorch_model.binconfig.json的原因。

3.2 train.py 中的训练循环与超参配置

训练部分在train.py里实现,这是整个项目最核心的脚本。需要做四件事:加载dataset/processed.pt、初始化模型、设置优化器、迭代训练。最容易被忽略的是学习率,BERT微调通常用很小的学习率,2e-5或3e-5,这个量级比重新训练Embedding的收敛慢,但能保留预训练知识,避免灾难性遗忘。

from transformers import AdamW, get_linear_schedule_with_warmup model = BertSpamClassifier(num_labels=2) optimizer = AdamW(model.parameters(), lr=2e-5, correct_bias=False) total_steps = len(train_loader) * epochs scheduler = get_linear_schedule_with_warmup( optimizer, num_warmup_steps=int(total_steps * 0.1), num_training_steps=total_steps )

correct_bias=False是transformers库对AdamW的要求,它让优化器不把bias项也做weight decay。warmup约占总step的10%,先让学习率从小到大爬升,再线性衰减到0,这是BERT系列模型比较通用的设置。推荐的超参参考下表:

参数数值调整建议
max_length128短信长度超128可增加到256,但训练时间变长
batch_size16显存不够就降到8,学习率也要同步降
learning_rate2e-5类别不平衡时尝试1e-5
epochs3数据集很小用3,抗曲折
dropout0.3过拟合上调到0.4,欠拟合下调到0.1

训练循环里要记得调用model.train()model.eval(),BatchNorm和Dropout在两种模式下的行为不同。我只在torch.no_grad()下做验证,否则反向传播的中间变量会一直累积,显存爆掉一两次才学会这个习惯。

3.3 训练过程的监控与权重保存

train.py训练结束不要只打印最后一个epoch的loss,我习惯每个batch后打印loss,每个epoch结束做一次验证,并保存验证loss最好的模型,而不是保存最后一次迭代的权重。这样可以避免最后一轮因为学习率过低导致过拟合。保存方式有两种:一是保存整个模型torch.save(model.state_dict(), "weights/message.pth"),二是用transformers的save_pretrained("model")保存到model文件夹。

torch.save({ "model_state_dict": model.state_dict(), "epoch": epoch, "val_loss": best_loss, }, "weights/message.pth")

这样保存的message.pth除了权重还包含优化器状态和验证loss,恢复训练时直接load进来继续跑。如果只想推理,只保存model.state_dict()就够了。注意weightsmodel这两个目录最好同时保留,前者是自定义pth权重,后者是完整tokenizer和bert配置,因为Gradio加载时两者都要用。

4. Gradio 界面:把BERT模型变成可以实时交互的Web应用

4.1 gradio_web.py 的界面设计思路

Gradio 的最大价值是把模型推理包装成Web界面而不需要写前端。项目里的gradio_web.py就是这样一个最简入口:用户输入一段短信,后端调用BERT模型返回“垃圾短信”或“正常短信”。和传统Django/Flask方案比,Gradio节省了路由、模板和HTTP通信的代码,同时自带了请求排队和并发处理,适合快速演示和毕业设计答辩。

设计界面时分两步走:先写一个predict函数,参数是输入短信的字符串,返回值是label和概率;再用gr.Interfacegr.Blocks组装界面。用gr.Blocks可以放阈值滑块,这是演示时很出效果的功能。界面不需要复杂,一个输入框、一个输出标签、一个阈值滑块就够:

import gradio as gr import torch from transformers import BertTokenizer, BertForSequenceClassification model = BertForSequenceClassification.from_pretrained("model") tokenizer = BertTokenizer.from_pretrained("model") def predict(text, threshold=0.5): inputs = tokenizer(text, return_tensors="pt", truncation=True, max_length=128) with torch.no_grad(): logits = model(**inputs).logits prob = torch.softmax(logits, dim=1)[0][1].item() label = "垃圾短信" if prob >= threshold else "正常短信" return label, f"垃圾概率: {prob:.4f}" demo = gr.Interface( fn=predict, inputs=[gr.Textbox(lines=4, label="输入短信内容"), gr.Slider(0, 1, value=0.5, label="阈值")], outputs=[gr.Label(label="判定结果"), gr.Textbox(label="概率")], title="垃圾短信过滤系统" ) demo.launch(server_name="0.0.0.0", server_port=7860)

这个代码里的threshold参数会传给预测函数,允许看着概率值调整判定边界。如果短信文案有明显诈骗特征但阈值设为0.7,就会被判成正常短信,演示时调低阈值立刻能看到效果变化。

4.2 模型加载与推理的细节

gradio_web.py里加载模型时,我用BertForSequenceClassification.from_pretrained("model")而不是BertSpamClassifier,是因为model目录里存的是transformers结构,Gradio不需要感知自定义的forward逻辑。如果你的自定义模型想加载到Gradio,就必须重新实例化BertSpamClassifier,再loadmessage.pth,而且要把模型文件和权重分开管理。

还有一个重要细节:predict函数内部每次都会调用tokenizer,这没问题;但模型如果每次请求都reset一遍,体验会很差。正确做法是把模型和tokenizer放在gradio_web.py的全局作用域里,只加载一次。另外,推理外面一定要包上torch.no_grad(),否则显存会随着请求次数持续增长。如果用户一次输入多段短信,可以在predict里对输入列表做循环,或者直接batch推理。

4.3 身份验证与部署时的常见做法

Gradio 的InterfaceBlocks都支持auth参数,接收一个字典或函数。只需要三个账号的场景,用字典最直接:

demo.launch(auth=("admin", "your_password"))

如果想做多用户校验,传一个函数,函数接收用户名和密码做数据库查询,返回True才能访问。这个功能在实际部署到公网时非常有用,避免任何人都能调用你的模型白白烧显卡。至于部署方式,常见做法是直接在服务器上python gradio_web.py,也可以用nohup放到后台,或者用systemd管理进程。注意Gradio默认只监听本机127.0.0.1,想远程访问必须显式设置server_name="0.0.0.0",否则会出现本地能打开、手机打不开的问题;这是初学者最常踩的坑。

5. 验证效果与避坑:让BERT模型在真实短信上可用

5.1 用混淆矩阵和阈值找到最佳分类点

训练完模型不要只看准确率。在垃圾短信场景里,正常短信占总量的90%以上,把全部判成正常准确率也高,但漏掉的诈骗短信带来的风险远大于误杀一条广告。我惯用的办法是算一遍混淆矩阵,看带权F1或recall。比如验证集有1000条正常、100条垃圾,模型把80条垃圾识别出来、20条漏掉,那垃圾短信召回率是0.8,准确率是0.8,F1也是0.8,这样评价比单纯准确率有价值得多。

阈值调整可以直接利用第4章在Gradio里加的那个滑块。如果希望尽量少漏,把阈值从0.5降到0.3,那么垃圾短信概率超过0.3就会被拦截。代价是正常短信被误判的比例上升。最常见的做法是画出ROC曲线,找到约登指数最大的点作为默认阈值,然后在界面上保留手调入口。这个项目的数据量不大,用sklearn的roc_curve只要几行代码就能算出来。

5.2 加载模型时报错的三个高频原因

这个项目拿到手最容易出错的位置集中在tokenizer.jsonpytorch_model.bin的路径上。用from_pretrained("model")时,model文件夹里必须有config.jsonvocab.txttokenizer.json和权重文件,缺一个都会报OSError: Can't load config。另一个常见问题是BertTokenizerBertModel版本不匹配,比如用transformers4.x 的tokenizer加载老版本的vocab.txt,有时需要强制指定do_lower_case=True。最后是Python版本导致的.pyc文件问题,项目里同时有model.cpython-38.pycmodel.cpython-312.pyc,如果你用Python 3.12跑,会优先读3.12的缓存,但如果缓存过期或损坏,直接删掉__pycache__让它重新生成即可。

5.3 提高单条推理速度的一个小技巧

洗不准的时候先看看attention_maskmax_length=128意味着每条样本都要padding 128个token,模型计算的attention也会处理128个位置,哪怕实际短信只有20个字。推理阶段可以把max_length改成动态截断到批内最长序列长度——因为Gradio默认单条请求,只需把真实长度传给tokenizer:

inputs = tokenizer(text, return_tensors="pt", truncation=True, max_length=128) model.eval() with torch.no_grad(): logits = model(**inputs).logits

这里动态计算长度后,20个字的短信不需要算128个位置,CPU上也能把延迟降到20ms左右。如果你后面用GPU批量过滤一批短信,建议固定max_length或者用pad_to_multiple_of=8,反而有更好的SIMD加速效果。这个技巧在答辩时主动提出来,比说“我用了最新框架”更能体现对模型原理的理解。

本文还有配套的精品资源,点击获取

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询