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//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