Papers
Topics
Authors
Recent
Search
2000 character limit reached

FlashMHF: Multi-Head FFN for Transformers

Updated 1 March 2026
  • FlashMHF is a multi-head feed-forward network that replaces standard FFNs in Transformers using dynamic routing and fused on-chip compute kernels.
  • It leverages dynamically weighted parallel sub-networks to balance intermediate and head dimensions, ensuring scalable expressivity and reduced memory usage.
  • Empirical evaluations show FlashMHF achieves lower perplexity, improved accuracy, faster inference, and greater hardware efficiency compared to standard SwiGLU.

Flash Multi-Head Feed-Forward Network (FlashMHF) is an architectural replacement for the position-wise feed-forward network (FFN) used in Transformer models. Motivated by the structural similarity between single-head attention and FFNs, and inspired by the benefits of multi-head mechanisms in increasing representational expressivity, FlashMHF introduces a multi-head design for FFNs. Two central innovations set FlashMHF apart: (1) an I/O-aware, fused compute kernel that computes outputs online in on-chip SRAM, and (2) a design utilizing dynamically weighted parallel sub-networks within each head to maintain a balanced ratio between intermediate and head dimensions as models scale. FlashMHF achieves improved perplexity, accuracy, and hardware efficiency while reducing peak activation memory and accelerating inference in practical settings (Zhang et al., 7 Dec 2025).

1. Motivation and Background

In the Transformer architecture, each layer includes an FFN after multi-head self-attention, usually expressed as

FFN(X)=W2 ϕ(W1X+b1)+b2,\mathrm{FFN}(X) = W_2\,\phi(W_1 X + b_1) + b_2,

where X∈RL×dmodelX \in \mathbb{R}^{L \times d_\mathrm{model}}, ϕ\phi is a nonlinearity, W1∈Rdf×dmodelW_1 \in \mathbb{R}^{d_f \times d_\mathrm{model}}, W2∈Rdmodel×dfW_2 \in \mathbb{R}^{d_\mathrm{model} \times d_f}. Geva et al. (2020) observed the computation can be reinterpreted as attention over parameters. The analogy to attention motivates extending feed-forward layers with a multi-head mechanism to enhance expressivity, expecting the decomposition of computation into parallel subspaces to yield richer modeling power.

However, a naïve multi-head FFN (MH-FFN) approach that replicates SwiGLU-style intermediate expansions headwise runs into two critical scaling issues:

  • Memory overhead scales with the product H×(L⋅df)H \times (L \cdot d_f), where HH is the number of heads and dfd_f the intermediate size.
  • As model scale increases and dfd_f grows (following scaling laws), the per-head dimension dhd_h remains fixed, leading to an imbalanced and inefficient ratio X∈RL×dmodelX \in \mathbb{R}^{L \times d_\mathrm{model}}0 that degrades both scalability and expressivity (Zhang et al., 7 Dec 2025).

2. Formal Construction and Mechanism

2.1 Standard FFN and Naïve MH-FFN

The classical FFN applies the transformation per token: X∈RL×dmodelX \in \mathbb{R}^{L \times d_\mathrm{model}}1 A naïve MH-FFN first projects the input to X∈RL×dmodelX \in \mathbb{R}^{L \times d_\mathrm{model}}2 heads (each dimension X∈RL×dmodelX \in \mathbb{R}^{L \times d_\mathrm{model}}3), then applies independently-parameterized intermediate expansions (e.g., SwiGLU) to each, and finally concatenates the outputs. Formally:

  • Input projection: X∈RL×dmodelX \in \mathbb{R}^{L \times d_\mathrm{model}}4.
  • Each head uses independent SwiGLU parameters X∈RL×dmodelX \in \mathbb{R}^{L \times d_\mathrm{model}}5, X∈RL×dmodelX \in \mathbb{R}^{L \times d_\mathrm{model}}6, X∈RL×dmodelX \in \mathbb{R}^{L \times d_\mathrm{model}}7.
  • Head output: X∈RL×dmodelX \in \mathbb{R}^{L \times d_\mathrm{model}}8.
  • Outputs concatenated across X∈RL×dmodelX \in \mathbb{R}^{L \times d_\mathrm{model}}9 heads and projected.

While expressive, this design requires storing ϕ\phi0 independent ϕ\phi1 activations and exacerbates the intermediate-to-head imbalance as models scale, leading to excessive memory usage and reduced functional capacity in each head.

