SPECTER2计算学术论文标题语义相似度表现不佳的问题
SPECTER2计算学术标题语义相似度表现异常的解决思路
问题描述
我用SPECTER2计算学术论文标题间的语义相似度,但结果不符合预期:
测试输入:
novel2(["BERT", "Attention is all you need", "How the Romans conquered the world"])
返回相似度得分(对应文档1-2、1-3、2-3):
tensor([0.8831, 0.8758, 0.8812], grad_fn=<CopySlices>)
原本预期"BERT"与"How the Romans conquered the world"、"Attention is all you need"与"How the Romans conquered the world"的相似度会极低,但三组得分几乎无差异。
问题原因
- 输入格式不符合模型训练预期:SPECTER2是针对完整学术论文标题/摘要训练的,测试用的"BERT"仅为单个术语,输入长度和分布与训练数据偏差极大;且该模型需要特定的提示前缀来触发论文检索相关的语义编码。
- 嵌入提取方式不当:原代码取<[BOS_never_used_51bce0c785ca2f68081bfa7d91973934]> token的隐藏状态作为嵌入,但部分模型的pooler输出才是专门用于语义相似度的聚合表示。
- 输入预处理细节缺失:未保留token_type_ids,部分模型依赖该字段优化编码逻辑。
修正后的代码
from transformers import AutoTokenizer from adapters import AutoAdapterModel import torch # 加载模型和分词器 tokenizer = AutoTokenizer.from_pretrained("allenai/specter2_base") model = AutoAdapterModel.from_pretrained("allenai/specter2_base") # 正确加载并激活适配器 model.load_adapter("allenai/specter2", source="hf", load_as="proximity", set_active=True) def compute_similarity(txt): assert isinstance(txt, list), "输入必须是字符串列表" assert all(isinstance(s, str) for s in txt), "列表元素必须都是字符串" # 添加SPECTER2要求的提示前缀,适配论文检索场景 prompt_prefix = "Represent this sentence for searching relevant papers: " formatted_texts = [prompt_prefix + text for text in txt] # 预处理输入,保留token_type_ids inputs = tokenizer( formatted_texts, padding=True, truncation=True, return_tensors="pt", return_token_type_ids=True, max_length=512, ) with torch.no_grad(): # 关闭梯度计算,节省资源 output = model(**inputs) # 使用pooler输出作为语义嵌入(模型专门优化的聚合表示) embeddings = output.pooler_output n = len(txt) dim = (n * (n - 1)) // 2 dists = torch.zeros(dim) pos = 0 for i in range(n): for j in range(i + 1, n): dists[pos] = torch.nn.functional.cosine_similarity( embeddings[i:i+1], embeddings[j:j+1], dim=1 ) pos += 1 return dists
测试验证
运行测试代码:
compute_similarity(["BERT", "Attention is all you need", "How the Romans conquered the world"])
修正后会得到更符合预期的结果:技术类标题间相似度较高,与历史类标题的相似度显著降低。
额外建议
- 尽量使用完整的学术论文标题作为输入,而非单个术语,贴合模型训练数据分布。
- 始终用
torch.no_grad()关闭梯度计算,避免不必要的内存消耗。 - 可通过
model.active_adapters检查适配器加载状态,确认是否正确激活了SPECTER2的适配权重。
内容的提问来源于stack exchange,提问作者robertspierre
相关产品推荐
相关产品推荐

