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

使用BERT计算两词余弦相似度时遇IndexError问题求助

问题原因及解决办法

错误根源

你遇到的IndexError是因为torch.cosine_similarity的参数设置问题:

  • 从BERT输出中提取的word1_embedding和word2_embedding是1维张量(形状为torch.Size([768]))
  • torch.cosine_similarity默认使用dim=1计算相似度,但1维张量只有dim=0这个维度,指定dim=1会超出范围,触发报错。

两种修复方案

方案1:指定计算维度为0

直接在调用cosine_similarity时显式设置dim=0,适配1维张量:

from transformers import BertTokenizer, BertModel
import torch

# Load the BERT model and tokenizer
model = BertModel.from_pretrained('bert-base-uncased')
tokenizer = BertTokenizer.from_pretrained('bert-base-uncased')

# Tokenize the input words
word1 = "cat"
word2 = "dog"
input_ids = torch.tensor([tokenizer.encode(word1, word2, add_special_tokens=True)])

# Get the BERT embeddings for the input words
output = model(input_ids)[0]

# Get the first and second word embeddings
word1_embedding = output[0, 1, :]
word2_embedding = output[0, 2, :]

# Calculate the cosine similarity between the two words
similarity = torch.cosine_similarity(word1_embedding, word2_embedding, dim=0)
print(similarity)

方案2:给词向量增加一个维度

通过unsqueeze(0)把1维张量转为2维(形状变为torch.Size([1, 768])),适配默认的dim=1参数:

from transformers import BertTokenizer, BertModel
import torch

# Load the BERT model and tokenizer
model = BertModel.from_pretrained('bert-base-uncased')
tokenizer = BertTokenizer.from_pretrained('bert-base-uncased')

# Tokenize the input words
word1 = "cat"
word2 = "dog"
input_ids = torch.tensor([tokenizer.encode(word1, word2, add_special_tokens=True)])

# Get the BERT embeddings for the input words
output = model(input_ids)[0]

# Get the first and second word embeddings, add extra dimension
word1_embedding = output[0, 1, :].unsqueeze(0)
word2_embedding = output[0, 2, :].unsqueeze(0)

# Calculate the cosine similarity between the two words
similarity = torch.cosine_similarity(word1_embedding, word2_embedding)
print(similarity)

额外说明

原代码中提取词向量的逻辑是正确的:BERT的输入会自动加上[CLS](索引0)和[SEP](索引3),所以cat对应索引1,dog对应索引2,这部分无需修改。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.05 19:50:11