SPB Git

spb/forge-studio Public

The Instruments of LLM training — a native macOS cockpit for Forge. Train language models from scratch on Apple Silicon without a terminal.

Swift 95.7% Shell 4.3%
12.9 KB · 296 lines swift
Raw Blame History
1// Author: Simon-Pierre Boucher — contact@spboucher.ai2//3// Codable mirror of Forge's JSON config (src/nn/config.h, RESEARCH.md §2).4// Field names are byte-compatible with the C++ parser; defaults match the5// C++ defaults so a partially-specified JSON round-trips identically.6import Foundation78struct ModelConfig: Codable, Equatable {9    var name = "model"10    var nLayers = 611    var dModel = 38412    var nHeads = 613    var nKvHeads = 614    var dFf = 102415    var vocabSize = 409616    var contextLength = 51217    var tiedEmbeddings = true18    var useRope = true19    var ropeTheta = 10000.020    var norm = "rmsnorm"          // rmsnorm | layernorm21    var normEps = 1e-622    var activation = "swiglu"     // swiglu | gelu | relu223    var dropout = 0.024    var quant = "none"            // none | int8 | ternary25    var qkNorm = false26    var finalSoftcap = 0.027    var scaleEmbeddings = false28    var attentionBias = false29    var headDimOverride = 0       // JSON key "head_dim"; 0 = d_model/n_heads30    var nopeEvery = 031    var normPlacement = "pre"     // pre | post | sandwich32    var ropeScaleFactor = 0.033    var ropeScaleLow = 1.034    var ropeScaleHigh = 4.035    var ropeScaleOrigCtx = 819236    var slidingWindow = 037    var slidingGlobalEvery = 038    var ropeThetaGlobal = 0.039    var attnSoftcap = 0.040    var nExperts = 041    var moeTopK = 242    var moeAuxWeight = 0.0143    var nSharedExperts = 044    var moeScoring = "softmax"    // softmax | sigmoid45    var moeNormTopk = true46    var routedScalingFactor = 1.047    var moeDFf = 048    var firstKDense = 049    var moeBiasGamma = 0.05051    enum CodingKeys: String, CodingKey {52        case name53        case nLayers = "n_layers", dModel = "d_model", nHeads = "n_heads"54        case nKvHeads = "n_kv_heads", dFf = "d_ff", vocabSize = "vocab_size"55        case contextLength = "context_length", tiedEmbeddings = "tied_embeddings"56        case useRope = "use_rope", ropeTheta = "rope_theta", norm57        case normEps = "norm_eps", activation, dropout, quant58        case qkNorm = "qk_norm", finalSoftcap = "final_softcap"59        case scaleEmbeddings = "scale_embeddings", attentionBias = "attention_bias"60        case headDimOverride = "head_dim", nopeEvery = "nope_every"61        case normPlacement = "norm_placement"62        case ropeScaleFactor = "rope_scale_factor", ropeScaleLow = "rope_scale_low"63        case ropeScaleHigh = "rope_scale_high", ropeScaleOrigCtx = "rope_scale_orig_ctx"64        case slidingWindow = "sliding_window"65        case slidingGlobalEvery = "sliding_global_every"66        case ropeThetaGlobal = "rope_theta_global", attnSoftcap = "attn_softcap"67        case nExperts = "n_experts", moeTopK = "moe_top_k"68        case moeAuxWeight = "moe_aux_weight", nSharedExperts = "n_shared_experts"69        case moeScoring = "moe_scoring", moeNormTopk = "moe_norm_topk"70        case routedScalingFactor = "routed_scaling_factor", moeDFf = "moe_d_ff"71        case firstKDense = "first_k_dense", moeBiasGamma = "moe_bias_gamma"72    }7374    // Forge tolerates missing keys everywhere — mirror that.75    init() {}76    init(from decoder: Decoder) throws {77        let c = try decoder.container(keyedBy: CodingKeys.self)78        func g<T: Decodable>(_ k: CodingKeys, _ d: T) -> T {79            (try? c.decodeIfPresent(T.self, forKey: k)) as? T ?? d80        }81        name = g(.name, name); nLayers = g(.nLayers, nLayers)82        dModel = g(.dModel, dModel); nHeads = g(.nHeads, nHeads)83        nKvHeads = g(.nKvHeads, nHeads); dFf = g(.dFf, dFf)84        vocabSize = g(.vocabSize, vocabSize)85        contextLength = g(.contextLength, contextLength)86        tiedEmbeddings = g(.tiedEmbeddings, tiedEmbeddings)87        useRope = g(.useRope, useRope); ropeTheta = g(.ropeTheta, ropeTheta)88        norm = g(.norm, norm); normEps = g(.normEps, normEps)89        activation = g(.activation, activation); dropout = g(.dropout, dropout)90        quant = g(.quant, quant); qkNorm = g(.qkNorm, qkNorm)91        finalSoftcap = g(.finalSoftcap, finalSoftcap)92        scaleEmbeddings = g(.scaleEmbeddings, scaleEmbeddings)93        attentionBias = g(.attentionBias, attentionBias)94        headDimOverride = g(.headDimOverride, headDimOverride)95        nopeEvery = g(.nopeEvery, nopeEvery)96        normPlacement = g(.normPlacement, normPlacement)97        ropeScaleFactor = g(.ropeScaleFactor, ropeScaleFactor)98        ropeScaleLow = g(.ropeScaleLow, ropeScaleLow)99        ropeScaleHigh = g(.ropeScaleHigh, ropeScaleHigh)100        ropeScaleOrigCtx = g(.ropeScaleOrigCtx, ropeScaleOrigCtx)101        slidingWindow = g(.slidingWindow, slidingWindow)102        slidingGlobalEvery = g(.slidingGlobalEvery, slidingGlobalEvery)103        ropeThetaGlobal = g(.ropeThetaGlobal, ropeThetaGlobal)104        attnSoftcap = g(.attnSoftcap, attnSoftcap)105        nExperts = g(.nExperts, nExperts); moeTopK = g(.moeTopK, moeTopK)106        moeAuxWeight = g(.moeAuxWeight, moeAuxWeight)107        nSharedExperts = g(.nSharedExperts, nSharedExperts)108        moeScoring = g(.moeScoring, moeScoring)109        moeNormTopk = g(.moeNormTopk, moeNormTopk)110        routedScalingFactor = g(.routedScalingFactor, routedScalingFactor)111        moeDFf = g(.moeDFf, moeDFf); firstKDense = g(.firstKDense, firstKDense)112        moeBiasGamma = g(.moeBiasGamma, moeBiasGamma)113    }114115    var headDim: Int { headDimOverride > 0 ? headDimOverride : dModel / max(nHeads, 1) }116117    // Same formula as ModelConfig::num_params() — cross-checked by118    // ParamCountTests against `forge info`.119    var paramCount: Int {120        let hd = headDim121        var attn = dModel * nHeads * hd + 2 * dModel * nKvHeads * hd122            + nHeads * hd * dModel123        if attentionBias { attn += (nHeads + 2 * nKvHeads) * hd }124        let actMats = activation == "swiglu" ? 3 : 2125        let mlpDense = actMats * dModel * dFf126        let expertDff = (nExperts > 0 && moeDFf > 0) ? moeDFf : dFf127        let mlpMoe = (nExperts + nSharedExperts) * actMats * dModel * expertDff128            + nExperts * dModel129        let nMoeLayers = nExperts > 0 ? nLayers - min(firstKDense, nLayers) : 0130        let mlpTotal = nMoeLayers * mlpMoe + (nLayers - nMoeLayers) * mlpDense131        let normsPerLayer = normPlacement == "sandwich" ? 4 : 2132        var norms = (norm == "layernorm" ? 2 : 1) * dModel133            * (normsPerLayer * nLayers + 1)134        if qkNorm { norms += 2 * hd * nLayers }135        var total = nLayers * attn + mlpTotal + norms + vocabSize * dModel136        if !tiedEmbeddings { total += vocabSize * dModel }137        if !useRope { total += contextLength * dModel }138        return total139    }140141    // Mirrors the C++ parse-time validation; returns human-actionable errors.142    var validationErrors: [String] {143        var e: [String] = []144        if headDimOverride == 0 && nHeads > 0 && dModel % nHeads != 0 {145            e.append("d_model doit être divisible par n_heads (ou fixer head_dim)")146        }147        if nKvHeads > 0 && nHeads % nKvHeads != 0 {148            e.append("n_heads doit être divisible par n_kv_heads")149        }150        if headDim % 2 != 0 { e.append("head_dim doit être pair (paires RoPE)") }151        if !["rmsnorm", "layernorm"].contains(norm) { e.append("norm invalide") }152        if !["swiglu", "gelu", "relu2"].contains(activation) {153            e.append("activation invalide")154        }155        if !["pre", "post", "sandwich"].contains(normPlacement) {156            e.append("norm_placement invalide")157        }158        if !["none", "int8", "ternary"].contains(quant) { e.append("quant invalide") }159        if nExperts > 0 && !(1...nExperts).contains(moeTopK) {160            e.append("moe_top_k doit être dans [1, n_experts]")161        }162        if nSharedExperts > 0 && nExperts == 0 {163            e.append("n_shared_experts requiert n_experts > 0")164        }165        if firstKDense < 0 || firstKDense > nLayers {166            e.append("first_k_dense doit être dans [0, n_layers]")167        }168        for (v, n) in [(nLayers, "n_layers"), (dModel, "d_model"), (nHeads, "n_heads"),169                       (dFf, "d_ff"), (vocabSize, "vocab_size"),170                       (contextLength, "context_length")] where v <= 0 {171            e.append("\(n) doit être > 0")172        }173        return e174    }175}176177struct TrainConfig: Codable, Equatable {178    var lr = 6e-4179    var minLrRatio = 0.1180    var warmupSteps = 2000181    var maxSteps = 100_000182    var schedule = "cosine"       // cosine | wsd183    var wsdDecayFrac = 0.15184    var optimizer = "adamw"       // adamw | muon185    var muonLr = 0.02186    var muonMomentum = 0.95187    var beta1 = 0.9188    var beta2 = 0.95189    var eps = 1e-8190    var weightDecay = 0.1191    var gradClip = 1.0192    var batchSize = 32193    var gradAccumSteps = 1194    var precision = "f32"195    var checkpointEvery = 1000196    var forgeSave = true197    var forgeDtype = "f32"198    var evalEvery = 500199    var evalBatches = 20200    var seed = 1337201    var deterministic = false202203    enum CodingKeys: String, CodingKey {204        case lr, minLrRatio = "min_lr_ratio", warmupSteps = "warmup_steps"205        case maxSteps = "max_steps", schedule, wsdDecayFrac = "wsd_decay_frac"206        case optimizer, muonLr = "muon_lr", muonMomentum = "muon_momentum"207        case beta1, beta2, eps, weightDecay = "weight_decay"208        case gradClip = "grad_clip", batchSize = "batch_size"209        case gradAccumSteps = "grad_accum_steps", precision210        case checkpointEvery = "checkpoint_every", forgeSave = "forge_save"211        case forgeDtype = "forge_dtype", evalEvery = "eval_every"212        case evalBatches = "eval_batches", seed, deterministic213    }214215    init() {}216    init(from decoder: Decoder) throws {217        let c = try decoder.container(keyedBy: CodingKeys.self)218        func g<T: Decodable>(_ k: CodingKeys, _ d: T) -> T {219            (try? c.decodeIfPresent(T.self, forKey: k)) as? T ?? d220        }221        lr = g(.lr, lr); minLrRatio = g(.minLrRatio, minLrRatio)222        warmupSteps = g(.warmupSteps, warmupSteps); maxSteps = g(.maxSteps, maxSteps)223        schedule = g(.schedule, schedule)224        wsdDecayFrac = g(.wsdDecayFrac, wsdDecayFrac)225        optimizer = g(.optimizer, optimizer); muonLr = g(.muonLr, muonLr)226        muonMomentum = g(.muonMomentum, muonMomentum)227        beta1 = g(.beta1, beta1); beta2 = g(.beta2, beta2); eps = g(.eps, eps)228        weightDecay = g(.weightDecay, weightDecay); gradClip = g(.gradClip, gradClip)229        batchSize = g(.batchSize, batchSize)230        gradAccumSteps = g(.gradAccumSteps, gradAccumSteps)231        precision = g(.precision, precision)232        checkpointEvery = g(.checkpointEvery, checkpointEvery)233        forgeSave = g(.forgeSave, forgeSave); forgeDtype = g(.forgeDtype, forgeDtype)234        evalEvery = g(.evalEvery, evalEvery); evalBatches = g(.evalBatches, evalBatches)235        seed = g(.seed, seed); deterministic = g(.deterministic, deterministic)236    }237238    var validationErrors: [String] {239        var e: [String] = []240        if lr <= 0 { e.append("lr doit être > 0") }241        if warmupSteps > maxSteps { e.append("warmup_steps ≤ max_steps requis") }242        if !["cosine", "wsd"].contains(schedule) { e.append("schedule invalide") }243        if !["adamw", "muon"].contains(optimizer) { e.append("optimizer invalide") }244        if !(0.0..<1.0).contains(wsdDecayFrac) || wsdDecayFrac <= 0 {245            e.append("wsd_decay_frac doit être dans (0, 1)")246        }247        if !["f32", "f16", "bf16"].contains(forgeDtype) {248            e.append("forge_dtype invalide")249        }250        if batchSize <= 0 || gradAccumSteps <= 0 || maxSteps <= 0 {251            e.append("batch/accum/max_steps doivent être > 0")252        }253        return e254    }255256    // LR schedule preview (same math as src/train/scheduler.h).257    func lrAt(step: Int) -> Double {258        let minLr = lr * minLrRatio259        if step < warmupSteps {260            return lr * Double(step + 1) / Double(warmupSteps + 1)261        }262        if schedule == "wsd" {263            let decaySteps = max(1, Int(Double(maxSteps) * wsdDecayFrac))264            let decayStart = maxSteps - decaySteps265            if step < decayStart { return lr }266            if step >= maxSteps { return minLr }267            let ratio = Double(step - decayStart) / Double(decaySteps)268            return minLr + (lr - minLr) * (1.0 - ratio.squareRoot())269        }270        if step >= maxSteps { return minLr }271        let ratio = Double(step - warmupSteps) / Double(maxSteps - warmupSteps)272        return minLr + 0.5 * (1.0 + cos(.pi * ratio)) * (lr - minLr)273    }274}275276struct ForgeConfig: Codable, Equatable {277    var model = ModelConfig()278    var train = TrainConfig()279280    var tokensPerStep: Int { train.batchSize * model.contextLength * train.gradAccumSteps }281    var totalTokens: Int { tokensPerStep * train.maxSteps }282    var validationErrors: [String] {283        model.validationErrors + train.validationErrors284    }285286    static func load(from url: URL) throws -> ForgeConfig {287        try JSONDecoder().decode(ForgeConfig.self, from: Data(contentsOf: url))288    }289290    func exportJSON() throws -> Data {291        let enc = JSONEncoder()292        enc.outputFormatting = [.prettyPrinted, .sortedKeys]293        return try enc.encode(self)294    }295}296