SPB Git

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%
7.6 KB · 197 lines c
Raw Blame History
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