ARTICLE · INTELLIGENCE

战地情报 · 详情页

来自尧图项目组的一线实战观察与深度解析

文本匹配算法源码解读:单塔与双塔模型选型及工程落地

文本匹配算法源码解读:单塔与双塔模型选型及工程落地 简介面向文本匹配与相似度计算任务基于Python的源码工程实现了PointWise单塔、DSSM双塔和Sentence BERT双塔三种主流算法并配套训练/推理脚本、数据集与详细使用说明适合NLP初学者、算法研究者以及需要完成毕设、课设或课程作业的高校学生。运行环境基于PyTorch和Transformers搭建安装依赖后即可直接使用。包体共34个文件压缩包7.86MB以18个Python脚本为核心覆盖模型定义、训练、评估与推理流程另有4个Shell一键运行脚本、2个TSV评测数据、4个TXT说明、2个Markdown文档和4张示意图片目录结构清晰。目前已有270人学习使用项目代码均测试通过答辩评审平均分达96分可放心下载。借助说明文档和脚本可快速复现三种匹配模型效果也可在此基础上修改扩展用于毕设、课设或项目初期演示。1. 文本匹配算法源码的三种模型从单塔到双塔该怎么选做文本匹配任务时最头疼的不是找不到模型而是同一份数据要试好几条技术路线。这套基于 Python 实现的文本匹配算法源码把 PointWise 单塔、DSSM 双塔、Sentence BERT 双塔三种方案一次性铺开还带了数据集、训练脚本、推理脚本和使用说明。简单说单塔管精排、双塔管召回这份资源把这两类核心场景都覆盖了。适合正在做毕设、课程设计的学生也适合刚接手搜索、客服意图识别、文本去重这类任务的工程师拿来做 baseline。源码里 train_pointwise.py、train_dssm.py、train_sentence_transformer.py 三个入口一一对应配套 .sh 启动脚本和 readme解压后能直接上手跑。2. 三种模型原理拆开看交互深度、向量化成本与选型判断文本匹配表面上比的是“两个句子像不像”工程上真正卡人的却是交互和表示的取舍。单塔模型让两个句子在编码器内部互相看见精度高但算力开销大双塔模型把两个句子各自编码成向量再用余弦相似度打分速度快但交互弱。这份资源把两条路线都给了实现所以拿到手第一件事不是急着跑训练而是先读 readme 和 model.py搞清楚每个脚本对应的模型结构再决定先跑哪一条。2.1 单塔 PointWise把两句话拼成一个序列用交互换精度单塔模型的输入很直白把 text_a 和 text_b 用 [SEP] 拼成一个序列前面加 [CLS]整个送进 BERT。这样做的好处是BERT 每一层 attention 都能同时看到两个句子的 token深层网络可以学到“苹果”和“iPhone”、“怎么查话费”和“话费查询方法”这种隐式对齐关系。模型结构通常长这样本资源的 model.py 也是这个思路class PointWiseModel(nn.Module): 单塔模型text_a 和 text_b 拼成一个序列送进编码器 def __init__(self, encoder, hidden_size768, num_labels2): super().__init__() self.encoder encoder self.classifier nn.Linear(hidden_size, num_labels) def forward(self, input_ids, attention_mask, token_type_ids): # token_type_ids 用来区分句子 A 和句子 B这是单塔的关键 outputs self.encoder( input_idsinput_ids, attention_maskattention_mask, token_type_idstoken_type_ids, return_dictTrue, ) # 取 [CLS] 位置的向量作为整句表示再接分类层 cls_vec outputs.last_hidden_state[:, 0] logits self.classifier(cls_vec) return logits这段代码里encoder 就是预训练 BERTclassifier 是一个线性层把 768 维的 [CLS] 向量映射成二分类 logits。token_type_ids 是单塔模型独有的输入它让 BERT 知道哪些 token 属于句子 A、哪些属于句子 B拼接时不传这个参数模型就分不清两句话的边界。训练时对应 train_pointwise.py推理时用 inference_pointwise.py输入一对文本输出一个相似度分数。PointWise 这个名字指“按点打分”一条样本就是一对文本加一个标签模型直接预测这对文本的匹配概率不构造正负样本 pair 做两两比较。它适合候选集很小但精度要求很高的场景比如客服对话里先召回 20 条相似问题再用单塔模型把这 20 条挨个精排。反过来的问题是如果候选集有十万条每一条都要和 query 拼接过一遍 BERT这个计算量在线基本扛不住。2.2 双塔 DSSM两句话各过各的把匹配变成向量运算DSSM 是另一条路线query 和 doc 分别过同一个编码器各自得到一个句向量最后算余弦相似度。整个过程中两个句子没有直接交互模型学到的其实是“把语义相近的句子映射到向量空间里相近的位置”。这种结构的巧妙之处在于doc 的向量可以离线全量算好存起来线上来一个 query 只需要算一次 query 向量然后去向量库里做检索。class DSSMModel(nn.Module): 双塔模型query 和 doc 各自过编码器最后算相似度 def __init__(self, encoder): super().__init__() # 两个塔共享同一份预训练权重效果比不共享更稳 self.encoder encoder def encode(self, input_ids, attention_mask): outputs self.encoder( input_idsinput_ids, attention_maskattention_mask, return_dictTrue, ) # 这里用 mean pooling 而不是 [CLS]句向量更平滑 vec outputs.last_hidden_state.mean(dim1) return vec def forward(self, query_input_ids, query_mask, doc_input_ids, doc_mask): query_vec self.encode(query_input_ids, query_mask) doc_vec self.encode(doc_input_ids, doc_mask) sim torch.cosine_similarity(query_vec, doc_vec, dim-1) return sim这段代码里encoder 是共享的也就是说 query 塔和 doc 塔用的是同一套权重。共享权重的好处是两个塔学习到的向量空间是一致的不会出现 query 向量和 doc 向量各在一个坐标系里、余弦相似度没意义的情况。encode 方法里我习惯用 mean pooling 取整句均值而不是直接拿 [CLS]原因是 BERT 的 [CLS] 向量在句向量任务上有各向异性的问题均值池化在双塔结构里更稳。实际训练对应 train_dssm.py脚本里加载数据和计算 loss 的细节都封装好了。DSSM 的短板也明显两个句子没有交互模型捕捉不到“字面不同但需要对齐”的细粒度关系。比如“明天会下雨吗”和“带伞的建议”单塔能通过 attention 建立起两边的联系双塔只能靠向量距离硬猜。所以双塔模型适合做召回阶段的粗筛不适合做最终打分的精排。2.3 Sentence BERT给双塔换一个更明确的句向量训练目标Sentence BERT 看着和 DSSM 一样是双塔但训练目标完全不同。DSSM 通常把二分类或相似度分数直接当监督信号Sentence BERT 用 siamese 结构专门优化句向量本身比如让相似句对的余弦距离更近、不相似句对更远。这解决了裸 BERT 句向量表现差的问题训练出来的向量可以直接作为语义 embedding 入库后续做聚类、检索、相似度计算都很顺手。资源里的 train_sentence_transformer.py 对应这条路线推理入口是 inference_sentence_transformer.py。还有一点值得注意资源根目录里有 unsupervised 子目录里面是 simcse 的实现。SIMCSE 是“无监督对比学习”的思路同一个句子过两次编码器因为 dropout 的随机性得到两个略有差异的向量把这两个向量当正例来拉近。这意味着在没有标注数据时也能训出语义向量。我一般会这样用有标注数据先跑 supervised 下的三个模型数据没标注就直接看 unsupervised/simcse两条路正好互补。2.4 选型判断这四种情况怎么选三种模型各有各的适用范围选错了再调参也是白费功夫。下面这个表格可以直接拿来当参考评估维度PointWise 单塔DSSM 双塔Sentence BERT 双塔交互能力强attention 能看到两侧弱两塔无交互弱两塔无交互在线成本每对都要重新算向量预计算线上只算 query向量预计算线上只算 query典型场景精排、小候选集打分召回、粗排、向量检索语义向量、聚类、相似度排序数据需求成对标注数据成对标注或构造负样本无标注时可用 simcse判断逻辑可以捋成四条。第一候选集超过一万条优先考虑双塔单塔的算力顶不住第二业务要的是“最终排序结果”而不是“先筛出几百条”可以双塔召回后再接一个单塔精排两层结构最常用第三手头有成对标注数据直接跑 supervised 下的三个训练脚本第四完全没有标注数据去 unsupervised 目录跑 simcse 拿到句向量这也是这份资源比一般文本匹配 demo 值钱的地方。3. 从解压到第一条 loss把训练管线完整跑一遍原理看明白了接下来就是把代码跑通。这一步最磨人的不是模型本身而是环境、数据格式、脚本路径这些细节。文本里写着pip install -r ../../requirements.txt很多人直接复制粘贴然后报“找不到文件”其实是因为依赖文件在项目上层目录相对路径取决于你当前所在位置。3.1 环境安装requirements.txt 的相对路径别踩先确认目录结构再动手。假设你把 zip 解压后得到了 text_match-master 目录脚本都在这个目录下而 requirements.txt 在它的上一级或更上层所以安装依赖之前先ls ../..看一眼文件到底在哪。常见做法是直接回到文件所在目录安装# 方法一回到项目父目录requirements.txt 就在那里 cd text_match-master/.. pip install -r requirements.txt # 方法二把依赖文件复制到当前目录再装省得反复切换 cp ../requirements.txt . pip install -r requirements.txt这段命令的逻辑很简单先让 pip 找到 requirements.txt再让它按文件内容安装依赖。方法一更干净因为训练脚本在 text_match-master 里跑依赖装在同一个环境里不会被路径干扰。方法二适合你不想破坏原目录结构的情况。项目基于 pytorch 和 transformersPython 建议 3.8 以上用 vscode 或 pycharm 都行关键是解释器要指向你实际安装依赖的那个虚拟环境不然跑起来会报 ModuleNotFoundError。额外提醒一句torch 分为 CPU 版和 CUDA 版如果你机器有显卡别用 CPU 版硬跑BERT 在 CPU 上训练慢到怀疑人生。装完依赖后可以跑一句python -c import torch; print(torch.cuda.is_available())输出 True 再继续。3.2 数据集先用命令确认 data 目录的真实格式data 目录里放的是训练和验证样本但具体是 txt、csv 还是 jsonl要以 readme.md 和 supervised/utils.py 里的读取逻辑为准。我拿到任何数据集第一件事不是直接开训而是先用命令看前几行ls -la data/ head -n 5 data/train.txt 2/dev/null || head -n 5 data/train.jsonl 2/dev/null第一行命令列出 data 目录下的文件第二行尝试读取文本格式的训练数据。如果 train.txt 存在就打印前 5 行不存在就尝试 train.jsonl避免你连文件名都没对上就硬跑。文本匹配数据最常见的格式是三列label、text_a、text_b行间用制表符或逗号分隔jsonl 格式则是每行一个 JSON 对象字段名通常是 label、text_a、text_b。具体看 utils.py 怎么解析这一步千万别跳。顺带说一句数据质量直接决定训练上限。我看过不少人把 label 和文本列弄反训练时 loss 一直不降还以为是模型问题。建议打开 readme.md 看数据集说明再用wc -l data/train.txt确认样本量太少了就先别上大模型。3.3 训练脚本和 .sh 启动文件参数逐个拆三个训练脚本对应三个 .sh 启动文件打开 train_dssm.sh 通常会看到类似下面的结构# 典型的训练启动脚本内容具体参数以解压出的文件为准 python train_dssm.py \ --data_dir data \ --output_dir output/dssm \ --model_name_or_path bert-base-chinese \ --batch_size 32 \ --learning_rate 2e-5 \ --num_epochs 3 \ --max_len 128这个脚本把训练参数全部放在命令行里改配置不用动 Python 代码。每个参数的含义如下data_dir 指定数据集目录output_dir 是模型保存位置model_name_or_path 指定预训练模型batch_size 是每步喂给 GPU 的样本数learning_rate 是 BERT 微调常用的 2e-5num_epochs 是训练轮数max_len 是序列最大长度。第一次跑建议把 num_epochs 改成 1确认全流程通了再完整训练。训练过程中的日志由 iTrainingLogger.py 负责它会记录每个 step 的 loss、学习率和耗时。我一般习惯把训练输出同时写到文件里方便事后回溯bash train_dssm.sh 21 | tee train.logtee 命令把终端输出同时打印到屏幕和 train.log训练结束后直接翻日志就能看到每一步的 loss 变化。这套脚本里还有一个 get_embedding.py是用来把语料库批量转成向量的工具后面讲到落地时会用到。3.4 跑通一次完整训练怎么确认它真的在学确认数据和脚本没问题后直接执行对应模型的启动脚本。这里以双塔 DSSM 为例bash train_dssm.sh跑起来之后重点看三点。第一终端或 train.log 里的 loss 是否在逐步下降BERT 微调任务通常前几百步会从零点几缓步往下走如果 loss 完全不动或者乱跳优先怀疑数据问题而不是模型第二output_dir 目录下是否按 epoch 保存了 checkpoint断点续传和后续推理都要靠这些文件第三GPU 显存占用是否正常如果报 CUDA out of memory回到 3.3 把 batch_size 调小。第一次做文本匹配实验别一上来就追求 SOTA先把 pipeline 跑通、看一次 loss 曲线、做一次推理这个最小闭环的价值比盲目调参大得多。跑完一个 epoch 后赶紧用 inference_dssm.py 拿两条句子试试相似度直观感受一下模型学成什么样。4. 避坑指南文本匹配源码最常见的五个翻车点模型代码本身不复杂出问题的地方几乎都在环境、数据和参数上。下面这五个坑是我跑这类项目时反复遇到的每一条都是现象、原因、解决三步说透。4.1 torch 报错 “Torch not compiled with CUDA enabled”现象训练脚本一跑就报错提示当前 torch 不支持 CUDA或者 GPU 显存永远显示 0。原因pip 默认装的是 CPU 版 torch或者 CUDA 版本和显卡驱动不匹配。BERT 类模型在 CPU 上虽然能跑但训练速度只有 GPU 的几十分之一基本等于跑不起来。解决先卸载重装对应 CUDA 版本的 torch以 CUDA 11.8 为例pip install torch --index-url https://download.pytorch.org/whl/cu118装完再用python -c import torch; print(torch.cuda.is_available())验证。如果下载慢把 pip 源切到国内镜像再装但注意 torch 的 CUDA 版本必须和你本机驱动匹配不是版本越高越好。4.2 bert-base-chinese 下载卡住或反复超时现象训练脚本启动后终端一直停在“Downloading bert-base-chinese”这一步进度条不动或者下载到一半报连接错误。原因transformers 默认从 Hugging Face 仓库下载预训练模型直连不稳定时很容易卡住。这个和训练代码本身没关系但能卡掉一大半新手。解决把模型下载地址切到国内镜像在训练脚本所在的终端里先执行一条命令export HF_ENDPOINThttps://hf-mirror.com bash train_dssm.sh还可以更彻底一点手动把 bert-base-chinese 下载到本地目录然后把 .sh 里的 model_name_or_path 改成绝对路径比如--model_name_or_path /data/pretrain/bert-base-chinese。这样训练完全离线不受网络影响之后换预训练模型也只是换个路径的事。4.3 显存不够CUDA out of memory现象训练跑到第几步突然报 CUDA out of memory进程直接退出前面算的都白算了。原因batch_size 太大或者 max_len 太长BERT 的中间激活值很吃显存。batch_size 设 32 甚至 64小显卡根本扛不住。解决先把 batch_size 降到 8 或 4把 max_len 从 128 降到 64确认能跑通后再往上加。修改方式直接改 .sh 里的命令行参数python train_pointwise.py \ --data_dir data \ --output_dir output/pointwise \ --batch_size 8 \ --max_len 64 \ --num_epochs 1如果显存还是不够用梯度累积来模拟大 batch代价是训练时间变长。我的血泪经验是宁可 batch_size 小一点也别让训练中途崩掉崩一次浪费的时间够调十次参数。4.4 loss 不降或准确率一直等于随机水平现象训练日志里 loss 基本不动或者验证准确率一直在 50% 上下跟抛硬币差不多。原因八成是数据读取出了问题。常见的有三类label 列和文本列解析反了模型拿文本当标签学标签分布严重不均衡且没做处理tokenizer 把文本都截断成了空序列。这些都属于数据管线的黑匣子问题光看训练日志很难定位。解决训练代码里临时加几行打印直接看模型吃进去的是什么for step, batch in enumerate(train_loader): if step 3: break # 打印 input_ids 的形状和标签分布确认数据读取没毛病 print(batch[input_ids].shape, batch[labels].unique(return_countsTrue))这段代码的作用是把前几个 batch 的形状和标签分布打出来。shape 应该是 [batch_size, max_len]labels 的 unique 结果应该同时包含 0 和 1。如果 labels 全是同一个值或者 shape 明显不对回去查 data 目录和 utils.py 的解析逻辑。另外BERT 微调学习率一定要用 2e-5 这个量级有人手滑改成 1e-3loss 直接飞了。4.5 双塔模型效果比单塔差很多别急着换模型现象DSSM 或 Sentence BERT 训完后在测试集上的准确率明显低于 PointWise感觉像模型坏了。原因这不是 bug双塔模型没有交互绝对精度本来就拼不过单塔。双塔的价值在召回效率和向量检索硬拿准确率去比单塔等于拿招回粗排去比精排比错对象了。解决把评测指标换成召回率或 Top-K 命中率。比如每个 query 有 100 个候选 doc其中 1 个是正确答案看双塔模型能不能把正确答案排进前 10。这个指标才能反映双塔的真实能力。另外双塔训练时检查两塔是否共享预训练权重、输出向量有没有做归一化这两个细节直接影响向量空间的质量。我一般会保证向量做了 L2 归一化再算余弦相似度否则长度信息会干扰排序。5. 从相似度分数到向量检索服务推理脚本与落地姿势模型训完只是第一步真正要让文本匹配能力用起来得看推理脚本怎么接。这份资源给了三个推理脚本还有 get_embedding.py正好覆盖了“打分”和“检索”两种使用方式。5.1 三个推理脚本的输入输出inference_pointwise.py、inference_dssm.py、inference_sentence_transformer.py 分别对应三个模型的推理。调用方式类似以双塔为例python inference_dssm.py \ --model_path output/dssm/best_model \ --text_a 如何查询话费 \ --text_b 话费查询方法这个命令会加载训练好的 checkpoint对两句话进行编码并计算余弦相似度最后打印一个 0 到 1 之间的分数。分数越高代表语义越接近。要注意 model_path 要指向保存模型权重和配置的目录而不是你随手建的输出根目录。三个脚本的差别在于模型结构不同单塔模型还要额外拼接 token_type_ids双塔模型只需要两个独立的输入所以推理脚本不完全通用。单塔推理有个明显的边界它是逐对打分的适合精排阶段对上一步召回的结果做细排。如果候选集有几万条每一条都要和 query 拼起来过一遍 BERT在线延迟和算力都吃不消。这种场景应该用双塔。5.2 get_embedding.py把语料库批量向量化双塔模型真正的打开方式是提前把候选语料全部转成向量存成向量文件线上只对 query 做一次编码。get_embedding.py 干的就是这件事python get_embedding.py \ --model_path output/dssm/best_model \ --input_file data/corpus.txt \ --output_file data/corpus_vec.npyinput_file 是候选语料每行一句话output_file 是输出的向量文件numpy 的 npy 格式形状大概是 [句子数, 768]BERT base 的 hidden_size 是 768。这个文件可以直接被后续检索程序读取。执行结束后可以顺手用 numpy 检查一下输出import numpy as np vecs np.load(data/corpus_vec.npy) print(vecs.shape) # 期望输出 (N, 768)N 等于语料条数如果形状对不上优先检查 corpus.txt 的每行是否都被正确编码空行会拉低 N 的实际质量。把语料库向量化存库之后query 侧的编码也用同一个模型保证两边的向量空间一致。5.3 向量检索验证40 行代码算出 Top-K 召回向量文件拿到手后最简单的验证就是给一条 query 算它和所有候选向量的相似度取前 K 个看结果。先用 numpy 跑通逻辑数据量大了再上 faissimport numpy as np # vecs 是候选向量矩阵query_vec 是单条 query 的向量 vecs np.load(data/corpus_vec.npy) query_vec np.load(data/query_vec.npy)[0] # 假设 query 向量已单独保存 # 先做 L2 归一化再算内积等价于余弦相似度 norm_vecs vecs / (np.linalg.norm(vecs, axis1, keepdimsTrue) 1e-9) norm_query query_vec / (np.linalg.norm(query_vec) 1e-9) sims norm_vecs norm_query # 取相似度最高的前 10 条 top_k_idx sims.argsort()[-10:][::-1] print(top_k_idx)这段代码的核心是归一化后做矩阵乘法一次算完 query 和全量候选的相似度。argsort 拿到的是按相似度排序后的索引下标[::-1] 是为了从高到低排列。numpy 在十万条以内的向量做暴力检索还能接受超过这个量级就要上 faiss 的 IVF 或 HNSW 索引了。实际生产环境里双塔召回加单塔精排是最常见的一套组合双塔从百万候选里筛出几百条单塔再对这几百条逐对精排既保性能又保精度。5.4 训练和推理保持一致tokenizer 别在两头各干各的推理阶段最容易忽略的是预处理一致性。训练时 max_len 设了 128推理脚本里如果没设或者设成了 512两边向量分布直接漂移线上效果掉一截。我一般会固定住三个东西max_len、tokenizer 的版本、add_special_tokens 开关。前两个都好理解第三个是指 BERT 输入要不要自动加 [CLS] 和 [SEP]。训练时加了这个开关推理时不加模型看到的序列结构完全变了分数当然不可信。检查起来很简单推理脚本里加载 tokenizer 之后打印一条预处理结果和训练时的样本对比一下即可。6. 再进一步换预训练模型与 Top-K 命中率验证跑通默认配置只是起点真正让模型贴合业务还得换预训练模型。把 .sh 里的 model_name_or_path 从 bert-base-chinese 换掉就行比如换成中文 RoBERTapython train_sentence_transformer.py \ --data_dir data \ --output_dir output/sbert_roberta \ --model_name_or_path hfl/chinese-roberta-wwm-ext \ --batch_size 16 \ --learning_rate 2e-5 \ --num_epochs 3 \ --max_len 128hfl/chinese-roberta-wwm-ext 是中文 NLP 里常用的预训练模型对中文语义的理解通常优于基础版 BERT。换模型唯一要注意的是显存和下载时间模型更大、加载更慢训练前先确认 output_dir 配置正确。如果公司内部有离线模型文件把 model_name_or_path 改成绝对路径还能绕开在线下载。换完模型后别光看 loss要看它到底能不能把正确答案排到前面。我一般会写一段很小的 Top-K 命中率验证脚本拿测试集的 query 去全量候选里检索统计正确答案出现在前 K 条里的比例import numpy as np # 假设已有全量候选向量 vecs 和测试集 vecs np.load(data/corpus_vec.npy) hits 0 total 0 for query_vec, gold_idx in test_set: norm_vecs vecs / (np.linalg.norm(vecs, axis1, keepdimsTrue) 1e-9) norm_query query_vec / (np.linalg.norm(query_vec) 1e-9) sims norm_vecs norm_query top_k sims.argsort()[-10:][::-1] if gold_idx in top_k: hits 1 total 1 print(fRecall10: {hits / total:.4f})这段代码统计的是 Recall10正好对应双塔模型该看的指标每一条 query 的正确答案被召回到前 10 就记一次命中最后算命中比例。单塔模型比的是准确率双塔模型比的就是这个。如果 Recall10 明显偏低先确认训练时有没有做困难负样本采样再检查归一化逻辑这两个是双塔召回效果最敏感的环节。以前我拿到新数据集就直接开训结果一半时间都在排数据格式和路径的错。现在每换一个数据源都强制先跑 20 条样本肉眼检查、再跑 50 步看 loss 下降趋势确认无误后才挂上完整训练。这个习惯帮我省掉的返工时间远超那几分钟检查成本。希望帮到你。本文还有配套的精品资源点击获取
RELATED READING

延伸阅读

更多一线实战笔记与深度复盘,助您持续精进