使用BERT模型计算词语相似度时遇张量维度错误的求助
BERT词相似度计算报错解决:too many indices for tensor of dimension 2
错误含义
这个错误表示你尝试对一个2维张量使用了超过2个维度的索引(比如三层索引),但该张量只有两个维度,无法匹配你的索引操作。
错误根源
你的代码存在两处核心问题:
张量维度冗余与索引逻辑错误
tokenizer.encode(..., return_tensors='pt')已经返回形状为[1, seq_len]的张量,你额外添加的unsqueeze(0)会让张量变成[1,1,seq_len],后续squeeze(0)操作依然会导致维度混乱。- 调用
model(input_ids=xxx)[1]['last_hidden_state']完全错误:AutoModel的输出是元组,[1]对应的是pooler_output(2维张量[batch_size, hidden_size]),它根本没有last_hidden_state这个键;而last_hidden_state是输出元组的第一个元素([0]),形状为[batch_size, seq_len, hidden_size]。
词嵌入提取逻辑不合理
直接取last_hidden_state[0]会拿到整个序列的嵌入,没有针对目标词做正确的聚合(比如取有效子词的平均嵌入)。
修正后的代码
import torch from transformers import AutoTokenizer, AutoModel # Load the BERT model model_name = 'bert-base-uncased' tokenizer = AutoTokenizer.from_pretrained(model_name) model = AutoModel.from_pretrained(model_name) # Encode the target word and the list of words target_word = "apple" word_list = ["blackberry", "iphone", "microsoft", "blueberry", "pineapple"] # 定义词嵌入获取函数:处理分词,返回子词平均后的嵌入 def get_word_embedding(word, tokenizer, model): inputs = tokenizer(word, return_tensors='pt', padding=True, truncation=True) with torch.no_grad(): outputs = model(**inputs) # 获取隐藏层输出,形状[1, seq_len, hidden_size] last_hidden = outputs.last_hidden_state # 过滤掉<SEP>和<SEP>的位置,只保留目标词的子词嵌入 mask = (inputs['input_ids'] != tokenizer.cls_token_id) & (inputs['input_ids'] != tokenizer.sep_token_id) # 对有效子词嵌入做平均,得到整个词的嵌入 word_emb = last_hidden[mask].mean(dim=0) return word_emb # 获取目标词与词列表的嵌入 target_emb = get_word_embedding(target_word, tokenizer, model) word_embs = [get_word_embedding(word, tokenizer, model) for word in word_list] # 计算余弦相似度 similarities = [torch.nn.functional.cosine_similarity(target_emb, emb, dim=0).item() for emb in word_embs] # 打印结果 for word, similarity in zip(word_list, similarities): print(f"Similarity between '{target_word}' and '{word}': {similarity:.2f}")
关键修改说明
- 移除冗余的
unsqueeze(0)和手动pad操作:利用tokenizer自带的padding=True自动处理长度对齐,减少维度混乱。 - 用
outputs.last_hidden_state直接获取隐藏层输出,避免索引错误。 - 新增子词聚合逻辑:针对BERT的分词特性,对目标词的所有有效子词嵌入取平均,得到更合理的整词嵌入。
- 函数化提取嵌入,代码更简洁易维护。
内容的提问来源于stack exchange,提问作者spain engy
相关产品推荐
相关产品推荐

