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+).
This commit is contained in:
Rocky
2026-06-23 00:46:58 +08:00
parent 5e5122f172
commit df1c5ff32c
160 changed files with 22080 additions and 492 deletions
@@ -0,0 +1,389 @@
#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