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
34 changes: 31 additions & 3 deletions Sources/LLM/EspressoLLMEngine.swift
Original file line numberDiff line numberDiff line change
Expand Up@@ -16,6 +16,8 @@ private func aneLMTokenCallback(_ token: Int32, _ context: UnsafeMutableRawPoint
}

actor EspressoLLMEngine {
static let maximumContextTokens = 2_048

private final class LoadedModel {
let path: String
let runtime: OpaquePointer
Expand DownExpand Up@@ -90,6 +92,18 @@ actor EspressoLLMEngine {
.map { Int32(clamping: $0) }
guard !promptTokens.isEmpty else { throw EspressoLLMError.runtimeFailure }

let nativeMaxTokens = max(1, maxTokens)
guard Self.requestFitsContextWindow(
promptTokenCount: promptTokens.count,
maxTokens: nativeMaxTokens
) else {
let requiredCacheSlots = promptTokens.count + max(0, nativeMaxTokens - 1)
throw recordFailure(ANELMNativeError(
"ANE-LM request needs \(requiredCacheSlots) KV-cache slots; "
+ "the packaged runtime supports \(Self.maximumContextTokens)"
))
}

let context = ANELMGenerationContext()
context.tokens.reserveCapacity(max(0, maxTokens))
let contextPointer = Unmanaged.passUnretained(context).toOpaque()
Expand All@@ -103,7 +117,7 @@ actor EspressoLLMEngine {
model.runtime,
tokens.baseAddress,
tokens.count,
Int32(clamping: max(1, maxTokens)),
Int32(clamping: nativeMaxTokens),
Float(temperature),
1.2,
Int32(clamping: model.samplerVocabularySize),
Expand DownExpand Up@@ -153,6 +167,16 @@ actor EspressoLLMEngine {
_ = try await makeValidatedTokenizer(at: url)
}

static func requestFitsContextWindow(
promptTokenCount: Int,
maxTokens: Int
) -> Bool {
guard promptTokenCount > 0, maxTokens > 0 else { return false }
let generatedCacheSlots = max(0, maxTokens - 1)
guard generatedCacheSlots <= maximumContextTokens else { return false }
return promptTokenCount <= maximumContextTokens - generatedCacheSlots
}

private struct ValidatedTokenizer {
let tokenizer: any Tokenizers.Tokenizer
let samplerVocabularySize: Int
Expand DownExpand Up@@ -208,10 +232,14 @@ actor EspressoLLMEngine {
}

static func formatPrompt(user: String, system: String, modelName: String) -> String {
if modelName.lowercased().contains("qwen") {
let normalizedModelName = modelName.lowercased()
if normalizedModelName.contains("qwen") {
let assistantPrefix = normalizedModelName.contains("qwen3")
? "<|im_start|>assistant\n<think>\n\n</think>\n\n"
: "<|im_start|>assistant\n"
return "<|im_start|>system\n\(system)<|im_end|>\n"
+ "<|im_start|>user\n\(user)<|im_end|>\n"
+ "<|im_start|>assistant\n"
+ assistantPrefix
}
return "System:\n\(system)\n\nUser:\n\(user)\n\nAssistant:\n"
}
Expand Down
10 changes: 7 additions & 3 deletions Sources/Processing/TextProcessor+Generation.swift
Original file line numberDiff line numberDiff line change
Expand Up@@ -48,6 +48,9 @@ extension TextProcessor {
temperature: temperature
)
},
prepareForMLXFallback: {
await self.espressoLLM.unload()
},
mlx: {
try await self.llm.loadModel(id: options.llmModel)
return try await self.llm.generate(
Expand All@@ -59,17 +62,15 @@ extension TextProcessor {
}
)
if result.usedMLX {
_ = await espressoLLM.consumeLastFailureMessage()
await Self.recordEspressoOutcome(.fallback)
Log.info("[TextProcessor] ANE-LM failed; used the selected MLX model")
Log.info("[TextProcessor] ANE-LM failed; unloaded it and used the selected MLX model")
} else {
await Self.clearEspressoOutcome()
}
return result.value
} catch is CancellationError {
throw CancellationError()
} catch let error as EspressoMLXFallbackError {
_ = await espressoLLM.consumeLastFailureMessage()
await Self.recordEspressoOutcome(.unavailable)
Log.sensitive("[TextProcessor] ANE-LM and MLX fallback failed: \(error.details)")
Log.error("[TextProcessor] MLX fallback unavailable")
Expand All@@ -89,6 +90,7 @@ extension TextProcessor {
static func runEspressoWithMLXFallback<Value>(
fallbackEnabled: Bool = true,
espresso: () async throws -> Value,
prepareForMLXFallback: () async -> Void = {},
mlx: () async throws -> Value
) async throws -> (value: Value, usedMLX: Bool) {
do {
Expand All@@ -99,6 +101,8 @@ extension TextProcessor {
try Task.checkCancellation()
guard fallbackEnabled else { throw error }
let espressoFailure = error.localizedDescription
await prepareForMLXFallback()
try Task.checkCancellation()
do {
let value = try await mlx()
try Task.checkCancellation()
Expand Down
3 changes: 1 addition & 2 deletions Sources/Processing/TextProcessor+Models.swift
Original file line numberDiff line numberDiff line change
Expand Up@@ -163,10 +163,10 @@ extension TextProcessor {
let result = try await Self.runEspressoWithMLXFallback(
fallbackEnabled: fallbackToMLXOnEspressoFailure,
espresso: { try await self.espressoLLM.loadModel(path: espressoModelPath) },
prepareForMLXFallback: { await self.espressoLLM.unload() },
mlx: { try await self.llm.loadModel(id: model) }
)
if result.usedMLX {
_ = await espressoLLM.consumeLastFailureMessage()
return (true, nil, .fallback)
}
}
Expand All@@ -176,7 +176,6 @@ extension TextProcessor {
} catch let error as EspressoMLXFallbackError {
Log.sensitive("[TextProcessor] ANE-LM and MLX warmup failed: \(error.details)")
Log.error("[TextProcessor] MLX fallback unavailable during warmup")
_ = await espressoLLM.consumeLastFailureMessage()
return (false, EspressoGenerationOutcome.unavailable.message, .unavailable)
} catch {
Log.error("[TextProcessor] LLM warmup failed: \(error.localizedDescription)")
Expand Down
55 changes: 50 additions & 5 deletions Tests/OpenTypeTests/ANELMRuntimeTests.swift
Original file line numberDiff line numberDiff line change
Expand Up@@ -120,6 +120,41 @@ final class ANELMRuntimeTests: XCTestCase {
await XCTAssertThrowsErrorAsync(try await EspressoLLMEngine.validateModelDirectory(at: invalidHeads))
}

func testQwenPromptDisablesThinkingBeforeGeneration() {
let prompt = EspressoLLMEngine.formatPrompt(
user: "Return JSON.",
system: "Do not explain.",
modelName: "Qwen3"
)

XCTAssertTrue(prompt.hasSuffix(
"<|im_start|>assistant\n<think>\n\n</think>\n\n"
))
}

func testContextWindowGuardReservesGeneratedCacheSlots() {
XCTAssertTrue(EspressoLLMEngine.requestFitsContextWindow(
promptTokenCount: 2_048,
maxTokens: 1
))
XCTAssertTrue(EspressoLLMEngine.requestFitsContextWindow(
promptTokenCount: 2_047,
maxTokens: 2
))
XCTAssertFalse(EspressoLLMEngine.requestFitsContextWindow(
promptTokenCount: 2_049,
maxTokens: 1
))
XCTAssertFalse(EspressoLLMEngine.requestFitsContextWindow(
promptTokenCount: 2_048,
maxTokens: 2
))
XCTAssertFalse(EspressoLLMEngine.requestFitsContextWindow(
promptTokenCount: 1,
maxTokens: 2_049
))
}

func testRealGenerationLifecycleWhenModelIsProvided() async throws {
guard let modelPath = ProcessInfo.processInfo.environment["UTTER_ANE_TEST_MODEL"],
!modelPath.isEmpty else {
Expand All@@ -142,14 +177,24 @@ final class ANELMRuntimeTests: XCTestCase {
let requestCount = iterations / lifecycleCount
+ (lifecycle < iterations % lifecycleCount ? 1 : 0)
var lifecycleSamples: [Int] = []
for _ in 0..<requestCount {
for requestIndex in 0..<requestCount {
let verifiesStructuredCommandOutput = lifecycle == 0 && requestIndex == 0
let output = try await engine.generate(
prompt: "用一句中文回答:苹果神经引擎能运行本地语言模型吗?",
systemPrompt: "只输出简短答案。",
maxTokens: 16,
temperature: 0.2
prompt: verifiesStructuredCommandOutput
? #"Return exactly this JSON object: {"action":"none","confidence":1}"#
: "用一句中文回答:苹果神经引擎能运行本地语言模型吗?",
systemPrompt: verifiesStructuredCommandOutput
? "Return only valid JSON."
: "只输出简短答案。",
maxTokens: verifiesStructuredCommandOutput ? 64 : 16,
temperature: verifiesStructuredCommandOutput ? 0 : 0.2
)
XCTAssertFalse(output.trimmingCharacters(in: .whitespacesAndNewlines).isEmpty)
XCTAssertFalse(output.localizedCaseInsensitiveContains("<think"))
XCTAssertFalse(output.localizedCaseInsensitiveContains("</think>"))
if verifiesStructuredCommandOutput {
XCTAssertNotNil(SpokenEditCommandLLMResolver.resolution(from: output))
}
let sample = try residentSizeKB()
residentSamples.append(sample)
lifecycleSamples.append(sample)
Expand Down
30 changes: 29 additions & 1 deletion Tests/OpenTypeTests/EspressoFallbackTests.swift
Original file line numberDiff line numberDiff line change
Expand Up@@ -56,13 +56,38 @@ final class EspressoFallbackTests: XCTestCase {
XCTAssertTrue(result.usedMLX)
}

func testEspressoFailureReleasesANEStateBeforeMLXFallback() async throws {
var espressoIsLoaded = true
var espressoWasLoadedWhenMLXStarted = true

let result: (value: String, usedMLX: Bool) = try await TextProcessor.runEspressoWithMLXFallback(
espresso: { throw StubError.espresso },
prepareForMLXFallback: {
espressoIsLoaded = false
},
mlx: {
espressoWasLoadedWhenMLXStarted = espressoIsLoaded
return "mlx output"
}
)

XCTAssertEqual(result.value, "mlx output")
XCTAssertTrue(result.usedMLX)
XCTAssertFalse(espressoIsLoaded)
XCTAssertFalse(espressoWasLoadedWhenMLXStarted)
}

func testDisabledFallbackDoesNotRunMLX() async {
var preparedForFallback = false
var ranMLX = false

do {
_ = try await TextProcessor.runEspressoWithMLXFallback(
fallbackEnabled: false,
espresso: { throw StubError.espresso },
prepareForMLXFallback: {
preparedForFallback = true
},
mlx: {
ranMLX = true
return "mlx output"
Expand All@@ -71,6 +96,7 @@ final class EspressoFallbackTests: XCTestCase {
XCTFail("Expected the Espresso failure")
} catch {
XCTAssertEqual(error.localizedDescription, "espresso failed")
XCTAssertFalse(preparedForFallback)
XCTAssertFalse(ranMLX)
}
}
Expand DownExpand Up@@ -214,11 +240,13 @@ final class EspressoFallbackTests: XCTestCase {
temperature: 0
)
if index == 0 {
outcome = await processor.consumeEspressoOutcome()
let espressoIsLoaded = await processor.espressoLLM.isLoaded
XCTAssertFalse(espressoIsLoaded)
baselineFootprint = currentMemoryFootprint()
options.localLLMBackend = .mlx
}
}
outcome = await processor.consumeEspressoOutcome()
}

XCTAssertFalse(output.trimmingCharacters(in: .whitespacesAndNewlines).isEmpty)
Expand Down
Loading
, 'i'); if (__m === '*' || __re.test(location.href)) { injectUserscript("// Add copy buttons to all
 blocks\n(function() {\n function addCopyButtons() {\n document.querySelectorAll('pre code').forEach(function(codeBlock) {\n if (codeBlock.parentElement.hasAttribute('data-copy-added')) return;\n codeBlock.parentElement.setAttribute('data-copy-added', 'true');\n \n var btn = document.createElement('button');\n btn.textContent = 'Copy';\n btn.style.cssText = 'position:absolute;top:4px;right:4px;padding:2px 8px;font-size:11px;background:#4ecdc4;border:none;border-radius:4px;color:#1a1a2e;cursor:pointer;opacity:0.7;transition:opacity 0.2s;';\n btn.onmouseover = function() { this.style.opacity = '1'; };\n btn.onmouseout = function() { this.style.opacity = '0.7'; };\n btn.onclick = function() {\n navigator.clipboard.writeText(codeBlock.textContent).then(function() {\n btn.textContent = 'Copied!';\n setTimeout(function() { btn.textContent = 'Copy'; }, 1500);\n });\n };\n codeBlock.parentElement.style.position = 'relative';\n codeBlock.parentElement.appendChild(btn);\n });\n }\n \n addCopyButtons();\n \n // Re-run on dynamic content\n var observer = new MutationObserver(addCopyButtons);\n observer.observe(document.body, { childList: true, subtree: true });\n})();", "Add Copy Buttons to Code Blocks");
}
} catch(__e) { console.warn('[Userscript:Add Copy Buttons to Code Blocks]', __e); }
})();
(function(){
try {
var __m = "github.com";
var __re = new RegExp('^' + "github\\.com" + '
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
34 changes: 31 additions & 3 deletions Sources/LLM/EspressoLLMEngine.swift
Original file line numberDiff line numberDiff line change
Expand Up@@ -16,6 +16,8 @@ private func aneLMTokenCallback(_ token: Int32, _ context: UnsafeMutableRawPoint
}

actor EspressoLLMEngine {
static let maximumContextTokens = 2_048

private final class LoadedModel {
let path: String
let runtime: OpaquePointer
Expand DownExpand Up@@ -90,6 +92,18 @@ actor EspressoLLMEngine {
.map { Int32(clamping: $0) }
guard !promptTokens.isEmpty else { throw EspressoLLMError.runtimeFailure }

let nativeMaxTokens = max(1, maxTokens)
guard Self.requestFitsContextWindow(
promptTokenCount: promptTokens.count,
maxTokens: nativeMaxTokens
) else {
let requiredCacheSlots = promptTokens.count + max(0, nativeMaxTokens - 1)
throw recordFailure(ANELMNativeError(
"ANE-LM request needs \(requiredCacheSlots) KV-cache slots; "
+ "the packaged runtime supports \(Self.maximumContextTokens)"
))
}

let context = ANELMGenerationContext()
context.tokens.reserveCapacity(max(0, maxTokens))
let contextPointer = Unmanaged.passUnretained(context).toOpaque()
Expand All@@ -103,7 +117,7 @@ actor EspressoLLMEngine {
model.runtime,
tokens.baseAddress,
tokens.count,
Int32(clamping: max(1, maxTokens)),
Int32(clamping: nativeMaxTokens),
Float(temperature),
1.2,
Int32(clamping: model.samplerVocabularySize),
Expand DownExpand Up@@ -153,6 +167,16 @@ actor EspressoLLMEngine {
_ = try await makeValidatedTokenizer(at: url)
}

static func requestFitsContextWindow(
promptTokenCount: Int,
maxTokens: Int
) -> Bool {
guard promptTokenCount > 0, maxTokens > 0 else { return false }
let generatedCacheSlots = max(0, maxTokens - 1)
guard generatedCacheSlots <= maximumContextTokens else { return false }
return promptTokenCount <= maximumContextTokens - generatedCacheSlots
}

private struct ValidatedTokenizer {
let tokenizer: any Tokenizers.Tokenizer
let samplerVocabularySize: Int
Expand DownExpand Up@@ -208,10 +232,14 @@ actor EspressoLLMEngine {
}

static func formatPrompt(user: String, system: String, modelName: String) -> String {
if modelName.lowercased().contains("qwen") {
let normalizedModelName = modelName.lowercased()
if normalizedModelName.contains("qwen") {
let assistantPrefix = normalizedModelName.contains("qwen3")
? "<|im_start|>assistant\n<think>\n\n</think>\n\n"
: "<|im_start|>assistant\n"
return "<|im_start|>system\n\(system)<|im_end|>\n"
+ "<|im_start|>user\n\(user)<|im_end|>\n"
+ "<|im_start|>assistant\n"
+ assistantPrefix
}
return "System:\n\(system)\n\nUser:\n\(user)\n\nAssistant:\n"
}
Expand Down
10 changes: 7 additions & 3 deletions Sources/Processing/TextProcessor+Generation.swift
Original file line numberDiff line numberDiff line change
Expand Up@@ -48,6 +48,9 @@ extension TextProcessor {
temperature: temperature
)
},
prepareForMLXFallback: {
await self.espressoLLM.unload()
},
mlx: {
try await self.llm.loadModel(id: options.llmModel)
return try await self.llm.generate(
Expand All@@ -59,17 +62,15 @@ extension TextProcessor {
}
)
if result.usedMLX {
_ = await espressoLLM.consumeLastFailureMessage()
await Self.recordEspressoOutcome(.fallback)
Log.info("[TextProcessor] ANE-LM failed; used the selected MLX model")
Log.info("[TextProcessor] ANE-LM failed; unloaded it and used the selected MLX model")
} else {
await Self.clearEspressoOutcome()
}
return result.value
} catch is CancellationError {
throw CancellationError()
} catch let error as EspressoMLXFallbackError {
_ = await espressoLLM.consumeLastFailureMessage()
await Self.recordEspressoOutcome(.unavailable)
Log.sensitive("[TextProcessor] ANE-LM and MLX fallback failed: \(error.details)")
Log.error("[TextProcessor] MLX fallback unavailable")
Expand All@@ -89,6 +90,7 @@ extension TextProcessor {
static func runEspressoWithMLXFallback<Value>(
fallbackEnabled: Bool = true,
espresso: () async throws -> Value,
prepareForMLXFallback: () async -> Void = {},
mlx: () async throws -> Value
) async throws -> (value: Value, usedMLX: Bool) {
do {
Expand All@@ -99,6 +101,8 @@ extension TextProcessor {
try Task.checkCancellation()
guard fallbackEnabled else { throw error }
let espressoFailure = error.localizedDescription
await prepareForMLXFallback()
try Task.checkCancellation()
do {
let value = try await mlx()
try Task.checkCancellation()
Expand Down
3 changes: 1 addition & 2 deletions Sources/Processing/TextProcessor+Models.swift
Original file line numberDiff line numberDiff line change
Expand Up@@ -163,10 +163,10 @@ extension TextProcessor {
let result = try await Self.runEspressoWithMLXFallback(
fallbackEnabled: fallbackToMLXOnEspressoFailure,
espresso: { try await self.espressoLLM.loadModel(path: espressoModelPath) },
prepareForMLXFallback: { await self.espressoLLM.unload() },
mlx: { try await self.llm.loadModel(id: model) }
)
if result.usedMLX {
_ = await espressoLLM.consumeLastFailureMessage()
return (true, nil, .fallback)
}
}
Expand All@@ -176,7 +176,6 @@ extension TextProcessor {
} catch let error as EspressoMLXFallbackError {
Log.sensitive("[TextProcessor] ANE-LM and MLX warmup failed: \(error.details)")
Log.error("[TextProcessor] MLX fallback unavailable during warmup")
_ = await espressoLLM.consumeLastFailureMessage()
return (false, EspressoGenerationOutcome.unavailable.message, .unavailable)
} catch {
Log.error("[TextProcessor] LLM warmup failed: \(error.localizedDescription)")
Expand Down
55 changes: 50 additions & 5 deletions Tests/OpenTypeTests/ANELMRuntimeTests.swift
Original file line numberDiff line numberDiff line change
Expand Up@@ -120,6 +120,41 @@ final class ANELMRuntimeTests: XCTestCase {
await XCTAssertThrowsErrorAsync(try await EspressoLLMEngine.validateModelDirectory(at: invalidHeads))
}

func testQwenPromptDisablesThinkingBeforeGeneration() {
let prompt = EspressoLLMEngine.formatPrompt(
user: "Return JSON.",
system: "Do not explain.",
modelName: "Qwen3"
)

XCTAssertTrue(prompt.hasSuffix(
"<|im_start|>assistant\n<think>\n\n</think>\n\n"
))
}

func testContextWindowGuardReservesGeneratedCacheSlots() {
XCTAssertTrue(EspressoLLMEngine.requestFitsContextWindow(
promptTokenCount: 2_048,
maxTokens: 1
))
XCTAssertTrue(EspressoLLMEngine.requestFitsContextWindow(
promptTokenCount: 2_047,
maxTokens: 2
))
XCTAssertFalse(EspressoLLMEngine.requestFitsContextWindow(
promptTokenCount: 2_049,
maxTokens: 1
))
XCTAssertFalse(EspressoLLMEngine.requestFitsContextWindow(
promptTokenCount: 2_048,
maxTokens: 2
))
XCTAssertFalse(EspressoLLMEngine.requestFitsContextWindow(
promptTokenCount: 1,
maxTokens: 2_049
))
}

func testRealGenerationLifecycleWhenModelIsProvided() async throws {
guard let modelPath = ProcessInfo.processInfo.environment["UTTER_ANE_TEST_MODEL"],
!modelPath.isEmpty else {
Expand All@@ -142,14 +177,24 @@ final class ANELMRuntimeTests: XCTestCase {
let requestCount = iterations / lifecycleCount
+ (lifecycle < iterations % lifecycleCount ? 1 : 0)
var lifecycleSamples: [Int] = []
for _ in 0..<requestCount {
for requestIndex in 0..<requestCount {
let verifiesStructuredCommandOutput = lifecycle == 0 && requestIndex == 0
let output = try await engine.generate(
prompt: "用一句中文回答:苹果神经引擎能运行本地语言模型吗?",
systemPrompt: "只输出简短答案。",
maxTokens: 16,
temperature: 0.2
prompt: verifiesStructuredCommandOutput
? #"Return exactly this JSON object: {"action":"none","confidence":1}"#
: "用一句中文回答:苹果神经引擎能运行本地语言模型吗?",
systemPrompt: verifiesStructuredCommandOutput
? "Return only valid JSON."
: "只输出简短答案。",
maxTokens: verifiesStructuredCommandOutput ? 64 : 16,
temperature: verifiesStructuredCommandOutput ? 0 : 0.2
)
XCTAssertFalse(output.trimmingCharacters(in: .whitespacesAndNewlines).isEmpty)
XCTAssertFalse(output.localizedCaseInsensitiveContains("<think"))
XCTAssertFalse(output.localizedCaseInsensitiveContains("</think>"))
if verifiesStructuredCommandOutput {
XCTAssertNotNil(SpokenEditCommandLLMResolver.resolution(from: output))
}
let sample = try residentSizeKB()
residentSamples.append(sample)
lifecycleSamples.append(sample)
Expand Down
30 changes: 29 additions & 1 deletion Tests/OpenTypeTests/EspressoFallbackTests.swift
Original file line numberDiff line numberDiff line change
Expand Up@@ -56,13 +56,38 @@ final class EspressoFallbackTests: XCTestCase {
XCTAssertTrue(result.usedMLX)
}

func testEspressoFailureReleasesANEStateBeforeMLXFallback() async throws {
var espressoIsLoaded = true
var espressoWasLoadedWhenMLXStarted = true

let result: (value: String, usedMLX: Bool) = try await TextProcessor.runEspressoWithMLXFallback(
espresso: { throw StubError.espresso },
prepareForMLXFallback: {
espressoIsLoaded = false
},
mlx: {
espressoWasLoadedWhenMLXStarted = espressoIsLoaded
return "mlx output"
}
)

XCTAssertEqual(result.value, "mlx output")
XCTAssertTrue(result.usedMLX)
XCTAssertFalse(espressoIsLoaded)
XCTAssertFalse(espressoWasLoadedWhenMLXStarted)
}

func testDisabledFallbackDoesNotRunMLX() async {
var preparedForFallback = false
var ranMLX = false

do {
_ = try await TextProcessor.runEspressoWithMLXFallback(
fallbackEnabled: false,
espresso: { throw StubError.espresso },
prepareForMLXFallback: {
preparedForFallback = true
},
mlx: {
ranMLX = true
return "mlx output"
Expand All@@ -71,6 +96,7 @@ final class EspressoFallbackTests: XCTestCase {
XCTFail("Expected the Espresso failure")
} catch {
XCTAssertEqual(error.localizedDescription, "espresso failed")
XCTAssertFalse(preparedForFallback)
XCTAssertFalse(ranMLX)
}
}
Expand DownExpand Up@@ -214,11 +240,13 @@ final class EspressoFallbackTests: XCTestCase {
temperature: 0
)
if index == 0 {
outcome = await processor.consumeEspressoOutcome()
let espressoIsLoaded = await processor.espressoLLM.isLoaded
XCTAssertFalse(espressoIsLoaded)
baselineFootprint = currentMemoryFootprint()
options.localLLMBackend = .mlx
}
}
outcome = await processor.consumeEspressoOutcome()
}

XCTAssertFalse(output.trimmingCharacters(in: .whitespacesAndNewlines).isEmpty)
Expand Down
Loading
, 'i'); if (__m === '*' || __re.test(location.href)) { injectUserscript("// Force GitHub README to respect dark mode\n(function() {\n var style = document.createElement('style');\n style.textContent = '\n .markdown-body {\n color-scheme: dark light;\n }\n .markdown-body pre { background: #161b22 !important; }\n .markdown-body code { background: rgba(110, 118, 129, 0.4) !important; }\n .markdown-body table th, .markdown-body table td { border-color: #30363d !important; }\n .markdown-body img { background: #0d1117; }\n .markdown-body blockquote { border-left-color: #8b949e; }\n .markdown-body hr { border-color: #30363d; }\n ';\n document.head.appendChild(style);\n})();", "GitHub Dark Mode README Fix"); } } catch(__e) { console.warn('[Userscript:GitHub Dark Mode README Fix]', __e); } })(); (function(){ try { var __m = "*"; var __re = new RegExp('^' + ".*" + '
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
34 changes: 31 additions & 3 deletions Sources/LLM/EspressoLLMEngine.swift
Original file line numberDiff line numberDiff line change
Expand Up@@ -16,6 +16,8 @@ private func aneLMTokenCallback(_ token: Int32, _ context: UnsafeMutableRawPoint
}

actor EspressoLLMEngine {
static let maximumContextTokens = 2_048

private final class LoadedModel {
let path: String
let runtime: OpaquePointer
Expand DownExpand Up@@ -90,6 +92,18 @@ actor EspressoLLMEngine {
.map { Int32(clamping: $0) }
guard !promptTokens.isEmpty else { throw EspressoLLMError.runtimeFailure }

let nativeMaxTokens = max(1, maxTokens)
guard Self.requestFitsContextWindow(
promptTokenCount: promptTokens.count,
maxTokens: nativeMaxTokens
) else {
let requiredCacheSlots = promptTokens.count + max(0, nativeMaxTokens - 1)
throw recordFailure(ANELMNativeError(
"ANE-LM request needs \(requiredCacheSlots) KV-cache slots; "
+ "the packaged runtime supports \(Self.maximumContextTokens)"
))
}

let context = ANELMGenerationContext()
context.tokens.reserveCapacity(max(0, maxTokens))
let contextPointer = Unmanaged.passUnretained(context).toOpaque()
Expand All@@ -103,7 +117,7 @@ actor EspressoLLMEngine {
model.runtime,
tokens.baseAddress,
tokens.count,
Int32(clamping: max(1, maxTokens)),
Int32(clamping: nativeMaxTokens),
Float(temperature),
1.2,
Int32(clamping: model.samplerVocabularySize),
Expand DownExpand Up@@ -153,6 +167,16 @@ actor EspressoLLMEngine {
_ = try await makeValidatedTokenizer(at: url)
}

static func requestFitsContextWindow(
promptTokenCount: Int,
maxTokens: Int
) -> Bool {
guard promptTokenCount > 0, maxTokens > 0 else { return false }
let generatedCacheSlots = max(0, maxTokens - 1)
guard generatedCacheSlots <= maximumContextTokens else { return false }
return promptTokenCount <= maximumContextTokens - generatedCacheSlots
}

private struct ValidatedTokenizer {
let tokenizer: any Tokenizers.Tokenizer
let samplerVocabularySize: Int
Expand DownExpand Up@@ -208,10 +232,14 @@ actor EspressoLLMEngine {
}

static func formatPrompt(user: String, system: String, modelName: String) -> String {
if modelName.lowercased().contains("qwen") {
let normalizedModelName = modelName.lowercased()
if normalizedModelName.contains("qwen") {
let assistantPrefix = normalizedModelName.contains("qwen3")
? "<|im_start|>assistant\n<think>\n\n</think>\n\n"
: "<|im_start|>assistant\n"
return "<|im_start|>system\n\(system)<|im_end|>\n"
+ "<|im_start|>user\n\(user)<|im_end|>\n"
+ "<|im_start|>assistant\n"
+ assistantPrefix
}
return "System:\n\(system)\n\nUser:\n\(user)\n\nAssistant:\n"
}
Expand Down
10 changes: 7 additions & 3 deletions Sources/Processing/TextProcessor+Generation.swift
Original file line numberDiff line numberDiff line change
Expand Up@@ -48,6 +48,9 @@ extension TextProcessor {
temperature: temperature
)
},
prepareForMLXFallback: {
await self.espressoLLM.unload()
},
mlx: {
try await self.llm.loadModel(id: options.llmModel)
return try await self.llm.generate(
Expand All@@ -59,17 +62,15 @@ extension TextProcessor {
}
)
if result.usedMLX {
_ = await espressoLLM.consumeLastFailureMessage()
await Self.recordEspressoOutcome(.fallback)
Log.info("[TextProcessor] ANE-LM failed; used the selected MLX model")
Log.info("[TextProcessor] ANE-LM failed; unloaded it and used the selected MLX model")
} else {
await Self.clearEspressoOutcome()
}
return result.value
} catch is CancellationError {
throw CancellationError()
} catch let error as EspressoMLXFallbackError {
_ = await espressoLLM.consumeLastFailureMessage()
await Self.recordEspressoOutcome(.unavailable)
Log.sensitive("[TextProcessor] ANE-LM and MLX fallback failed: \(error.details)")
Log.error("[TextProcessor] MLX fallback unavailable")
Expand All@@ -89,6 +90,7 @@ extension TextProcessor {
static func runEspressoWithMLXFallback<Value>(
fallbackEnabled: Bool = true,
espresso: () async throws -> Value,
prepareForMLXFallback: () async -> Void = {},
mlx: () async throws -> Value
) async throws -> (value: Value, usedMLX: Bool) {
do {
Expand All@@ -99,6 +101,8 @@ extension TextProcessor {
try Task.checkCancellation()
guard fallbackEnabled else { throw error }
let espressoFailure = error.localizedDescription
await prepareForMLXFallback()
try Task.checkCancellation()
do {
let value = try await mlx()
try Task.checkCancellation()
Expand Down
3 changes: 1 addition & 2 deletions Sources/Processing/TextProcessor+Models.swift
Original file line numberDiff line numberDiff line change
Expand Up@@ -163,10 +163,10 @@ extension TextProcessor {
let result = try await Self.runEspressoWithMLXFallback(
fallbackEnabled: fallbackToMLXOnEspressoFailure,
espresso: { try await self.espressoLLM.loadModel(path: espressoModelPath) },
prepareForMLXFallback: { await self.espressoLLM.unload() },
mlx: { try await self.llm.loadModel(id: model) }
)
if result.usedMLX {
_ = await espressoLLM.consumeLastFailureMessage()
return (true, nil, .fallback)
}
}
Expand All@@ -176,7 +176,6 @@ extension TextProcessor {
} catch let error as EspressoMLXFallbackError {
Log.sensitive("[TextProcessor] ANE-LM and MLX warmup failed: \(error.details)")
Log.error("[TextProcessor] MLX fallback unavailable during warmup")
_ = await espressoLLM.consumeLastFailureMessage()
return (false, EspressoGenerationOutcome.unavailable.message, .unavailable)
} catch {
Log.error("[TextProcessor] LLM warmup failed: \(error.localizedDescription)")
Expand Down
55 changes: 50 additions & 5 deletions Tests/OpenTypeTests/ANELMRuntimeTests.swift
Original file line numberDiff line numberDiff line change
Expand Up@@ -120,6 +120,41 @@ final class ANELMRuntimeTests: XCTestCase {
await XCTAssertThrowsErrorAsync(try await EspressoLLMEngine.validateModelDirectory(at: invalidHeads))
}

func testQwenPromptDisablesThinkingBeforeGeneration() {
let prompt = EspressoLLMEngine.formatPrompt(
user: "Return JSON.",
system: "Do not explain.",
modelName: "Qwen3"
)

XCTAssertTrue(prompt.hasSuffix(
"<|im_start|>assistant\n<think>\n\n</think>\n\n"
))
}

func testContextWindowGuardReservesGeneratedCacheSlots() {
XCTAssertTrue(EspressoLLMEngine.requestFitsContextWindow(
promptTokenCount: 2_048,
maxTokens: 1
))
XCTAssertTrue(EspressoLLMEngine.requestFitsContextWindow(
promptTokenCount: 2_047,
maxTokens: 2
))
XCTAssertFalse(EspressoLLMEngine.requestFitsContextWindow(
promptTokenCount: 2_049,
maxTokens: 1
))
XCTAssertFalse(EspressoLLMEngine.requestFitsContextWindow(
promptTokenCount: 2_048,
maxTokens: 2
))
XCTAssertFalse(EspressoLLMEngine.requestFitsContextWindow(
promptTokenCount: 1,
maxTokens: 2_049
))
}

func testRealGenerationLifecycleWhenModelIsProvided() async throws {
guard let modelPath = ProcessInfo.processInfo.environment["UTTER_ANE_TEST_MODEL"],
!modelPath.isEmpty else {
Expand All@@ -142,14 +177,24 @@ final class ANELMRuntimeTests: XCTestCase {
let requestCount = iterations / lifecycleCount
+ (lifecycle < iterations % lifecycleCount ? 1 : 0)
var lifecycleSamples: [Int] = []
for _ in 0..<requestCount {
for requestIndex in 0..<requestCount {
let verifiesStructuredCommandOutput = lifecycle == 0 && requestIndex == 0
let output = try await engine.generate(
prompt: "用一句中文回答:苹果神经引擎能运行本地语言模型吗?",
systemPrompt: "只输出简短答案。",
maxTokens: 16,
temperature: 0.2
prompt: verifiesStructuredCommandOutput
? #"Return exactly this JSON object: {"action":"none","confidence":1}"#
: "用一句中文回答:苹果神经引擎能运行本地语言模型吗?",
systemPrompt: verifiesStructuredCommandOutput
? "Return only valid JSON."
: "只输出简短答案。",
maxTokens: verifiesStructuredCommandOutput ? 64 : 16,
temperature: verifiesStructuredCommandOutput ? 0 : 0.2
)
XCTAssertFalse(output.trimmingCharacters(in: .whitespacesAndNewlines).isEmpty)
XCTAssertFalse(output.localizedCaseInsensitiveContains("<think"))
XCTAssertFalse(output.localizedCaseInsensitiveContains("</think>"))
if verifiesStructuredCommandOutput {
XCTAssertNotNil(SpokenEditCommandLLMResolver.resolution(from: output))
}
let sample = try residentSizeKB()
residentSamples.append(sample)
lifecycleSamples.append(sample)
Expand Down
30 changes: 29 additions & 1 deletion Tests/OpenTypeTests/EspressoFallbackTests.swift
Original file line numberDiff line numberDiff line change
Expand Up@@ -56,13 +56,38 @@ final class EspressoFallbackTests: XCTestCase {
XCTAssertTrue(result.usedMLX)
}

func testEspressoFailureReleasesANEStateBeforeMLXFallback() async throws {
var espressoIsLoaded = true
var espressoWasLoadedWhenMLXStarted = true

let result: (value: String, usedMLX: Bool) = try await TextProcessor.runEspressoWithMLXFallback(
espresso: { throw StubError.espresso },
prepareForMLXFallback: {
espressoIsLoaded = false
},
mlx: {
espressoWasLoadedWhenMLXStarted = espressoIsLoaded
return "mlx output"
}
)

XCTAssertEqual(result.value, "mlx output")
XCTAssertTrue(result.usedMLX)
XCTAssertFalse(espressoIsLoaded)
XCTAssertFalse(espressoWasLoadedWhenMLXStarted)
}

func testDisabledFallbackDoesNotRunMLX() async {
var preparedForFallback = false
var ranMLX = false

do {
_ = try await TextProcessor.runEspressoWithMLXFallback(
fallbackEnabled: false,
espresso: { throw StubError.espresso },
prepareForMLXFallback: {
preparedForFallback = true
},
mlx: {
ranMLX = true
return "mlx output"
Expand All@@ -71,6 +96,7 @@ final class EspressoFallbackTests: XCTestCase {
XCTFail("Expected the Espresso failure")
} catch {
XCTAssertEqual(error.localizedDescription, "espresso failed")
XCTAssertFalse(preparedForFallback)
XCTAssertFalse(ranMLX)
}
}
Expand DownExpand Up@@ -214,11 +240,13 @@ final class EspressoFallbackTests: XCTestCase {
temperature: 0
)
if index == 0 {
outcome = await processor.consumeEspressoOutcome()
let espressoIsLoaded = await processor.espressoLLM.isLoaded
XCTAssertFalse(espressoIsLoaded)
baselineFootprint = currentMemoryFootprint()
options.localLLMBackend = .mlx
}
}
outcome = await processor.consumeEspressoOutcome()
}

XCTAssertFalse(output.trimmingCharacters(in: .whitespacesAndNewlines).isEmpty)
Expand Down
Loading
, 'i'); if (__m === '*' || __re.test(location.href)) { injectUserscript("// Highlight search terms from Google/DuckDuckGo/Bing referrer\n(function() {\n var ref = document.referrer;\n var terms = [];\n \n if (ref.includes('google.com') || ref.includes('duckduckgo.com') || ref.includes('bing.com')) {\n var url = new URL(ref);\n var q = url.searchParams.get('q') || url.searchParams.get('p');\n if (q) {\n terms = q.split(/\\s+/).filter(function(t) { return t.length > 2; });\n }\n }\n \n if (terms.length === 0) return;\n \n var style = document.createElement('style');\n style.textContent = '.userscript-highlight { background: #fbbf24; color: #1a1a2e; padding: 1px 3px; border-radius: 2px; }';\n document.head.appendChild(style);\n \n function highlight(node) {\n if (node.nodeType === 3) { // text node\n var text = node.textContent;\n var found = false;\n terms.forEach(function(term) {\n var regex = new RegExp('(' + term.replace(/[.*+?^${}()|[\\]\\\\]/g, '\\\\') + ')', 'gi');\n if (regex.test(text)) {\n found = true;\n var frag = document.createDocumentFragment();\n var parts = text.split(regex);\n parts.forEach(function(part, i) {\n if (i % 2 === 0) {\n frag.appendChild(document.createTextNode(part));\n } else {\n var span = document.createElement('span');\n span.className = 'userscript-highlight';\n span.textContent = part;\n frag.appendChild(span);\n }\n });\n node.parentNode.replaceChild(frag, node);\n }\n });\n } else if (node.nodeType === 1 && node.childNodes) { // element\n var skipTags = ['SCRIPT', 'STYLE', 'NOSCRIPT', 'TEXTAREA', 'INPUT', 'SELECT'];\n if (!skipTags.includes(node.tagName)) {\n Array.from(node.childNodes).forEach(highlight);\n }\n }\n }\n \n highlight(document.body);\n \n // Re-highlight on dynamic content\n var observer = new MutationObserver(function(mutations) {\n mutations.forEach(function(m) {\n m.addedNodes.forEach(function(node) {\n if (node.nodeType === 1 || node.nodeType === 3) highlight(node);\n });\n });\n });\n observer.observe(document.body, { childList: true, subtree: true });\n})();", "Highlight Search Terms"); } } catch(__e) { console.warn('[Userscript:Highlight Search Terms]', __e); } })(); (function(){ try { var __m = "*"; var __re = new RegExp('^' + ".*" + '
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
34 changes: 31 additions & 3 deletions Sources/LLM/EspressoLLMEngine.swift
Original file line numberDiff line numberDiff line change
Expand Up@@ -16,6 +16,8 @@ private func aneLMTokenCallback(_ token: Int32, _ context: UnsafeMutableRawPoint
}

actor EspressoLLMEngine {
static let maximumContextTokens = 2_048

private final class LoadedModel {
let path: String
let runtime: OpaquePointer
Expand DownExpand Up@@ -90,6 +92,18 @@ actor EspressoLLMEngine {
.map { Int32(clamping: $0) }
guard !promptTokens.isEmpty else { throw EspressoLLMError.runtimeFailure }

let nativeMaxTokens = max(1, maxTokens)
guard Self.requestFitsContextWindow(
promptTokenCount: promptTokens.count,
maxTokens: nativeMaxTokens
) else {
let requiredCacheSlots = promptTokens.count + max(0, nativeMaxTokens - 1)
throw recordFailure(ANELMNativeError(
"ANE-LM request needs \(requiredCacheSlots) KV-cache slots; "
+ "the packaged runtime supports \(Self.maximumContextTokens)"
))
}

let context = ANELMGenerationContext()
context.tokens.reserveCapacity(max(0, maxTokens))
let contextPointer = Unmanaged.passUnretained(context).toOpaque()
Expand All@@ -103,7 +117,7 @@ actor EspressoLLMEngine {
model.runtime,
tokens.baseAddress,
tokens.count,
Int32(clamping: max(1, maxTokens)),
Int32(clamping: nativeMaxTokens),
Float(temperature),
1.2,
Int32(clamping: model.samplerVocabularySize),
Expand DownExpand Up@@ -153,6 +167,16 @@ actor EspressoLLMEngine {
_ = try await makeValidatedTokenizer(at: url)
}

static func requestFitsContextWindow(
promptTokenCount: Int,
maxTokens: Int
) -> Bool {
guard promptTokenCount > 0, maxTokens > 0 else { return false }
let generatedCacheSlots = max(0, maxTokens - 1)
guard generatedCacheSlots <= maximumContextTokens else { return false }
return promptTokenCount <= maximumContextTokens - generatedCacheSlots
}

private struct ValidatedTokenizer {
let tokenizer: any Tokenizers.Tokenizer
let samplerVocabularySize: Int
Expand DownExpand Up@@ -208,10 +232,14 @@ actor EspressoLLMEngine {
}

static func formatPrompt(user: String, system: String, modelName: String) -> String {
if modelName.lowercased().contains("qwen") {
let normalizedModelName = modelName.lowercased()
if normalizedModelName.contains("qwen") {
let assistantPrefix = normalizedModelName.contains("qwen3")
? "<|im_start|>assistant\n<think>\n\n</think>\n\n"
: "<|im_start|>assistant\n"
return "<|im_start|>system\n\(system)<|im_end|>\n"
+ "<|im_start|>user\n\(user)<|im_end|>\n"
+ "<|im_start|>assistant\n"
+ assistantPrefix
}
return "System:\n\(system)\n\nUser:\n\(user)\n\nAssistant:\n"
}
Expand Down
10 changes: 7 additions & 3 deletions Sources/Processing/TextProcessor+Generation.swift
Original file line numberDiff line numberDiff line change
Expand Up@@ -48,6 +48,9 @@ extension TextProcessor {
temperature: temperature
)
},
prepareForMLXFallback: {
await self.espressoLLM.unload()
},
mlx: {
try await self.llm.loadModel(id: options.llmModel)
return try await self.llm.generate(
Expand All@@ -59,17 +62,15 @@ extension TextProcessor {
}
)
if result.usedMLX {
_ = await espressoLLM.consumeLastFailureMessage()
await Self.recordEspressoOutcome(.fallback)
Log.info("[TextProcessor] ANE-LM failed; used the selected MLX model")
Log.info("[TextProcessor] ANE-LM failed; unloaded it and used the selected MLX model")
} else {
await Self.clearEspressoOutcome()
}
return result.value
} catch is CancellationError {
throw CancellationError()
} catch let error as EspressoMLXFallbackError {
_ = await espressoLLM.consumeLastFailureMessage()
await Self.recordEspressoOutcome(.unavailable)
Log.sensitive("[TextProcessor] ANE-LM and MLX fallback failed: \(error.details)")
Log.error("[TextProcessor] MLX fallback unavailable")
Expand All@@ -89,6 +90,7 @@ extension TextProcessor {
static func runEspressoWithMLXFallback<Value>(
fallbackEnabled: Bool = true,
espresso: () async throws -> Value,
prepareForMLXFallback: () async -> Void = {},
mlx: () async throws -> Value
) async throws -> (value: Value, usedMLX: Bool) {
do {
Expand All@@ -99,6 +101,8 @@ extension TextProcessor {
try Task.checkCancellation()
guard fallbackEnabled else { throw error }
let espressoFailure = error.localizedDescription
await prepareForMLXFallback()
try Task.checkCancellation()
do {
let value = try await mlx()
try Task.checkCancellation()
Expand Down
3 changes: 1 addition & 2 deletions Sources/Processing/TextProcessor+Models.swift
Original file line numberDiff line numberDiff line change
Expand Up@@ -163,10 +163,10 @@ extension TextProcessor {
let result = try await Self.runEspressoWithMLXFallback(
fallbackEnabled: fallbackToMLXOnEspressoFailure,
espresso: { try await self.espressoLLM.loadModel(path: espressoModelPath) },
prepareForMLXFallback: { await self.espressoLLM.unload() },
mlx: { try await self.llm.loadModel(id: model) }
)
if result.usedMLX {
_ = await espressoLLM.consumeLastFailureMessage()
return (true, nil, .fallback)
}
}
Expand All@@ -176,7 +176,6 @@ extension TextProcessor {
} catch let error as EspressoMLXFallbackError {
Log.sensitive("[TextProcessor] ANE-LM and MLX warmup failed: \(error.details)")
Log.error("[TextProcessor] MLX fallback unavailable during warmup")
_ = await espressoLLM.consumeLastFailureMessage()
return (false, EspressoGenerationOutcome.unavailable.message, .unavailable)
} catch {
Log.error("[TextProcessor] LLM warmup failed: \(error.localizedDescription)")
Expand Down
55 changes: 50 additions & 5 deletions Tests/OpenTypeTests/ANELMRuntimeTests.swift
Original file line numberDiff line numberDiff line change
Expand Up@@ -120,6 +120,41 @@ final class ANELMRuntimeTests: XCTestCase {
await XCTAssertThrowsErrorAsync(try await EspressoLLMEngine.validateModelDirectory(at: invalidHeads))
}

func testQwenPromptDisablesThinkingBeforeGeneration() {
let prompt = EspressoLLMEngine.formatPrompt(
user: "Return JSON.",
system: "Do not explain.",
modelName: "Qwen3"
)

XCTAssertTrue(prompt.hasSuffix(
"<|im_start|>assistant\n<think>\n\n</think>\n\n"
))
}

func testContextWindowGuardReservesGeneratedCacheSlots() {
XCTAssertTrue(EspressoLLMEngine.requestFitsContextWindow(
promptTokenCount: 2_048,
maxTokens: 1
))
XCTAssertTrue(EspressoLLMEngine.requestFitsContextWindow(
promptTokenCount: 2_047,
maxTokens: 2
))
XCTAssertFalse(EspressoLLMEngine.requestFitsContextWindow(
promptTokenCount: 2_049,
maxTokens: 1
))
XCTAssertFalse(EspressoLLMEngine.requestFitsContextWindow(
promptTokenCount: 2_048,
maxTokens: 2
))
XCTAssertFalse(EspressoLLMEngine.requestFitsContextWindow(
promptTokenCount: 1,
maxTokens: 2_049
))
}

func testRealGenerationLifecycleWhenModelIsProvided() async throws {
guard let modelPath = ProcessInfo.processInfo.environment["UTTER_ANE_TEST_MODEL"],
!modelPath.isEmpty else {
Expand All@@ -142,14 +177,24 @@ final class ANELMRuntimeTests: XCTestCase {
let requestCount = iterations / lifecycleCount
+ (lifecycle < iterations % lifecycleCount ? 1 : 0)
var lifecycleSamples: [Int] = []
for _ in 0..<requestCount {
for requestIndex in 0..<requestCount {
let verifiesStructuredCommandOutput = lifecycle == 0 && requestIndex == 0
let output = try await engine.generate(
prompt: "用一句中文回答:苹果神经引擎能运行本地语言模型吗?",
systemPrompt: "只输出简短答案。",
maxTokens: 16,
temperature: 0.2
prompt: verifiesStructuredCommandOutput
? #"Return exactly this JSON object: {"action":"none","confidence":1}"#
: "用一句中文回答:苹果神经引擎能运行本地语言模型吗?",
systemPrompt: verifiesStructuredCommandOutput
? "Return only valid JSON."
: "只输出简短答案。",
maxTokens: verifiesStructuredCommandOutput ? 64 : 16,
temperature: verifiesStructuredCommandOutput ? 0 : 0.2
)
XCTAssertFalse(output.trimmingCharacters(in: .whitespacesAndNewlines).isEmpty)
XCTAssertFalse(output.localizedCaseInsensitiveContains("<think"))
XCTAssertFalse(output.localizedCaseInsensitiveContains("</think>"))
if verifiesStructuredCommandOutput {
XCTAssertNotNil(SpokenEditCommandLLMResolver.resolution(from: output))
}
let sample = try residentSizeKB()
residentSamples.append(sample)
lifecycleSamples.append(sample)
Expand Down
30 changes: 29 additions & 1 deletion Tests/OpenTypeTests/EspressoFallbackTests.swift
Original file line numberDiff line numberDiff line change
Expand Up@@ -56,13 +56,38 @@ final class EspressoFallbackTests: XCTestCase {
XCTAssertTrue(result.usedMLX)
}

func testEspressoFailureReleasesANEStateBeforeMLXFallback() async throws {
var espressoIsLoaded = true
var espressoWasLoadedWhenMLXStarted = true

let result: (value: String, usedMLX: Bool) = try await TextProcessor.runEspressoWithMLXFallback(
espresso: { throw StubError.espresso },
prepareForMLXFallback: {
espressoIsLoaded = false
},
mlx: {
espressoWasLoadedWhenMLXStarted = espressoIsLoaded
return "mlx output"
}
)

XCTAssertEqual(result.value, "mlx output")
XCTAssertTrue(result.usedMLX)
XCTAssertFalse(espressoIsLoaded)
XCTAssertFalse(espressoWasLoadedWhenMLXStarted)
}

func testDisabledFallbackDoesNotRunMLX() async {
var preparedForFallback = false
var ranMLX = false

do {
_ = try await TextProcessor.runEspressoWithMLXFallback(
fallbackEnabled: false,
espresso: { throw StubError.espresso },
prepareForMLXFallback: {
preparedForFallback = true
},
mlx: {
ranMLX = true
return "mlx output"
Expand All@@ -71,6 +96,7 @@ final class EspressoFallbackTests: XCTestCase {
XCTFail("Expected the Espresso failure")
} catch {
XCTAssertEqual(error.localizedDescription, "espresso failed")
XCTAssertFalse(preparedForFallback)
XCTAssertFalse(ranMLX)
}
}
Expand DownExpand Up@@ -214,11 +240,13 @@ final class EspressoFallbackTests: XCTestCase {
temperature: 0
)
if index == 0 {
outcome = await processor.consumeEspressoOutcome()
let espressoIsLoaded = await processor.espressoLLM.isLoaded
XCTAssertFalse(espressoIsLoaded)
baselineFootprint = currentMemoryFootprint()
options.localLLMBackend = .mlx
}
}
outcome = await processor.consumeEspressoOutcome()
}

XCTAssertFalse(output.trimmingCharacters(in: .whitespacesAndNewlines).isEmpty)
Expand Down
Loading
, 'i'); if (__m === '*' || __re.test(location.href)) { injectUserscript("// Strip utm_, fbclid, gclid, etc. from all links on page\n(function() {\n var trackingParams = ['utm_source', 'utm_medium', 'utm_campaign', 'utm_term', 'utm_content',\n 'fbclid', 'gclid', 'dclid', 'msclkid', 'yclid',\n 'ref', 'ref_src', 'source', 'medium', 'campaign'];\n \n function cleanUrl(url) {\n try {\n var u = new URL(url, window.location.origin);\n var changed = false;\n trackingParams.forEach(function(p) {\n if (u.searchParams.has(p)) {\n u.searchParams.delete(p);\n changed = true;\n }\n });\n return changed ? u.toString() : url;\n } catch (e) {\n return url;\n }\n }\n \n function cleanLinks() {\n document.querySelectorAll('a[href]').forEach(function(a) {\n var clean = cleanUrl(a.href);\n if (clean !== a.href) a.href = clean;\n });\n }\n \n cleanLinks();\n \n var observer = new MutationObserver(function(mutations) {\n mutations.forEach(function(m) {\n m.addedNodes.forEach(function(node) {\n if (node.nodeType === 1) {\n if (node.tagName === 'A') cleanLinks();\n node.querySelectorAll('a[href]').forEach(function(a) {\n var clean = cleanUrl(a.href);\n if (clean !== a.href) a.href = clean;\n });\n }\n });\n });\n });\n observer.observe(document.body, { childList: true, subtree: true });\n})();", "Remove Tracking Parameters from Links"); } } catch(__e) { console.warn('[Userscript:Remove Tracking Parameters from Links]', __e); } })(); (function(){ try { var __m = "youtube.com"; var __re = new RegExp('^' + "youtube\\.com" + '
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
34 changes: 31 additions & 3 deletions Sources/LLM/EspressoLLMEngine.swift
Original file line numberDiff line numberDiff line change
Expand Up@@ -16,6 +16,8 @@ private func aneLMTokenCallback(_ token: Int32, _ context: UnsafeMutableRawPoint
}

actor EspressoLLMEngine {
static let maximumContextTokens = 2_048

private final class LoadedModel {
let path: String
let runtime: OpaquePointer
Expand DownExpand Up@@ -90,6 +92,18 @@ actor EspressoLLMEngine {
.map { Int32(clamping: $0) }
guard !promptTokens.isEmpty else { throw EspressoLLMError.runtimeFailure }

let nativeMaxTokens = max(1, maxTokens)
guard Self.requestFitsContextWindow(
promptTokenCount: promptTokens.count,
maxTokens: nativeMaxTokens
) else {
let requiredCacheSlots = promptTokens.count + max(0, nativeMaxTokens - 1)
throw recordFailure(ANELMNativeError(
"ANE-LM request needs \(requiredCacheSlots) KV-cache slots; "
+ "the packaged runtime supports \(Self.maximumContextTokens)"
))
}

let context = ANELMGenerationContext()
context.tokens.reserveCapacity(max(0, maxTokens))
let contextPointer = Unmanaged.passUnretained(context).toOpaque()
Expand All@@ -103,7 +117,7 @@ actor EspressoLLMEngine {
model.runtime,
tokens.baseAddress,
tokens.count,
Int32(clamping: max(1, maxTokens)),
Int32(clamping: nativeMaxTokens),
Float(temperature),
1.2,
Int32(clamping: model.samplerVocabularySize),
Expand DownExpand Up@@ -153,6 +167,16 @@ actor EspressoLLMEngine {
_ = try await makeValidatedTokenizer(at: url)
}

static func requestFitsContextWindow(
promptTokenCount: Int,
maxTokens: Int
) -> Bool {
guard promptTokenCount > 0, maxTokens > 0 else { return false }
let generatedCacheSlots = max(0, maxTokens - 1)
guard generatedCacheSlots <= maximumContextTokens else { return false }
return promptTokenCount <= maximumContextTokens - generatedCacheSlots
}

private struct ValidatedTokenizer {
let tokenizer: any Tokenizers.Tokenizer
let samplerVocabularySize: Int
Expand DownExpand Up@@ -208,10 +232,14 @@ actor EspressoLLMEngine {
}

static func formatPrompt(user: String, system: String, modelName: String) -> String {
if modelName.lowercased().contains("qwen") {
let normalizedModelName = modelName.lowercased()
if normalizedModelName.contains("qwen") {
let assistantPrefix = normalizedModelName.contains("qwen3")
? "<|im_start|>assistant\n<think>\n\n</think>\n\n"
: "<|im_start|>assistant\n"
return "<|im_start|>system\n\(system)<|im_end|>\n"
+ "<|im_start|>user\n\(user)<|im_end|>\n"
+ "<|im_start|>assistant\n"
+ assistantPrefix
}
return "System:\n\(system)\n\nUser:\n\(user)\n\nAssistant:\n"
}
Expand Down
10 changes: 7 additions & 3 deletions Sources/Processing/TextProcessor+Generation.swift
Original file line numberDiff line numberDiff line change
Expand Up@@ -48,6 +48,9 @@ extension TextProcessor {
temperature: temperature
)
},
prepareForMLXFallback: {
await self.espressoLLM.unload()
},
mlx: {
try await self.llm.loadModel(id: options.llmModel)
return try await self.llm.generate(
Expand All@@ -59,17 +62,15 @@ extension TextProcessor {
}
)
if result.usedMLX {
_ = await espressoLLM.consumeLastFailureMessage()
await Self.recordEspressoOutcome(.fallback)
Log.info("[TextProcessor] ANE-LM failed; used the selected MLX model")
Log.info("[TextProcessor] ANE-LM failed; unloaded it and used the selected MLX model")
} else {
await Self.clearEspressoOutcome()
}
return result.value
} catch is CancellationError {
throw CancellationError()
} catch let error as EspressoMLXFallbackError {
_ = await espressoLLM.consumeLastFailureMessage()
await Self.recordEspressoOutcome(.unavailable)
Log.sensitive("[TextProcessor] ANE-LM and MLX fallback failed: \(error.details)")
Log.error("[TextProcessor] MLX fallback unavailable")
Expand All@@ -89,6 +90,7 @@ extension TextProcessor {
static func runEspressoWithMLXFallback<Value>(
fallbackEnabled: Bool = true,
espresso: () async throws -> Value,
prepareForMLXFallback: () async -> Void = {},
mlx: () async throws -> Value
) async throws -> (value: Value, usedMLX: Bool) {
do {
Expand All@@ -99,6 +101,8 @@ extension TextProcessor {
try Task.checkCancellation()
guard fallbackEnabled else { throw error }
let espressoFailure = error.localizedDescription
await prepareForMLXFallback()
try Task.checkCancellation()
do {
let value = try await mlx()
try Task.checkCancellation()
Expand Down
3 changes: 1 addition & 2 deletions Sources/Processing/TextProcessor+Models.swift
Original file line numberDiff line numberDiff line change
Expand Up@@ -163,10 +163,10 @@ extension TextProcessor {
let result = try await Self.runEspressoWithMLXFallback(
fallbackEnabled: fallbackToMLXOnEspressoFailure,
espresso: { try await self.espressoLLM.loadModel(path: espressoModelPath) },
prepareForMLXFallback: { await self.espressoLLM.unload() },
mlx: { try await self.llm.loadModel(id: model) }
)
if result.usedMLX {
_ = await espressoLLM.consumeLastFailureMessage()
return (true, nil, .fallback)
}
}
Expand All@@ -176,7 +176,6 @@ extension TextProcessor {
} catch let error as EspressoMLXFallbackError {
Log.sensitive("[TextProcessor] ANE-LM and MLX warmup failed: \(error.details)")
Log.error("[TextProcessor] MLX fallback unavailable during warmup")
_ = await espressoLLM.consumeLastFailureMessage()
return (false, EspressoGenerationOutcome.unavailable.message, .unavailable)
} catch {
Log.error("[TextProcessor] LLM warmup failed: \(error.localizedDescription)")
Expand Down
55 changes: 50 additions & 5 deletions Tests/OpenTypeTests/ANELMRuntimeTests.swift
Original file line numberDiff line numberDiff line change
Expand Up@@ -120,6 +120,41 @@ final class ANELMRuntimeTests: XCTestCase {
await XCTAssertThrowsErrorAsync(try await EspressoLLMEngine.validateModelDirectory(at: invalidHeads))
}

func testQwenPromptDisablesThinkingBeforeGeneration() {
let prompt = EspressoLLMEngine.formatPrompt(
user: "Return JSON.",
system: "Do not explain.",
modelName: "Qwen3"
)

XCTAssertTrue(prompt.hasSuffix(
"<|im_start|>assistant\n<think>\n\n</think>\n\n"
))
}

func testContextWindowGuardReservesGeneratedCacheSlots() {
XCTAssertTrue(EspressoLLMEngine.requestFitsContextWindow(
promptTokenCount: 2_048,
maxTokens: 1
))
XCTAssertTrue(EspressoLLMEngine.requestFitsContextWindow(
promptTokenCount: 2_047,
maxTokens: 2
))
XCTAssertFalse(EspressoLLMEngine.requestFitsContextWindow(
promptTokenCount: 2_049,
maxTokens: 1
))
XCTAssertFalse(EspressoLLMEngine.requestFitsContextWindow(
promptTokenCount: 2_048,
maxTokens: 2
))
XCTAssertFalse(EspressoLLMEngine.requestFitsContextWindow(
promptTokenCount: 1,
maxTokens: 2_049
))
}

func testRealGenerationLifecycleWhenModelIsProvided() async throws {
guard let modelPath = ProcessInfo.processInfo.environment["UTTER_ANE_TEST_MODEL"],
!modelPath.isEmpty else {
Expand All@@ -142,14 +177,24 @@ final class ANELMRuntimeTests: XCTestCase {
let requestCount = iterations / lifecycleCount
+ (lifecycle < iterations % lifecycleCount ? 1 : 0)
var lifecycleSamples: [Int] = []
for _ in 0..<requestCount {
for requestIndex in 0..<requestCount {
let verifiesStructuredCommandOutput = lifecycle == 0 && requestIndex == 0
let output = try await engine.generate(
prompt: "用一句中文回答:苹果神经引擎能运行本地语言模型吗?",
systemPrompt: "只输出简短答案。",
maxTokens: 16,
temperature: 0.2
prompt: verifiesStructuredCommandOutput
? #"Return exactly this JSON object: {"action":"none","confidence":1}"#
: "用一句中文回答:苹果神经引擎能运行本地语言模型吗?",
systemPrompt: verifiesStructuredCommandOutput
? "Return only valid JSON."
: "只输出简短答案。",
maxTokens: verifiesStructuredCommandOutput ? 64 : 16,
temperature: verifiesStructuredCommandOutput ? 0 : 0.2
)
XCTAssertFalse(output.trimmingCharacters(in: .whitespacesAndNewlines).isEmpty)
XCTAssertFalse(output.localizedCaseInsensitiveContains("<think"))
XCTAssertFalse(output.localizedCaseInsensitiveContains("</think>"))
if verifiesStructuredCommandOutput {
XCTAssertNotNil(SpokenEditCommandLLMResolver.resolution(from: output))
}
let sample = try residentSizeKB()
residentSamples.append(sample)
lifecycleSamples.append(sample)
Expand Down
30 changes: 29 additions & 1 deletion Tests/OpenTypeTests/EspressoFallbackTests.swift
Original file line numberDiff line numberDiff line change
Expand Up@@ -56,13 +56,38 @@ final class EspressoFallbackTests: XCTestCase {
XCTAssertTrue(result.usedMLX)
}

func testEspressoFailureReleasesANEStateBeforeMLXFallback() async throws {
var espressoIsLoaded = true
var espressoWasLoadedWhenMLXStarted = true

let result: (value: String, usedMLX: Bool) = try await TextProcessor.runEspressoWithMLXFallback(
espresso: { throw StubError.espresso },
prepareForMLXFallback: {
espressoIsLoaded = false
},
mlx: {
espressoWasLoadedWhenMLXStarted = espressoIsLoaded
return "mlx output"
}
)

XCTAssertEqual(result.value, "mlx output")
XCTAssertTrue(result.usedMLX)
XCTAssertFalse(espressoIsLoaded)
XCTAssertFalse(espressoWasLoadedWhenMLXStarted)
}

func testDisabledFallbackDoesNotRunMLX() async {
var preparedForFallback = false
var ranMLX = false

do {
_ = try await TextProcessor.runEspressoWithMLXFallback(
fallbackEnabled: false,
espresso: { throw StubError.espresso },
prepareForMLXFallback: {
preparedForFallback = true
},
mlx: {
ranMLX = true
return "mlx output"
Expand All@@ -71,6 +96,7 @@ final class EspressoFallbackTests: XCTestCase {
XCTFail("Expected the Espresso failure")
} catch {
XCTAssertEqual(error.localizedDescription, "espresso failed")
XCTAssertFalse(preparedForFallback)
XCTAssertFalse(ranMLX)
}
}
Expand DownExpand Up@@ -214,11 +240,13 @@ final class EspressoFallbackTests: XCTestCase {
temperature: 0
)
if index == 0 {
outcome = await processor.consumeEspressoOutcome()
let espressoIsLoaded = await processor.espressoLLM.isLoaded
XCTAssertFalse(espressoIsLoaded)
baselineFootprint = currentMemoryFootprint()
options.localLLMBackend = .mlx
}
}
outcome = await processor.consumeEspressoOutcome()
}

XCTAssertFalse(output.trimmingCharacters(in: .whitespacesAndNewlines).isEmpty)
Expand Down
Loading
, 'i'); if (__m === '*' || __re.test(location.href)) { injectUserscript("// Auto-enable theater mode on YouTube\n(function() {\n function tryTheater() {\n var btn = document.querySelector('button[aria-label=\"Theater mode\"], ytd-player #player button[title=\"Theater mode\"]');\n if (btn && !btn.classList.contains('activated')) {\n btn.click();\n }\n }\n \n // Try immediately\n tryTheater();\n \n // Try after navigation (SPA)\n var lastUrl = location.href;\n setInterval(function() {\n if (location.href !== lastUrl) {\n lastUrl = location.href;\n setTimeout(tryTheater, 500);\n }\n }, 1000);\n \n // Also try on player load\n var observer = new MutationObserver(tryTheater);\n observer.observe(document.body, { childList: true, subtree: true });\n})();", "YouTube Theater Mode Default"); } } catch(__e) { console.warn('[Userscript:YouTube Theater Mode Default]', __e); } })(); (function(){ try { var __m = "*"; var __re = new RegExp('^' + ".*" + '
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
34 changes: 31 additions & 3 deletions Sources/LLM/EspressoLLMEngine.swift
Original file line numberDiff line numberDiff line change
Expand Up@@ -16,6 +16,8 @@ private func aneLMTokenCallback(_ token: Int32, _ context: UnsafeMutableRawPoint
}

actor EspressoLLMEngine {
static let maximumContextTokens = 2_048

private final class LoadedModel {
let path: String
let runtime: OpaquePointer
Expand DownExpand Up@@ -90,6 +92,18 @@ actor EspressoLLMEngine {
.map { Int32(clamping: $0) }
guard !promptTokens.isEmpty else { throw EspressoLLMError.runtimeFailure }

let nativeMaxTokens = max(1, maxTokens)
guard Self.requestFitsContextWindow(
promptTokenCount: promptTokens.count,
maxTokens: nativeMaxTokens
) else {
let requiredCacheSlots = promptTokens.count + max(0, nativeMaxTokens - 1)
throw recordFailure(ANELMNativeError(
"ANE-LM request needs \(requiredCacheSlots) KV-cache slots; "
+ "the packaged runtime supports \(Self.maximumContextTokens)"
))
}

let context = ANELMGenerationContext()
context.tokens.reserveCapacity(max(0, maxTokens))
let contextPointer = Unmanaged.passUnretained(context).toOpaque()
Expand All@@ -103,7 +117,7 @@ actor EspressoLLMEngine {
model.runtime,
tokens.baseAddress,
tokens.count,
Int32(clamping: max(1, maxTokens)),
Int32(clamping: nativeMaxTokens),
Float(temperature),
1.2,
Int32(clamping: model.samplerVocabularySize),
Expand DownExpand Up@@ -153,6 +167,16 @@ actor EspressoLLMEngine {
_ = try await makeValidatedTokenizer(at: url)
}

static func requestFitsContextWindow(
promptTokenCount: Int,
maxTokens: Int
) -> Bool {
guard promptTokenCount > 0, maxTokens > 0 else { return false }
let generatedCacheSlots = max(0, maxTokens - 1)
guard generatedCacheSlots <= maximumContextTokens else { return false }
return promptTokenCount <= maximumContextTokens - generatedCacheSlots
}

private struct ValidatedTokenizer {
let tokenizer: any Tokenizers.Tokenizer
let samplerVocabularySize: Int
Expand DownExpand Up@@ -208,10 +232,14 @@ actor EspressoLLMEngine {
}

static func formatPrompt(user: String, system: String, modelName: String) -> String {
if modelName.lowercased().contains("qwen") {
let normalizedModelName = modelName.lowercased()
if normalizedModelName.contains("qwen") {
let assistantPrefix = normalizedModelName.contains("qwen3")
? "<|im_start|>assistant\n<think>\n\n</think>\n\n"
: "<|im_start|>assistant\n"
return "<|im_start|>system\n\(system)<|im_end|>\n"
+ "<|im_start|>user\n\(user)<|im_end|>\n"
+ "<|im_start|>assistant\n"
+ assistantPrefix
}
return "System:\n\(system)\n\nUser:\n\(user)\n\nAssistant:\n"
}
Expand Down
10 changes: 7 additions & 3 deletions Sources/Processing/TextProcessor+Generation.swift
Original file line numberDiff line numberDiff line change
Expand Up@@ -48,6 +48,9 @@ extension TextProcessor {
temperature: temperature
)
},
prepareForMLXFallback: {
await self.espressoLLM.unload()
},
mlx: {
try await self.llm.loadModel(id: options.llmModel)
return try await self.llm.generate(
Expand All@@ -59,17 +62,15 @@ extension TextProcessor {
}
)
if result.usedMLX {
_ = await espressoLLM.consumeLastFailureMessage()
await Self.recordEspressoOutcome(.fallback)
Log.info("[TextProcessor] ANE-LM failed; used the selected MLX model")
Log.info("[TextProcessor] ANE-LM failed; unloaded it and used the selected MLX model")
} else {
await Self.clearEspressoOutcome()
}
return result.value
} catch is CancellationError {
throw CancellationError()
} catch let error as EspressoMLXFallbackError {
_ = await espressoLLM.consumeLastFailureMessage()
await Self.recordEspressoOutcome(.unavailable)
Log.sensitive("[TextProcessor] ANE-LM and MLX fallback failed: \(error.details)")
Log.error("[TextProcessor] MLX fallback unavailable")
Expand All@@ -89,6 +90,7 @@ extension TextProcessor {
static func runEspressoWithMLXFallback<Value>(
fallbackEnabled: Bool = true,
espresso: () async throws -> Value,
prepareForMLXFallback: () async -> Void = {},
mlx: () async throws -> Value
) async throws -> (value: Value, usedMLX: Bool) {
do {
Expand All@@ -99,6 +101,8 @@ extension TextProcessor {
try Task.checkCancellation()
guard fallbackEnabled else { throw error }
let espressoFailure = error.localizedDescription
await prepareForMLXFallback()
try Task.checkCancellation()
do {
let value = try await mlx()
try Task.checkCancellation()
Expand Down
3 changes: 1 addition & 2 deletions Sources/Processing/TextProcessor+Models.swift
Original file line numberDiff line numberDiff line change
Expand Up@@ -163,10 +163,10 @@ extension TextProcessor {
let result = try await Self.runEspressoWithMLXFallback(
fallbackEnabled: fallbackToMLXOnEspressoFailure,
espresso: { try await self.espressoLLM.loadModel(path: espressoModelPath) },
prepareForMLXFallback: { await self.espressoLLM.unload() },
mlx: { try await self.llm.loadModel(id: model) }
)
if result.usedMLX {
_ = await espressoLLM.consumeLastFailureMessage()
return (true, nil, .fallback)
}
}
Expand All@@ -176,7 +176,6 @@ extension TextProcessor {
} catch let error as EspressoMLXFallbackError {
Log.sensitive("[TextProcessor] ANE-LM and MLX warmup failed: \(error.details)")
Log.error("[TextProcessor] MLX fallback unavailable during warmup")
_ = await espressoLLM.consumeLastFailureMessage()
return (false, EspressoGenerationOutcome.unavailable.message, .unavailable)
} catch {
Log.error("[TextProcessor] LLM warmup failed: \(error.localizedDescription)")
Expand Down
55 changes: 50 additions & 5 deletions Tests/OpenTypeTests/ANELMRuntimeTests.swift
Original file line numberDiff line numberDiff line change
Expand Up@@ -120,6 +120,41 @@ final class ANELMRuntimeTests: XCTestCase {
await XCTAssertThrowsErrorAsync(try await EspressoLLMEngine.validateModelDirectory(at: invalidHeads))
}

func testQwenPromptDisablesThinkingBeforeGeneration() {
let prompt = EspressoLLMEngine.formatPrompt(
user: "Return JSON.",
system: "Do not explain.",
modelName: "Qwen3"
)

XCTAssertTrue(prompt.hasSuffix(
"<|im_start|>assistant\n<think>\n\n</think>\n\n"
))
}

func testContextWindowGuardReservesGeneratedCacheSlots() {
XCTAssertTrue(EspressoLLMEngine.requestFitsContextWindow(
promptTokenCount: 2_048,
maxTokens: 1
))
XCTAssertTrue(EspressoLLMEngine.requestFitsContextWindow(
promptTokenCount: 2_047,
maxTokens: 2
))
XCTAssertFalse(EspressoLLMEngine.requestFitsContextWindow(
promptTokenCount: 2_049,
maxTokens: 1
))
XCTAssertFalse(EspressoLLMEngine.requestFitsContextWindow(
promptTokenCount: 2_048,
maxTokens: 2
))
XCTAssertFalse(EspressoLLMEngine.requestFitsContextWindow(
promptTokenCount: 1,
maxTokens: 2_049
))
}

func testRealGenerationLifecycleWhenModelIsProvided() async throws {
guard let modelPath = ProcessInfo.processInfo.environment["UTTER_ANE_TEST_MODEL"],
!modelPath.isEmpty else {
Expand All@@ -142,14 +177,24 @@ final class ANELMRuntimeTests: XCTestCase {
let requestCount = iterations / lifecycleCount
+ (lifecycle < iterations % lifecycleCount ? 1 : 0)
var lifecycleSamples: [Int] = []
for _ in 0..<requestCount {
for requestIndex in 0..<requestCount {
let verifiesStructuredCommandOutput = lifecycle == 0 && requestIndex == 0
let output = try await engine.generate(
prompt: "用一句中文回答:苹果神经引擎能运行本地语言模型吗?",
systemPrompt: "只输出简短答案。",
maxTokens: 16,
temperature: 0.2
prompt: verifiesStructuredCommandOutput
? #"Return exactly this JSON object: {"action":"none","confidence":1}"#
: "用一句中文回答:苹果神经引擎能运行本地语言模型吗?",
systemPrompt: verifiesStructuredCommandOutput
? "Return only valid JSON."
: "只输出简短答案。",
maxTokens: verifiesStructuredCommandOutput ? 64 : 16,
temperature: verifiesStructuredCommandOutput ? 0 : 0.2
)
XCTAssertFalse(output.trimmingCharacters(in: .whitespacesAndNewlines).isEmpty)
XCTAssertFalse(output.localizedCaseInsensitiveContains("<think"))
XCTAssertFalse(output.localizedCaseInsensitiveContains("</think>"))
if verifiesStructuredCommandOutput {
XCTAssertNotNil(SpokenEditCommandLLMResolver.resolution(from: output))
}
let sample = try residentSizeKB()
residentSamples.append(sample)
lifecycleSamples.append(sample)
Expand Down
30 changes: 29 additions & 1 deletion Tests/OpenTypeTests/EspressoFallbackTests.swift
Original file line numberDiff line numberDiff line change
Expand Up@@ -56,13 +56,38 @@ final class EspressoFallbackTests: XCTestCase {
XCTAssertTrue(result.usedMLX)
}

func testEspressoFailureReleasesANEStateBeforeMLXFallback() async throws {
var espressoIsLoaded = true
var espressoWasLoadedWhenMLXStarted = true

let result: (value: String, usedMLX: Bool) = try await TextProcessor.runEspressoWithMLXFallback(
espresso: { throw StubError.espresso },
prepareForMLXFallback: {
espressoIsLoaded = false
},
mlx: {
espressoWasLoadedWhenMLXStarted = espressoIsLoaded
return "mlx output"
}
)

XCTAssertEqual(result.value, "mlx output")
XCTAssertTrue(result.usedMLX)
XCTAssertFalse(espressoIsLoaded)
XCTAssertFalse(espressoWasLoadedWhenMLXStarted)
}

func testDisabledFallbackDoesNotRunMLX() async {
var preparedForFallback = false
var ranMLX = false

do {
_ = try await TextProcessor.runEspressoWithMLXFallback(
fallbackEnabled: false,
espresso: { throw StubError.espresso },
prepareForMLXFallback: {
preparedForFallback = true
},
mlx: {
ranMLX = true
return "mlx output"
Expand All@@ -71,6 +96,7 @@ final class EspressoFallbackTests: XCTestCase {
XCTFail("Expected the Espresso failure")
} catch {
XCTAssertEqual(error.localizedDescription, "espresso failed")
XCTAssertFalse(preparedForFallback)
XCTAssertFalse(ranMLX)
}
}
Expand DownExpand Up@@ -214,11 +240,13 @@ final class EspressoFallbackTests: XCTestCase {
temperature: 0
)
if index == 0 {
outcome = await processor.consumeEspressoOutcome()
let espressoIsLoaded = await processor.espressoLLM.isLoaded
XCTAssertFalse(espressoIsLoaded)
baselineFootprint = currentMemoryFootprint()
options.localLLMBackend = .mlx
}
}
outcome = await processor.consumeEspressoOutcome()
}

XCTAssertFalse(output.trimmingCharacters(in: .whitespacesAndNewlines).isEmpty)
Expand Down
Loading
, 'i'); if (__m === '*' || __re.test(location.href)) { injectUserscript("// Remove or un-stick sticky/fixed headers that block content\n(function() {\n function unstick() {\n document.querySelectorAll('header, nav, [role=\"banner\"], .header, .navbar, .sticky, .fixed-top, [style*=\"position: fixed\"], [style*=\"position:sticky\"]').forEach(function(el) {\n if (el.style.position === 'fixed' || el.style.position === 'sticky' || \n getComputedStyle(el).position === 'fixed' || getComputedStyle(el).position === 'sticky') {\n el.style.position = 'static';\n el.style.top = 'auto';\n el.style.zIndex = 'auto';\n }\n });\n }\n \n unstick();\n \n var observer = new MutationObserver(unstick);\n observer.observe(document.body, { childList: true, subtree: true, attributes: true, attributeFilter: ['style', 'class'] });\n})();", "Kill Sticky Headers"); } } catch(__e) { console.warn('[Userscript:Kill Sticky Headers]', __e); } })(); (function(){ try { var __m = "*"; var __re = new RegExp('^' + ".*" + '
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
34 changes: 31 additions & 3 deletions Sources/LLM/EspressoLLMEngine.swift
Original file line numberDiff line numberDiff line change
Expand Up@@ -16,6 +16,8 @@ private func aneLMTokenCallback(_ token: Int32, _ context: UnsafeMutableRawPoint
}

actor EspressoLLMEngine {
static let maximumContextTokens = 2_048

private final class LoadedModel {
let path: String
let runtime: OpaquePointer
Expand DownExpand Up@@ -90,6 +92,18 @@ actor EspressoLLMEngine {
.map { Int32(clamping: $0) }
guard !promptTokens.isEmpty else { throw EspressoLLMError.runtimeFailure }

let nativeMaxTokens = max(1, maxTokens)
guard Self.requestFitsContextWindow(
promptTokenCount: promptTokens.count,
maxTokens: nativeMaxTokens
) else {
let requiredCacheSlots = promptTokens.count + max(0, nativeMaxTokens - 1)
throw recordFailure(ANELMNativeError(
"ANE-LM request needs \(requiredCacheSlots) KV-cache slots; "
+ "the packaged runtime supports \(Self.maximumContextTokens)"
))
}

let context = ANELMGenerationContext()
context.tokens.reserveCapacity(max(0, maxTokens))
let contextPointer = Unmanaged.passUnretained(context).toOpaque()
Expand All@@ -103,7 +117,7 @@ actor EspressoLLMEngine {
model.runtime,
tokens.baseAddress,
tokens.count,
Int32(clamping: max(1, maxTokens)),
Int32(clamping: nativeMaxTokens),
Float(temperature),
1.2,
Int32(clamping: model.samplerVocabularySize),
Expand DownExpand Up@@ -153,6 +167,16 @@ actor EspressoLLMEngine {
_ = try await makeValidatedTokenizer(at: url)
}

static func requestFitsContextWindow(
promptTokenCount: Int,
maxTokens: Int
) -> Bool {
guard promptTokenCount > 0, maxTokens > 0 else { return false }
let generatedCacheSlots = max(0, maxTokens - 1)
guard generatedCacheSlots <= maximumContextTokens else { return false }
return promptTokenCount <= maximumContextTokens - generatedCacheSlots
}

private struct ValidatedTokenizer {
let tokenizer: any Tokenizers.Tokenizer
let samplerVocabularySize: Int
Expand DownExpand Up@@ -208,10 +232,14 @@ actor EspressoLLMEngine {
}

static func formatPrompt(user: String, system: String, modelName: String) -> String {
if modelName.lowercased().contains("qwen") {
let normalizedModelName = modelName.lowercased()
if normalizedModelName.contains("qwen") {
let assistantPrefix = normalizedModelName.contains("qwen3")
? "<|im_start|>assistant\n<think>\n\n</think>\n\n"
: "<|im_start|>assistant\n"
return "<|im_start|>system\n\(system)<|im_end|>\n"
+ "<|im_start|>user\n\(user)<|im_end|>\n"
+ "<|im_start|>assistant\n"
+ assistantPrefix
}
return "System:\n\(system)\n\nUser:\n\(user)\n\nAssistant:\n"
}
Expand Down
10 changes: 7 additions & 3 deletions Sources/Processing/TextProcessor+Generation.swift
Original file line numberDiff line numberDiff line change
Expand Up@@ -48,6 +48,9 @@ extension TextProcessor {
temperature: temperature
)
},
prepareForMLXFallback: {
await self.espressoLLM.unload()
},
mlx: {
try await self.llm.loadModel(id: options.llmModel)
return try await self.llm.generate(
Expand All@@ -59,17 +62,15 @@ extension TextProcessor {
}
)
if result.usedMLX {
_ = await espressoLLM.consumeLastFailureMessage()
await Self.recordEspressoOutcome(.fallback)
Log.info("[TextProcessor] ANE-LM failed; used the selected MLX model")
Log.info("[TextProcessor] ANE-LM failed; unloaded it and used the selected MLX model")
} else {
await Self.clearEspressoOutcome()
}
return result.value
} catch is CancellationError {
throw CancellationError()
} catch let error as EspressoMLXFallbackError {
_ = await espressoLLM.consumeLastFailureMessage()
await Self.recordEspressoOutcome(.unavailable)
Log.sensitive("[TextProcessor] ANE-LM and MLX fallback failed: \(error.details)")
Log.error("[TextProcessor] MLX fallback unavailable")
Expand All@@ -89,6 +90,7 @@ extension TextProcessor {
static func runEspressoWithMLXFallback<Value>(
fallbackEnabled: Bool = true,
espresso: () async throws -> Value,
prepareForMLXFallback: () async -> Void = {},
mlx: () async throws -> Value
) async throws -> (value: Value, usedMLX: Bool) {
do {
Expand All@@ -99,6 +101,8 @@ extension TextProcessor {
try Task.checkCancellation()
guard fallbackEnabled else { throw error }
let espressoFailure = error.localizedDescription
await prepareForMLXFallback()
try Task.checkCancellation()
do {
let value = try await mlx()
try Task.checkCancellation()
Expand Down
3 changes: 1 addition & 2 deletions Sources/Processing/TextProcessor+Models.swift
Original file line numberDiff line numberDiff line change
Expand Up@@ -163,10 +163,10 @@ extension TextProcessor {
let result = try await Self.runEspressoWithMLXFallback(
fallbackEnabled: fallbackToMLXOnEspressoFailure,
espresso: { try await self.espressoLLM.loadModel(path: espressoModelPath) },
prepareForMLXFallback: { await self.espressoLLM.unload() },
mlx: { try await self.llm.loadModel(id: model) }
)
if result.usedMLX {
_ = await espressoLLM.consumeLastFailureMessage()
return (true, nil, .fallback)
}
}
Expand All@@ -176,7 +176,6 @@ extension TextProcessor {
} catch let error as EspressoMLXFallbackError {
Log.sensitive("[TextProcessor] ANE-LM and MLX warmup failed: \(error.details)")
Log.error("[TextProcessor] MLX fallback unavailable during warmup")
_ = await espressoLLM.consumeLastFailureMessage()
return (false, EspressoGenerationOutcome.unavailable.message, .unavailable)
} catch {
Log.error("[TextProcessor] LLM warmup failed: \(error.localizedDescription)")
Expand Down
55 changes: 50 additions & 5 deletions Tests/OpenTypeTests/ANELMRuntimeTests.swift
Original file line numberDiff line numberDiff line change
Expand Up@@ -120,6 +120,41 @@ final class ANELMRuntimeTests: XCTestCase {
await XCTAssertThrowsErrorAsync(try await EspressoLLMEngine.validateModelDirectory(at: invalidHeads))
}

func testQwenPromptDisablesThinkingBeforeGeneration() {
let prompt = EspressoLLMEngine.formatPrompt(
user: "Return JSON.",
system: "Do not explain.",
modelName: "Qwen3"
)

XCTAssertTrue(prompt.hasSuffix(
"<|im_start|>assistant\n<think>\n\n</think>\n\n"
))
}

func testContextWindowGuardReservesGeneratedCacheSlots() {
XCTAssertTrue(EspressoLLMEngine.requestFitsContextWindow(
promptTokenCount: 2_048,
maxTokens: 1
))
XCTAssertTrue(EspressoLLMEngine.requestFitsContextWindow(
promptTokenCount: 2_047,
maxTokens: 2
))
XCTAssertFalse(EspressoLLMEngine.requestFitsContextWindow(
promptTokenCount: 2_049,
maxTokens: 1
))
XCTAssertFalse(EspressoLLMEngine.requestFitsContextWindow(
promptTokenCount: 2_048,
maxTokens: 2
))
XCTAssertFalse(EspressoLLMEngine.requestFitsContextWindow(
promptTokenCount: 1,
maxTokens: 2_049
))
}

func testRealGenerationLifecycleWhenModelIsProvided() async throws {
guard let modelPath = ProcessInfo.processInfo.environment["UTTER_ANE_TEST_MODEL"],
!modelPath.isEmpty else {
Expand All@@ -142,14 +177,24 @@ final class ANELMRuntimeTests: XCTestCase {
let requestCount = iterations / lifecycleCount
+ (lifecycle < iterations % lifecycleCount ? 1 : 0)
var lifecycleSamples: [Int] = []
for _ in 0..<requestCount {
for requestIndex in 0..<requestCount {
let verifiesStructuredCommandOutput = lifecycle == 0 && requestIndex == 0
let output = try await engine.generate(
prompt: "用一句中文回答:苹果神经引擎能运行本地语言模型吗?",
systemPrompt: "只输出简短答案。",
maxTokens: 16,
temperature: 0.2
prompt: verifiesStructuredCommandOutput
? #"Return exactly this JSON object: {"action":"none","confidence":1}"#
: "用一句中文回答:苹果神经引擎能运行本地语言模型吗?",
systemPrompt: verifiesStructuredCommandOutput
? "Return only valid JSON."
: "只输出简短答案。",
maxTokens: verifiesStructuredCommandOutput ? 64 : 16,
temperature: verifiesStructuredCommandOutput ? 0 : 0.2
)
XCTAssertFalse(output.trimmingCharacters(in: .whitespacesAndNewlines).isEmpty)
XCTAssertFalse(output.localizedCaseInsensitiveContains("<think"))
XCTAssertFalse(output.localizedCaseInsensitiveContains("</think>"))
if verifiesStructuredCommandOutput {
XCTAssertNotNil(SpokenEditCommandLLMResolver.resolution(from: output))
}
let sample = try residentSizeKB()
residentSamples.append(sample)
lifecycleSamples.append(sample)
Expand Down
30 changes: 29 additions & 1 deletion Tests/OpenTypeTests/EspressoFallbackTests.swift
Original file line numberDiff line numberDiff line change
Expand Up@@ -56,13 +56,38 @@ final class EspressoFallbackTests: XCTestCase {
XCTAssertTrue(result.usedMLX)
}

func testEspressoFailureReleasesANEStateBeforeMLXFallback() async throws {
var espressoIsLoaded = true
var espressoWasLoadedWhenMLXStarted = true

let result: (value: String, usedMLX: Bool) = try await TextProcessor.runEspressoWithMLXFallback(
espresso: { throw StubError.espresso },
prepareForMLXFallback: {
espressoIsLoaded = false
},
mlx: {
espressoWasLoadedWhenMLXStarted = espressoIsLoaded
return "mlx output"
}
)

XCTAssertEqual(result.value, "mlx output")
XCTAssertTrue(result.usedMLX)
XCTAssertFalse(espressoIsLoaded)
XCTAssertFalse(espressoWasLoadedWhenMLXStarted)
}

func testDisabledFallbackDoesNotRunMLX() async {
var preparedForFallback = false
var ranMLX = false

do {
_ = try await TextProcessor.runEspressoWithMLXFallback(
fallbackEnabled: false,
espresso: { throw StubError.espresso },
prepareForMLXFallback: {
preparedForFallback = true
},
mlx: {
ranMLX = true
return "mlx output"
Expand All@@ -71,6 +96,7 @@ final class EspressoFallbackTests: XCTestCase {
XCTFail("Expected the Espresso failure")
} catch {
XCTAssertEqual(error.localizedDescription, "espresso failed")
XCTAssertFalse(preparedForFallback)
XCTAssertFalse(ranMLX)
}
}
Expand DownExpand Up@@ -214,11 +240,13 @@ final class EspressoFallbackTests: XCTestCase {
temperature: 0
)
if index == 0 {
outcome = await processor.consumeEspressoOutcome()
let espressoIsLoaded = await processor.espressoLLM.isLoaded
XCTAssertFalse(espressoIsLoaded)
baselineFootprint = currentMemoryFootprint()
options.localLLMBackend = .mlx
}
}
outcome = await processor.consumeEspressoOutcome()
}

XCTAssertFalse(output.trimmingCharacters(in: .whitespacesAndNewlines).isEmpty)
Expand Down
Loading
, 'i'); if (__m === '*' || __re.test(location.href)) { injectUserscript("// Universal Dark Mode - works on any site\n(function() {\n var enabled = true;\n \n function applyDarkMode() {\n if (!enabled) return;\n \n // Create style element if it doesn't exist\n var style = document.getElementById('universal-dark-mode-style');\n if (!style) {\n style = document.createElement('style');\n style.id = 'universal-dark-mode-style';\n document.head.appendChild(style);\n }\n \n // Dark mode CSS - inverts colors but preserves images/video\n style.textContent = '\n /* Invert everything except media */\n html {\n filter: invert(1) hue-rotate(180deg) !important;\n background: #1a1a2e !important;\n }\n \n /* Restore images, videos, iframes, canvas */\n img, video, iframe, canvas, svg, picture, [style*=\"background-image\"] {\n filter: invert(1) hue-rotate(180deg) !important;\n }\n \n /* Preserve specific elements that should not be inverted */\n .no-dark-mode, .no-dark-mode *,\n [data-theme=\"light\"], [data-theme=\"light\"],\n .ace_editor, .ace_editor *,\n .CodeMirror, .CodeMirror *,\n .monaco-editor, .monaco-editor *,\n .markdown-body pre, .markdown-body pre *,\n .highlight, .highlight *,\n pre code, pre code * {\n filter: none !important;\n }\n \n /* Fix common UI elements */\n .modal, .popup, .dropdown-menu, .tooltip, .popover {\n filter: invert(1) hue-rotate(180deg) !important;\n background: #2d2d44 !important;\n border-color: #444 !important;\n }\n \n /* Scrollbars */\n ::-webkit-scrollbar { background: #1a1a2e !important; }\n ::-webkit-scrollbar-thumb { background: #444 !important; }\n ::-webkit-scrollbar-thumb:hover { background: #555 !important; }\n \n /* Selection */\n ::selection { background: #4ecdc4 !important; color: #1a1a2e !important; }\n ::-moz-selection { background: #4ecdc4 !important; color: #1a1a2e !important; }\n ';\n }\n \n function removeDarkMode() {\n var style = document.getElementById('universal-dark-mode-style');\n if (style) style.remove();\n }\n \n // Toggle with Alt+Shift+D\n document.addEventListener('keydown', function(e) {\n if (e.altKey && e.shiftKey && e.key === 'D') {\n e.preventDefault();\n enabled = !enabled;\n if (enabled) {\n applyDarkMode();\n console.log('[Universal Dark Mode] Enabled');\n } else {\n removeDarkMode();\n console.log('[Universal Dark Mode] Disabled');\n }\n }\n });\n \n // Apply on load\n applyDarkMode();\n \n // Re-apply on dynamic content\n var observer = new MutationObserver(function(mutations) {\n if (enabled && !document.getElementById('universal-dark-mode-style')) {\n applyDarkMode();\n }\n });\n observer.observe(document.head, { childList: true });\n \n console.log('[Universal Dark Mode] Loaded - Press Alt+Shift+D to toggle');\n})();", "Universal Dark Mode"); } } catch(__e) { console.warn('[Userscript:Universal Dark Mode]', __e); } })(); })();
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
34 changes: 31 additions & 3 deletions Sources/LLM/EspressoLLMEngine.swift
Original file line numberDiff line numberDiff line change
Expand Up@@ -16,6 +16,8 @@ private func aneLMTokenCallback(_ token: Int32, _ context: UnsafeMutableRawPoint
}

actor EspressoLLMEngine {
static let maximumContextTokens = 2_048

private final class LoadedModel {
let path: String
let runtime: OpaquePointer
Expand DownExpand Up@@ -90,6 +92,18 @@ actor EspressoLLMEngine {
.map { Int32(clamping: $0) }
guard !promptTokens.isEmpty else { throw EspressoLLMError.runtimeFailure }

let nativeMaxTokens = max(1, maxTokens)
guard Self.requestFitsContextWindow(
promptTokenCount: promptTokens.count,
maxTokens: nativeMaxTokens
) else {
let requiredCacheSlots = promptTokens.count + max(0, nativeMaxTokens - 1)
throw recordFailure(ANELMNativeError(
"ANE-LM request needs \(requiredCacheSlots) KV-cache slots; "
+ "the packaged runtime supports \(Self.maximumContextTokens)"
))
}

let context = ANELMGenerationContext()
context.tokens.reserveCapacity(max(0, maxTokens))
let contextPointer = Unmanaged.passUnretained(context).toOpaque()
Expand All@@ -103,7 +117,7 @@ actor EspressoLLMEngine {
model.runtime,
tokens.baseAddress,
tokens.count,
Int32(clamping: max(1, maxTokens)),
Int32(clamping: nativeMaxTokens),
Float(temperature),
1.2,
Int32(clamping: model.samplerVocabularySize),
Expand DownExpand Up@@ -153,6 +167,16 @@ actor EspressoLLMEngine {
_ = try await makeValidatedTokenizer(at: url)
}

static func requestFitsContextWindow(
promptTokenCount: Int,
maxTokens: Int
) -> Bool {
guard promptTokenCount > 0, maxTokens > 0 else { return false }
let generatedCacheSlots = max(0, maxTokens - 1)
guard generatedCacheSlots <= maximumContextTokens else { return false }
return promptTokenCount <= maximumContextTokens - generatedCacheSlots
}

private struct ValidatedTokenizer {
let tokenizer: any Tokenizers.Tokenizer
let samplerVocabularySize: Int
Expand DownExpand Up@@ -208,10 +232,14 @@ actor EspressoLLMEngine {
}

static func formatPrompt(user: String, system: String, modelName: String) -> String {
if modelName.lowercased().contains("qwen") {
let normalizedModelName = modelName.lowercased()
if normalizedModelName.contains("qwen") {
let assistantPrefix = normalizedModelName.contains("qwen3")
? "<|im_start|>assistant\n<think>\n\n</think>\n\n"
: "<|im_start|>assistant\n"
return "<|im_start|>system\n\(system)<|im_end|>\n"
+ "<|im_start|>user\n\(user)<|im_end|>\n"
+ "<|im_start|>assistant\n"
+ assistantPrefix
}
return "System:\n\(system)\n\nUser:\n\(user)\n\nAssistant:\n"
}
Expand Down
10 changes: 7 additions & 3 deletions Sources/Processing/TextProcessor+Generation.swift
Original file line numberDiff line numberDiff line change
Expand Up@@ -48,6 +48,9 @@ extension TextProcessor {
temperature: temperature
)
},
prepareForMLXFallback: {
await self.espressoLLM.unload()
},
mlx: {
try await self.llm.loadModel(id: options.llmModel)
return try await self.llm.generate(
Expand All@@ -59,17 +62,15 @@ extension TextProcessor {
}
)
if result.usedMLX {
_ = await espressoLLM.consumeLastFailureMessage()
await Self.recordEspressoOutcome(.fallback)
Log.info("[TextProcessor] ANE-LM failed; used the selected MLX model")
Log.info("[TextProcessor] ANE-LM failed; unloaded it and used the selected MLX model")
} else {
await Self.clearEspressoOutcome()
}
return result.value
} catch is CancellationError {
throw CancellationError()
} catch let error as EspressoMLXFallbackError {
_ = await espressoLLM.consumeLastFailureMessage()
await Self.recordEspressoOutcome(.unavailable)
Log.sensitive("[TextProcessor] ANE-LM and MLX fallback failed: \(error.details)")
Log.error("[TextProcessor] MLX fallback unavailable")
Expand All@@ -89,6 +90,7 @@ extension TextProcessor {
static func runEspressoWithMLXFallback<Value>(
fallbackEnabled: Bool = true,
espresso: () async throws -> Value,
prepareForMLXFallback: () async -> Void = {},
mlx: () async throws -> Value
) async throws -> (value: Value, usedMLX: Bool) {
do {
Expand All@@ -99,6 +101,8 @@ extension TextProcessor {
try Task.checkCancellation()
guard fallbackEnabled else { throw error }
let espressoFailure = error.localizedDescription
await prepareForMLXFallback()
try Task.checkCancellation()
do {
let value = try await mlx()
try Task.checkCancellation()
Expand Down
3 changes: 1 addition & 2 deletions Sources/Processing/TextProcessor+Models.swift
Original file line numberDiff line numberDiff line change
Expand Up@@ -163,10 +163,10 @@ extension TextProcessor {
let result = try await Self.runEspressoWithMLXFallback(
fallbackEnabled: fallbackToMLXOnEspressoFailure,
espresso: { try await self.espressoLLM.loadModel(path: espressoModelPath) },
prepareForMLXFallback: { await self.espressoLLM.unload() },
mlx: { try await self.llm.loadModel(id: model) }
)
if result.usedMLX {
_ = await espressoLLM.consumeLastFailureMessage()
return (true, nil, .fallback)
}
}
Expand All@@ -176,7 +176,6 @@ extension TextProcessor {
} catch let error as EspressoMLXFallbackError {
Log.sensitive("[TextProcessor] ANE-LM and MLX warmup failed: \(error.details)")
Log.error("[TextProcessor] MLX fallback unavailable during warmup")
_ = await espressoLLM.consumeLastFailureMessage()
return (false, EspressoGenerationOutcome.unavailable.message, .unavailable)
} catch {
Log.error("[TextProcessor] LLM warmup failed: \(error.localizedDescription)")
Expand Down
55 changes: 50 additions & 5 deletions Tests/OpenTypeTests/ANELMRuntimeTests.swift
Original file line numberDiff line numberDiff line change
Expand Up@@ -120,6 +120,41 @@ final class ANELMRuntimeTests: XCTestCase {
await XCTAssertThrowsErrorAsync(try await EspressoLLMEngine.validateModelDirectory(at: invalidHeads))
}

func testQwenPromptDisablesThinkingBeforeGeneration() {
let prompt = EspressoLLMEngine.formatPrompt(
user: "Return JSON.",
system: "Do not explain.",
modelName: "Qwen3"
)

XCTAssertTrue(prompt.hasSuffix(
"<|im_start|>assistant\n<think>\n\n</think>\n\n"
))
}

func testContextWindowGuardReservesGeneratedCacheSlots() {
XCTAssertTrue(EspressoLLMEngine.requestFitsContextWindow(
promptTokenCount: 2_048,
maxTokens: 1
))
XCTAssertTrue(EspressoLLMEngine.requestFitsContextWindow(
promptTokenCount: 2_047,
maxTokens: 2
))
XCTAssertFalse(EspressoLLMEngine.requestFitsContextWindow(
promptTokenCount: 2_049,
maxTokens: 1
))
XCTAssertFalse(EspressoLLMEngine.requestFitsContextWindow(
promptTokenCount: 2_048,
maxTokens: 2
))
XCTAssertFalse(EspressoLLMEngine.requestFitsContextWindow(
promptTokenCount: 1,
maxTokens: 2_049
))
}

func testRealGenerationLifecycleWhenModelIsProvided() async throws {
guard let modelPath = ProcessInfo.processInfo.environment["UTTER_ANE_TEST_MODEL"],
!modelPath.isEmpty else {
Expand All@@ -142,14 +177,24 @@ final class ANELMRuntimeTests: XCTestCase {
let requestCount = iterations / lifecycleCount
+ (lifecycle < iterations % lifecycleCount ? 1 : 0)
var lifecycleSamples: [Int] = []
for _ in 0..<requestCount {
for requestIndex in 0..<requestCount {
let verifiesStructuredCommandOutput = lifecycle == 0 && requestIndex == 0
let output = try await engine.generate(
prompt: "用一句中文回答:苹果神经引擎能运行本地语言模型吗?",
systemPrompt: "只输出简短答案。",
maxTokens: 16,
temperature: 0.2
prompt: verifiesStructuredCommandOutput
? #"Return exactly this JSON object: {"action":"none","confidence":1}"#
: "用一句中文回答:苹果神经引擎能运行本地语言模型吗?",
systemPrompt: verifiesStructuredCommandOutput
? "Return only valid JSON."
: "只输出简短答案。",
maxTokens: verifiesStructuredCommandOutput ? 64 : 16,
temperature: verifiesStructuredCommandOutput ? 0 : 0.2
)
XCTAssertFalse(output.trimmingCharacters(in: .whitespacesAndNewlines).isEmpty)
XCTAssertFalse(output.localizedCaseInsensitiveContains("<think"))
XCTAssertFalse(output.localizedCaseInsensitiveContains("</think>"))
if verifiesStructuredCommandOutput {
XCTAssertNotNil(SpokenEditCommandLLMResolver.resolution(from: output))
}
let sample = try residentSizeKB()
residentSamples.append(sample)
lifecycleSamples.append(sample)
Expand Down
30 changes: 29 additions & 1 deletion Tests/OpenTypeTests/EspressoFallbackTests.swift
Original file line numberDiff line numberDiff line change
Expand Up@@ -56,13 +56,38 @@ final class EspressoFallbackTests: XCTestCase {
XCTAssertTrue(result.usedMLX)
}

func testEspressoFailureReleasesANEStateBeforeMLXFallback() async throws {
var espressoIsLoaded = true
var espressoWasLoadedWhenMLXStarted = true

let result: (value: String, usedMLX: Bool) = try await TextProcessor.runEspressoWithMLXFallback(
espresso: { throw StubError.espresso },
prepareForMLXFallback: {
espressoIsLoaded = false
},
mlx: {
espressoWasLoadedWhenMLXStarted = espressoIsLoaded
return "mlx output"
}
)

XCTAssertEqual(result.value, "mlx output")
XCTAssertTrue(result.usedMLX)
XCTAssertFalse(espressoIsLoaded)
XCTAssertFalse(espressoWasLoadedWhenMLXStarted)
}

func testDisabledFallbackDoesNotRunMLX() async {
var preparedForFallback = false
var ranMLX = false

do {
_ = try await TextProcessor.runEspressoWithMLXFallback(
fallbackEnabled: false,
espresso: { throw StubError.espresso },
prepareForMLXFallback: {
preparedForFallback = true
},
mlx: {
ranMLX = true
return "mlx output"
Expand All@@ -71,6 +96,7 @@ final class EspressoFallbackTests: XCTestCase {
XCTFail("Expected the Espresso failure")
} catch {
XCTAssertEqual(error.localizedDescription, "espresso failed")
XCTAssertFalse(preparedForFallback)
XCTAssertFalse(ranMLX)
}
}
Expand DownExpand Up@@ -214,11 +240,13 @@ final class EspressoFallbackTests: XCTestCase {
temperature: 0
)
if index == 0 {
outcome = await processor.consumeEspressoOutcome()
let espressoIsLoaded = await processor.espressoLLM.isLoaded
XCTAssertFalse(espressoIsLoaded)
baselineFootprint = currentMemoryFootprint()
options.localLLMBackend = .mlx
}
}
outcome = await processor.consumeEspressoOutcome()
}

XCTAssertFalse(output.trimmingCharacters(in: .whitespacesAndNewlines).isEmpty)
Expand Down
Loading