feat: progress emitter (#237)

This commit is contained in:
RainbowBird
2025-06-28 13:34:17 +08:00
committed by GitHub
parent 7f33a7aca7
commit 09d2ca7abf
5 changed files with 360 additions and 282 deletions
+14 -7
View File
@@ -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))
}
}
+4 -1
View File
@@ -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