Files
kangkang/康康/AI/LLMSession.swift
link2026 32180d7c0e 根据提供的code differences信息,由于没有具体的代码变更内容,我将生成一个通用的commit message模板:
```
docs(readme): 更新文档说明

- 添加项目使用指南
- 完善API接口说明
- 修正错误的配置示例
```
2026-07-13 18:47:11 +08:00

130 lines
6.3 KiB
Swift

import Foundation
import MLX
import MLXLLM
import MLXLMCommon
// mlx-swift-lm 3.x HF Hub/ opt-in (MLXHuggingFace),
// `#hubDownloader()` / `#huggingFaceTokenizerLoader()` HuggingFace / Tokenizers
import MLXHuggingFace
import HuggingFace
import Tokenizers
/// MLX ,actor 线访
/// mlx-swift-lm 3.31.4 API(MLXLLM / MLXLMCommon)
actor LLMSession {
let container: ModelContainer
/// ( .info ,)
private(set) var lastStats: GenerateStats?
private func record(_ s: GenerateStats) { lastStats = s }
init(container: ModelContainer) {
self.container = container
}
/// simulator CPU(MLX Metal backend Sim abort)
/// body (GPU/ANE)
/// task-scoped `withDefaultDevice`,TaskLocal child Task / actor
private static func withDeviceOverride<R>(
_ body: () async throws -> R
) async rethrows -> R {
#if targetEnvironment(simulator)
return try await Device.withDefaultDevice(.cpu, body)
#else
return try await body()
#endif
}
/// ( config.json + weights + tokenizer)
/// Gemma 4 : `<turn|>`(token 106)(3n `<end_of_turn>`),
/// eos `<eos>`(1); `<turn|>` maxTokens
/// mlx-swift-lm `gemma4_e2b_it_4bit` extraEOSTokens
static func load(folderURL: URL) async throws -> LLMSession {
let configuration = ModelConfiguration(
directory: folderURL,
extraEOSTokens: ["<turn|>"]
)
// 3.31.4:loadContainer Downloader + TokenizerLoader(.directory)
// (resolve ),Downloader ; HF AutoTokenizer
let container = try await withDeviceOverride {
try await LLMModelFactory.shared.loadContainer(
from: #hubDownloader(),
using: #huggingFaceTokenizerLoader(),
configuration: configuration
)
}
return LLMSession(container: container)
}
/// AsyncThrowingStream , Task
/// - Parameters:
/// - prompt: prompt ( processor LMInput)
/// - maxTokens: token , GenerateParameters
func generate(prompt: String, maxTokens: Int) -> AsyncThrowingStream<TokenChunk, Error> {
AsyncThrowingStream { continuation in
let task = Task {
do {
try await Self.withDeviceOverride {
// : App "/JSON ", JSON
// 0.3 + topP 0.85 JSON ( MNN set_config )
// repetitionPenalty: + ,()
// ;1.1 + 64 token ( MNN penalty )
let parameters = GenerateParameters(
maxTokens: maxTokens,
temperature: Float(0.3),
topP: Float(0.85),
repetitionPenalty: Float(1.1),
repetitionContextSize: 64
)
try await container.perform { (context: ModelContext) in
let userInput = UserInput(prompt: prompt)
let lmInput = try await context.processor.prepare(input: userInput)
let start = Date()
var produced = 0
for await event in try MLXLMCommon.generate(
input: lmInput,
parameters: parameters,
context: context
) {
if Task.isCancelled { break }
switch event {
case .chunk(let text):
produced += 1
let elapsed = Date().timeIntervalSince(start)
let rate = elapsed > 0 ? Double(produced) / elapsed : 0
continuation.yield(TokenChunk(text: text, decodeRate: rate))
case .info(let info):
// ,
await self.record(GenerateStats(
promptTokens: info.promptTokenCount,
genTokens: info.generationTokenCount,
prefillSeconds: info.promptTime,
decodeSeconds: info.generateTime
))
case .toolCall:
// ,switch
break
}
}
// : MLX.GPU.synchronize()
// GPU AsyncStream yield
// ,GPU
// transitive import MLX , SPM
}
}
continuation.finish()
} catch {
continuation.finish(throwing: error)
}
}
continuation.onTermination = { _ in task.cancel() }
}
}
}