-
Notifications
You must be signed in to change notification settings - Fork 1.7k
feat: add is_final support for streaming Paraformer #3282
New issue
Have a question about this project? Sign up for a free GitHub account to open an issue and contact its maintainers and the community.
By clicking “Sign up for GitHub”, you agree to our terms of service and privacy statement. We’ll occasionally send you account related emails.
Already on GitHub? Sign in to your account
Changes from all commits
File filter
Filter by extension
Conversations
Jump to
Diff view
Diff view
There are no files selected for viewing
| Original file line number | Diff line number | Diff line change | ||||||||||||||||
|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|
|
|
@@ -156,7 +156,15 @@ class OnlineRecognizerParaformerImpl : public OnlineRecognizerImpl { | |||||||||||||||||
| } | ||||||||||||||||||
|
|
||||||||||||||||||
| bool IsReady(OnlineStream *s) const override { | ||||||||||||||||||
| return s->GetNumProcessedFrames() + chunk_size_ < s->NumFramesReady(); | ||||||||||||||||||
| if (s->GetNumProcessedFrames() + chunk_size_ < s->NumFramesReady()) { | ||||||||||||||||||
| return true; | ||||||||||||||||||
| } | ||||||||||||||||||
| // is_final: accept short chunks (less than chunk_size_ frames) | ||||||||||||||||||
| if (s->IsParaformerFinalChunk() && | ||||||||||||||||||
| s->GetNumProcessedFrames() < s->NumFramesReady()) { | ||||||||||||||||||
| return true; | ||||||||||||||||||
| } | ||||||||||||||||||
| return false; | ||||||||||||||||||
| } | ||||||||||||||||||
|
|
||||||||||||||||||
| void DecodeStreams(OnlineStream **ss, int32_t n) const override { | ||||||||||||||||||
|
|
@@ -212,8 +220,30 @@ class OnlineRecognizerParaformerImpl : public OnlineRecognizerImpl { | |||||||||||||||||
| private: | ||||||||||||||||||
| void DecodeStream(OnlineStream *s) const { | ||||||||||||||||||
| const auto num_processed_frames = s->GetNumProcessedFrames(); | ||||||||||||||||||
| std::vector<float> frames = s->GetFrames(num_processed_frames, chunk_size_); | ||||||||||||||||||
| s->GetNumProcessedFrames() += chunk_size_ - 1; | ||||||||||||||||||
|
|
||||||||||||||||||
| // is_final: accept short chunks, pad with zeros if needed | ||||||||||||||||||
| int32_t available_frames = s->NumFramesReady() - num_processed_frames; | ||||||||||||||||||
| int32_t actual_chunk_size = chunk_size_; | ||||||||||||||||||
| if (s->IsParaformerFinalChunk() && available_frames < chunk_size_) { | ||||||||||||||||||
| actual_chunk_size = available_frames; | ||||||||||||||||||
| } | ||||||||||||||||||
|
|
||||||||||||||||||
| std::vector<float> frames = | ||||||||||||||||||
| s->GetFrames(num_processed_frames, actual_chunk_size); | ||||||||||||||||||
|
|
||||||||||||||||||
| // Pad to chunk_size_ if short chunk | ||||||||||||||||||
| if (actual_chunk_size < chunk_size_) { | ||||||||||||||||||
| int32_t feat_dim_raw = config_.feat_config.feature_dim; | ||||||||||||||||||
| frames.resize(chunk_size_ * feat_dim_raw, 0.0f); | ||||||||||||||||||
| } | ||||||||||||||||||
|
|
||||||||||||||||||
| // For non-final chunks the original code uses chunk_size_ - 1 to create | ||||||||||||||||||
| // 1-frame overlap. For the final short chunk we consume all frames. | ||||||||||||||||||
| if (s->IsParaformerFinalChunk()) { | ||||||||||||||||||
| s->GetNumProcessedFrames() += actual_chunk_size; | ||||||||||||||||||
| } else { | ||||||||||||||||||
| s->GetNumProcessedFrames() += chunk_size_ - 1; | ||||||||||||||||||
| } | ||||||||||||||||||
|
|
||||||||||||||||||
| frames = ApplyLFR(frames); | ||||||||||||||||||
| ApplyCMVN(&frames); | ||||||||||||||||||
|
|
@@ -315,6 +345,16 @@ class OnlineRecognizerParaformerImpl : public OnlineRecognizerImpl { | |||||||||||||||||
|
|
||||||||||||||||||
| alpha_cache[0] = integrate; | ||||||||||||||||||
|
|
||||||||||||||||||
| // is_final: tail flush — force-fire residual token if integrate is | ||||||||||||||||||
| // high enough (token was mostly accumulated, just shy of threshold). | ||||||||||||||||||
| if (s->IsParaformerFinalChunk() && integrate >= kCifTailFlushMinAlpha) { | ||||||||||||||||||
| acoustic_embedding.insert(acoustic_embedding.end(), | ||||||||||||||||||
| initial_hidden.begin(), initial_hidden.end()); | ||||||||||||||||||
| integrate = 0.0f; | ||||||||||||||||||
| std::fill(initial_hidden.begin(), initial_hidden.end(), 0.0f); | ||||||||||||||||||
| alpha_cache[0] = integrate; | ||||||||||||||||||
| } | ||||||||||||||||||
|
|
||||||||||||||||||
| if (acoustic_embedding.empty()) { | ||||||||||||||||||
| return; | ||||||||||||||||||
| } | ||||||||||||||||||
|
|
@@ -372,7 +412,8 @@ class OnlineRecognizerParaformerImpl : public OnlineRecognizerImpl { | |||||||||||||||||
|
|
||||||||||||||||||
| for (int32_t i = 0; i != num_tokens; ++i) { | ||||||||||||||||||
| int32_t t = p_sample_ids[i]; | ||||||||||||||||||
| if (t == 0) { | ||||||||||||||||||
| if (t == 0 || t == 1 || t == 2) { | ||||||||||||||||||
| // skip blank(0), sos(1), eos(2) | ||||||||||||||||||
| continue; | ||||||||||||||||||
| } | ||||||||||||||||||
|
Comment on lines
+415
to
418
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. Using magic numbers for special token IDs makes the code less readable and harder to maintain. It's better to define them as named constants. You could add these constants to the static constexpr int32_t kBlankId = 0;
static constexpr int32_t kSosId = 1;
static constexpr int32_t kEosId = 2;Then, you can use them here to make the condition more explicit.
Suggested change
|
||||||||||||||||||
|
|
||||||||||||||||||
|
|
@@ -469,6 +510,11 @@ class OnlineRecognizerParaformerImpl : public OnlineRecognizerImpl { | |||||||||||||||||
|
|
||||||||||||||||||
| int32_t left_chunk_size_ = 5; | ||||||||||||||||||
| int32_t right_chunk_size_ = 3; | ||||||||||||||||||
|
|
||||||||||||||||||
| // Minimum CIF residual alpha to force-fire a tail token on final chunk. | ||||||||||||||||||
| // Empirically tuned: below 0.6 causes hallucinated extra tokens; | ||||||||||||||||||
| // above 0.6 misses legitimate partial tokens. | ||||||||||||||||||
| static constexpr float kCifTailFlushMinAlpha = 0.6f; | ||||||||||||||||||
| }; | ||||||||||||||||||
|
|
||||||||||||||||||
| } // namespace sherpa_onnx | ||||||||||||||||||
|
|
||||||||||||||||||
There was a problem hiding this comment.
Choose a reason for hiding this comment
The reason will be displayed to describe this comment to others. Learn more.
For consistency with the C++ (
SetParaformerFinalChunk) and Python (set_paraformer_final_chunk) APIs, and to make it clear this function is specific to Paraformer models, consider renaming this function toSherpaOnnxOnlineStreamSetParaformerFinalChunk.This change would also need to be applied to:
sherpa-onnx/c-api/c-api.hsherpa-onnx/c-api/sherpa-onnx-symbols-c.exp