You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

使用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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.29 11:58:12