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

如何调整代码实现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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.01 22:16:17