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:
@@ -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
|
||||
Reference in New Issue
Block a user