十年匠心定制 · 商业建站与技术教学双线并行 咨询热线:400-886-1026 service@lmnt.cn
ARTICLE DETAIL

资讯详情

深耕网站建设与运营推广的一线实战洞察。

用BERT微调实现垃圾短信过滤:从数据预处理到Gradio部署

用BERT微调实现垃圾短信过滤:从数据预处理到Gradio部署 简介这是一份面向人工智能、深度学习的毕业设计项目资源聚焦垃圾短信过滤场景基于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中读取统一字段为label和text。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, engineopenpyxl) df df[[label, text]].copy() df[text] df[text].astype(str) df.to_csv(message.csv, indexFalse, encodingutf-8-sig) print(df[label].value_counts())这段代码把Excel中的原始数据转成csvencodingutf-8-sig是为了让Excel直接打开不乱码。处理完一定要打印类别分布如果垃圾短信占比远低于正常短信后面要做类别加权或过采样。2.2 分词、编码与截断data_process.py 的完整流程data_process.py的主要任务是把中文短信转换成BERT的输入格式。BERT不能直接吃字符串需要借助bert-base-chinese配套的vocab.txt和tokenizer.json做分词再转成input_ids、attention_mask和token_type_ids。中文BERT默认按字切分所以不需要引入jieba。我一般会限制max_length128因为垃圾短信一般不会超过这个长度过长反而会把一些广告的尾部干扰带进来。代码结构大致如下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_lengthmax_len, paddingmax_length, truncationTrue, return_tensorspt ) 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)注意paddingmax_length会把不足128的token全部补成[PAD]这会让推理时计算量增大。如果追求效率可以改为paddinglongest等训练完再统一处理。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_size0.2, stratifydf[label], random_state42 )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_labels2, dropout0.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_idsinput_ids, attention_maskattention_mask) pooled outputs.pooler_output return self.classifier(self.dropout(pooled))pooler_output是BERT对[CLS]这个token做线性变换和tanh激活后的结果可以理解为整个句子的语义向量。这里的dropout0.3是防过拟合用的对短文本分类来说0.2到0.4之间都是合理选择。如果你的训练数据很少不加载bert-base-chinese的预训练权重而是随机初始化模型效果会差很多这也是这个项目必须带pytorch_model.bin和config.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_labels2) optimizer AdamW(model.parameters(), lr2e-5, correct_biasFalse) total_steps len(train_loader) * epochs scheduler get_linear_schedule_with_warmup( optimizer, num_warmup_stepsint(total_steps * 0.1), num_training_stepstotal_steps )correct_biasFalse是transformers库对AdamW的要求它让优化器不把bias项也做weight decay。warmup约占总step的10%先让学习率从小到大爬升再线性衰减到0这是BERT系列模型比较通用的设置。推荐的超参参考下表参数数值调整建议max_length128短信长度超128可增加到256但训练时间变长batch_size16显存不够就降到8学习率也要同步降learning_rate2e-5类别不平衡时尝试1e-5epochs3数据集很小用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()就够了。注意weights和model这两个目录最好同时保留前者是自定义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.Interface或gr.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, threshold0.5): inputs tokenizer(text, return_tensorspt, truncationTrue, max_length128) with torch.no_grad(): logits model(**inputs).logits prob torch.softmax(logits, dim1)[0][1].item() label 垃圾短信 if prob threshold else 正常短信 return label, f垃圾概率: {prob:.4f} demo gr.Interface( fnpredict, inputs[gr.Textbox(lines4, label输入短信内容), gr.Slider(0, 1, value0.5, label阈值)], outputs[gr.Label(label判定结果), gr.Textbox(label概率)], title垃圾短信过滤系统 ) demo.launch(server_name0.0.0.0, server_port7860)这个代码里的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 的Interface和Blocks都支持auth参数接收一个字典或函数。只需要三个账号的场景用字典最直接demo.launch(auth(admin, your_password))如果想做多用户校验传一个函数函数接收用户名和密码做数据库查询返回True才能访问。这个功能在实际部署到公网时非常有用避免任何人都能调用你的模型白白烧显卡。至于部署方式常见做法是直接在服务器上python gradio_web.py也可以用nohup放到后台或者用systemd管理进程。注意Gradio默认只监听本机127.0.0.1想远程访问必须显式设置server_name0.0.0.0否则会出现本地能打开、手机打不开的问题这是初学者最常踩的坑。5. 验证效果与避坑让BERT模型在真实短信上可用5.1 用混淆矩阵和阈值找到最佳分类点训练完模型不要只看准确率。在垃圾短信场景里正常短信占总量的90%以上把全部判成正常准确率也高但漏掉的诈骗短信带来的风险远大于误杀一条广告。我惯用的办法是算一遍混淆矩阵看带权F1或recall。比如验证集有1000条正常、100条垃圾模型把80条垃圾识别出来、20条漏掉那垃圾短信召回率是0.8准确率是0.8F1也是0.8这样评价比单纯准确率有价值得多。阈值调整可以直接利用第4章在Gradio里加的那个滑块。如果希望尽量少漏把阈值从0.5降到0.3那么垃圾短信概率超过0.3就会被拦截。代价是正常短信被误判的比例上升。最常见的做法是画出ROC曲线找到约登指数最大的点作为默认阈值然后在界面上保留手调入口。这个项目的数据量不大用sklearn的roc_curve只要几行代码就能算出来。5.2 加载模型时报错的三个高频原因这个项目拿到手最容易出错的位置集中在tokenizer.json和pytorch_model.bin的路径上。用from_pretrained(model)时model文件夹里必须有config.json、vocab.txt、tokenizer.json和权重文件缺一个都会报OSError: Cant load config。另一个常见问题是BertTokenizer和BertModel版本不匹配比如用transformers4.x 的tokenizer加载老版本的vocab.txt有时需要强制指定do_lower_caseTrue。最后是Python版本导致的.pyc文件问题项目里同时有model.cpython-38.pyc和model.cpython-312.pyc如果你用Python 3.12跑会优先读3.12的缓存但如果缓存过期或损坏直接删掉__pycache__让它重新生成即可。5.3 提高单条推理速度的一个小技巧洗不准的时候先看看attention_maskmax_length128意味着每条样本都要padding 128个token模型计算的attention也会处理128个位置哪怕实际短信只有20个字。推理阶段可以把max_length改成动态截断到批内最长序列长度——因为Gradio默认单条请求只需把真实长度传给tokenizerinputs tokenizer(text, return_tensorspt, truncationTrue, max_length128) model.eval() with torch.no_grad(): logits model(**inputs).logits这里动态计算长度后20个字的短信不需要算128个位置CPU上也能把延迟降到20ms左右。如果你后面用GPU批量过滤一批短信建议固定max_length或者用pad_to_multiple_of8反而有更好的SIMD加速效果。这个技巧在答辩时主动提出来比说“我用了最新框架”更能体现对模型原理的理解。本文还有配套的精品资源点击获取
返回列表