// // CLIFoundry.swift // Zyquo MLX // // Author: Simon-Pierre Boucher // Mail: contact@spboucher.ai // import Foundation /// CLI drivers for the Phase 3 foundry services (training, fusion, datasets). /// /// ZyquoMLX --train --model --dataset [--iters N] [--method lora|qlora|dora|full] /// ZyquoMLX --fuse --model --adapter --out [--dequantize] /// ZyquoMLX --validate-dataset extension CLI { static var shouldRunFoundry: Bool { let args = CommandLine.arguments return args.contains("--train") || args.contains("--fuse") || args.contains("--quantize") || args.contains("--validate-dataset") || args.contains("--download") || args.contains("--transcribe") } static func runFoundry() async -> Int32 { do { let args = CommandLine.arguments if args.contains("--train") { try await train(args: args) return 0 } if args.contains("--fuse") { try await fuse(args: args) return 0 } if args.contains("--quantize") { try await quantize(args: args) return 0 } if let path = value(after: "--validate-dataset", in: args) { try await validateDataset(path: path) return 0 } if let repo = value(after: "--download", in: args) { try await download(repo: repo) return 0 } if let modelPath = value(after: "--transcribe", in: args), let audioPath = value(after: "--audio", in: args) { let model = try await ModelStore.shared.describe( directory: URL(fileURLWithPath: (modelPath as NSString).expandingTildeInPath)) let result = try await SpeechService.shared.transcribe( model: model, audio: URL(fileURLWithPath: (audioPath as NSString).expandingTildeInPath)) print("language: \(result.language ?? "?") segments: \(result.segments) took \(String(format: "%.1f", result.duration))s") print(result.text) return 0 } return 2 } catch { FileHandle.standardError.write(Data("error: \(error.localizedDescription)\n".utf8)) return 1 } } // MARK: - Train private static func train(args: [String]) async throws { guard let modelPath = value(after: "--model", in: args), let dataPath = value(after: "--dataset", in: args) else { throw CLIError.usage("--train requires --model and --dataset ") } let modelURL = URL(fileURLWithPath: (modelPath as NSString).expandingTildeInPath) let model = try await ModelStore.shared.describe(directory: modelURL) // Accept a raw JSONL (imported on the fly) or an existing dataset dir. let dataURL = URL(fileURLWithPath: (dataPath as NSString).expandingTildeInPath) let dataset: Dataset if dataURL.pathExtension == "jsonl" { let report = try await DatasetService.shared.importJSONL( from: dataURL, name: dataURL.deletingPathExtension().lastPathComponent) dataset = report.dataset print("dataset: \(dataset.name) [\(dataset.format.displayName)] " + "train=\(dataset.trainCount) valid=\(dataset.validCount)" + (report.issues.isEmpty ? "" : " (\(report.issues.count) rows skipped)")) for issue in report.issues.prefix(5) { print(" line \(issue.line): \(issue.problem) → \(issue.fix)") } } else { dataset = try PersistenceService.loadJSON( Dataset.self, from: dataURL.appendingPathComponent("dataset.json")) } var hp = HyperParams() hp.iterations = value(after: "--iters", in: args).flatMap(Int.init) ?? 100 hp.batchSize = value(after: "--batch-size", in: args).flatMap(Int.init) ?? 2 hp.saveEvery = value(after: "--save-every", in: args).flatMap(Int.init) ?? 50 hp.stepsPerEval = value(after: "--steps-per-eval", in: args).flatMap(Int.init) ?? 50 hp.stepsPerReport = 5 let methodName = value(after: "--method", in: args) ?? (model.quantization != nil ? "qlora" : "lora") guard let method = FineTuneMethod(rawValue: methodName) else { throw CLIError.usage("unknown --method \(methodName)") } let service = TrainingService.shared let resume: Bool let run: TrainingRun if let resumeID = value(after: "--resume-run", in: args) { guard let id = UUID(uuidString: resumeID) else { throw CLIError.usage("bad run id") } run = try await RunStore.shared.load(id: id) resume = true } else { run = try await service.createRun( name: "cli-\(model.name)-\(method.rawValue)", baseModel: model, dataset: dataset, method: method, hyperParams: hp) resume = false } print("run: \(run.id.uuidString)\(resume ? " (warm resume)" : "")") print("method: \(run.method.displayName) on \(model.name)") print("verdict: \(MemoryAdvisor.trainingVerdict(for: model, method: run.method, params: run.hyperParams).displayName)\n") // Test hook mirroring the UI's Cancel button (drives the same // TrainingService.cancel() path). if let cancelAfter = value(after: "--cancel-after", in: args).flatMap(Double.init) { Task { try? await Task.sleep(for: .seconds(cancelAfter)) print("\n[--cancel-after \(cancelAfter)s] cancelling…") await service.cancel() } } let events = try await service.start(run: run, baseModel: model, dataset: dataset, resume: resume) for await event in events { switch event { case .started(let model, let iterations): print("training \(model) for \(iterations) iterations…") case .metric(let m): if let loss = m.trainLoss { let speed = m.tokensPerSecond.map { String(format: "%.0f tok/s", $0) } ?? "" let mem = m.peakMemoryGB.map { String(format: "peak %.1f GB", $0) } ?? "" print("iter \(m.iteration): train loss \(String(format: "%.3f", loss)) \(speed) \(mem)") } if let loss = m.valLoss { print("iter \(m.iteration): VAL loss \(String(format: "%.3f", loss))") } case .checkpointSaved(let name, _): print("checkpoint saved: \(name)") case .finished: print("\ntraining finished ✅") case .failed(let message): print("\ntraining FAILED: \(message)") } } let finished = try await RunStore.shared.load(id: run.id) print("state: \(finished.state.rawValue), adapters: \(await RunStore.shared.latestAdapter(for: finished)?.path ?? "none")") } // MARK: - Fuse private static func fuse(args: [String]) async throws { guard let modelPath = value(after: "--model", in: args), let adapterPath = value(after: "--adapter", in: args), let outName = value(after: "--out", in: args) else { throw CLIError.usage("--fuse requires --model , --adapter , --out ") } let model = try await ModelStore.shared.describe( directory: URL(fileURLWithPath: (modelPath as NSString).expandingTildeInPath)) let adapter = URL(fileURLWithPath: (adapterPath as NSString).expandingTildeInPath) // Fusing into a quantized base re-quantizes the merged weights, which // rounds away small LoRA deltas (verified empirically — the adapter's // behavior vanished). Default to de-quantizing for quantized bases; // --no-dequantize opts back into the lossy compact form. let dequantize = args.contains("--dequantize") || (model.quantization != nil && !args.contains("--no-dequantize")) if dequantize && model.quantization != nil { print("note: de-quantizing while fusing (quantized base) — re-quantize afterwards in Convert if needed") } let events = try await ConversionService.shared.fuse( baseModel: model, adapterDirectory: adapter, outputName: outName, dequantize: dequantize) for await event in events { switch event { case .stage(let stage, _): print("stage: \(stage)") case .finished(let output): print("fused model → \(output.path) ✅") case .failed(let message): throw CLIError.usage("fuse failed: \(message)") } } } // MARK: - Quantize (Swift-native, docs/MLX-RESEARCH.md §2) private static func quantize(args: [String]) async throws { guard let modelPath = value(after: "--quantize", in: args), let outName = value(after: "--out", in: args) else { throw CLIError.usage("--quantize --out [--bits N] [--group-size G]") } let model = try await ModelStore.shared.describe( directory: URL(fileURLWithPath: (modelPath as NSString).expandingTildeInPath)) var config = QuantConfig() config.bits = value(after: "--bits", in: args).flatMap(Int.init) ?? 4 config.groupSize = value(after: "--group-size", in: args).flatMap(Int.init) ?? 64 if let params = model.parameterCount { let predicted = config.predictedWeightBytes(parameterCount: params) print("quantizing \(model.name) → \(config.label)") print("size preview: \(ByteCountFormatter.string(fromByteCount: model.weightsSize, countStyle: .file)) → ~\(ByteCountFormatter.string(fromByteCount: predicted, countStyle: .file))") } let events = await ConversionService.shared.quantize( model: model, config: config, outputName: outName) for await event in events { switch event { case .stage(let stage, let fraction): let pct = fraction.map { String(format: " %.0f%%", $0 * 100) } ?? "" print("stage: \(stage)\(pct)") case .finished(let output): let result = try await ModelStore.shared.describe(directory: output) print("quantized model → \(output.path)") print("actual size: \(ByteCountFormatter.string(fromByteCount: result.weightsSize, countStyle: .file)) ✅") case .failed(let message): throw CLIError.usage("quantize failed: \(message)") } } } // MARK: - Download (drives HubService + DownloadManager) private static func download(repo: String) async throws { print("searching hub for \(repo)…") let hits = try await HubService.shared.search(query: repo, author: nil, limit: 3) for hit in hits.prefix(3) { print(" \(hit.id) [\(hit.modelType.displayName)] \(hit.downloads.formatted()) downloads") } let events = await DownloadManager.shared.download(repo: repo) var lastPercent = -1 for await event in events { switch event { case .progress(let progress): let percent = Int(progress.fraction * 100) if percent != lastPercent, percent % 10 == 0 { print("\(percent)% (\(progress.currentFile))") lastPercent = percent } case .finished(let model): print("installed \(model.name) [\(model.type.displayName)] — \(ByteCountFormatter.string(fromByteCount: model.diskSize, countStyle: .file)) ✅") case .failed(let message): throw CLIError.usage("download failed: \(message)") } } } // MARK: - Dataset validation private static func validateDataset(path: String) async throws { let url = URL(fileURLWithPath: (path as NSString).expandingTildeInPath) let report = try await DatasetService.shared.importJSONL( from: url, name: url.deletingPathExtension().lastPathComponent + "-validated") let d = report.dataset print("format: \(d.format.displayName)") print("valid rows: \(d.sampleCount) (train \(d.trainCount) / valid \(d.validCount) / test \(d.testCount))") print("est. tokens: \(d.totalTokens ?? 0) total, longest sample ≈ \(d.maxSequenceTokens ?? 0)") if report.issues.isEmpty { print("issues: none ✅") } else { print("issues: \(report.issues.count) rows skipped") for issue in report.issues.prefix(10) { print(" line \(issue.line): \(issue.problem) → \(issue.fix)") } } print("written to: \(d.directory.path)") } } enum CLIError: LocalizedError { case usage(String) var errorDescription: String? { if case .usage(let message) = self { return message } return nil } }