feat: progress emitter (#237)
This commit is contained in:
@@ -97,7 +97,9 @@ async fn stop_click_through(window: tauri::Window) -> Result<(), String> {
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn load_whisper_model() -> Result<whisper::WhisperProcessor, anyhow::Error> {
|
||||
fn load_whisper_model(window: tauri::Window) -> anyhow::Result<()> {
|
||||
let progress_manager = crate::whisper::progress::ModelLoadProgressEmitterManager::new(window);
|
||||
|
||||
// Determine device to use
|
||||
let device = if candle_core::utils::cuda_is_available() {
|
||||
candle_core::Device::new_cuda(0)?
|
||||
@@ -109,12 +111,19 @@ fn load_whisper_model() -> Result<whisper::WhisperProcessor, anyhow::Error> {
|
||||
|
||||
info!("Using device: {device:?}");
|
||||
|
||||
let whisper_model = whisper::WhichWhisperModel::Tiny;
|
||||
let whisper_model = whisper::whisper::WhichWhisperModel::Tiny;
|
||||
|
||||
info!("Loading whisper model: {:?}", whisper_model);
|
||||
|
||||
let model = whisper::WhisperProcessor::new(whisper_model, device.clone())?;
|
||||
Ok(model)
|
||||
let _ = whisper::whisper::WhisperProcessor::new(whisper_model, device.clone(), progress_manager)?;
|
||||
Ok(())
|
||||
}
|
||||
|
||||
#[tauri::command]
|
||||
async fn load_models(window: tauri::Window) -> Result<(), String> {
|
||||
println!("load_models");
|
||||
load_whisper_model(window).unwrap();
|
||||
Ok(())
|
||||
}
|
||||
|
||||
#[cfg_attr(mobile, tauri::mobile_entry_point)]
|
||||
@@ -162,9 +171,6 @@ pub fn run() {
|
||||
)?;
|
||||
}
|
||||
|
||||
// Load whisper model
|
||||
load_whisper_model()?;
|
||||
|
||||
// TODO: i18n
|
||||
let quit_item = MenuItem::with_id(app, "quit", "Quit", true, None::<&str>)?;
|
||||
let settings_item = MenuItem::with_id(app, "settings", "Settings", true, None::<&str>)?;
|
||||
@@ -231,6 +237,7 @@ pub fn run() {
|
||||
stop_monitor,
|
||||
start_click_through,
|
||||
stop_click_through,
|
||||
load_models,
|
||||
])
|
||||
.build(tauri::generate_context!())
|
||||
.expect("error while building tauri application")
|
||||
|
||||
@@ -1,274 +1,2 @@
|
||||
// https://github.com/proj-airi/candle-examples/blob/3a6783a4788333e7a117f4f8548f02b2fbfec2ed/apps/silero-vad-whisper-realtime/src/whisper.rs
|
||||
use anyhow::Result;
|
||||
use byteorder::{ByteOrder, LittleEndian};
|
||||
use candle_core::{Device, IndexOp, Tensor};
|
||||
use candle_nn::VarBuilder;
|
||||
use candle_transformers::models::whisper::{self as whisper_model, Config, audio};
|
||||
use clap::ValueEnum;
|
||||
use hf_hub::{
|
||||
Repo,
|
||||
RepoType,
|
||||
api::sync::{Api, ApiBuilder},
|
||||
};
|
||||
use tokenizers::Tokenizer;
|
||||
|
||||
pub enum WhisperModel {
|
||||
Normal(whisper_model::model::Whisper),
|
||||
}
|
||||
|
||||
impl WhisperModel {
|
||||
pub fn encoder_forward(
|
||||
&mut self,
|
||||
x: &Tensor,
|
||||
flush: bool,
|
||||
) -> candle_core::Result<Tensor> {
|
||||
match self {
|
||||
Self::Normal(model) => model.encoder.forward(x, flush),
|
||||
}
|
||||
}
|
||||
|
||||
pub fn decoder_forward(
|
||||
&mut self,
|
||||
x: &Tensor,
|
||||
encoder_out: &Tensor,
|
||||
flush: bool,
|
||||
) -> candle_core::Result<Tensor> {
|
||||
match self {
|
||||
Self::Normal(model) => model.decoder.forward(x, encoder_out, flush),
|
||||
}
|
||||
}
|
||||
|
||||
pub fn decoder_final_linear(
|
||||
&self,
|
||||
x: &Tensor,
|
||||
) -> candle_core::Result<Tensor> {
|
||||
match self {
|
||||
Self::Normal(model) => model.decoder.final_linear(x),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone, Copy, Debug, PartialEq, Eq, ValueEnum)]
|
||||
pub enum WhichWhisperModel {
|
||||
Tiny,
|
||||
#[value(name = "tiny.en")]
|
||||
TinyEn,
|
||||
Base,
|
||||
#[value(name = "base.en")]
|
||||
BaseEn,
|
||||
Small,
|
||||
#[value(name = "small.en")]
|
||||
SmallEn,
|
||||
Medium,
|
||||
#[value(name = "medium.en")]
|
||||
MediumEn,
|
||||
Large,
|
||||
LargeV2,
|
||||
LargeV3,
|
||||
LargeV3Turbo,
|
||||
#[value(name = "distil-medium.en")]
|
||||
DistilMediumEn,
|
||||
#[value(name = "distil-large-v2")]
|
||||
DistilLargeV2,
|
||||
}
|
||||
|
||||
impl WhichWhisperModel {
|
||||
pub const fn model_and_revision(self) -> (&'static str, &'static str) {
|
||||
match self {
|
||||
Self::Tiny => ("openai/whisper-tiny", "main"),
|
||||
Self::TinyEn => ("openai/whisper-tiny.en", "refs/pr/15"),
|
||||
Self::Base => ("openai/whisper-base", "refs/pr/22"),
|
||||
Self::BaseEn => ("openai/whisper-base.en", "refs/pr/13"),
|
||||
Self::Small => ("openai/whisper-small", "main"),
|
||||
Self::SmallEn => ("openai/whisper-small.en", "refs/pr/10"),
|
||||
Self::Medium => ("openai/whisper-medium", "main"),
|
||||
Self::MediumEn => ("openai/whisper-medium.en", "main"),
|
||||
Self::Large => ("openai/whisper-large", "refs/pr/36"),
|
||||
Self::LargeV2 => ("openai/whisper-large-v2", "refs/pr/57"),
|
||||
Self::LargeV3 => ("openai/whisper-large-v3", "main"),
|
||||
Self::LargeV3Turbo => ("openai/whisper-large-v3-turbo", "main"),
|
||||
Self::DistilMediumEn => ("distil-whisper/distil-medium.en", "main"),
|
||||
Self::DistilLargeV2 => ("distil-whisper/distil-large-v2", "main"),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub struct WhisperProcessor {
|
||||
pub model: WhisperModel,
|
||||
pub tokenizer: Tokenizer,
|
||||
pub config: Config,
|
||||
pub mel_filters: Vec<f32>,
|
||||
pub device: Device,
|
||||
}
|
||||
|
||||
pub struct WhisperProgressEmitter {
|
||||
filename: String,
|
||||
}
|
||||
|
||||
impl WhisperProgressEmitter {
|
||||
fn new(filename: &str) -> Self {
|
||||
Self {
|
||||
filename: filename.to_string(),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl hf_hub::api::Progress for WhisperProgressEmitter {
|
||||
fn init(
|
||||
&mut self,
|
||||
size: usize,
|
||||
filename: &str,
|
||||
) {
|
||||
todo!();
|
||||
}
|
||||
|
||||
fn update(
|
||||
&mut self,
|
||||
size: usize,
|
||||
) {
|
||||
todo!();
|
||||
}
|
||||
|
||||
fn finish(&mut self) {
|
||||
todo!();
|
||||
}
|
||||
}
|
||||
|
||||
impl WhisperProcessor {
|
||||
pub fn new(
|
||||
model: WhichWhisperModel,
|
||||
device: Device,
|
||||
) -> Result<Self> {
|
||||
// Load the Whisper model based on the provided model type
|
||||
let api = ApiBuilder::new().with_progress(false).build()?;
|
||||
let (model_id, revision) = model.model_and_revision();
|
||||
let repo = api.repo(Repo::with_revision(
|
||||
model_id.to_string(),
|
||||
RepoType::Model,
|
||||
revision.to_string(),
|
||||
));
|
||||
|
||||
let config_filename =
|
||||
repo.download_with_progress("config.json", WhisperProgressEmitter::new("config.json"))?;
|
||||
// TODO: for RainbowBird
|
||||
let tokenizer_filename = repo.get("tokenizer.json")?;
|
||||
// TODO: for RainbowBird
|
||||
let model_filename = repo.get("model.safetensors")?;
|
||||
|
||||
let config: Config = serde_json::from_str(&std::fs::read_to_string(config_filename)?)?;
|
||||
let tokenizer = Tokenizer::from_file(tokenizer_filename).map_err(anyhow::Error::msg)?;
|
||||
|
||||
println!("Loading Whisper model from: {:?}", model_filename.display());
|
||||
|
||||
// SAFETY: This is safe because we are using a mmaped file and the safetensors library guarantees that the data is valid.
|
||||
let var_builder = unsafe {
|
||||
VarBuilder::from_mmaped_safetensors(&[model_filename], whisper_model::DTYPE, &device)?
|
||||
};
|
||||
|
||||
let model = WhisperModel::Normal(whisper_model::model::Whisper::load(
|
||||
&var_builder,
|
||||
config.clone(),
|
||||
)?);
|
||||
|
||||
let mel_bytes = match config.num_mel_bins {
|
||||
80 => include_bytes!("./melfilters.bytes").as_slice(),
|
||||
128 => include_bytes!("./melfilters128.bytes").as_slice(),
|
||||
num_mel_bins => anyhow::bail!("Unsupported number of mel bins: {}", num_mel_bins),
|
||||
};
|
||||
|
||||
let mut mel_filters = vec![0f32; mel_bytes.len() / 4];
|
||||
<LittleEndian as ByteOrder>::read_f32_into(mel_bytes, &mut mel_filters);
|
||||
|
||||
Ok(Self {
|
||||
model,
|
||||
tokenizer,
|
||||
config,
|
||||
mel_filters,
|
||||
device,
|
||||
})
|
||||
}
|
||||
|
||||
pub fn transcribe(
|
||||
&mut self,
|
||||
audio: &[f32],
|
||||
) -> Result<String> {
|
||||
// Convert PCM to mel spectrogram
|
||||
let mel = audio::pcm_to_mel(&self.config, audio, &self.mel_filters);
|
||||
let mel_len = mel.len();
|
||||
let mel = Tensor::from_vec(
|
||||
mel,
|
||||
(
|
||||
1,
|
||||
self.config.num_mel_bins,
|
||||
mel_len / self.config.num_mel_bins,
|
||||
),
|
||||
&self.device,
|
||||
)?;
|
||||
|
||||
// Run inference
|
||||
let audio_features = self.model.encoder_forward(&mel, true)?;
|
||||
// Simple greedy decoding
|
||||
let tokens = self.decode_greedy(&audio_features)?;
|
||||
|
||||
let text = self
|
||||
.tokenizer
|
||||
.decode(&tokens, true)
|
||||
.map_err(anyhow::Error::msg)?;
|
||||
|
||||
Ok(text)
|
||||
}
|
||||
|
||||
fn decode_greedy(
|
||||
&mut self,
|
||||
audio_features: &Tensor,
|
||||
) -> Result<Vec<u32>> {
|
||||
let mut tokens = vec![
|
||||
self.token_id(whisper_model::SOT_TOKEN)?,
|
||||
self.token_id(whisper_model::TRANSCRIBE_TOKEN)?,
|
||||
self.token_id(whisper_model::NO_TIMESTAMPS_TOKEN)?,
|
||||
];
|
||||
|
||||
let max_len = 50; // Short sequence for real-time processing
|
||||
|
||||
for i in 0..max_len {
|
||||
let tokens_t = Tensor::new(tokens.as_slice(), &self.device)?.unsqueeze(0)?;
|
||||
let ys = self
|
||||
.model
|
||||
.decoder_forward(&tokens_t, audio_features, i == 0)?;
|
||||
|
||||
let (_, seq_len, _) = ys.dims3()?;
|
||||
let logits = self
|
||||
.model
|
||||
.decoder_final_linear(&ys.i((..1, seq_len - 1..))?)?
|
||||
.i(0)?
|
||||
.i(0)?;
|
||||
|
||||
// Get most likely token
|
||||
let logits_v: Vec<f32> = logits.to_vec1()?;
|
||||
let next_token = logits_v
|
||||
.iter()
|
||||
.enumerate()
|
||||
.max_by(|(_, a), (_, b)| a.total_cmp(b))
|
||||
.map(|(i, _)| u32::try_from(i).unwrap())
|
||||
.unwrap();
|
||||
|
||||
if next_token == self.token_id(whisper_model::EOT_TOKEN)? {
|
||||
break;
|
||||
}
|
||||
|
||||
tokens.push(next_token);
|
||||
}
|
||||
|
||||
Ok(tokens)
|
||||
}
|
||||
|
||||
fn token_id(
|
||||
&self,
|
||||
token: &str,
|
||||
) -> Result<u32> {
|
||||
self
|
||||
.tokenizer
|
||||
.token_to_id(token)
|
||||
.ok_or_else(|| anyhow::anyhow!("Token not found: {}", token))
|
||||
}
|
||||
}
|
||||
pub mod progress;
|
||||
pub mod whisper;
|
||||
|
||||
@@ -0,0 +1,93 @@
|
||||
use log::error;
|
||||
use tauri::Emitter;
|
||||
|
||||
#[derive(Clone)]
|
||||
pub struct ModelLoadProgressEmitterManager {
|
||||
window: tauri::Window,
|
||||
}
|
||||
|
||||
impl ModelLoadProgressEmitterManager {
|
||||
pub fn new(window: tauri::Window) -> Self {
|
||||
Self { window }
|
||||
}
|
||||
|
||||
pub fn new_processor(
|
||||
self,
|
||||
filename: &str,
|
||||
) -> ModelLoadProgressEmitter {
|
||||
ModelLoadProgressEmitter::new(self.window, filename)
|
||||
}
|
||||
}
|
||||
|
||||
pub struct ModelLoadProgressEmitter {
|
||||
filename: String,
|
||||
total_size: usize,
|
||||
progress: f32,
|
||||
window: tauri::Window,
|
||||
}
|
||||
|
||||
impl ModelLoadProgressEmitter {
|
||||
fn new(
|
||||
window: tauri::Window,
|
||||
filename: &str,
|
||||
) -> Self {
|
||||
Self {
|
||||
filename: filename.to_string(),
|
||||
total_size: 0,
|
||||
progress: 0.0,
|
||||
window,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl hf_hub::api::Progress for ModelLoadProgressEmitter {
|
||||
fn init(
|
||||
&mut self,
|
||||
size: usize,
|
||||
_: &str,
|
||||
) {
|
||||
self.total_size = size;
|
||||
self.progress = 0.0;
|
||||
self
|
||||
.window
|
||||
.emit(
|
||||
"tauri-app:model-load-progress",
|
||||
(self.filename.clone(), self.progress),
|
||||
)
|
||||
.map_err(|err| {
|
||||
error!("Failed to emit model-load-progress: {:?}", err);
|
||||
})
|
||||
.unwrap();
|
||||
}
|
||||
|
||||
fn update(
|
||||
&mut self,
|
||||
size: usize,
|
||||
) {
|
||||
self.progress = size as f32 / self.total_size as f32 * 100.0;
|
||||
self
|
||||
.window
|
||||
.emit(
|
||||
"tauri-app:model-load-progress",
|
||||
(self.filename.clone(), self.progress),
|
||||
)
|
||||
.map_err(|err| {
|
||||
error!("Failed to emit model-load-progress: {:?}", err);
|
||||
})
|
||||
.unwrap();
|
||||
}
|
||||
|
||||
fn finish(&mut self) {
|
||||
self.progress = 100.0;
|
||||
self
|
||||
.window
|
||||
.emit(
|
||||
"tauri-app:model-load-progress",
|
||||
(self.filename.clone(), self.progress),
|
||||
)
|
||||
.map_err(|err| {
|
||||
error!("Failed to emit model-load-progress: {:?}", err);
|
||||
})
|
||||
.unwrap();
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,247 @@
|
||||
// https://github.com/proj-airi/candle-examples/blob/3a6783a4788333e7a117f4f8548f02b2fbfec2ed/apps/silero-vad-whisper-realtime/src/whisper.rs
|
||||
use anyhow::Result;
|
||||
use byteorder::{ByteOrder, LittleEndian};
|
||||
use candle_core::{Device, IndexOp, Tensor};
|
||||
use candle_nn::VarBuilder;
|
||||
use candle_transformers::models::whisper::{self as whisper_model, Config, audio};
|
||||
use clap::ValueEnum;
|
||||
use hf_hub::{Repo, RepoType, api::sync::ApiBuilder};
|
||||
use tokenizers::Tokenizer;
|
||||
|
||||
use crate::whisper::progress;
|
||||
|
||||
pub enum WhisperModel {
|
||||
Normal(whisper_model::model::Whisper),
|
||||
}
|
||||
|
||||
impl WhisperModel {
|
||||
pub fn encoder_forward(
|
||||
&mut self,
|
||||
x: &Tensor,
|
||||
flush: bool,
|
||||
) -> candle_core::Result<Tensor> {
|
||||
match self {
|
||||
Self::Normal(model) => model.encoder.forward(x, flush),
|
||||
}
|
||||
}
|
||||
|
||||
pub fn decoder_forward(
|
||||
&mut self,
|
||||
x: &Tensor,
|
||||
encoder_out: &Tensor,
|
||||
flush: bool,
|
||||
) -> candle_core::Result<Tensor> {
|
||||
match self {
|
||||
Self::Normal(model) => model.decoder.forward(x, encoder_out, flush),
|
||||
}
|
||||
}
|
||||
|
||||
pub fn decoder_final_linear(
|
||||
&self,
|
||||
x: &Tensor,
|
||||
) -> candle_core::Result<Tensor> {
|
||||
match self {
|
||||
Self::Normal(model) => model.decoder.final_linear(x),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone, Copy, Debug, PartialEq, Eq, ValueEnum)]
|
||||
pub enum WhichWhisperModel {
|
||||
Tiny,
|
||||
#[value(name = "tiny.en")]
|
||||
TinyEn,
|
||||
Base,
|
||||
#[value(name = "base.en")]
|
||||
BaseEn,
|
||||
Small,
|
||||
#[value(name = "small.en")]
|
||||
SmallEn,
|
||||
Medium,
|
||||
#[value(name = "medium.en")]
|
||||
MediumEn,
|
||||
Large,
|
||||
LargeV2,
|
||||
LargeV3,
|
||||
LargeV3Turbo,
|
||||
#[value(name = "distil-medium.en")]
|
||||
DistilMediumEn,
|
||||
#[value(name = "distil-large-v2")]
|
||||
DistilLargeV2,
|
||||
}
|
||||
|
||||
impl WhichWhisperModel {
|
||||
pub const fn model_and_revision(self) -> (&'static str, &'static str) {
|
||||
match self {
|
||||
Self::Tiny => ("openai/whisper-tiny", "main"),
|
||||
Self::TinyEn => ("openai/whisper-tiny.en", "refs/pr/15"),
|
||||
Self::Base => ("openai/whisper-base", "refs/pr/22"),
|
||||
Self::BaseEn => ("openai/whisper-base.en", "refs/pr/13"),
|
||||
Self::Small => ("openai/whisper-small", "main"),
|
||||
Self::SmallEn => ("openai/whisper-small.en", "refs/pr/10"),
|
||||
Self::Medium => ("openai/whisper-medium", "main"),
|
||||
Self::MediumEn => ("openai/whisper-medium.en", "main"),
|
||||
Self::Large => ("openai/whisper-large", "refs/pr/36"),
|
||||
Self::LargeV2 => ("openai/whisper-large-v2", "refs/pr/57"),
|
||||
Self::LargeV3 => ("openai/whisper-large-v3", "main"),
|
||||
Self::LargeV3Turbo => ("openai/whisper-large-v3-turbo", "main"),
|
||||
Self::DistilMediumEn => ("distil-whisper/distil-medium.en", "main"),
|
||||
Self::DistilLargeV2 => ("distil-whisper/distil-large-v2", "main"),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub struct WhisperProcessor {
|
||||
pub model: WhisperModel,
|
||||
pub tokenizer: Tokenizer,
|
||||
pub config: Config,
|
||||
pub mel_filters: Vec<f32>,
|
||||
pub device: Device,
|
||||
}
|
||||
|
||||
impl WhisperProcessor {
|
||||
pub fn new(
|
||||
model: WhichWhisperModel,
|
||||
device: Device,
|
||||
manager: progress::ModelLoadProgressEmitterManager,
|
||||
) -> Result<Self> {
|
||||
// Load the Whisper model based on the provided model type
|
||||
let api = ApiBuilder::new().with_progress(false).build()?;
|
||||
let (model_id, revision) = model.model_and_revision();
|
||||
let repo = api.repo(Repo::with_revision(
|
||||
model_id.to_string(),
|
||||
RepoType::Model,
|
||||
revision.to_string(),
|
||||
));
|
||||
|
||||
let config_filename =
|
||||
repo.download_with_progress("config.json", manager.clone().new_processor("config.json"))?;
|
||||
println!("config_filename: {:?}", config_filename.display());
|
||||
let tokenizer_filename = repo.download_with_progress(
|
||||
"tokenizer.json",
|
||||
manager.clone().new_processor("tokenizer.json"),
|
||||
)?;
|
||||
println!("tokenizer_filename: {:?}", tokenizer_filename.display());
|
||||
let model_filename = repo.download_with_progress(
|
||||
"model.safetensors",
|
||||
manager.clone().new_processor("model.safetensors"),
|
||||
)?;
|
||||
println!("model_filename: {:?}", model_filename.display());
|
||||
|
||||
let config: Config = serde_json::from_str(&std::fs::read_to_string(config_filename)?)?;
|
||||
let tokenizer = Tokenizer::from_file(tokenizer_filename).map_err(anyhow::Error::msg)?;
|
||||
|
||||
println!("Loading Whisper model from: {:?}", model_filename.display());
|
||||
|
||||
// SAFETY: This is safe because we are using a mmaped file and the safetensors library guarantees that the data is valid.
|
||||
let var_builder = unsafe {
|
||||
VarBuilder::from_mmaped_safetensors(&[model_filename], whisper_model::DTYPE, &device)?
|
||||
};
|
||||
|
||||
let model = WhisperModel::Normal(whisper_model::model::Whisper::load(
|
||||
&var_builder,
|
||||
config.clone(),
|
||||
)?);
|
||||
|
||||
let mel_bytes = match config.num_mel_bins {
|
||||
80 => include_bytes!("./melfilters.bytes").as_slice(),
|
||||
128 => include_bytes!("./melfilters128.bytes").as_slice(),
|
||||
num_mel_bins => anyhow::bail!("Unsupported number of mel bins: {}", num_mel_bins),
|
||||
};
|
||||
|
||||
let mut mel_filters = vec![0f32; mel_bytes.len() / 4];
|
||||
<LittleEndian as ByteOrder>::read_f32_into(mel_bytes, &mut mel_filters);
|
||||
|
||||
Ok(Self {
|
||||
model,
|
||||
tokenizer,
|
||||
config,
|
||||
mel_filters,
|
||||
device,
|
||||
})
|
||||
}
|
||||
|
||||
pub fn transcribe(
|
||||
&mut self,
|
||||
audio: &[f32],
|
||||
) -> Result<String> {
|
||||
// Convert PCM to mel spectrogram
|
||||
let mel = audio::pcm_to_mel(&self.config, audio, &self.mel_filters);
|
||||
let mel_len = mel.len();
|
||||
let mel = Tensor::from_vec(
|
||||
mel,
|
||||
(
|
||||
1,
|
||||
self.config.num_mel_bins,
|
||||
mel_len / self.config.num_mel_bins,
|
||||
),
|
||||
&self.device,
|
||||
)?;
|
||||
|
||||
// Run inference
|
||||
let audio_features = self.model.encoder_forward(&mel, true)?;
|
||||
// Simple greedy decoding
|
||||
let tokens = self.decode_greedy(&audio_features)?;
|
||||
|
||||
let text = self
|
||||
.tokenizer
|
||||
.decode(&tokens, true)
|
||||
.map_err(anyhow::Error::msg)?;
|
||||
|
||||
Ok(text)
|
||||
}
|
||||
|
||||
fn decode_greedy(
|
||||
&mut self,
|
||||
audio_features: &Tensor,
|
||||
) -> Result<Vec<u32>> {
|
||||
let mut tokens = vec![
|
||||
self.token_id(whisper_model::SOT_TOKEN)?,
|
||||
self.token_id(whisper_model::TRANSCRIBE_TOKEN)?,
|
||||
self.token_id(whisper_model::NO_TIMESTAMPS_TOKEN)?,
|
||||
];
|
||||
|
||||
let max_len = 50; // Short sequence for real-time processing
|
||||
|
||||
for i in 0..max_len {
|
||||
let tokens_t = Tensor::new(tokens.as_slice(), &self.device)?.unsqueeze(0)?;
|
||||
let ys = self
|
||||
.model
|
||||
.decoder_forward(&tokens_t, audio_features, i == 0)?;
|
||||
|
||||
let (_, seq_len, _) = ys.dims3()?;
|
||||
let logits = self
|
||||
.model
|
||||
.decoder_final_linear(&ys.i((..1, seq_len - 1..))?)?
|
||||
.i(0)?
|
||||
.i(0)?;
|
||||
|
||||
// Get most likely token
|
||||
let logits_v: Vec<f32> = logits.to_vec1()?;
|
||||
let next_token = logits_v
|
||||
.iter()
|
||||
.enumerate()
|
||||
.max_by(|(_, a), (_, b)| a.total_cmp(b))
|
||||
.map(|(i, _)| u32::try_from(i).unwrap())
|
||||
.unwrap();
|
||||
|
||||
if next_token == self.token_id(whisper_model::EOT_TOKEN)? {
|
||||
break;
|
||||
}
|
||||
|
||||
tokens.push(next_token);
|
||||
}
|
||||
|
||||
Ok(tokens)
|
||||
}
|
||||
|
||||
fn token_id(
|
||||
&self,
|
||||
token: &str,
|
||||
) -> Result<u32> {
|
||||
self
|
||||
.tokenizer
|
||||
.token_to_id(token)
|
||||
.ok_or_else(|| anyhow::anyhow!("Token not found: {}", token))
|
||||
}
|
||||
}
|
||||
@@ -1,6 +1,5 @@
|
||||
<script setup lang="ts">
|
||||
import { WidgetStage } from '@proj-airi/stage-ui/components'
|
||||
import { useAppRuntime } from '@proj-airi/stage-ui/composables'
|
||||
import { useMcpStore } from '@proj-airi/stage-ui/stores'
|
||||
import { connectServer } from '@proj-airi/tauri-plugin-mcp'
|
||||
import { invoke } from '@tauri-apps/api/core'
|
||||
@@ -8,6 +7,7 @@ import { listen } from '@tauri-apps/api/event'
|
||||
import { storeToRefs } from 'pinia'
|
||||
import { computed, onMounted, onUnmounted, ref } from 'vue'
|
||||
|
||||
import { useAppRuntime } from '../composables/runtime'
|
||||
import { useWindowShortcuts } from '../composables/window-shortcuts'
|
||||
import { useWindowControlStore } from '../stores/window-controls'
|
||||
import { WindowControlMode } from '../types/window-controls'
|
||||
@@ -96,6 +96,9 @@ function onTauriPositionCursorAndWindowFrameEvent(event: { payload: [Point, Wind
|
||||
onMounted(async () => {
|
||||
// Listen for click-through state changes
|
||||
unListenFuncs.push(await listen('tauri-app:window-click-through:position-cursor-and-window-frame', onTauriPositionCursorAndWindowFrameEvent))
|
||||
invoke('load_models')
|
||||
// eslint-disable-next-line no-console
|
||||
unListenFuncs.push(await listen('tauri-app:model-load-progress', console.log))
|
||||
|
||||
if (connected.value)
|
||||
return
|
||||
|
||||
Reference in New Issue
Block a user