
gte-base-zh模型蒸馏教程打造轻量级学生模型你是不是遇到过这种情况一个AI模型效果特别好但就是太大了跑起来慢还特别占资源想把它塞到手机或者小设备里根本不可能。这时候模型蒸馏技术就能派上大用场了。简单来说模型蒸馏就像一位经验丰富的老师大模型在教一个聪明的学生小模型。老师把自己多年积累的“知识”和“解题思路”传授给学生让学生虽然体量小但也能达到接近老师的水平。今天我们就来手把手教你如何把强大的gte-base-zh模型老师的知识蒸馏到一个轻量级的学生模型上让它变得又快又小还能在边缘设备上欢快地跑起来。通过这篇教程即使你之前没接触过模型蒸馏也能跟着步骤一步步完成。我们会从最基础的概念讲起准备好数据设计好训练方法最后得到一个可以部署的轻量模型。整个过程我们都会用代码和例子来说明保证你能看懂、能操作。1. 准备工作理解蒸馏与搭建环境在开始动手之前我们得先搞清楚两件事我们要用的“老师”和“学生”是谁以及我们需要准备什么样的工具和环境。1.1 认识我们的“老师”与“学生”教师模型 (Teacher Model):gte-base-zh这是我们知识的主要来源。gte-base-zh是一个在中文文本上训练的高质量文本嵌入模型它擅长理解句子的含义并把句子转换成一组有意义的数字向量。它效果很好但相应的模型参数也多计算量较大。学生模型 (Student Model)这是我们要训练的目标。为了能在资源有限的设备上运行学生模型通常结构更简单、参数更少。比如我们可以选择一个更小的BERT变体如bert-mini、bert-tiny或者像MiniLM这样专门为蒸馏设计的架构。它的目标是模仿老师的行为输出相似的句子向量。知识蒸馏的核心思想我们不仅仅让学生模型去学习原始的训练数据比如判断两个句子是否相似这个“标准答案”更重要的是让它去学习老师模型输出的“软标签”。什么是软标签呢想象一下老师做选择题他不仅告诉你正确答案是A还会告诉你“A的可能性是80%B是15%C是5%”。这种概率分布包含了老师更丰富的知识学生学这个往往比只学“答案是A”这个硬标签效果更好。1.2 环境配置与安装我们需要一个Python环境以及一些关键的机器学习库。推荐使用conda创建一个独立的环境。# 创建并激活一个新的conda环境可选但推荐 conda create -n model-distillation python3.8 conda activate model-distillation # 安装核心库PyTorch请根据你的CUDA版本选择安装命令 # 例如对于CUDA 11.3 pip install torch torchvision torchaudio --index-url https://download.pytorch.org/whl/cu113 # 安装Hugging Face Transformers和Datasets库这是我们操作模型和数据的主力 pip install transformers datasets # 安装用于评估和训练的辅助库 pip install scikit-learn tqdm安装完成后我们可以用下面几行代码快速验证一下环境并看看如何加载我们的老师模型。from transformers import AutoModel, AutoTokenizer # 加载教师模型 gte-base-zh 及其分词器 teacher_model_name thenlper/gte-base-zh teacher_tokenizer AutoTokenizer.from_pretrained(teacher_model_name) teacher_model AutoModel.from_pretrained(teacher_model_name) print(f教师模型加载成功: {teacher_model_name}) print(f模型结构概览: {teacher_model})2. 知识蒸馏的核心损失函数与数据准备环境好了接下来就是准备“教材”数据和设计“教学方法”损失函数了。2.1 设计蒸馏损失函数在蒸馏过程中学生模型的损失通常由两部分组成蒸馏损失 (Distillation Loss)让学生模型的输出logits或嵌入向量尽量靠近老师模型的输出。这让学生学习老师的“知识”。任务损失 (Task Loss)让学生模型在原始任务如文本相似度上的表现也不要太差。这确保学生不偏离基本目标。对于gte-base-zh这种产出句子向量的模型一个常见的做法是使用余弦嵌入损失或均方误差损失来对齐师生模型的输出向量。import torch import torch.nn as nn import torch.nn.functional as F class DistillationLoss(nn.Module): def __init__(self, alpha0.5, temperature2.0): Args: alpha: 控制蒸馏损失和任务损失权重的超参数0alpha1。 alpha越大越依赖老师蒸馏损失权重高。 temperature: “温度”参数用于软化概率分布让知识更容易传递。 super().__init__() self.alpha alpha self.temperature temperature self.mse_loss nn.MSELoss() # 用于对齐向量 # 假设我们还有一个原始任务损失例如对比学习损失 # 这里我们用余弦相似度损失作为任务损失的简化示例 self.task_loss nn.CosineEmbeddingLoss() def forward(self, student_vec, teacher_vec, labelsNone): student_vec: 学生模型输出的句子向量 teacher_vec: 教师模型输出的句子向量 labels: 原始任务标签例如句子对是否相似 # 1. 计算蒸馏损失让学生向量逼近老师向量使用MSE或余弦距离 # 我们这里使用MSE简单有效 loss_distill self.mse_loss(student_vec, teacher_vec.detach()) # detach老师不更新老师参数 if labels is not None: # 2. 计算任务损失学生模型在原始任务上的表现 # 例如对于相似度任务我们希望相似句子的向量接近不相似句子的向量远离 # 这里简化处理计算学生向量与一个理想目标可以是教师向量也可以是其他的余弦损失 # 更复杂的任务需要设计对应的损失函数 loss_task self.task_loss(student_vec, teacher_vec.detach(), labels) # 3. 组合损失 total_loss self.alpha * loss_distill (1 - self.alpha) * loss_task return total_loss, loss_distill, loss_task else: # 如果只有蒸馏例如无标签数据则只使用蒸馏损失 return loss_distill2.2 准备训练数据数据是训练的基础。对于蒸馏我们既可以利用有标签的监督数据同时优化任务损失也可以利用大量无标签的数据只优化蒸馏损失让学到的向量空间更通用。方案一使用公开数据集有监督例如我们可以使用中文语义相似度数据集如ATEC、BQ Corpus或LCQMC。这些数据提供了句子对和它们的相似度标签。from datasets import load_dataset # 以LCQMC数据集为例 dataset load_dataset(shibing624/nli_zh, LCQMC) print(dataset[train][0]) # 查看一条数据{sentence1: 句子1, sentence2: 句子2, label: 1相似或0不相似}方案二构建无标签数据我们可以从维基百科、新闻文章或特定领域文本中收集大量句子无需标注。蒸馏的目标就是让学生模型为这些句子生成的向量与老师模型生成的向量尽可能一致。# 假设我们有一个文本文件每行一个句子 with open(unlabeled_sentences.txt, r, encodingutf-8) as f: unlabeled_sentences [line.strip() for line in f.readlines()[:10000]] # 取前1万句示例 # 接下来我们需要用老师和学生的分词器分别处理这些句子并得到它们的向量表示。 # 这个过程通常需要批量进行我们会在训练循环中看到。3. 分步实践构建与训练学生模型现在让我们把“老师”、“学生”、“教材”和“教学方法”组合起来开始真正的训练。3.1 构建学生模型我们选择bert-mini作为学生模型它比gte-base-zh小很多。from transformers import AutoModel, AutoConfig student_model_name nghuyong/bert-mini-zh # 一个小的中文BERT模型 student_config AutoConfig.from_pretrained(student_model_name) # 关键修改学生模型的输出维度使其与教师模型匹配 # gte-base-zh 的隐藏层大小是1024输出池化后也是1024维向量 # bert-mini 原始隐藏层是256我们需要加一个投影层来匹配 teacher_hidden_size 1024 student_hidden_size student_config.hidden_size # bert-mini 是256 class StudentModelForDistillation(nn.Module): def __init__(self, base_student_model): super().__init__() self.bert base_student_model # 加载预训练的bert-mini # 一个简单的投影层将学生模型的256维输出映射到老师模型的1024维 self.projection nn.Linear(student_hidden_size, teacher_hidden_size) def forward(self, input_ids, attention_mask): # 获取学生模型的序列输出 outputs self.bert(input_idsinput_ids, attention_maskattention_mask) # 通常取 [CLS] 位置的输出作为句子表示这里简化处理 # 实际中gte-base-zh可能使用不同的池化策略需要对齐 sequence_output outputs.last_hidden_state # [batch, seq_len, hidden] cls_output sequence_output[:, 0, :] # 取[CLS] token # 投影到与老师相同的空间 sentence_embedding self.projection(cls_output) # 通常还会进行归一化与老师模型保持一致 sentence_embedding F.normalize(sentence_embedding, p2, dim1) return sentence_embedding # 初始化学生模型 base_student AutoModel.from_pretrained(student_model_name) student_model StudentModelForDistillation(base_student) print(f学生模型参数量: {sum(p.numel() for p in student_model.parameters()):,})3.2 编写训练循环这是最核心的部分我们将在一个循环中完成前向传播、损失计算和反向传播。from torch.utils.data import DataLoader, Dataset from tqdm.auto import tqdm import numpy as np # 1. 构建一个简单的数据集类 class DistillationDataset(Dataset): def __init__(self, sentences, tokenizer, max_len128): self.sentences sentences self.tokenizer tokenizer self.max_len max_len def __len__(self): return len(self.sentences) def __getitem__(self, idx): encoding self.tokenizer( self.sentences[idx], truncationTrue, paddingmax_length, max_lengthself.max_len, return_tensorspt ) # 去掉batch维度因为DataLoader会添加 return {key: val.squeeze(0) for key, val in encoding.items()} # 假设我们使用无标签句子 train_sentences unlabeled_sentences[:8000] # 训练集 val_sentences unlabeled_sentences[8000:] # 验证集 train_dataset DistillationDataset(train_sentences, student_tokenizer) val_dataset DistillationDataset(val_sentences, student_tokenizer) train_loader DataLoader(train_dataset, batch_size32, shuffleTrue) val_loader DataLoader(val_dataset, batch_size32) # 2. 初始化模型、损失函数和优化器 device torch.device(cuda if torch.cuda.is_available() else cpu) teacher_model.to(device).eval() # 老师模型固定不训练 student_model.to(device) criterion DistillationLoss(alpha0.7, temperature2.0) # 调整alpha和temperature optimizer torch.optim.AdamW(student_model.parameters(), lr2e-5) # 3. 训练循环 num_epochs 3 for epoch in range(num_epochs): student_model.train() total_loss 0 progress_bar tqdm(train_loader, descfEpoch {epoch1}) for batch in progress_bar: # 将数据移到设备 input_ids batch[input_ids].to(device) attention_mask batch[attention_mask].to(device) # 前向传播 with torch.no_grad(): # 不计算老师模型的梯度 teacher_outputs teacher_model(input_idsinput_ids, attention_maskattention_mask) # 获取教师模型的句子向量具体方法需参考gte-base-zh的用法 # 这里假设通过pooler_output获得 teacher_embeddings teacher_outputs.pooler_output teacher_embeddings F.normalize(teacher_embeddings, p2, dim1) student_embeddings student_model(input_idsinput_ids, attention_maskattention_mask) # 计算损失无标签只计算蒸馏损失 loss criterion(student_embeddings, teacher_embeddings) # 反向传播 optimizer.zero_grad() loss.backward() optimizer.step() total_loss loss.item() progress_bar.set_postfix({loss: loss.item()}) avg_train_loss total_loss / len(train_loader) print(fEpoch {epoch1} 平均训练损失: {avg_train_loss:.4f}) # 简单验证 student_model.eval() val_loss 0 with torch.no_grad(): for batch in val_loader: input_ids batch[input_ids].to(device) attention_mask batch[attention_mask].to(device) teacher_outputs teacher_model(input_idsinput_ids, attention_maskattention_mask) teacher_embeddings teacher_outputs.pooler_output teacher_embeddings F.normalize(teacher_embeddings, p2, dim1) student_embeddings student_model(input_idsinput_ids, attention_maskattention_mask) loss criterion(student_embeddings, teacher_embeddings) val_loss loss.item() avg_val_loss val_loss / len(val_loader) print(fEpoch {epoch1} 平均验证损失: {avg_val_loss:.4f}) print(训练完成)4. 模型评估与部署训练完成后我们得看看这个“学生”学得怎么样然后把它打包带走放到小设备上去用。4.1 评估学生模型性能评估嵌入模型的一个常见方法是看它在下游任务上的表现比如语义相似度计算。我们可以用斯皮尔曼相关系数来衡量模型预测的相似度与人工标注的相似度之间的相关性。from sklearn.metrics.pairwise import cosine_similarity import scipy.stats def evaluate_on_sts_dataset(student_model, tokenizer, dataset, device): 在语义文本相似度数据集上评估模型 similarities_pred [] similarities_true [] student_model.eval() with torch.no_grad(): for example in dataset: sent1, sent2, true_score example[sentence1], example[sentence2], example[label] # 编码句子 enc1 tokenizer(sent1, return_tensorspt, paddingTrue, truncationTrue, max_length128).to(device) enc2 tokenizer(sent2, return_tensorspt, paddingTrue, truncationTrue, max_length128).to(device) # 获取向量 vec1 student_model(**enc1).cpu().numpy() vec2 student_model(**enc2).cpu().numpy() # 计算余弦相似度 pred_score cosine_similarity(vec1, vec2)[0][0] similarities_pred.append(pred_score) similarities_true.append(true_score) # 计算斯皮尔曼相关系数 spearman_corr scipy.stats.spearmanr(similarities_true, similarities_pred)[0] return spearman_corr # 加载一个小的测试集例如STS-B的中文部分或LCQMC的测试集 test_dataset load_dataset(shibing624/nli_zh, LCQMC, splittest) # 取前100条进行评估 sample_test [test_dataset[i] for i in range(100)] spearman evaluate_on_sts_dataset(student_model, student_tokenizer, sample_test, device) print(f学生模型在测试集上的斯皮尔曼相关系数: {spearman:.4f}) # 可以和教师模型在同一测试集上的结果对比看看性能保留了多少4.2 模型保存与轻量化部署评估满意后我们就可以保存模型并考虑进一步的优化以便部署。# 保存学生模型和分词器 save_path ./distilled_gte_mini student_model.bert.save_pretrained(save_path) # 保存底层BERT student_tokenizer.save_pretrained(save_path) # 注意投影层需要单独保存或者将其合并到底层模型中这里简化处理 torch.save(student_model.projection.state_dict(), f{save_path}/projection_layer.pth) print(f模型已保存至 {save_path}) # 对于部署我们还可以考虑 # 1. 模型量化使用PyTorch的量化工具减少模型大小提升推理速度 # 2. 转换为ONNX格式获得更好的跨平台推理性能 # 3. 使用移动端推理框架如TensorFlow Lite、PyTorch Mobile、MNN等 # 示例动态量化非常简单的示例实际需调整 quantized_model torch.quantization.quantize_dynamic( student_model, {torch.nn.Linear}, dtypetorch.qint8 ) print(模型量化完成示例。)5. 总结与进阶思考走完这一整套流程你应该已经成功地将gte-base-zh的知识蒸馏到了一个更小的模型里。回顾一下我们先是理解了蒸馏就是让小学生模仿大学生解题思路的道理然后准备好了计算环境、定义了怎么算“模仿得像”损失函数、收集了练习题数据接着一步步搭建了小学生的模型架构并开始了训练。最后我们还检查了这个小学生的考试成绩并把它打包成了可以随时带走的“学习卡片”。整个过程里有几个地方特别关键也值得你后续多琢磨一是损失函数里那个平衡老师和原始任务的alpha参数还有让知识变“软”好吸收的temperature参数调一调它们效果可能会不一样。二是学生模型结构的设计我们这里简单加了个投影层其实可以更精巧。三是训练数据无标签数据用起来方便但如果能混合一些有标签的监督数据学生可能学得更扎实。这个蒸馏出来的小模型已经可以在很多对延迟和资源敏感的场景里发挥作用了比如手机App里的语义搜索、智能客服的实时匹配或者嵌入式设备上的文本分类。当然如果你想追求极致的性能还可以探索更高级的蒸馏技巧比如只蒸馏中间某几层的知识或者让多个老师一起教一个学生。希望这篇教程能帮你打开模型压缩和加速的一扇门。动手试试吧调整参数换换不同的学生模型结构说不定你能蒸馏出一个更出色的“小学霸”。获取更多AI镜像想探索更多AI镜像和应用场景访问 CSDN星图镜像广场提供丰富的预置镜像覆盖大模型推理、图像生成、视频生成、模型微调等多个领域支持一键部署。