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.9 KB · 191 lines cpp
Raw Blame History
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