From 093c297ff969995392421c659f6ccf58350160c5 Mon Sep 17 00:00:00 2001 From: Eason WaveKat Date: Sun, 17 May 2026 10:56:00 +1200 Subject: [PATCH] feat: backend-agnostic HF download helper with byte progress MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Adds `wavekat_asr::download::download_files_with_progress(repo_id, files, dest_dir, on_progress)` and `DownloadProgress` — a backend- agnostic helper for streaming a fixed list of files from HuggingFace into a caller-chosen directory while a `FnMut` closure receives byte- level progress updates. Gated behind a new `download` Cargo feature that the `sherpa-onnx` feature now turns on transitively, so any future backend whose model weights live on HF Hub (Whisper, Qwen3-ASR-local, etc.) can enable `download` and reuse the same helper instead of reimplementing it. The sherpa-onnx backend keeps its convenience wrapper: `backends::sherpa_onnx::download_preset_with_progress(preset, dest_dir, on_progress)`. It's now a 4-line shim over the generic helper that turns a `ModelPreset` into a `(model_id, &[files])` pair via the also-new `ModelPreset::files()`. Behavior notes preserved from the first iteration: - Streams byte-level progress per file (file_index/file_count, bytes_done, bytes_total). - Copies each file into a caller-chosen `dest_dir` rather than only populating hf-hub's cache, so consumers can ship a "model lives here" path users can introspect, move, or delete. - Synthesizes a final 100% event in `finish()` if hf-hub's last `update` short-counted, so a UI progress bar can't get stuck at 99%. Implementation note on the closure-capture trick: hf-hub's `Progress` trait wants ownership of the reporter per `download_with_progress` call, so the adapter holds a `&mut F` and is rebuilt per file — that lets one `FnMut(DownloadProgress)` closure span all files in a batch without `Rc>`. Tests (all without network access): - `download::tests::callback_progress_reports_init_update_and_finish` drives the hf-hub `Progress` bridge directly and asserts the init+update sequence reaches the user closure intact. - `download::tests::callback_progress_finish_synthesizes_completion` covers the short-count safeguard. - `backends::sherpa_onnx::tests::preset_files_transducer_includes_joiner` / `preset_files_paraformer_omits_joiner` cover the preset-shape glue. Verified locally: cargo fmt --check, cargo clippy --all-features -D warnings, cargo test --features sherpa-onnx (7 unit tests + 1 doc-test pass), cargo doc --all-features, cargo check on each of: default features, `--features download` alone, `--features sherpa-onnx`. Step 2 of 3 toward wavekat-asr 0.0.5. The wavekat-core 0.0.11 bump (#8) is the other prerequisite and is merged. Co-Authored-By: Claude Opus 4.7 (1M context) --- crates/wavekat-asr/Cargo.toml | 7 +- .../wavekat-asr/src/backends/sherpa_onnx.rs | 60 +++++ crates/wavekat-asr/src/download.rs | 243 ++++++++++++++++++ crates/wavekat-asr/src/lib.rs | 4 + 4 files changed, 313 insertions(+), 1 deletion(-) create mode 100644 crates/wavekat-asr/src/download.rs diff --git a/crates/wavekat-asr/Cargo.toml b/crates/wavekat-asr/Cargo.toml index 6adfe27..a8c1f72 100644 --- a/crates/wavekat-asr/Cargo.toml +++ b/crates/wavekat-asr/Cargo.toml @@ -14,10 +14,15 @@ exclude = ["CHANGELOG.md"] [features] default = [] +# Generic HuggingFace download helper with byte-level progress +# reporting. Used by any backend whose model files live on HF Hub; +# enabled transitively by such backends so callers usually don't +# turn it on directly. +download = ["dep:hf-hub"] # Local streaming backend wrapping sherpa-onnx around a streaming # Zipformer transducer. Auto-downloads the default bilingual EN+ZH # model from HuggingFace on first use. -sherpa-onnx = ["dep:sherpa-onnx", "dep:hf-hub"] +sherpa-onnx = ["dep:sherpa-onnx", "download"] [dependencies] wavekat-core = "0.0.11" diff --git a/crates/wavekat-asr/src/backends/sherpa_onnx.rs b/crates/wavekat-asr/src/backends/sherpa_onnx.rs index 5406639..3fadfdb 100644 --- a/crates/wavekat-asr/src/backends/sherpa_onnx.rs +++ b/crates/wavekat-asr/src/backends/sherpa_onnx.rs @@ -75,6 +75,22 @@ pub struct ModelPreset { pub tokens: &'static str, } +impl ModelPreset { + /// Files this preset pulls from HuggingFace, in download order: encoder, + /// decoder, joiner (transducer only), tokens. Used by + /// [`download_preset_with_progress`] to drive the per-file callback. + pub fn files(&self) -> Vec<&'static str> { + let mut files = Vec::with_capacity(4); + files.push(self.encoder); + files.push(self.decoder); + if let Some(joiner) = self.joiner { + files.push(joiner); + } + files.push(self.tokens); + files + } +} + /// Bilingual EN+ZH streaming Zipformer (default). Handles mixed-language /// speech but can over-produce Hanzi for English-only audio. pub const BILINGUAL_ZH_EN: ModelPreset = ModelPreset { @@ -470,6 +486,31 @@ fn download_from_hf(config: &SherpaOnnxConfig) -> Result { }) } +/// Re-exported here so callers reaching for the symbol via +/// `backends::sherpa_onnx::DownloadProgress` get the same type they would +/// from the crate root or from a future backend's preset downloader. +pub use crate::download::DownloadProgress; + +/// Download every file `preset` needs from HuggingFace into `dest_dir`, +/// reporting byte progress as it goes. +/// +/// Thin wrapper over [`crate::download::download_files_with_progress`] +/// that knows how to turn a [`ModelPreset`] into a `(repo_id, files)` +/// pair. On success, `dest_dir` contains all files named by +/// [`ModelPreset::files`] and is directly loadable via +/// `SherpaOnnxConfig { model_dir: Some(dest_dir.into()), .. }`. +pub fn download_preset_with_progress( + preset: ModelPreset, + dest_dir: &Path, + on_progress: F, +) -> Result<(), AsrError> +where + F: FnMut(DownloadProgress), +{ + let files = preset.files(); + crate::download::download_files_with_progress(preset.model_id, &files, dest_dir, on_progress) +} + fn path_to_string(path: &Path) -> Result { path.to_str() .map(|s| s.to_string()) @@ -500,6 +541,25 @@ mod tests { assert!(cfg.joiner_filename.is_none()); } + #[test] + fn preset_files_transducer_includes_joiner() { + let files = BILINGUAL_ZH_EN.files(); + assert_eq!(files.len(), 4); + assert_eq!(files[0], BILINGUAL_ZH_EN.encoder); + assert_eq!(files[1], BILINGUAL_ZH_EN.decoder); + assert_eq!(files[2], BILINGUAL_ZH_EN.joiner.unwrap()); + assert_eq!(files[3], BILINGUAL_ZH_EN.tokens); + } + + #[test] + fn preset_files_paraformer_omits_joiner() { + let files = PARAFORMER_ZH.files(); + assert_eq!(files.len(), 3); + assert_eq!(files[0], PARAFORMER_ZH.encoder); + assert_eq!(files[1], PARAFORMER_ZH.decoder); + assert_eq!(files[2], PARAFORMER_ZH.tokens); + } + #[test] fn load_from_dir_errors_on_missing_files() { let cfg = SherpaOnnxConfig::default(); diff --git a/crates/wavekat-asr/src/download.rs b/crates/wavekat-asr/src/download.rs new file mode 100644 index 0000000..4d7a57b --- /dev/null +++ b/crates/wavekat-asr/src/download.rs @@ -0,0 +1,243 @@ +//! Generic HuggingFace download helper with byte-level progress. +//! +//! This module is backend-agnostic: any backend whose model weights live +//! on HuggingFace Hub can call [`download_files_with_progress`] to fetch +//! them into a caller-chosen directory while a `FnMut` closure receives +//! [`DownloadProgress`] updates as bytes arrive. +//! +//! Today only the [`crate::backends::sherpa_onnx`] backend uses this — +//! `sherpa_onnx::download_preset_with_progress` is a thin wrapper that +//! turns a `ModelPreset` into a `(repo_id, &[&str])` call into here. A +//! future Whisper / Qwen3-ASR / etc. backend can do the same. +//! +//! Gated behind the `download` Cargo feature; the `sherpa-onnx` feature +//! turns it on transitively. + +use std::path::Path; + +use crate::AsrError; + +/// Byte-level progress for one file inside a multi-file download. +/// +/// Passed to the callback of [`download_files_with_progress`]. Fires at +/// the start of each file (`bytes_done = 0`), repeatedly during streaming +/// as bytes arrive, and once on completion (`bytes_done == bytes_total`). +#[derive(Debug, Clone)] +pub struct DownloadProgress { + /// Filename inside the HuggingFace repo. + pub file: String, + /// 1-indexed position of `file` within the batch, paired with + /// [`file_count`](Self::file_count) so a UI can render + /// "file 2 of 4". + pub file_index: usize, + /// Total files in the batch. + pub file_count: usize, + /// Bytes downloaded so far for the current file. Resets per file. + pub bytes_done: u64, + /// Total bytes for the current file, reported by HuggingFace before + /// streaming begins. `None` only for the very first call before + /// metadata is known. + pub bytes_total: Option, +} + +/// Download every file in `files` from `repo_id` on HuggingFace Hub into +/// `dest_dir`, reporting byte progress as it goes. +/// +/// On success, `dest_dir` contains every filename in `files` and is +/// directly loadable by a backend that consumes the files from a flat +/// directory. `dest_dir` is created if missing; existing files with the +/// same name are overwritten so a partial previous run retries cleanly. +/// +/// `on_progress` runs synchronously on the calling thread inside the +/// hf-hub download loop — keep it cheap (channel send, atomic store) and +/// don't block on it. +/// +/// hf-hub's own cache (controlled by the `HF_HOME` env var) still gets +/// populated as a side effect; callers that only want the files in +/// `dest_dir` can ignore it. +pub fn download_files_with_progress( + repo_id: &str, + files: &[&str], + dest_dir: &Path, + mut on_progress: F, +) -> Result<(), AsrError> +where + F: FnMut(DownloadProgress), +{ + use hf_hub::api::sync::Api; + + std::fs::create_dir_all(dest_dir).map_err(|e| { + AsrError::Backend(format!( + "creating model directory {}: {e}", + dest_dir.display() + )) + })?; + + let file_count = files.len(); + + let api = Api::new().map_err(|e| AsrError::Backend(format!("hf-hub init failed: {e}")))?; + let repo = api.model(repo_id.to_string()); + + for (idx, name) in files.iter().enumerate() { + let file_index = idx + 1; + + // Fresh adapter per file: it borrows `on_progress` for the + // duration of one `download_with_progress` call, so the next + // iteration can re-borrow it for the next file. + let adapter = CallbackProgress { + file: (*name).to_string(), + file_index, + file_count, + bytes_done: 0, + bytes_total: None, + on_progress: &mut on_progress, + }; + + tracing::debug!( + repo_id, + file = name, + file_index, + file_count, + "fetching from HuggingFace with progress" + ); + + let src = repo + .download_with_progress(name, adapter) + .map_err(|e| AsrError::Backend(format!("hf-hub download of {name} failed: {e}")))?; + + let dest = dest_dir.join(name); + // copy, not rename — the hf-hub blob is shared with its cache; + // the user gets their own copy under `dest_dir` so they can move + // / delete it without breaking the cache. + std::fs::copy(&src, &dest).map_err(|e| { + AsrError::Backend(format!( + "copying {} to {}: {e}", + src.display(), + dest.display() + )) + })?; + } + + Ok(()) +} + +/// Bridge between hf-hub's `Progress` trait and our +/// `FnMut(DownloadProgress)` callback. Borrows the user's closure so a +/// single closure can be reused across every file in the batch. +struct CallbackProgress<'a, F: FnMut(DownloadProgress)> { + file: String, + file_index: usize, + file_count: usize, + bytes_done: u64, + bytes_total: Option, + on_progress: &'a mut F, +} + +impl CallbackProgress<'_, F> { + fn emit(&mut self) { + (self.on_progress)(DownloadProgress { + file: self.file.clone(), + file_index: self.file_index, + file_count: self.file_count, + bytes_done: self.bytes_done, + bytes_total: self.bytes_total, + }); + } +} + +impl hf_hub::api::Progress for CallbackProgress<'_, F> { + fn init(&mut self, size: usize, _filename: &str) { + self.bytes_total = Some(size as u64); + self.bytes_done = 0; + self.emit(); + } + + fn update(&mut self, size: usize) { + // hf-hub passes the chunk size, not the cumulative total. + self.bytes_done = self.bytes_done.saturating_add(size as u64); + self.emit(); + } + + fn finish(&mut self) { + // Make sure the last emitted value is the file's full size; + // some backends short the final `update` call. + if let Some(total) = self.bytes_total { + if self.bytes_done < total { + self.bytes_done = total; + self.emit(); + } + } + } +} + +#[cfg(test)] +mod tests { + use super::*; + + /// Exercises the hf-hub `Progress` bridge directly so we can assert + /// the callback contract without a real network download. + #[test] + fn callback_progress_reports_init_update_and_finish() { + use hf_hub::api::Progress; + + let events = std::cell::RefCell::new(Vec::::new()); + let mut on_progress = |p: DownloadProgress| events.borrow_mut().push(p); + + let mut adapter = CallbackProgress { + file: "encoder.onnx".to_string(), + file_index: 1, + file_count: 4, + bytes_done: 0, + bytes_total: None, + on_progress: &mut on_progress, + }; + + adapter.init(1000, "encoder.onnx"); + adapter.update(400); + adapter.update(600); + adapter.finish(); + + let got = events.into_inner(); + // 1 init + 2 updates = 3 emissions. `finish` should NOT add a + // 4th since bytes_done already equals bytes_total. + assert_eq!(got.len(), 3); + + assert_eq!(got[0].bytes_done, 0); + assert_eq!(got[0].bytes_total, Some(1000)); + assert_eq!(got[0].file_index, 1); + assert_eq!(got[0].file_count, 4); + assert_eq!(got[0].file, "encoder.onnx"); + + assert_eq!(got[1].bytes_done, 400); + assert_eq!(got[2].bytes_done, 1000); + } + + /// If hf-hub's last `update` short-counts (we've seen this on retried + /// downloads), `finish` should synthesize a final 100% event so the UI + /// can't get stuck at 99%. + #[test] + fn callback_progress_finish_synthesizes_completion() { + use hf_hub::api::Progress; + + let events = std::cell::RefCell::new(Vec::::new()); + let mut on_progress = |p: DownloadProgress| events.borrow_mut().push(p); + + let mut adapter = CallbackProgress { + file: "tokens.txt".to_string(), + file_index: 4, + file_count: 4, + bytes_done: 0, + bytes_total: None, + on_progress: &mut on_progress, + }; + + adapter.init(500, "tokens.txt"); + adapter.update(400); // short + adapter.finish(); // should push a synthetic 500/500 + + let got = events.into_inner(); + assert_eq!(got.len(), 3); + assert_eq!(got.last().unwrap().bytes_done, 500); + assert_eq!(got.last().unwrap().bytes_total, Some(500)); + } +} diff --git a/crates/wavekat-asr/src/lib.rs b/crates/wavekat-asr/src/lib.rs index 0ec93e2..a9b6ce2 100644 --- a/crates/wavekat-asr/src/lib.rs +++ b/crates/wavekat-asr/src/lib.rs @@ -17,8 +17,12 @@ //! auto-downloads its model from HuggingFace on first use. pub mod backends; +#[cfg(feature = "download")] +pub mod download; pub mod error; +#[cfg(feature = "download")] +pub use download::DownloadProgress; pub use error::AsrError; pub use wavekat_core::AudioFrame;