Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 1 addition & 1 deletion Package.resolved

Some generated files are not rendered by default. Learn more about how customized files appear on GitHub.

12 changes: 12 additions & 0 deletions Sources/AnyLanguageModel/LanguageModel.swift
Original file line number Diff line number Diff line change
Expand Up @@ -17,6 +17,11 @@ public protocol LanguageModel: Sendable {
promptPrefix: Prompt?
)

func prewarm(
for session: LanguageModelSession,
promptPrefix: Prompt?
) async throws

func respond<Content>(
within session: LanguageModelSession,
to prompt: Prompt,
Expand Down Expand Up @@ -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,
Expand Down
4 changes: 4 additions & 0 deletions Sources/AnyLanguageModel/LanguageModelSession.swift
Original file line number Diff line number Diff line change
Expand Up @@ -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() }
Expand Down
118 changes: 71 additions & 47 deletions Sources/AnyLanguageModel/Models/MLXLanguageModel.swift
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand All @@ -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
)
}
}
}
Expand Down
Loading