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%
13.5 KB · 348 lines cpp
Raw Blame History
1// Author: Simon-Pierre Boucher — contact@spboucher.ai2//3// forge CLI:4//   forge train    --config configs/gpt-25m.json --data data/tinystories --out runs/exp15//                  [--resume runs/exp1/ckpt_latest.bin] [--backend metal|cpu]6//   forge generate --checkpoint runs/exp1/ckpt_latest.bin --tokenizer data/tinystories/tok4096.model7//                  --prompt "Once upon a time" [--temp 0.8] [--top-k 40] [--max-tokens 256]8//   forge eval     --checkpoint runs/exp1/ckpt_latest.bin --data data/tinystories/val.bin9//                  [--batches 50]10//   forge info     [--config <json>]11#include <Foundation/Foundation.hpp>12#include <Metal/Metal.hpp>1314#include "core/device.h"15#include "core/fmodel.h"16#include "nn/config.h"17#include "nn/transformer.h"18#include "ops/metal/metal_ops.h"19#include "ops/ops.h"20#include "tokenizer/bpe.h"21#include "train/checkpoint.h"22#include "train/dataloader.h"23#include "train/trainer.h"2425#include <cmath>26#include <cstdio>27#include <filesystem>28#include <fstream>29#include <map>30#include <memory>31#include <random>32#include <sstream>33#include <string>34#include <unordered_set>35#include <vector>3637namespace {3839void print_usage() {40    std::printf(41        "usage: forge <command> [options]\n"42        "\n"43        "commands:\n"44        "  train     --config <json> --data <dir> --out <dir> [--resume <ckpt>] [--backend metal|cpu]\n"45        "  generate  --checkpoint <ckpt> --tokenizer <model> --prompt <text>\n"46        "            [--temp t] [--top-k k] [--max-tokens n] [--seed s]\n"47        "  eval      --checkpoint <ckpt> --data <val.bin> [--batches n]\n"48        "  export    --checkpoint <ckpt.bin> --out <model.forge> [--dtype f32|f16|bf16]\n"49        "            [--shard-mb n] [--tag label]   (git-style delta repo, zero-copy load)\n"50        "  info      [--config <json>]\n"51        "\n"52        "generate/eval also accept a .forge repo (or one of its manifests) as --checkpoint.\n");53}5455std::map<std::string, std::string> parse_flags(int argc, char** argv, int start) {56    std::map<std::string, std::string> flags;57    for (int i = start; i < argc; ++i) {58        std::string arg = argv[i];59        if (arg.rfind("--", 0) != 0) continue;60        std::string key = arg.substr(2);61        if (i + 1 < argc && std::string(argv[i + 1]).rfind("--", 0) != 0) {62            flags[key] = argv[++i];63        } else {64            flags[key] = "true";65        }66    }67    return flags;68}6970std::string flag(const std::map<std::string, std::string>& flags, const std::string& key,71                 const std::string& fallback = "") {72    auto it = flags.find(key);73    return it == flags.end() ? fallback : it->second;74}7576void select_backend(const std::string& name) {77    if (name == "cpu") {78        forge::ops::set_backend(forge::ops::Backend::CPU);79    } else {80        forge::ops::set_backend(forge::ops::Backend::Metal);81        forge::Device::get(); // fail fast if Metal init is broken82    }83}8485int cmd_info(const std::map<std::string, std::string>& flags) {86    forge::Device& dev = forge::Device::get();87    std::printf("device: %s\n", dev.name().c_str());88    std::printf("recommended working set: %.1f GB\n",89                double(dev.recommended_working_set()) / (1024.0 * 1024.0 * 1024.0));90    if (auto path = flag(flags, "config"); !path.empty()) {91        forge::Config cfg = forge::load_config(path);92        const auto& m = cfg.model;93        std::printf("\nmodel: %s\n", m.name.c_str());94        std::printf("  layers=%lld d_model=%lld heads=%lld kv_heads=%lld head_dim=%lld\n",95                    m.n_layers, m.d_model, m.n_heads, m.n_kv_heads, m.head_dim());96        std::printf("  d_ff=%lld vocab=%lld context=%lld\n", m.d_ff, m.vocab_size,97                    m.context_length);98        std::printf("  norm=%s act=%s rope=%d tied=%d\n", m.norm.c_str(),99                    m.activation.c_str(), m.use_rope, m.tied_embeddings);100        std::printf("  parameters: %.2fM\n", double(m.num_params()) / 1e6);101    }102    return 0;103}104105int cmd_train(const std::map<std::string, std::string>& flags) {106    const std::string config_path = flag(flags, "config");107    const std::string data_dir = flag(flags, "data");108    const std::string out_dir = flag(flags, "out");109    if (config_path.empty() || data_dir.empty() || out_dir.empty()) {110        print_usage();111        return 1;112    }113    select_backend(flag(flags, "backend", "metal"));114115    std::ifstream in(config_path);116    std::stringstream ss;117    ss << in.rdbuf();118    const std::string config_json = ss.str();119120    forge::Config cfg = forge::load_config(config_path);121    forge::train::Trainer trainer(cfg, data_dir, out_dir, config_json);122    trainer.train(flag(flags, "resume"));123    return 0;124}125126// True when the path is a .forge repo directory or a manifest json inside one.127bool is_forge_repo(const std::string& path) {128    namespace fs = std::filesystem;129    if (fs::is_directory(path)) return fs::exists(fs::path(path) / "manifest-latest.json");130    return path.size() > 5 && path.rfind(".json") == path.size() - 5;131}132133// Build the model described by a .forge snapshot and alias its weights.134// f32 repos are ZERO-COPY: parameters point straight into the mmapped,135// page-aligned shards (unified memory — the file cache feeds the GPU).136std::unique_ptr<forge::nn::Transformer> load_forge_model(const std::string& path,137                                                         forge::Config* out_cfg) {138    forge::fmodel::Snapshot snap = forge::fmodel::Snapshot::open(path);139    nlohmann::json j = nlohmann::json::parse(snap.config_json());140    forge::Config cfg;141    if (j.contains("model")) j.at("model").get_to(cfg.model);142    if (j.contains("train")) j.at("train").get_to(cfg.train);143    auto model = std::make_unique<forge::nn::Transformer>(cfg.model, 0);144145    std::unordered_set<const void*> seen;146    for (const auto& [name, p] : model->named_parameters()) {147        if (!seen.insert(p.id()).second) continue; // tied param: already aliased148        if (!snap.has(name)) {149            std::fprintf(stderr, "forge: %s missing from %s\n", name.c_str(),150                         path.c_str());151            std::exit(1);152        }153        if (snap.shape(name) != p.value().shape()) {154            std::fprintf(stderr, "forge: shape mismatch for %s\n", name.c_str());155            std::exit(1);156        }157        p.value() = snap.tensor_f32(name);158    }159    if (out_cfg) *out_cfg = cfg;160    std::printf("loaded %s (step %lld)\n", path.c_str(),161                static_cast<long long>(snap.step()));162    return model;163}164165int cmd_export(const std::map<std::string, std::string>& flags) {166    const std::string ckpt = flag(flags, "checkpoint");167    const std::string out = flag(flags, "out");168    if (ckpt.empty() || out.empty()) {169        print_usage();170        return 1;171    }172    select_backend(flag(flags, "backend", "metal"));173174    const std::string config_json = forge::train::read_checkpoint_config(ckpt);175    nlohmann::json j = nlohmann::json::parse(config_json);176    forge::ModelConfig mc;177    j.at("model").get_to(mc);178    forge::nn::Transformer model(mc, 0);179    forge::train::CheckpointData meta =180        forge::train::load_checkpoint(ckpt, model.named_parameters(), nullptr);181182    forge::fmodel::SaveOptions opts;183    const std::string dt = flag(flags, "dtype", "f32");184    opts.dtype = dt == "f16" ? forge::DType::F16185               : dt == "bf16" ? forge::DType::BF16186                              : forge::DType::F32;187    opts.shard_mb = std::stoll(flag(flags, "shard-mb", "95"));188    opts.tag = flag(flags, "tag");189    opts.step = meta.step;190    forge::fmodel::save(out, config_json, model.named_parameters(), opts);191    return 0;192}193194int cmd_generate(const std::map<std::string, std::string>& flags) {195    const std::string ckpt = flag(flags, "checkpoint");196    const std::string tok_path = flag(flags, "tokenizer");197    if (ckpt.empty() || tok_path.empty()) {198        print_usage();199        return 1;200    }201    select_backend(flag(flags, "backend", "metal"));202    const float temp = std::stof(flag(flags, "temp", "0.8"));203    const int64_t top_k = std::stoll(flag(flags, "top-k", "40"));204    const int64_t max_tokens = std::stoll(flag(flags, "max-tokens", "256"));205    const uint64_t seed = std::stoull(flag(flags, "seed", "1234"));206207    // model config comes from the checkpoint / .forge repo itself208    std::unique_ptr<forge::nn::Transformer> model_ptr;209    forge::ModelConfig mc;210    if (is_forge_repo(ckpt)) {211        forge::Config cfg;212        model_ptr = load_forge_model(ckpt, &cfg);213        mc = cfg.model;214    } else {215        nlohmann::json j =216            nlohmann::json::parse(forge::train::read_checkpoint_config(ckpt));217        j.at("model").get_to(mc);218        model_ptr = std::make_unique<forge::nn::Transformer>(mc, 0);219        forge::train::load_checkpoint(ckpt, model_ptr->named_parameters(), nullptr);220    }221    forge::nn::Transformer& model = *model_ptr;222223    forge::tok::BPETokenizer tokenizer;224    tokenizer.load(tok_path);225226    std::vector<int32_t> ctx = tokenizer.encode(flag(flags, "prompt", "Once upon a time"));227    if (ctx.empty()) ctx.push_back(0);228    std::printf("%s", tokenizer.decode(ctx).c_str());229    std::fflush(stdout);230231    std::mt19937_64 rng(seed);232    forge::NoGrad ng;233    for (int64_t n = 0; n < max_tokens; ++n) {234        // full-context recompute each token (KV cache lands with M5)235        const int64_t T =236            std::min<int64_t>(int64_t(ctx.size()), mc.context_length);237        forge::Tensor ids = forge::Tensor::empty({1, T}, forge::DType::I32);238        for (int64_t t = 0; t < T; ++t)239            ids.data<int32_t>()[t] = ctx[ctx.size() - size_t(T) + size_t(t)];240241        forge::Var logits = model.forward(ids); // [T, V]242        if (forge::ops::backend() == forge::ops::Backend::Metal) forge::metal::sync();243244        const float* row = logits.value().data<float>() + (T - 1) * mc.vocab_size;245        std::vector<std::pair<float, int32_t>> cand(size_t(mc.vocab_size));246        for (int64_t v = 0; v < mc.vocab_size; ++v) cand[size_t(v)] = {row[v], int32_t(v)};247        const size_t k = size_t(std::min<int64_t>(top_k > 0 ? top_k : mc.vocab_size,248                                                  mc.vocab_size));249        std::partial_sort(cand.begin(), cand.begin() + long(k), cand.end(),250                          [](auto& a, auto& b) { return a.first > b.first; });251252        // temperature softmax over the top-k253        float m = cand[0].first;254        double sum = 0.0;255        std::vector<double> probs(k);256        const float tinv = temp > 0.0f ? 1.0f / temp : 1.0f;257        for (size_t i = 0; i < k; ++i) {258            probs[i] = std::exp(double((cand[i].first - m) * tinv));259            sum += probs[i];260        }261        double r = std::uniform_real_distribution<double>(0.0, sum)(rng);262        int32_t next = cand[0].second;263        for (size_t i = 0; i < k; ++i) {264            r -= probs[i];265            if (r <= 0.0) {266                next = cand[i].second;267                break;268            }269        }270        ctx.push_back(next);271        std::printf("%s", tokenizer.decode({next}).c_str());272        std::fflush(stdout);273    }274    std::printf("\n");275    return 0;276}277278int cmd_eval(const std::map<std::string, std::string>& flags) {279    const std::string ckpt = flag(flags, "checkpoint");280    const std::string data = flag(flags, "data");281    if (ckpt.empty() || data.empty()) {282        print_usage();283        return 1;284    }285    select_backend(flag(flags, "backend", "metal"));286    const int64_t batches = std::stoll(flag(flags, "batches", "50"));287288    std::unique_ptr<forge::nn::Transformer> model_ptr;289    forge::Config cfg;290    if (is_forge_repo(ckpt)) {291        model_ptr = load_forge_model(ckpt, &cfg);292    } else {293        nlohmann::json j =294            nlohmann::json::parse(forge::train::read_checkpoint_config(ckpt));295        if (j.contains("model")) j.at("model").get_to(cfg.model);296        if (j.contains("train")) j.at("train").get_to(cfg.train);297        model_ptr = std::make_unique<forge::nn::Transformer>(cfg.model, 0);298        forge::train::load_checkpoint(ckpt, model_ptr->named_parameters(), nullptr);299    }300    forge::nn::Transformer& model = *model_ptr;301302    forge::train::DataLoader loader(data, cfg.model.context_length, 0);303    const int64_t B = cfg.train.batch_size, T = cfg.model.context_length;304305    forge::NoGrad ng;306    double total = 0.0;307    for (int64_t i = 0; i < batches; ++i) {308        forge::Tensor ids = forge::Tensor::empty({B, T}, forge::DType::I32);309        forge::Tensor targets = forge::Tensor::empty({B * T}, forge::DType::I32);310        loader.seq_batch(i, ids, targets);311        forge::Var loss = model.loss(ids, targets);312        if (forge::ops::backend() == forge::ops::Backend::Metal) forge::metal::sync();313        total += double(loss.value().data<float>()[0]);314    }315    const double mean = total / double(batches);316    std::printf("val loss: %.4f | ppl: %.2f (%lld batches of %lldx%lld)\n", mean,317                std::exp(mean), static_cast<long long>(batches),318                static_cast<long long>(B), static_cast<long long>(T));319    return 0;320}321322} // namespace323324int main(int argc, char** argv) {325    NS::AutoreleasePool* pool = NS::AutoreleasePool::alloc()->init();326327    int rc = 0;328    if (argc < 2) {329        print_usage();330        rc = 1;331    } else {332        const std::string command = argv[1];333        const auto flags = parse_flags(argc, argv, 2);334        if (command == "info") rc = cmd_info(flags);335        else if (command == "train") rc = cmd_train(flags);336        else if (command == "generate") rc = cmd_generate(flags);337        else if (command == "eval") rc = cmd_eval(flags);338        else if (command == "export") rc = cmd_export(flags);339        else {340            print_usage();341            rc = 1;342        }343    }344345    pool->drain();346    return rc;347}348