diff --git a/README.md b/README.md index 4c9b8a2..2d1b4f4 100644 --- a/README.md +++ b/README.md @@ -4,9 +4,9 @@ # Local Wisper -Local speech-to-text for Linux, built around Sway. Press a shortcut, speak, and -the transcript appears in the focused window without sending your audio to the -cloud. +English speech-to-text for Linux, built around Sway and NVIDIA's +Parakeet Unified EN 0.6B model. Press a shortcut, speak, and the transcript +appears in the focused window without sending your audio to the cloud. ## Requirements diff --git a/baml_src/app.baml b/baml_src/app.baml index 1dce857..e7fb3c1 100644 --- a/baml_src/app.baml +++ b/baml_src/app.baml @@ -84,8 +84,8 @@ function parse_options(args: string[]) -> AppOptions { }, "--model" => { let value = required_value(args, index + 1, arg); - if (value != "nvidia/parakeet-tdt-0.6b-v3") { - invalid_argument("only --model nvidia/parakeet-tdt-0.6b-v3 is supported") + if (value != "nvidia/parakeet-unified-en-0.6b") { + invalid_argument("only --model nvidia/parakeet-unified-en-0.6b is supported") } index += 2 }, @@ -342,6 +342,13 @@ test "bare arguments select interactive recording" { assert.equal(parse_options([]).command, AppCommand.Record) } +test "CLI accepts the Parakeet Unified model name" { + assert.equal( + parse_options(["--model", "nvidia/parakeet-unified-en-0.6b"]).command, + AppCommand.Record, + ) +} + test "post-processing model is user configurable" { let options = parse_options(["--post-process-model", "gpt-5.6-luna-next"]); assert.equal(options.post_process_model, "gpt-5.6-luna-next") diff --git a/baml_src/model.baml b/baml_src/model.baml index f48ba4b..56b0833 100644 --- a/baml_src/model.baml +++ b/baml_src/model.baml @@ -1,5 +1,5 @@ enum ModelVariant { - Fp16, + Fp32, Int8, } @@ -11,45 +11,57 @@ class ModelAsset { } function model_assets(variant: ModelVariant) -> ModelAsset[] { - let vocabulary = ModelAsset { - remote_name: "vocab.txt", - local_name: "vocab.txt", - size: 102132, - sha256: "ba8e4007c65f4bb4358ffe2ecc13d9ccc7a10351151065242b5c3a943e685742", + let tokenizer = ModelAsset { + remote_name: "tokenizer.model", + local_name: "tokenizer.model", + size: 251056, + sha256: "07d4e5a63840a53ab2d4d106d2874768143fb3fbdd47938b3910d2da05bfb0a9", }; match (variant) { - ModelVariant.Fp16 => { + ModelVariant.Fp32 => { [ ModelAsset { - remote_name: "encoder-model.fp16.onnx", - local_name: "encoder-model.onnx", - size: 1238960452, - sha256: "a2bdeeb99cb7e5548818e823127b33854dd0c26f5d0c8da91effdd895ea0e717", + remote_name: "encoder.onnx", + local_name: "encoder.onnx", + size: 41901475, + sha256: "50083dd0b2af503b87ea8fcf9fd7d7dc15dda3130e4113017c8b37ed00418ba4", }, ModelAsset { - remote_name: "decoder_joint-model.fp16.onnx", - local_name: "decoder_joint-model.onnx", - size: 36266140, - sha256: "b33a73b7c1d71b9d5a0911f5cb478be3dcbf79f53355c531ab1cd1dcd68ad8ef", + remote_name: "encoder.onnx.data", + local_name: "encoder.onnx.data", + size: 2437091328, + sha256: "c054b2932ee13aa39ed84982a399822f3193311fb6fd8e1e204566722614bd87", }, - vocabulary, + ModelAsset { + remote_name: "decoder_joint.onnx", + local_name: "decoder_joint.onnx", + size: 35779240, + sha256: "64648c91935ea4819e9c31ea8a20f011d3b1c960be3861cb8bc74df0259cd998", + }, + tokenizer, ] }, ModelVariant.Int8 => { [ ModelAsset { - remote_name: "encoder-model.int8.onnx", - local_name: "encoder-model.onnx", - size: 652183999, - sha256: "6139d2fa7e1b086097b277c7149725edbab89cc7c7ae64b23c741be4055aff09", + remote_name: "encoder.int8.onnx", + local_name: "encoder.int8.onnx", + size: 42606669, + sha256: "c81adfab77634e00c1668a221a14f244c5fb3409e7c14eeebaf6ac963425910f", + }, + ModelAsset { + remote_name: "encoder.int8.onnx.data", + local_name: "encoder.int8.onnx.data", + size: 611491584, + sha256: "3d54dd04646c15677bd2844a84df3770b12cc1ce183481f7b6e0def31c92114a", }, ModelAsset { - remote_name: "decoder_joint-model.int8.onnx", - local_name: "decoder_joint-model.onnx", - size: 18202004, - sha256: "eea7483ee3d1a30375daedc8ed83e3960c91b098812127a0d99d1c8977667a70", + remote_name: "decoder_joint.int8.onnx", + local_name: "decoder_joint.int8.onnx", + size: 8995064, + sha256: "7f76ad5f35035f25630075699c6c942a2c0c05ff42cb398f966f3c256d148e1e", }, - vocabulary, + tokenizer, ] }, } @@ -57,15 +69,15 @@ function model_assets(variant: ModelVariant) -> ModelAsset[] { function model_variant_name(variant: ModelVariant) -> string { match (variant) { - ModelVariant.Fp16 => "FP16", + ModelVariant.Fp32 => "FP32", ModelVariant.Int8 => "INT8", } } function model_variant_cache_dir(variant: ModelVariant) -> string { match (variant) { - ModelVariant.Fp16 => "parakeet-tdt-0.6b-v3-fp16-f88260fa", - ModelVariant.Int8 => "parakeet-tdt-0.6b-v3-int8-f88260fa", + ModelVariant.Fp32 => "parakeet-unified-en-0.6b-fp32-09e90603", + ModelVariant.Int8 => "parakeet-unified-en-0.6b-int8-09e90603", } } @@ -95,8 +107,8 @@ function verify_sha256(path: string, expected: string) -> bool { } function download_model_asset(model_dir: string, asset: ModelAsset) -> null { - let repository = "ysdede/parakeet-tdt-0.6b-v3-onnx"; - let revision = "f88260fa0777fe0868dda6df85d1a98f012a4a7a"; + let repository = "bobNight/parakeet-unified-en-0.6b-onnx"; + let revision = "09e9060322d99c5f070010724786e6ee090fd51d"; let destination = join_path(model_dir, asset.local_name); let part = `${destination}.part`; let url = `https://huggingface.co/${repository}/resolve/${revision}/${asset.remote_name}`; @@ -141,8 +153,8 @@ function prepare_model_variant(variant: ModelVariant) -> string { } download_model_asset(model_dir, asset); } - let repository = "ysdede/parakeet-tdt-0.6b-v3-onnx"; - let revision = "f88260fa0777fe0868dda6df85d1a98f012a4a7a"; + let repository = "bobNight/parakeet-unified-en-0.6b-onnx"; + let revision = "09e9060322d99c5f070010724786e6ee090fd51d"; atomic_write(marker, `${repository}@${revision} ${model_variant_name(variant)}\n`); model_dir } @@ -162,7 +174,7 @@ function load_cuda_model( variant: ModelVariant, ) -> null throws baml.errors.HostCallable, ) -> null { - native_load_model(prepare_model_variant(ModelVariant.Fp16), ModelVariant.Fp16) + native_load_model(prepare_model_variant(ModelVariant.Fp32), ModelVariant.Fp32) } function try_load_cuda_model( @@ -202,13 +214,22 @@ function initialize_model( native_load_model(prepare_model_variant(ModelVariant.Int8), ModelVariant.Int8) } -test "model variants use the filenames expected by Parakeet" { - for (let variant in [ModelVariant.Fp16, ModelVariant.Int8]) { - assert.equal( - model_assets(variant).map((asset) -> { - asset.local_name - }), - ["encoder-model.onnx", "decoder_joint-model.onnx", "vocab.txt"], - ) - } +test "model variants use the filenames expected by Parakeet Unified" { + assert.equal( + model_assets(ModelVariant.Fp32).map((asset) -> { + asset.local_name + }), + ["encoder.onnx", "encoder.onnx.data", "decoder_joint.onnx", "tokenizer.model"], + ); + assert.equal( + model_assets(ModelVariant.Int8).map((asset) -> { + asset.local_name + }), + [ + "encoder.int8.onnx", + "encoder.int8.onnx.data", + "decoder_joint.int8.onnx", + "tokenizer.model", + ], + ) } diff --git a/install.sh b/install.sh index d2d70c2..4845e8e 100755 --- a/install.sh +++ b/install.sh @@ -148,4 +148,4 @@ echo "Caching the BAML runtime..." "${lw_target}" sway-cancel echo "Installed ${lw_target}" -echo "Run 'lw preload' to select the best available runtime and load Parakeet." +echo "Run 'lw preload' to select the best available runtime and load Parakeet Unified." diff --git a/src/model.rs b/src/model.rs index 6311338..82de24c 100644 --- a/src/model.rs +++ b/src/model.rs @@ -5,12 +5,12 @@ use std::time::Instant; use anyhow::{Context, Result, bail}; use fs2::FileExt; -use parakeet_rs::{ExecutionConfig, ParakeetTDT, TimestampMode, Transcriber}; +use parakeet_rs::{ExecutionConfig, ParakeetUnified, TimestampMode, Transcriber}; use crate::{paths, runtime}; struct Model { - inner: ParakeetTDT, + inner: ParakeetUnified, } #[derive(Default)] @@ -49,7 +49,7 @@ impl ModelHost { .map_err(|_| anyhow::anyhow!("model lock holder was poisoned"))? .is_none() { - bail!("refusing to load Parakeet without the per-user model lock") + bail!("refusing to load Parakeet Unified without the per-user model lock") } let mut slot = self .model @@ -59,22 +59,21 @@ impl ModelHost { bail!("the resident model is already loaded") } let (variant_name, device, config) = match variant { - baml_sdk::ModelVariant::Fp16 => { + baml_sdk::ModelVariant::Fp32 => { runtime::prepare_cuda()?; - ("FP16", "CUDA", strict_cuda_config()) + ("FP32", "CUDA", strict_cuda_config()) } baml_sdk::ModelVariant::Int8 => ("INT8", "CPU", ExecutionConfig::new()), }; let started = Instant::now(); - let inner = ParakeetTDT::from_pretrained(Path::new(&model_dir), Some(config)).with_context( - || { + let inner = ParakeetUnified::from_pretrained(Path::new(&model_dir), Some(config)) + .with_context(|| { format!( - "failed to load Parakeet {variant_name} with the {device} execution provider from {model_dir}" + "failed to load Parakeet Unified {variant_name} with the {device} execution provider from {model_dir}" ) - }, - )?; + })?; eprintln!( - "Parakeet {variant_name} loaded on {device} in {:.2?}", + "Parakeet Unified {variant_name} loaded on {device} in {:.2?}", started.elapsed() ); *slot = Some(Model { inner }); @@ -88,10 +87,12 @@ impl ModelHost { .map_err(|_| anyhow::anyhow!("model holder was poisoned"))?; let model = slot.as_mut().context("resident model is not loaded")?; let started = Instant::now(); - let result = model - .inner - .transcribe_file(Path::new(&audio_path), Some(TimestampMode::Sentences)) - .with_context(|| format!("failed to transcribe {audio_path}"))?; + let result = Transcriber::transcribe_file( + &mut model.inner, + Path::new(&audio_path), + Some(TimestampMode::Sentences), + ) + .with_context(|| format!("failed to transcribe {audio_path}"))?; eprintln!("transcribed {audio_path} in {:.2?}", started.elapsed()); Ok(result.text.trim().to_owned()) }