如何给SentenceTransformer添加新词汇并解决编码与保存问题
问题背景
刚入门NLP,搭建了基于用户输入检索相似文本的推荐系统(用SentenceTransformer+FAISS实现),代码可正常运行,但遇到模型词汇库未收录的新词(如Covid-19)或特殊代码时检索效果差。目标是支持SQL-like筛选条件的k近邻检索(例:输入gender = MAN and sports in (football),推荐gender = MAN and sports in (football, baseball))。尝试给模型添加新词时触发报错:AttributeError: 'BertModel' object has no attribute 'encode',需解决以下三个核心问题:
1. 添加新词后如何实现文本编码
报错原因
你用transformers.AutoModel加载的是原生BERT模型,它没有encode方法——这个方法是SentenceTransformer封装的便捷接口。要实现新词适配后的编码,有两种可行方案:
方案一:自定义编码函数适配原生模型
直接基于修改后的AutoModel和AutoTokenizer实现编码逻辑,模仿SentenceTransformer的句子嵌入生成方式(取<[BOS_never_used_51bce0c785ca2f68081bfa7d91973934]>token的输出):
import pandas as pd import numpy as np import torch from transformers import AutoTokenizer, AutoModel import faiss # 加载原模型组件 model_name = "sentence-transformers/all-MiniLM-L6-v2" tokenizer = AutoTokenizer.from_pretrained(model_name) base_model = AutoModel.from_pretrained(model_name) # 添加新词并调整词嵌入层 new_tokens = ["Covid-19", "XYZ-789"] # 过滤已存在于词汇表的词 new_tokens = [t for t in new_tokens if t not in tokenizer.vocab] if new_tokens: tokenizer.add_tokens(new_tokens) # 扩展词嵌入层,新嵌入随机初始化 base_model.resize_token_embeddings(len(tokenizer)) # 自定义编码函数 def encode_texts(texts, model, tokenizer): device = "cuda" if torch.cuda.is_available() else "cpu" model.to(device).eval() embeddings = [] with torch.no_grad(): for text in texts: inputs = tokenizer( text, return_tensors="pt", padding=True, truncation=True, max_length=512 ).to(device) outputs = model(**inputs) # 取<[BOS_never_used_51bce0c785ca2f68081bfa7d91973934]>token的输出作为句子嵌入 cls_embed = outputs.last_hidden_state[:, 0, :].cpu().numpy() embeddings.append(cls_embed) return np.concatenate(embeddings, axis=0) # 加载数据并编码 df = pd.read_excel("data/data.xlsx", "data") document_embeddings = encode_texts(df['combined'], base_model, tokenizer) # FAISS索引构建及搜索逻辑(与原代码一致) d = document_embeddings.shape[1] index = faiss.IndexFlatL2(d) index.add(document_embeddings) def search_query(input_text, top_k): xq = encode_texts([input_text], base_model, tokenizer) D, I = index.search(xq, top_k) print(f'Input: {input_text}') for i in range(top_k): print(f'Similar Content: {df.iloc[I[0][i]]["combined"]}') search_query("Covid-19 prevention", 10)
方案二:将修改后的模型封装为SentenceTransformer
如果想继续使用SentenceTransformer的encode方法,可以把修改后的模型重新封装为其兼容的组件:
from sentence_transformers import SentenceTransformer, models # 保存修改后的tokenizer和模型(先执行方案一中的修改步骤) tokenizer.save_pretrained("./modified_minilm") base_model.save_pretrained("./modified_minilm") # 用SentenceTransformer加载封装 word_embedding_model = models.Transformer("./modified_minilm") # 配置池化层(取<[BOS_never_used_51bce0c785ca2f68081bfa7d91973934]>token作为句子嵌入) pooling_model = models.Pooling( word_embedding_model.get_word_embedding_dimension(), pooling_mode_cls_token=True, pooling_mode_mean_tokens=False ) model = SentenceTransformer(modules=[word_embedding_model, pooling_model]) # 现在可以直接用model.encode() document_embeddings = model.encode(df['combined'])
2. 如何保存修改后的模型供后续使用
需要同时保存tokenizer和模型权重,因为两者是绑定的:
# 保存修改后的tokenizer和模型 tokenizer.save_pretrained("./custom_minilm_model") base_model.save_pretrained("./custom_minilm_model") # 后续加载使用 from transformers import AutoTokenizer, AutoModel tokenizer = AutoTokenizer.from_pretrained("./custom_minilm_model") base_model = AutoModel.from_pretrained("./custom_minilm_model") # 如果用SentenceTransformer加载,参考方案二的封装逻辑
保存的目录会包含以下文件:config.json、pytorch_model.bin(模型权重)、vocab.txt(更新后的词汇表)等,后续直接加载即可复用修改后的词汇和模型。
3. 模型可添加的新词数量上限及超量解决方案
数量上限
理论上没有硬性限制,但受两个核心因素约束:
- 显存容量:每个新词对应一个嵌入向量(MiniLM的嵌入维度为384),每添加1000个词会增加约384KB的参数(float32精度),大量添加会快速消耗显存。
- 模型设计:部分预训练模型的词汇表有默认上限(如BERT-base的默认词汇表大小为30522),但
resize_token_embeddings方法可以突破这个限制,只是显存压力会随之增大。
超量后的解决方案
如果需要添加的新词数量极大(如数万级),建议采用以下替代方案:
- 依赖子词拆分:大部分现代tokenizer(如BERT的WordPiece)会自动将未知词拆分为已有子词,无需手动添加。对于特殊代码/术语,可以提前做标准化处理(如把
XYZ-123拆分为XYZ和123)。 - Prompt提示法:无需修改模型,将新词用自然语言描述后加入输入(如把
Covid-19替换为新冠病毒(Covid-19)),让模型通过上下文理解语义。 - 轻量级参数微调:用LoRA(Low-Rank Adaptation)等方法,只微调模型的部分参数,无需扩展整个词嵌入层,大幅降低显存占用。
- 更换大词汇量模型:选择针对领域优化的预训练模型(如Roberta-large、GPT系列),这类模型的词汇表更大,对未知词的兼容性更好。
- 字符级模型:改用字符级Transformer模型,完全不依赖预训练词汇表,对任意新词/特殊字符都能直接处理。
内容的提问来源于stack exchange,提问作者0o0SkY0o0

