Support dynamic decoder layers in canary model runtime - #3268
Conversation
Summary of ChangesHello, 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 enhances the Highlights
Changelog
Activity
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
|
📝 WalkthroughWalkthroughAdded two new metadata fields to OfflineCanaryModelMetaData and updated initialization to use those fields for decoder layer count and hidden size instead of hardcoded values. Changes
Sequence DiagramsequenceDiagram
participant Client
participant InitEncoder
participant Metadata
participant GetInitialDecoderStates
participant Decoder
Client->>InitEncoder: Initialize model
InitEncoder->>Metadata: Read "num_decoder_layers"
Metadata-->>InitEncoder: Value or default (6)
InitEncoder->>Metadata: Read "decoder_hidden_size"
Metadata-->>InitEncoder: Value or default (1024)
InitEncoder->>GetInitialDecoderStates: Invoke with updated meta_
GetInitialDecoderStates->>GetInitialDecoderStates: Loop num_decoder_layers times
GetInitialDecoderStates->>Decoder: Allocate per-layer tensor (shape uses decoder_hidden_size)
Decoder-->>GetInitialDecoderStates: Return tensor
GetInitialDecoderStates-->>InitEncoder: Return initial decoder states
Estimated code review effort🎯 3 (Moderate) | ⏱️ ~20 minutes 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 |
There was a problem hiding this comment.
Code Review
This pull request effectively removes hardcoded values for decoder layers and hidden size in the canary model, allowing for dynamic configuration from ONNX metadata, which is a great improvement for model flexibility. However, the implementation lacks proper validation for the values read from the model's metadata, specifically num_decoder_layers and decoder_hidden_size. This could lead to a Denial of Service (DoS) if a malformed or malicious model file is loaded due to unchecked negative or excessively large values. My review includes suggestions to enhance the robustness of the metadata parsing by using more specific exception handling and adding validation for these parsed values.
| try { | ||
| auto num_layers_str = meta_data.LookupCustomMetadataMapAllocated( | ||
| "num_decoder_layers", alloc); | ||
| if (num_layers_str) { | ||
| meta_.num_decoder_layers = std::stoi(num_layers_str.get()); | ||
| } | ||
| } catch (...) { | ||
| // Use default (6) if not present | ||
| } | ||
| try { | ||
| auto hidden_size_str = meta_data.LookupCustomMetadataMapAllocated( | ||
| "decoder_hidden_size", alloc); | ||
| if (hidden_size_str) { | ||
| meta_.decoder_hidden_size = std::stoll(hidden_size_str.get()); | ||
| } | ||
| } catch (...) { | ||
| // Use default (1024) if not present | ||
| } | ||
| } |
There was a problem hiding this comment.
The application reads num_decoder_layers from the ONNX model's custom metadata without validating its range. This is a critical security vulnerability: negative values can cause std::vector::reserve to throw std::length_error (Denial of Service), and excessively large values can lead to memory exhaustion (std::bad_alloc). The current catch(...) block is too broad; it should specifically catch const std::exception & for parsing errors from std::stoi. The parsed num_decoder_layers must be explicitly validated to be a positive number before assignment, similar to how other metadata fields are handled by SHERPA_ONNX_READ_META_DATA.
try {
auto num_layers_str = meta_data.LookupCustomMetadataMapAllocated(
"num_decoder_layers", alloc);
if (num_layers_str) {
int32_t num_layers = std::stoi(num_layers_str.get());
if (num_layers > 0) {
meta_.num_decoder_layers = num_layers;
}
}
} catch (const std::exception &) {
// Use default if not present or on parsing error
}| try { | ||
| auto hidden_size_str = meta_data.LookupCustomMetadataMapAllocated( | ||
| "decoder_hidden_size", alloc); | ||
| if (hidden_size_str) { | ||
| meta_.decoder_hidden_size = std::stoll(hidden_size_str.get()); | ||
| } | ||
| } catch (...) { | ||
| // Use default (1024) if not present | ||
| } |
There was a problem hiding this comment.
Similar to the num_decoder_layers parsing, catch(...) should be replaced with catch (const std::exception &) for more specific error handling. The parsed decoder_hidden_size should also be validated to be positive.
try {
auto hidden_size_str = meta_data.LookupCustomMetadataMapAllocated(
"decoder_hidden_size", alloc);
if (hidden_size_str) {
int64_t hidden_size = std::stoll(hidden_size_str.get());
if (hidden_size > 0) {
meta_.decoder_hidden_size = hidden_size;
}
}
} catch (const std::exception &) {
// Use default if not present or on parsing error
}ac83733 to
71e0750
Compare
| SHERPA_ONNX_READ_META_DATA(meta_.subsampling_factor, "subsampling_factor"); | ||
| SHERPA_ONNX_READ_META_DATA(meta_.feat_dim, "feat_dim"); | ||
|
|
||
| // Read decoder architecture metadata (with defaults for backward compat) |
There was a problem hiding this comment.
Please use
sherpa-onnx/sherpa-onnx/csrc/macros.h
Line 72 in 33554f7
See its usage at
sherpa-onnx/sherpa-onnx/csrc/vocos-vocoder.cc
Line 162 in 33554f7
There was a problem hiding this comment.
thanks, updated to use the existing macro
There was a problem hiding this comment.
Actionable comments posted: 1
🧹 Nitpick comments (1)
sherpa-onnx/csrc/offline-canary-model-meta-data.h (1)
17-18: Clarify whatnum_decoder_layerscounts.
GetInitialDecoderStates()consumes this as the number of decoder state tensors, but the PR rationale distinguishes 8 decoder blocks from 10 memory states. Please add an inline comment here (or rename the field if the metadata key is still flexible) so model exporters do not write the transformer-layer count by mistake.🤖 Prompt for AI Agents
Verify each finding against the current code and only fix it if needed. In `@sherpa-onnx/csrc/offline-canary-model-meta-data.h` around lines 17 - 18, The field num_decoder_layers is ambiguous: GetInitialDecoderStates() treats it as the number of decoder state tensors, not the count of transformer blocks, so add a clarifying inline comment (or rename the field to something explicit like num_decoder_state_tensors) above num_decoder_layers to state: "number of decoder state tensors consumed by GetInitialDecoderStates(), not the number of transformer layers/blocks." Update any documentation or exporter guidance if the metadata key is changed.
🤖 Prompt for all review comments with AI agents
Verify each finding against the current code and only fix it if needed.
Inline comments:
In `@sherpa-onnx/csrc/offline-canary-model.cc`:
- Around line 184-211: The decoder/encoder metadata must be validated against
the loaded decoder session: after parsing meta_.num_decoder_layers and
meta_.decoder_hidden_size, call decoder_sess_->GetInputCount() and
decoder_sess_->GetOutputCount() and verify they equal (3 +
meta_.num_decoder_layers) and (1 + meta_.num_decoder_layers) respectively; if
not, abort initialization with a clear error. Additionally, inspect the decoder
session input/output type/shape (using decoder_sess_->GetInputTypeInfo /
GetOutputTypeInfo and Ort::TensorTypeAndShapeInfo) for the state tensors and
confirm their hidden-dimension matches meta_.decoder_hidden_size, failing
initialization with a descriptive exception if it differs. Ensure errors
reference the symbols meta_.num_decoder_layers, meta_.decoder_hidden_size and
decoder_sess_ so the failure message makes the mismatch obvious.
---
Nitpick comments:
In `@sherpa-onnx/csrc/offline-canary-model-meta-data.h`:
- Around line 17-18: The field num_decoder_layers is ambiguous:
GetInitialDecoderStates() treats it as the number of decoder state tensors, not
the count of transformer blocks, so add a clarifying inline comment (or rename
the field to something explicit like num_decoder_state_tensors) above
num_decoder_layers to state: "number of decoder state tensors consumed by
GetInitialDecoderStates(), not the number of transformer layers/blocks." Update
any documentation or exporter guidance if the metadata key is changed.
ℹ️ Review info
⚙️ Run configuration
Configuration used: defaults
Review profile: CHILL
Plan: Pro
Run ID: 83735955-339f-4a6e-a36d-7eff7e86854c
📒 Files selected for processing (2)
sherpa-onnx/csrc/offline-canary-model-meta-data.hsherpa-onnx/csrc/offline-canary-model.cc
| // Read decoder architecture metadata (with defaults for backward compat) | ||
| { | ||
| Ort::AllocatorWithDefaultOptions alloc; | ||
| try { | ||
| auto num_layers_str = meta_data.LookupCustomMetadataMapAllocated( | ||
| "num_decoder_layers", alloc); | ||
| if (num_layers_str) { | ||
| int32_t num_layers = std::stoi(num_layers_str.get()); | ||
| if (num_layers > 0) { | ||
| meta_.num_decoder_layers = num_layers; | ||
| } | ||
| } | ||
| } catch (const std::exception &) { | ||
| // Use default (6) if not present or on parsing error | ||
| } | ||
| try { | ||
| auto hidden_size_str = meta_data.LookupCustomMetadataMapAllocated( | ||
| "decoder_hidden_size", alloc); | ||
| if (hidden_size_str) { | ||
| int64_t hidden_size = std::stoll(hidden_size_str.get()); | ||
| if (hidden_size > 0) { | ||
| meta_.decoder_hidden_size = hidden_size; | ||
| } | ||
| } | ||
| } catch (const std::exception &) { | ||
| // Use default (1024) if not present or on parsing error | ||
| } | ||
| } |
There was a problem hiding this comment.
Fail fast when encoder metadata and the decoder graph disagree.
These parsed values now drive the decoder-state contract, but nothing checks that the separately loaded decoder actually exposes the same number of state inputs/outputs. A mismatched encoder/decoder pair or malformed metadata will load successfully and only fail later inside decoder_sess_->Run(). Please validate this during initialization (e.g. 3 + num_decoder_layers inputs and 1 + num_decoder_layers outputs, plus the state hidden size if you can inspect it) and abort with a clear error when it does not match.
🤖 Prompt for AI Agents
Verify each finding against the current code and only fix it if needed.
In `@sherpa-onnx/csrc/offline-canary-model.cc` around lines 184 - 211, The
decoder/encoder metadata must be validated against the loaded decoder session:
after parsing meta_.num_decoder_layers and meta_.decoder_hidden_size, call
decoder_sess_->GetInputCount() and decoder_sess_->GetOutputCount() and verify
they equal (3 + meta_.num_decoder_layers) and (1 + meta_.num_decoder_layers)
respectively; if not, abort initialization with a clear error. Additionally,
inspect the decoder session input/output type/shape (using
decoder_sess_->GetInputTypeInfo / GetOutputTypeInfo and
Ort::TensorTypeAndShapeInfo) for the state tensors and confirm their
hidden-dimension matches meta_.decoder_hidden_size, failing initialization with
a descriptive exception if it differs. Ensure errors reference the symbols
meta_.num_decoder_layers, meta_.decoder_hidden_size and decoder_sess_ so the
failure message makes the mismatch obvious.
There was a problem hiding this comment.
♻️ Duplicate comments (1)
sherpa-onnx/csrc/offline-canary-model.cc (1)
184-187:⚠️ Potential issue | 🟠 MajorFail fast on invalid or mismatched decoder-state metadata.
meta_.num_decoder_layersandmeta_.decoder_hidden_sizenow define how many state tensors you allocate and pass intodecoder_sess_, but initialization still accepts0-valued metadata and never checks that the loaded decoder exposes the same number/shape of state inputs and outputs. A malformed metadata block or mismatched encoder/decoder pair will still load and only fail later inRun(). Please validate both fields are> 0and verify the decoder I/O contract during initialization.🤖 Prompt for AI Agents
Verify each finding against the current code and only fix it if needed. In `@sherpa-onnx/csrc/offline-canary-model.cc` around lines 184 - 187, Validate decoder metadata immediately after reading it: check meta_.num_decoder_layers and meta_.decoder_hidden_size are > 0 (instead of allowing 0) and fail fast (return error/throw/log) if not; then inspect decoder_sess_’s input and output signatures during initialization (e.g., in the same init function that calls SHERPA_ONNX_READ_META_DATA_WITH_DEFAULT) and verify the number and shapes of state inputs/outputs match what you will allocate/use (use meta_.num_decoder_layers and meta_.decoder_hidden_size to compute expected tensors) and return a clear initialization error if there is any mismatch so Run() cannot proceed with wrong decoder I/O.
🤖 Prompt for all review comments with AI agents
Verify each finding against the current code and only fix it if needed.
Duplicate comments:
In `@sherpa-onnx/csrc/offline-canary-model.cc`:
- Around line 184-187: Validate decoder metadata immediately after reading it:
check meta_.num_decoder_layers and meta_.decoder_hidden_size are > 0 (instead of
allowing 0) and fail fast (return error/throw/log) if not; then inspect
decoder_sess_’s input and output signatures during initialization (e.g., in the
same init function that calls SHERPA_ONNX_READ_META_DATA_WITH_DEFAULT) and
verify the number and shapes of state inputs/outputs match what you will
allocate/use (use meta_.num_decoder_layers and meta_.decoder_hidden_size to
compute expected tensors) and return a clear initialization error if there is
any mismatch so Run() cannot proceed with wrong decoder I/O.
ℹ️ Review info
⚙️ Run configuration
Configuration used: defaults
Review profile: CHILL
Plan: Pro
Run ID: 2930c052-fbb0-48c7-8fa9-cfe9383f911b
📒 Files selected for processing (2)
sherpa-onnx/csrc/offline-canary-model-meta-data.hsherpa-onnx/csrc/offline-canary-model.cc
csukuangfj
left a comment
There was a problem hiding this comment.
Thank you for your contribution!
|
@mm65x Thank you so much for your effort supporting canary-1b-v2. Do you know by any chance if the model will be available for download? I tried to generate ONNX models for canary-1b-flash or canary-1b-v2, but I don’t have enough skills and RAM to make it happen. |
Enables canary-1b-v2 which uses 10 decoder memory states (8 transformer layers + initial + final layer norm).
Related to #3190 and #3193.
Problem:
GetInitialDecoderStates()in the canary model runtime hardcodes 6 decoder layers and 1024 hidden size. This works for canary-180m-flash (6 layers) but fails for larger canary models e.g. canary-1b-v2 which has 10 decoder memory states (8 transformer layers + initial + final layer norm).This PR makes it read
num_decoder_layersanddecoder_hidden_sizefrom the ONNX encoder metadata instead of hardcoding. Defaults to 6/1024 for backward compatibility with existing 180m models that don't have these metadata fields.The PR introduces changes to the following files:
sherpa-onnx/csrc/offline-canary-model-meta-data.h: addnum_decoder_layers,decoder_hidden_sizefields with defaultssherpa-onnx/csrc/offline-canary-model.cc: read from metadata inInitEncoder(), use inGetInitialDecoderStates()Summary by CodeRabbit