使用Google Flan-T5-Large构建语义搜索引擎时嵌入相似度异常求助
问题分析与修正方案
你的代码存在几个关键问题,导致生成的嵌入无法正确反映语义相似度:
错误调用encoder-decoder模型的生成模式:Flan-T5是encoder-decoder架构,你调用模型时同时传入
input_ids和decoder_input_ids=input_ids,这会触发模型的文本生成逻辑(把输入文本重新生成一遍),此时得到的last_hidden_state是decoder的输出,而非适合语义匹配的encoder编码结果。decoder输出偏向生成任务的token表示,不适合做语义嵌入。Tokenizer使用不规范:直接用
tokenizer.encode没有生成attention_mask,模型无法区分真实token和padding(即便文本短,规范处理也很重要),且没有统一截断/padding逻辑,可能导致嵌入不稳定。池化逻辑忽略attention mask:你直接对整个
last_hidden_state取均值,没有排除padding的0向量,会拉低嵌入的准确性。
修正后的代码
import torch from transformers import AutoTokenizer, AutoModel from sklearn.metrics.pairwise import cosine_similarity # 加载模型和tokenizer tokenizer = AutoTokenizer.from_pretrained('google/flan-t5-large') model = AutoModel.from_pretrained('google/flan-t5-large') def generate_embeddings(text_list): all_embeddings = [] for text in text_list: # 用tokenizer的标准调用方式生成包含attention_mask的输入 inputs = tokenizer(text, return_tensors='pt', padding=True, truncation=True, max_length=512) with torch.no_grad(): # 只传入encoder所需参数,获取encoder的输出 outputs = model(**inputs) last_hidden = outputs.last_hidden_state # 结合attention_mask过滤padding,计算均值池化 mask = inputs['attention_mask'].unsqueeze(-1).expand(last_hidden.size()) mean_embedding = torch.sum(last_hidden * mask, dim=1) / torch.clamp(mask.sum(dim=1), min=1e-9) all_embeddings.append((mean_embedding, text)) return all_embeddings def run_query(query, corpus): # 统一处理查询的输入格式 inputs = tokenizer(query, return_tensors='pt', padding=True, truncation=True, max_length=512) with torch.no_grad(): outputs = model(**inputs) last_hidden = outputs.last_hidden_state mask = inputs['attention_mask'].unsqueeze(-1).expand(last_hidden.size()) query_embedding = torch.sum(last_hidden * mask, dim=1) / torch.clamp(mask.sum(dim=1), min=1e-9) # 用余弦相似度计算语义相关性(值越大越相似) similarity = [] for embed, text in corpus: sim = cosine_similarity(embed.numpy(), query_embedding.numpy())[0][0] similarity.append((text, float(sim))) # 按相似度从高到低排序 similarity.sort(key=lambda x: x[1], reverse=True) return similarity # 测试示例 text_corpus = ['some sad song', 'a very happy song'] corpus_embeddings = generate_embeddings(text_corpus) query_text = "I'm feeling so sad rn" results = run_query(query_text, corpus_embeddings) for item in results: print(f"文本: {item[0]}, 余弦相似度: {item[1]:.4f}")
额外说明
如果还是觉得语义匹配效果不够理想,Flan-T5本身并非专门为语义嵌入设计的模型,你可以尝试专门的嵌入模型(如Sentence-BERT系列),这类模型在语义相似度任务上的表现会更针对性。同时要保证所有文本的处理逻辑完全一致,否则嵌入会失去可比性。
内容的提问来源于stack exchange,提问作者Affan Mir
相关产品推荐
相关产品推荐

