← All Posts
From Scratch · C++ · Transformers · Attention Variants· Part 13 of 13

Interleaved Head Attention (IHA) in C++ (LibTorch)

Interleaved Head Attention (IHA) is a flexible design pattern where each attention head can operate in a different mode: causal (autoregressive), sliding-window (local), or global (all-to-all). This enables models to mix attention patterns within a single layer—some heads attend causally, others see a local context, and a few act as global summarizers. We implement this in C++ via an enum-based router and per-head masking.

1 · Motivation: Multi-Pattern Attention

Standard transformers apply the same attention pattern (e.g., causal or unrestricted) to all heads. IHA decouples this: each head can choose its connectivity. Benefits:

2 · Attention Mode Enum

Define an enum for attention patterns:

enum class AttentionMode { Causal, // Standard autoregressive: attend to past only Sliding, // Local window: attend to ±window_size tokens Global, // All-to-all: unrestricted };

Each head in the layer has an associated mode. During forward, we apply the corresponding mask before softmax.

3 · LibTorch Implementation

struct InterleavedHeadAttentionImpl : public torch::nn::Module { int n_heads, d_model, d_k, window_size; std::vector<AttentionMode> head_modes; torch::nn::Linear W_q{nullptr}, W_k{nullptr}, W_v{nullptr}, W_o{nullptr}; InterleavedHeadAttentionImpl(int h, int d, int ws, std::vector<AttentionMode> modes) : n_heads(h), d_model(d), d_k(d / h), window_size(ws), head_modes(modes) { TORCH_CHECK(modes.size() == h, "One mode per head"); W_q = register_module("W_q", torch::nn::Linear(d, d)); W_k = register_module("W_k", torch::nn::Linear(d, d)); W_v = register_module("W_v", torch::nn::Linear(d, d)); W_o = register_module("W_o", torch::nn::Linear(d, d)); } torch::Tensor build_mask(AttentionMode mode, int time) { torch::Tensor mask = torch::zeros({time, time}); if (mode == AttentionMode::Causal) { // Lower triangular: (i, j) valid iff j <= i mask = torch::triu(torch::ones({time, time}), 1) * -1e9; } else if (mode == AttentionMode::Sliding) { // Sliding window: attend within ±window_size for (int i = 0; i < time; i++) { for (int j = 0; j < time; j++) { if (std::abs(i - j) > window_size) { mask[i][j] = -1e9; } } } } // Global: no mask (all zeros) return mask; } torch::Tensor forward(torch::Tensor x) { int batch = x.size(0), time = x.size(1); // Project Q, K, V torch::Tensor Q = W_q->(forward)(x).view({batch, time, n_heads, d_k}).transpose(1, 2); torch::Tensor K = W_k->(forward)(x).view({batch, time, n_heads, d_k}).transpose(1, 2); torch::Tensor V = W_v->(forward)(x).view({batch, time, n_heads, d_k}).transpose(1, 2); // Compute attention per head with mode-specific masking torch::Tensor out_list = torch::zeros({batch, n_heads, time, d_k}); for (int h = 0; h < n_heads; h++) { torch::Tensor Q_h = Q.select(1, h); // (batch, time, d_k) torch::Tensor K_h = K.select(1, h); torch::Tensor V_h = V.select(1, h); torch::Tensor scores = torch::matmul(Q_h, K_h.transpose(-1, -2)) / std::sqrt(d_k); torch::Tensor mask = build_mask(head_modes[h], time).to(scores.device()); scores = scores + mask.unsqueeze(0); // (batch, time, time) scores = torch::softmax(scores, -1); torch::Tensor head_out = torch::matmul(scores, V_h); // (batch, time, d_k) out_list.select(1, h).copy_(head_out); } // Transpose and reshape torch::Tensor out = out_list.transpose(1, 2); out = out.contiguous().view({batch, time, d_model}); return W_o->(forward)(out); } }; TORCH_MODULE(InterleavedHeadAttention);

4 · Mask Patterns

Mode Pattern Sparsity Use Case
Causal Lower triangular: (i, j) valid iff j ≤ i ~50% non-zero Autoregressive generation
Sliding Banded: attend within ±window_size O(seq × window) Local context (long sequences)
Global Dense: all entries non-zero 100% dense Summary/global reasoning

5 · Per-Head Mode Configuration

Example: 12 heads, 4 in each mode.

std::vector<AttentionMode> modes; // Heads 0–3: causal for (int i = 0; i < 4; i++) modes.push_back(AttentionMode::Causal); // Heads 4–7: sliding window (local) for (int i = 0; i < 4; i++) modes.push_back(AttentionMode::Sliding); // Heads 8–11: global for (int i = 0; i < 4; i++) modes.push_back(AttentionMode::Global); auto iha_layer = std::make_shared<InterleadedHeadAttention>( 12, // n_heads 768, // d_model 64, // window_size for sliding modes );

6 · Head Mode Visualization

Interleaved Head Attention: Mixed Modes

Step through heads, each with its own attention pattern. Green = causal, blue = sliding, orange = global.

7 · Computational & Memory Benefits

For a sequence of length N with h heads, mix modes as follows:

Dense global heads: O(h_global × N²) Sliding window heads: O(h_sliding × N × window_size) Causal heads: O(h_causal × N²) [but triangular, slightly cheaper] Example (seq=4096, h=12, window=256): All causal: 12 × 4096² / 2 ≈ 100M ops 4 causal + 4 sliding + 4 global: 4 × 4096² / 2 + 4 × 4096 × 256 + 4 × 4096² ≈ 33M + 4M + 67M ≈ 104M ops Mixed modes offer modest speedup but great flexibility for task-specific tuning.

8 · Advanced Extensions

You can extend IHA further:

Practical note: While IHA is flexible, most production transformers use uniform attention (all MHA). IHA shines in specialized architectures: document long-range reasoning (global + sliding), hierarchical models (causal for local, global for inter-block).

9 · Series Conclusion

We've now covered the full spectrum of attention mechanisms from scratch in C++ using LibTorch:

  1. Self-Attention (Part 1)
  2. Softmax Stabilization (Part 2)
  3. Causal Masking (Part 3)
  4. Cross-Attention (Part 4)
  5. Sliding-Window Attention (Part 5)
  6. Global (All-to-All) Attention (Part 6)
  7. Linear Attention (Part 7)
  8. FlashAttention (Part 8)
  9. Multi-Head Attention (Part 9)
  10. Multi-Query Attention (Part 10)
  11. Grouped-Query Attention (Part 11)
  12. Multi-Head Latent Attention (Part 12)
  13. Interleaved Head Attention (Part 13)

From basic dot-product attention to advanced latent compression and flexible per-head routing, you now have a mental model and implementation reference for every major attention variant used in modern transformers. Build on these foundations to create novel architectures!