Files
OSGKeyboard/OSGKeyboard/ThirdParty/Qwen3Speech/Sources/Qwen3ASR/CoreMLASRModel.swift
T
Rocky df1c5ff32c feat: migrate on-device Qwen3 ASR to CoreML for background Flow dictation
Replace MLX GPU inference with CoreML bundles so transcription continues
while the host app is backgrounded. Adds model download and warm-up,
vendored Qwen3Speech, and updates onboarding, settings, and copy for the
~1.6 GB CoreML package (iOS 18+).
2026-06-23 00:46:58 +08:00

390 lines
16 KiB
Swift
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
#if canImport(CoreML)
import CoreML
import Foundation
import MLX
import AudioCommon
/// Full CoreML ASR model: CoreML encoder + CoreML text decoder.
///
/// Runs the entire Qwen3-ASR pipeline on CoreML (Neural Engine + CPU),
/// eliminating the MLX GPU dependency. Requires macOS 15+ / iOS 18+
/// for MLState KV cache support.
public class CoreMLASRModel {
public let encoder: CoreMLASREncoder
public let decoder: CoreMLTextDecoder
public let featureExtractor: WhisperFeatureExtractor
private var tokenizer: Qwen3Tokenizer?
public init(encoder: CoreMLASREncoder, decoder: CoreMLTextDecoder) {
self.encoder = encoder
self.decoder = decoder
self.featureExtractor = WhisperFeatureExtractor()
}
/// Load full CoreML ASR from HuggingFace.
///
/// Downloads encoder and decoder models from `aufklarer/Qwen3-ASR-CoreML`.
///
/// Compute units are split per-component because the encoder and decoder
/// have different optimal backends: the encoder defaults to `.all` (the
/// 30 s fixed-shape graph runs well on GPU on most Macs), while the
/// decoder defaults to `.cpuAndNeuralEngine` (the autoregressive
/// MLState path is ANE-friendly and ~7× faster there than GPU per the
/// rebuilt-encoder PR). A single `computeUnits` parameter that
/// propagated into both calls silently overrode the decoder's stated
/// `.cpuAndNeuralEngine` default with `.all`, costing real ANE
/// throughput on the full-pipeline path.
public static func fromPretrained(
encoderModelId: String = CoreMLASREncoder.defaultModelId,
decoderModelId: String = CoreMLASREncoder.defaultModelId,
tokenizerModelId: String = "aufklarer/Qwen3-ASR-0.6B-MLX-4bit",
encoderComputeUnits: MLComputeUnits = CoreMLComputeUnitsResolver.resolved(default: .all),
decoderComputeUnits: MLComputeUnits = CoreMLComputeUnitsResolver.resolved(default: .cpuAndNeuralEngine),
cacheDir: URL? = nil,
offlineMode: Bool = false,
progressHandler: ((Double, String) -> Void)? = nil
) async throws -> CoreMLASRModel {
// Download encoder (0-30%)
progressHandler?(0.0, "Loading CoreML encoder...")
let enc = try await CoreMLASREncoder.fromPretrained(
modelId: encoderModelId,
computeUnits: encoderComputeUnits,
cacheDir: cacheDir,
offlineMode: offlineMode
) { p, msg in
progressHandler?(p * 0.3, msg)
}
// Download decoder (30-80%)
progressHandler?(0.3, "Loading CoreML decoder...")
let dec = try await CoreMLTextDecoder.fromPretrained(
modelId: decoderModelId,
computeUnits: decoderComputeUnits,
cacheDir: cacheDir,
offlineMode: offlineMode
) { p, msg in
progressHandler?(0.3 + p * 0.5, msg)
}
// Download tokenizer (80-90%)
progressHandler?(0.8, "Loading tokenizer...")
let tokenizerDir = try cacheDir ?? HuggingFaceDownloader.getCacheDirectory(for: tokenizerModelId)
try await HuggingFaceDownloader.downloadWeights(
modelId: tokenizerModelId,
to: tokenizerDir,
additionalFiles: ["vocab.json", "merges.txt", "tokenizer_config.json"],
offlineMode: offlineMode
)
let model = CoreMLASRModel(encoder: enc, decoder: dec)
let vocabPath = tokenizerDir.appendingPathComponent("vocab.json")
if FileManager.default.fileExists(atPath: vocabPath.path) {
let tokenizer = Qwen3Tokenizer()
try tokenizer.load(from: vocabPath)
model.tokenizer = tokenizer
}
progressHandler?(1.0, "Ready")
return model
}
/// Warm up both encoder and decoder.
public func warmUp() throws {
try encoder.warmUp()
try decoder.warmUp()
}
/// Transcribe audio to text using full CoreML pipeline.
///
/// The entire inference runs on CoreML (Neural Engine + CPU) without MLX GPU.
public func transcribe(
audio: [Float],
sampleRate: Int = 16000,
language: String? = nil,
maxTokens: Int = 448
) throws -> String {
let profile = ProcessInfo.processInfo.environment["COREML_ASR_PROFILE"] == "1"
let t0 = CFAbsoluteTimeGetCurrent()
let melFeatures = featureExtractor.process(audio, sampleRate: sampleRate)
let t1 = CFAbsoluteTimeGetCurrent()
// The encoder pads mel to a fixed 30 s shape and reports the real
// (un-padded) audio-token count via ``output_length``. We feed only
// the first ``numAudioTokens`` of the padded embeddings to the
// decoder so trailing zero-derived tokens don't pollute attention.
let (audioEmbeds, numAudioTokens) = try encoder.encode(melFeatures)
let t2 = CFAbsoluteTimeGetCurrent()
decoder.resetCache()
// Build chat template token sequence
let imStartId: Int32 = 151644
let imEndId: Int32 = 151645
let audioStartId: Int32 = 151669
let audioEndId: Int32 = 151670
let asrTextId: Int32 = 151704
let newlineId: Int32 = 198
let systemId: Int32 = 8948
let userId: Int32 = 872
let assistantId: Int32 = 77091
// <|im_start|>system\n<|im_end|>\n
var prefixTokens: [Int32] = [imStartId, systemId, newlineId, imEndId, newlineId]
// <|im_start|>user\n<|audio_start|>
prefixTokens += [imStartId, userId, newlineId, audioStartId]
// <|audio_end|><|im_end|>\n<|im_start|>assistant\n
var suffixTokens: [Int32] = [audioEndId, imEndId, newlineId, imStartId, assistantId, newlineId]
// Language hint + <|asr_text|>
if let lang = language, let tokenizer = tokenizer {
let langPrefix = "language \(lang)"
let langTokens = tokenizer.encode(langPrefix)
suffixTokens += langTokens.map { Int32($0) }
}
suffixTokens.append(asrTextId)
// ── Prefill: process all prefix tokens (one batched call) ──
var lastLogits: MLMultiArray?
lastLogits = try decoder.decoderPrefillTokens(prefixTokens)
let t3 = CFAbsoluteTimeGetCurrent()
// ── Prefill: process audio embeddings ──
// Bulk-extract the MLX audio embeddings once (single Metal sync)
// then feed them to the decoder in batched chunks of
// ``prefillBatchSize`` tokens. The fixed-T CoreML decoder packs
// T tokens per ANE dispatch, so a 20 s / 250-token clip becomes
// ~2 calls instead of 250 (each ANE dispatch costs ~30 ms — the
// dispatch overhead dominated the per-step cost in profiling).
let _ = audioEmbeds.dim(2) // sanity-check hidden dim
let audioEmbedsFlat: [Float] = audioEmbeds.asArray(Float.self)
let chunk = decoder.prefillBatchSize
var consumed = 0
while consumed < numAudioTokens {
let n = min(chunk, numAudioTokens - consumed)
lastLogits = try decoder.decoderPrefill(
flatEmbeddings: audioEmbedsFlat,
offset: consumed,
realCount: n,
)
consumed += n
}
let t4 = CFAbsoluteTimeGetCurrent()
// ── Prefill: process suffix tokens (one batched call) ──
lastLogits = try decoder.decoderPrefillTokens(suffixTokens)
let t5 = CFAbsoluteTimeGetCurrent()
// ── Autoregressive generation ──
guard var logits = lastLogits else {
return "[CoreML decoder: no output]"
}
// Known WER issue with this path (multi-sentence utterances on
// LibriSpeech test-clean): the CoreML port emits ``<|im_end|>``
// after the first sentence-final period with a wide logit margin
// (~6+ nats over the runner-up). The MLX path at the same
// effective bit width keeps generating. We tried both a
// logit-margin guard and a force-first-EOS-suppression; neither
// helps — the model's runner-up at the truncation point is also
// wrong (e.g. " The" instead of " on"), so substituting it just
// trades deletions for substitutions. The root cause is upstream
// — encoder INT8 quantization / mel padding leakage into audio
// embeddings, or position drift in chunked prefill. Tracked as
// a separate fix that requires model re-export, not a sampler
// change. The ``argmax(skipping:)`` / ``logit(_:at:)`` helpers
// stay for that future work.
var generatedTokens: [Int32] = []
var nextToken = decoder.argmax(logits: logits)
// First-token EOS would mean the model thinks the audio yielded an
// empty transcript — never right on real speech. Cheap to guard.
if nextToken == imEndId {
nextToken = decoder.argmax(logits: logits, skipping: imEndId)
}
generatedTokens.append(nextToken)
for _ in 1..<maxTokens {
if nextToken == imEndId { break }
let embedding = try decoder.embed(tokenId: nextToken)
logits = try decoder.decoderStep(embedding: embedding)
nextToken = decoder.argmax(logits: logits)
generatedTokens.append(nextToken)
}
if profile {
let audioDuration = Double(audio.count) / Double(sampleRate)
print("[COREML-ASR-PROFILE] audio=\(audioDuration)s gen=\(generatedTokens.count)")
}
let t6 = CFAbsoluteTimeGetCurrent()
if profile {
let ms = { (a: CFAbsoluteTime, b: CFAbsoluteTime) in (b - a) * 1000 }
print(String(format: "[COREML-ASR-PROFILE] mel=%.0fms encoder=%.0fms prefix=%.0fms audio_prefill=%.0fms(%dtok→%.1fms/tok) suffix=%.0fms gen=%.0fms(%dtok→%.1fms/tok) total=%.0fms",
ms(t0, t1), ms(t1, t2), ms(t2, t3),
ms(t3, t4), numAudioTokens, ms(t3, t4) / Double(max(numAudioTokens, 1)),
ms(t4, t5),
ms(t5, t6), generatedTokens.count, ms(t5, t6) / Double(max(generatedTokens.count, 1)),
ms(t0, t6)))
}
// Decode tokens
if let tokenizer = tokenizer {
let rawText = tokenizer.decode(tokens: generatedTokens.map { Int($0) })
if let range = rawText.range(of: "<asr_text>") {
return String(rawText[range.upperBound...]).trimmingCharacters(in: .whitespaces)
}
return rawText
} else {
return generatedTokens.map { String($0) }.joined(separator: " ")
}
}
// MARK: - MLX-Free Transcription
/// Transcribe audio to text without any MLX/Metal dependency.
///
/// Uses `featureExtractor.processRaw()` (CPU via Accelerate) and
/// `encoder.encode(melData:melBins:timeFrames:)` (CoreML) to produce
/// MLMultiArray embeddings, then decodes using `audioEmbeddingFromMultiArray()`.
///
/// This method is safe for iOS background execution where Metal GPU eval
/// (triggered by MLXArray operations) would cause a crash.
///
/// - Note: Requires `processRaw()` on WhisperFeatureExtractor and
/// `encode(melData:melBins:timeFrames:)` on CoreMLASREncoder, both added by T2.
public func transcribeWithoutMLX(
audio: [Float],
sampleRate: Int = 16000,
language: String? = nil,
maxTokens: Int = 448
) throws -> String {
// 1. Extract mel features (pure CPU via Accelerate — no MLXArray)
let melFeatures = featureExtractor.processRaw(audio, sampleRate: sampleRate)
// 2. Encode audio → MLMultiArray embeddings + real (un-padded)
// audio-token count from the encoder's ``output_length``.
let encoded = try encoder.encode(
melData: melFeatures.data,
melBins: melFeatures.melBins,
timeFrames: melFeatures.timeFrames
)
let audioEmbeds = encoded.embeddings
let numAudioTokens = encoded.outputLength
// 3. Reset decoder KV cache
decoder.resetCache()
// 4. Build chat template token sequence (identical to transcribe())
let imStartId: Int32 = 151644
let imEndId: Int32 = 151645
let audioStartId: Int32 = 151669
let audioEndId: Int32 = 151670
let asrTextId: Int32 = 151704
let newlineId: Int32 = 198
let systemId: Int32 = 8948
let userId: Int32 = 872
let assistantId: Int32 = 77091
// <|im_start|>system\n<|im_end|>\n
var prefixTokens: [Int32] = [imStartId, systemId, newlineId, imEndId, newlineId]
// <|im_start|>user\n<|audio_start|>
prefixTokens += [imStartId, userId, newlineId, audioStartId]
// <|audio_end|><|im_end|>\n<|im_start|>assistant\n
var suffixTokens: [Int32] = [audioEndId, imEndId, newlineId, imStartId, assistantId, newlineId]
// Language hint + <|asr_text|>
if let lang = language, let tokenizer = tokenizer {
let langPrefix = "language \(lang)"
let langTokens = tokenizer.encode(langPrefix)
suffixTokens += langTokens.map { Int32($0) }
}
suffixTokens.append(asrTextId)
// 5. Prefill: process all prefix tokens
var lastLogits: MLMultiArray?
for token in prefixTokens {
let embedding = try decoder.embed(tokenId: token)
lastLogits = try decoder.decoderStep(embedding: embedding)
}
// 6. Prefill: process audio embeddings (MLX-free path)
for i in 0..<numAudioTokens {
let audioEmbed = try decoder.audioEmbeddingFromMultiArray(audioEmbeds, at: i)
lastLogits = try decoder.decoderStep(embedding: audioEmbed)
}
// Prefill: process suffix tokens
for token in suffixTokens {
let embedding = try decoder.embed(tokenId: token)
lastLogits = try decoder.decoderStep(embedding: embedding)
}
// 7. Autoregressive generation (same EOS note as `transcribe()` —
// see that path for the background).
guard var logits = lastLogits else {
return "[CoreML decoder: no output]"
}
var generatedTokens: [Int32] = []
var nextToken = decoder.argmax(logits: logits)
if nextToken == imEndId {
nextToken = decoder.argmax(logits: logits, skipping: imEndId)
}
generatedTokens.append(nextToken)
for _ in 1..<maxTokens {
if nextToken == imEndId { break }
let embedding = try decoder.embed(tokenId: nextToken)
logits = try decoder.decoderStep(embedding: embedding)
nextToken = decoder.argmax(logits: logits)
generatedTokens.append(nextToken)
}
// Decode tokens
if let tokenizer = tokenizer {
let rawText = tokenizer.decode(tokens: generatedTokens.map { Int($0) })
if let range = rawText.range(of: "<asr_text>") {
return String(rawText[range.upperBound...]).trimmingCharacters(in: .whitespaces)
}
return rawText
} else {
return generatedTokens.map { String($0) }.joined(separator: " ")
}
}
}
// MARK: - SpeechRecognitionModel
extension CoreMLASRModel: SpeechRecognitionModel {
public var inputSampleRate: Int { 16000 }
public func transcribe(audio: [Float], sampleRate: Int, language: String?) -> String {
do {
return try transcribe(audio: audio, sampleRate: sampleRate, language: language, maxTokens: 448)
} catch {
return "[CoreML error: \(error.localizedDescription)]"
}
}
}
// MARK: - Background-Safe Transcription
extension CoreMLASRModel {
/// Background-safe transcription (no MLX/Metal dependency).
///
/// Uses `transcribeWithoutMLX()` which avoids all MLXArray operations
/// that would trigger Metal GPU eval. Safe to call from iOS background
/// audio processing where GPU access is prohibited.
public func transcribeBackgroundSafe(audio: [Float], sampleRate: Int, language: String?) -> String {
do {
return try transcribeWithoutMLX(audio: audio, sampleRate: sampleRate, language: language)
} catch {
return "[CoreML error: \(error.localizedDescription)]"
}
}
}
#endif