From 443220e7ea08fcb82f14663f0da42be3be8b3fcb Mon Sep 17 00:00:00 2001 From: David Koski Date: Mon, 20 Jul 2026 10:02:50 -0700 Subject: [PATCH 1/2] update for MaterializedArray --- Applications/LLMBasic/ChatModel.swift | 10 +-- .../LoRATrainingExample/ContentView.swift | 81 ++++++++--------- Tools/llm-tool/Chat.swift | 12 ++- Tools/llm-tool/LLMTool.swift | 20 +++-- Tools/llm-tool/LoraCommands.swift | 88 ++++++++----------- 5 files changed, 95 insertions(+), 116 deletions(-) diff --git a/Applications/LLMBasic/ChatModel.swift b/Applications/LLMBasic/ChatModel.swift index 280fbb09..89b34c1b 100644 --- a/Applications/LLMBasic/ChatModel.swift +++ b/Applications/LLMBasic/ChatModel.swift @@ -24,8 +24,8 @@ private let generateParameters = GenerateParameters(temperature: 0.5) enum State { case idle - case loading(Task) - case loaded(ModelContainer) + case loading(Task) + case loaded(ModelContext) } public var progress = 0.0 @@ -38,12 +38,12 @@ private let generateParameters = GenerateParameters(temperature: 0.5) private var state = State.idle - public func model() async throws -> ModelContainer { + public func model() async throws -> ModelContext { switch self.state { case .idle: let task = Task { // download and report progress - try await #huggingFaceLoadModelContainer( + try await #huggingFaceLoadModel( configuration: modelConfiguration ) { value in Task { @MainActor in @@ -79,7 +79,7 @@ private let generateParameters = GenerateParameters(temperature: 0.5) task != nil } - public init(model: ModelContainer) { + public init(model: ModelContext) { self.session = ChatSession( model, instructions: instructions, diff --git a/Applications/LoRATrainingExample/ContentView.swift b/Applications/LoRATrainingExample/ContentView.swift index 8c833257..d825ddf9 100644 --- a/Applications/LoRATrainingExample/ContentView.swift +++ b/Applications/LoRATrainingExample/ContentView.swift @@ -114,9 +114,9 @@ class LoRAEvaluator { case failed(String) } - enum ModelState: Sendable { + enum ModelState { case idle - case loaded(ModelContainer) + case loaded(TrainableModelContext) } var state = State.idle @@ -135,7 +135,7 @@ class LoRAEvaluator { private let evaluateShowEvery = 8 private let maxTokens = 200 - private func loadModel() async throws -> ModelContainer { + private func loadModel() async throws -> TrainableModelContext { switch self.model { case .idle: let name = modelConfiguration.name @@ -143,7 +143,7 @@ class LoRAEvaluator { progress = .init(title: "Loading \(name)", current: 0, limit: 1) } - let modelContainer = try await #huggingFaceLoadModelContainer( + let context = try await #huggingFaceLoadTrainabledModel( configuration: modelConfiguration ) { progress in @@ -153,8 +153,8 @@ class LoRAEvaluator { limit: 1.0) } } - self.model = .loaded(modelContainer) - return modelContainer + self.model = .loaded(context) + return context case .loaded(let modelContainer): return modelContainer @@ -185,15 +185,13 @@ class LoRAEvaluator { } // load the model - let modelContainer = try await loadModel() + let context = try await loadModel() // apply LoRA adapters and train - let _ = try await modelContainer.perform { context in - try LoRAContainer.from( - model: context.model, - configuration: LoRAConfiguration(numLayers: loraLayers) - ) - } + let _ = try LoRAContainer.from( + model: context.model, + configuration: LoRAConfiguration(numLayers: loraLayers) + ) let train = try loadLoRAData(name: "train") let valid = try loadLoRAData(name: "valid") @@ -202,29 +200,27 @@ class LoRAEvaluator { return } - try await modelContainer.perform { context in - let optimizer = Adam(learningRate: learningRate) - try LoRATrain.train( - model: context.model, train: train, validate: valid, optimizer: optimizer, - tokenizer: context.tokenizer, - parameters: parameters - ) { progress in - Task { @MainActor in - switch progress { - case .train(let i, _, _, _): - self.progress = .init( - title: "Train", current: Double(i), limit: Double(parameters.iterations) - ) - case .validation: - output += "\n" - default: - break - } - output += progress.description + "\n" + let optimizer = Adam(learningRate: learningRate) + try LoRATrain.train( + model: context.model, train: train, validate: valid, optimizer: optimizer, + tokenizer: context.tokenizer, + parameters: parameters + ) { progress in + Task { @MainActor in + switch progress { + case .train(let i, _, _, _): + self.progress = .init( + title: "Train", current: Double(i), limit: Double(parameters.iterations) + ) + case .validation: + output += "\n" + default: + break } - - return .more + output += progress.description + "\n" } + + return .more } // done training, test @@ -234,11 +230,9 @@ class LoRAEvaluator { return } - let loss = await modelContainer.perform { context in - LoRATrain.evaluate( - model: context.model, dataset: test, - tokenizer: context.tokenizer, batchSize: 1, batchCount: 0) - } + let loss = LoRATrain.evaluate( + model: context.model, dataset: test, + tokenizer: context.tokenizer, batchSize: 1, batchCount: 0) self.progress = nil self.output += "\n" @@ -262,15 +256,16 @@ class LoRAEvaluator { MLXRandom.seed(UInt64(Date.timeIntervalSinceReferenceDate * 1000)) - let modelContainer = try await loadModel() + let context = try await loadModel() // evaluate - let input = try await modelContainer.processor.prepare(input: .init(prompt: prompt)) + let input = try await context.processor.prepare(input: .init(prompt: prompt)) + let evaluationContext = ModelContext(context) var count = 0 var output = "" - for try await item in try await modelContainer.generate( - input: input, parameters: generateParameters + for try await item in try generate( + input: input, parameters: generateParameters, context: evaluationContext ) { switch item { case .chunk(let string): diff --git a/Tools/llm-tool/Chat.swift b/Tools/llm-tool/Chat.swift index a33d58d2..ddeea04f 100644 --- a/Tools/llm-tool/Chat.swift +++ b/Tools/llm-tool/Chat.swift @@ -22,7 +22,7 @@ struct ChatCommand: AsyncParsableCommand { let defaultModel = MLXLLM.LLMRegistry.mistral7B4bit // Load the model - let modelContainer = try await memory.start { [args] in + var context = try await memory.start { [args] in do { return try await args.load( defaultModel: defaultModel.name, modelFactory: VLMModelFactory.shared) @@ -33,16 +33,14 @@ struct ChatCommand: AsyncParsableCommand { } // update the context/configuration with any command line parameters - await modelContainer.update { [generate] context in - generate.prepare(&context) - } + generate.prepare(&context) - try await chat(modelContainer: modelContainer) + try await chat(context: context) } - func chat(modelContainer: ModelContainer) async throws { + func chat(context: ModelContext) async throws { let session = ChatSession( - modelContainer, + context, instructions: generate.system, generateParameters: generate.generateParameters, processing: media.processing diff --git a/Tools/llm-tool/LLMTool.swift b/Tools/llm-tool/LLMTool.swift index a424a58c..e4c2361e 100644 --- a/Tools/llm-tool/LLMTool.swift +++ b/Tools/llm-tool/LLMTool.swift @@ -43,7 +43,12 @@ struct ModelArguments: ParsableArguments, Sendable { } @Sendable - func load(defaultModel: String, modelFactory: any ModelFactory) async throws -> ModelContainer { + func load(defaultModel: String, modelFactory: any ModelFactory) async throws -> ModelContext { + ModelContext(try await loadTrainable(defaultModel: defaultModel, modelFactory: modelFactory)) + } + + @Sendable + func loadTrainable(defaultModel: String, modelFactory: any ModelFactory) async throws -> TrainableModelContext { let modelConfiguration: ModelConfiguration let modelName = self.model ?? defaultModel @@ -58,11 +63,12 @@ struct ModelArguments: ParsableArguments, Sendable { modelConfiguration = modelFactory.configuration(id: modelName) } - return try await modelFactory.loadContainer( + return try await modelFactory.loadTrainable( from: self.downloader, using: #huggingFaceTokenizerLoader(), configuration: modelConfiguration) } + } struct PromptArguments: ParsableArguments, Sendable { @@ -334,17 +340,15 @@ struct EvaluateCommand: AsyncParsableCommand { } // Load the model - let modelContainer = try await memory.start { [args] in + var context = try await memory.start { [args] in try await args.load(defaultModel: defaultModel.name, modelFactory: modelFactory) } // update the context/configuration with any command line parameters - await modelContainer.update { [generate] context in - generate.prepare(&context) - } + generate.prepare(&context) // Get the resolved configuration (this has the default prompt) - let modelConfiguration = await modelContainer.configuration + let modelConfiguration = context.configuration let prompt = (try? self.prompt.resolvePrompt(configuration: modelConfiguration)) @@ -355,7 +359,7 @@ struct EvaluateCommand: AsyncParsableCommand { } let session = ChatSession( - modelContainer, + context, instructions: generate.system, generateParameters: generate.generateParameters, processing: media.processing, diff --git a/Tools/llm-tool/LoraCommands.swift b/Tools/llm-tool/LoraCommands.swift index c66b1e52..5689f12d 100644 --- a/Tools/llm-tool/LoraCommands.swift +++ b/Tools/llm-tool/LoraCommands.swift @@ -42,8 +42,8 @@ struct LoRAModelArguments: ParsableArguments, Sendable { func load( defaultModel: String = defaultModel, modelFactory: any ModelFactory = LLMModelFactory.shared - ) async throws -> (ModelContainer, ModelAdapter) { - let modelContainer = try await args.load( + ) async throws -> (TrainableModelContext, ModelAdapter) { + let modelContext = try await args.loadTrainable( defaultModel: defaultModel, modelFactory: modelFactory) // Load LoRA adapter from directory or create a new one @@ -51,13 +51,11 @@ struct LoRAModelArguments: ParsableArguments, Sendable { do { modelAdapter = try LoRAContainer.from(directory: adapter) } catch { - modelAdapter = try await modelContainer.perform { context in - return try LoRAContainer.from( - model: context.model, configuration: LoRAConfiguration(numLayers: loraLayers)) - } + modelAdapter = try! LoRAContainer.from( + model: modelContext.model, configuration: LoRAConfiguration(numLayers: loraLayers)) } - return (modelContainer, modelAdapter) + return (modelContext, modelAdapter) } func describe(model: Module) { @@ -125,18 +123,14 @@ struct LoRATrainCommand: AsyncParsableCommand { @MainActor mutating func run() async throws { - let (modelContainer, modelAdapter) = try await args.load() - await modelContainer.perform { [args] context in - args.describe(model: context.model) - } + let (modelContext, modelAdapter) = try await args.load() + args.describe(model: modelContext.model) memory.start() if resume { print("Loading pretrained adapters from \(args.adapter.path())") - try await modelContainer.perform { context in - try context.model.load(adapter: modelAdapter) - } + try modelContext.model.load(adapter: modelAdapter) } // load the train/validation data @@ -151,18 +145,16 @@ struct LoRATrainCommand: AsyncParsableCommand { } // train - try await modelContainer.perform { [args, parameters, learningRate] context in - let optimizer = Adam(learningRate: learningRate) - try LoRATrain.train( - model: context.model, train: train, validate: valid, optimizer: optimizer, - tokenizer: context.tokenizer, - parameters: parameters - ) { progress in - print(progress) - return .more - } - try LoRATrain.saveLoRAWeights(model: context.model, url: args.adapter) + let optimizer = Adam(learningRate: learningRate) + try LoRATrain.train( + model: modelContext.model, train: train, validate: valid, optimizer: optimizer, + tokenizer: modelContext.tokenizer, + parameters: parameters + ) { progress in + print(progress) + return .more } + try LoRATrain.saveLoRAWeights(model: modelContext.model, url: args.adapter) } } @@ -202,15 +194,13 @@ struct LoRAFuseCommand: AsyncParsableCommand { outputURL = cache.repoDirectory(repo: repo, kind: .model) } - let (modelContainer, modelAdapter) = try await args.load() + let (modelContext, modelAdapter) = try await args.load() // fuse LoRA layers back into Linear/QuantizedLinear - try await modelContainer.perform { context in - try context.model.fuse(with: modelAdapter) - } + try modelContext.model.fuse(with: modelAdapter) let resolved = try await resolve( - configuration: modelContainer.configuration, + configuration: modelContext.configuration, from: args.args.downloader, useLatest: false, progressHandler: { _ in }) @@ -230,10 +220,8 @@ struct LoRAFuseCommand: AsyncParsableCommand { } // write them back out - try await modelContainer.perform { context in - let weights = Dictionary(uniqueKeysWithValues: context.model.parameters().flattened()) - try save(arrays: weights, url: outputURL.appending(component: "weights.safetensors")) - } + let weights = Dictionary(uniqueKeysWithValues: modelContext.model.parameters().flattened()) + try save(arrays: weights, url: outputURL.appending(component: "weights.safetensors")) print("Fused weights written to \(outputURL.path())") print("Use with:\n\tllm-tool eval --model \(output)") @@ -259,20 +247,16 @@ struct LoRATestCommand: AsyncParsableCommand { @MainActor mutating func run() async throws { - let (modelContainer, _) = try await args.load() - await modelContainer.perform { [args] context in - args.describe(model: context.model) - } + let (modelContext, _) = try await args.load() + args.describe(model: modelContext.model) memory.start() let test = try loadLoRAData(directory: data, name: "test") - let loss = await modelContainer.perform { [batchSize] context in - LoRATrain.evaluate( - model: context.model, dataset: test, - tokenizer: context.tokenizer, batchSize: batchSize, - batchCount: 0) - } + let loss = LoRATrain.evaluate( + model: modelContext.model, dataset: test, + tokenizer: modelContext.tokenizer, batchSize: batchSize, + batchCount: 0) print("Test loss \(loss.formatted()), ppl \(exp(loss).formatted())") } @@ -293,14 +277,12 @@ struct LoRAEvalCommand: AsyncParsableCommand { @MainActor mutating func run() async throws { - let (modelContainer, _) = try await args.load() - await modelContainer.perform { [args] context in - args.describe(model: context.model) - } + let (modelContext, _) = try await args.load() + args.describe(model: modelContext.model) memory.start() - let defaultPrompt = await modelContainer.configuration.defaultPrompt + let defaultPrompt = modelContext.configuration.defaultPrompt let prompt = prompt.prompt ?? defaultPrompt if !generate.quiet { @@ -309,10 +291,10 @@ struct LoRAEvalCommand: AsyncParsableCommand { } // generate and print the result - let (result, _) = try await modelContainer.perform { [generate] context in - let input = try await context.processor.prepare(input: .init(prompt: prompt)) - return try await generate.generate(input: input, context: context) - } + let input = try await modelContext.processor.prepare(input: .init(prompt: prompt)) + + let evaluationContext = ModelContext(modelContext) + let (result, _) = try await generate.generate(input: input, context: evaluationContext) if !generate.quiet { print("------") From 46f53c389756a1e6775531206b31c4e42bb2e66a Mon Sep 17 00:00:00 2001 From: David Koski Date: Mon, 20 Jul 2026 15:06:52 -0700 Subject: [PATCH 2/2] fix more warnings --- .../LLMEval/ViewModels/LLMEvaluator.swift | 25 ++--- .../LoRATrainingExample/ContentView.swift | 2 +- .../MLXChatExample/Services/MLXService.swift | 38 ++++--- .../EmbedderRuntime+Embedding.swift | 98 +++++++++---------- Tools/embedder-tool/EmbedderTool.swift | 6 +- Tools/embedder-tool/ModelArguments.swift | 6 +- Tools/llm-tool/LLMTool.swift | 17 +++- Tools/llm-tool/LoraCommands.swift | 4 +- 8 files changed, 106 insertions(+), 90 deletions(-) diff --git a/Applications/LLMEval/ViewModels/LLMEvaluator.swift b/Applications/LLMEval/ViewModels/LLMEvaluator.swift index 2aacbd28..f3eade31 100644 --- a/Applications/LLMEval/ViewModels/LLMEvaluator.swift +++ b/Applications/LLMEval/ViewModels/LLMEvaluator.swift @@ -63,7 +63,7 @@ class LLMEvaluator { enum LoadState { case idle case loading - case loaded(ModelContainer) + case loaded(ModelContext) } var loadState = LoadState.idle @@ -81,7 +81,7 @@ class LLMEvaluator { } /// Load and return the model. Can be called multiple times; subsequent calls return the cached model. - func load() async throws -> ModelContainer { + func load() async throws -> ModelContext { while true { switch loadState { case .idle: @@ -91,13 +91,13 @@ class LLMEvaluator { // Already loading, wait and retry try await Task.sleep(for: .milliseconds(100)) - case .loaded(let modelContainer): - return modelContainer + case .loaded(let model): + return model } } } - private func performLoad() async throws -> ModelContainer { + private func performLoad() async throws -> ModelContext { loadState = .loading modelInfo = "Downloading \(modelName)..." downloadProgress = 0.0 @@ -137,16 +137,16 @@ class LLMEvaluator { downloadProgress = nil totalSize = nil - let modelContainer = try await LLMModelFactory.shared.loadContainer( + let context = try await LLMModelFactory.shared.load( from: resolved.modelDirectory, using: #huggingFaceTokenizerLoader()) - let numParams = await modelContainer.perform { $0.model.numParameters() } + let numParams = context.model.parameterCount self.prompt = PresetPrompts.all[0].prompt self.modelInfo = formatModelInfo(name: modelConfiguration.name, parameters: numParams) - loadState = .loaded(modelContainer) - return modelContainer + loadState = .loaded(context) + return context } catch { resetLoadingState() @@ -240,7 +240,7 @@ class LLMEvaluator { ) do { - let modelContainer = try await load() + let context = try await load() // Capture parameters on MainActor before entering perform block let parameters = generateParameters @@ -248,10 +248,11 @@ class LLMEvaluator { // Seed random generator to ensure varied output each generation MLXRandom.seed(UInt64(Date.timeIntervalSinceReferenceDate * 1000)) - let lmInput = try await modelContainer.prepare(input: userInput) + let lmInput = try await context.processor.prepare(input: userInput) let promptTokenCount = lmInput.text.tokens.size let start = Date.timeIntervalSinceReferenceDate - let stream = try await modelContainer.generate(input: lmInput, parameters: parameters) + let stream = try MLXLMCommon.generate( + input: lmInput, parameters: parameters, context: context) var iterator = stream.makeAsyncIterator() if let first = await iterator.next() { diff --git a/Applications/LoRATrainingExample/ContentView.swift b/Applications/LoRATrainingExample/ContentView.swift index d825ddf9..977aa61c 100644 --- a/Applications/LoRATrainingExample/ContentView.swift +++ b/Applications/LoRATrainingExample/ContentView.swift @@ -143,7 +143,7 @@ class LoRAEvaluator { progress = .init(title: "Loading \(name)", current: 0, limit: 1) } - let context = try await #huggingFaceLoadTrainabledModel( + let context = try await #huggingFaceLoadTrainableModel( configuration: modelConfiguration ) { progress in diff --git a/Applications/MLXChatExample/Services/MLXService.swift b/Applications/MLXChatExample/Services/MLXService.swift index b9049970..c665e1fb 100644 --- a/Applications/MLXChatExample/Services/MLXService.swift +++ b/Applications/MLXChatExample/Services/MLXService.swift @@ -39,8 +39,16 @@ class MLXService { LMModel(name: "gemma3n:E4B", configuration: LLMRegistry.gemma3n_E4B_it_lm_4bit, type: .llm), ] + fileprivate final class ContextBox: Sendable { + let context: ModelContext + + init(_ context: ModelContext) { + self.context = context + } + } + /// Cache to store loaded model containers to avoid reloading. - private let modelCache = NSCache() + private let modelCache = NSCache() /// Tracks the current model download progress. /// Access this property to monitor model download status. @@ -51,16 +59,16 @@ class MLXService { /// - Parameter model: The model configuration to load /// - Returns: A ModelContainer instance containing the loaded model /// - Throws: Errors that might occur during model loading - private func load(model: LMModel) async throws -> ModelContainer { + private func load(model: LMModel) async throws -> ModelContext { // Set GPU memory limit to prevent out of memory issues Memory.cacheLimit = 20 * 1024 * 1024 // Return cached model if available to avoid reloading - if let container = modelCache.object(forKey: model.name as NSString) { - return container + if let box = modelCache.object(forKey: model.name as NSString) { + return box.context } else { // Select appropriate factory based on model type - let factory: ModelFactory = + let factory: any ModelFactory = switch model.type { case .llm: LLMModelFactory.shared @@ -72,7 +80,7 @@ class MLXService { let loader = #huggingFaceTokenizerLoader() // Load model and track download progress - let container = try await factory.loadContainer( + let context = try await factory.load( from: downloader, using: loader, configuration: model.configuration @@ -83,9 +91,9 @@ class MLXService { } // Cache the loaded model for future use - modelCache.setObject(container, forKey: model.name as NSString) + modelCache.setObject(.init(context), forKey: model.name as NSString) - return container + return context } } @@ -97,7 +105,7 @@ class MLXService { /// - Throws: Errors that might occur during generation func generate(messages: [Message], model: LMModel) async throws -> AsyncStream { // Load or retrieve model from cache - let modelContainer = try await load(model: model) + let context = try await load(model: model) // Exclude trailing empty assistant message so the chat template // leaves the assistant turn open for generation (matching ChatSession behavior) @@ -131,13 +139,11 @@ class MLXService { chat: chat, processing: .init(resize: .init(width: 1024, height: 1024))) // Generate response using the model - return try await modelContainer.perform { context in - let lmInput = try await context.processor.prepare(input: userInput) - // Set temperature for response randomness (0.7 provides good balance) - let parameters = GenerateParameters(temperature: 0.7) + let lmInput = try await context.processor.prepare(input: userInput) + // Set temperature for response randomness (0.7 provides good balance) + let parameters = GenerateParameters(temperature: 0.7) - return try MLXLMCommon.generate( - input: lmInput, parameters: parameters, context: context) - } + return try MLXLMCommon.generate( + input: lmInput, parameters: parameters, context: context) } } diff --git a/Tools/embedder-tool/EmbedderRuntime+Embedding.swift b/Tools/embedder-tool/EmbedderRuntime+Embedding.swift index 4c98fcaa..3f6d2e15 100644 --- a/Tools/embedder-tool/EmbedderRuntime+Embedding.swift +++ b/Tools/embedder-tool/EmbedderRuntime+Embedding.swift @@ -26,66 +26,64 @@ extension EmbedderRuntime { embeddings: [], skippedIndices: [], fallbackDescription: nil) } - return try await container.perform { context in - var skippedIndices: [Int] = [] + var skippedIndices: [Int] = [] - let tokenizer = context.tokenizer - let encoded = texts.enumerated().compactMap { index, text -> (Int, [Int])? in - let tokens = tokenizer.encode(text: text, addSpecialTokens: true) - guard !tokens.isEmpty else { - skippedIndices.append(index) - return nil - } - return (index, tokens) - } - - guard !encoded.isEmpty else { - return RuntimeEmbeddingResult( - embeddings: [], - skippedIndices: skippedIndices, - fallbackDescription: nil - ) + let tokenizer = context.tokenizer + let encoded = texts.enumerated().compactMap { index, text -> (Int, [Int])? in + let tokens = tokenizer.encode(text: text, addSpecialTokens: true) + guard !tokens.isEmpty else { + skippedIndices.append(index) + return nil } + return (index, tokens) + } - // [PAD] (BERT standard), EOS (autoregressive like Qwen) - let padToken = tokenizer.convertTokenToId("[PAD]") ?? tokenizer.eosTokenId ?? 0 + guard !encoded.isEmpty else { + return RuntimeEmbeddingResult( + embeddings: [], + skippedIndices: skippedIndices, + fallbackDescription: nil + ) + } - let maxLength = encoded.map { $0.1.count }.max() ?? 0 + // [PAD] (BERT standard), EOS (autoregressive like Qwen) + let padToken = tokenizer.convertTokenToId("[PAD]") ?? tokenizer.eosTokenId ?? 0 - let padded = stacked( - encoded.map { _, tokens in - MLXArray(tokens + Array(repeating: padToken, count: maxLength - tokens.count)) - }) - let mask = (padded .!= padToken) - let tokenTypes = MLXArray.zeros(like: padded) + let maxLength = encoded.map { $0.1.count }.max() ?? 0 - let outputs = context.model( - padded, - positionIds: nil, - tokenTypeIds: tokenTypes, - attentionMask: mask - ) + let padded = stacked( + encoded.map { _, tokens in + MLXArray(tokens + Array(repeating: padToken, count: maxLength - tokens.count)) + }) + let mask = (padded .!= padToken) + let tokenTypes = MLXArray.zeros(like: padded) - let poolingModule = resolvedPooler(for: context.pooling) - let pooled = poolingModule( - outputs, - mask: mask, - normalize: self.normalize, - applyLayerNorm: self.applyLayerNorm - ) - pooled.eval() + let outputs = context.model( + padded, + positionIds: nil, + tokenTypeIds: tokenTypes, + attentionMask: mask + ) - let extraction = try extractVectors(from: pooled, expectedCount: encoded.count) + let poolingModule = resolvedPooler(for: context.pooling) + let pooled = poolingModule( + outputs, + mask: mask, + normalize: self.normalize, + applyLayerNorm: self.applyLayerNorm + ) + pooled.eval() - let embeddings = zip(encoded.map { $0.0 }, extraction.vectors).map { index, vector in - (index: index, vector: vector) - } + let extraction = try extractVectors(from: pooled, expectedCount: encoded.count) - return RuntimeEmbeddingResult( - embeddings: embeddings, - skippedIndices: skippedIndices, - fallbackDescription: extraction.fallbackDescription - ) + let embeddings = zip(encoded.map { $0.0 }, extraction.vectors).map { index, vector in + (index: index, vector: vector) } + + return RuntimeEmbeddingResult( + embeddings: embeddings, + skippedIndices: skippedIndices, + fallbackDescription: extraction.fallbackDescription + ) } } diff --git a/Tools/embedder-tool/EmbedderTool.swift b/Tools/embedder-tool/EmbedderTool.swift index d6c85d64..e8e68611 100644 --- a/Tools/embedder-tool/EmbedderTool.swift +++ b/Tools/embedder-tool/EmbedderTool.swift @@ -43,11 +43,11 @@ struct EmbedderTool: AsyncParsableCommand { -> EmbedderRuntime { let loadedModel = try await model.load(default: defaultModelConfiguration) - let baseStrategy = await loadedModel.container.poolingStrategy + let baseStrategy = loadedModel.context.pooling.strategy return EmbedderRuntime( configuration: loadedModel.configuration, - container: loadedModel.container, + context: loadedModel.context, baseStrategy: baseStrategy, strategyOverride: pooling.strategyOverride, normalize: pooling.normalize, @@ -58,7 +58,7 @@ struct EmbedderTool: AsyncParsableCommand { struct EmbedderRuntime { let configuration: ModelConfiguration - let container: EmbedderModelContainer + let context: EmbedderModelContext let baseStrategy: Pooling.Strategy let strategyOverride: Pooling.Strategy? let normalize: Bool diff --git a/Tools/embedder-tool/ModelArguments.swift b/Tools/embedder-tool/ModelArguments.swift index a25fa6d6..681b389e 100644 --- a/Tools/embedder-tool/ModelArguments.swift +++ b/Tools/embedder-tool/ModelArguments.swift @@ -38,7 +38,7 @@ struct ModelArguments: ParsableArguments { struct LoadedEmbedderModel { let configuration: ModelConfiguration - let container: EmbedderModelContainer + let context: EmbedderModelContext } extension ModelArguments { @@ -51,7 +51,7 @@ extension ModelArguments { print("Loading model \(configuration.name)...") - let container = try await EmbedderModelFactory.shared.loadContainer( + let context = try await EmbedderModelFactory.shared.load( from: hub, using: loader, configuration: configuration, @@ -65,7 +65,7 @@ extension ModelArguments { } ) - return LoadedEmbedderModel(configuration: configuration, container: container) + return LoadedEmbedderModel(configuration: configuration, context: context) } var downloader: any Downloader { diff --git a/Tools/llm-tool/LLMTool.swift b/Tools/llm-tool/LLMTool.swift index e4c2361e..593cf14b 100644 --- a/Tools/llm-tool/LLMTool.swift +++ b/Tools/llm-tool/LLMTool.swift @@ -44,11 +44,14 @@ struct ModelArguments: ParsableArguments, Sendable { @Sendable func load(defaultModel: String, modelFactory: any ModelFactory) async throws -> ModelContext { - ModelContext(try await loadTrainable(defaultModel: defaultModel, modelFactory: modelFactory)) + ModelContext( + try await loadTrainable(defaultModel: defaultModel, modelFactory: modelFactory)) } - + @Sendable - func loadTrainable(defaultModel: String, modelFactory: any ModelFactory) async throws -> TrainableModelContext { + func loadTrainable(defaultModel: String, modelFactory: any ModelFactory) async throws + -> TrainableModelContext + { let modelConfiguration: ModelConfiguration let modelName = self.model ?? defaultModel @@ -209,6 +212,14 @@ struct GenerateArguments: ParsableArguments, Sendable { } } + func prepare( + _ context: inout ModelContext + ) { + if let extraEosToken { + context.configuration.extraEOSTokens.insert(extraEosToken) + } + } + func generate( input: LMInput, context: ModelContextProviding ) async throws -> (GenerateCompletionInfo, String) { diff --git a/Tools/llm-tool/LoraCommands.swift b/Tools/llm-tool/LoraCommands.swift index 5689f12d..ea93a8c8 100644 --- a/Tools/llm-tool/LoraCommands.swift +++ b/Tools/llm-tool/LoraCommands.swift @@ -59,7 +59,7 @@ struct LoRAModelArguments: ParsableArguments, Sendable { } func describe(model: Module) { - let totalParameterCount = model.numParameters() + let totalParameterCount = model.parameterCount let trainableParameterCount = model.trainableParameters() .flattenedValues().map { $0.size }.reduce(0, +) @@ -292,7 +292,7 @@ struct LoRAEvalCommand: AsyncParsableCommand { // generate and print the result let input = try await modelContext.processor.prepare(input: .init(prompt: prompt)) - + let evaluationContext = ModelContext(modelContext) let (result, _) = try await generate.generate(input: input, context: evaluationContext)