SPB Git

spb/zyquo-mlx Public MIT

The local MLX foundry for your Mac — run, fine-tune, quantize, and ship models. Nothing leaves your machine.

Swift 93.4% Python 3.8% Makefile 2.2% Shell 0.5%
13.2 KB · 294 lines swift
Raw Blame History
1//2//  CLIFoundry.swift3//  Zyquo MLX4//5//  Author: Simon-Pierre Boucher6//  Mail: contact@spboucher.ai7//89import Foundation1011/// CLI drivers for the Phase 3 foundry services (training, fusion, datasets).12///13///     ZyquoMLX --train --model <dir> --dataset <jsonl-or-dataset-dir> [--iters N] [--method lora|qlora|dora|full]14///     ZyquoMLX --fuse --model <dir> --adapter <dir> --out <name> [--dequantize]15///     ZyquoMLX --validate-dataset <file.jsonl>16extension CLI {1718    static var shouldRunFoundry: Bool {19        let args = CommandLine.arguments20        return args.contains("--train") || args.contains("--fuse")21            || args.contains("--quantize") || args.contains("--validate-dataset")22            || args.contains("--download") || args.contains("--transcribe")23    }2425    static func runFoundry() async -> Int32 {26        do {27            let args = CommandLine.arguments28            if args.contains("--train") {29                try await train(args: args)30                return 031            }32            if args.contains("--fuse") {33                try await fuse(args: args)34                return 035            }36            if args.contains("--quantize") {37                try await quantize(args: args)38                return 039            }40            if let path = value(after: "--validate-dataset", in: args) {41                try await validateDataset(path: path)42                return 043            }44            if let repo = value(after: "--download", in: args) {45                try await download(repo: repo)46                return 047            }48            if let modelPath = value(after: "--transcribe", in: args),49                let audioPath = value(after: "--audio", in: args)50            {51                let model = try await ModelStore.shared.describe(52                    directory: URL(fileURLWithPath: (modelPath as NSString).expandingTildeInPath))53                let result = try await SpeechService.shared.transcribe(54                    model: model,55                    audio: URL(fileURLWithPath: (audioPath as NSString).expandingTildeInPath))56                print("language: \(result.language ?? "?")  segments: \(result.segments)  took \(String(format: "%.1f", result.duration))s")57                print(result.text)58                return 059            }60            return 261        } catch {62            FileHandle.standardError.write(Data("error: \(error.localizedDescription)\n".utf8))63            return 164        }65    }6667    // MARK: - Train6869    private static func train(args: [String]) async throws {70        guard let modelPath = value(after: "--model", in: args),71            let dataPath = value(after: "--dataset", in: args)72        else {73            throw CLIError.usage("--train requires --model <dir> and --dataset <jsonl|dir>")74        }7576        let modelURL = URL(fileURLWithPath: (modelPath as NSString).expandingTildeInPath)77        let model = try await ModelStore.shared.describe(directory: modelURL)7879        // Accept a raw JSONL (imported on the fly) or an existing dataset dir.80        let dataURL = URL(fileURLWithPath: (dataPath as NSString).expandingTildeInPath)81        let dataset: Dataset82        if dataURL.pathExtension == "jsonl" {83            let report = try await DatasetService.shared.importJSONL(84                from: dataURL, name: dataURL.deletingPathExtension().lastPathComponent)85            dataset = report.dataset86            print("dataset: \(dataset.name) [\(dataset.format.displayName)] "87                + "train=\(dataset.trainCount) valid=\(dataset.validCount)"88                + (report.issues.isEmpty ? "" : "  (\(report.issues.count) rows skipped)"))89            for issue in report.issues.prefix(5) {90                print("  line \(issue.line): \(issue.problem)\(issue.fix)")91            }92        } else {93            dataset = try PersistenceService.loadJSON(94                Dataset.self, from: dataURL.appendingPathComponent("dataset.json"))95        }9697        var hp = HyperParams()98        hp.iterations = value(after: "--iters", in: args).flatMap(Int.init) ?? 10099        hp.batchSize = value(after: "--batch-size", in: args).flatMap(Int.init) ?? 2100        hp.saveEvery = value(after: "--save-every", in: args).flatMap(Int.init) ?? 50101        hp.stepsPerEval = value(after: "--steps-per-eval", in: args).flatMap(Int.init) ?? 50102        hp.stepsPerReport = 5103104        let methodName = value(after: "--method", in: args) ?? (model.quantization != nil ? "qlora" : "lora")105        guard let method = FineTuneMethod(rawValue: methodName) else {106            throw CLIError.usage("unknown --method \(methodName)")107        }108109        let service = TrainingService.shared110        let resume: Bool111        let run: TrainingRun112        if let resumeID = value(after: "--resume-run", in: args) {113            guard let id = UUID(uuidString: resumeID) else { throw CLIError.usage("bad run id") }114            run = try await RunStore.shared.load(id: id)115            resume = true116        } else {117            run = try await service.createRun(118                name: "cli-\(model.name)-\(method.rawValue)",119                baseModel: model, dataset: dataset, method: method, hyperParams: hp)120            resume = false121        }122        print("run:     \(run.id.uuidString)\(resume ? " (warm resume)" : "")")123        print("method:  \(run.method.displayName) on \(model.name)")124        print("verdict: \(MemoryAdvisor.trainingVerdict(for: model, method: run.method, params: run.hyperParams).displayName)\n")125126        // Test hook mirroring the UI's Cancel button (drives the same127        // TrainingService.cancel() path).128        if let cancelAfter = value(after: "--cancel-after", in: args).flatMap(Double.init) {129            Task {130                try? await Task.sleep(for: .seconds(cancelAfter))131                print("\n[--cancel-after \(cancelAfter)s] cancelling…")132                await service.cancel()133            }134        }135136        let events = try await service.start(run: run, baseModel: model, dataset: dataset, resume: resume)137        for await event in events {138            switch event {139            case .started(let model, let iterations):140                print("training \(model) for \(iterations) iterations…")141            case .metric(let m):142                if let loss = m.trainLoss {143                    let speed = m.tokensPerSecond.map { String(format: "%.0f tok/s", $0) } ?? ""144                    let mem = m.peakMemoryGB.map { String(format: "peak %.1f GB", $0) } ?? ""145                    print("iter \(m.iteration): train loss \(String(format: "%.3f", loss))  \(speed)  \(mem)")146                }147                if let loss = m.valLoss {148                    print("iter \(m.iteration): VAL loss \(String(format: "%.3f", loss))")149                }150            case .checkpointSaved(let name, _):151                print("checkpoint saved: \(name)")152            case .finished:153                print("\ntraining finished ✅")154            case .failed(let message):155                print("\ntraining FAILED: \(message)")156            }157        }158159        let finished = try await RunStore.shared.load(id: run.id)160        print("state: \(finished.state.rawValue), adapters: \(await RunStore.shared.latestAdapter(for: finished)?.path ?? "none")")161    }162163    // MARK: - Fuse164165    private static func fuse(args: [String]) async throws {166        guard let modelPath = value(after: "--model", in: args),167            let adapterPath = value(after: "--adapter", in: args),168            let outName = value(after: "--out", in: args)169        else {170            throw CLIError.usage("--fuse requires --model <dir>, --adapter <dir>, --out <name>")171        }172        let model = try await ModelStore.shared.describe(173            directory: URL(fileURLWithPath: (modelPath as NSString).expandingTildeInPath))174        let adapter = URL(fileURLWithPath: (adapterPath as NSString).expandingTildeInPath)175176        // Fusing into a quantized base re-quantizes the merged weights, which177        // rounds away small LoRA deltas (verified empirically — the adapter's178        // behavior vanished). Default to de-quantizing for quantized bases;179        // --no-dequantize opts back into the lossy compact form.180        let dequantize =181            args.contains("--dequantize")182            || (model.quantization != nil && !args.contains("--no-dequantize"))183        if dequantize && model.quantization != nil {184            print("note: de-quantizing while fusing (quantized base) — re-quantize afterwards in Convert if needed")185        }186        let events = try await ConversionService.shared.fuse(187            baseModel: model, adapterDirectory: adapter, outputName: outName,188            dequantize: dequantize)189        for await event in events {190            switch event {191            case .stage(let stage, _):192                print("stage: \(stage)")193            case .finished(let output):194                print("fused model → \(output.path) ✅")195            case .failed(let message):196                throw CLIError.usage("fuse failed: \(message)")197            }198        }199    }200201    // MARK: - Quantize (Swift-native, docs/MLX-RESEARCH.md §2)202203    private static func quantize(args: [String]) async throws {204        guard let modelPath = value(after: "--quantize", in: args),205            let outName = value(after: "--out", in: args)206        else {207            throw CLIError.usage("--quantize <model-dir> --out <name> [--bits N] [--group-size G]")208        }209        let model = try await ModelStore.shared.describe(210            directory: URL(fileURLWithPath: (modelPath as NSString).expandingTildeInPath))211212        var config = QuantConfig()213        config.bits = value(after: "--bits", in: args).flatMap(Int.init) ?? 4214        config.groupSize = value(after: "--group-size", in: args).flatMap(Int.init) ?? 64215216        if let params = model.parameterCount {217            let predicted = config.predictedWeightBytes(parameterCount: params)218            print("quantizing \(model.name)\(config.label)")219            print("size preview: \(ByteCountFormatter.string(fromByteCount: model.weightsSize, countStyle: .file)) → ~\(ByteCountFormatter.string(fromByteCount: predicted, countStyle: .file))")220        }221222        let events = await ConversionService.shared.quantize(223            model: model, config: config, outputName: outName)224        for await event in events {225            switch event {226            case .stage(let stage, let fraction):227                let pct = fraction.map { String(format: " %.0f%%", $0 * 100) } ?? ""228                print("stage: \(stage)\(pct)")229            case .finished(let output):230                let result = try await ModelStore.shared.describe(directory: output)231                print("quantized model → \(output.path)")232                print("actual size: \(ByteCountFormatter.string(fromByteCount: result.weightsSize, countStyle: .file)) ✅")233            case .failed(let message):234                throw CLIError.usage("quantize failed: \(message)")235            }236        }237    }238239    // MARK: - Download (drives HubService + DownloadManager)240241    private static func download(repo: String) async throws {242        print("searching hub for \(repo)…")243        let hits = try await HubService.shared.search(query: repo, author: nil, limit: 3)244        for hit in hits.prefix(3) {245            print("  \(hit.id)  [\(hit.modelType.displayName)]  \(hit.downloads.formatted()) downloads")246        }247        let events = await DownloadManager.shared.download(repo: repo)248        var lastPercent = -1249        for await event in events {250            switch event {251            case .progress(let progress):252                let percent = Int(progress.fraction * 100)253                if percent != lastPercent, percent % 10 == 0 {254                    print("\(percent)%  (\(progress.currentFile))")255                    lastPercent = percent256                }257            case .finished(let model):258                print("installed \(model.name) [\(model.type.displayName)] — \(ByteCountFormatter.string(fromByteCount: model.diskSize, countStyle: .file)) ✅")259            case .failed(let message):260                throw CLIError.usage("download failed: \(message)")261            }262        }263    }264265    // MARK: - Dataset validation266267    private static func validateDataset(path: String) async throws {268        let url = URL(fileURLWithPath: (path as NSString).expandingTildeInPath)269        let report = try await DatasetService.shared.importJSONL(270            from: url, name: url.deletingPathExtension().lastPathComponent + "-validated")271        let d = report.dataset272        print("format:      \(d.format.displayName)")273        print("valid rows:  \(d.sampleCount) (train \(d.trainCount) / valid \(d.validCount) / test \(d.testCount))")274        print("est. tokens: \(d.totalTokens ?? 0) total, longest sample ≈ \(d.maxSequenceTokens ?? 0)")275        if report.issues.isEmpty {276            print("issues:      none ✅")277        } else {278            print("issues:      \(report.issues.count) rows skipped")279            for issue in report.issues.prefix(10) {280                print("  line \(issue.line): \(issue.problem)\(issue.fix)")281            }282        }283        print("written to:  \(d.directory.path)")284    }285}286287enum CLIError: LocalizedError {288    case usage(String)289    var errorDescription: String? {290        if case .usage(let message) = self { return message }291        return nil292    }293}294