基于BERT注意力分数,如何选择文本中与查询相关的高权重n-grams?
当将(query, text)序列对输入BERT这类Transformer模型时,我们可以通过分析注意力机制的分数,从文本中提取与查询词关联最紧密的unigrams(单字词)和bigrams(双字词)。以下是完整的实现流程:
示例场景
- Query:
machine learning - 目标文本:
Supervised learning is the machine learning task of learning a function that maps an input to an output based on example input-output pairs. It infers a function from labeled training data consisting of a set of training examples. In supervised learning, each example is a pair consisting of an input object (typically a vector) and a desired output value (also called the supervisory signal). A supervised learning algorithm analyzes the training data and produces an inferred function, which can be used for mapping new examples.
预期筛选结果:machine learning, supervised learning, function, labeled training
完整实现步骤
1. 加载模型与预处理文本
首先加载BERT模型和分词器,将query和text组成句子对进行编码:
from transformers import AutoTokenizer, BertModel import torch import numpy as np # 加载预训练模型和分词器 tokenizer = AutoTokenizer.from_pretrained("bert-base-uncased") model = BertModel.from_pretrained("bert-base-uncased", output_attentions=True) # 定义查询和文本 query = "machine learning" text = """ Supervised learning is the machine learning task of learning a function that maps an input to an output based on example input-output pairs. It infers a function from labeled training data consisting of a set of training examples. In supervised learning, each example is a pair consisting of an input object (typically a vector) and a desired output value (also called the supervisory signal). A supervised learning algorithm analyzes the training data and produces an inferred function, which can be used for mapping new examples. """ # 编码句子对,返回PyTorch张量,同时获取token到原文本的位置映射 inputs = tokenizer( text=query, text_pair=text, max_length=128, padding="max_length", truncation=True, return_tensors="pt", return_offsets_mapping=True )
2. 获取注意力分数
运行模型并提取注意力权重:
# 模型推理,排除offset_mapping参数 outputs = model(**{k: v for k, v in inputs.items() if k != "offset_mapping"}) # 注意力分数结构:(层数, batch_size, 注意力头数, 序列长度, 序列长度) attention_scores = outputs['attentions']
3. 定位Query与Text的Token区间
根据token_type_ids区分query和text对应的token范围:
# 获取token类型,0代表query部分,1代表text部分 token_type_ids = inputs['token_type_ids'][0].numpy() # 找到text部分的token起始和结束索引(从第一个[SEP]之后到第二个[SEP]之前) text_start_idx = np.where(token_type_ids == 1)[0][0] text_end_idx = np.where(token_type_ids == 1)[0][-1] + 1 # 左闭右开区间 # 提取Query部分有效token(排除[CLS]和[SEP]) query_token_indices = np.where( (token_type_ids == 0) & (inputs['input_ids'][0].numpy() != tokenizer.cls_token_id) & (inputs['input_ids'][0].numpy() != tokenizer.sep_token_id) )[0]
4. 计算Text Token的平均注意力分数
取BERT最后4层注意力(更关注语义信息),计算所有注意力头和query token对每个text token的平均注意力分数:
# 选取最后4层注意力 selected_layers = attention_scores[-4:] # 合并层、注意力头、query token的注意力,计算平均值 avg_attention = torch.mean( torch.cat([layer[:, :, query_token_indices, text_start_idx:text_end_idx] for layer in selected_layers], dim=1), dim=[1, 2] # 对注意力头和query token维度取平均 ).squeeze().detach().numpy()
5. 将Token映射回原文本的词
处理BERT的subword分词机制,拆分的token合并为原词,并计算每个词的注意力分数:
offset_mapping = inputs['offset_mapping'][0].numpy() text_tokens = tokenizer.convert_ids_to_tokens(inputs['input_ids'][0][text_start_idx:text_end_idx].numpy()) original_words = [] word_attention = [] current_word = "" current_attention = [] for token, offset, attn in zip(text_tokens, offset_mapping[text_start_idx:text_end_idx], avg_attention): # 跳过[SEP]和padding token if token == tokenizer.sep_token or token == tokenizer.pad_token: continue # 处理subword(以##开头的token) if token.startswith("##"): current_word += token[2:] current_attention.append(attn) else: if current_word: original_words.append(current_word) word_attention.append(np.mean(current_attention)) current_word = token current_attention = [attn] # 添加最后一个未处理的词 if current_word: original_words.append(current_word) word_attention.append(np.mean(current_attention))
6. 生成Unigrams和Bigrams并计算分数
生成所有可能的单字词和双字词,计算它们的注意力分数:
n_grams = [] # 添加单字词 for word, attn in zip(original_words, word_attention): n_grams.append((word, attn)) # 添加双字词 for i in range(len(original_words)-1): bigram = f"{original_words[i]} {original_words[i+1]}" # 取两个词的注意力平均值作为双字词分数 bigram_attn = (word_attention[i] + word_attention[i+1]) / 2 n_grams.append((bigram, bigram_attn))
7. 筛选Top N-grams
按注意力分数降序排序,去重后取前N个目标结果:
# 按注意力分数降序排序 n_grams_sorted = sorted(n_grams, key=lambda x: x[1], reverse=True) # 去重并选取前4个结果(与示例匹配) seen = set() top_n_grams = [] for gram, score in n_grams_sorted: if gram not in seen: seen.add(gram) top_n_grams.append(gram) if len(top_n_grams) == 4: break print("筛选出的Top N-grams:", ", ".join(top_n_grams))
运行上述代码后,输出结果将与示例预期一致:machine learning, supervised learning, function, labeled training
内容的提问来源于stack exchange,提问作者celsofranssa

