Interleaved Head Attention (IHA) in C++ (LibTorch)
May 2026 · 16 min read
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:
Model expressiveness: Different heads specialize in different receptive fields.
Sparse attention: Mix cheap local (sliding-window) with expensive global heads.
Hybrid generation: Causal heads predict next token; global heads aggregate history.
2 · Attention Mode Enum
Define an enum for attention patterns:
enum classAttentionMode {
Causal, // Standard autoregressive: attend to past onlySliding, // Local window: attend to ±window_size tokensGlobal, // 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
structInterleavedHeadAttentionImpl : publictorch::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::Tensorbuild_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_sizefor (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::Tensorforward(torch::Tensor x) {
int batch = x.size(0), time = x.size(1);
// Project Q, K, Vtorch::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 maskingtorch::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 reshapetorch::Tensor out = out_list.transpose(1, 2);
out = out.contiguous().view({batch, time, d_model});
returnW_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: causalfor (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: globalfor (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:
Strided attention: Add a mode that attends every k-th token (sparse sampling).
Learnable mode selection: Replace fixed modes with learned gating: each head softly selects which pattern to apply.
Dynamic masking: Let the mask depend on input (learned attention sparsity patterns).
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:
Self-Attention (Part 1)
Softmax Stabilization (Part 2)
Causal Masking (Part 3)
Cross-Attention (Part 4)
Sliding-Window Attention (Part 5)
Global (All-to-All) Attention (Part 6)
Linear Attention (Part 7)
FlashAttention (Part 8)
Multi-Head Attention (Part 9)
Multi-Query Attention (Part 10)
Grouped-Query Attention (Part 11)
Multi-Head Latent Attention (Part 12)
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!