Skip to content

Instantly share code, notes, and snippets.

@peterdresslar
Created April 15, 2026 19:45
Show Gist options
  • Select an option

  • Save peterdresslar/24e85de5b31e0c32aaf92b12df3fa9ae to your computer and use it in GitHub Desktop.

Select an option

Save peterdresslar/24e85de5b31e0c32aaf92b12df3fa9ae to your computer and use it in GitHub Desktop.
Quick example of attention head entropy
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