2.2 FlashMHF: Balanced and Dynamic Multi-Head Design

FlashMHF restructures each head’s intermediate pathway into ϕ\phi2 smaller, dynamically weighted sub-networks of width ϕ\phi3, the SwiGLU-optimal ratio. Specifically, the intermediate size ϕ\phi4 is partitioned as ϕ\phi5. Each head’s query ϕ\phi6 feeds into:

  • A gating matrix ϕ\phi7 that computes logits ϕ\phi8 per token: ϕ\phi9.
  • Gating probabilities are normalized to sum to one via sigmoid and scaling:

W1∈Rdf×dmodelW_1 \in \mathbb{R}^{d_f \times d_\mathrm{model}}0

  • Each sub-network W1∈Rdf×dmodelW_1 \in \mathbb{R}^{d_f \times d_\mathrm{model}}1 has parameters W1∈Rdf×dmodelW_1 \in \mathbb{R}^{d_f \times d_\mathrm{model}}2, W1∈Rdf×dmodelW_1 \in \mathbb{R}^{d_f \times d_\mathrm{model}}3, W1∈Rdf×dmodelW_1 \in \mathbb{R}^{d_f \times d_\mathrm{model}}4.
  • The final output for each head is a weighted sum over its sub-networks:

W1∈Rdf×dmodelW_1 \in \mathbb{R}^{d_f \times d_\mathrm{model}}5

Outputs from all heads are concatenated as usual and projected out.

This restores the optimal ratio between intermediate and head dimensions regardless of W1∈Rdf×dmodelW_1 \in \mathbb{R}^{d_f \times d_\mathrm{model}}6, enabling scalable expressivity and memory efficiency at all model scales (Zhang et al., 7 Dec 2025).

3. Fused Kernel and I/O-Aware Computation

FlashMHF’s kernel (“FlashFFN”) is architected to avoid the memory bottleneck of materializing the full W1∈Rdf×dmodelW_1 \in \mathbb{R}^{d_f \times d_\mathrm{model}}7 activation tensor. Analogous to FlashAttention, it exploits SRAM by streaming computations over the W1∈Rdf×dmodelW_1 \in \mathbb{R}^{d_f \times d_\mathrm{model}}8 dimension in tiled blocks of size W1∈Rdf×dmodelW_1 \in \mathbb{R}^{d_f \times d_\mathrm{model}}9, with all operations for a tile (associated W2∈Rdmodel×dfW_2 \in \mathbb{R}^{d_\mathrm{model} \times d_f}0) computed online before proceeding to the next. The pseudo-code logic is:

  • For each block W2∈Rdmodel×dfW_2 \in \mathbb{R}^{d_\mathrm{model} \times d_f}1 (W2∈Rdmodel×dfW_2 \in \mathbb{R}^{d_\mathrm{model} \times d_f}2):
    • Load W2∈Rdmodel×dfW_2 \in \mathbb{R}^{d_\mathrm{model} \times d_f}3 W2∈Rdmodel×dfW_2 \in \mathbb{R}^{d_\mathrm{model} \times d_f}4 for the tile,
    • Compute W2∈Rdmodel×dfW_2 \in \mathbb{R}^{d_\mathrm{model} \times d_f}5, W2∈Rdmodel×dfW_2 \in \mathbb{R}^{d_\mathrm{model} \times d_f}6,
    • Apply SiLU and gating: W2∈Rdmodel×dfW_2 \in \mathbb{R}^{d_\mathrm{model} \times d_f}7,
    • Apply router weights W2∈Rdmodel×dfW_2 \in \mathbb{R}^{d_\mathrm{model} \times d_f}8,
    • Accumulate W2∈Rdmodel×dfW_2 \in \mathbb{R}^{d_\mathrm{model} \times d_f}9.
  • Final output is written only once per head.

Each tile fits entirely in SRAM, and the total layerwise peak activation memory is reduced from H×(L⋅df)H \times (L \cdot d_f)0 (SwiGLU) to H×(L⋅df)H \times (L \cdot d_f)1—a practical reduction of 3–5H×(L⋅df)H \times (L \cdot d_f)2.

