//
// 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
}
}