Files
EdisonVoice/Sources/EdisonCore/Transcription/ParakeetEngine.swift
T

163 lines
7.1 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.
import AVFoundation
import FluidAudio
import Foundation
/// NVIDIA Parakeet TDT 0.6B, compiled to CoreML and run on the Neural Engine via FluidAudio.
///
/// **Batch, not streaming.** Audio is accumulated while the key is held and transcribed in
/// one pass on release. That's a deliberate trade: at ~100× realtime a 30-second utterance
/// resolves in roughly a third of a second, which is imperceptible for push-to-talk — but
/// it means no live text in the HUD while you speak, unlike Apple's engine.
/// FluidAudio's `SlidingWindowAsrManager` would restore live partials at the cost of a
/// more complex integration; see the note in `docs`.
public actor ParakeetEngine: TranscriptionEngine {
private var samples: [Float] = []
private var continuation: AsyncThrowingStream<TranscriptionChunk, Error>.Continuation?
/// Defaults to 16 kHz mono float32 — exactly what Parakeet is trained on.
private let converter = AudioConverter()
public init() {}
public func preferredInputFormat() async -> AVAudioFormat? {
// Parakeet is trained on 16 kHz mono; AudioCapture converts to whatever we ask for.
AVAudioFormat(commonFormat: .pcmFormatFloat32, sampleRate: 16_000, channels: 1, interleaved: false)
}
public func start() async throws -> AsyncThrowingStream<TranscriptionChunk, Error> {
samples.removeAll(keepingCapacity: true)
let (stream, continuation) = AsyncThrowingStream<TranscriptionChunk, Error>.makeStream()
self.continuation = continuation
// Force the (possibly very slow) first load to happen here rather than on release,
// so the user waits before speaking instead of losing an utterance to a timeout.
_ = try await ParakeetModels.shared.manager()
return stream
}
public func feed(_ chunk: AudioChunk) async {
let buffer = chunk.buffer
guard buffer.frameLength > 0 else { return }
// Delegated to FluidAudio's own converter rather than hand-rolled, for one reason
// that matters more than tidiness: `AsrManager.transcribe(_ samples: [Float])`
// performs **no resampling and no rate validation**. Feed it the wrong sample rate
// and it doesn't throw — it silently transcribes garbage.
//
// That's a live risk here. In compare mode the capture format is dictated by
// Apple's analyzer, and `bestAvailableAudioFormat` may legitimately return 8 kHz
// as well as 16 kHz. `resampleBuffer` normalizes whatever arrives to the 16 kHz
// mono float32 the model expects, and its Int16→Float path is bit-identical to
// dividing by 32768, so nothing is lost versus doing it by hand.
do {
samples.append(contentsOf: try converter.resampleBuffer(buffer))
} catch {
Log.speech.error("Parakeet: audio conversion failed — \(error.localizedDescription)")
}
}
public func finish() async {
defer {
continuation?.finish()
continuation = nil
samples.removeAll(keepingCapacity: true)
}
// Parakeet's encoder needs a minimum window; a stray tap of the key isn't speech.
// Logged rather than silent — an unexpected drop to zero here is how the
// format bug above disguised itself as a fast, empty result.
guard samples.count >= 1_600 else {
Log.speech.info("Parakeet: skipped — only \(self.samples.count) samples captured")
return
}
do {
let manager = try await ParakeetModels.shared.manager()
var decoderState = try TdtDecoderState()
let started = Date()
let result = try await manager.transcribe(samples, decoderState: &decoderState)
let elapsed = Date().timeIntervalSince(started)
let audioSeconds = Double(samples.count) / 16_000
Log.speech.info("""
Parakeet: \(audioSeconds, format: .fixed(precision: 1))s audio in \
\(elapsed, format: .fixed(precision: 2))s (\(audioSeconds / max(elapsed, 0.0001), format: .fixed(precision: 0))× realtime)
""")
continuation?.yield(
TranscriptionChunk(
text: result.text.trimmingCharacters(in: .whitespacesAndNewlines),
isFinal: true
)
)
} catch {
Log.speech.error("Parakeet failed: \(error.localizedDescription)")
continuation?.finish(throwing: error)
continuation = nil
}
}
}
/// Process-wide model cache.
///
/// Loading is expensive — ~470 MB downloaded on first ever run, then a few seconds from
/// disk per process — and the models are immutable once loaded, so every dictation shares
/// one instance rather than paying that per utterance. Its own actor because `static var`
/// on `ParakeetEngine` would be unprotected global mutable state under Swift 6.
public actor ParakeetModels {
public static let shared = ParakeetModels()
/// Whether the models are already on disk, checked without loading them.
///
/// `nonisolated` and filesystem-based on purpose: the menu needs this synchronously
/// while drawing, and an in-memory "have I loaded yet" flag would wrongly report
/// "not downloaded" on every fresh launch.
nonisolated public static var isDownloaded: Bool {
let support = FileManager.default.urls(for: .applicationSupportDirectory, in: .userDomainMask)[0]
let encoder = support
.appendingPathComponent("FluidAudio/Models/parakeet-tdt-0.6b-v3/Encoder.mlmodelc")
return FileManager.default.fileExists(atPath: encoder.path)
}
private var loaded: AsrManager?
private var loadTask: Task<AsrManager, Error>?
var isLoaded: Bool { loaded != nil }
/// Loads once; concurrent callers await the same task rather than racing to download.
public func manager() async throws -> AsrManager {
if let loaded { return loaded }
if let loadTask { return try await loadTask.value }
let task = Task<AsrManager, Error> {
// Built as a value first: os.Logger requires a literal interpolation, so a
// ternary can't be passed directly as the argument.
let stage = Self.isDownloaded
? "loading models from disk"
: "downloading models (~470 MB, one time)"
Log.speech.info("Parakeet: \(stage, privacy: .public)")
let started = Date()
let models = try await AsrModels.downloadAndLoad(version: .v3, encoderPrecision: .int8)
let manager = AsrManager(config: .default)
try await manager.loadModels(models)
Log.speech.info("Parakeet: ready in \(Date().timeIntervalSince(started), format: .fixed(precision: 1))s")
return manager
}
loadTask = task
do {
let manager = try await task.value
loaded = manager
return manager
} catch {
// Don't cache a failed load — a transient download error shouldn't wedge the
// engine for the rest of the session.
loadTask = nil
throw error
}
}
}