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%
5.8 KB · 150 lines swift
Raw Blame History
1//2//  CLI.swift3//  Zyquo MLX4//5//  Author: Simon-Pierre Boucher6//  Mail: contact@spboucher.ai7//89import Foundation10import MLXLMCommon1112/// Command-line proof-of-concept mode (Phase 2 gate): run inference on local13/// model directories without the UI.14///15///     ZyquoMLX --infer <model-dir> [--prompt "…"] [--max-tokens N] [--image <path>]16///     ZyquoMLX --embed <model-dir> --text "…" [--text "…"]…17enum CLI {1819    static var shouldRun: Bool {20        let args = CommandLine.arguments21        return args.contains("--infer") || args.contains("--embed")22    }2324    static func run() async -> Int32 {25        do {26            let args = CommandLine.arguments27            if let dir = value(after: "--infer", in: args) {28                try await infer(29                    directory: dir,30                    prompt: value(after: "--prompt", in: args)31                        ?? "Explain in one short sentence what MLX is.",32                    maxTokens: value(after: "--max-tokens", in: args).flatMap(Int.init) ?? 256,33                    imagePath: value(after: "--image", in: args)34                )35                return 036            }37            if let dir = value(after: "--embed", in: args) {38                let texts = values(after: "--text", in: args)39                try await embed(40                    directory: dir,41                    texts: texts.isEmpty42                        ? ["The quick brown fox", "A fast auburn fox", "Quarterly revenue grew 4%"]43                        : texts)44                return 045            }46            FileHandle.standardError.write(Data("usage: ZyquoMLX --infer <dir> | --embed <dir>\n".utf8))47            return 248        } catch {49            FileHandle.standardError.write(Data("error: \(error.localizedDescription)\n".utf8))50            return 151        }52    }5354    // MARK: - Subcommands5556    private static func infer(directory: String, prompt: String, maxTokens: Int, imagePath: String?) async throws {57        let url = URL(fileURLWithPath: (directory as NSString).expandingTildeInPath)58        let model = try await ModelStore.shared.describe(directory: url)59        print("model:   \(model.name) [\(model.type.displayName)\(model.quantization.map { ", \($0.label)" } ?? "")]")60        print("verdict: \(MemoryAdvisor.inferenceVerdict(for: model).displayName)")6162        let engine = InferenceEngine.shared63        let loadStart = Date()64        try await engine.load(model: model)65        print("loaded in \(String(format: "%.2f", Date().timeIntervalSince(loadStart)))s\n")6667        var message = Chat.Message.user(prompt)68        if let imagePath {69            let imageURL = URL(fileURLWithPath: (imagePath as NSString).expandingTildeInPath)70            message = Chat.Message.user(prompt, images: [.url(imageURL)])71        }7273        var params = GenerationParams()74        params.maxTokens = maxTokens7576        let stream = try await engine.generate(messages: [message], params: params)77        var stats: InferenceStats?78        for try await event in stream {79            switch event {80            case .chunk(let text):81                print(text, terminator: "")82                fflush(stdout)83            case .finished(let s):84                stats = s85            }86        }87        print("\n")88        if let stats {89            print("── stats ──────────────────────────────")90            print("prompt tokens:    \(stats.promptTokens)")91            print("generated tokens: \(stats.generatedTokens)")92            print(String(format: "ttft:             %.2fs", stats.ttft))93            print(String(format: "speed:            %.1f tok/s", stats.tokensPerSecond))94            print("stop reason:      \(stats.stopReason)")95        }96        let freed = try await engine.unload()97        print("unloaded (freed \(ByteCountFormatter.string(fromByteCount: freed, countStyle: .memory)))")98    }99100    private static func embed(directory: String, texts: [String]) async throws {101        let url = URL(fileURLWithPath: (directory as NSString).expandingTildeInPath)102        let model = try await ModelStore.shared.describe(directory: url)103        print("model: \(model.name) [\(model.type.displayName)]")104105        let engine = InferenceEngine.shared106        try await engine.load(model: model)107108        let start = Date()109        let vectors = try await engine.embed(texts: texts)110        let elapsed = Date().timeIntervalSince(start)111112        for (text, vector) in zip(texts, vectors) {113            let preview = vector.prefix(4).map { String(format: "%+.4f", $0) }.joined(separator: ", ")114            print("dim=\(vector.count)  [\(preview), …]  \"\(text)\"")115        }116        if vectors.count >= 2 {117            print("\n── cosine similarity ──────────────────")118            for i in 0..<vectors.count {119                for j in (i + 1)..<vectors.count {120                    let sim = InferenceEngine.cosineSimilarity(vectors[i], vectors[j])121                    print(String(format: "%.4f  \"%@\"\"%@\"", sim, texts[i], texts[j]))122                }123            }124        }125        print(String(format: "\nembedded %d texts in %.2fs", texts.count, elapsed))126        try await engine.unload()127    }128129    // MARK: - Arg parsing (shared with CLIFoundry)130131    static func value(after flag: String, in args: [String]) -> String? {132        guard let index = args.firstIndex(of: flag), index + 1 < args.count else { return nil }133        return args[index + 1]134    }135136    private static func values(after flag: String, in args: [String]) -> [String] {137        var out: [String] = []138        var i = 0139        while i < args.count {140            if args[i] == flag, i + 1 < args.count {141                out.append(args[i + 1])142                i += 2143            } else {144                i += 1145            }146        }147        return out148    }149}150