Conversation
…custom vocab log probs API - Resolved conflicts in offline-stream.cc to support both ys_log_probs (upstream) and token_log_probs (custom) - Updated pubspec.yaml to use local path for development - Preserved custom SherpaOnnxVocabLogProbs API for full vocabulary log probabilities - Maintained compatibility with upstream's ys_log_probs field
|
Important Review skippedDraft detected. Please check the settings in the CodeRabbit UI or the ⚙️ Run configurationConfiguration used: defaults Review profile: CHILL Plan: Pro Run ID: You can disable this status message by setting the Use the checkbox below for a quick retry:
Note Reviews pausedIt looks like this branch is under active development. To avoid overwhelming you with review comments due to an influx of new commits, CodeRabbit has automatically paused this review. You can configure this behavior by changing the Use the following commands to manage reviews:
Use the checkboxes below for quick actions:
📝 WalkthroughWalkthroughAdds token- and vocabulary-level log-probability support across the codebase: new C API struct/functions for vocab log-probs, Dart FFI bindings and stream methods to fetch them, result types extended with token/ys/vocab probs, decoders updated to compute log-softmax distributions, and JSON/serialization wiring updated. Changes
Sequence Diagram(s)sequenceDiagram
participant App as Dart App
participant DartStream as Dart Online/Offline Stream
participant Bindings as SherpaOnnxBindings
participant Native as Native Lib (C/C++)
participant Heap as Native Heap
App->>DartStream: getVocabLogProbs()
DartStream->>Bindings: check get*VocabLogProbs ptr
Bindings-->>DartStream: function ptr
DartStream->>Native: call get*VocabLogProbs(stream)
activate Native
Native->>Heap: allocate SherpaOnnxVocabLogProbs + flattened array
Native-->>DartStream: Pointer<SherpaOnnxVocabLogProbs>
deactivate Native
DartStream->>DartStream: read numTokens, vocabSize, flatten -> Map<String,List<double>>
DartStream->>Bindings: destroyVocabLogProbs(ptr)
Bindings->>Native: free array + struct
Native-->>Bindings: done
DartStream-->>App: return Map or null
Estimated code review effort🎯 4 (Complex) | ⏱️ ~60 minutes Possibly related PRs
Suggested reviewers
Poem
🚥 Pre-merge checks | ✅ 2 | ❌ 1❌ Failed checks (1 warning)
✅ Passed checks (2 passed)
✏️ Tip: You can configure your own custom pre-merge checks in the settings. ✨ Finishing Touches🧪 Generate unit tests (beta)
Thanks for using CodeRabbit! It's free for OSS, and your support helps us grow. If you like it, consider giving us a shout-out. Comment |
Summary of ChangesHello @Dokotela, I'm Gemini Code Assist1! I'm currently reviewing this pull request and will post my feedback shortly. In the meantime, here's a summary to help you and other reviewers quickly get up to speed! This pull request introduces a significant enhancement by exposing the full vocabulary log probabilities for each decoded token position across both online and offline speech recognition streams. This new capability is crucial for advanced applications such as confidence scoring, uncertainty estimation, and hypothesis fusion algorithms, providing a deeper insight into model predictions beyond just the most likely token. Highlights
Using Gemini Code AssistThe full guide for Gemini Code Assist can be found on our documentation page, here are some quick tips. Invoking Gemini You can request assistance from Gemini at any point by creating a comment using either
Customization To customize Gemini Code Assist for GitHub experience, repository maintainers can create a configuration file and/or provide a custom code review style guide (such as PEP-8 for Python) by creating and adding files to a Limitations & Feedback Gemini Code Assist may make mistakes. Please leave feedback on any instances where its feedback is incorrect or counter productive. You can react with 👍 and 👎 on @gemini-code-assist comments. If you're interested in giving your feedback about your experience with Gemini Code Assist for Github and other Google products, sign up here. You can also get AI-powered code generation, chat, as well as code reviews directly in the IDE at no cost with the Gemini Code Assist IDE Extension. Footnotes
|
There was a problem hiding this comment.
Code Review
This pull request introduces a new API to access vocabulary log probabilities, which is a valuable feature for advanced use cases like model fusion and uncertainty estimation. The changes span across the C-API, core C++ logic, and Dart/Flutter bindings.
My review has identified a few areas for improvement:
- There are some inconsistencies in propagating the new
vocab_log_probsdata in different recognizer implementations. - A critical copy-paste error was found in one of the decoders.
- Some opportunities for code cleanup and performance optimization were also noted.
Overall, this is a great addition. Addressing the feedback will improve the robustness and consistency of the new API.
There was a problem hiding this comment.
Actionable comments posted: 6
🧹 Nitpick comments (9)
sherpa-onnx/csrc/offline-moonshine-greedy-search-decoder.cc (2)
92-96: Consider removing the "ADD THIS" comment.The inline comment
// ADD THISappears to be a development artifact and should be removed before merging.- // Store full vocabulary distribution (already log-softmaxed) + // Store full vocabulary log-probability distribution std::vector<float> full_vocab_probs(vocab_size);
119-120: Remove development comment.The comment
// ADD THISon line 120 should be removed.ans.token_log_probs = std::move(token_log_probs); - ans.vocab_log_probs = std::move(vocab_log_probs); // ADD THIS + ans.vocab_log_probs = std::move(vocab_log_probs);sherpa-onnx/csrc/offline-moonshine-decoder.h (1)
17-19: Consider adding documentation forvocab_log_probs.The
token_log_probsfield has a descriptive comment. Consider adding a similar comment forvocab_log_probsto document its shape and semantics./// Token-level log probabilities (confidence scores) std::vector<float> token_log_probs; + /// Full vocabulary log-probability distribution per token [num_tokens][vocab_size] std::vector<std::vector<float>> vocab_log_probs;sherpa-onnx/c-api/c-api.h (1)
676-689: Add documentation for the new API functions.The struct definition has good inline documentation, but the three new functions lack documentation comments explaining:
- Return value semantics (caller owns the pointer, returns NULL if no vocab log probs available)
- Memory management (must call
SherpaOnnxDestroyVocabLogProbsto avoid memory leak)SHERPA_ONNX_API typedef struct SherpaOnnxVocabLogProbs { const float *log_probs; // Flattened 2D array [num_tokens][vocab_size] int32_t num_tokens; int32_t vocab_size; } SherpaOnnxVocabLogProbs; +/// Get vocabulary log probabilities for an online stream. +/// +/// @param stream A pointer returned by SherpaOnnxCreateOnlineStream(). +/// @return A pointer to vocab log probs, or NULL if not available. +/// The user must invoke SherpaOnnxDestroyVocabLogProbs() to free +/// the returned pointer to avoid memory leak. SHERPA_ONNX_API const SherpaOnnxVocabLogProbs * SherpaOnnxOnlineStreamGetVocabLogProbs(const SherpaOnnxOnlineStream *stream); +/// Get vocabulary log probabilities for an offline stream. +/// +/// @param stream A pointer returned by SherpaOnnxCreateOfflineStream(). +/// @return A pointer to vocab log probs, or NULL if not available. +/// The user must invoke SherpaOnnxDestroyVocabLogProbs() to free +/// the returned pointer to avoid memory leak. SHERPA_ONNX_API const SherpaOnnxVocabLogProbs * SherpaOnnxOfflineStreamGetVocabLogProbs(const SherpaOnnxOfflineStream *stream); +/// Free a pointer returned by SherpaOnnxOnlineStreamGetVocabLogProbs() or +/// SherpaOnnxOfflineStreamGetVocabLogProbs(). +/// +/// @param log_probs A pointer returned by the getter functions. SHERPA_ONNX_API void SherpaOnnxDestroyVocabLogProbs( const SherpaOnnxVocabLogProbs *log_probs);sherpa-onnx/csrc/offline-stream.cc (2)
445-473: Suggest consistent precision for log probability fields.The serialization uses different precisions for log probabilities:
ys_log_probs: precision 6 (line 454)token_log_probs: precision 4 (line 469)Since both represent log probabilities, consider using consistent precision (e.g., 6 for both) to maintain uniformity in the JSON output.
Apply this diff for consistency:
for (auto prob : token_log_probs) { - os << sep << std::fixed << std::setprecision(4) << prob; + os << sep << std::fixed << std::setprecision(6) << prob; sep = ", "; }
472-486: Trailing comma assumes words is always serialized.Line 472 adds a trailing comma after
token_log_probs, which is fine becausewordsis always serialized next (lines 477-486). However, if in the futurewordsserialization becomes conditional, this would produce invalid JSON whentoken_log_probsis present butwordsis empty.Consider this pattern if words ever becomes optional:
// Don't add trailing comma after token_log_probs os << "]"; // Later, add comma before words if token_log_probs was present if (!token_log_probs.empty() && !words.empty()) { os << ", "; }flutter/sherpa_onnx/lib/src/online_stream.dart (1)
40-73: Consider usingList<List<double>>instead ofMap<String, List<double>>for better performance and semantics.The current implementation has two concerns:
Data structure choice: Using a
Mapwith string keys ("token_0","token_1", etc.) for inherently sequential data is unusual. AList<List<double>>would be more natural, efficient, and consistent with the underlying C array structure.Performance: The nested loop with per-element access is O(numTokens × vocabSize). For large vocabularies, using
asTypedListfor bulk copy would be significantly faster.- Map<String, List<double>>? getVocabLogProbs() { + List<List<double>>? getVocabLogProbs() { final getFunc = SherpaOnnxBindings.getOnlineStreamVocabLogProbs; final destroyFunc = SherpaOnnxBindings.destroyVocabLogProbs; if (getFunc == null || destroyFunc == null) { return null; } final ptr = getFunc(this.ptr); if (ptr == nullptr) { return null; } final vocabLogProbs = ptr.ref; final numTokens = vocabLogProbs.numTokens; final vocabSize = vocabLogProbs.vocabSize; - final Map<String, List<double>> result = {}; - - for (int tokenIdx = 0; tokenIdx < numTokens; tokenIdx++) { - final List<double> tokenProbs = []; - - for (int vocabIdx = 0; vocabIdx < vocabSize; vocabIdx++) { - final index = tokenIdx * vocabSize + vocabIdx; - final logProb = vocabLogProbs.logProbs[index]; - tokenProbs.add(logProb); - } - - result['token_$tokenIdx'] = tokenProbs; - } + final List<List<double>> result = []; + final allProbs = vocabLogProbs.logProbs.asTypedList(numTokens * vocabSize); + + for (int tokenIdx = 0; tokenIdx < numTokens; tokenIdx++) { + final start = tokenIdx * vocabSize; + result.add(allProbs.sublist(start, start + vocabSize).toList()); + } destroyFunc(ptr); return result; }sherpa-onnx/csrc/offline-whisper-greedy-search-decoder.cc (1)
156-168: Consider reusing CalculateLogProb logic to reduce duplication.The manual log-softmax computation here (lines 159-167) duplicates logic from the
CalculateLogProbhelper function. Consider extracting a shared helper to compute the full log-softmax distribution.For example:
// Add a helper function to compute full log-softmax static void ComputeLogSoftmax(const float *logits, int32_t vocab_size, std::vector<float> &out) { out.resize(vocab_size); float max_logit = *std::max_element(logits, logits + vocab_size); double sum_exp = 0.0; for (int32_t i = 0; i < vocab_size; ++i) { sum_exp += std::exp(logits[i] - max_logit); } float log_sum = max_logit + std::log(sum_exp); for (int32_t i = 0; i < vocab_size; ++i) { out[i] = logits[i] - log_sum; } } // Then in the loop: std::vector<float> full_vocab_probs; ComputeLogSoftmax(current_logits, vocab_size, full_vocab_probs); predicted_vocab_log_probs.push_back(std::move(full_vocab_probs));sherpa-onnx/csrc/offline-recognizer-canary-impl.h (1)
171-209: Consider optimizing to reduce redundant passes.The method makes multiple passes over the logits array:
- Lines 179: Find max_logit
- Lines 183-185: Compute sum_exp
- Lines 191-193: Fill full_distribution (if requested)
- Lines 200-206: Find max token
Consider combining passes for better cache locality and performance:
std::pair<int32_t, float> GetMaxTokenIdWithConfidence( Ort::Value *logits, std::vector<float> *full_distribution = nullptr) const { auto meta = model_->GetModelMetadata(); const float *p_logits = logits->GetTensorData<float>(); // Find max for numerical stability float max_logit = *std::max_element(p_logits, p_logits + meta.vocab_size); // Single pass: compute sum_exp and find max token float sum_exp = 0.0f; int32_t max_token_id = 0; float max_log_prob = -std::numeric_limits<float>::infinity(); for (int32_t i = 0; i < meta.vocab_size; ++i) { float exp_val = std::exp(p_logits[i] - max_logit); sum_exp += exp_val; float log_prob = p_logits[i] - max_logit; // will subtract log_sum later if (log_prob > max_log_prob) { max_log_prob = log_prob; max_token_id = i; } } float log_sum = max_logit + std::log(sum_exp); max_log_prob = max_log_prob + max_logit - log_sum; // correct the log_prob // Fill distribution if requested if (full_distribution != nullptr) { full_distribution->resize(meta.vocab_size); for (int32_t i = 0; i < meta.vocab_size; ++i) { (*full_distribution)[i] = p_logits[i] - log_sum; } } return {max_token_id, max_log_prob}; }
📜 Review details
Configuration used: CodeRabbit UI
Review profile: CHILL
Plan: Pro
📒 Files selected for processing (26)
flutter/sherpa_onnx/lib/src/offline_recognizer.dart(6 hunks)flutter/sherpa_onnx/lib/src/offline_stream.dart(1 hunks)flutter/sherpa_onnx/lib/src/online_recognizer.dart(3 hunks)flutter/sherpa_onnx/lib/src/online_stream.dart(1 hunks)flutter/sherpa_onnx/lib/src/sherpa_onnx_bindings.dart(4 hunks)flutter/sherpa_onnx/pubspec.yaml(1 hunks)sherpa-onnx/c-api/c-api.cc(1 hunks)sherpa-onnx/c-api/c-api.h(1 hunks)sherpa-onnx/csrc/offline-ctc-decoder.h(1 hunks)sherpa-onnx/csrc/offline-moonshine-decoder.h(1 hunks)sherpa-onnx/csrc/offline-moonshine-greedy-search-decoder.cc(4 hunks)sherpa-onnx/csrc/offline-recognizer-canary-impl.h(9 hunks)sherpa-onnx/csrc/offline-recognizer-ctc-impl.h(2 hunks)sherpa-onnx/csrc/offline-recognizer-moonshine-impl.h(1 hunks)sherpa-onnx/csrc/offline-recognizer-whisper-impl.h(2 hunks)sherpa-onnx/csrc/offline-stream.cc(1 hunks)sherpa-onnx/csrc/offline-stream.h(1 hunks)sherpa-onnx/csrc/offline-whisper-decoder.h(1 hunks)sherpa-onnx/csrc/offline-whisper-greedy-search-decoder.cc(4 hunks)sherpa-onnx/csrc/online-recognizer-transducer-nemo-impl.h(2 hunks)sherpa-onnx/csrc/online-recognizer.h(1 hunks)sherpa-onnx/csrc/online-transducer-decoder.cc(2 hunks)sherpa-onnx/csrc/online-transducer-decoder.h(1 hunks)sherpa-onnx/csrc/online-transducer-greedy-search-decoder.cc(2 hunks)sherpa-onnx/csrc/online-transducer-greedy-search-nemo-decoder.cc(3 hunks)sherpa-onnx/csrc/online-transducer-greedy-search-nemo-decoder.h(1 hunks)
🧰 Additional context used
🧠 Learnings (2)
📚 Learning: 2025-08-06T04:23:50.237Z
Learnt from: litongjava
Repo: k2-fsa/sherpa-onnx PR: 2440
File: sherpa-onnx/java-api/src/main/java/com/k2fsa/sherpa/onnx/core/Core.java:4-6
Timestamp: 2025-08-06T04:23:50.237Z
Learning: The sherpa-onnx JNI library files are stored in Hugging Face repository at https://huggingface.co/csukuangfj/sherpa-onnx-libs under versioned directories like jni/1.12.7/, and the actual Windows JNI library filename is "sherpa-onnx-jni.dll" as defined in Core.java constants.
Applied to files:
flutter/sherpa_onnx/pubspec.yamlflutter/sherpa_onnx/lib/src/sherpa_onnx_bindings.dart
📚 Learning: 2025-08-06T04:18:47.981Z
Learnt from: litongjava
Repo: k2-fsa/sherpa-onnx PR: 2440
File: sherpa-onnx/java-api/src/main/java/com/k2fsa/sherpa/onnx/core/Core.java:4-6
Timestamp: 2025-08-06T04:18:47.981Z
Learning: In sherpa-onnx Java API, the native library names in Core.java (WIN_NATIVE_LIBRARY_NAME = "sherpa-onnx-jni.dll", UNIX_NATIVE_LIBRARY_NAME = "libsherpa-onnx-jni.so", MACOS_NATIVE_LIBRARY_NAME = "libsherpa-onnx-jni.dylib") are copied directly from the compiled binary filenames and should not be changed to match other libraries' naming conventions.
Applied to files:
flutter/sherpa_onnx/pubspec.yamlflutter/sherpa_onnx/lib/src/sherpa_onnx_bindings.dart
🧬 Code graph analysis (11)
sherpa-onnx/csrc/online-recognizer.h (1)
sherpa-onnx/csrc/math.h (1)
float(52-72)
sherpa-onnx/csrc/online-transducer-greedy-search-nemo-decoder.cc (2)
sherpa-onnx/csrc/offline-moonshine-greedy-search-decoder.cc (1)
full_vocab_probs(92-92)sherpa-onnx/csrc/offline-whisper-greedy-search-decoder.cc (1)
full_vocab_probs(158-158)
sherpa-onnx/csrc/offline-recognizer-moonshine-impl.h (1)
sherpa-onnx/csrc/offline-stream.cc (2)
r(257-257)r(257-257)
sherpa-onnx/c-api/c-api.h (1)
sherpa-onnx/c-api/c-api.cc (6)
SherpaOnnxOnlineStreamGetVocabLogProbs(757-780)SherpaOnnxOnlineStreamGetVocabLogProbs(757-758)SherpaOnnxOfflineStreamGetVocabLogProbs(782-806)SherpaOnnxOfflineStreamGetVocabLogProbs(782-783)SherpaOnnxDestroyVocabLogProbs(808-814)SherpaOnnxDestroyVocabLogProbs(808-809)
sherpa-onnx/csrc/offline-moonshine-greedy-search-decoder.cc (1)
sherpa-onnx/csrc/offline-whisper-greedy-search-decoder.cc (1)
full_vocab_probs(158-158)
sherpa-onnx/csrc/offline-recognizer-ctc-impl.h (1)
sherpa-onnx/csrc/offline-stream.cc (2)
r(257-257)r(257-257)
sherpa-onnx/csrc/online-transducer-decoder.h (1)
sherpa-onnx/csrc/math.h (1)
float(52-72)
sherpa-onnx/csrc/offline-whisper-greedy-search-decoder.cc (1)
sherpa-onnx/csrc/offline-moonshine-greedy-search-decoder.cc (1)
full_vocab_probs(92-92)
sherpa-onnx/csrc/online-transducer-greedy-search-decoder.cc (2)
sherpa-onnx/csrc/offline-moonshine-greedy-search-decoder.cc (1)
full_vocab_probs(92-92)sherpa-onnx/csrc/offline-whisper-greedy-search-decoder.cc (1)
full_vocab_probs(158-158)
sherpa-onnx/csrc/offline-stream.h (1)
sherpa-onnx/csrc/math.h (1)
float(52-72)
sherpa-onnx/csrc/offline-recognizer-canary-impl.h (4)
sherpa-onnx/csrc/offline-moonshine-greedy-search-decoder.cc (1)
full_vocab_probs(92-92)sherpa-onnx/csrc/offline-whisper-greedy-search-decoder.cc (1)
full_vocab_probs(158-158)sherpa-onnx/csrc/onnx-utils.cc (2)
View(188-227)View(188-188)sherpa-onnx/csrc/offline-stream.cc (2)
r(257-257)r(257-257)
🪛 Cppcheck (2.18.0)
sherpa-onnx/csrc/offline-moonshine-greedy-search-decoder.cc
[warning] 72-72: Invalid std
(invalidFunctionArg)
sherpa-onnx/csrc/offline-whisper-greedy-search-decoder.cc
[warning] 44-44: Invalid std
(invalidFunctionArg)
⏰ Context from checks skipped due to timeout of 90000ms. You can increase the timeout in your CodeRabbit configuration to a maximum of 15 minutes (900000ms). (20)
- GitHub Check: ubuntu-24.04 3.11
- GitHub Check: ubuntu-24.04 3.9
- GitHub Check: Release shared tts-ON
- GitHub Check: ubuntu-24.04 3.12
- GitHub Check: ubuntu-24.04 3.10
- GitHub Check: Release static tts-ON
- GitHub Check: Release shared tts-OFF
- GitHub Check: Debug static tts-ON
- GitHub Check: Release shared-OFF tts-OFF
- GitHub Check: Release shared-OFF tts-ON
- GitHub Check: ubuntu-24.04 3.8
- GitHub Check: Debug shared tts-OFF
- GitHub Check: Debug shared-OFF tts-ON
- GitHub Check: rknn shared ON
- GitHub Check: ubuntu-24.04 3.13
- GitHub Check: Release shared-ON tts-OFF
- GitHub Check: Debug shared-ON tts-OFF
- GitHub Check: Debug shared-OFF tts-OFF
- GitHub Check: Debug shared-ON tts-ON
- GitHub Check: Release shared-ON tts-ON
🔇 Additional comments (28)
sherpa-onnx/csrc/online-transducer-decoder.cc (1)
43-43: LGTM!The copy and move assignment operators correctly propagate
vocab_log_probsfollowing the same pattern used forys_probs,lm_probs, andcontext_scores.Also applies to: 71-71
sherpa-onnx/csrc/offline-moonshine-greedy-search-decoder.cc (1)
65-83: Log-softmax implementation is correct but has minor redundancy.The numerical stability approach (subtracting max_logit) is appropriate. However, the loop at lines 77-83 recomputes
p[j] - log_sumwhich was already computed forp[0]on line 75. Consider combining the max-finding with the log-prob computation in a single pass after computinglog_sum.The static analysis warning about
std::log(sum_exp)at line 72 is a false positive sincesum_expis always positive whenvocab_size > 0(asexp()returns positive values).sherpa-onnx/c-api/c-api.cc (2)
808-814: LGTM!The destructor correctly handles the null case and properly frees both the inner array (
log_probs) withdelete[]and the struct itself withdelete, matching the allocation pattern in the getter functions.
757-780: No robustness issue here. Thevocab_log_probsdata structure is populated by decoders where each inner vector is created with exactlyvocab_sizeelements from model output, which is a fixed property per model. Since all tokens use the same model with the same output shape, all inner vectors are guaranteed to have identical sizes by design. Thestd::copyoperation is safe—if any vector had a different size, it would indicate data corruption in the decoder itself, which would be a more fundamental problem than a buffer overflow in this code.sherpa-onnx/csrc/offline-recognizer-moonshine-impl.h (1)
30-30: LGTM!The assignment correctly propagates token_log_probs from the decoder result to the recognition result, consistent with how other fields are copied.
sherpa-onnx/csrc/online-recognizer.h (1)
45-45: LGTM!The vocab_log_probs field addition enables per-token vocabulary distribution access for online recognition. The 2D structure (tokens × vocab_size) is appropriate for storing full distributions.
sherpa-onnx/csrc/online-transducer-decoder.h (1)
33-36: LGTM!The vocab_log_probs field addition is well-documented with clear shape specification and usage conditions. The 2D vector structure appropriately stores per-token vocabulary distributions.
sherpa-onnx/csrc/offline-stream.h (1)
52-54: LGTM!The addition of token_log_probs (per-token confidence) and vocab_log_probs (full vocabulary distributions) to OfflineRecognitionResult enables comprehensive confidence scoring and uncertainty estimation as described in the PR objectives.
flutter/sherpa_onnx/lib/src/offline_stream.dart (1)
36-69: LGTM!The implementation correctly:
- Validates function pointers and returned data
- Calculates flattened array indices (line 59)
- Cleans up native resources (line 67)
- Handles null cases gracefully
The Map<String, List> structure with "token_i" keys provides clear token identification.
sherpa-onnx/csrc/offline-recognizer-ctc-impl.h (1)
34-62: LGTM!The implementation correctly maintains alignment between
r.tokensandr.token_log_probs. When tokens are filtered (e.g., SIL tokens at line 41), both the token and its corresponding log probability are skipped together, preserving the one-to-one correspondence.sherpa-onnx/csrc/offline-recognizer-whisper-impl.h (1)
157-172: The token/log_prob alignment is correctly preserved through filtering.The loop correctly maintains 1:1 correspondence between
r.tokensandr.token_log_probs: when a token is filtered by!sym_table.Contains(i), the loop continues without appending to either vector, so both are skipped symmetrically. The guard on line 169 is defensive but unnecessary given the decoder's guarantee of equal-sized source vectors. No alignment issue exists.sherpa-onnx/csrc/offline-whisper-decoder.h (1)
16-24: LGTM!The new
token_log_probsandvocab_log_probsfields are correctly typed and appropriately placed in the result struct to capture per-token and full-vocabulary log probability distributions.sherpa-onnx/csrc/online-recognizer-transducer-nemo-impl.h (2)
52-59: LGTM!The
temperature_scaleparameter is correctly plumbed fromconfig_to the decoder constructor, consistent with the updatedOnlineTransducerGreedySearchNeMoDecodersignature.
72-79: LGTM!The templated constructor path correctly mirrors the primary constructor's handling of
temperature_scale.sherpa-onnx/csrc/online-transducer-greedy-search-decoder.cc (1)
142-160: LGTM!The per-token log probability export and full vocabulary distribution storage are correctly implemented with a single pass through temperature scaling and
LogSoftmax. This serves as the correct pattern that other decoders (like the NeMo decoder) should follow.flutter/sherpa_onnx/lib/src/offline_recognizer.dart (3)
610-613: LGTM!The JSON parsing for
token_log_probscorrectly handles null values with a safe default empty list and properly converts numeric values to doubles.
843-869: LGTM!The
getResultmethod correctly handlestokenLogProbsin both the error path (empty result) and the success path (parsed from JSON), ensuring consistency.
597-604: All instantiations ofOfflineRecognizerResultalready provide thetokenLogProbsparameter. The three usages in the file (fromJson factory at line 607, empty result at line 848, and parsedJson construction at line 862) all include this parameter. No breaking changes were introduced.flutter/sherpa_onnx/lib/src/online_recognizer.dart (2)
329-330: LGTM!The
ysProbsparameter is optional with a sensible default value, which avoids breaking existing code while supporting the new feature.
508-510: LGTM!The parsing of
ysProbsingetResultcorrectly handles null values with a safe default empty list.sherpa-onnx/csrc/offline-whisper-greedy-search-decoder.cc (2)
18-45: LGTM with a note on the static analysis hint.The
CalculateLogProbfunction correctly implements log-softmax using the log-sum-exp trick for numerical stability.Note: The static analysis tool flags line 44 as "Invalid std", but this appears to be a false positive.
std::logis a valid function from<cmath>(included on line 8).
204-208: LGTM!The addition of
IsMultiLingual()guard before accessingid2langis good defensive programming and prevents potential issues with non-multilingual models.sherpa-onnx/csrc/offline-recognizer-canary-impl.h (3)
60-64: LGTM!The early return for empty encoder output is good defensive programming and prevents potential crashes downstream.
146-152: LGTM!The
Convertmethod signature update to includetoken_log_probsis consistent with changes in other files, and the assignment is straightforward.
218-220: LGTM!The addition of an early return when the feature vector is empty prevents potential issues downstream and aligns with the empty check in
DecodeStream.Also applies to: 235-237
flutter/sherpa_onnx/lib/src/sherpa_onnx_bindings.dart (3)
680-689: LGTM!The
SherpaOnnxVocabLogProbsstruct is correctly defined for FFI with appropriate types. The flattened 2D array design (viaPointer<Float>) is the right approach for C interop.
1433-1445: LGTM!The typedefs for vocab log probs functions are correctly defined with appropriate Native and Dart versions.
2325-2341: LGTM!The initialization of vocab log probs bindings follows the existing pattern in the file and correctly uses
lookupFunctionto bind the native functions.
There was a problem hiding this comment.
Actionable comments posted: 1
Caution
Some comments are outside the diff and can’t be posted inline due to platform limitations.
⚠️ Outside diff range comments (1)
sherpa-onnx/csrc/offline-recognizer-canary-impl.h (1)
150-173: Critical: Token log probabilities misalignment when tokens are skipped.The method assigns
token_log_probswholesale at line 155, but then iterates through tokens and skips some (line 162continuewhen!symbol_table_.Contains(token_id)). This breaks the correspondence betweenr.tokensandr.token_log_probsindices. The same issue affectsvocab_log_probsalignment (set at line 131).Apply this diff to fix the alignment:
OfflineRecognitionResult Convert( const std::vector<int32_t> &tokens, const std::vector<float> &token_log_probs) const { OfflineRecognitionResult r; r.tokens.reserve(tokens.size()); - r.token_log_probs = token_log_probs; std::string text; for (size_t idx = 0; idx < tokens.size(); ++idx) { int32_t token_id = tokens[idx]; if (!symbol_table_.Contains(token_id)) { continue; } const auto &s = symbol_table_[token_id]; text += s; r.tokens.push_back(s); + + if (idx < token_log_probs.size()) { + r.token_log_probs.push_back(token_log_probs[idx]); + } } r.text = std::move(text); return r; }Note: You'll also need to filter
vocab_log_probssimilarly inDecodeStreambefore assigning tor.vocab_log_probsat line 131, or ensure tokens are never skipped by guaranteeing all generated tokens exist in the symbol table.
🧹 Nitpick comments (4)
sherpa-onnx/csrc/offline-whisper-greedy-search-decoder.cc (1)
121-145: Log-softmax implementation is correct and addresses previous review feedback.The current implementation computes log-softmax once per iteration and then extracts the log probability for the selected token, which addresses the redundant computation concern raised in the previous review.
Optional refactoring opportunities:
Max finding consistency: Lines 123-128 manually find the max logit, while lines 101-102 use
std::max_element. Consider usingstd::max_elementhere as well for consistency:float max_logit = *std::max_element(current_logits, current_logits + vocab_size);Allocation efficiency:
full_vocab_probsis allocated in every loop iteration (line 122). For large vocabularies, consider pre-allocating it outside the loop:std::vector<float> full_vocab_probs(vocab_size); for (int32_t i = 0; i < num_possible_tokens; ++i) { // ... reuse full_vocab_probs in each iterationNote on static analysis: The Cppcheck warning on line 134 about "Invalid std" is a false positive—
std::logis a valid function from<cmath>.sherpa-onnx/csrc/offline-moonshine-greedy-search-decoder.cc (1)
65-96: Consider optimizing the duplicate log-softmax computation.The log-sum-exp implementation is correct and numerically stable. However, the log-softmax is computed twice: once while finding the max token (lines 68-83) and again while storing the full vocabulary distribution (lines 92-96).
Apply this diff to compute log-softmax once:
- // Compute log-softmax and find max token with confidence - float max_logit = *std::max_element(p, p + vocab_size); - - float sum_exp = 0.0f; - for (int32_t j = 0; j < vocab_size; ++j) { - sum_exp += std::exp(p[j] - max_logit); - } - float log_sum = max_logit + std::log(sum_exp); - - int32_t max_token_id = 0; - float max_log_prob = p[0] - log_sum; - - for (int32_t j = 1; j < vocab_size; ++j) { - float log_prob = p[j] - log_sum; - if (log_prob > max_log_prob) { - max_log_prob = log_prob; - max_token_id = j; - } - } - - if (max_token_id == eos) { - break; - } - tokens.push_back(max_token_id); - token_log_probs.push_back(max_log_prob); - - // Store full vocabulary distribution (already log-softmaxed) - std::vector<float> full_vocab_probs(vocab_size); - for (int32_t j = 0; j < vocab_size; ++j) { - full_vocab_probs[j] = p[j] - log_sum; - } - vocab_log_probs.push_back(std::move(full_vocab_probs)); + // Compute log-softmax once for both max selection and storage + float max_logit = *std::max_element(p, p + vocab_size); + + float sum_exp = 0.0f; + for (int32_t j = 0; j < vocab_size; ++j) { + sum_exp += std::exp(p[j] - max_logit); + } + float log_sum = max_logit + std::log(sum_exp); + + // Compute log-softmax for all tokens and find max in single pass + std::vector<float> full_vocab_probs(vocab_size); + int32_t max_token_id = 0; + float max_log_prob = p[0] - log_sum; + full_vocab_probs[0] = max_log_prob; + + for (int32_t j = 1; j < vocab_size; ++j) { + float log_prob = p[j] - log_sum; + full_vocab_probs[j] = log_prob; + if (log_prob > max_log_prob) { + max_log_prob = log_prob; + max_token_id = j; + } + } + + if (max_token_id == eos) { + break; + } + tokens.push_back(max_token_id); + token_log_probs.push_back(max_log_prob); + vocab_log_probs.push_back(std::move(full_vocab_probs));flutter/sherpa_onnx/lib/src/online_stream.dart (1)
53-66: Consider adding defensive validation for native values.The method reads
numTokensandvocabSizefrom native code and immediately uses them in nested loops without validation. If native code returns negative or unexpectedly large values, this could cause issues.Apply this diff to add defensive checks:
final vocabLogProbs = vocabPtr.ref; final numTokens = vocabLogProbs.numTokens; final vocabSize = vocabLogProbs.vocabSize; + +if (numTokens < 0 || vocabSize < 0 || numTokens > 10000 || vocabSize > 100000) { + destroyFunc(vocabPtr); + return null; +} final Map<String, List<double>> result = {};sherpa-onnx/csrc/offline-recognizer-canary-impl.h (1)
101-119: Remove unused variabletokens_generated.The variable
tokens_generatedis declared and incremented but never read. It serves no purpose in the current implementation.Apply this diff:
if (max_token_id != eos) { - int32_t tokens_generated = 0; for (int32_t i = decoder_input.size() + 1; i <= decoder_input.size() + num_tokens; ++i) { if (tokens.back() == eos) { break; } std::tie(logits, decoder_states) = RunDecoder(tokens.back(), i, std::move(decoder_states), View(&enc_states), View(&enc_mask)); std::vector<float> next_full_vocab_probs; auto [next_token_id, next_confidence] = GetMaxTokenIdWithConfidence(&logits, &next_full_vocab_probs); tokens.push_back(next_token_id); token_log_probs.push_back(next_confidence); vocab_log_probs.push_back(std::move(next_full_vocab_probs)); - tokens_generated++; } }
📜 Review details
Configuration used: CodeRabbit UI
Review profile: CHILL
Plan: Pro
📒 Files selected for processing (10)
flutter/sherpa_onnx/lib/src/online_recognizer.dart(3 hunks)flutter/sherpa_onnx/lib/src/online_stream.dart(1 hunks)flutter/sherpa_onnx/pubspec.yaml(1 hunks)sherpa-onnx/csrc/offline-ctc-decoder.h(1 hunks)sherpa-onnx/csrc/offline-moonshine-greedy-search-decoder.cc(4 hunks)sherpa-onnx/csrc/offline-recognizer-canary-impl.h(9 hunks)sherpa-onnx/csrc/offline-recognizer-moonshine-impl.h(1 hunks)sherpa-onnx/csrc/offline-recognizer-whisper-impl.h(2 hunks)sherpa-onnx/csrc/offline-whisper-greedy-search-decoder.cc(3 hunks)sherpa-onnx/csrc/online-transducer-greedy-search-nemo-decoder.cc(3 hunks)
🚧 Files skipped from review as they are similar to previous changes (5)
- sherpa-onnx/csrc/offline-recognizer-moonshine-impl.h
- sherpa-onnx/csrc/offline-ctc-decoder.h
- flutter/sherpa_onnx/lib/src/online_recognizer.dart
- sherpa-onnx/csrc/offline-recognizer-whisper-impl.h
- flutter/sherpa_onnx/pubspec.yaml
🧰 Additional context used
🧬 Code graph analysis (4)
sherpa-onnx/csrc/offline-moonshine-greedy-search-decoder.cc (1)
sherpa-onnx/csrc/offline-whisper-greedy-search-decoder.cc (1)
full_vocab_probs(122-122)
sherpa-onnx/csrc/offline-recognizer-canary-impl.h (3)
sherpa-onnx/csrc/offline-canary-model.cc (2)
tokens(83-117)tokens(83-85)sherpa-onnx/csrc/offline-recognizer-ctc-impl.h (2)
Convert(26-323)- `` (193-199)
sherpa-onnx/csrc/offline-stream.cc (2)
r(257-257)r(257-257)
sherpa-onnx/csrc/offline-whisper-greedy-search-decoder.cc (1)
sherpa-onnx/csrc/offline-moonshine-greedy-search-decoder.cc (1)
full_vocab_probs(92-92)
sherpa-onnx/csrc/online-transducer-greedy-search-nemo-decoder.cc (3)
sherpa-onnx/csrc/online-stream.cc (8)
r(68-68)r(68-68)r(72-74)r(72-72)r(93-93)r(93-93)r(95-97)r(95-95)sherpa-onnx/csrc/offline-whisper-greedy-search-decoder.cc (1)
full_vocab_probs(122-122)sherpa-onnx/csrc/offline-moonshine-greedy-search-decoder.cc (1)
full_vocab_probs(92-92)
🪛 Cppcheck (2.18.0)
sherpa-onnx/csrc/offline-moonshine-greedy-search-decoder.cc
[warning] 72-72: Invalid std
(invalidFunctionArg)
sherpa-onnx/csrc/offline-whisper-greedy-search-decoder.cc
[warning] 134-134: Invalid std
(invalidFunctionArg)
⏰ Context from checks skipped due to timeout of 90000ms. You can increase the timeout in your CodeRabbit configuration to a maximum of 15 minutes (900000ms). (20)
- GitHub Check: Release shared tts-OFF
- GitHub Check: Debug static tts-OFF
- GitHub Check: Release shared tts-ON
- GitHub Check: Debug shared-OFF tts-ON
- GitHub Check: Debug shared-ON tts-ON
- GitHub Check: Debug shared-OFF tts-OFF
- GitHub Check: Release shared-ON tts-OFF
- GitHub Check: Debug shared-ON tts-OFF
- GitHub Check: Release shared-ON tts-ON
- GitHub Check: Release shared-OFF tts-OFF
- GitHub Check: Release shared-OFF tts-ON
- GitHub Check: swift (macos-latest)
- GitHub Check: rknn shared ON
- GitHub Check: rknn shared OFF
- GitHub Check: ubuntu-24.04 3.13
- GitHub Check: ubuntu-24.04 3.12
- GitHub Check: ubuntu-24.04 3.11
- GitHub Check: ubuntu-24.04 3.9
- GitHub Check: ubuntu-24.04 3.8
- GitHub Check: ubuntu-24.04 3.10
🔇 Additional comments (16)
sherpa-onnx/csrc/online-transducer-greedy-search-nemo-decoder.cc (2)
90-99: Looks good: single-pass temperature + LogSoftmax, and full-vocab log-probs capture.
This also appears to address the previously-reported duplicate “export per-token log scores” block (only one remains now).
138-139: Call site threadingtemperature_scale_intoDecodeOne()is consistent.
Just ensuretemperature_scale_is validated at config boundaries too (so all decoders behave consistently).sherpa-onnx/csrc/offline-whisper-greedy-search-decoder.cc (5)
8-9: LGTM!The new includes are necessary for the log-softmax computation (
std::log,std::exp) and max logit initialization (std::numeric_limits).
92-103: LGTM!The variable renaming (
logits_tensor,current_token_id) improves code clarity, and the initialization logic correctly identifies the first token usingstd::max_element.
108-110: LGTM!The new data structures for tracking per-token log probabilities and full vocabulary distributions are properly declared and will be populated in the decoding loop.
147-176: LGTM!The loop continuation logic correctly:
- Creates the next token input tensor with the selected token
- Forwards through the decoder with proper state propagation
- Updates the offset and checks context boundaries
- Computes the next token for the subsequent iteration
178-189: LGTM!The result assembly correctly:
- Guards language lookup with
IsMultiLingual()check- Populates the predicted tokens
- Populates the new
token_log_probsandvocab_log_probsfields usingstd::movefor efficiencysherpa-onnx/csrc/offline-moonshine-greedy-search-decoder.cc (4)
8-8: LGTM!The
<cmath>include is necessary forstd::logandstd::expused in the log-softmax computation.
42-43: LGTM!The vector declarations appropriately support the PR's objective of exposing per-token log-probabilities and full vocabulary distributions.
119-120: LGTM!The result population correctly uses
std::moveto efficiently transfer the computed log-probability data to the result structure.
66-72: The cppcheck warning "Invalid std (invalidFunctionArg)" at line 72 is a false positive.The
std::log(sum_exp)call is valid—sum_expis guaranteed to be strictly positive as it's a sum ofstd::exp()values. This log-sum-exp pattern is used consistently throughout the codebase without issues. No action is required.sherpa-onnx/csrc/offline-recognizer-canary-impl.h (5)
60-64: LGTM! Good defensive guard.The early return for empty encoder output prevents potential crashes or undefined behavior in downstream processing.
79-92: LGTM! Vocab log probs now properly collected.This addresses the previous review feedback by capturing the full vocabulary distribution for the first token. The three parallel vectors (tokens, token_log_probs, vocab_log_probs) are correctly initialized.
123-128: LGTM! Proper cleanup of EOS token.Correctly removes the EOS token and its associated confidence and vocab distribution from all three parallel vectors, maintaining alignment.
175-213: LGTM! Correct log-softmax implementation.The method correctly computes log-softmax using the log-sum-exp trick for numerical stability. The optional full distribution capture and max token identification are both implemented correctly.
222-224: LGTM! Good defensive guard.The early return for empty feature vectors prevents potential downstream errors and is consistent with the empty encoder output guard in
DecodeStream.
|
I don't know what a style check is, but otherwise I think this PR is ready. |
csukuangfj
left a comment
There was a problem hiding this comment.
You have a lot of redundant code. Please try to use existing functions to simplify it.
| double sum_exp = 0.0; | ||
| for (int32_t j = 0; j < vocab_size; ++j) { | ||
| sum_exp += std::exp(p[j] - max_logit); | ||
| } | ||
| float log_sum = max_logit + static_cast<float>(std::log(sum_exp)); | ||
|
|
||
| // Compute log-softmax for all tokens and find max in single pass | ||
| std::vector<float> full_vocab_probs(vocab_size); | ||
| int32_t max_token_id = 0; | ||
| float max_log_prob = p[0] - log_sum; | ||
| full_vocab_probs[0] = max_log_prob; | ||
|
|
||
| for (int32_t j = 1; j < vocab_size; ++j) { | ||
| float log_prob = p[j] - log_sum; | ||
| full_vocab_probs[j] = log_prob; | ||
| if (log_prob > max_log_prob) { | ||
| max_log_prob = log_prob; | ||
| max_token_id = j; | ||
| } | ||
| } |
There was a problem hiding this comment.
Too complicated. Please use existing functions to simplify it.
| OfflineRecognitionResult empty_result; | ||
| s->SetResult(empty_result); |
There was a problem hiding this comment.
| OfflineRecognitionResult empty_result; | |
| s->SetResult(empty_result); |
Just return is enough. Every stream has a default initialized result.
Co-Authored-By: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
Use existing LogSoftmax() and MaxElementIndex() from math.h instead of manual implementations in moonshine, whisper, and canary decoders. Extract FlattenVocabLogProbs helper in C API to deduplicate identical flatten logic. Remove unnecessary reserve calls, redundant emptiness guards, and dead default parameters. Co-Authored-By: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
| auto max_it = std::max_element( | ||
| static_cast<const float *>(p_log_probs), | ||
| std::max_element( | ||
| static_cast<const float *>(p_log_probs), | ||
| static_cast<const float *>(p_log_probs) + vocab_size))); | ||
| static_cast<const float *>(p_log_probs) + vocab_size); | ||
| auto y = static_cast<int64_t>(std::distance( | ||
| static_cast<const float *>(p_log_probs), max_it)); | ||
| float log_prob = *max_it; |
There was a problem hiding this comment.
Can you use the following function
sherpa-onnx/sherpa-onnx/csrc/math.h
Line 170 in b95ec53
Once you get the max element index, you can use p_log_probs[index] to get the log prob.
The current code is a bit complex.
- Use MaxElementIndex instead of std::max_element in CTC decoder - Add size validation in FlattenVocabLogProbs to prevent OOB on mismatched vectors - Remove redundant size guards in Convert functions (vectors built in lockstep) - Simplify Canary GetMaxTokenIdWithConfidence (remove unused nullptr default) - Deduplicate Dart getVocabLogProbs into shared readAndFreeVocabLogProbs helper - Change getVocabLogProbs return type from Map<String, List> to List<List> - Fix missing tokenLogProbs in offline recognizer error path Co-Authored-By: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
Co-Authored-By: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
OfflineCtcFstDecoder does not populate token_log_probs, so the Convert function must check bounds before indexing. Co-Authored-By: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
MoonshineV2 greedy search decoder does not populate token_log_probs or vocab_log_probs, so the shared Convert function must check bounds. Co-Authored-By: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
| SHERPA_ONNX_API const SherpaOnnxVocabLogProbs * | ||
| SherpaOnnxOnlineStreamGetVocabLogProbs(const SherpaOnnxOnlineStream *stream); | ||
|
|
||
| SHERPA_ONNX_API const SherpaOnnxVocabLogProbs * | ||
| SherpaOnnxOfflineStreamGetVocabLogProbs(const SherpaOnnxOfflineStream *stream); | ||
|
|
||
| SHERPA_ONNX_API void SherpaOnnxDestroyVocabLogProbs( | ||
| const SherpaOnnxVocabLogProbs *log_probs); |
There was a problem hiding this comment.
Can you put the log_probs inside the result struct, like how we return timestamps?
Captures full vocabulary distribution for CTC (SenseVoice, MedASR, paraformer, zipformer) and offline NeMo transducer (parakeet-tdt, nemotron, nemo-fastconformer) decoders. Propagates through Convert functions and exposes via Python pybind11 binding. Addresses reviewer feedback to put log_probs inside the result struct. Co-Authored-By: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
- Add vocab_log_probs and vocab_size fields to both SherpaOnnxOnlineRecognizerResult and SherpaOnnxOfflineRecognizerResult, following the same pattern as timestamps - Remove the separate SherpaOnnxVocabLogProbs struct and its getter/destroy functions from the C API - Update Dart bindings to read vocab_log_probs from the result struct via recognizer.getVocabLogProbs(stream) instead of stream.getVocabLogProbs() Co-Authored-By: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
Remove duplicate token_log_probs field from OfflineRecognitionResult. All decoder convert functions now write to ys_log_probs (the field exposed by Python/Dart bindings). CTC, SenseVoice, Whisper, Moonshine, and Canary all use the same field now. JSON serialization still outputs "token_log_probs" key for backward compatibility with Dart bindings. Co-Authored-By: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
Co-Authored-By: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
OfflineRecognitionResult had both ys_log_probs (upstream) and token_log_probs (our addition) for the same data. Removed the duplicate. All convert functions now write to ys_log_probs only. JSON and Dart also use ys_log_probs key. Verified: SenseVoice returns 960 tokens with ys_log_probs and vocab_log_probs populated correctly after clean rebuild. Co-Authored-By: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
|
|
||
| final class SherpaOnnxSpokenLanguageIdentification extends Opaque {} | ||
|
|
||
| /// FFI struct for the online recognizer result (mirrors C API struct layout). |
There was a problem hiding this comment.
You don't need to add a new class.
Please use our existing method to parse the result, i.e., parse the json result.
|
Closing this. I built this API to support confidence-weighted ROVER fusion for clinical ASR — the idea being that richer per-token distributions should give better voting than flat equal-weight voting. After testing seven fusion methods on clinical conversation audio (naive ROVER, CNC with epsilon tuning, confidence-weighted, Shannon/Tsallis entropy-weighted, margin-weighted, and a learned per-word GBM classifier with 5-fold cross-validation), the methods that actually consume vocab log-probs all landed within ~0.1pp of naive equal-weight ROVER (13.30–13.41% vs. 13.41% WER on my datasets). The best practical result came from a method that doesn't use vocab log-probs at all (epsilon-tuned CNC, 13.05%). The oracle upper bound is 9.07% — a ~4pp gap no voting scheme we tested could close, because roughly two-thirds of the best model's deletions are also errors in every other model in our set. The residual appears to be a structural floor rather than something fusion can fix. Now I am using this for medical audio, and I suppose other datasets could give different results. However, since I can't demonstrate a measurable benefit for my own motivating use case, I decided to close this. I'll keep the branch in the very off chance someone else finds something about it useful. Thanks to the maintainers for the review attention along the way. |
Preserving in-progress changes to the vocab_log_probs API before archiving the branch. PR k2-fsa#2897 closed after testing showed confidence-weighted and learned ROVER variants do not meaningfully exceed naive equal-weight ROVER for clinical ASR fusion.
Vocabulary Log Probabilities API
Summary
I needed this for fusion from two different models, and its worked ok for me so far. I think there was also another PR/Issues that was asking about this. I tried to add API for accessing full vocabulary log probability, to use with advanced confidence scoring, entropy-based uncertainty estimation, and hypothesis fusion algorithms (e.g., weighted ROVER). This exposes the complete probability distribution over the entire vocabulary for each decoded token position.
Changes
Core API (
sherpa-onnx/c-api/)New C API Functions:
SherpaOnnxOnlineStreamGetVocabLogProbs()- full vocab log probs online streamsSherpaOnnxOfflineStreamGetVocabLogProbs()- full vocab log probs offline streamsSherpaOnnxDestroyVocabLogProbs()- Memory management for vocab log probsNew Data Structure:
Model Support
Online Models:
Offline Models:
Implementation Details
Core Changes:
vocab_log_probsfield toOfflineRecognitionResultandOnlineRecognizerResultstructuresoffline-whisper-greedy-search-decoder.cc- Captures joiner output logitsoffline-moonshine-greedy-search-decoder.cc- Captures model output distributionsonline-transducer-greedy-search-decoder.cc- Captures joiner logits per frameonline-transducer-greedy-search-nemo-decoder.cc- NeMo-specific implementationoffline-recognizer-canary-impl.h- Canary model supportoffline-recognizer-ctc-impl.h- CTC model supportFlutter/Dart Bindings:
getVocabLogProbs()methods toOnlineStreamandOfflineStreamsherpa_onnx_bindings.dartJSON Serialization:
token_log_probsfield to JSON output (for backward compatibility)ys_log_probsfieldUse - I'm using it for Confidence-Weighted ROVER Fusion mostly (but can also be used for Entropy-Based Uncertainty Quantification).
Testing - I'm not sure how you run tests on this repo, I couldn't really find a test suite, so I'm not sure if I broke anything (in my local testing it didn't seem to, but I certainly didn't test it with all of the platforms and programming languages that you use)
Summary by CodeRabbit