如何调整代码实现BERT模型Token注意力权重对齐与可视化
问题
我基于Transformer模型(PyTorch+HuggingFace,GPU运行)进行文本分类任务,模型及训练循环已正常运行。为分析误分类样本的预测逻辑,我希望查看误分类样本中各Token对应的注意力权重,但当前实现存在两个问题:
- 每个误分类样本输出的权重数量冗余(推测是输出了所有Transformer层的权重)
- 注意力权重未与对应Token关联,无法直观解释,期望实现按权重深浅为Token着色的可视化效果
现有核心代码片段如下:
model = AutoModelForSequenceClassification.from_pretrained(checkpoint, num_labels=len(label_dict), output_attentions=True, output_hidden_states=True)
def get_wrong_predictions(predictions, true_vals): wrong_predictions = [] for i in range(len(true_vals)): if np.argmax(predictions[i]) != true_vals[i]: wrong_predictions.append((true_vals[i], np.argmax(predictions[i]))) return wrong_predictions
seed_val = 17 random.seed(seed_val) np.random.seed(seed_val) torch.manual_seed(seed_val) torch.cuda.manual_seed_all(seed_val) def evaluate(dataloader_val): model.eval() loss_val_total = 0 predictions, true_vals = [], [] attentions = [] for batch in dataloader_val: batch = tuple(b.to(device) for b in batch) inputs = {'input_ids': batch[0], 'attention_mask': batch[1], 'labels': batch[2], } with torch.no_grad(): outputs = model(**inputs) loss = outputs[0] logits = outputs[1] attention_scores = outputs.attentions[-1].detach().cpu().numpy() # extract the attention scores from the last layer loss_val_total += loss.item() logits = logits.detach().cpu().numpy() label_ids = inputs['labels'].cpu().numpy() predictions.append(logits) true_vals.append(label_ids) attentions.append(attention_scores) loss_val_avg = loss_val_total/len(dataloader_val) predictions = np.concatenate(predictions, axis=0) true_vals = np.concatenate(true_vals, axis=0) attentions = np.concatenate(attentions, axis=0) preds = np.argmax(predictions, axis=1) misclassified = np.where(preds != true_vals)[0] texts = df["description"].tolist() global misclassified_examples misclassified_examples = [] for idx in misclassified[:min(len(attention_scores), len(misclassified))]: text = texts[idx] true_label = true_vals[idx] pred_label = preds[idx] attention_weights = attentions[idx] # get the attention weights for the misclassified examples misclassified_examples.append({ 'text': text, 'true_label': true_label, 'pred_label': pred_label, 'attention_weights' : attention_weights }) return loss_val_avg, predictions, true_vals train_losses = [] #to plot later val_losses = [] overall_accuracy = [] for epoch in tqdm(range(1, epochs+1)): model.train() loss_train_total = 0 progress_bar = tqdm(dataloader_train, desc='Epoch {:1d}'.format(epoch), leave=False, disable=False) for batch in progress_bar: model.zero_grad() batch = tuple(b.to(device) for b in batch) inputs = {'input_ids': batch[0], 'attention_mask': batch[1], 'labels': batch[2], } outputs = model(**inputs) loss = outputs[0] loss_train_total += loss.item() loss.backward() torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0) optimizer.step() scheduler.step() progress_bar.set_postfix({'training_loss': '{:.3f}'.format(loss.item()/len(batch))}) loss_train_avg = loss_train_total/len(dataloader_train) #to plot later train_losses.append(loss_train_avg) # to plot later tqdm.write(f'\nEpoch {epoch}') loss_train_avg = loss_train_total/len(dataloader_train) tqdm.write(f'Training loss: {loss_train_avg}') val_loss, predictions, true_vals = evaluate(dataloader_validation) val_f1 = f1_score_func(predictions, true_vals) tqdm.write(f'Validation loss: {val_loss}') tqdm.write(f'F1 Score (Weighted): {val_f1}') val_overallaccuracy = overallaccuracy(true_vals, predictions) tqdm.write(f'Overall Accuracy: {val_overallaccuracy}') val_mcc = mcc_score_func(true_vals, predictions) tqdm.write(f'MCC: {val_mcc}') val_acc = accuracy_per_class(predictions, true_vals) tqdm.write(f'Accuracy per class: {val_acc}') val_losses.append(val_loss) overall_accuracy.append(val_overallaccuracy)
解决方案
1. 精简注意力权重(解决冗余问题)
Transformer的注意力输出包含多层多头的权重(形状为[样本数, 层数, 头数, 序列长度, 序列长度]),你当前只取了最后一层,但仍保留了多头维度。针对分类任务,通常可以:
- 取
[CLS]token(序列第一个位置)对其他所有token的注意力(因为分类任务的最终决策依赖[CLS]的表征) - 对多头注意力取平均值,得到每个token对应的单一权重值
2. 关联Token与注意力权重
需要用HuggingFace的Tokenizer将input_ids转换为可读的Token,同时过滤掉padding token(避免无效权重干扰)。
3. 代码调整细节
步骤1:保存Input IDs用于Token转换
修改evaluate函数,新增存储input_ids的列表,后续用于映射为Token:
def evaluate(dataloader_val): model.eval() loss_val_total = 0 predictions, true_vals = [], [] attentions = [] input_ids_list = [] # 新增:保存所有样本的input_ids for batch in dataloader_val: batch = tuple(b.to(device) for b in batch) inputs = {'input_ids': batch[0], 'attention_mask': batch[1], 'labels': batch[2], } with torch.no_grad(): outputs = model(**inputs) loss = outputs[0] logits = outputs[1] # 取最后一层的注意力权重,形状:[batch_size, num_heads, seq_len, seq_len] attention_scores = outputs.attentions[-1].detach().cpu().numpy() loss_val_total += loss.item() logits = logits.detach().cpu().numpy() label_ids = inputs['labels'].cpu().numpy() predictions.append(logits) true_vals.append(label_ids) attentions.append(attention_scores) input_ids_list.append(inputs['input_ids'].detach().cpu().numpy()) # 保存当前batch的input_ids # 拼接所有batch的数据 predictions = np.concatenate(predictions, axis=0) true_vals = np.concatenate(true_vals, axis=0) attentions = np.concatenate(attentions, axis=0) input_ids_list_flattened = np.concatenate(input_ids_list, axis=0) # 拼接所有样本的input_ids preds = np.argmax(predictions, axis=1) misclassified = np.where(preds != true_vals)[0] texts = df["description"].tolist() # 提前加载Tokenizer(确保和模型使用的checkpoint一致) tokenizer = AutoTokenizer.from_pretrained(checkpoint) global misclassified_examples misclassified_examples = [] for idx in misclassified: text = texts[idx] true_label = true_vals[idx] pred_label = preds[idx] # 1. 转换Input IDs为Token,并过滤padding input_ids = input_ids_list_flattened[idx] tokens = tokenizer.convert_ids_to_tokens(input_ids) valid_indices = input_ids != tokenizer.pad_token_id # 过滤padding token filtered_tokens = [tokens[i] for i in range(len(tokens)) if valid_indices[i]] # 2. 处理注意力权重:取[CLS]对所有token的注意力,多头取平均 attention_weights = attentions[idx] # 形状:[num_heads, seq_len, seq_len] # 取[CLS]位置(索引0)对所有token的注意力,形状:[num_heads, seq_len] cls_attentions = attention_weights[:, 0, :] # 对多头取平均,得到每个token的权重,形状:[seq_len] avg_cls_attentions = np.mean(cls_attentions, axis=0) # 过滤padding对应的权重 filtered_weights = avg_cls_attentions[valid_indices] misclassified_examples.append({ 'text': text, 'true_label': true_label, 'pred_label': pred_label, 'tokens': filtered_tokens, 'attention_weights': filtered_weights }) return loss_val_avg, predictions, true_vals
步骤2:实现Token着色可视化
添加一个生成HTML可视化的函数,根据权重深浅为Token添加背景色:
def visualize_token_attention(tokens, weights): # 归一化权重到0-1范围,方便映射颜色深浅 norm_weights = (weights - weights.min()) / (weights.max() - weights.min()) html_content = '<div style="font-family: monospace; line-height: 1.8;">' for token, weight in zip(tokens, weights): # 用蓝色系表示权重,权重越高颜色越深 bg_color = f'rgba(100, 180, 255, {weight})' html_content += f'<span style="background-color: {bg_color}; padding: 3px 6px; margin: 0 2px; border-radius: 3px;">{token}</span>' html_content += '</div>' return html_content
使用示例:
# 取第一个误分类样本 example = misclassified_examples[0] visual_html = visualize_token_attention(example['tokens'], example['attention_weights']) # 可以将HTML保存到文件查看,或者在Notebook中直接显示 with open('attention_visual.html', 'w', encoding='utf-8') as f: f.write(visual_html)
关键说明
- 若你想分析特定注意力头的权重,可跳过多头平均步骤,直接选取对应头的权重(比如
cls_attentions[head_idx, :]) - 部分模型的Tokenizer可能会产生子词(比如BERT的
##xxx),可视化时可以保留原始形式,或合并为完整单词(需额外处理)
内容的提问来源于stack exchange,提问作者Lorra
相关产品推荐
相关产品推荐

