Skip to content
Snula
Curriculum
RU Open

Curriculum

Week 19. Other architectures: RNN, SSM, MoE

Phase 5. Expansion · week 19 of 24

Learn in the app: tutor, coding problems →

Core: RNNs and vanishing gradients, the idea of SSMs and Mamba, MoE: routing and balancing · Depth: LSTM in detail, linear attention, Gaussian processes (05-ГЛУБИНА) · ≈ 11 h core / 23 h total

The whole week answers one question: how do you pay for context? A transformer stores the entire history (the KV cache is O(S)) but trains in parallel. An RNN compresses the history into a fixed-size state: O(1) per token during generation, but training is sequential and distant things get forgotten. SSMs and hybrids try to get both.

Step 1. Vanilla RNN and decay over time

In plain terms. The gradient flows back through 40 steps and gets multiplied by the same number at each one. Multiplier 0.9: 0.9⁴⁰ ≈ 0.015, one and a half percent of the signal remains. Multiplier 1.1: 1.1⁴⁰ ≈ 45, the signal has exploded. In an RNN the number is replaced by a matrix W_h, and its spectral radius (the largest absolute value of an eigenvalue) plays the role of the multiplier. You cannot hold it at exactly 1, and the derivative of tanh (at most 1) only makes the multiplier smaller.

  • h_t = tanh(W_h h_{t−1} + W_x x_t). Unrolling over time (a copy of the cell for each step) gives a network of depth T with shared weights
  • Derive: ∂h_T/∂h_k = Π_{t=k+1..T} diag(1 − h_t²) W_h. This is a product of T − k Jacobians, which behaves like ρ(W_h)^{T−k}. Less than one means vanishing, greater than one means exploding
  • Takeaway: the problem is not tanh but repeated multiplication by the same matrix. Exploding is fixed by clipping (week 3); vanishing is fixed only by changing the architecture

Step 2. LSTM: additive memory

In plain terms. In an LSTM, along the memory path c, the gradient is multiplied not by W but by the forget gate, a number between 0 and 1 that the network sets itself at each step. A gate of 0.99 over 40 steps gives 0.99⁴⁰ ≈ 0.67: two thirds of the signal gets through. A gate of 0.9 would give the same one and a half percent as the RNN. The difference is that the network can keep this multiplier near one without breaking everything else.

  • c_t = f_t ⊙ c_{t−1} + i_t ⊙ g_t. Along the direct path ∂c_t/∂c_{t−1} = diag(f_t): there is no multiplication by W, and with f_t ≈ 1 the gradient travels hundreds of steps. That is why the forget gate bias is initialized to one
  • This is the same idea as the residual stream (week 6): an additive update gives the gradient a path around the nonlinearities

Step 3. SSM and Mamba

In plain terms. A one-dimensional linear recurrence h_t = 0.5·h_{t−1} + x_t, output y_t = h_t, input (2, 4, 6). Recurrently: h₁ = 2, h₂ = 1 + 4 = 5, h₃ = 2.5 + 6 = 8.5. As a convolution with kernel (1, 0.5, 0.25): y₃ = 1·6 + 0.5·4 + 0.25·2 = 8.5. Same number, but now all y_t can be computed at once. If the coefficient 0.5 depends on the input, as in Mamba, each position has its own kernel, and there is no longer a single convolution.

  • SSM (state space model, a model with a hidden state and a linear recurrence). The linear recurrence h_t = A h_{t−1} + B x_t, y_t = C h_t unrolls into a convolution with kernel (CB, CAB, CA²B, …). The main takeaway of the week: linearity buys both modes: parallel training as a convolution, and O(1) recurrent generation
  • Mamba: B, C and the step Δ depend on the input, so the model decides what to write and what to forget. The kernel stops being constant, and the convolution no longer works. It is replaced by a parallel scan (computing all prefixes of an associative operation in O(log T) parallel steps), which is possible because the composition of linear updates is associative
  • Linear attention: without softmax, (QKᵀ)V = Q(KᵀV), with a state KᵀV of size H×H. The cost of both approaches: worse exact retrieval from distant context. Hence hybrids that alternate attention and SSM layers

