如何解决瑞典语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
相关产品推荐
相关产品推荐

