Created
April 15, 2026 19:45
-
-
Save peterdresslar/24e85de5b31e0c32aaf92b12df3fa9ae to your computer and use it in GitHub Desktop.
Quick example of attention head entropy
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
| import torch | |
| attentions = [ | |
| # ========================== LAYER 0 ========================== | |
| torch.tensor([ | |
| [ # <--- Start of Batch Item 0 | |
| # --------------------------------------------------- | |
| # HEAD 0 | |
| # A "focused" head: The last token looks mostly at "cat" | |
| # --------------------------------------------------- | |
| [ # Key tokens: | |
| # "The" "cat" "sat" | |
| [ 1.00, 0.00, 0.00 ], # Query: "The" (can only see itself) | |
| [ 0.70, 0.30, 0.00 ], # Query: "cat" (can see "The", "cat") | |
| [ 0.10, 0.80, 0.10 ] # Query: "sat" (can see all 3) | |
| # ^^^^^^^^^^^^^^^^^^ <--- THIS row is `attn[0, 0, -1, :]` | |
| ], | |
| # --------------------------------------------------- | |
| # HEAD 1 | |
| # A "diffuse" head: The last token looks at everything equally | |
| # --------------------------------------------------- | |
| [ # Key tokens: | |
| # "The" "cat" "sat" | |
| [ 1.00, 0.00, 0.00 ], # Query: "The" | |
| [ 0.50, 0.50, 0.00 ], # Query: "cat" | |
| [ 0.33, 0.33, 0.34 ] # Query: "sat" | |
| # ^^^^^^^^^^^^^^^^^^ <--- THIS row is `attn[0, 1, -1, :]` | |
| ] | |
| ] | |
| ]) | |
| ] | |
| def get_attn_head_entropy(attentions): | |
| """ | |
| Compute the Shannon entropy of the attention weights for the final token across all heads, | |
| and return a dictionary with the entropy for each head. | |
| Example output (for a 3-layer model with 4 heads per layer): | |
| { | |
| "attn_entropy_per_head_final": [ | |
| [0.1, 0.2, 0.3, 0.4], # Layer 0 | |
| [0.5, 0.6, 0.7, 0.8], # Layer 1 | |
| [0.9, 1.0, 1.1, 1.2] # Layer 2 | |
| ] | |
| } | |
| Args: | |
| attentions (list): A list of attention weights for each layer. | |
| Returns: | |
| dict: A dictionary with the entropy for each head. | |
| """ | |
| result = {"attn_entropy_per_head_final": []} | |
| for layer_pos, attn in enumerate(attentions): | |
| last_token_attn = attn[0, :, -1, :].float() # we just want the last token: slice to get it from tensor | |
| # Compute Shannon entropy: H(X) = -Sum(P(x) * log2(P(x))) | |
| # PyTorch is a bit like numpy (for gpus). Here we use a torch function, log2 | |
| # Since it is a torch function it applies to the *entire* tensor slice at once. | |
| ent = -(last_token_attn * torch.log2(last_token_attn + 1e-10)).sum(dim=-1) | |
| # And we're done with this layer. We simply append the result to our entropy list | |
| # (see example output) | |
| # In our Qwen example this would actually be 16 head entropy values, | |
| # and our list would have 8 layers --> 128 values | |
| result["attn_entropy_per_head_final"].append(ent.tolist()) | |
| return result |
Sign up for free
to join this conversation on GitHub.
Already have an account?
Sign in to comment