Skip to content
← All tracks

Flax

JAX

Modules, layers from scratch, attention, transformer architectures, training loops, lifted transforms. Production model code in JAX (Linen API).

This track is written for JAX, which isn't the mode you're browsing in.

0 / 100 solved
  1. 1. Not solved yet. Module with @nn.compact
  2. 2. Not solved yet. Module with setup() (alternative to compact)
  3. 3. Not solved yet. Custom Parameter Initializer
  4. 4. Not solved yet. init() and apply() Round-Trip
  5. 5. Not solved yet. nn.Sequential Composition
  6. 6. Not solved yet. Module with Multiple Named Sub-Modules
  7. 7. Not solved yet. Module That Branches on a Config Flag
  8. 8. Not solved yet. Three Levels of Module Nesting
  9. 9. Not solved yet. Multiple PRNG Streams (params, dropout)
  10. 10. Not solved yet. Train vs Eval Branches via train Flag
  11. 11. Not solved yet. Implement Dense from Scratch
  12. 12. Not solved yet. Implement Conv1D from Scratch
  13. 13. Not solved yet. Implement Conv2D with Stride and Padding
  14. 14. Not solved yet. Implement Transposed Convolution
  15. 15. Not solved yet. Implement Depthwise-Separable Convolution
  16. 16. Not solved yet. Implement LayerNorm with γ/β
  17. 17. Not solved yet. Implement BatchNorm with Mutable batch_stats
  18. 18. Not solved yet. Implement GroupNorm
  19. 19. Not solved yet. Implement RMSNorm (Modern LLM Norm)
  20. 20. Not solved yet. Implement Dropout with RNG Threading
  21. 21. Not solved yet. Scaled Dot-Product Attention
  22. 22. Not solved yet. Multi-Head Self-Attention with Flax
  23. 23. Not solved yet. Causal Multi-Head Self-Attention
  24. 24. Not solved yet. Cross-Attention with Flax MHA
  25. 25. Not solved yet. Multi-Head Attention with KV Cache
  26. 26. Not solved yet. Grouped-Query Attention (GQA)
  27. 27. Not solved yet. Multi-Query Attention (MQA)
  28. 28. Not solved yet. Sliding-Window Attention (Mistral-style)
  29. 29. Not solved yet. ALiBi: Attention with Linear Biases
  30. 30. Not solved yet. Block-Diagonal Attention Mask
  31. 31. Not solved yet. Token Embedding with Flax
  32. 32. Not solved yet. Sinusoidal Position Encoding
  33. 33. Not solved yet. Learned Position Embedding
  34. 34. Not solved yet. Rotary Position Embedding (RoPE)
  35. 35. Not solved yet. ALiBi Bias Matrix
  36. 36. Not solved yet. T5 Relative Position Bucketing
  37. 37. Not solved yet. Tied Input/Output Embedding
  38. 38. Not solved yet. ViT Patch Embedding
  39. 39. Not solved yet. Transformer Encoder Block (Pre-LN)
  40. 40. Not solved yet. Transformer Decoder Block (Pre-LN)
  41. 41. Not solved yet. Pre-LN vs Post-LN Residual Pattern
  42. 42. Not solved yet. Mini GPT — Decoder-Only Language Model
  43. 43. Not solved yet. Mini BERT — Encoder-Only Hidden States
  44. 44. Not solved yet. Mini T5 — Encoder-Decoder with RMSNorm and Tied Embeddings
  45. 45. Not solved yet. Vision Transformer (Mean-Pool Variant)
  46. 46. Not solved yet. Vision Transformer with [CLS] Token
  47. 47. Not solved yet. DeiT — Data-Efficient Image Transformer
  48. 48. Not solved yet. SwiGLU Feed-Forward Network
  49. 49. Not solved yet. ResNet Basic Block
  50. 50. Not solved yet. ResNet Bottleneck Block
  51. 51. Not solved yet. Tiny ResNet Classifier
  52. 52. Not solved yet. Tiny U-Net
  53. 53. Not solved yet. GRU Cell Step
  54. 54. Not solved yet. LSTM Cell Step
  55. 55. Not solved yet. Bidirectional RNN (Flax)
  56. 56. Not solved yet. Mixture-of-Experts FFN
  57. 57. Not solved yet. Squeeze-and-Excitation Block (Flax)
  58. 58. Not solved yet. Vision-Language Fusion
  59. 59. Not solved yet. TrainState — One Step
  60. 60. Not solved yet. train_step with value_and_grad
  61. 61. Not solved yet. eval_step — Forward + Metrics
  62. 62. Not solved yet. Label-Smoothed Cross-Entropy
  63. 63. Not solved yet. Mixed-Precision Training Step
  64. 64. Not solved yet. Train with Mutable batch_stats
  65. 65. Not solved yet. Multi-Task Two-Head Loss
  66. 66. Not solved yet. Sharded Eval Loss
  67. 67. Not solved yet. Warmup-Cosine LR at Step
  68. 68. Not solved yet. Gradient Accumulation Step
  69. 69. Not solved yet. EMA of Parameters
  70. 70. Not solved yet. Orbax Save (Tree-Leaf Count)
  71. 71. Not solved yet. Orbax Load (Restore via Template)
  72. 72. Not solved yet. HF Weight Load (Kernel Transpose)
  73. 73. Not solved yet. Pre-train then Fine-tune (Frozen Trunk)
  74. 74. Not solved yet. Per-Param Weight Decay Mask
  75. 75. Not solved yet. Per-Param Learning Rate Multipliers
  76. 76. Not solved yet. Param Freezing via Grad Zeroing
  77. 77. Not solved yet. Test-Time Augmentation Aggregation
  78. 78. Not solved yet. Distributed Checkpoint (Sharding Math)
  79. 79. Not solved yet. nn.scan over an RNN cell
  80. 80. Not solved yet. nn.scan over layers
  81. 81. Not solved yet. nn.vmap with shared params
  82. 82. Not solved yet. nn.checkpoint (gradient checkpointing)
  83. 83. Not solved yet. nn.jit (Flax-aware JIT lift)
  84. 84. Not solved yet. nn.remat with checkpoint policies
  85. 85. Not solved yet. Composed lifts: nn.scan + nn.vmap
  86. 86. Not solved yet. Batched init via jax.vmap
  87. 87. Not solved yet. jax.lax.scan inside a Flax Module
  88. 88. Not solved yet. Custom lift: roll your own ensemble
  89. 89. Not solved yet. PartitionSpec Layout
  90. 90. Not solved yet. with_sharding_constraint Annotation
  91. 91. Not solved yet. nn.with_partitioning Annotation
  92. 92. Not solved yet. flax.struct.dataclass — Pytree-Friendly State
  93. 93. Not solved yet. Param Surgery — Kernel Replace
  94. 94. Not solved yet. Param Surgery — Zero Last Layer
  95. 95. Not solved yet. Param Surgery — Freeze First Dense
  96. 96. Not solved yet. Partial Init — Warm-Start From Smaller Checkpoint
  97. 97. Not solved yet. Multiple Mutable Collections
  98. 98. Not solved yet. Param Sharing — One Module, Two Call Sites
  99. 99. Not solved yet. shard_map Simulation — Manual SPMD
  100. 100. Not solved yet. Mini LM Capstone — Putting It All Together

Check yourself

4 questions · one attempt each

These do not count toward finishing the track. They are here to catch the things that are easy to read past.

0 / 4

A Linen model with BatchNorm is applied in training mode. Why does it raise without mutable=?

BN().apply(v, x, train=True)
BN().apply(v, x, train=True, mutable=["batch_stats"])
Question 1 of 4