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/Sources/AnyLanguageModel/LanguageModel.swift b/Sources/AnyLanguageModel/LanguageModel.swift index 635de68c..03c04ac7 100644 --- a/Sources/AnyLanguageModel/LanguageModel.swift +++ b/Sources/AnyLanguageModel/LanguageModel.swift @@ -17,6 +17,11 @@ public protocol LanguageModel: Sendable { promptPrefix: Prompt? ) + func prewarm( + for session: LanguageModelSession, + promptPrefix: Prompt? + ) async throws + func respond( within session: LanguageModelSession, to prompt: Prompt, @@ -59,6 +64,13 @@ extension LanguageModel { return } + public func prewarm( + for session: LanguageModelSession, + promptPrefix: Prompt? = nil + ) async throws { + return + } + public func logFeedbackAttachment( within session: LanguageModelSession, sentiment: LanguageModelFeedback.Sentiment? = nil, diff --git a/Sources/AnyLanguageModel/LanguageModelSession.swift b/Sources/AnyLanguageModel/LanguageModelSession.swift index a2e94c62..c5b1d56e 100644 --- a/Sources/AnyLanguageModel/LanguageModelSession.swift +++ b/Sources/AnyLanguageModel/LanguageModelSession.swift @@ -112,6 +112,10 @@ public final class LanguageModelSession: @unchecked Sendable { 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 95b42405..b40dc613 100644 --- a/Sources/AnyLanguageModel/Models/MLXLanguageModel.swift +++ b/Sources/AnyLanguageModel/Models/MLXLanguageModel.swift @@ -1300,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, @@ -1310,53 +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 + ) } } }