Diffusion Model Attention Heatmaps: Unlocking Interpretability in Text-to-Image Generation
This article introduces three complementary attention heatmap visualizations — image-to-text, text-to-image, and image-to-image — for interpreting text-to-image diffusion models, detailing aggregated attention computation, bilinear interpolation rendering, and an interactive Flask web service to diagnose prompt adherence and model biases.
Computing Aggregated Attention Score Matrices
Text-to-image diffusion models typically consist of three components: a CLIP text encoder that converts prompts into token embeddings, a U-Net that progressively denoises latent images, and a VAE that decodes the final latent representation into pixels. Model interpretability aims to explain why a trained network produces a given output, helping locate attribute binding errors, missing objects, and revealing what the model focuses on, thereby surfacing hidden biases.
Attention is well-suited for this task because most hidden activations are hard to interpret directly, while attention itself is a normalized, human-readable "what attends to what" relationship map that provides rare insight into Transformer inference at minimal cost.
The analysis starts from attention probability tensors. In each attention layer of the U-Net, queries Q, keys K, and values V are constructed, and attention is computed as: A = softmax(Q Kᵀ / √d) where d is the per-head dimension and softmax is applied along the key axis. Modern diffusers implementations compute softmax inside fused kernels and do not expose the attention matrix A; therefore, attention must be recomputed classically to retain the matrices. The pre-softmax logits S = Q Kᵀ / √d are also saved, as they record raw scores before normalization.
Two forms of attention exist:
Self-attention : both queries and keys come from the image latent representation. A maps each patch to a distribution over other patches, yielding a P×P matrix.
Cross-attention : queries come from image latents, keys from CLIP text tokens. A maps P image patches to distributions over T tokens, yielding a P×T matrix.
Single-step, single-layer results are noisy, so averaging is required. The U-Net computes attention at multiple feature-map resolutions; only matrices of the same size can be averaged together. Aggregation is performed per resolution r, averaging over attention heads H, denoising steps 𝒯, and layers Lᵣ at that resolution:
Āᵣ = (1 / |H|·|𝒯|·|Lᵣ|) Σ_{h∈H, t∈𝒯, l∈Lᵣ} A_{h,t,l}For cross-attention, Āᵣ (and similarly S̄ᵣ) has shape r²×T; for self-attention, shape r²×r². Different resolutions are not merged; a specific resolution r is selected for downstream analysis.
Method 1: Image-to-Text
Fix an image patch p₀ and read the corresponding row of Āᵣ to obtain the distribution over text tokens for that location. Only true word tokens are retained. The distribution is then normalized across word tokens and each token label is colored by its weight. This answers: which words did the model reference when rendering this patch?
Example with prompt "A cat sitting on a red sofa": clicking the cat's chest highlights "cat" as the hottest token; clicking the sofa cushion highlights "sofa"; clicking the cat's lower body near the sofa edge still shows high heat for "cat" and "sitting" even though the click lands on the cushion. Token colors indicate influence strength on the selected region, making attribute binding clearly visible.
Method 2: Text-to-Image
Fix a token t₀ and scan all image patches along its column — effectively a transpose of Method 1. Here the pre-softmax logits S̄ᵣ are used instead of Āᵣ because softmax normalizes per query patch; with a fixed token, raw logits better compare different patches. The resulting slice holds one value per latent unit, reshaped to an r×r grid. After min-max normalization to [0,1], bilinear interpolation upsamples to the image resolution N×N, and the heatmap is blended onto the image. This produces the familiar "word 'cat' lights up the cat" heatmaps. Token attention concentrates in local spatial regions, making results easy to interpret.
Example: token "cat" highlights the cat's body and outline; token "red" lights up the sofa area; token "sitting" covers the cat and the sofa surface beneath it.
Method 3: Image-to-Image
Self-attention operates purely within the image. Select a query patch p₀, read its row from the self-attention map, and map it back to image space. The slice is again one value per latent unit, rendered as a spatial heatmap exactly like Method 1. This reveals the model's internal structural judgments: clicking one eye often lights up the other eye; clicking a wall lights up other wall regions while the cat and sofa stay cool. Self-attention groups patches belonging to the same object or texture.
Example: selecting a background point illuminates the entire wall or background, while cat and sofa remain low; selecting one eye lights up both eyes; selecting one ear lights up both ears.
Rendering Heatmaps
All numerical values — scores, per-resolution averages Āᵣ and S̄ᵣ, and min-max normalized results — are computed at the native latent resolution without any scaling. The upsampling pipeline used for text-to-image and image-to-image views works as follows: the slice g (one value per latent unit) is arranged into an r×r grid M. To overlay M onto the full N×N image, bilinear interpolation enlarges M by factor s = N/r. Each output pixel (i, j) maps back to (i/s, j/s) in the small grid, and the four surrounding cells are blended with weights derived from the fractional parts fᵢ = i/s − ⌊i/s⌋ and fⱼ = j/s − ⌊j/s⌋. The four weights sum to 1, so each pixel is a convex combination of the four nearest cells, then colored with a colormap. Upsampling adds no detail; it only smooths between cell centers. Low-resolution layers produce softer, blockier maps; high-resolution layers yield sharper, edge-aligned results.
Web Service
A small Flask application enables interactive exploration. In Image → Text mode, clicking an image location recolors token labels by attention weight. In Text → Image mode, clicking a token label displays its heatmap. In Image → Image mode, clicking a location highlights related regions.
Summary
Making these visualizations visible requires modest effort but yields concrete diagnostic evidence. When prompt generation fails — whether colors bind to wrong objects or required objects are missing — heatmaps point directly to the responsible tokens or regions, turning vague "generation went wrong" into specific, actionable issues. They also expose what information the model relies on. Before deploying generative models to users, such intuitive, inspectable evidence forms a foundation for trust.
Code repository:
https://github.com/Oxotall/Text2ImageAttnVisualizationSigned-in readers can open the original source through BestHub's protected redirect.
This article has been distilled and summarized from source material, then republished for learning and reference. If you believe it infringes your rights, please contactand we will review it promptly.
Data Party THU
Official platform of Tsinghua Big Data Research Center, sharing the team's latest research, teaching updates, and big data news.
How this landed with the community
Was this worth your time?
0 Comments
Thoughtful readers leave field notes, pushback, and hard-won operational detail here.
