π» Modern ArchitecturesLEVEL 3 Β· ADVANCEDChapter 9
Mixture of Experts (MoE) & Multi-Head Latent Attention (MLA)
DeepSeek-V3 architecture: Top-2 gating routing, auxiliary-loss-free balancing, and 93% KV cache compression.
π Inspect Architecture: DeepSeek-V3 and R1 reasoning MoEπ‘ 1. Core Intuition & Concepts
Frontier reasoning architectures like DeepSeek-V3 scale model capacity while keeping inference FLOPs constant using Sparse Mixture-of-Experts (MoE) and Multi-Head Latent Attention (MLA). MLA compresses keys and values into a shared low-rank latent representation.
π 2. Mathematical Formulations & Derivations
MoE Top-k Softmax Routing
Tokens route dynamically to k of E experts with weighted combination.
Multi-Head Latent Attention (MLA) Low-Rank Compression
Compresses key-value cache by up to 93% using low-rank latent vector c_t^{KV}.
βοΈ 3. Step-by-Step Computational Mechanism
1
Router Gating Score
Compute token affinity scores s = Softmax(x W_g) across E available expert networks.
2
Top-k Selection
Select indices of k highest affinity experts (e.g. Top-2 or Top-8) and normalize routing weights.
3
Expert Dispatch & Accumulation
Execute selected expert FFNs and compute weighted sum sum s_i E_i(x).
4
MLA Latent Decompression
Project compressed latent vector c_t^{KV} into multi-head keys and values on the fly.
π» 4. Code from Scratch (python)
moe_mla.py
import torch
import torch.nn as nn
import torch.nn.functional as F
class SparseMoELayer(nn.Module):
"""Top-2 Sparse Mixture-of-Experts Layer with Gating in PyTorch."""
def __init__(self, dim: int, hidden_dim: int, num_experts: int = 8, top_k: int = 2):
super().__init__()
self.num_experts = num_experts
self.top_k = top_k
self.router = nn.Linear(dim, num_experts, bias=False)
self.experts = nn.ModuleList([
nn.Sequential(
nn.Linear(dim, hidden_dim),
nn.SiLU(),
nn.Linear(hidden_dim, dim)
) for _ in range(num_experts)
])
def forward(self, x: torch.Tensor) -> torch.Tensor:
B, T, D = x.shape
x_flat = x.view(-1, D)
router_logits = self.router(x_flat)
top_k_logits, top_k_indices = torch.topk(router_logits, self.top_k, dim=-1)
top_k_weights = F.softmax(top_k_logits, dim=-1)
out = torch.zeros_like(x_flat)
for expert_idx in range(self.num_experts):
mask = (top_k_indices == expert_idx)
if not mask.any(): continue
token_indices, k_positions = torch.where(mask)
expert_inputs = x_flat[token_indices]
expert_outputs = self.experts[expert_idx](expert_inputs)
weights = top_k_weights[token_indices, k_positions].unsqueeze(-1)
out.index_add_(0, token_indices, expert_outputs * weights)
return out.view(B, T, D)π§ 5. Comprehension Checkpoint
Answer all 1 questions correctly to complete the chapter Β· 0 / 1 done
Q1/1 What is the key advantage of a 671B parameter MoE model that activates 37B parameters per token over a dense 671B model?
In the catalog
Finished this chapter?