// Author: Simon-Pierre Boucher — contact@spboucher.ai #include "train/checkpoint.h" #include #include #include #include #include namespace forge::train { namespace { constexpr uint32_t kMagic = 0x45475246; // "FRGE" constexpr uint32_t kVersion = 1; [[noreturn]] void die(const std::string& msg) { std::fprintf(stderr, "forge/checkpoint: %s\n", msg.c_str()); std::abort(); } void write_bytes(FILE* f, const void* p, size_t n) { if (std::fwrite(p, 1, n, f) != n) die("write failed"); } void read_bytes(FILE* f, void* p, size_t n) { if (std::fread(p, 1, n, f) != n) die("read failed (truncated checkpoint?)"); } template void write_pod(FILE* f, T v) { write_bytes(f, &v, sizeof(T)); } template T read_pod(FILE* f) { T v; read_bytes(f, &v, sizeof(T)); return v; } void write_str(FILE* f, const std::string& s) { write_pod(f, uint32_t(s.size())); write_bytes(f, s.data(), s.size()); } std::string read_str(FILE* f) { const uint32_t n = read_pod(f); std::string s(n, '\0'); read_bytes(f, s.data(), n); return s; } // Tied params appear under several names; store each storage once (first // name wins) and load by name into whichever Var carries it. std::vector> dedupe( const std::vector>& named) { std::vector> out; std::unordered_set seen; for (const auto& [name, p] : named) if (seen.insert(p.id()).second) out.emplace_back(name, p); return out; } } // namespace void save_checkpoint(const std::string& path, const std::vector>& named_params, AdamW* opt, const CheckpointData& meta) { const std::string tmp = path + ".tmp"; FILE* f = std::fopen(tmp.c_str(), "wb"); if (!f) die("cannot open " + tmp); write_pod(f, kMagic); write_pod(f, kVersion); write_pod(f, meta.step); write_pod(f, opt ? opt->t() : 0); write_pod(f, meta.rng_state); write_str(f, meta.config_json); const auto params = dedupe(named_params); write_pod(f, uint32_t(params.size())); for (const auto& [name, p] : params) { write_str(f, name); const auto& shape = p.value().shape(); write_pod(f, uint32_t(shape.size())); for (int64_t d : shape) write_pod(f, d); write_bytes(f, p.value().raw(), p.value().nbytes()); } const bool has_opt = opt && !opt->m().empty(); write_pod(f, has_opt ? uint32_t(opt->m().size()) : 0u); if (has_opt) { for (size_t i = 0; i < opt->m().size(); ++i) { write_bytes(f, opt->m()[i].raw(), opt->m()[i].nbytes()); write_bytes(f, opt->v()[i].raw(), opt->v()[i].nbytes()); } } std::fclose(f); if (std::rename(tmp.c_str(), path.c_str()) != 0) die("rename failed: " + path); } CheckpointData load_checkpoint(const std::string& path, const std::vector>& named_params, AdamW* opt) { FILE* f = std::fopen(path.c_str(), "rb"); if (!f) die("cannot open " + path); if (read_pod(f) != kMagic) die("bad magic: " + path); if (read_pod(f) != kVersion) die("unsupported version: " + path); CheckpointData meta; meta.step = read_pod(f); const int64_t opt_t = read_pod(f); meta.rng_state = read_pod(f); meta.config_json = read_str(f); std::unordered_map by_name; for (const auto& [name, p] : dedupe(named_params)) by_name.emplace(name, p); const uint32_t n_params = read_pod(f); if (n_params != by_name.size()) die("parameter count mismatch"); for (uint32_t i = 0; i < n_params; ++i) { const std::string name = read_str(f); const uint32_t ndim = read_pod(f); std::vector shape(ndim); for (auto& d : shape) d = read_pod(f); auto it = by_name.find(name); if (it == by_name.end()) die("unknown parameter in checkpoint: " + name); if (it->second.value().shape() != shape) die("shape mismatch: " + name); read_bytes(f, it->second.value().raw(), it->second.value().nbytes()); } const uint32_t n_opt = read_pod(f); if (n_opt > 0 && opt) { opt->ensure_state(); if (opt->m().size() != n_opt) die("optimizer state count mismatch"); for (uint32_t i = 0; i < n_opt; ++i) { read_bytes(f, opt->m()[i].raw(), opt->m()[i].nbytes()); read_bytes(f, opt->v()[i].raw(), opt->v()[i].nbytes()); } opt->set_t(opt_t); } else if (n_opt > 0) { // skip optimizer state std::fseek(f, 0, SEEK_END); } std::fclose(f); return meta; } std::string read_checkpoint_config(const std::string& path) { FILE* f = std::fopen(path.c_str(), "rb"); if (!f) die("cannot open " + path); if (read_pod(f) != kMagic) die("bad magic: " + path); if (read_pod(f) != kVersion) die("unsupported version: " + path); read_pod(f); read_pod(f); read_pod(f); const std::string cfg = read_str(f); std::fclose(f); return cfg; } } // namespace forge::train