Conversation
Summary of ChangesHello @ssteo, 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 significantly enhances the ONNX export capabilities for Canary models. It addresses a critical bug in the existing 180M-flash export script, ensuring its proper functionality. Concurrently, it expands support to the larger 1B-flash model by introducing dedicated export and testing infrastructure, allowing for efficient deployment and inference of this model in ONNX format. 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
|
📝 WalkthroughWalkthroughAdds ONNX export, quantization, and validation pipelines for NVIDIA Canary ASR (180M and 1B‑flash), new encoder/decoder wrappers and attention‑mask utilities, expanded post‑export verification including encoder ONNX files, and a runtime ONNX test harness with Kaldi fbank feature extraction. Changes
Sequence Diagram(s)sequenceDiagram
participant User
participant ExportScript as "Export Script\n(export_onnx_1b_flash.py)"
participant PyTorch as "PyTorch Model\n(nvidia/canary-1b-flash)"
participant ONNXExport as "torch.onnx.export"
participant Quantizer as "onnxruntime.quantize_dynamic"
participant ONNXRuntime as "ONNXRuntime\n(test_1b_flash.py)"
User->>ExportScript: run export_onnx_1b_flash.py
ExportScript->>PyTorch: load model
ExportScript->>ExportScript: wrap encoder/decoder (EncoderWrapper, DecoderWrapper)
ExportScript->>ONNXExport: export encoder.onnx, decoder.onnx
ONNXExport-->>ExportScript: write .onnx files
ExportScript->>Quantizer: quantize .onnx -> .int8.onnx
Quantizer-->>ExportScript: write .int8.onnx files
ExportScript->>ExportScript: add_meta_data, adjust IR
User->>ONNXRuntime: run test_1b_flash.py with audio
ONNXRuntime->>ONNXRuntime: extract fbank features
ONNXRuntime->>ONNXRuntime: run_encoder(features) => enc_states, enc_mask
loop decoding
ONNXRuntime->>ONNXRuntime: run_decoder(input_ids, mems, enc_states, enc_mask) => logits, updated_mems
ONNXRuntime->>ONNXRuntime: append token, check EOS
end
ONNXRuntime-->>User: print transcription
Estimated code review effort🎯 4 (Complex) | ⏱️ ~45 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 |
There was a problem hiding this comment.
Code Review
This pull request fixes an export error for the 180m-flash model and introduces export functionality for the 1b-flash model. The changes involve adjusting the ONNX export parameters and handling of decoder memory lists to ensure compatibility with ONNX runtime. A new script for exporting and testing the 1b-flash model has been added, along with an update to use uv pip install for dependency management in the run_180m_flash.sh script. Overall, the changes improve the ONNX export process and expand model support.
| "next_decoder_mem_list_5", | ||
| ], | ||
| dynamic_axes={ | ||
| "decoder_input_ids": {1: "num_tokens"}, |
There was a problem hiding this comment.
The decoder_input_ids was removed from dynamic_axes. If the second dimension (num_tokens) of decoder_input_ids can vary during inference, this could lead to issues with the exported ONNX model. Please confirm that decoder_input_ids will always have a fixed size for its second dimension in the ONNX graph, or if it should remain dynamic.
There was a problem hiding this comment.
Actionable comments posted: 5
🤖 Fix all issues with AI agents
Verify each finding against the current code and only fix it if needed.
In `@scripts/nemo/canary/export_onnx_180m_flash.py`:
- Line 263: The call to torch.onnx.export currently passes dynamo=False which
breaks on PyTorch <2.0; update the export call to detect and conditionally
include the dynamo kwarg by creating an inspect.signature of torch.onnx.export
(e.g., export_sig = inspect.signature(torch.onnx.export)), build a kwargs dict
and only set kwargs["dynamo"]=False when "dynamo" is in export_sig.parameters,
then pass those options into torch.onnx.export via **kwargs instead of
hardcoding dynamo=False.
In `@scripts/nemo/canary/export_onnx_1b_flash.py`:
- Around line 297-307: In export_tokens, accessing s[0] after s =
canary_model.tokenizer.ids_to_text([i]) can raise IndexError if ids_to_text
returns an empty string; update the logic to first ensure s is non-empty (e.g.,
check if s and len(s) > 0) before testing s[0], and handle the empty case
(replace with underline + "" or a safe placeholder, or write the token ID alone)
so the loop never indexes into an empty string while still applying the
underline substitution for leading spaces.
- Around line 1-384: This file duplicates many utilities from
export_onnx_180m_flash.py; extract the shared functions and classes
(fixed_form_attention_mask, add_meta_data, lens_to_mask, EncoderWrapper,
DecoderWrapper, export_encoder, export_decoder, export_tokens) into a new common
module (e.g., canary_export_utils.py) and import them here, leaving only
model-specific configuration (the EncDecMultiTaskModel.from_pretrained call,
vocab/URL/meta_data values, and the load_external_data flag) in this script;
replace the duplicated definitions with imports and update main() to call the
shared export_*.py utilities and pass model-specific parameters (model name/URL,
load_external_data, meta_data) so both export_onnx_1b_flash.py and
export_onnx_180m_flash.py share the same implementation.
In `@scripts/nemo/canary/run_180m_flash.sh`:
- Around line 15-24: The script uses the literal command "uv pip install" but
doesn't verify that the "uv" wrapper is present; update the script to either
check for the uv executable and fall back to plain "pip install" (or
install/notify about uv) before invoking "uv pip install", or add a clear
prerequisite comment; locate the invocation of "uv pip install" and add a guard
like a command-existence check for "uv" with a fallback path or an explanatory
note so the script won't fail in environments without "uv".
In `@scripts/nemo/canary/test_1b_flash.py`:
- Around line 50-55: The two f-strings printing model sections are inconsistent:
one uses "{model} Input" and the other "{model }Output", causing the Output line
to lack a separating space; update the print f-string that uses "{model }Output"
so the space is outside the braces (match the format used in the first print),
i.e., make the Output line use the same " {model} Output " spacing pattern
around the model variable in the print statement that iterates over
sess.get_outputs().
- Line 217: Remove the unused timing variable or use it to log elapsed time:
either delete the assignment "start = time.time()" or keep it and compute
elapsed = time.time() - start at the end of main() and log or print the duration
(e.g., use process/logger or print) so the variable is read; reference the
"start" assignment and the "main()" function to locate where to remove or add
the elapsed-time logging.
- Around line 249-261: The code appends f-strings with no interpolation to
decoder_input_ids which triggers Ruff/Flake8 F541; replace the unnecessary
f-quoted literals with plain string literals where token2id is indexed
(references: decoder_input_ids, token2id, args.target_lang, args.use_pnc) —
change f"<|en|>", f"<|{args.target_lang}|>" when target_lang is a literal choice
resolved, f"<|pnc|>", f"<|nopnc|>", and f"<|noitn|>" to regular quoted strings
(remove the leading f) so the tokens are plain string keys when doing
token2id[...] lookups.
🧹 Nitpick comments (2)
🤖 Fix all nitpicks with AI agents
Verify each finding against the current code and only fix it if needed. In `@scripts/nemo/canary/export_onnx_1b_flash.py`: - Around line 1-384: This file duplicates many utilities from export_onnx_180m_flash.py; extract the shared functions and classes (fixed_form_attention_mask, add_meta_data, lens_to_mask, EncoderWrapper, DecoderWrapper, export_encoder, export_decoder, export_tokens) into a new common module (e.g., canary_export_utils.py) and import them here, leaving only model-specific configuration (the EncDecMultiTaskModel.from_pretrained call, vocab/URL/meta_data values, and the load_external_data flag) in this script; replace the duplicated definitions with imports and update main() to call the shared export_*.py utilities and pass model-specific parameters (model name/URL, load_external_data, meta_data) so both export_onnx_1b_flash.py and export_onnx_180m_flash.py share the same implementation. In `@scripts/nemo/canary/run_180m_flash.sh`: - Around line 15-24: The script uses the literal command "uv pip install" but doesn't verify that the "uv" wrapper is present; update the script to either check for the uv executable and fall back to plain "pip install" (or install/notify about uv) before invoking "uv pip install", or add a clear prerequisite comment; locate the invocation of "uv pip install" and add a guard like a command-existence check for "uv" with a fallback path or an explanatory note so the script won't fail in environments without "uv".scripts/nemo/canary/export_onnx_1b_flash.py (1)
1-384: Significant code duplication withexport_onnx_180m_flash.py.This file is nearly identical to
export_onnx_180m_flash.py— the functionsfixed_form_attention_mask,add_meta_data,lens_to_mask,EncoderWrapper,DecoderWrapper,export_encoder,export_decoder, andexport_tokensare all verbatim copies. Only the model name, URL, andload_external_dataflag differ.Consider extracting the shared logic into a common module (e.g.,
canary_export_utils.py) and having each model-specific script import and configure it. This would reduce the maintenance burden when fixing bugs (like the one this PR addresses for 180m).🤖 Prompt for AI Agents
Verify each finding against the current code and only fix it if needed. In `@scripts/nemo/canary/export_onnx_1b_flash.py` around lines 1 - 384, This file duplicates many utilities from export_onnx_180m_flash.py; extract the shared functions and classes (fixed_form_attention_mask, add_meta_data, lens_to_mask, EncoderWrapper, DecoderWrapper, export_encoder, export_decoder, export_tokens) into a new common module (e.g., canary_export_utils.py) and import them here, leaving only model-specific configuration (the EncDecMultiTaskModel.from_pretrained call, vocab/URL/meta_data values, and the load_external_data flag) in this script; replace the duplicated definitions with imports and update main() to call the shared export_*.py utilities and pass model-specific parameters (model name/URL, load_external_data, meta_data) so both export_onnx_1b_flash.py and export_onnx_180m_flash.py share the same implementation.scripts/nemo/canary/run_180m_flash.sh (1)
15-24:uvis now required but not checked or installed.Switching from
pip installtouv pip installassumesuvis available in the execution environment. If this script is run outside the expected CI environment, it will fail. Consider adding a guard or a comment noting the prerequisite. This is consistent withrun_1b_flash.sh, so presumably intentional.🤖 Prompt for AI Agents
Verify each finding against the current code and only fix it if needed. In `@scripts/nemo/canary/run_180m_flash.sh` around lines 15 - 24, The script uses the literal command "uv pip install" but doesn't verify that the "uv" wrapper is present; update the script to either check for the uv executable and fall back to plain "pip install" (or install/notify about uv) before invoking "uv pip install", or add a clear prerequisite comment; locate the invocation of "uv pip install" and add a guard like a command-existence check for "uv" with a fallback path or an explanatory note so the script won't fail in environments without "uv".
| def export_tokens(canary_model): | ||
| underline = "▁" | ||
| with open("./tokens.txt", "w", encoding="utf-8") as f: | ||
| for i in range(canary_model.tokenizer.vocab_size): | ||
| s = canary_model.tokenizer.ids_to_text([i]) | ||
|
|
||
| if s[0] == " ": | ||
| s = underline + s[1:] | ||
|
|
||
| f.write(f"{s} {i}\n") | ||
| print("Saved to tokens.txt") |
There was a problem hiding this comment.
Potential IndexError if ids_to_text returns an empty string.
Line 303 accesses s[0] without checking that s is non-empty. If the tokenizer returns an empty string for any token ID, this will raise an IndexError.
Proposed fix
- if s[0] == " ":
+ if s and s[0] == " ":🤖 Prompt for AI Agents
Verify each finding against the current code and only fix it if needed.
In `@scripts/nemo/canary/export_onnx_1b_flash.py` around lines 297 - 307, In
export_tokens, accessing s[0] after s = canary_model.tokenizer.ids_to_text([i])
can raise IndexError if ids_to_text returns an empty string; update the logic to
first ensure s is non-empty (e.g., check if s and len(s) > 0) before testing
s[0], and handle the empty case (replace with underline + "" or a safe
placeholder, or write the token ID alone) so the loop never indexes into an
empty string while still applying the underline substitution for leading spaces.
| print(f"=========={model} Input==========") | ||
| for i in sess.get_inputs(): | ||
| print(i) | ||
| print(f"=========={model }Output==========") | ||
| for i in sess.get_outputs(): | ||
| print(i) |
There was a problem hiding this comment.
Minor formatting inconsistency in display output.
Line 50 uses {model} Input (space outside braces), while line 53 uses {model }Output (space inside braces). The space inside {model } is part of the expression whitespace and gets ignored by Python, so the Output line will print without a separating space (e.g., ==========encoderOutput==========).
Proposed fix
- print(f"=========={model }Output==========")
+ print(f"=========={model} Output==========")📝 Committable suggestion
‼️ IMPORTANT
Carefully review the code before committing. Ensure that it accurately replaces the highlighted code, contains no missing lines, and has no issues with indentation. Thoroughly test & benchmark the code to ensure it meets the requirements.
| print(f"=========={model} Input==========") | |
| for i in sess.get_inputs(): | |
| print(i) | |
| print(f"=========={model }Output==========") | |
| for i in sess.get_outputs(): | |
| print(i) | |
| print(f"=========={model} Input==========") | |
| for i in sess.get_inputs(): | |
| print(i) | |
| print(f"=========={model} Output==========") | |
| for i in sess.get_outputs(): | |
| print(i) |
🤖 Prompt for AI Agents
Verify each finding against the current code and only fix it if needed.
In `@scripts/nemo/canary/test_1b_flash.py` around lines 50 - 55, The two f-strings
printing model sections are inconsistent: one uses "{model} Input" and the other
"{model }Output", causing the Output line to lack a separating space; update the
print f-string that uses "{model }Output" so the space is outside the braces
(match the format used in the first print), i.e., make the Output line use the
same " {model} Output " spacing pattern around the model variable in the print
statement that iterates over sess.get_outputs().
|
|
||
| fbank = create_fbank() | ||
|
|
||
| start = time.time() |
There was a problem hiding this comment.
Unused variable start.
start = time.time() is assigned but never read. Looks like leftover timing code — either remove it or add elapsed-time logging at the end of main().
Proposed fix
- start = time.time()
audio, sample_rate = sf.read(args.wav, dtype="float32", always_2d=True)Or, if timing was intended:
start = time.time()
# ... (end of main, before final print) ...
+ elapsed = time.time() - start
+ print(f"Elapsed time: {elapsed:.3f}s")📝 Committable suggestion
‼️ IMPORTANT
Carefully review the code before committing. Ensure that it accurately replaces the highlighted code, contains no missing lines, and has no issues with indentation. Thoroughly test & benchmark the code to ensure it meets the requirements.
| start = time.time() | |
| audio, sample_rate = sf.read(args.wav, dtype="float32", always_2d=True) |
🧰 Tools
🪛 Flake8 (7.3.0)
[error] 217-217: local variable 'start' is assigned to but never used
(F841)
🪛 Ruff (0.15.0)
[error] 217-217: Local variable start is assigned to but never used
Remove assignment to unused variable start
(F841)
🤖 Prompt for AI Agents
Verify each finding against the current code and only fix it if needed.
In `@scripts/nemo/canary/test_1b_flash.py` at line 217, Remove the unused timing
variable or use it to log elapsed time: either delete the assignment "start =
time.time()" or keep it and compute elapsed = time.time() - start at the end of
main() and log or print the duration (e.g., use process/logger or print) so the
variable is read; reference the "start" assignment and the "main()" function to
locate where to remove or add the elapsed-time logging.
| decoder_input_ids.append(token2id[f"<|en|>"]) | ||
|
|
||
| if args.target_lang in ("en", "es", "de", "fr"): | ||
| decoder_input_ids.append(token2id[f"<|{args.target_lang}|>"]) | ||
| else: | ||
| decoder_input_ids.append(token2id[f"<|en|>"]) | ||
|
|
||
| if args.use_pnc: | ||
| decoder_input_ids.append(token2id[f"<|pnc|>"]) | ||
| else: | ||
| decoder_input_ids.append(token2id[f"<|nopnc|>"]) | ||
|
|
||
| decoder_input_ids.append(token2id[f"<|noitn|>"]) |
There was a problem hiding this comment.
Remove extraneous f prefixes on string literals without placeholders.
Lines 249, 254, 257, 259, and 261 use f"..." but contain no {...} interpolation. These should be plain strings. Flagged by both Ruff (F541) and Flake8 (F541).
Proposed fix
- decoder_input_ids.append(token2id[f"<|en|>"])
+ decoder_input_ids.append(token2id["<|en|>"])Apply the same pattern to lines 254, 257, 259, and 261.
📝 Committable suggestion
‼️ IMPORTANT
Carefully review the code before committing. Ensure that it accurately replaces the highlighted code, contains no missing lines, and has no issues with indentation. Thoroughly test & benchmark the code to ensure it meets the requirements.
| decoder_input_ids.append(token2id[f"<|en|>"]) | |
| if args.target_lang in ("en", "es", "de", "fr"): | |
| decoder_input_ids.append(token2id[f"<|{args.target_lang}|>"]) | |
| else: | |
| decoder_input_ids.append(token2id[f"<|en|>"]) | |
| if args.use_pnc: | |
| decoder_input_ids.append(token2id[f"<|pnc|>"]) | |
| else: | |
| decoder_input_ids.append(token2id[f"<|nopnc|>"]) | |
| decoder_input_ids.append(token2id[f"<|noitn|>"]) | |
| decoder_input_ids.append(token2id["<|en|>"]) | |
| if args.target_lang in ("en", "es", "de", "fr"): | |
| decoder_input_ids.append(token2id[f"<|{args.target_lang}|>"]) | |
| else: | |
| decoder_input_ids.append(token2id["<|en|>"]) | |
| if args.use_pnc: | |
| decoder_input_ids.append(token2id["<|pnc|>"]) | |
| else: | |
| decoder_input_ids.append(token2id["<|nopnc|>"]) | |
| decoder_input_ids.append(token2id["<|noitn|>"]) |
🧰 Tools
🪛 Flake8 (7.3.0)
[error] 249-249: f-string is missing placeholders
(F541)
[error] 254-254: f-string is missing placeholders
(F541)
[error] 257-257: f-string is missing placeholders
(F541)
[error] 259-259: f-string is missing placeholders
(F541)
[error] 261-261: f-string is missing placeholders
(F541)
🪛 Ruff (0.15.0)
[error] 249-249: f-string without any placeholders
Remove extraneous f prefix
(F541)
[error] 254-254: f-string without any placeholders
Remove extraneous f prefix
(F541)
[error] 257-257: f-string without any placeholders
Remove extraneous f prefix
(F541)
[error] 259-259: f-string without any placeholders
Remove extraneous f prefix
(F541)
[error] 261-261: f-string without any placeholders
Remove extraneous f prefix
(F541)
🤖 Prompt for AI Agents
Verify each finding against the current code and only fix it if needed.
In `@scripts/nemo/canary/test_1b_flash.py` around lines 249 - 261, The code
appends f-strings with no interpolation to decoder_input_ids which triggers
Ruff/Flake8 F541; replace the unnecessary f-quoted literals with plain string
literals where token2id is indexed (references: decoder_input_ids, token2id,
args.target_lang, args.use_pnc) — change f"<|en|>", f"<|{args.target_lang}|>"
when target_lang is a literal choice resolved, f"<|pnc|>", f"<|nopnc|>", and
f"<|noitn|>" to regular quoted strings (remove the leading f) so the tokens are
plain string keys when doing token2id[...] lookups.
Co-authored-by: gemini-code-assist[bot] <176961590+gemini-code-assist[bot]@users.noreply.github.com>
There was a problem hiding this comment.
Actionable comments posted: 3
🤖 Prompt for all review comments with AI agents
Verify each finding against the current code and only fix it if needed.
Inline comments:
In `@scripts/nemo/canary/export_onnx_1b_flash.py`:
- Around line 12-13: The file imports an unused symbol "os" which triggers
Flake8 F401; remove the unused import from the top-level import list so only
needed names remain (e.g., keep "from typing import Dict, Tuple" and delete
"import os") or alternatively use "os" where intended; update the import line(s)
around the top of export_onnx_1b_flash.py (referencing the current import of os
and the typing imports) to eliminate the unused import.
- Line 380: The module-level call to subprocess.run(["ls", "-lh", "*.onnx"])
will raise NameError because subprocess is not imported, runs at import time
instead of only when main() executes, and won't expand the glob; fix by adding
an import subprocess at the top of the file, move the subprocess.run call into
the main() function (or whichever entrypoint runs export_onnx_1b_flash), and
replace the incorrect glob usage with either a shell string (e.g.,
subprocess.run("ls -lh *.onnx", shell=True, check=True)) or, preferably, use
Python's glob.glob to list ONNX files and iterate over them (e.g., import glob
and use glob.glob("*.onnx") then print or run ls on each file) so shell
expansion issues are avoided.
- Around line 136-208: DecoderWrapper incorrectly hardcodes six decoder mems;
change the forward signature of DecoderWrapper.forward to accept only
decoder_mems_list_0 through decoder_mems_list_3, pass a list of those four mems
into self.decoder.decoder.forward, unpack decoder_mems_list into (out_mems_0,
out_mems_1, out_mems_2, out_mems_3), compute logits from out_mems_3 (e.g.
self.log_softmax(hidden_states=out_mems_3[:, -1:])), and return only logits plus
out_mems_0..out_mems_3; also update export_decoder() to create/pass four dummy
tensors and rename all ONNX input/output names and dynamic_axes to reference
decoder_mems_list_0..decoder_mems_list_3 instead of 0..5 so the ONNX schema
matches the 4-layer model.
---
Duplicate comments:
In `@scripts/nemo/canary/export_onnx_1b_flash.py`:
- Around line 297-307: In export_tokens, avoid IndexError by checking that the
result of canary_model.tokenizer.ids_to_text([i]) is a non-empty string before
accessing s[0]; update the logic around the call to tokenizer.ids_to_text in
export_tokens so it tests e.g. "if s and s[0] == ' ':" (or treats empty strings
explicitly, e.g. replace empty s with a placeholder) and then write the safe
string to tokens.txt; ensure you reference export_tokens and
canary_model.tokenizer.ids_to_text([i]) when making the change.
| class DecoderWrapper(torch.nn.Module): | ||
| def __init__(self, m): | ||
| super().__init__() | ||
| self.decoder = m.transf_decoder | ||
| self.log_softmax = m.log_softmax | ||
|
|
||
| # We use only greedy search, so there is no need to compute log_softmax | ||
| self.log_softmax.mlp.log_softmax = False | ||
|
|
||
| def forward( | ||
| self, | ||
| decoder_input_ids: torch.Tensor, | ||
| decoder_mems_list_0: torch.Tensor, | ||
| decoder_mems_list_1: torch.Tensor, | ||
| decoder_mems_list_2: torch.Tensor, | ||
| decoder_mems_list_3: torch.Tensor, | ||
| decoder_mems_list_4: torch.Tensor, | ||
| decoder_mems_list_5: torch.Tensor, | ||
| enc_states: torch.Tensor, | ||
| enc_mask: torch.Tensor, | ||
| ): | ||
| """ | ||
| Args: | ||
| decoder_input_ids: (N, num_tokens), torch.int32 | ||
| decoder_mems_list_i: (N, num_tokens, 1024) | ||
| enc_states: (N, T, 1024) | ||
| enc_mask: (N, T) | ||
| Returns: | ||
| - logits: (N, 1, vocab_size) | ||
| - decoder_mems_list_i: (N, num_tokens_2, 1024) | ||
| """ | ||
| pos = decoder_input_ids[0][-1].item() | ||
| decoder_input_ids = decoder_input_ids[:, :-1] | ||
|
|
||
| decoder_hidden_states = self.decoder.embedding.forward( | ||
| decoder_input_ids, start_pos=pos | ||
| ) | ||
| decoder_input_mask = torch.ones_like(decoder_input_ids).float() | ||
|
|
||
| decoder_mems_list, _xatt_scores = self.decoder.decoder.forward( | ||
| decoder_hidden_states, | ||
| decoder_input_mask, | ||
| enc_states, | ||
| enc_mask, | ||
| [ | ||
| decoder_mems_list_0, | ||
| decoder_mems_list_1, | ||
| decoder_mems_list_2, | ||
| decoder_mems_list_3, | ||
| decoder_mems_list_4, | ||
| decoder_mems_list_5, | ||
| ], | ||
| return_mems=True, | ||
| ) | ||
| ( | ||
| out_mems_0, | ||
| out_mems_1, | ||
| out_mems_2, | ||
| out_mems_3, | ||
| out_mems_4, | ||
| out_mems_5, | ||
| ) = decoder_mems_list | ||
| logits = self.log_softmax(hidden_states=out_mems_5[:, -1:]) | ||
|
|
||
| return ( | ||
| logits, | ||
| out_mems_0, | ||
| out_mems_1, | ||
| out_mems_2, | ||
| out_mems_3, | ||
| out_mems_4, | ||
| out_mems_5, | ||
| ) |
There was a problem hiding this comment.
🧩 Analysis chain
🏁 Script executed:
# First, let's explore the repository structure
find . -type f -name "*.py" | grep -E "(config|model|export)" | head -20Repository: k2-fsa/sherpa-onnx
Length of output: 952
🏁 Script executed:
# Look for the export_onnx files mentioned
fd "export_onnx.*\.py$" | head -10Repository: k2-fsa/sherpa-onnx
Length of output: 519
🏁 Script executed:
# Check the actual file in question
wc -l scripts/nemo/canary/export_onnx_1b_flash.pyRepository: k2-fsa/sherpa-onnx
Length of output: 109
🏁 Script executed:
# Read the 1b_flash export file to see the full context
cat -n scripts/nemo/canary/export_onnx_1b_flash.pyRepository: k2-fsa/sherpa-onnx
Length of output: 14679
🏁 Script executed:
# Also read the 180m_flash export file for comparison
cat -n scripts/nemo/canary/export_onnx_180m_flash.pyRepository: k2-fsa/sherpa-onnx
Length of output: 14610
🏁 Script executed:
# Look for model configs or documentation about Canary models
find . -type f \( -name "*.md" -o -name "*.yaml" -o -name "*.json" \) | xargs grep -l "canary\|Canary" 2>/dev/null | head -10Repository: k2-fsa/sherpa-onnx
Length of output: 404
🏁 Script executed:
# Search for any comments or documentation about model architecture in the export files
grep -n "layer\|hidden\|decoder\|dimension" scripts/nemo/canary/export_onnx_*.pyRepository: k2-fsa/sherpa-onnx
Length of output: 14017
🏁 Script executed:
# Check if there are any README or docs specifically for these models
find . -path "*nemo*canary*" -type fRepository: k2-fsa/sherpa-onnx
Length of output: 959
🏁 Script executed:
# Check the workflow file for export configuration
cat .github/workflows/export-nemo-canary-180m-flash.yamlRepository: k2-fsa/sherpa-onnx
Length of output: 3282
🏁 Script executed:
# Check test files to understand expected model behavior
head -100 scripts/nemo/canary/test_180m_flash.py
head -100 scripts/nemo/canary/test_1b_flash.pyRepository: k2-fsa/sherpa-onnx
Length of output: 5411
🌐 Web query:
nvidia canary-1b-flash model architecture decoder layers hidden dimension
💡 Result:
- Decoder layers: 4 (Transformer decoder) [1]
- Hidden dimension (model width):
d_model = 1024[2]
Sources: [1] NVIDIA canary-1b-flash model card (Hugging Face). [2] config.json for nvidia/canary-1b-flash (Hugging Face).
🌐 Web query:
nvidia canary-180m-flash model decoder configuration specs
💡 Result:
canary-180m-flash — decoder (Transformer) configuration specs
High-level
- Model type: FastConformer encoder + Transformer decoder (encoder-decoder AED).
- Layer counts: 17 encoder layers, 4 decoder layers. [1]
Transformer decoder (model hyperparams)
From NeMo’s Canary/fast-conformer AED base config (and its Canary variant notes):
num_layers: 4 (for canary-180m-flash variant) [2]hidden_size(LM/decoder hidden): 1024 [2]num_attention_heads: 8 [2]inner_size(FFN size): 4 × hidden = 4096 [2]max_sequence_length: 1024 (for canary-180m-flash variant) [2]hidden_act: relu [2]- Dropouts:
embedding_dropout: 0.1 [2]ffn_dropout: 0.1 [2]attn_score_dropout: 0.1 [2]attn_layer_dropout: 0.1 [2]
- LayerNorm style:
pre_ln: true,pre_ln_final_layer_norm: true[2] - Positional encodings:
learn_positional_encodings: false[2] vocab_size: set at runtime by the model (Nonein config). [2]
Decoding (inference) config defaults
strategy: beam [2]beam.beam_size: 1 [2] (also shown in the model card usage snippet) [1]beam.len_pen: 0.0 [2]beam.max_generation_delta: 50 [2]
Sources
[1] NVIDIA Hugging Face model card for nvidia/canary-180m-flash (huggingface.co)
[2] NeMo base config fast-conformer_aed.yaml (includes Canary variant notes + decoder/decoding fields) (raw.githubusercontent.com)
Citations:
- 1: https://huggingface.co/nvidia/canary-180m-flash
- 2: https://raw.githubusercontent.com/NVIDIA-NeMo/NeMo/main/examples/asr/conf/speech_multitask/fast-conformer_aed.yaml
Fix hardcoded decoder layer count from 6 to 4 to match actual model architecture.
The DecoderWrapper class incorrectly hardcodes 6 decoder memory tensors (decoder_mems_list_0 through decoder_mems_list_5), but both the canary-1b-flash and canary-180m-flash models have only 4 Transformer decoder layers. This causes a mismatch between the ONNX export schema and the actual model, resulting in silent failures or runtime crashes.
Required changes:
- Remove
decoder_mems_list_4anddecoder_mems_list_5from the forward method signature (lines 145-156) - Update the mems list passed to
self.decoder.decoder.forward()to include only indices 0-3 (lines 180-187) - Change unpacking from 6 variables to 4 variables (lines 190-197)
- Change line 198 from
out_mems_5toout_mems_3to use the last decoder layer's output - Update return statement to return only 4 mems (lines 200-208)
- Update
export_decoder()to create and pass only 4 dummy tensors (lines 238-243, 252-257) - Update all ONNX input/output names and dynamic axes to reference only
decoder_mems_list_0throughdecoder_mems_list_3(lines 264-290)
The hidden dimension of 1024 in the docstring is correct and requires no changes.
🤖 Prompt for AI Agents
Verify each finding against the current code and only fix it if needed.
In `@scripts/nemo/canary/export_onnx_1b_flash.py` around lines 136 - 208,
DecoderWrapper incorrectly hardcodes six decoder mems; change the forward
signature of DecoderWrapper.forward to accept only decoder_mems_list_0 through
decoder_mems_list_3, pass a list of those four mems into
self.decoder.decoder.forward, unpack decoder_mems_list into (out_mems_0,
out_mems_1, out_mems_2, out_mems_3), compute logits from out_mems_3 (e.g.
self.log_softmax(hidden_states=out_mems_3[:, -1:])), and return only logits plus
out_mems_0..out_mems_3; also update export_decoder() to create/pass four dummy
tensors and rename all ONNX input/output names and dynamic_axes to reference
decoder_mems_list_0..decoder_mems_list_3 instead of 0..5 so the ONNX schema
matches the 4-layer model.
| print("new", model.ir_version) | ||
| onnx.save(model, filename) | ||
|
|
||
| subprocess.run(["ls", "-lh", "*.onnx"], check=True) |
There was a problem hiding this comment.
subprocess is not imported — this line will raise NameError at runtime.
Three problems here:
- Missing import:
subprocessis never imported (Ruff F821). - Outside
main(): This statement is at module level, so it executes on import rather than only whenmain()is called. It should be moved insidemain(). - Glob won't expand: Passing
"*.onnx"as a list element tosubprocess.runbypasses shell expansion, solswill literally look for a file named*.onnx. Useshell=Truewith a string command, or useglob.glob.
Proposed fix
Add the import at the top of the file:
+import subprocessThen move the call inside main() and fix the glob issue:
-subprocess.run(["ls", "-lh", "*.onnx"], check=True)
-
-
-if __name__ == "__main__":
- main()
+ import glob
+ onnx_files = glob.glob("*.onnx")
+ if onnx_files:
+ subprocess.run(["ls", "-lh"] + sorted(onnx_files), check=True)
+
+
+if __name__ == "__main__":
+ main()🧰 Tools
🪛 Flake8 (7.3.0)
[error] 380-380: undefined name 'subprocess'
(F821)
🪛 Ruff (0.15.0)
[error] 380-380: Undefined name subprocess
(F821)
🤖 Prompt for AI Agents
Verify each finding against the current code and only fix it if needed.
In `@scripts/nemo/canary/export_onnx_1b_flash.py` at line 380, The module-level
call to subprocess.run(["ls", "-lh", "*.onnx"]) will raise NameError because
subprocess is not imported, runs at import time instead of only when main()
executes, and won't expand the glob; fix by adding an import subprocess at the
top of the file, move the subprocess.run call into the main() function (or
whichever entrypoint runs export_onnx_1b_flash), and replace the incorrect glob
usage with either a shell string (e.g., subprocess.run("ls -lh *.onnx",
shell=True, check=True)) or, preferably, use Python's glob.glob to list ONNX
files and iterate over them (e.g., import glob and use glob.glob("*.onnx") then
print or run ls on each file) so shell expansion issues are avoided.
There was a problem hiding this comment.
Actionable comments posted: 1
🤖 Prompt for all review comments with AI agents
Verify each finding against the current code and only fix it if needed.
Inline comments:
In `@scripts/nemo/canary/test_1b_flash.py`:
- Around line 265-289: The auto-regressive generation loop uses incorrect
position indices causing collisions with the prompt: when seeding the prompt you
use positions via enumerate(decoder_input_ids) and then start generation with
for i in range(1, 200) which reuses position 1 (colliding with the prompt's
second token). Fix the loop so the first generated token starts at position
len(decoder_input_ids) (i.e., use range(len(decoder_input_ids),
len(decoder_input_ids) + max_gen) or compute start_pos = len(decoder_input_ids)
and increment from there), ensuring the call to model.run_decoder that builds
decoder_input_ids and the position fed to run_decoder advances from start_pos;
update references to decoder_input_ids, tokens, decoder_mems_list, and the
run_decoder call accordingly.
---
Duplicate comments:
In `@scripts/nemo/canary/export_onnx_1b_flash.py`:
- Around line 381-385: The subprocess.run call at module level uses an
unexpanded glob and runs on import; move the subprocess.run invocation into the
existing main() function (so it only executes when __main__ runs) and expand the
glob before calling subprocess.run (e.g., use glob.glob("*.onnx") to collect
files and pass the file list to subprocess.run or build a safe argument list).
Update the call referencing subprocess.run in export_onnx_1b_flash.py to operate
on the expanded list (or skip calling ls entirely if unnecessary) so the command
isn't given a literal "*.onnx" and does not execute during import.
- Line 12: Remove the unused import "os" from the top of the file
export_onnx_1b_flash.py; locate the import statement (import os) and delete it
so the module no longer imports an unused symbol and fixes Flake8 F401, ensuring
no other references to os (e.g., os.path or os.getenv) exist in functions like
any export or model conversion helpers before committing.
- Around line 298-308: In export_tokens, avoid indexing s[0] when
ids_to_text([i]) may return an empty string: after calling s =
canary_model.tokenizer.ids_to_text([i]) ensure s is a non-empty string (e.g., if
not s: set s = underline or a safe placeholder), then apply the
space-to-underline transform (if s and s[0] == " " -> prepend underline to
s[1:]) before writing to file; update any logic that assumes s is non-empty so
ids_to_text returning "" cannot cause an IndexError.
- Around line 137-209: DecoderWrapper currently accepts 6 decoder mem tensors
but the canary-1b-flash model uses 4 decoder layers; update
DecoderWrapper.forward to accept only decoder_mems_list_0 through
decoder_mems_list_3, update the unpacking to (out_mems_0, out_mems_1,
out_mems_2, out_mems_3) and call self.log_softmax with out_mems_3[:, -1:], and
propagate the same 4-mem change to export_decoder() and the test script; also
verify the model's actual layer count (e.g., model.config.num_layers or
n_layers) before finalizing to avoid mismatch.
In `@scripts/nemo/canary/test_1b_flash.py`:
- Around line 49-55: The f-string in the display function prints the model name
and "Output" without the intended space because it uses "{model }Output"; update
that f-string in display to place the space outside the braces (e.g., "{model}
Output" or " {model} Output" as appropriate) so the header prints
"=========={model} Output==========" consistently; keep the first header format
("=========={model} Input==========") and mirror its spacing for the output
header.
- Line 217: The variable 'start' is assigned from time.time() but never used;
either remove the unused assignment (delete the line start = time.time()) or use
it to log elapsed time by capturing end = time.time() and logging the difference
(e.g., in the test function surrounding where start is set, reference 'start'
and compute elapsed). Update the function that contains this assignment (look
for start = time.time() in test_1b_flash.py) accordingly to eliminate the
unused-variable warning.
- Around line 249-261: The code uses unnecessary f-strings where there is no
interpolation (causing F541); replace f"..." literals with plain string literals
when appending tokens to decoder_input_ids — e.g., in the blocks that call
token2id[f"<|en|>"], token2id[f"<|{args.target_lang}|>"] (when args.target_lang
is one of the fixed tags), token2id[f"<|pnc|>"], token2id[f"<|nopnc|>"], and
token2id[f"<|noitn|>"] convert those f-strings to normal strings
(token2id["<|en|>"], token2id[f"<|{args.target_lang}|>"] remains an f-string
only when it actually contains {args.target_lang}, otherwise use
token2id["<|en|>"]) so that decoder_input_ids and token2id calls no longer use
extraneous f prefixes.
| decoder_mems_list = [np.zeros((1, 0, 1024), dtype=np.float32) for _ in range(6)] | ||
|
|
||
| for pos, decoder_input_id in enumerate(decoder_input_ids): | ||
| logits, decoder_mems_list = model.run_decoder( | ||
| np.array([[decoder_input_id, pos]], dtype=np.int32), | ||
| decoder_mems_list, | ||
| enc_states, | ||
| enc_masks, | ||
| ) | ||
| tokens = [logits.argmax()] | ||
| print("decoder_input_ids", decoder_input_ids) | ||
| eos = token2id["<|endoftext|>"] | ||
|
|
||
| for i in range(1, 200): | ||
| decoder_input_ids = [tokens[-1], i] | ||
| logits, decoder_mems_list = model.run_decoder( | ||
| np.array([decoder_input_ids], dtype=np.int32), | ||
| decoder_mems_list, | ||
| enc_states, | ||
| enc_masks, | ||
| ) | ||
| t = logits.argmax() | ||
| if t == eos: | ||
| break | ||
| tokens.append(t) |
There was a problem hiding this comment.
🧩 Analysis chain
🏁 Script executed:
#!/bin/bash
# First, let's find and examine the 180m flash test file
fd "test_180m_flash.py" scripts/nemo/canary/ --exec cat {} \; 2>/dev/null | grep -n -A10 -B2 "range(1," | head -50Repository: k2-fsa/sherpa-onnx
Length of output: 491
🏁 Script executed:
#!/bin/bash
# Search for start_pos usage in the decoder and understand its semantics
rg -n "start_pos" scripts/nemo/canary/ -A3 -B3Repository: k2-fsa/sherpa-onnx
Length of output: 1244
🏁 Script executed:
#!/bin/bash
# Check the run_decoder method signature and implementation
rg -n "def run_decoder" scripts/nemo/canary/ -A20Repository: k2-fsa/sherpa-onnx
Length of output: 3242
🏁 Script executed:
#!/bin/bash
# Look at the model implementation to understand position encoding
rg -n "positional\|position" scripts/nemo/canary/test_1b_flash.py -iRepository: k2-fsa/sherpa-onnx
Length of output: 44
🏁 Script executed:
#!/bin/bash
# Check what the prompt_ids and prompt_len are in the test
rg -n "prompt_ids\|prompt_len\|decoder_input_ids\s*=" scripts/nemo/canary/test_1b_flash.py -B2 -A2 | head -60Repository: k2-fsa/sherpa-onnx
Length of output: 44
🏁 Script executed:
#!/bin/bash
# Look at the full context of the initial prompt processing loop
sed -n '260,280p' scripts/nemo/canary/test_1b_flash.pyRepository: k2-fsa/sherpa-onnx
Length of output: 832
🏁 Script executed:
#!/bin/bash
# Check if there are any comments or documentation about position handling
rg -n "position\|pos\|absolute\|relative" scripts/nemo/canary/test_1b_flash.py -i -B2 -A2 | head -80Repository: k2-fsa/sherpa-onnx
Length of output: 44
🏁 Script executed:
#!/bin/bash
# Check the prompt initialization to understand what decoder_input_ids contains at line 267
sed -n '255,275p' scripts/nemo/canary/test_1b_flash.pyRepository: k2-fsa/sherpa-onnx
Length of output: 812
🏁 Script executed:
#!/bin/bash
# Look for test validation or assertions that might indicate expected behavior
rg -n "assert\|test\|expected\|check" scripts/nemo/canary/test_1b_flash.py -i | tail -20Repository: k2-fsa/sherpa-onnx
Length of output: 44
🏁 Script executed:
#!/bin/bash
# Check if there's a main/test entry point that shows intended usage
sed -n '1,50p' scripts/nemo/canary/test_1b_flash.pyRepository: k2-fsa/sherpa-onnx
Length of output: 1314
🏁 Script executed:
#!/bin/bash
# Look for any README or documentation about position encoding
fd "README|USAGE|GUIDE" scripts/nemo/canary/ --exec cat {} \;Repository: k2-fsa/sherpa-onnx
Length of output: 44
🏁 Script executed:
#!/bin/bash
# Check if there's any code that uses start_pos to understand the semantics
rg -n "start_pos" --type python -B5 -A5 | grep -E "start_pos|def |class " | head -40Repository: k2-fsa/sherpa-onnx
Length of output: 92
Position values must be corrected in auto-regressive decoding loop.
The initial prompt loop (lines 267–273) assigns positions 0 through len(decoder_input_ids)-1 to each token. The auto-regressive loop then starts at position 1 via range(1, 200), creating a collision: the first generated token receives position 1, which matches the prompt's second token. After the prompt phase completes, the next position should be len(prompt), not 1.
This same issue appears in test_180m_flash.py, suggesting a systematic error. Update the auto-regressive loop to start at len(decoder_input_ids) to ensure generated tokens use correct positional indices.
🤖 Prompt for AI Agents
Verify each finding against the current code and only fix it if needed.
In `@scripts/nemo/canary/test_1b_flash.py` around lines 265 - 289, The
auto-regressive generation loop uses incorrect position indices causing
collisions with the prompt: when seeding the prompt you use positions via
enumerate(decoder_input_ids) and then start generation with for i in range(1,
200) which reuses position 1 (colliding with the prompt's second token). Fix the
loop so the first generated token starts at position len(decoder_input_ids)
(i.e., use range(len(decoder_input_ids), len(decoder_input_ids) + max_gen) or
compute start_pos = len(decoder_input_ids) and increment from there), ensuring
the call to model.run_decoder that builds decoder_input_ids and the position fed
to run_decoder advances from start_pos; update references to decoder_input_ids,
tokens, decoder_mems_list, and the run_decoder call accordingly.
|
closing this PR since it's not merged after so long. |
Fixed the error when exporting 180m-flash and added export for 1b-flash
Summary by CodeRabbit
New Features
Chores
Bug Fixes