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

如何解决瑞典语BERT匹配Pandas列相似字符串的RuntimeError?

瑞典语文本相似度匹配的RuntimeError解决方法

问题背景

我尝试从Pandas数据框的列中,找到与瑞典语输入字符串最相似的匹配项,通过编码文本并计算余弦相似度实现,代码如下:

# Load the pre-trained BERT model for Swedish language
tokenizer = AutoTokenizer.from_pretrained("KBLab/sentence-bert-swedish-cased")
model = AutoModel.from_pretrained("KBLab/sentence-bert-swedish-cased")

def find_similar(input_text, column_name, df):
    # Encode the input text and the text in the column
    encoded_input = tokenizer.encode(input_text, return_tensors='pt', truncation=True)
    input_ids = encoded_input[0]
    column_text = df[column_name].tolist()
    column_ids = [tokenizer.encode(text, return_tensors='pt') for text in column_text]

    # Calculate the similarity score between the input text and each text in the column
    with torch.no_grad():
        input_embeddings = model(encoded_input).last_hidden_state
        column_embeddings = [model(column_id).last_hidden_state for column_id in column_ids]
        similarity_scores = [cosine_similarity(input_embeddings, embedding).item() for embedding in column_embeddings]
    
    # Find the index of the text with the highest similarity score
    max_index = similarity_scores.index(max(similarity_scores))
    
    # Return the most similar text
    return column_text[max_index]

input_text = test
column_name = "DESC"
most_similar_text = find_similar(input_text, column_name, df)
print(most_similar_text)

报错信息

运行时抛出RuntimeError,错误详情:

RuntimeError                              Traceback (most recent call last)
Cell In [49], line 3
      1 input_text = test
      2 column_name = "DESC"
----> 3 most_similar_text = find_similar(input_text, column_name, df)
      4 print(most_similar_text)

Cell In [48], line 12, in find_similar(input_text, column_name, df)
     10     input_embeddings = model(encoded_input).last_hidden_state
     11     column_embeddings = [model(column_id).last_hidden_state for column_id in column_ids]
---> 12     similarity_scores = [cosine_similarity(input_embeddings, embedding).item() for embedding in column_embeddings]
     14 # Find the index of the text with the highest similarity score
     15 max_index = similarity_scores.index(max(similarity_scores))

Cell In [48], line 12, in <listcomp>(.0)
     10     input_embeddings = model(encoded_input).last_hidden_state
     11     column_embeddings = [model(column_id).last_hidden_state for column_id in column_ids]
---> 12     similarity_scores = [cosine_similarity(input_embeddings, embedding).item() for embedding in column_embeddings]
     14 # Find the index of the text with the highest similarity score
     15 max_index = similarity_scores.index(max(similarity_scores))

RuntimeError: The size of tensor a (19) must match the size of tensor b (10) at non-singleton dimension 1

解决方法

报错核心原因:BERT输出的last_hidden_state是序列级嵌入(形状为[batch_size, seq_len, hidden_size]),输入文本与列中文本的序列长度不一致,导致余弦相似度计算时维度不匹配。

需要将序列嵌入转换为句子级嵌入,同时优化代码效率,修正后的完整代码如下:

import torch
from sklearn.metrics.pairwise import cosine_similarity
from transformers import AutoTokenizer, AutoModel

# 加载瑞典语预训练BERT模型
tokenizer = AutoTokenizer.from_pretrained("KBLab/sentence-bert-swedish-cased")
model = AutoModel.from_pretrained("KBLab/sentence-bert-swedish-cased")

def mean_pooling(model_output, attention_mask):
    # 均值池化生成句子嵌入,忽略padding部分
    token_embeddings = model_output.last_hidden_state
    input_mask = attention_mask.unsqueeze(-1).expand(token_embeddings.size()).float()
    return torch.sum(token_embeddings * input_mask, 1) / torch.clamp(input_mask.sum(1), min=1e-9)

def find_similar(input_text, column_name, df):
    column_text = df[column_name].tolist()
    # 批量编码输入文本和所有列文本,自动处理padding和截断
    all_texts = [input_text] + column_text
    encoded_input = tokenizer(all_texts, padding=True, truncation=True, return_tensors='pt')
    
    with torch.no_grad():
        model_output = model(**encoded_input)
        # 生成统一维度的句子嵌入
        sentence_embeddings = mean_pooling(model_output, encoded_input['attention_mask'])
    
    # 计算余弦相似度
    input_embedding = sentence_embeddings[0].reshape(1, -1)
    column_embeddings = sentence_embeddings[1:]
    similarity_scores = cosine_similarity(input_embedding, column_embeddings)[0]
    
    # 获取相似度最高的文本索引
    max_index = similarity_scores.argmax()
    
    return column_text[max_index]

# 示例调用(替换为实际的瑞典语文本)
input_text = "din svenska testtext"
column_name = "DESC"
most_similar_text = find_similar(input_text, column_name, df)
print(most_similar_text)

关键修正说明

  • 均值池化处理:通过mean_pooling函数将序列级嵌入转换为固定维度的句子嵌入,彻底解决维度不匹配问题
  • 批量编码优化:避免循环编码每个列文本,大幅提升运行效率
  • 索引获取优化:用argmax直接获取最大相似度索引,比原方法更高效且避免重复值问题
  • 变量修正:原代码中input_text = test未定义变量,示例中补充了实际瑞典语文本提示

内容的提问来源于stack exchange,提问作者Vai

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.02 01:50:29