Files
gart/code/Hex-main/Hex/Clients/TranscriptionClient.swift

429 lines
16 KiB
Swift
Executable File
Raw Permalink 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.
//
// TranscriptionClient.swift
// Hex
//
// Created by Kit Langton on 1/24/25.
//
import AVFoundation
import Dependencies
import DependenciesMacros
import Foundation
import HexCore
import WhisperKit
private let transcriptionLogger = HexLog.transcription
private let modelsLogger = HexLog.models
private let parakeetLogger = HexLog.parakeet
/// A client that downloads and loads WhisperKit models, then transcribes audio files using the loaded model.
/// Exposes progress callbacks to report overall download-and-load percentage and transcription progress.
@DependencyClient
struct TranscriptionClient {
/// Transcribes an audio file at the specified `URL` using the named `model`.
/// Reports transcription progress via `progressCallback`.
var transcribe: @Sendable (URL, String, DecodingOptions, @escaping (Progress) -> Void) async throws -> String
/// Ensures a model is downloaded (if missing) and loaded into memory, reporting progress via `progressCallback`.
var downloadModel: @Sendable (String, @escaping (Progress) -> Void) async throws -> Void
/// Deletes a model from disk if it exists
var deleteModel: @Sendable (String) async throws -> Void
/// Checks if a named model is already downloaded on this system.
var isModelDownloaded: @Sendable (String) async -> Bool = { _ in false }
/// Fetches a recommended set of models for the user's hardware from Hugging Face's `argmaxinc/whisperkit-coreml`.
var getRecommendedModels: @Sendable () async throws -> ModelSupport
/// Lists all model variants found in `argmaxinc/whisperkit-coreml`.
var getAvailableModels: @Sendable () async throws -> [String]
}
extension TranscriptionClient: DependencyKey {
static var liveValue: Self {
let live = TranscriptionClientLive()
return Self(
transcribe: { try await live.transcribe(url: $0, model: $1, options: $2, progressCallback: $3) },
downloadModel: { try await live.downloadAndLoadModel(variant: $0, progressCallback: $1) },
deleteModel: { try await live.deleteModel(variant: $0) },
isModelDownloaded: { await live.isModelDownloaded($0) },
getRecommendedModels: { await live.getRecommendedModels() },
getAvailableModels: { try await live.getAvailableModels() }
)
}
}
extension DependencyValues {
var transcription: TranscriptionClient {
get { self[TranscriptionClient.self] }
set { self[TranscriptionClient.self] = newValue }
}
}
/// An `actor` that manages WhisperKit models by downloading (from Hugging Face),
// loading them into memory, and then performing transcriptions.
actor TranscriptionClientLive {
// MARK: - Stored Properties
/// The current in-memory `WhisperKit` instance, if any.
private var whisperKit: WhisperKit?
/// The name of the currently loaded model, if any.
private var currentModelName: String?
private var parakeet: ParakeetClient = ParakeetClient()
/// The base folder under which we store model data (e.g., ~/Library/Application Support/...).
private lazy var modelsBaseFolder: URL = {
do {
return try URL.hexModelsDirectory
} catch {
fatalError("Could not create Application Support folder: \(error)")
}
}()
// MARK: - Public Methods
/// Ensures the given `variant` model is downloaded and loaded, reporting
/// overall progress (0%–50% for downloading, 50%–100% for loading).
func downloadAndLoadModel(variant: String, progressCallback: @escaping (Progress) -> Void) async throws {
// If Parakeet, use Parakeet client path
if isParakeet(variant) {
try await parakeet.ensureLoaded(modelName: variant, progress: progressCallback)
currentModelName = variant
return
}
// Resolve wildcard patterns (e.g., "distil*large-v3") to a concrete variant
let variant = await resolveVariant(variant)
// Special handling for corrupted or malformed variant names
if variant.isEmpty {
throw NSError(
domain: "TranscriptionClient",
code: -3,
userInfo: [
NSLocalizedDescriptionKey: "Cannot download model: Empty model name",
]
)
}
let overallProgress = Progress(totalUnitCount: 100)
overallProgress.completedUnitCount = 0
progressCallback(overallProgress)
modelsLogger.info("Preparing model download and load for \(variant)")
// 1) Model download phase (0-50% progress)
if !(await isModelDownloaded(variant)) {
try await downloadModelIfNeeded(variant: variant) { downloadProgress in
let fraction = downloadProgress.fractionCompleted * 0.5
overallProgress.completedUnitCount = Int64(fraction * 100)
progressCallback(overallProgress)
}
} else {
// Skip download phase if already downloaded
overallProgress.completedUnitCount = 50
progressCallback(overallProgress)
}
// 2) Model loading phase (50-100% progress)
try await loadWhisperKitModel(variant) { loadingProgress in
let fraction = 0.5 + (loadingProgress.fractionCompleted * 0.5)
overallProgress.completedUnitCount = Int64(fraction * 100)
progressCallback(overallProgress)
}
// Final progress update
overallProgress.completedUnitCount = 100
progressCallback(overallProgress)
}
/// Deletes a model from disk if it exists
func deleteModel(variant: String) async throws {
if isParakeet(variant) {
try await parakeet.deleteCaches(modelName: variant)
if currentModelName == variant { unloadCurrentModel() }
return
}
let modelFolder = modelPath(for: variant)
// Check if the model exists
guard FileManager.default.fileExists(atPath: modelFolder.path) else {
// Model doesn't exist, nothing to delete
return
}
// If this is the currently loaded model, unload it first
if currentModelName == variant {
unloadCurrentModel()
}
// Delete the model directory
try FileManager.default.removeItem(at: modelFolder)
modelsLogger.info("Deleted model \(variant)")
}
/// Returns `true` if the model is already downloaded to the local folder.
/// Performs a thorough check to ensure the model files are actually present and usable.
func isModelDownloaded(_ modelName: String) async -> Bool {
if isParakeet(modelName) {
let available = await parakeet.isModelAvailable(modelName)
parakeetLogger.debug("Parakeet available? \(available)")
return available
}
let modelFolderPath = modelPath(for: modelName).path
let fileManager = FileManager.default
// First, check if the basic model directory exists
guard fileManager.fileExists(atPath: modelFolderPath) else {
// Don't print logs that would spam the console
return false
}
do {
// Check if the directory has actual model files in it
let contents = try fileManager.contentsOfDirectory(atPath: modelFolderPath)
// Model should have multiple files and certain key components
guard !contents.isEmpty else {
return false
}
// Check for specific model structure - need both tokenizer and model files
let hasModelFiles = contents.contains { $0.hasSuffix(".mlmodelc") || $0.contains("model") }
let tokenizerFolderPath = tokenizerPath(for: modelName).path
let hasTokenizer = fileManager.fileExists(atPath: tokenizerFolderPath)
// Both conditions must be true for a model to be considered downloaded
return hasModelFiles && hasTokenizer
} catch {
return false
}
}
/// Returns a list of recommended models based on current device hardware.
func getRecommendedModels() async -> ModelSupport {
await WhisperKit.recommendedRemoteModels()
}
/// Lists all model variants available in the `argmaxinc/whisperkit-coreml` repository.
func getAvailableModels() async throws -> [String] {
var names = try await WhisperKit.fetchAvailableModels()
#if canImport(FluidAudio)
for model in ParakeetModel.allCases.reversed() {
if !names.contains(model.identifier) { names.insert(model.identifier, at: 0) }
}
#endif
return names
}
/// Transcribes the audio file at `url` using a `model` name.
/// If the model is not yet loaded (or if it differs from the current model), it is downloaded and loaded first.
/// Transcription progress can be monitored via `progressCallback`.
func transcribe(
url: URL,
model: String,
options: DecodingOptions,
progressCallback: @escaping (Progress) -> Void
) async throws -> String {
let startAll = Date()
if isParakeet(model) {
transcriptionLogger.notice("Transcribing with Parakeet model=\(model) file=\(url.lastPathComponent)")
let startLoad = Date()
try await downloadAndLoadModel(variant: model) { p in
progressCallback(p)
}
transcriptionLogger.info("Parakeet ensureLoaded took \(String(format: "%.2f", Date().timeIntervalSince(startLoad)))s")
let preparedClip = try ParakeetClipPreparer.ensureMinimumDuration(url: url, logger: parakeetLogger)
defer { preparedClip.cleanup() }
let startTx = Date()
let text = try await parakeet.transcribe(preparedClip.url)
transcriptionLogger.info("Parakeet transcription took \(String(format: "%.2f", Date().timeIntervalSince(startTx)))s")
transcriptionLogger.info("Parakeet request total elapsed \(String(format: "%.2f", Date().timeIntervalSince(startAll)))s")
return text
}
let model = await resolveVariant(model)
// Load or switch to the required model if needed.
if whisperKit == nil || model != currentModelName {
unloadCurrentModel()
let startLoad = Date()
try await downloadAndLoadModel(variant: model) { p in
// Debug logging, or scale as desired:
progressCallback(p)
}
let loadDuration = Date().timeIntervalSince(startLoad)
transcriptionLogger.info("WhisperKit ensureLoaded model=\(model) took \(String(format: "%.2f", loadDuration))s")
}
guard let whisperKit = whisperKit else {
throw NSError(
domain: "TranscriptionClient",
code: -1,
userInfo: [
NSLocalizedDescriptionKey: "Failed to initialize WhisperKit for model: \(model)",
]
)
}
// Perform the transcription.
transcriptionLogger.notice("Transcribing with WhisperKit model=\(model) file=\(url.lastPathComponent)")
let startTx = Date()
let results = try await whisperKit.transcribe(audioPath: url.path, decodeOptions: options)
transcriptionLogger.info("WhisperKit transcription took \(String(format: "%.2f", Date().timeIntervalSince(startTx)))s")
transcriptionLogger.info("WhisperKit request total elapsed \(String(format: "%.2f", Date().timeIntervalSince(startAll)))s")
// Concatenate results from all segments.
let text = results.map(\.text).joined(separator: " ")
return text
}
// MARK: - Private Helpers
/// Resolve wildcard patterns (e.g. "distil*large-v3") to a concrete model name.
/// Preference: downloaded > non-turbo > any match.
private func resolveVariant(_ variant: String) async -> String {
guard variant.contains("*") || variant.contains("?") else { return variant }
let names: [String]
do { names = try await WhisperKit.fetchAvailableModels() } catch { return variant }
// Build tuple array with download status for matching models
var models: [(name: String, isDownloaded: Bool)] = []
for name in names where ModelPatternMatcher.matches(variant, name) {
models.append((name, await isModelDownloaded(name)))
}
return ModelPatternMatcher.resolvePattern(variant, from: models) ?? variant
}
private func isParakeet(_ name: String) -> Bool {
ParakeetModel(rawValue: name) != nil
}
/// Creates or returns the local folder (on disk) for a given `variant` model.
private func modelPath(for variant: String) -> URL {
// Remove any possible path traversal or invalid characters from variant name
let sanitizedVariant = variant.components(separatedBy: CharacterSet(charactersIn: "./\\")).joined(separator: "_")
return modelsBaseFolder
.appendingPathComponent("argmaxinc")
.appendingPathComponent("whisperkit-coreml")
.appendingPathComponent(sanitizedVariant, isDirectory: true)
}
/// Creates or returns the local folder for the tokenizer files of a given `variant`.
private func tokenizerPath(for variant: String) -> URL {
modelPath(for: variant).appendingPathComponent("tokenizer", isDirectory: true)
}
// Unloads any currently loaded model (clears `whisperKit` and `currentModelName`).
private func unloadCurrentModel() {
whisperKit = nil
currentModelName = nil
}
/// Downloads the model to a temporary folder (if it isn't already on disk),
/// then moves it into its final folder in `modelsBaseFolder`.
private func downloadModelIfNeeded(
variant: String,
progressCallback: @escaping (Progress) -> Void
) async throws {
let modelFolder = modelPath(for: variant)
// If the model folder exists but isn't a complete model, clean it up
let isDownloaded = await isModelDownloaded(variant)
if FileManager.default.fileExists(atPath: modelFolder.path), !isDownloaded {
try FileManager.default.removeItem(at: modelFolder)
}
// If model is already fully downloaded, we're done
if isDownloaded {
return
}
modelsLogger.info("Downloading model \(variant)")
// Create parent directories
let parentDir = modelFolder.deletingLastPathComponent()
try FileManager.default.createDirectory(at: parentDir, withIntermediateDirectories: true)
do {
// Download directly using the exact variant name provided
// WhisperKit 0.15.0 changed downloader params: passing
// "argmaxinc/whisperkit-coreml" to a parameter interpreted as a host
// yields NSURLErrorCannotFindHost in production builds that need
// to fetch models for the first time. Let WhisperKit use its
// default repo/host (Hugging Face) by omitting the repo/host arg.
let tempFolder = try await WhisperKit.download(
variant: variant,
downloadBase: nil,
useBackgroundSession: false,
progressCallback: { progress in
progressCallback(progress)
}
)
// Ensure target folder exists
try FileManager.default.createDirectory(at: modelFolder, withIntermediateDirectories: true)
// Move the downloaded snapshot to the final location
try moveContents(of: tempFolder, to: modelFolder)
modelsLogger.info("Downloaded model to \(modelFolder.path)")
} catch {
// Clean up any partial download if an error occurred
FileManager.default.removeItemIfExists(at: modelFolder)
// Rethrow the original error
modelsLogger.error("Error downloading model \(variant): \(error.localizedDescription)")
throw error
}
}
/// Loads a local model folder via `WhisperKitConfig`, optionally reporting load progress.
private func loadWhisperKitModel(
_ modelName: String,
progressCallback: @escaping (Progress) -> Void
) async throws {
let loadingProgress = Progress(totalUnitCount: 100)
loadingProgress.completedUnitCount = 0
progressCallback(loadingProgress)
let modelFolder = modelPath(for: modelName)
let tokenizerFolder = tokenizerPath(for: modelName)
// Use WhisperKit's config to load the model
let config = WhisperKitConfig(
model: modelName,
modelFolder: modelFolder.path,
tokenizerFolder: tokenizerFolder,
// verbose: true,
// logLevel: .debug,
prewarm: false,
load: true
)
// The initializer automatically calls `loadModels`.
whisperKit = try await WhisperKit(config)
currentModelName = modelName
// Finalize load progress
loadingProgress.completedUnitCount = 100
progressCallback(loadingProgress)
modelsLogger.info("Loaded WhisperKit model \(modelName)")
}
/// Moves all items from `sourceFolder` into `destFolder` (shallow move of directory contents).
private func moveContents(of sourceFolder: URL, to destFolder: URL) throws {
let fileManager = FileManager.default
let items = try fileManager.contentsOfDirectory(atPath: sourceFolder.path)
for item in items {
let src = sourceFolder.appendingPathComponent(item)
let dst = destFolder.appendingPathComponent(item)
try fileManager.moveItem(at: src, to: dst)
}
}
}