Elementwise kernels in C: RMSNorm, RoPE, SiLU-mul, embedding, add, argmax
Overview
Section titled “Overview”| Module | L9.6 · side · C · Pass 6 · 3 to 4 h |
| You build | c/src/kernels/elementwise.c: tl_rmsnorm_f32, tl_rope_f32 (half and interleaved layouts, partial rotary, attention scaling), tl_silu_mul_f32, tl_embedding_f32, tl_add_f32, tl_argmax_f32 |
| Contract | course/contracts/c/include/tinyllm/elementwise.h · rules: c/ABI.md |
| Tests | course/tests/L9.6/: test_elementwise.c (C, under ASan and UBSan) and shared file fixtures (what they check: section 4) |
| Needs | rt.02 the loader · M09.5 your tl_rsqrtf · M09.6 your tl_expf · L7.1 your RMSNorm · L7.3 your rope_cos_sin and apply_rope (or --ref-deps). Reading: L7.2 (SwiGLU) |
| Used by | These standalone C routines are useful as independent examples; the Python and Rust engine paths retain their own implementations. |
| Milestone | MS-L9 (the C backend generates the same tokens as numpy) |
| Optional depth | Zhang and Sennrich, “Root Mean Square Layer Normalization” (2019); Su et al., “RoFormer” (2021), section 3.4; Shazeer, “GLU Variants Improve Transformer” (2020) |
Key Takeaways
Section titled “Key Takeaways”- Each kernel touches every element once, so it is bound by memory bandwidth, not arithmetic: one pass, no temporaries, and outputs that may alias their first input (
rmsnorm_rows_are_independent_and_in_place,add_in_place). - RMSNorm puts inside the root, , and its sum of squares runs in a fixed order per row, so a row’s bits never depend on its batch (
rmsnorm_eps_is_inside_the_root). - RoPE is a rotation of pairs, and which entries form a pair is the layout: for HF Llama, for Meta’s code; both members are read before either is written (
rope_interleaved_layout_and_position_zero,rope_is_a_rotation). - SiLU saturates without NaN when written : at the exponential overflows to and the quotient is (
silu_mul_saturates_without_nan). - Greedy argmax breaks ties to the lowest index and skips NaN, the rule the Python sampler and the Rust engine share (
argmax_ties_nan_and_empty).
How to work this chapter
Section titled “How to work this chapter”ol start L9.6 # stubs c/src/kernels/elementwise.c into your repool tests L9.6 # read the test catalog firstol check L9.6 # exit code is the verdictol check L9.6 --ref-deps # only if rt.02, M09.5, M09.6, L7.1, or L7.3 is not passing yetol diff L9.6 # after passing: your code against the reference1. Why now
Section titled “1. Why now”Elementwise and row-wise operations include embedding lookup, normalization, rotary position encoding, gated activation, residual addition, and greedy selection. Each operation has details that silently change model output when implemented incorrectly: where goes, which entries RoPE pairs, how SiLU behaves at large gates, and how argmax breaks ties. This optional module implements the C versions as standalone routines, with the Python definitions serving as the behavioral reference.
2. Principles
Section titled “2. Principles”| Symbol | Meaning | Type / shape |
|---|---|---|
| one row (a token’s hidden state) | float[d] | |
| the RMSNorm gain | float[d] | |
| a small constant that keeps the root away from 0 | float | |
| , the mean of squares | scalar | |
tokens, heads, head width of a RoPE input [T, H, D] | int64_t | |
d_rot, the rotated width (even, ) | int64_t | |
| the absolute position of token | int32_t | |
inv_freq[i], the frequency of pair , | float[r/2] | |
| the rotation angle of pair at token | float | |
attn_scaling, YaRN’s factor on cos and sin (1 otherwise) | float | |
| the logistic sigmoid |
2.1 RMSNorm
Section titled “2.1 RMSNorm”RMSNorm (L7.1) rescales a row to unit root mean square, then applies a learned gain:
In C that is one pass to accumulate in float32 in increasing , one call to your tl_rsqrtf (M09.5, Newton’s method) for , and one pass to write . Because the order of the sum is fixed per row, a row’s result is the same whether the call normalizes 1 row or 64 (batch invariance, c/ABI.md rule 10). Because the sum is complete before the first write, y == x works. goes inside the root: a zero row then gives , never , and Llama’s checkpoints were trained that way.
2.2 RoPE
Section titled “2.2 RoPE”RoPE (L7.3) rotates pairs of a query or key vector by an angle proportional to its position. Pair of token is rotated by :
A rotation keeps the pair’s length (times ), and the dot product of a query rotated by and a key rotated by depends only on : attention sees relative position. Three details of the contract:
- The layout says which entries form pair :
layout 0(“half”, HF Llama, SmolLM2) pairs ;layout 1(“interleaved”, Meta’s code, the RoPE paper) pairs . Both are the same rotation on a permuted vector, which is why HF’s conversion script permutes the rows ofq_projandk_proj. - Partial rotary: only the first entries of each head rotate; entries to pass through.
- The angle is formed once per (token, pair) in double precision and rounded to float, which equals float32 float32 for every position below : the same angle your Python computes. One angle serves all heads, so the loop order is token, pair, head.
Both members of a pair must be read before either is written; computing from the new is a different (wrong) map.
2.3 SiLU-mul (SwiGLU)
Section titled “2.3 SiLU-mul (SwiGLU)”Llama’s MLP (L7.2) computes . The middle step is elementwise:
Written this way it is safe at both ends. For , and . For , overflows to (your tl_expf saturates above 88.7), and , the right limit. The algebraically equal computes at .
2.4 Embedding, add, argmax
Section titled “2.4 Embedding, add, argmax”Embedding copies row ids[t] of a [V, d] table: the source starts at element ids[t] * d. The contract makes the caller check ids against (the kernel cannot know ). Add is the residual connection, , usually in place. Argmax picks the greedy token (spec/sampling.md, temperature 0): the largest value, the lowest index among equal values (a strict > while scanning forward), NaN entries skipped, and when there is no number at all, so the caller reports an error instead of emitting token 0.
3. Worked example by hand
Section titled “3. Worked example by hand”RMSNorm of , , : , , , so .
RoPE (half layout) of , , , position 1, . Pairs are at and at :
| Pair | ||||
|---|---|---|---|---|
| , | 0.5403023 | 0.8414710 | ||
| , | 0.9999500 | 0.0099998 |
Written back to positions 0, 2 and 1, 3: . With the interleaved layout the pairs are and and the result is (rope_interleaved_layout_and_position_zero).
SiLU-mul of gate and up : , , ; times up: .
Argmax of : 7 first appears at index 1, and index 2 is not strictly greater, so the answer is 1.
All four are the first test, hand_example; RMSNorm and argmax repeat through the shared fixture in test_hand_example.
4. The interface
Section titled “4. The interface”/* tinyllm/elementwise.h: void, the caller guarantees valid pointers and sizes */void tl_rmsnorm_f32(const float *x, const float *w, float *y, int64_t rows, int64_t d, float eps);void tl_rope_f32(float *x, const int32_t *pos, int64_t T, int64_t H, int64_t D, int64_t d_rot, const float *inv_freq, float attn_scaling, int layout /* 0 half, 1 interleaved */);void tl_silu_mul_f32(const float *gate, const float *up, float *y, int64_t n);void tl_embedding_f32(const float *table, const int32_t *ids, float *out, int64_t n, int64_t d);void tl_add_f32(const float *a, const float *b, float *y, int64_t n);int32_t tl_argmax_f32(const float *x, int64_t n); /* lowest index on ties; NaN skipped; -1 if none */From Python, declare each with restype None (or c_int32 for argmax) on your loader.
What the tests check
Section titled “What the tests check”| Test | KIND | Checks | Why it matters downstream |
|---|---|---|---|
hand_example | unit, smoke | section 3 for RMSNorm, RoPE, SiLU-mul, argmax | you and the tests agree on the definitions |
rmsnorm_eps_is_inside_the_root | boundary | with gives 0.7071; a zero row gives zeros | tiny-norm rows, Llama’s trained form |
rmsnorm_rows_are_independent_and_in_place | property | every row alone equals the row in a batch of up to 12, bit for bit; y == x | batched decode in L10.2 |
rope_interleaved_layout_and_position_zero | unit | the interleaved worked example; position 0 is the identity | loading Meta-layout weights |
rope_partial_rotary_and_scaling | unit | the tail past d_rot untouched; doubles the pair; [T, H, D] indexing | GPT-NeoX style partial rotary, YaRN |
rope_is_a_rotation | property | every pair keeps its length, random shapes and positions, both layouts | no half-updated pairs, no wrong partners |
silu_mul_saturates_without_nan | boundary | gates , give the right limits | large activations in trained MLPs |
embedding_gathers_rows | unit | repeated and out-of-order ids | the first op of the forward |
add_in_place | unit | with y == a | the residual stream |
argmax_ties_nan_and_empty | boundary | ties to the lowest index, NaN skipped, for all NaN or empty, all gives 0 | greedy decoding parity |
hand_example | unit, smoke | RMSNorm and argmax of section 3 in the C test harness | checks standalone routines |
5. Pitfalls
Section titled “5. Pitfalls”| Pitfall | Symptom | Caught by |
|---|---|---|
| outside the root, | correct on ordinary rows, wrong on tiny ones; drifts from the trained model | rmsnorm_eps_is_inside_the_root (mutant s01) |
| the sum of squares not divided by | every output shrinks by | rmsnorm_eps_is_inside_the_root (mutant s02) |
| forgetting the gain | right shape, wrong scale | hand_example (mutant s03) |
| the sum of squares declared once per call, not per row | row 2 is normalized by rows 1 and 2 together; results depend on the batch | rmsnorm_rows_are_independent_and_in_place (mutant s15) |
| the layouts swapped | plausible numbers, wrong attention pattern on HF weights | rope_interleaved_layout_and_position_zero (mutant s04) |
| rotating by (a sign flipped) | relative positions mirrored; the model’s outputs degrade | rope_interleaved_layout_and_position_zero (mutant s05) |
| computing from the new | lengths change; not a rotation | rope_is_a_rotation (mutant s06) |
rotating all entries when d_rot < D | the pass-through tail is scrambled (and inv_freq is read past its end) | rope_partial_rotary_and_scaling (mutant s07) |
ignoring attn_scaling | YaRN-extended models lose their temperature correction | rope_partial_rotary_and_scaling (mutant s08) |
| SiLU as a sigmoid, or with instead of | wrong MLP; or instead of at large gates | silu_mul_saturates_without_nan (mutants s09, s10) |
embedding row at ids[t] instead of ids[t] * d | every token reads a slice of row 0 or 1 | embedding_gathers_rows (mutant s11) |
argmax with >= | ties go to the last index; greedy output differs from Python | argmax_ties_nan_and_empty (mutant s12) |
| argmax without the NaN check | a NaN in the first place wins forever | argmax_ties_nan_and_empty (mutant s13) |
| returning 0 for an empty or all-NaN row | the engine emits token 0 instead of reporting the error | argmax_ties_nan_and_empty (mutant s14) |
6. Where it’s used next
Section titled “6. Where it’s used next”| Direction | Module | How it uses this |
|---|---|---|
| Back | rt.02 | the shared status, error, and allocator support |
| Back | M09.5 | tl_rsqrtf: the of RMSNorm |
| Back | M09.6 | tl_expf: the of SiLU |
| Back | L7.1 | RMSNorm, the specification of tl_rmsnorm_f32 |
| Back | L7.3 | rope_cos_sin and apply_rope, the specification of tl_rope_f32 |
| Forward | the standalone Rust engine | the Rust forward calls the same six kernels in its independent Rust implementation |
Going further
Section titled “Going further”| Your piece | Production equivalent | What it adds | Where to look |
|---|---|---|---|
| separate RMSNorm and add | fused add + RMSNorm (vLLM fused_add_rms_norm) | one pass that adds the residual and normalizes, halving memory traffic | vLLM csrc/layernorm_kernels.cu |
RoPE with cos per call | cached cos/sin tables and fused QK rotation | precomputed tables per position, rotation fused into the QKV projection epilogue | llama.cpp ggml_rope_ext; FlashInfer rope.cuh |
| scalar SiLU | vectorized SwiGLU with a polynomial exp | SIMD over 8 or 16 lanes | ggml ggml_vec_swiglu_f32 |
| argmax | fused sampling kernels | argmax and top-k inside the final matmul’s epilogue | FlashInfer sampling.cuh |