Skip to content
Closed
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
5 changes: 5 additions & 0 deletions sherpa-onnx/c-api/c-api.cc
Original file line number Diff line number Diff line change
Expand Up @@ -351,6 +351,11 @@ void SherpaOnnxOnlineStreamInputFinished(const SherpaOnnxOnlineStream *stream) {
stream->impl->InputFinished();
}

void SherpaOnnxOnlineStreamSetFinalChunk(
const SherpaOnnxOnlineStream *stream) {
stream->impl->SetParaformerFinalChunk(true);
}
Comment on lines +354 to +357

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

medium

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 to SherpaOnnxOnlineStreamSetParaformerFinalChunk.

This change would also need to be applied to:

  • The function declaration in sherpa-onnx/c-api/c-api.h
  • The exported symbol in sherpa-onnx/c-api/sherpa-onnx-symbols-c.exp
void SherpaOnnxOnlineStreamSetParaformerFinalChunk(
    const SherpaOnnxOnlineStream *stream) {
  stream->impl->SetParaformerFinalChunk(true);
}


int32_t SherpaOnnxOnlineStreamIsEndpoint(
const SherpaOnnxOnlineRecognizer *recognizer,
const SherpaOnnxOnlineStream *stream) {
Expand Down
11 changes: 11 additions & 0 deletions sherpa-onnx/c-api/c-api.h
Original file line number Diff line number Diff line change
Expand Up @@ -380,6 +380,17 @@ SHERPA_ONNX_API void SherpaOnnxOnlineStreamReset(
SHERPA_ONNX_API void SherpaOnnxOnlineStreamInputFinished(
const SherpaOnnxOnlineStream *stream);

/// Signal that the current chunk is the final chunk for streaming Paraformer.
/// This enables:
/// 1. Short chunk acceptance (less than chunk_size frames)
/// 2. CIF tail token flush for residual accumulated alpha
///
/// Call this BEFORE the last InputFinished() + DecodeStream() cycle.
///
/// @param stream A pointer returned by SherpaOnnxCreateOnlineStream()
SHERPA_ONNX_API void SherpaOnnxOnlineStreamSetFinalChunk(
const SherpaOnnxOnlineStream *stream);

/// Return 1 if an endpoint has been detected.
///
/// @param recognizer A pointer returned by SherpaOnnxCreateOnlineRecognizer()
Expand Down
1 change: 1 addition & 0 deletions sherpa-onnx/c-api/sherpa-onnx-symbols-c.exp
Original file line number Diff line number Diff line change
Expand Up @@ -111,6 +111,7 @@ _SherpaOnnxOnlinePunctuationAddPunct
_SherpaOnnxOnlinePunctuationFreeText
_SherpaOnnxOnlineStreamAcceptWaveform
_SherpaOnnxOnlineStreamInputFinished
_SherpaOnnxOnlineStreamSetFinalChunk
_SherpaOnnxOnlineStreamIsEndpoint
_SherpaOnnxOnlineStreamReset
_SherpaOnnxPrint
Expand Down
54 changes: 50 additions & 4 deletions sherpa-onnx/csrc/online-recognizer-paraformer-impl.h
Original file line number Diff line number Diff line change
Expand Up @@ -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 {
Expand Down Expand Up @@ -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);
Expand Down Expand Up @@ -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;
}
Expand Down Expand Up @@ -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

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

medium

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 OnlineRecognizerParaformerImpl class, for instance, near kCifTailFlushMinAlpha:

  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
if (t == 0 || t == 1 || t == 2) {
// skip blank(0), sos(1), eos(2)
continue;
}
if (t == kBlankId || t == kSosId || t == kEosId) {
// skip blank(0), sos(1), eos(2)
continue;
}


Expand Down Expand Up @@ -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
Expand Down
18 changes: 18 additions & 0 deletions sherpa-onnx/csrc/online-stream.cc
Original file line number Diff line number Diff line change
Expand Up @@ -48,6 +48,7 @@ class OnlineStream::Impl {
// we don't reset the feature extractor
start_frame_index_ += num_processed_frames_;
num_processed_frames_ = 0;
paraformer_is_final_ = false;
}

int32_t &GetNumProcessedFrames() {
Expand Down Expand Up @@ -128,6 +129,14 @@ class OnlineStream::Impl {
return paraformer_alpha_cache_;
}

void SetParaformerFinalChunk(bool is_final) {
paraformer_is_final_ = is_final;
}

bool IsParaformerFinalChunk() const {
return paraformer_is_final_;
}

void SetFasterDecoder(std::unique_ptr<kaldi_decoder::FasterDecoder> decoder) {
faster_decoder_ = std::move(decoder);
}
Expand Down Expand Up @@ -159,6 +168,7 @@ class OnlineStream::Impl {
std::vector<float> paraformer_encoder_out_cache_;
std::vector<float> paraformer_alpha_cache_;
OnlineParaformerDecoderResult paraformer_result_;
bool paraformer_is_final_ = false;
std::unique_ptr<kaldi_decoder::FasterDecoder> faster_decoder_;
int32_t faster_decoder_processed_frames_ = 0;
};
Expand Down Expand Up @@ -282,4 +292,12 @@ std::vector<float> &OnlineStream::GetParaformerAlphaCache() {
return impl_->GetParaformerAlphaCache();
}

void OnlineStream::SetParaformerFinalChunk(bool is_final) {
impl_->SetParaformerFinalChunk(is_final);
}

bool OnlineStream::IsParaformerFinalChunk() const {
return impl_->IsParaformerFinalChunk();
}

} // namespace sherpa_onnx
4 changes: 4 additions & 0 deletions sherpa-onnx/csrc/online-stream.h
Original file line number Diff line number Diff line change
Expand Up @@ -111,6 +111,10 @@ class OnlineStream {
std::vector<float> &GetParaformerEncoderOutCache();
std::vector<float> &GetParaformerAlphaCache();

// for streaming paraformer is_final support
void SetParaformerFinalChunk(bool is_final);
bool IsParaformerFinalChunk() const;

private:
class Impl;
std::unique_ptr<Impl> impl_;
Expand Down
3 changes: 3 additions & 0 deletions sherpa-onnx/python/csrc/online-stream.cc
Original file line number Diff line number Diff line change
Expand Up @@ -52,6 +52,9 @@ void PybindOnlineStream(py::module *m) {
py::call_guard<py::gil_scoped_release>())
.def("input_finished", &PyClass::InputFinished,
py::call_guard<py::gil_scoped_release>())
.def("set_paraformer_final_chunk", &PyClass::SetParaformerFinalChunk,
py::arg("is_final") = true,
py::call_guard<py::gil_scoped_release>())
.def("get_frames", &PyClass::GetFrames,
py::arg("frame_index"), py::arg("n"), kGetFramesUsage,
py::call_guard<py::gil_scoped_release>());
Expand Down