Repository navigation
Expose token + vocab log probabilities from offline CTC decoder #3630
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
base: master
Are you sure you want to change the base?
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 |
|---|---|---|
|
|
@@ -865,6 +865,24 @@ const SherpaOnnxOfflineRecognizerResult *SherpaOnnxGetOfflineStreamResult( | |
| r->ys_log_probs = nullptr; | ||
| } | ||
|
|
||
| // Copy vocab_log_probs (flattened row-major: count * vocab_size) | ||
| if (!result.vocab_log_probs.empty() && | ||
| static_cast<int32_t>(result.vocab_log_probs.size()) == r->count && | ||
| !result.vocab_log_probs[0].empty()) { | ||
| int32_t vocab_size = | ||
| static_cast<int32_t>(result.vocab_log_probs[0].size()); | ||
| r->vocab_size = vocab_size; | ||
| float *flat = new float[r->count * vocab_size]; | ||
| for (int32_t i = 0; i < r->count; ++i) { | ||
| std::copy(result.vocab_log_probs[i].begin(), | ||
| result.vocab_log_probs[i].end(), flat + i * vocab_size); | ||
| } | ||
| r->vocab_log_probs = flat; | ||
|
Comment on lines
+875
to
+880
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. Potential integer overflow in the calculation of the buffer size and pointer arithmetic. r->count and vocab_size are both int32_t, so their product is calculated as a 32-bit signed integer. If the product exceeds INT32_MAX (approx. 2.1 billion), it will overflow, leading to an incorrect allocation size and undefined behavior during indexing in the subsequent loop. This is a realistic scenario for long audio files or large vocabularies. float *flat = new float[static_cast<size_t>(r->count) * vocab_size];
for (int32_t i = 0; i < r->count; ++i) {
std::copy(result.vocab_log_probs[i].begin(),
result.vocab_log_probs[i].end(),
flat + static_cast<size_t>(i) * vocab_size);
}
r->vocab_log_probs = flat; |
||
| } else { | ||
| r->vocab_log_probs = nullptr; | ||
| r->vocab_size = 0; | ||
| } | ||
|
|
||
| // Copy segment-level timestamps (from Whisper with segment timestamps) | ||
| auto segment_count = result.segment_texts.size(); | ||
| if (segment_count > 0 && result.segment_timestamps.size() == segment_count && | ||
|
|
@@ -908,6 +926,8 @@ const SherpaOnnxOfflineRecognizerResult *SherpaOnnxGetOfflineStreamResult( | |
| r->segment_texts_arr = nullptr; | ||
| } | ||
|
|
||
| // NB: vocab_log_probs / vocab_size are set above (before the segment block). | ||
|
|
||
| return r; | ||
| } | ||
|
|
||
|
|
@@ -928,6 +948,7 @@ void SherpaOnnxDestroyOfflineRecognizerResult( | |
| delete[] r->segment_durations; | ||
| delete[] r->segment_texts; | ||
| delete[] r->segment_texts_arr; | ||
| delete[] r->vocab_log_probs; | ||
| delete r; | ||
| } | ||
| } | ||
|
|
||
| Original file line number | Diff line number | Diff line change |
|---|---|---|
|
|
@@ -37,12 +37,15 @@ std::vector<OfflineCtcDecoderResult> OfflineCtcGreedySearchDecoder::Decode( | |
| std::max_element( | ||
| static_cast<const float *>(p_log_probs), | ||
| static_cast<const float *>(p_log_probs) + vocab_size))); | ||
| p_log_probs += vocab_size; | ||
| float log_prob = p_log_probs[y]; | ||
|
|
||
| if (y != blank_id_ && y != prev_id) { | ||
| r.tokens.push_back(y); | ||
| r.timestamps.push_back(t); | ||
| r.token_log_probs.push_back(log_prob); | ||
| r.vocab_log_probs.emplace_back(p_log_probs, p_log_probs + vocab_size); | ||
|
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. Unconditionally populating vocab_log_probs for every emitted token can lead to excessive memory consumption, especially for long audio inputs and models with large vocabularies. For instance, a 1-hour recording with 10,000 tokens and a 50,000-word vocabulary would require approximately 2GB of additional memory. Since this data is only needed for specific use cases like entropy-based confidence estimation, it should be optional. Additionally, using std::vector<std::vector> results in many small allocations; a single flat vector would be more efficient. |
||
| } | ||
| p_log_probs += vocab_size; | ||
| prev_id = y; | ||
| } // for (int32_t t = 0; ...) | ||
|
|
||
|
|
||
| Original file line number | Diff line number | Diff line change |
|---|---|---|
|
|
@@ -77,6 +77,16 @@ OfflineRecognitionResult Convert(const OfflineCtcDecoderResult &src, | |
|
|
||
| r.words = std::move(src.words); | ||
|
|
||
| if (!src.token_log_probs.empty() && | ||
| src.token_log_probs.size() == src.tokens.size()) { | ||
| r.ys_log_probs = src.token_log_probs; | ||
| } | ||
|
|
||
| if (!src.vocab_log_probs.empty() && | ||
| src.vocab_log_probs.size() == src.tokens.size()) { | ||
| r.vocab_log_probs = src.vocab_log_probs; | ||
| } | ||
|
Comment on lines
+80
to
+88
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. Keep log-prob arrays aligned with filtered output tokens
💡 Suggested fix@@
OfflineRecognitionResult r;
r.tokens.reserve(src.tokens.size());
r.timestamps.reserve(src.timestamps.size());
+ std::vector<int32_t> kept_indices;
+ kept_indices.reserve(src.tokens.size());
@@
auto sym = sym_table[src.tokens[i]];
@@
r.tokens.push_back(std::move(sym));
+ kept_indices.push_back(i);
}
@@
- if (!src.token_log_probs.empty() &&
- src.token_log_probs.size() == src.tokens.size()) {
- r.ys_log_probs = src.token_log_probs;
+ if (!src.token_log_probs.empty() &&
+ src.token_log_probs.size() == src.tokens.size()) {
+ r.ys_log_probs.reserve(kept_indices.size());
+ for (auto idx : kept_indices) {
+ r.ys_log_probs.push_back(src.token_log_probs[idx]);
+ }
}
if (!src.vocab_log_probs.empty() &&
src.vocab_log_probs.size() == src.tokens.size()) {
- r.vocab_log_probs = src.vocab_log_probs;
+ r.vocab_log_probs.reserve(kept_indices.size());
+ for (auto idx : kept_indices) {
+ r.vocab_log_probs.push_back(src.vocab_log_probs[idx]);
+ }
}🤖 Prompt for AI Agents |
||
|
|
||
| return r; | ||
| } | ||
|
|
||
|
|
||
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.
Validate every row width before flattening to avoid buffer corruption.
At Line 877, the copy assumes every
result.vocab_log_probs[i]hasvocab_sizeelements (from Line 873), but only row 0 is validated. A wider row can overrunflat; a shorter row yields partial/uninitialized output. Also, Line 875 multipliesint32_tdimensions without overflow guarding.Suggested fix
🤖 Prompt for AI Agents