spb/forge Public MIT
Forge — LLM training from scratch in pure C++20 + Metal on Apple Silicon.
C++ 61.2%
C 23%
Python 7.6%
TeX 7.2%
CMake 1.1%
1// Author: Simon-Pierre Boucher — contact@spboucher.ai2#pragma once34#include "nn/attention.h"5#include "nn/config.h"6#include "nn/embedding.h"7#include "nn/mlp.h"89#include <cmath>10#include <memory>11#include <vector>1213namespace forge::nn {1415// Norm wrapper so blocks don't branch on the norm flavour.16class Norm : public Module {17public:18 Norm(const ModelConfig& cfg) : layernorm_(cfg.norm == "layernorm"), eps_(cfg.norm_eps) {19 w_ = register_param("weight", Tensor::full({cfg.d_model}, 1.0f));20 if (layernorm_) b_ = register_param("bias", Tensor::zeros({cfg.d_model}));21 }2223 Var forward(const Var& x) const {24 return layernorm_ ? ops::layernorm(x, w_, b_, eps_) : ops::rmsnorm(x, w_, eps_);25 }2627private:28 bool layernorm_;29 float eps_;30 Var w_, b_;31};3233// Residual block; norm placement is config-selected:34// pre (default): x += attn(norm1(x)); x += mlp(norm2(x))35// post (OLMo 2): x += norm1(attn(x)); x += norm2(mlp(x))36// sandwich (Gemma): x += norm1p(attn(norm1(x))); x += norm2p(mlp(norm2(x)))37class TransformerBlock : public Module {38public:39 TransformerBlock(const ModelConfig& cfg, float proj_std, std::mt19937_64& rng,40 int64_t layer_idx = 0)41 : placement_(cfg.norm_placement == "pre" ? Placement::Pre42 : cfg.norm_placement == "post" ? Placement::Post43 : Placement::Sandwich),44 norm1_(std::make_unique<Norm>(cfg)),45 attn_(std::make_unique<CausalSelfAttention>(cfg, proj_std, rng, layer_idx)),46 norm2_(std::make_unique<Norm>(cfg)) {47 if (cfg.n_experts > 0 && layer_idx >= cfg.first_k_dense)48 mlp_ = std::make_unique<MoEMLP>(cfg, proj_std, rng);49 else50 mlp_ = std::make_unique<MLP>(cfg, proj_std, rng);51 absorb("norm1", *norm1_);52 absorb("attn", *attn_);53 absorb("norm2", *norm2_);54 absorb("mlp", *mlp_);55 if (placement_ == Placement::Sandwich) {56 norm1_post_ = std::make_unique<Norm>(cfg);57 norm2_post_ = std::make_unique<Norm>(cfg);58 absorb("norm1_post", *norm1_post_);59 absorb("norm2_post", *norm2_post_);60 }61 }6263 Var forward(const Var& x) const {64 const int64_t B = x.value().size(0), T = x.value().size(1), C = x.value().size(2);65 Var a;66 switch (placement_) {67 case Placement::Pre: a = attn_->forward(norm1_->forward(x)); break;68 case Placement::Post: a = norm1_->forward(attn_->forward(x)); break;69 case Placement::Sandwich:70 a = norm1_post_->forward(attn_->forward(norm1_->forward(x)));71 break;72 }73 Var h = ops::add(x, a);74 Var m;75 switch (placement_) {76 case Placement::Pre:77 m = mlp_->forward(norm2_->forward(h).reshaped({B * T, C}))78 .reshaped({B, T, C});79 break;80 case Placement::Post:81 m = norm2_->forward(82 mlp_->forward(h.reshaped({B * T, C})).reshaped({B, T, C}));83 break;84 case Placement::Sandwich:85 m = norm2_post_->forward(86 mlp_->forward(norm2_->forward(h).reshaped({B * T, C}))87 .reshaped({B, T, C}));88 break;89 }90 return ops::add(h, m);91 }9293 // MoE load-balance loss of the last forward (undefined for dense MLP).94 Var moe_aux() const { return mlp_->aux(); }95 void moe_bias_update(float gamma) { mlp_->bias_update(gamma); }9697private:98 enum class Placement { Pre, Post, Sandwich };99 Placement placement_;100 std::unique_ptr<Norm> norm1_;101 std::unique_ptr<AttentionBase> attn_;102 std::unique_ptr<Norm> norm2_;103 std::unique_ptr<MLPBase> mlp_;104 std::unique_ptr<Norm> norm1_post_, norm2_post_; // sandwich only105};106107// Decoder-only transformer, entirely shaped by ModelConfig.108class Transformer : public Module {109public:110 Transformer(const ModelConfig& cfg, uint64_t seed) : cfg_(cfg) {111 std::mt19937_64 rng(seed);112 const float std = 0.02f;113 // Residual projections scaled by 1/sqrt(2L): one attn + one mlp114 // residual add per layer (nanoGPT).115 const float proj_std = std / std::sqrt(2.0f * float(cfg.n_layers));116117 tok_emb_ = std::make_unique<Embedding>(cfg.vocab_size, cfg.d_model, std, rng);118 absorb("tok_emb", *tok_emb_);119 if (!cfg.use_rope) {120 pos_emb_ = std::make_unique<Embedding>(cfg.context_length, cfg.d_model, std, rng);121 absorb("pos_emb", *pos_emb_);122 }123 for (int64_t i = 0; i < cfg.n_layers; ++i) {124 blocks_.push_back(std::make_unique<TransformerBlock>(cfg, proj_std, rng, i));125 absorb("blocks." + std::to_string(i), *blocks_.back());126 }127 final_norm_ = std::make_unique<Norm>(cfg);128 absorb("final_norm", *final_norm_);129130 if (cfg.tied_embeddings) {131 lm_head_weight_ = register_param("lm_head.weight", tok_emb_->weight());132 } else {133 lm_head_weight_ = register_param(134 "lm_head.weight", normal_init({cfg.vocab_size, cfg.d_model}, std, rng));135 }136 }137138 // ids: [B, T] u16/i32 → logits [B*T, V]139 Var forward(const Tensor& ids) const {140 const int64_t B = ids.shape()[0], T = ids.shape()[1];141 Var x = tok_emb_->forward(ids);142 if (cfg_.scale_embeddings)143 x = ops::scale(x, std::sqrt(float(cfg_.d_model)));144 if (pos_emb_) {145 Tensor pos = Tensor::empty({1, T}, DType::I32);146 for (int64_t t = 0; t < T; ++t) pos.data<int32_t>()[t] = int32_t(t);147 Var p = pos_emb_->forward(pos); // [1, T, C]148 // Broadcast over batch: viewed as [B, T*C] rows + a [T*C] "bias",149 // add_bias sums the positional grad over the batch — exactly right.150 Var xb = x.reshaped({B, T * cfg_.d_model});151 Var pflat = p.reshaped({T * cfg_.d_model});152 x = ops::add_bias(xb, pflat).reshaped({B, T, cfg_.d_model});153 }154 for (const auto& blk : blocks_) x = blk->forward(x);155 x = final_norm_->forward(x);156 Var x2d = x.reshaped({B * T, cfg_.d_model});157 Var logits = ops::matmul(x2d, lm_head_weight_, false, true); // [B*T, V]158 if (cfg_.final_softcap > 0.0f)159 logits = ops::softcap(logits, cfg_.final_softcap);160 return logits;161 }162163 // targets: [B, T] with ignore_index=-1 → scalar mean CE loss. With MoE,164 // each block's load-balance term is added (also during eval — its scale165 // is ~moe_aux_weight, negligible next to CE but keeps train/val166 // comparable).167 Var loss(const Tensor& ids, const Tensor& targets) const {168 Var logits = forward(ids);169 Tensor tflat = targets;170 Var l = ops::cross_entropy(logits, tflat.view({targets.numel()}));171 if (cfg_.n_experts > 0 && cfg_.moe_aux_weight > 0.0f) {172 for (const auto& blk : blocks_) {173 Var a = blk->moe_aux();174 if (a.defined()) l = ops::add(l, a);175 }176 }177 return l;178 }179180 const ModelConfig& config() const { return cfg_; }181182 // noaux balancing step (V3): call after the optimizer sync, CPU-side.183 void update_moe_bias(float gamma) {184 for (auto& blk : blocks_) blk->moe_bias_update(gamma);185 }186187private:188 ModelConfig cfg_;189 std::unique_ptr<Embedding> tok_emb_;190 std::unique_ptr<Embedding> pos_emb_; // null when RoPE191 std::vector<std::unique_ptr<TransformerBlock>> blocks_;192 std::unique_ptr<Norm> final_norm_;193 Var lm_head_weight_;194};195196} // namespace forge::nn197