From e9e88ef6f5bbcae7ca74bb8564068bd4282fd6d8 Mon Sep 17 00:00:00 2001 From: Taylor Lineman Date: Tue, 25 Aug 2026 23:50:27 -0400 Subject: [PATCH 1/2] Added async throws prewarm Signed-off-by: Taylor Lineman --- Package.resolved | 2 +- Package.swift | 16 +- Sources/AnyLanguageModel/LanguageModel.swift | 13 ++ .../LanguageModelSession.swift | 5 + .../Models/MLXLanguageModel.swift | 150 +++++++++++------- 5 files changed, 120 insertions(+), 66 deletions(-) diff --git a/Package.resolved b/Package.resolved index c21ad81b..2ba11832 100644 --- a/Package.resolved +++ b/Package.resolved @@ -1,5 +1,5 @@ { - "originHash" : "4d76b6542f6103b20f1463e64cc20cc693e2deed4721637f851e01ea691c9bc2", + "originHash" : "d3648f331b52e336310a641dfb95851cb857fa1a6cd5f54c184e0506743506e6", "pins" : [ { "identity" : "eventsource", diff --git a/Package.swift b/Package.swift index 23460de8..d1ff645d 100644 --- a/Package.swift +++ b/Package.swift @@ -12,7 +12,7 @@ let package = Package( .iOS(.v17), .tvOS(.v17), .watchOS(.v10), - .visionOS(.v1) + .visionOS(.v1), ], products: [ @@ -26,7 +26,7 @@ let package = Package( .trait(name: "MLX"), .trait(name: "Llama"), .trait(name: "AsyncHTTPClient"), - .default(enabledTraits: []) + .default(enabledTraits: ["MLX"]), ], dependencies: [ .package(url: "https://github.com/huggingface/swift-transformers", from: "1.0.0"), @@ -36,7 +36,7 @@ let package = Package( from: "1.3.0", traits: [ .defaults, - .trait(name: "AsyncHTTPClient", condition: .when(traits: ["AsyncHTTPClient"])) + .trait(name: "AsyncHTTPClient", condition: .when(traits: ["AsyncHTTPClient"])), ] ), .package(url: "https://github.com/mattt/JSONSchema", from: "1.3.0"), @@ -46,7 +46,7 @@ let package = Package( // .package(url: "https://github.com/ml-explore/mlx-swift-lm", from: "3.0.0"), .package(url: "https://github.com/impel-intelligence/mlx-swift-lm", from: "1.0.1"), .package(url: "https://github.com/swiftlang/swift-syntax", from: "602.0.0"), - .package(url: "https://github.com/swift-server/async-http-client.git", from: "1.24.0") + .package(url: "https://github.com/swift-server/async-http-client.git", from: "1.24.0"), ], targets: [ .target( @@ -100,7 +100,7 @@ let package = Package( name: "AsyncHTTPClient", package: "async-http-client", condition: .when(traits: ["AsyncHTTPClient"]) - ) + ), ] ), .macro( @@ -109,7 +109,7 @@ let package = Package( .product(name: "SwiftCompilerPlugin", package: "swift-syntax"), .product(name: "SwiftSyntax", package: "swift-syntax"), .product(name: "SwiftSyntaxBuilder", package: "swift-syntax"), - .product(name: "SwiftSyntaxMacros", package: "swift-syntax") + .product(name: "SwiftSyntaxMacros", package: "swift-syntax"), ] ), .testTarget( @@ -120,8 +120,8 @@ let package = Package( name: "AsyncHTTPClient", package: "async-http-client", condition: .when(traits: ["AsyncHTTPClient"]) - ) + ), ], - ) + ), ] ) diff --git a/Sources/AnyLanguageModel/LanguageModel.swift b/Sources/AnyLanguageModel/LanguageModel.swift index 635de68c..34584347 100644 --- a/Sources/AnyLanguageModel/LanguageModel.swift +++ b/Sources/AnyLanguageModel/LanguageModel.swift @@ -16,6 +16,11 @@ public protocol LanguageModel: Sendable { for session: LanguageModelSession, promptPrefix: Prompt? ) + + func prewarm( + for session: LanguageModelSession, + promptPrefix: Prompt? + ) async throws func respond( within session: LanguageModelSession, @@ -58,6 +63,14 @@ extension LanguageModel { ) { return } + + public func prewarm( + for session: LanguageModelSession, + promptPrefix: Prompt? = nil + ) async throws { + return + } + public func logFeedbackAttachment( within session: LanguageModelSession, diff --git a/Sources/AnyLanguageModel/LanguageModelSession.swift b/Sources/AnyLanguageModel/LanguageModelSession.swift index a2e94c62..20fa4547 100644 --- a/Sources/AnyLanguageModel/LanguageModelSession.swift +++ b/Sources/AnyLanguageModel/LanguageModelSession.swift @@ -111,6 +111,11 @@ public final class LanguageModelSession: @unchecked Sendable { public func prewarm(promptPrefix: Prompt? = nil) { model.prewarm(for: self, promptPrefix: promptPrefix) } + + public func prewarm(promptPrefix: Prompt? = nil) async throws { + try await model.prewarm(for: self, promptPrefix: promptPrefix) + } + nonisolated private func beginResponding() { withMutation(keyPath: \.isResponding) { diff --git a/Sources/AnyLanguageModel/Models/MLXLanguageModel.swift b/Sources/AnyLanguageModel/Models/MLXLanguageModel.swift index 27b55107..091637a3 100644 --- a/Sources/AnyLanguageModel/Models/MLXLanguageModel.swift +++ b/Sources/AnyLanguageModel/Models/MLXLanguageModel.swift @@ -53,7 +53,8 @@ import Foundation ) async throws -> ModelContext { let cacheKey = key as NSString if let cached = cache.object(forKey: cacheKey), - case .loaded(let context) = cached.value { + case .loaded(let context) = cached.value + { return context } @@ -287,7 +288,7 @@ import Foundation var additionalContextForUserInput: [String: any Sendable]? { additionalContext?.mapValues { $0.toSendable() } } - + /// The nucleus-sampling probability threshold. Only tokens whose cumulative /// probability mass falls within the top `topP` fraction are considered during /// sampling. `nil` lets MLX use its model default. @@ -311,7 +312,7 @@ import Foundation /// appeared in the output. Values greater than `1.0` reduce repetition; values less /// than `1.0` encourage it. `nil` lets MLX use its model default. public var repetitionPenalty: Float? - + /// Creates MLX-specific generation options. /// /// - Parameters: @@ -853,7 +854,8 @@ import Foundation let existingEntry = getSessionCache(for: session) if let existingEntry, - isCacheHit(entry: existingEntry, currentTokens: fullTokens, signature: signature, lmInput: lmInput) { + isCacheHit(entry: existingEntry, currentTokens: fullTokens, signature: signature, lmInput: lmInput) + { let cachedCount = existingEntry.prefillTokenCount let newTokens = lmInput.text.tokens[cachedCount...] let newMask = lmInput.text.mask?[cachedCount...] @@ -1298,6 +1300,70 @@ import Foundation return LanguageModelSession.ResponseStream(stream: stream) } + public func _prewarm( + for session: LanguageModelSession, + promptPrefix: Prompt?, + modelID: String, + hub: HubClient?, + directory: URL? + ) async throws { + guard Self.acquireGenerationSlot(for: session) else { + return + } + defer { Self.releaseGenerationSlot(for: session) } + + let generationScope = beginGenerationScope() + defer { endGenerationScope(generationScope) } + + let context = try await loadContext(modelId: modelId, hub: hub, directory: directory) + guard let instructions = session.instructions?.description, !instructions.isEmpty else { + return + } + + let toolSpecs = mlxToolSpecs(for: session) + + let params = toGenerateParameters(.init()) + let newCache = context.model.newCache(parameters: params) + let userInput = MLXLMCommon.UserInput( + chat: [.init(role: .system, content: instructions)], + processing: .init(resize: nil), + tools: toolSpecs + ) + let lmInput = try await context.processor.prepare(input: userInput) + + let state: MLXLMCommon.LMOutput.State? = nil + let prepareResult = try context.model.prepare(lmInput, cache: newCache, state: state, windowSize: params.prefill.stepSize) + switch prepareResult { + case .tokens(let tokensToProcess): + _ = context.model(tokensToProcess[text: .newAxis], cache: newCache, state: state) + case .logits: + break + } + storeSessionCache( + cache: newCache, + fullTokens: tokens(from: lmInput), + generateParameters: params, + session: session + ) + } + + public func prewarm( + for session: LanguageModelSession, + promptPrefix: Prompt? + ) async throws { + let modelId = self.modelId + let hub = self.hub + let directory = self.directory + + try await _prewarm( + for: session, + promptPrefix: promptPrefix, + modelID: modelId, + hub: hub, + directory: directory + ) + } + /// Prewarms the model public func prewarm( for session: LanguageModelSession, @@ -1308,48 +1374,13 @@ import Foundation let directory = self.directory Task { - guard Self.acquireGenerationSlot(for: session) else { - return - } - defer { Self.releaseGenerationSlot(for: session) } - - let generationScope = beginGenerationScope() - defer { endGenerationScope(generationScope) } - - do { - let context = try await loadContext(modelId: modelId, hub: hub, directory: directory) - guard let instructions = session.instructions?.description, !instructions.isEmpty else { - return - } - - let toolSpecs = mlxToolSpecs(for: session) - - let params = toGenerateParameters(.init()) - let newCache = context.model.newCache(parameters: params) - let userInput = MLXLMCommon.UserInput( - chat: [.init(role: .system, content: instructions)], - processing: .init(resize: nil), - tools: toolSpecs - ) - let lmInput = try await context.processor.prepare(input: userInput) - - let state: MLXLMCommon.LMOutput.State? = nil - let prepareResult = try context.model.prepare(lmInput, cache: newCache, state: state, windowSize: params.prefill.stepSize) - switch prepareResult { - case .tokens(let tokensToProcess): - _ = context.model(tokensToProcess[text: .newAxis], cache: newCache, state: state) - case .logits: - break - } - storeSessionCache( - cache: newCache, - fullTokens: tokens(from: lmInput), - generateParameters: params, - session: session - ) - } catch { - // Ignore errors during prewarm - } + try await _prewarm( + for: session, + promptPrefix: promptPrefix, + modelID: modelId, + hub: hub, + directory: directory + ) } } } @@ -1361,7 +1392,7 @@ import Foundation /// - Returns: Temperature, topP, and topK. Temperature is a double to match GenerationOptions.temperature private func parametersFromSampling(sampling: GenerationOptions.SamplingMode?) -> (temperature: Double?, topP: Float?, topK: Int?) { guard let sampling else { return (nil, nil, nil) } - + switch sampling.mode { case .greedy: return (1.0, nil, nil) @@ -1370,13 +1401,13 @@ import Foundation case .nucleus(let topP, _): return (nil, Float(topP), nil) } - + } private func toGenerateParameters(_ options: GenerationOptions) -> MLXLMCommon.GenerateParameters { let custom = options[custom: MLXLanguageModel.self] let sampling = parametersFromSampling(sampling: options.sampling) - + return MLXLMCommon.GenerateParameters( maxTokens: options.maximumResponseTokens, maxKVSize: custom?.kvCache.maxSize, @@ -1426,7 +1457,8 @@ import Foundation // Add instructions from session if present and not in transcript if !hasInstructionsInTranscript, let instructions = session.instructions?.description, - !instructions.isEmpty { + !instructions.isEmpty + { chat.append(.init(role: .system, content: instructions)) } @@ -1489,12 +1521,14 @@ import Foundation case .data(let data, _): #if canImport(UIKit) if let uiImage = UIKit.UIImage(data: data), - let ciImage = CIImage(image: uiImage) { + let ciImage = CIImage(image: uiImage) + { images.append(.ciImage(ciImage)) } #elseif canImport(AppKit) if let nsImage = AppKit.NSImage(data: data), - let cgImage = nsImage.cgImage(forProposedRect: nil, context: nil, hints: nil) { + let cgImage = nsImage.cgImage(forProposedRect: nil, context: nil, hints: nil) + { let ciImage = CIImage(cgImage: cgImage) images.append(.ciImage(ciImage)) } @@ -1533,12 +1567,12 @@ import Foundation let functionSpec: [String: any Sendable] = [ "name": tool.name, "description": tool.description, - "parameters": parametersDict + "parameters": parametersDict, ] let toolSpec: ToolSpec = [ "type": "function", - "function": functionSpec + "function": functionSpec, ] return toolSpec @@ -1548,7 +1582,7 @@ import Foundation [ "type": "object", "properties": [String: any Sendable](), - "required": [String]() + "required": [String](), ] } @@ -1653,11 +1687,13 @@ import Foundation if let constValue = jsonSchema.const, let data = try? encoder.encode(constValue), - let constString = String(data: data, encoding: .utf8) { + let constString = String(data: data, encoding: .utf8) + { header += ". Expected value: \(constString)" } else if let enumValues = jsonSchema.enum, !enumValues.isEmpty, let data = try? encoder.encode(enumValues), - let enumString = String(data: data, encoding: .utf8) { + let enumString = String(data: data, encoding: .utf8) + { header += ". Allowed values: \(enumString)" } From 7fc8b924569144cb0f47e3cd49471bf65c5a974a Mon Sep 17 00:00:00 2001 From: Taylor Lineman Date: Tue, 25 Aug 2026 23:55:59 -0400 Subject: [PATCH 2/2] Ran swiflint Signed-off-by: Taylor Lineman --- Package.swift | 16 +++---- Sources/AnyLanguageModel/LanguageModel.swift | 5 +-- .../LanguageModelSession.swift | 3 +- .../Models/MLXLanguageModel.swift | 43 ++++++++----------- 4 files changed, 29 insertions(+), 38 deletions(-) diff --git a/Package.swift b/Package.swift index d1ff645d..7a682953 100644 --- a/Package.swift +++ b/Package.swift @@ -12,7 +12,7 @@ let package = Package( .iOS(.v17), .tvOS(.v17), .watchOS(.v10), - .visionOS(.v1), + .visionOS(.v1) ], products: [ @@ -26,7 +26,7 @@ let package = Package( .trait(name: "MLX"), .trait(name: "Llama"), .trait(name: "AsyncHTTPClient"), - .default(enabledTraits: ["MLX"]), + .default(enabledTraits: ["MLX"]) ], dependencies: [ .package(url: "https://github.com/huggingface/swift-transformers", from: "1.0.0"), @@ -36,7 +36,7 @@ let package = Package( from: "1.3.0", traits: [ .defaults, - .trait(name: "AsyncHTTPClient", condition: .when(traits: ["AsyncHTTPClient"])), + .trait(name: "AsyncHTTPClient", condition: .when(traits: ["AsyncHTTPClient"])) ] ), .package(url: "https://github.com/mattt/JSONSchema", from: "1.3.0"), @@ -46,7 +46,7 @@ let package = Package( // .package(url: "https://github.com/ml-explore/mlx-swift-lm", from: "3.0.0"), .package(url: "https://github.com/impel-intelligence/mlx-swift-lm", from: "1.0.1"), .package(url: "https://github.com/swiftlang/swift-syntax", from: "602.0.0"), - .package(url: "https://github.com/swift-server/async-http-client.git", from: "1.24.0"), + .package(url: "https://github.com/swift-server/async-http-client.git", from: "1.24.0") ], targets: [ .target( @@ -100,7 +100,7 @@ let package = Package( name: "AsyncHTTPClient", package: "async-http-client", condition: .when(traits: ["AsyncHTTPClient"]) - ), + ) ] ), .macro( @@ -109,7 +109,7 @@ let package = Package( .product(name: "SwiftCompilerPlugin", package: "swift-syntax"), .product(name: "SwiftSyntax", package: "swift-syntax"), .product(name: "SwiftSyntaxBuilder", package: "swift-syntax"), - .product(name: "SwiftSyntaxMacros", package: "swift-syntax"), + .product(name: "SwiftSyntaxMacros", package: "swift-syntax") ] ), .testTarget( @@ -120,8 +120,8 @@ let package = Package( name: "AsyncHTTPClient", package: "async-http-client", condition: .when(traits: ["AsyncHTTPClient"]) - ), + ) ], - ), + ) ] ) diff --git a/Sources/AnyLanguageModel/LanguageModel.swift b/Sources/AnyLanguageModel/LanguageModel.swift index 34584347..03c04ac7 100644 --- a/Sources/AnyLanguageModel/LanguageModel.swift +++ b/Sources/AnyLanguageModel/LanguageModel.swift @@ -16,7 +16,7 @@ public protocol LanguageModel: Sendable { for session: LanguageModelSession, promptPrefix: Prompt? ) - + func prewarm( for session: LanguageModelSession, promptPrefix: Prompt? @@ -63,7 +63,7 @@ extension LanguageModel { ) { return } - + public func prewarm( for session: LanguageModelSession, promptPrefix: Prompt? = nil @@ -71,7 +71,6 @@ extension LanguageModel { return } - public func logFeedbackAttachment( within session: LanguageModelSession, sentiment: LanguageModelFeedback.Sentiment? = nil, diff --git a/Sources/AnyLanguageModel/LanguageModelSession.swift b/Sources/AnyLanguageModel/LanguageModelSession.swift index 20fa4547..c5b1d56e 100644 --- a/Sources/AnyLanguageModel/LanguageModelSession.swift +++ b/Sources/AnyLanguageModel/LanguageModelSession.swift @@ -111,12 +111,11 @@ public final class LanguageModelSession: @unchecked Sendable { public func prewarm(promptPrefix: Prompt? = nil) { model.prewarm(for: self, promptPrefix: promptPrefix) } - + public func prewarm(promptPrefix: Prompt? = nil) async throws { try await model.prewarm(for: self, promptPrefix: promptPrefix) } - nonisolated private func beginResponding() { withMutation(keyPath: \.isResponding) { state.withLock { $0.beginResponding() } diff --git a/Sources/AnyLanguageModel/Models/MLXLanguageModel.swift b/Sources/AnyLanguageModel/Models/MLXLanguageModel.swift index 091637a3..41897c70 100644 --- a/Sources/AnyLanguageModel/Models/MLXLanguageModel.swift +++ b/Sources/AnyLanguageModel/Models/MLXLanguageModel.swift @@ -53,8 +53,7 @@ import Foundation ) async throws -> ModelContext { let cacheKey = key as NSString if let cached = cache.object(forKey: cacheKey), - case .loaded(let context) = cached.value - { + case .loaded(let context) = cached.value { return context } @@ -288,7 +287,7 @@ import Foundation var additionalContextForUserInput: [String: any Sendable]? { additionalContext?.mapValues { $0.toSendable() } } - + /// The nucleus-sampling probability threshold. Only tokens whose cumulative /// probability mass falls within the top `topP` fraction are considered during /// sampling. `nil` lets MLX use its model default. @@ -312,7 +311,7 @@ import Foundation /// appeared in the output. Values greater than `1.0` reduce repetition; values less /// than `1.0` encourage it. `nil` lets MLX use its model default. public var repetitionPenalty: Float? - + /// Creates MLX-specific generation options. /// /// - Parameters: @@ -854,8 +853,7 @@ import Foundation let existingEntry = getSessionCache(for: session) if let existingEntry, - isCacheHit(entry: existingEntry, currentTokens: fullTokens, signature: signature, lmInput: lmInput) - { + isCacheHit(entry: existingEntry, currentTokens: fullTokens, signature: signature, lmInput: lmInput) { let cachedCount = existingEntry.prefillTokenCount let newTokens = lmInput.text.tokens[cachedCount...] let newMask = lmInput.text.mask?[cachedCount...] @@ -1330,7 +1328,7 @@ import Foundation tools: toolSpecs ) let lmInput = try await context.processor.prepare(input: userInput) - + let state: MLXLMCommon.LMOutput.State? = nil let prepareResult = try context.model.prepare(lmInput, cache: newCache, state: state, windowSize: params.prefill.stepSize) switch prepareResult { @@ -1346,7 +1344,7 @@ import Foundation session: session ) } - + public func prewarm( for session: LanguageModelSession, promptPrefix: Prompt? @@ -1363,7 +1361,7 @@ import Foundation directory: directory ) } - + /// Prewarms the model public func prewarm( for session: LanguageModelSession, @@ -1392,7 +1390,7 @@ import Foundation /// - Returns: Temperature, topP, and topK. Temperature is a double to match GenerationOptions.temperature private func parametersFromSampling(sampling: GenerationOptions.SamplingMode?) -> (temperature: Double?, topP: Float?, topK: Int?) { guard let sampling else { return (nil, nil, nil) } - + switch sampling.mode { case .greedy: return (1.0, nil, nil) @@ -1401,13 +1399,13 @@ import Foundation case .nucleus(let topP, _): return (nil, Float(topP), nil) } - + } private func toGenerateParameters(_ options: GenerationOptions) -> MLXLMCommon.GenerateParameters { let custom = options[custom: MLXLanguageModel.self] let sampling = parametersFromSampling(sampling: options.sampling) - + return MLXLMCommon.GenerateParameters( maxTokens: options.maximumResponseTokens, maxKVSize: custom?.kvCache.maxSize, @@ -1457,8 +1455,7 @@ import Foundation // Add instructions from session if present and not in transcript if !hasInstructionsInTranscript, let instructions = session.instructions?.description, - !instructions.isEmpty - { + !instructions.isEmpty { chat.append(.init(role: .system, content: instructions)) } @@ -1521,14 +1518,12 @@ import Foundation case .data(let data, _): #if canImport(UIKit) if let uiImage = UIKit.UIImage(data: data), - let ciImage = CIImage(image: uiImage) - { + let ciImage = CIImage(image: uiImage) { images.append(.ciImage(ciImage)) } #elseif canImport(AppKit) if let nsImage = AppKit.NSImage(data: data), - let cgImage = nsImage.cgImage(forProposedRect: nil, context: nil, hints: nil) - { + let cgImage = nsImage.cgImage(forProposedRect: nil, context: nil, hints: nil) { let ciImage = CIImage(cgImage: cgImage) images.append(.ciImage(ciImage)) } @@ -1567,12 +1562,12 @@ import Foundation let functionSpec: [String: any Sendable] = [ "name": tool.name, "description": tool.description, - "parameters": parametersDict, + "parameters": parametersDict ] let toolSpec: ToolSpec = [ "type": "function", - "function": functionSpec, + "function": functionSpec ] return toolSpec @@ -1582,7 +1577,7 @@ import Foundation [ "type": "object", "properties": [String: any Sendable](), - "required": [String](), + "required": [String]() ] } @@ -1687,13 +1682,11 @@ import Foundation if let constValue = jsonSchema.const, let data = try? encoder.encode(constValue), - let constString = String(data: data, encoding: .utf8) - { + let constString = String(data: data, encoding: .utf8) { header += ". Expected value: \(constString)" } else if let enumValues = jsonSchema.enum, !enumValues.isEmpty, let data = try? encoder.encode(enumValues), - let enumString = String(data: data, encoding: .utf8) - { + let enumString = String(data: data, encoding: .utf8) { header += ". Allowed values: \(enumString)" }