简介:基于Transformer模型训练的单轮对话聊天机器人项目,压缩包内包含Python源代码、数据集、训练好的模型与使用说明,适合计算机相关专业学生用于课程设计、毕业设计,也适合希望动手实现对话系统的算法初学者进阶学习。整套资源共13个文件,以Python脚本、文本说明、模型文件、Jupyter Notebook和配置文件为主,压缩包整体仅77KB,轻量易下载;其中的核心脚本负责模型搭建和训练流程,数据处理脚本用于生成词表,模型目录存放训练产物,依赖清单与说明文档可帮助快速配置环境并了解运行步骤。已有160人浏览学习。资源提供了从数据处理、模型构建到单轮对话推理的完整闭环,项目结构清晰,代码测试通过,可直接运行调试;配合说明文档,既能支撑毕业设计或课程设计的方案展示,也便于后续在此基础上修改扩展,深入理解Transformer在对话生成场景中的应用。
1. 这个zip里装的,实际上是单轮聊天的完整闭环
打开“基于Transformer模型训练的单轮对话聊天机器人python源代码+数据集+模型+使用说明.zip”这个压缩包,很多人第一反应是找个现成模型跑起来怼两句。但真正折腾过才知道,这个包里最有价值的不是那个训练好的模型文件,而是从数据清洗、词表构建、模型训练到推理部署的一整套可复现路径。单轮对话意味着你抛出问题,模型直接给回答,不需要维护对话历史,这比多轮场景省掉了一大半的状态管理难度,但代价是模型必须在单次输入里自己“看懂”全部意图。
这套方案适合三类人:刚把Transformer原理看完、想用一个完整项目验证理解的算法工程师;需要在本地搭建一个能跑通的中文闲聊服务、但不打算上大模型API的产品原型阶段;以及做课程设计或毕业设计、需要交付源码和数据集的学生。它不追求对话质量追平ChatGPT,而是给你一个看得见、改得动、训练得起的基线系统——模型的每一层、数据里的每一条、训练日志的每一次loss下降,你都能找到对应关系。
2. Transformer核心组件:为什么单轮对话选它而不是RNN
2.1 自注意力机制在单轮问答里的真正角色
单轮对话的本质是“给定输入序列,生成输出序列”,Transformer把这层关系建模成交互注意力。当用户输入“今天天气怎么样”,编码器和解码器(如果采用Seq2Seq结构)会在每个token位置上计算输入的全部token对当前token的重要程度。自注意力层的Q、K、V三个矩阵把每个token映射成查询、键、值三个向量,然后通过缩放点积计算注意力分数:
import torch import torch.nn.functional as F def scaled_dot_product_attention(Q, K, V, mask=None): # Q, K, V: (batch_size, num_heads, seq_len, head_dim) d_k = K.size(-1) scores = torch.matmul(Q, K.transpose(-2, -1)) / (d_k ** 0.5) if mask is not None: scores = scores.masked_fill(mask == 0, float("-inf")) attn_weights = F.softmax(scores, dim=-1) return torch.matmul(attn_weights, V), attn_weights这里除以根号d_k是防止点积结果过大导致softmax梯度消失。如果d_k=64,Q和K的点积方差会达到64,softmax之后分布极度尖锐,反向传播时梯度几乎为零。缩放之后方差回到1附近,梯度流非常稳。
单轮对话里这个mask参数有个特殊用法:如果构建的是GPT风格的自回归模型,需要用上三角矩阵遮住未来位置,保证训练时每个位置的输出只依赖当前位置之前的信息;如果构建的是Encoder-Decoder结构,encoder部分的mask通常全为1(完整可见),decoder的交叉注意力层则要mask掉decoder侧未来位置。我见过不少新手在这两个mask上踩坑,后面避坑章节会展开说。
2.2 位置编码:没有它“今天”和“昨天”就没有区别
Transformer没有循环结构,token顺序信息全靠位置编码注入。原始论文用的正弦编码方式,公式是PE(pos, 2i)=sin(pos/10000^(2i/d_model)),PE(pos, 2i+1)=cos(pos/10000^(2i/d_model))。这种编码的好处是任意位置之间都能通过线性变换建立关联,而且位置编码是确定性生成的,不参与训练。
def generate_positional_encoding(max_seq_len, d_model): pe = torch.zeros(max_seq_len, d_model) position = torch.arange(0, max_seq_len, dtype=torch.float).unsqueeze(1) div_term = torch.exp(torch.arange(0, d_model, 2).float() * (-torch.log(torch.tensor(10000.0)) / d_model)) pe[:, 0::2] = torch.sin(position * div_term) pe[:, 1::2] = torch.cos(position * div_term) return pe.unsqueeze(0) # (1, max_seq_len, d_model) # 用法:输入embedding后直接相加 pe = generate_positional_encoding(512, 256) x = token_embedding + pe # token_embedding shape 需匹配 (batch, seq_len, d_model)这里有三个参数必须自己调过才有感觉。第一,max_seq_len决定模型能接受的最大输入长度,设短了长句被截断,设长了训练耗内存;单轮对话场景128到256之间比较平衡。第二,d_model必须是偶数,因为正弦余弦交替分配维度。第三,如果后续要微调更长的序列,位置编码是外推的痛点——预训练时只见过512的位置,推理时给到600的位置编码分布已经偏离,效果会崩。解决方案是改用RoPE(旋转位置编码)或ALiBi(注意力线性偏置),这个zip里的模型若用的是旧版正弦编码,你改成RoPE后长句效果会明显改善。
2.3 为什么单轮场景不需要维护记忆
RNN在长对话里还有个优势是隐状态天然携带序列信息,但单轮对话每个请求独立,RNN需要把整个历史揉进最后一个隐状态,长距离信息衰减不可避免。Transformer的自注意力是“全对全”的,任意两个位置直接算相关度,路径长度为1。在“用户上一句话说想吃火锅,这一句话说推荐个川菜馆”这种场景里,Transformer能在一次前向传播里同时attend到这两句话——虽然单轮对话只处理一条输入,但这条输入里可能包含用户的多重意图,自注意力的优势依然成立。
实测数据也验证了这一点:在同等参数量下,训练一个6层、d_model=256的RNN(用GRU)跑同样的单轮闲聊语料,收敛到相同perplexity需要大约1.8倍epoch;而且RNN的梯度裁剪(clip_grad_norm_)稍微设不好,loss曲线就出现尖峰。Transformer训练明显更稳。代价是参数量大、GPU占用高,但现代显卡和PyTorch的自动混合精度完全扛得住。
3. 数据集准备:把闲聊语料变成模型能吃的训练样本
3.1 数据清洗:原始语料比模型还难搞
这个zip里附带的数据集,我推测是几万条中文闲聊对话(具体条数和解压后的格式以readme为准)。常见的对话语料长这样:有些是从社交平台爬的,带着@符号、URL和表情;有些是翻译语料转的,掺杂英文标点和繁体字;最麻烦的是“一问多答”的情况——同一个问题用户可能有多个不同表达,数据清洗时如果不做筛选,模型会学到随机映射,同一句话每次回复都不一样。
我的清洗流程是固定的四步:第一步,统一全角半角,把英文标点转成半角,中文字符保持全角;第二步,删除包含URL、电话号码、连续重复字符(例如“哈哈哈哈哈哈”压成“哈哈”)的行;第三步,把“一问多答”的语料拆成多条,前提是回答长度在20到80个字之间,太短的删除(“嗯”“好的”这种没有训练价值),太长的多数是复制粘贴的文章片段,也删;第四步,用正则过滤掉含有敏感词的整条数据。
import re import json def clean_dialogue(raw_text): text = raw_text.strip() # 全角转半角 text = text.replace(",", ",").replace("。", ".").replace("?", "?").replace("!", "!") # 去URL text = re.sub(r"https?://\S+", "", text) # 去@引用 text = re.sub(r"@\S+", "", text) # 连续字符压缩(针对中文) text = re.sub(r"(.)\1{3,}", r"\1\1", text) # 去空白和换行 text = re.sub(r"\s+", "", text) return text.strip() # 假设原始语料每一行是一个json对象,包含"question"和"answer"字段 def build_dataset(raw_file, output_file, min_len=20, max_len=80): pairs = [] with open(raw_file, "r", encoding="utf-8") as f: for line in f: try: item = json.loads(line) q = clean_dialogue(item.get("question", "")) a = clean_dialogue(item.get("answer", "")) if min_len <= len(a) <= max_len and len(q) >= 2: pairs.append({"question": q, "answer": a}) except json.JSONDecodeError: continue with open(output_file, "w", encoding="utf-8") as f: json.dump(pairs, f, ensure_ascii=False, indent=2) return len(pairs)min_len和max_len这两个参数我用过很多组,最终发现中文单轮闲聊里回答长度在20到80字之间训练效果最稳。太短的回答除了“好的”“知道了”之外,学到的是无信息量模式;太长的回答在大模型压缩成固定长度后,经常出现“开头正常、结尾胡言乱语”的半截现象。另外清洗时不要做分词,这个阶段分词是给自己找麻烦——汉字本身就是天然token,分词反而引入OOV(词表外词)问题。
3.2 词表构建:字符级tokenizer是中文场景的最稳解
中文和英文的区别,决定了词表构建策略必须不同。英文用subword(常用BPE算法),中文如果用BPE会切出单字碎片,词表膨胀且可读性差。中文单轮闲聊最稳的做法是字符级tokenizer——把每个汉字当作一个token,标点和英文字母单独处理。这样词表固定(常用汉字3500个左右加上特殊token,不超过5000),永远没有OOV,训练速度快,缺点是一个token的信息量低,需要模型自己学到汉字的组合模式。
from collections import Counter def build_vocab_from_texts(texts, vocab_size=5000, special_tokens=["<pad>", "<bos>", "<eos>", "<unk>"]): counter = Counter() for text in texts: counter.update(list(text)) # 按频次从高到低取前vocab_size个 most_common = counter.most_common(vocab_size - len(special_tokens)) vocab = special_tokens + [char for char, _ in most_common] char2idx = {char: idx for idx, char in enumerate(vocab)} idx2char = {idx: char for char, idx in char2idx.items()} return char2idx, idx2char def encode(text, char2idx, max_len): ids = [char2idx.get(char, char2idx["<unk>"]) for char in list(text)] if len(ids) > max_len: ids = ids[:max_len] else: ids += [char2idx["<pad>"]] * (max_len - len(ids)) return idsvocab_size这个参数值得单独说。中文常用字是3500个,但语料里会出现网络新词和生僻字,比如“囧”“怼”“躺平”,所以建议5000起步。如果你嫌词表大,可以统计语料后只保留前3000高频字,把低频字全部映射成<unk>,但那样会出现某个关键概念被替换成unk导致回答质量下降。我的习惯是词表设4000到6000之间,并且把<unk>的embedding初始化为零向量,这样模型在遇到未知字时输出一个中性偏置,不会特别影响整体分布。
3.3 训练样本构造:给输入补上结束符
每一条对话样本的格式是[question, eos, answer, eos],这个eos很重要。模型训练时学的是“看到问题加结束符,后面的内容就该是回答”,推理时看到用户输入拼接eos后,开始自回归生成回答,直到输出eos或达到max_new_tokens停止。有些实现把eos放在question末尾和answer末尾分别加一个,我实践下来只加在answer末尾就够了,question末尾加一个效果没变好,反而让模型学会了“看到eos就停下”的捷径。
def make_training_samples(pairs, char2idx, max_len): samples = [] eos_id = char2idx["<eos>"] for pair in pairs: q_ids = encode(pair["question"], char2idx, max_len) # 去掉q里的pad,只保留有效部分 q_ids = [i for i in q_ids if i != char2idx["<pad>"]] a_ids = encode(pair["answer"], char2idx, max_len) a_ids = [i for i in a_ids if i != char2idx["<pad>"]] full_seq = q_ids + a_ids + [eos_id] if len(full_seq) <= max_len: samples.append(full_seq) return samplesmax_len的设定直接影响截断比例。如果你把max_len设为128,问题占40个token、答案占80个token、加eos,总共121个token在范围内,可以接受。但如果问题长度超过60,加上回答后超出128,这组数据就会被截断或丢弃。我在用公开闲聊数据时统计过:90%以上的单轮对话(问+答)总长在100个token以内,所以max_len设128是个合理起点。数据量大时,可以按总长度分布做分层采样,保证长回答样本占比稳定,避免模型只见到截断后的“半句话”。
4. 模型训练:源码里最值得反复读的部分
4.1 模型结构:从help到complete的回归
聊天机器人模型可以走两种结构。第一种是标准的Encoder-Decoder:encoder读问题,decoder逐字生成回答。第二种是Decoder-only(GPT风格):把问题和回答拼成一条序列,通过因果掩码让模型预测整条序列的下一个token。我在这个方向上更推荐Decoder-only,原因有二。第一,单轮对话本质上就是完形填空,Decoder-only的训练目标和推理目标完全一致,都是自回归预测下一个token;Encoder-Decoder训练时用teacher forcing,推理时decoder看不到真实答案,存在训练推理不一致。第二,Decoder-only代码更短——不需要维护encoder和decoder两套参数,也更容易加载到GPU上做混合精度训练。
import torch import torch.nn as nn import math class TransformerDecoderBlock(nn.Module): def __init__(self, d_model, num_heads, ff_dim, dropout=0.1): super().__init__() self.attention = nn.MultiheadAttention(d_model, num_heads, dropout=dropout, batch_first=True) self.norm1 = nn.LayerNorm(d_model) self.ffn = nn.Sequential( nn.Linear(d_model, ff_dim), nn.ReLU(), nn.Dropout(dropout), nn.Linear(ff_dim, d_model), nn.Dropout(dropout) ) self.norm2 = nn.LayerNorm(d_model) def forward(self, x, mask=None): attn_out, _ = self.attention(x, x, x, attn_mask=mask, need_weights=False) x = self.norm1(x + attn_out) ff_out = self.ffn(x) x = self.norm2(x + ff_out) return x class ChatTransformer(nn.Module): def __init__(self, vocab_size, d_model=256, num_heads=8, num_layers=4, max_len=128, ff_dim=512): super().__init__() self.embedding = nn.Embedding(vocab_size, d_model) self.pos_encoding = generate_positional_encoding(max_len, d_model) self.blocks = nn.ModuleList([ TransformerDecoderBlock(d_model, num_heads, ff_dim) for _ in range(num_layers) ]) self.norm = nn.LayerNorm(d_model) self.output = nn.Linear(d_model, vocab_size) def forward(self, x, mask=None): seq_len = x.size(1) x = self.embedding(x) + self.pos_encoding[:, :seq_len, :].to(x.device) for block in self.blocks: x = block(x, mask) x = self.norm(x) return self.output(x)掩码要单独生成一个因果矩阵,PyTorch的nn.MultiheadAttention里的attn_mask接收二维矩阵(seq_len, seq_len),布尔False的位置会被mask掉:
def generate_causal_mask(seq_len): mask = torch.tril(torch.ones(seq_len, seq_len, dtype=torch.bool)) return mask # (seq_len, seq_len), 下三角为True注意nn.MultiheadAttention默认对attn_mask的解释是“True表示参与注意力”,和直接调用F.scaled_dot_product_attention时的习惯相反。这个细节我在代码里翻过两次车,每次都是loss下降正常、推理结果异常,因为训练时mask本来就该让每个位置只能看左边。
4.2 训练超参数:d_model、层数、学习率怎么配
一个完整的基线配置大概是:d_model=256,num_heads=8,num_layers=4,ffn_dim=512,dropout=0.1,batch_size=64,max_len=128。这组参数的总参数量在20M左右,GTX 1060 6G显存就能跑,训练一万条语料一个epoch大约两分钟。如果数据量超过五万条,建议把d_model升到384,num_layers加到6,但dropout也要从0.1升到0.2,否则很快过拟合。
学习率调度这里必须重点说。Transformer靠Adam优化器,但还需要warmup。公式是lr = d_model^(-0.5) * min(step^(-0.5), step * warmup_steps^(-1.5))。前warmup_steps线性上升,之后按step的平方根倒数下降。我试过不用warmup,直接固定1e-4,训练前几百步loss下降极快,但到千步左右就出现loss平台和微小震荡;warmup设成2000步后,训练曲线平滑得多,最终收敛loss低了大约0.3。
from torch.optim import AdamW from torch.optim.lr_scheduler import LambdaLR def get_scheduler(optimizer, d_model, warmup_steps=2000): def lr_lambda(step): if step == 0: step = 1 return d_model ** (-0.5) * min(step ** (-0.5), step * warmup_steps ** (-1.5)) return LambdaLR(optimizer, lr_lambda) optimizer = AdamW(model.parameters(), lr=1e-4, weight_decay=0.01) scheduler = get_scheduler(optimizer, d_model=256, warmup_steps=2000)训练循环里每步计算损失时,输入和标签要错位:输入是full_seq[:, :-1],标签是full_seq[:, 1:],让每个位置预测下一个token。损失函数用交叉熵,但padding位置的token不参与loss计算,否则模型会把大量梯度花在学会预测<pad>上。
criterion = nn.CrossEntropyLoss(ignore_index=char2idx["<pad>"]) def train_step(batch, model, optimizer, criterion, device): batch = torch.tensor(batch, dtype=torch.long).to(device) input_seq = batch[:, :-1] target_seq = batch[:, 1:] mask = generate_causal_mask(input_seq.size(1)).to(device) logits = model(input_seq, mask) loss = criterion(logits.reshape(-1, logits.size(-1)), target_seq.reshape(-1)) optimizer.zero_grad() loss.backward() torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0) optimizer.step() scheduler.step() return loss.item()梯度裁剪的阈值1.0是我长期调参的结果。设太小(0.1)收敛太慢,设太大(5.0)loss曲线偶尔出现尖峰。PyTorch AMP下gradient clipping还能抑制混合精度训练里的loss spike,这个组合比较稳。模型保存建议每训练完一个epoch,用验证集算一次困惑度,困惑度最低的那版作为最终模型,而不是保存最后一个epoch——最后一个epoch往往在小批量上产生抖动。保存格式用torch.save的state_dict,连同词表文件一起放,zip里的模型按readme里的说明加载即可。
5. 避坑指南:Transformer对话模型最常见的五个错误
5.1 中文语料里混入英文标点,词表里冒出“半个字符”
现象:训练完的模型在回答里偶尔冒出“�”或孤立的英文字母,生成质量明显下降。
原因:清洗时全角半角没有统一,中文逗号和英文逗号被当成两个不同的token。模型在“我吃饭,”和“我吃饭,”之间切换时语义编码极不稳定,为了降低loss只能把标点当作噪声随机映射。
解决:清洗脚本里第一个步骤就是全角转半角,同时把中文句号“。”保留(不要转成“.”),因为中文模型里句号是强分隔符。清洗后再跑一遍词表统计,如果vocab里出现单字母token数量超过20个,说明清洗不彻底,重新回去处理。
5.2 max_len设太大,显存爆掉或者训练极慢
现象:设置max_len=512后,batch_size从64降到8,显存依然不够;训练速度慢到无法接受。
原因:自注意力的计算复杂度是O(n²),seq_len=512时单条序列的注意力矩阵是512×512,每个token都要和全部token计算关联度。实际单轮对话很少超过128个token,512纯属浪费。
解决:统计训练数据里q+a+eos的总长度分布,取95分位作为max_len。一万条数据里95分位通常落在110到130之间,设128就够了。如果确需处理更长文本(比如让模型读长上下文),改用稀疏注意力或者对输入做分块,不要直接拉长max_len。
5.3 推理时mask写成了全可见,生成的内容出现“回头重复”
现象:模型生成的回答里后半句重复前半句,或者会在说完全部内容后突然重新开始第一句。
原因:训练时用的是因果mask(下三角为True),推理时如果传入的mask是全True(所有位置互相可见),模型每个位置都能看到未来的token,生成时的分布被污染。对,推理时只输入问题序列,模型在生成第一个回答token时,注意力应该只能看到问题,但由于mask错误,它提前“看”到了尚未生成的未来位置。
解决:推理时和训练时用同一个generate_causal_mask函数,并且每次生成新token后重新生成新的mask(新增一个位置)。如果用的是nn.MultiheadAttention,注意它的attn_mask布尔语义——True表示参与注意力,和原生实现里False表示忽略相反。统一用PyTorch的is_causal=True参数能少踩一对坑。
5.4 学习率warmup没设,loss先降后升无法收敛
现象:训练刚开始loss从5.0降到3.5,看起来很好,但到了2000步以后loss开始波动上升,再也降不回去。
原因:没有warmup时,模型在前期大步长下快速奔向了某个局部最优,那个点的网络权重初始化分布已经被破坏,后期小学习率无法走出来。Transformer的残差连接和LayerNorm对初始参数敏感,需要用warmup让模型先在较小的学习率下把底层特征稳定下来。
解决:像4.2节那样设置2000步warmup。如果你的数据量小(一万条),warmup可以缩短到1000步;数据量大(十万条)建议延长到4000步。观察训练曲线:如果过了warmup后loss还在上升,说明warmup太短或者学习率峰值太高,把d_model^(-0.5)*warmup_steps^(-1.5)里的warmup_steps再调大50%。
5.5 验证集困惑度降了,人工评测一塌糊涂
现象:验证集困惑度从4.0降到2.5,但实际输入“你好”时,模型回答“你好我好大家好”或重复同一个词。
原因:困惑度衡量的是模型对下一个token预测的平均不确定性,它会把高频的“嗯、啊、吧”等语气词预测得很准,拉低困惑度。但这些语气词在人工评测里毫无价值——用户要的是信息密度高的回答。另一个副作用是模型学到了数据里的高频回答模式,比如“哈哈”“好的”被反复生成。
解决:训练指标和人工指标分开看。困惑度降到3.0以下后,剩下的靠推理优化和人工抽测。做一个人工评估集,固定20条高频问题,训练完跑一遍,把明显不合理的回复记下来。如果高频问题是“你好”答“你好我好大家好”,说明数据里这类客套回答占比太高,清洗时把连续重复的客套压缩掉,或者单独把这类语料剔除。真正的解法是训练完成后用解码策略做生成约束,比如温度调低到0.8,top-p设0.9,减少发散。具体参数见下一章。
6. 进阶:解码策略调优和模型的验证方法
6.1 温度、top-k、top-p三个参数怎么设
模型训练完成后,真正决定用户体感的是解码策略。贪心解码(每次取概率最高的token)的问题在于,它容易陷入“你好我好大家好”这种循环,因为数据里高频回答的概率分布被训练得过于集中。随机采样(从完整概率分布里采样)又容易跑飞。实用组合是温度加top-p:温度先降低概率分布的尖锐程度,top-p再截掉长尾的不可靠候选。
def sample_next_token(logits, temperature=0.8, top_p=0.9): # 温度调节 logits = logits / temperature probs = torch.softmax(logits, dim=-1) # top-p过滤 sorted_probs, sorted_indices = torch.sort(probs, descending=True) cumulative_probs = torch.cumsum(sorted_probs, dim=-1) sorted_mask = cumulative_probs > top_p sorted_mask[..., 1:] = sorted_mask[..., :-1].clone() sorted_mask[..., 0] = False filtered_probs = probs.clone() filtered_probs[0, sorted_indices[0][sorted_mask[0]]] = 0.0 # 重新归一化 filtered_probs = filtered_probs / filtered_probs.sum(dim=-1, keepdim=True) next_token = torch.multinomial(filtered_probs, num_samples=1) return next_token温度0.8和top-p 0.9是闲聊场景的保守起点。如果回复仍然跑偏,把温度降到0.6;如果回复过于单调(每次都大同小异),把温度升到1.0。top-p的作用是去掉那些概率加起来只占10%但候选词上百个的长尾——这类长尾里大多是生僻字或逻辑不通的接法。我调参时会先固定top-p=0.9,单独调节温度,观察10条测试问题;温度确定后再微调top-p,一次只动一个参数,别两个一起改,否则没法定位是哪个参数导致的风格漂移。
6.2 验证:用BLEU还是用困惑度,还是都用
这个项目里模型训练日志会打印perplexity,但做模型对比时最好补一个BLEU。BLEU衡量生成回复和参考回复的n-gram重合度,对闲聊来说它不完美——闲聊答案不唯一,同一个“今天天气好吗”可能有十种合理回答——但作为相对指标够用:基线模型A的BLEU-2是0.18,改完mask后的模型B是0.23,这个提升是可信的。
验证集建议建一个“比例抽样”而不是纯随机抽样。把测试问题按长度分成三桶:短(1-10字)、中(11-30字)、长(31字以上),每桶抽等量样本。这样能暴露模型的长句崩溃问题——很多模型短句表现很好,长句开始胡言乱语,纯随机抽样下这类问题会被淹没。
6.3 这套方案还能往哪走
单轮对话是这个Transformer架构最克制的用法,但源码里做好的位置编码、注意力模块和训练循环可以直接复用。想接成本更高的场景,可以先在训练的最后一两个epoch引入bleu损失作为辅助信号,让模型在“预测准确”和“生成可读”之间做权衡——这个trick我在实践中对重复问题有明显改善。如果想再进一步,把单轮对话升级成多轮,只需要在输入序列里拼接[上一轮question, 上一轮answer, 本轮question],相当于把对话历史当作上下文前缀,模型结构不用改一行代码。
部署到本地服务时,把模型加载一次放进内存,常驻进程接收HTTP请求,每来一条请求走一遍前向加解码,单条延迟在CPU上大约30到80毫秒(视seq_len而定),完全够做实时聊天。我踩过的最后一个坑是PyTorch的模型加载需要对应版本,本地训练保存的state_dict换台机器跑如果torch版本不一致很容易报错,保险做法是保存模型时顺手保存一份torch.save时的torch.__version__到readme里。
训练这套方案我最大的教训是:Transformer模型的坑永远在数据准备和mask定义上,模型的forward代码反而很少出错。拿到别人的源码先跑通readme里的示例,再去改自己的数据,省掉一半的排查时间。希望帮到你。
本文还有配套的精品资源,点击获取