Aleph Alpha
Warm orange gradient with translucent stepped blocks in the lower left and upper right corners and a soft pale glow in the centre
Research

Pit NeitemeierAlessio Serra

Designing Kolibri: Architecture Trade-Offs from First Principles

TL;DR


Introduction

Architecture Configurations

Total
78.1B
Active / token
3.5B
FLOPs / token · 16K
28.3G
State · 16K
189M
Backbone
Embedding tying
Attention
KV heads4
Experts

Parameter Allocation

Total and Active Parameters per Model

78.1B3.5B4.4%
30.6B3.3B10.7%
31.6B3.2B10.2%
120.7B12.2B10.1%
549.3B55.0B10.0%
125.1B6.0B4.8%
1.57T48.8B3.1%
551.5B15.8B2.9%
2.78T104.2B3.7%
308.8B14.8B4.8%
116.8B5.1B4.4%
20.9B3.6B17.3%
743.4B40.3B5.4%
426.2B24.7B5.8%

Total counts every stored weight of the text decoder. Active / token counts the weights each token passes through: sequence mixer, the top-k routed and shared experts, dense FFN layers, router, norms and LM head; the input embedding is a lookup and is left out.

Figure 1. Total and active parameters per model, active parameters exclude the input embedding. Sort by any column and click a model to highlight it in every chart.

FFN: Dense Layers, Routed and Shared Experts

Prouted=Lmoe×3×d×de×{EtotalkactiveP_{\mathrm{routed}} = L_{\mathrm{moe}} \times 3 \times d \times d_{e} \times \begin{cases}E & \text{total} \\ k & \text{active}\end{cases}
Pshared=Lmoe×3×d×dsP_{\mathrm{shared}} = L_{\mathrm{moe}} \times 3 \times d \times d_{s}
Pdense=Ldense×3×d×dffP_{\mathrm{dense}} = L_{\mathrm{dense}} \times 3 \times d \times d_{\mathrm{ff}}
Prouter=Lmoe×E×dP_{\mathrm{router}} = L_{\mathrm{moe}} \times E \times d

Sequence Mixer: Attention Projections

PQKVO=Lattn×[ddq+2ddkv+dqd]=2Lattnd(dq+dkv)P_{\mathrm{QKVO}} = L_{\mathrm{attn}} \times [d d_{q} + 2d d_{\mathrm{kv}} + d_{q} d] = 2L_{\mathrm{attn}} d(d_{q} + d_{\mathrm{kv}})

Norms

Pnorms=(4L+1)×dP_{\mathrm{norms}} = (4L + 1) \times d

Embedding and LM Head

Plookup=Phead=d×VP_{\mathrm{lookup}} = P_{\mathrm{head}} = d \times V

Active Parameter Composition per Token

sequence mixerFFN · active experts + dense layersLM headnorms
Figure 2. Active parameters per token by component. Choose which component is stacked first and how the models are sorted.

Parameter Allocation Across Architectures

Pseq=2Lattn d (dq+dkv)Pffn=3d[Lmoe(k de+ds)+Ldense dff]\begin{aligned} P_{\text{seq}} &= 2L_{\text{attn}}\,d\,(d_q+d_{\text{kv}}) \\[4pt] P_{\text{ffn}} &= 3d\left[L_{\text{moe}}(k\,d_e+d_s)+L_{\text{dense}}\,d_{\text{ff}}\right] \end{aligned}
PseqPffn=2Lattn(dq+dkv)3[Lmoe(k de+ds)+Ldense dff]\frac{P_{\text{seq}}}{P_{\text{ffn}}} = \frac{2L_{\text{attn}}(d_q+d_{\text{kv}})} {3\left[L_{\text{moe}}(k\,d_e+d_s)+L_{\text{dense}}\,d_{\text{ff}}\right]}

FLOP Allocation

Ftrain/ token ≈F_{\mathrm{train}}/\text{ token } \approx 6×Pmat6 \times P_{\mathrm{mat}} parameter FLOPs ++ 12×Lfull×dq×n/212 \times L_{\mathrm{full}} \times d_{q} \times n/2 full attention ++ 12×Lswa×dq×min⁡(n,w)12 \times L_{\mathrm{swa}} \times d_{q} \times \operatorname{min}(n,w) sliding window
Y=XW ⁣:2TdindoutY = XW\colon 2T d_{\mathrm{in}} d_{\mathrm{out}}∂X=∂YW⊤ ⁣:2Tdindout\partial X = \partial Y W^{\top}\colon 2T d_{\mathrm{in}} d_{\mathrm{out}}∂W=X⊤∂Y ⁣:2Tdindout\partial W = X^{\top}\partial Y\colon 2T d_{\mathrm{in}} d_{\mathrm{out}}

Beyond Full and Sliding-Window Attention

FFN

Fffn/ token =6×3×d×[Lmoe×(k×de+ds)+Ldense×dff]F_{\mathrm{ffn}}/\text{ token } = 6 \times 3 \times d \times [L_{\mathrm{moe}} \times (k \times d_{e} + d_{s}) + L_{\mathrm{dense}} \times d_{\mathrm{ff}}]
Frouter/ token =6×Lmoe×E×dF_{\mathrm{router}}/\text{ token } = 6 \times L_{\mathrm{moe}} \times E \times d

Sequence-Mixer Projections

Fproj/ token =6×2×Lattn×d×(dq+dkv)F_{\mathrm{proj}}/\text{ token } = 6 \times 2 \times L_{\mathrm{attn}} \times d \times (d_{q} + d_{\mathrm{kv}})

Pure Attention

Fattn/ token =12×Lfull×dq×n/2F_{\mathrm{attn}}/\text{ token } = 12 \times L_{\mathrm{full}} \times d_{q} \times n/2
Fattn/ token =12×Lswa×dq×min⁡(n,w)F_{\mathrm{attn}}/\text{ token } = 12 \times L_{\mathrm{swa}} \times d_{q} \times \operatorname{min}(n,w)

LM Head

Fhead/ token =6×d×VF_{\mathrm{head}}/\text{ token } = 6 \times d \times V

Training FLOP Composition at 16K

pure attention + state updatessequence-mixer linearsFFNLM headnorms
Figure 3. Training FLOPs per token by component at the selected sequence length. Move the slider to see the mix shift towards pure attention.

FLOP Allocation vs. Context Length

Training FLOPs Over Sequence Length

CandidatesReferences
Figure 4. Model training FLOPs per token. Lower numerical precision does not change the FLOP count, but it can increase hardware throughput by reducing the number of bits used to represent each value.

FLOP Scaling vs. Context Length

Sequence-Mixer State Size

Mstate=c×B×b×dkv×Lkv×nretM_{\mathrm{state}} = c \times B \times b \times d_{\mathrm{kv}} \times L_{\mathrm{kv}} \times n_{\mathrm{ret}}

Cached Representation

Batch

Bytes per Element

KV Width

Attention Layers

Retained Tokens

Recurrent Layers

Sequence-Mixer State Over Sequence Length

CandidatesReferences
Figure 5. State elements for one sequence (Mstate with b = 1): total in the cache (solid) and read by one decode step (dashed, where it differs). Every full-attention layer grows with context; bounded layers cap at their window or hold a fixed recurrent state.

Sequence-Mixer State Across Architectures

Conclusion

Acknowledgements

References