Sidebar. A window of observations as the state

In plain terms. A pendulum where you see only its displacement, rounded: 0, 7, 10, 7, 0, −7, −10, −7, 0. The value 7 appears twice: on the way out from the center and on the way back. From one number you cannot tell where the pendulum is heading. The pair "previous, current" already tells them apart: (0, 7) means "moving away", (10, 7) means "coming back". The difference of neighboring numbers gives the velocity, and position plus velocity together are the full state of the pendulum.

  • Steps 1–3 went from state to observations: RNNs and SSMs compress history into h_t. For a deterministic system the converse also holds (Takens' theorem, 1981): a window X_t = (x_t, x_{t−τ}, …, x_{t−(E−1)τ}) of values of one observed quantity with delay step τ generically determines the state one-to-one if E > 2d (d: the dimension of the attractor, the set the system moves on)
  • For a frictionless pendulum the trajectory is closed, d = 1, and the theorem guarantees success at E = 3. In the example two were enough: the guarantee has slack
  • The bridge to LLMs: the context window can also serve as the state, and a transformer with window S does not have to compress anything. But text is not a low-dimensional deterministic system, so this is an analogy, not a theorem
  • Do not confuse the homonyms. "State space" in SSM comes from control theory: the state h_t is linear, hidden, and defined by the model's equation. "State space reconstruction" in nonlinear dynamics means recovering the unknown state of someone else's system from a window of observations. Same root, different meaning

Step 4. Mixture of Experts

In plain terms. A block's FFN has 100M parameters. Replace it with 8 experts of the same size and a router (a small linear layer that picks experts for a token) with top-2. The FFN now has 800M parameters, but each token passes through only two experts, that is, through 200M. Compute doubled rather than growing 8 times. Collapse: with top-2 out of 8, each expert should get 25% of the tokens. If at the start expert 3 happens to get 40%, it learns faster, gets better, and the router sends it even more.

  • FFN → E experts, and the router picks top-k. Parameters grow by a factor of E, FLOPs per token by roughly k. Parameter count and compute are decoupled
  • Without balancing, routing collapses: popular experts learn faster and become more popular. Below are two ways to arrange experts differently and three ways to keep the load even

Shared and routed experts. Some experts are always on (shared), and the router picks the rest (routed): y = x + Σ_{j∈shared} FFN_j(x) + Σ_{i∈top-k} gᵢ·FFNᵢ(x). Shared experts handle what every token needs (grammar, frequent words), so the routed ones do not have to duplicate that knowledge in each of the E experts (DeepSeekMoE).

Fine-grained experts. In plain terms. 16 experts of width F, top-2: a token passes through 2F neurons, and there are C(16, 2) = 120 distinct pairs of experts. Cut each one into 4 narrow experts of width F/4: that makes 64, and you take top-8. A token passes through the same 8·F/4 = 2F neurons, the parameters are the same 64·F/4 = 16F, but there are C(64, 8) ≈ 4.4·10⁹ combinations. FLOPs and memory are unchanged, yet there are 37 million times more ways to assemble a "custom" expert for a token.

Auxiliary loss (Switch Transformer). In plain terms. E = 4, the fractions of assignments in the batch are f = (0.5, 0.25, 0.125, 0.125), the mean router probabilities are P = (0.4, 0.3, 0.15, 0.15). α·E·Σ fᵢ·Pᵢ = α·4·(0.2 + 0.075 + 0.01875 + 0.01875) = 1.25·α. With an even load f = P = (¼, ¼, ¼, ¼) you would get 1.0·α.

  • L_aux = α·E·Σᵢ fᵢ·Pᵢ. fᵢ (the fraction of tokens sent to expert i) comes from the top-k choice and has no gradient, so the gradient flows through Pᵢ (the batch-mean router probability for i): ∂L_aux/∂Pᵢ = α·E·fᵢ, which pushes hardest on the probability of the overloaded expert. The factor E makes the value at balance equal to 1 for any E. α is small, on the order of 10⁻²: a strong loss fights with the main task

Expert capacity. In plain terms. A batch has 96 tokens, E = 8, top-2, capacity factor (capacity headroom) CF = 1.25: each expert has 1.25·96·2/8 = 30 slots. An expert received 37 assignments: it does not process the 7 extra ones, and they continue only through the residual.

  • capacity = CF · B·T · k / E. Capacity is needed because tensor shapes on the accelerator are fixed. A larger CF means fewer dropped tokens, but more empty slots and more memory

Balancing without an auxiliary loss (router bias, DeepSeek-V3). In plain terms. A token's scores are s = (0.70, 0.55, 0.52, 0.30), so top-2 would pick experts 1 and 2. But 1 and 2 have been overloaded lately, 3 is underloaded, and biases b = (−0.02, −0.02, +0.04, 0) have accumulated. The choice is made on s + b = (0.68, 0.53, 0.56, 0.30): experts 1 and 3. The mixing weights come from the original scores: 0.70 / (0.70 + 0.52) ≈ 0.57 and 0.52 / 1.22 ≈ 0.43.

  • The bias bᵢ decides only which experts to pick and does not enter the weights gᵢ, so the output is not distorted by a meaningless offset. After each step, bᵢ is decreased by γ for overloaded experts and increased by γ for underloaded ones (γ: a small update rate). b has no gradient, and the main loss does not see it
  • Details for DeepSeek-V3. The gate is sigmoid: sᵢ = sigmoid(xᵀeᵢ) (eᵢ is the expert's vector), so each expert is scored independently, with no softmax over all of them; the weights are normalized over the selected ones, gᵢ = sᵢ / Σ_{j∈top-k} sⱼ. And "no loss at all" is not accurate: a sequence-level balancing loss with a very small α remains, so that a single sequence does not pile onto a couple of experts
  • Expert parallelism: two all-to-alls (each GPU sends each other GPU its share of tokens) per layer (week 14). The volume per layer is proportional to tokens × D × top-k; how many milliseconds that is within a node and across nodes is worked out in back-of-the-envelope problem 5 of week 14. The trainer has a problem on top-k routing: a convenient place to start

MTP as an auxiliary training objective (DeepSeek-V3). How the depth modules work is covered in week 11 (step 3); here is the training-side view. In plain terms. A sequence of 6 tokens. The main head gets 5 targets (each position predicts the next token), the depth-1 module gets 4 more (each position predicts the token after next). The main head's loss is 2.0, the module's is 3.0 (guessing two steps ahead is harder), weight λ = 0.3: L = 2.0 + 0.3·3.0 = 2.9. The shared trunk receives gradient from both objectives, with the auxiliary one weighted by 0.3.

  • Why: more signal per position, and the representation has to "plan" one step ahead. The gain is larger for big models and on code; for a small model the extra objective can hurt
  • Cost: one transformer block per depth and its FLOPs during training. The embedding and output head are shared with the model, so few parameters are added. λ is kept small and lowered toward the end of training
  • At inference the module can be dropped, and the model works as usual. Or it can become a draft for speculative decoding: if the second token is accepted 85% of the time, a pass yields 1 + 0.85 = 1.85 tokens on average
  • A common bug in code: the target of depth module j at position t is tokens[t + j + 1], not tokens[t + j]. The tail with no target is masked; with a single depth, the main objective must match ordinary shifted CE

Code → nanolm/rnn.py: VanillaRNN, LSTM (checked against nn.RNN and nn.LSTM) and gradient_norm_over_time, the norm of the gradient of the last step's loss with respect to each h_t. The exercise is in exercises_en/rnn.py; check it with NANOLM_IMPL=exercises_en pytest tests/test_rnn.py -v. Code → nanolm/moe.py: the router route (softmax or sigmoid gate, with the bias affecting only the choice), load_balancing_loss, bias_update, the capacity expert_capacity, the count moe_param_counts, and the MoE layer with shared and routed experts. The exercise is in exercises_en/moe.py: NANOLM_IMPL=exercises_en pytest tests/test_moe.py -v. Check by hand that with a skewed load the balancing loss is greater than α, and that the bias evens out the load without touching the mixing weights. Demo: python scripts/rnn_vanishing.py. For the RNN the norm falls exponentially; for the LSTM with forget gate bias 1 it falls more slowly, and with bias 3 it barely falls. As an extra, repeat the measurement for an RNN with spectral radius of W_h equal to 0.9, 1.0 and 1.1, and compare the slope of log‖∂L/∂h_t‖ with log ρ. Vanishing stops being a word and becomes the slope of a line.

Math (Track D): D29: vanishing and exploding gradients in RNNs; D30: stationary distribution and return time.

Interview question of the week: "Mamba is recurrent. Why does it train in parallel, and what does it pay for that?" A 3-minute structure: linear recurrence = convolution → selectivity breaks the convolution → a scan replaces it thanks to associativity → the cost: a fixed-size state, weaker exact retrieval → hence hybrids.

Sources: Hochreiter & Schmidhuber, LSTM (1997); Pascanu et al., *On the difficulty of training recurrent neural networks* (2013); Gu et al., S4 (2021); Gu & Dao, Mamba (2023); Fedus et al., Switch Transformers (2021); Lepikhin et al., GShard (2020); Dai et al., DeepSeekMoE (2024); Wang et al., Auxiliary-Loss-Free Load Balancing Strategy for Mixture-of-Experts (2024); DeepSeek-AI, DeepSeek-V3 Technical Report (2024); Takens, Detecting strange attractors in turbulence (1981).

Deeper: 05-ГЛУБИНА, sections "Weeks 19 and 22. Gaussian process regression" and "Small additions".

Week outcomes

  • I can derive ∂h_T/∂h_k for an RNN and explain vanishing and exploding through the spectral radius.
  • I can show through ∂c_t/∂c_{t−1} why an LSTM does not vanish, and connect this to the residual stream.
  • I can implement an RNN and an LSTM from scratch and show the difference in vanishing on a plot.
  • I can explain in 3 minutes why a linear SSM trains as a convolution while Mamba trains with a scan.
  • I can compute, for an MoE with E experts and top-k, the growth in parameters and in FLOPs per token.
  • I can compute the auxiliary loss and the expert capacity, and explain why the router bias does not enter the mixing weights.
  • I can compute the loss with MTP on a small example and explain how an MTP module becomes a draft at inference.

Self-check

  1. Why does gradient clipping save you from exploding but not from vanishing?
  2. What is "selective" in Mamba, and why does that rule out training as a convolution?
  3. What happens to an MoE without load balancing, and how is it fixed?
  4. Fine-grained experts do not change FLOPs or the parameter count. What do they gain, then?

Mock interviews of the week (1 and 2 of 12). From this week on, mocks come two a week; the protocols of sessions A–E are collected in week 23, scoring uses the rubrics. (1) ML coding, 45 minutes following session B: attention with a mask and GQA from an empty file. (2) Rapid-fire, session A: 30 questions from the self-checks of weeks 1–18. No partner: record your answers and score them against the rubric a day later.

In the app each week has skills to rate yourself on, questions with answer checking, Python coding problems and a tutor grounded in the course.

Learn in the app: tutor, coding problems
← PreviousWeek 18. Evaluation Next →Week 20. Multimodality and long context

Snula
Snula: LLMs from scratch

  • Home
  • Curriculum
  • App
  • Privacy
  • Terms

The course text is licensed under CC BY-NC-SA 4.0, nanolm code under Apache-2.0.