Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
6 changes: 3 additions & 3 deletions README.md
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand Down
11 changes: 9 additions & 2 deletions baml_src/app.baml
Original file line number Diff line number Diff line change
Expand Up @@ -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
},
Expand Down Expand Up @@ -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")
Expand Down
105 changes: 63 additions & 42 deletions baml_src/model.baml
Original file line number Diff line number Diff line change
@@ -1,5 +1,5 @@
enum ModelVariant {
Fp16,
Fp32,
Int8,
}

Expand All @@ -11,61 +11,73 @@ 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,
]
},
}
}

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",
}
}

Expand Down Expand Up @@ -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}`;
Expand Down Expand Up @@ -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
}
Expand All @@ -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(
Expand Down Expand Up @@ -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",
],
)
}
2 changes: 1 addition & 1 deletion install.sh
Original file line number Diff line number Diff line change
Expand Up @@ -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."
31 changes: 16 additions & 15 deletions src/model.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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)]
Expand Down Expand Up @@ -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
Expand All @@ -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 });
Expand All @@ -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())
}
Expand Down