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#include "train/trainer.h"34#include "core/fmodel.h"5#include "ops/metal/metal_ops.h"6#include "ops/ops.h"7#include "train/checkpoint.h"8#include "train/scheduler.h"910#include <Foundation/Foundation.hpp>1112#include <algorithm>13#include <chrono>14#include <cstdio>15#include <filesystem>1617namespace forge::train {1819Trainer::Trainer(Config cfg, const std::string& data_dir, const std::string& out_dir,20 const std::string& config_json)21 : cfg_(cfg), out_dir_(out_dir), config_json_(config_json) {22 std::filesystem::create_directories(out_dir_);23 model_ = std::make_unique<nn::Transformer>(cfg_.model, cfg_.train.seed);2425 Optimizer::Options opts;26 opts.kind = cfg_.train.optimizer;27 opts.beta1 = cfg_.train.beta1;28 opts.beta2 = cfg_.train.beta2;29 opts.eps = cfg_.train.eps;30 opts.weight_decay = cfg_.train.weight_decay;31 opts.muon_momentum = cfg_.train.muon_momentum;32 opts.muon_lr_ratio = cfg_.train.muon_lr / cfg_.train.lr;33 opt_ = std::make_unique<Optimizer>(model_->named_parameters(), opts);3435 train_data_ = std::make_unique<DataLoader>(data_dir + "/train.bin",36 cfg_.model.context_length,37 cfg_.train.seed);38 const std::string val_path = data_dir + "/val.bin";39 if (std::filesystem::exists(val_path)) {40 val_data_ = std::make_unique<DataLoader>(val_path, cfg_.model.context_length,41 cfg_.train.seed + 1);42 }43}4445float Trainer::eval_loss() {46 if (!val_data_) return -1.0f;47 NoGrad ng;48 const int64_t B = cfg_.train.batch_size, T = cfg_.model.context_length;49 double total = 0.0;50 for (int64_t i = 0; i < cfg_.train.eval_batches; ++i) {51 Tensor ids = Tensor::empty({B, T}, DType::I32);52 Tensor targets = Tensor::empty({B * T}, DType::I32);53 val_data_->seq_batch(i, ids, targets);54 Var loss = model_->loss(ids, targets);55 if (ops::backend() == ops::Backend::Metal) metal::sync();56 total += double(loss.value().data<float>()[0]);57 }58 return float(total / double(cfg_.train.eval_batches));59}6061void Trainer::save(int64_t step) {62 CheckpointData meta;63 meta.step = step;64 meta.rng_state = uint64_t(step); // dataloader reseeded from step on resume65 meta.config_json = config_json_;66 char name[64];67 std::snprintf(name, sizeof(name), "/ckpt_%06lld.bin", static_cast<long long>(step));68 save_checkpoint(out_dir_ + name, model_->named_parameters(), opt_.get(), meta);69 save_checkpoint(out_dir_ + "/ckpt_latest.bin", model_->named_parameters(), opt_.get(),70 meta);71 std::printf("checkpoint saved: %s\n", (out_dir_ + name).c_str());7273 // Native .forge commit: the run's whole weight history lives in one74 // git-style repo; content addressing means an unchanged tensor (frozen,75 // masked, tied) is never rewritten.76 if (cfg_.train.forge_save) {77 fmodel::SaveOptions fo;78 fo.dtype = cfg_.train.forge_dtype == "f16" ? DType::F1679 : cfg_.train.forge_dtype == "bf16" ? DType::BF1680 : DType::F32;81 char tag[32];82 std::snprintf(tag, sizeof(tag), "step-%06lld", static_cast<long long>(step));83 fo.tag = tag;84 fo.step = step;85 fmodel::save(out_dir_ + "/model.forge", config_json_,86 model_->named_parameters(), fo);87 }88}8990void Trainer::train(const std::string& resume_from) {91 const auto& tc = cfg_.train;92 const int64_t B = tc.batch_size, T = cfg_.model.context_length;9394 int64_t start_step = 0;95 if (!resume_from.empty()) {96 CheckpointData meta =97 load_checkpoint(resume_from, model_->named_parameters(), opt_.get());98 start_step = meta.step;99 std::printf("resumed from %s at step %lld\n", resume_from.c_str(),100 static_cast<long long>(start_step));101 }102103 FILE* csv = std::fopen((out_dir_ + "/log.csv").c_str(),104 start_step > 0 ? "ab" : "wb");105 if (csv && start_step == 0)106 std::fprintf(csv, "step,loss,lr,grad_norm,tokens_per_sec,val_loss,elapsed_s\n");107 const auto run_t0 = std::chrono::steady_clock::now();108109 const int64_t min_lr_steps = tc.max_steps;110 const float min_lr = tc.lr * tc.min_lr_ratio;111 const int64_t tokens_per_step = B * T * tc.grad_accum_steps;112 const bool wsd = tc.schedule == "wsd";113 const int64_t wsd_decay_steps =114 std::max<int64_t>(1, int64_t(float(tc.max_steps) * tc.wsd_decay_frac));115116 std::printf("training %s: %lld params, %lld steps, %lld tokens/step, backend=%s, "117 "opt=%s, sched=%s\n",118 cfg_.model.name.c_str(),119 static_cast<long long>(cfg_.model.num_params()),120 static_cast<long long>(tc.max_steps),121 static_cast<long long>(tokens_per_step),122 ops::backend() == ops::Backend::Metal ? "metal" : "cpu",123 tc.optimizer.c_str(), tc.schedule.c_str());124125 for (int64_t step = start_step; step < tc.max_steps; ++step) {126 NS::AutoreleasePool* pool = NS::AutoreleasePool::alloc()->init();127 const auto t0 = std::chrono::steady_clock::now();128 const float lr = wsd129 ? lr_wsd(step, tc.lr, min_lr, tc.warmup_steps, tc.max_steps, wsd_decay_steps)130 : lr_at(step, tc.lr, min_lr, tc.warmup_steps, min_lr_steps);131132 opt_->zero_grad();133 const bool on_gpu = ops::backend() == ops::Backend::Metal;134 double loss_val = 0.0;135 for (int64_t micro = 0; micro < tc.grad_accum_steps; ++micro) {136 Tensor ids = Tensor::empty({B, T}, DType::I32);137 Tensor targets = Tensor::empty({B * T}, DType::I32);138 train_data_->next_batch(ids, targets);139 Var loss = model_->loss(ids, targets);140 // scale so accumulated grads average over micro-batches141 Var scaled = ops::scale(loss, 1.0f / float(tc.grad_accum_steps));142 Tape::get().backward(scaled);143 // Sync per micro-batch, not per step: pooled buffers freed while a144 // command buffer is open sit on the allocator's retire list until145 // the next sync, so without this every micro-batch's activations146 // stay resident and peak memory scales with grad_accum_steps.147 if (on_gpu) metal::sync();148 loss_val += double(loss.value().data<float>()[0]) /149 double(tc.grad_accum_steps);150 }151152 const float grad_norm = opt_->step_with_clip(lr, tc.grad_clip);153 if (cfg_.model.n_experts > 0 && cfg_.model.moe_bias_gamma > 0.0f)154 model_->update_moe_bias(cfg_.model.moe_bias_gamma);155156 const auto t1 = std::chrono::steady_clock::now();157 const double dt = std::chrono::duration<double>(t1 - t0).count();158 const double tps = double(tokens_per_step) / dt;159160 float val = -1.0f;161 if (tc.eval_every > 0 && (step + 1) % tc.eval_every == 0) val = eval_loss();162163 if (step < 10 || step % 10 == 0 || val >= 0.0f) {164 std::printf("step %6lld | loss %.4f | lr %.2e | gnorm %.3f | %.0f tok/s",165 static_cast<long long>(step), loss_val, double(lr),166 double(grad_norm), tps);167 if (val >= 0.0f) std::printf(" | val %.4f", double(val));168 std::printf("\n");169 std::fflush(stdout);170 }171 if (csv) {172 const double elapsed =173 std::chrono::duration<double>(t1 - run_t0).count();174 std::fprintf(csv, "%lld,%.6f,%.6e,%.6f,%.1f,%.6f,%.3f\n",175 static_cast<long long>(step), loss_val, double(lr),176 double(grad_norm), tps, double(val), elapsed);177 std::fflush(csv);178 }179180 if (tc.checkpoint_every > 0 && (step + 1) % tc.checkpoint_every == 0)181 save(step + 1);182183 pool->drain();184 }185186 save(tc.max_steps);187 if (csv) std::fclose(csv);188}189190} // namespace forge::train191