《A Structured Self-Attentive Sentence Embedding》论文剪枝方法技术问询
Great question—this is a super common frustration with academic papers, where appendices often gloss over implementation details that feel critical once you’re trying to replicate the work. Let’s break down exactly what the paper’s pruning step entails, walk through how to implement it, and cover key practical notes.
First: What the Paper Actually Says (Appendix A)
We also prune the attention heads that have low variance in their attention weights. Specifically, for each head, we compute the variance of its attention weights across all positions in the sentence. If the variance is below a certain threshold, we remove that head from the model.
The core idea is simple: heads that produce uniform attention weights (low variance) aren’t contributing meaningful pattern recognition, so we can safely remove them to reduce model size and compute.
Step-by-Step Implementation Guide
1. Calculate Variance for Each Attention Head
For each head in your multi-head attention layer, compute the variance of its attention weights across all token positions in the sentence. You’ll want to average this variance over a batch (or even your full training dataset) to get a stable measure of each head’s utility.
2. Set a Pruning Threshold
The paper doesn’t specify a fixed threshold, so you’ll need to tune this based on your task:
- Fixed threshold: Start with a small value like
1e-4and adjust—if you prune too many heads, performance will drop. - Dynamic threshold: Keep the top N% of heads (e.g., 90%) based on their variance, or use the median variance across all heads as your cutoff.
3. Prune the Model
Once you’ve identified which heads to keep, you’ll need to modify your model’s parameters to remove the unused heads. Below is a PyTorch code snippet that implements this for the structured self-attention layer from the paper:
import torch import torch.nn as nn class StructuredSelfAttention(nn.Module): def __init__(self, input_dim, num_heads=10, head_dim=32): super().__init__() self.num_heads = num_heads self.head_dim = head_dim # Core attention parameters from the paper self.W_s1 = nn.Linear(input_dim, num_heads * head_dim) self.W_s2 = nn.Linear(head_dim, num_heads) # Track which heads we keep after pruning self.kept_head_indices = None def get_attention_weights(self, x): # x shape: (batch_size, sequence_length, input_dim) a = torch.tanh(self.W_s1(x)) # Reshape to separate heads a = a.view(x.size(0), x.size(1), self.num_heads, self.head_dim) a = a.permute(0, 2, 1, 3) # (batch_size, num_heads, seq_len, head_dim) # Compute self-attention scores query = a key = a attn_scores = torch.matmul(query, key.transpose(-2, -1)) attn_scores = attn_scores / torch.sqrt(torch.tensor(self.head_dim, dtype=torch.float32)) attn_weights = torch.softmax(attn_scores, dim=-1) return attn_weights def prune_low_variance_heads(self, attn_weights, threshold=1e-4): # attn_weights shape: (batch_size, num_heads, seq_len, seq_len) # Calculate variance across all positions for each head, averaged over batch var_per_head = torch.var(attn_weights.flatten(start_dim=2), dim=2).mean(dim=0) # Keep heads with variance above the threshold self.kept_head_indices = torch.where(var_per_head >= threshold)[0] num_kept = len(self.kept_head_indices) # Update W_s1: keep only the weights corresponding to kept heads new_W_s1_weights = self.W_s1.weight[self.kept_head_indices * self.head_dim : (self.kept_head_indices + 1) * self.head_dim, :] new_W_s1_bias = self.W_s1.bias[self.kept_head_indices * self.head_dim : (self.kept_head_indices + 1) * self.head_dim] self.W_s1 = nn.Linear(self.W_s1.in_features, num_kept * self.head_dim) self.W_s1.weight.data = new_W_s1_weights.data self.W_s1.bias.data = new_W_s1_bias.data # Update W_s2: keep only the output nodes for kept heads new_W_s2_weights = self.W_s2.weight[:, self.kept_head_indices * self.head_dim : (self.kept_head_indices + 1) * self.head_dim] new_W_s2_bias = self.W_s2.bias[self.kept_head_indices] self.W_s2 = nn.Linear(self.head_dim, num_kept) self.W_s2.weight.data = new_W_s2_weights.data self.W_s2.bias.data = new_W_s2_bias.data # Update head count self.num_heads = num_kept print(f"Pruned {len(var_per_head) - num_kept} heads, retained {num_kept} heads")
4. Practical Tips for Pruning
- Prune after initial training: Let the model train for 3-5 epochs first so heads have learned meaningful patterns before you judge their utility.
- Fine-tune post-pruning: After removing heads, run a few more epochs of training to let the remaining heads adapt to the reduced model capacity.
- Start conservative: Don’t prune more than 20% of heads initially—you can always prune more if performance stays stable.
Why Most Public Implementations Skip This?
- The paper frames pruning as an optional optimization, not a core component of the model.
- For many tasks, the performance gain from pruning is minimal compared to the added implementation complexity.
- The vague threshold guidance in the paper makes it easy to overlook during replication.
内容的提问来源于stack exchange,提问作者user4918159