Hardware implementations leverage kernel fusion in Triton and Hopper/TK, incorporating asynchronous producer–consumer staging, warp-group specialization, and multi-stage buffering for maximal bandwidth utilization on NVIDIA H100 GPUs (Zhang et al., 7 Dec 2025).

4. Dynamical Routing via Parallel Sub-Networks

A distinctive feature of FlashMHF is the dynamic, per-head, per-token routing afforded by small gating networks. For each input token, the gating mechanism determines a weight over H×(L⋅df)H \times (L \cdot d_f)3 sub-networks, analogous to an internal soft MoE, but without the capacity collapse or static routing issues frequently present in such modules. All sub-networks remain active and contribute to each output, ensuring full differentiability and the ability to emphasize different “reasoning sub-paths” at a fine granularity.

By tethering the sub-network width H×(L⋅df)H \times (L \cdot d_f)4 directly to H×(L⋅df)H \times (L \cdot d_f)5 via H×(L⋅df)H \times (L \cdot d_f)6, FlashMHF permits scalable parallelization and expressivity. Empirical ablations demonstrate that tying H×(L⋅df)H \times (L \cdot d_f)7 to H×(L⋅df)H \times (L \cdot d_f)8 is critical—dense routing without the multi-head decomposition (i.e., H×(L⋅df)H \times (L \cdot d_f)9) underperforms, affirming the necessity of the combined multi-head and dynamic sub-network approach (Zhang et al., 7 Dec 2025).

5. Implementation and Integration in Transformers

FlashMHF serves as a drop-in replacement for SwiGLU or standard FFN layers within the Transformer block. The data flow, including LayerNorm, residual connections, and attention mechanisms, remains unchanged upstream and downstream. The hardware memory pattern involves:

  • Each HH0 head read into SRAM per tile.
  • Sub-network parameters (for HH1) double-buffered and streamed from high-bandwidth memory.
  • Activation accumulation and storage performed entirely in SRAM.

This strategy further reduces memory traffic by avoiding the allocation of large intermediate tensors, particularly beneficial for long context lengths or large model deployments on memory-constrained GPUs.

6. Empirical Evaluation

Benchmarks span models ranging from 128M to 1.3B parameters, with primary evaluation on PG19 validation (60–100B token pretraining) and downstream tasks including HellaSwag, SIQA, PIQA, OBQA, Winogrande, and RACE. Key results include:

  • Consistently improved perplexity versus SwiGLU:
    • 128M: FlashMHF-128hdim achieves lower eval loss.
    • 370M: SwiGLU loss 3.030, FlashMHF (d_h=128) loss 3.014; naïve MH-FFN fails to scale.
    • 1.3B: SwiGLU loss 2.843, FlashMHF (d_h=128) loss 2.793 (HH2 –0.050).
  • The highest average and per-task accuracy occurs with FlashMHF, especially for HH3.
  • Inference latency on NVIDIA H100 is up to 1.08HH4 faster than SwiGLU, with mean speedup HH51.05HH6; memory footprint per layer reduces by 3–5HH7 across varying sequence lengths.
  • Ablations indicate HH8 as optimal, with HH9 underfitting and dfd_f0 exhibiting diminishing returns due to decreasing head diversity.

7. Limitations and Trade-Offs

While FlashMHF delivers substantial memory and modest inference time improvements, several caveats are observed:

  • Kernel engineering complexity is significant, with performance hinging on careful tiling, warp group design, and latency-sensitive staging.
  • The latency advantage is less pronounced than FlashAttention due to already highly optimized baseline SwiGLU kernels.
  • Optimal dfd_f1 and dfd_f2 must be tuned per compute regime, balancing head diversity against routing and memory overhead.
  • The approach is most advantageous for settings involving long input contexts or large models deployed on memory-constrained GPUs.
  • For very small models or scenarios where kernel launch overhead dominates, established baselines such as SwiGLU may remain preferable (Zhang et al., 7 Dec 2025).

FlashMHF establishes multi-head, dynamically routed FFNs as an expressive and hardware-efficient principle for Transformer design, providing a scalable alternative for next-generation architectures.

Definition Search Book Streamline Icon: https://streamlinehq.com
References (1)

Topic to Video (Beta)

No one has generated a video about this topic yet.

Whiteboard

No one has generated a whiteboard explanation for this topic yet.

Follow Topic

Get notified by email when new papers are published related to Flash Multi-Head Feed-Forward Network (FlashMHF).