Skip to content
Merged
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
26 changes: 16 additions & 10 deletions Tests/AnyLanguageModelTests/MLXLanguageModelTests.swift
Original file line number Diff line number Diff line change
Expand Up @@ -40,6 +40,9 @@ import Testing
directory: ProcessInfo.processInfo.environment["MLX_MODEL_DIRECTORY"].map { URL(fileURLWithPath: $0) }
)
let visionModel = MLXLanguageModel(modelId: "mlx-community/Qwen2-VL-2B-Instruct-4bit")
// Text generation has no default token limit, and a model that never emits an end token
// would keep generating. Bound every response so that one can't stall the suite.
let boundedOptions = GenerationOptions(maximumResponseTokens: 512)

@Test func availabilityBecomesAvailableAfterSuccessfulLoad() async throws {
await model.removeFromCache()
Expand All @@ -48,7 +51,7 @@ import Testing
#expect(model.isAvailable == false)

let session = LanguageModelSession(model: model)
let response = try await session.respond(to: "Say hello")
let response = try await session.respond(to: "Say hello", options: boundedOptions)
#expect(!response.content.isEmpty)

#expect(model.availability == .available)
Expand All @@ -58,14 +61,14 @@ import Testing
@Test func basicResponse() async throws {
let session = LanguageModelSession(model: model)

let response = try await session.respond(to: "Say hello")
let response = try await session.respond(to: "Say hello", options: boundedOptions)
#expect(!response.content.isEmpty)
}

@Test func streamingResponse() async throws {
let session = LanguageModelSession(model: model)

let stream = session.streamResponse(to: "Count to 5")
let stream = session.streamResponse(to: "Count to 5", options: boundedOptions)
var chunks: [String] = []

for try await response in stream {
Expand Down Expand Up @@ -265,10 +268,13 @@ import Testing

@Test func multiTurnSameSession() async throws {
let session = LanguageModelSession(model: model)
let first = try await session.respond(to: "Say hello in one sentence.")
let first = try await session.respond(to: "Say hello in one sentence.", options: boundedOptions)
#expect(!first.content.isEmpty)

let second = try await session.respond(to: "Now answer with one more short sentence.")
let second = try await session.respond(
to: "Now answer with one more short sentence.",
options: boundedOptions
)
#expect(!second.content.isEmpty)
}

Expand Down Expand Up @@ -323,7 +329,7 @@ import Testing

let response = try await session.respond(
to: "How's the weather in San Francisco?",
options: GenerationOptions(sampling: .greedy)
options: GenerationOptions(sampling: .greedy, maximumResponseTokens: 512)
)

var foundToolOutput = false
Expand Down Expand Up @@ -354,7 +360,7 @@ import Testing

let stream = session.streamResponse(
to: "How's the weather in San Francisco?",
options: GenerationOptions(sampling: .greedy)
options: GenerationOptions(sampling: .greedy, maximumResponseTokens: 512)
)

// Iterate the stream, keeping the last snapshot as the final state.
Expand Down Expand Up @@ -399,7 +405,7 @@ import Testing
)
])
let session = LanguageModelSession(model: visionModel, transcript: transcript)
var options = GenerationOptions()
var options = boundedOptions
var mlxOptions = MLXLanguageModel.CustomGenerationOptions.default
mlxOptions.userInputProcessing = .resize(to: CGSize(width: 512, height: 512))
options[custom: MLXLanguageModel.self] = mlxOptions
Expand All @@ -417,7 +423,7 @@ import Testing
)
])
let session = LanguageModelSession(model: visionModel, transcript: transcript)
var options = GenerationOptions()
var options = boundedOptions
var mlxOptions = MLXLanguageModel.CustomGenerationOptions.default
mlxOptions.userInputProcessing = .resize(to: CGSize(width: 512, height: 512))
options[custom: MLXLanguageModel.self] = mlxOptions
Expand Down Expand Up @@ -553,7 +559,7 @@ import Testing
@Test func removeAllFromCacheThenRespond() async throws {
await MLXLanguageModel.removeAllFromCache()
let session = LanguageModelSession(model: model)
let response = try await session.respond(to: "Say hello after cache clear")
let response = try await session.respond(to: "Say hello after cache clear", options: boundedOptions)
#expect(!response.content.isEmpty)
}
}
Expand Down
Loading