Skip to content

fixed 180m-flash export and added 1b-flash export - #3193

Closed
jackalxyz wants to merge 6 commits into
k2-fsa:masterfrom
jackalxyz:master
Closed

jackalxyz wants to merge 6 commits into
k2-fsa:masterfrom
jackalxyz:master

Conversation

@jackalxyz

@jackalxyz jackalxyz commented Feb 16, 2026 •

Copy link
Copy Markdown

Fixed the error when exporting 180m-flash and added export for 1b-flash

Summary by CodeRabbit

  • New Features

    • Added ONNX export and test tooling for canary models (including 1B and 180M variants), plus runnable scripts to export, quantize, and validate FP32/INT8 encoder+decoder models and produce transcriptions.
  • Chores

    • Included encoder ONNXs in validation, attached metadata, and adjusted IR compatibility and attention-mask handling for ONNXRuntime.
  • Bug Fixes

    • Fixed attention-mask compatibility; corrected an install command typo that could break the run script.

@dosubot dosubot Bot added the size:XL This PR changes 500-999 lines, ignoring generated files. label Feb 16, 2026
@gemini-code-assist

Copy link
Copy Markdown

Summary of Changes

Hello @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

  • Fixed 180M-Flash ONNX Export: Resolved issues in the export_onnx_180m_flash.py script, including correcting the decoder's forward pass return values, adjusting ONNX export parameters (dynamo and dynamic axes), and ensuring all exported ONNX models (encoder and decoder, FP32 and INT8) have their IR version set to 9 for compatibility.
  • Added 1B-Flash ONNX Export: Introduced a new script, export_onnx_1b_flash.py, to enable ONNX export for the canary-1b-flash model. This includes custom wrappers for the encoder and decoder, a fix for the attention mask to address ONNXRuntime Trilu issues, and functionality for dynamic quantization and metadata addition.
  • New Run and Test Scripts for 1B-Flash: Provided run_1b_flash.sh to automate the export process and test_1b_flash.py to validate the exported 1B-flash ONNX models (both FP32 and INT8) across different language pairs, ensuring correct functionality post-export.
  • Dependency Management Update: Updated the run_180m_flash.sh script to use uv pip install for dependency management, aligning with modern Python packaging practices.
Changelog
  • scripts/nemo/canary/export_onnx_180m_flash.py
    • Modified decoder's forward method to correctly unpack cross-attention scores and return individual memory states.
    • Updated torch.onnx.export parameters for the decoder, disabling dynamo and removing decoder_input_ids from dynamic axes.
    • Added canary_model.cpu() call to ensure model is on CPU before export.
    • Extended ONNX IR version fix to include encoder models alongside decoder models.
  • scripts/nemo/canary/export_onnx_1b_flash.py
    • Added new script for exporting the canary-1b-flash model to ONNX format.
    • Implemented fixed_form_attention_mask to resolve ONNXRuntime Trilu operator compatibility issues.
    • Introduced EncoderWrapper and DecoderWrapper classes to prepare the model for ONNX export.
    • Included functions for exporting encoder, decoder, and token vocabulary, along with dynamic quantization and metadata embedding.
  • scripts/nemo/canary/run_180m_flash.sh
    • Updated pip install command to uv pip install for dependency installation.
  • scripts/nemo/canary/run_1b_flash.sh
    • Added new shell script to automate the ONNX export and testing workflow for the canary-1b-flash model.
    • Includes steps for installing dependencies, running the export script, and testing both FP32 and INT8 ONNX models.
  • scripts/nemo/canary/test_1b_flash.py
    • Added new Python script for performing inference and testing on the exported canary-1b-flash ONNX models.
    • Features argument parsing for model paths, token file, language settings, and WAV input.
    • Includes OnnxModel class for managing ONNXRuntime sessions and create_fbank/compute_features for audio preprocessing.
Activity
  • The pull request was created by ssteo.
  • Initial code changes were pushed to fix an existing ONNX export script and introduce new functionality for another model.
  • New scripts were added to facilitate the export and testing of the 1B-flash model.
Using Gemini Code Assist

The 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 /gemini <command> or @gemini-code-assist <command>. Below is a summary of the supported commands on the current page.

Feature Command Description
Code Review /gemini review Performs a code review for the current pull request in its current state.
Pull Request Summary /gemini summary Provides a summary of the current pull request in its current state.
Comment @gemini-code-assist Responds in comments when explicitly tagged, both in pull request comments and review comments.
Help /gemini help Displays a list of available commands.

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 .gemini/ folder in the base of the repository. Detailed instructions can be found here.

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

  1. Review the Privacy Notices, Generative AI Prohibited Use Policy, Terms of Service, and learn how to configure Gemini Code Assist in GitHub here. Gemini can make mistakes, so double check it and use code with caution. ↩

@coderabbitai

coderabbitai Bot commented Feb 16, 2026 •

Copy link
Copy Markdown
📝 Walkthrough

Walkthrough

Adds 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

Cohort / File(s) Summary
Export (180M changes)
scripts/nemo/canary/export_onnx_180m_flash.py
Decoder forward now unpacks and returns multiple memory tensors (out_mems_0…out_mems_5); logits computed from out_mems_5. Decoder export disables dynamo, removes external_data, adjusts dynamic_axes, and expands post‑export validation to include encoder ONNX files and IR handling.
New Export (1B)
scripts/nemo/canary/export_onnx_1b_flash.py
New ONNX export script for canary-1b-flash: adds fixed_form_attention_mask, lens_to_mask, add_meta_data; introduces EncoderWrapper/DecoderWrapper; exports encoder/decoder/tokens; applies dynamic quantization and IR/version adjustments; writes metadata.
Runtime Test Harness
scripts/nemo/canary/test_1b_flash.py
New ONNX Runtime test harness: OnnxModel loads encoder/decoder sessions, exposes run_encoder/run_decoder handling multi-output mem states; includes Kaldi-native fbank feature extraction, token mapping, and greedy decoding loop for inference.
Run Scripts / Automation
scripts/nemo/canary/run_1b_flash.sh, scripts/nemo/canary/run_180m_flash.sh
Adds run_1b_flash.sh to run export, quantize, and test matrix for FP32/INT8; run_180m_flash.sh modified (accidental uv prefix added before pip install).
Generated artifacts (post-export)
*.onnx, *.int8.onnx (generated)
Post‑export flow now produces and validates encoder/decoder FP32 and INT8 ONNX files, converts/resaves with compatible IR versions, and attaches metadata to encoder artifacts.

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
Loading

Estimated code review effort

🎯 4 (Complex) | ⏱️ ~45 minutes

Possibly related PRs

Suggested reviewers

  • csukuangfj

Poem

🐰 I hop through code with curious cheer,
I export encoders far and near.
I stash the mems and count the hops,
Quantize carrots, skip the stops.
Tokens tumble — tests appear!

🚥 Pre-merge checks | ✅ 2 | ❌ 1

❌ Failed checks (1 warning)

Check name Status Explanation Resolution
Docstring Coverage ⚠️ Warning Docstring coverage is 30.77% which is insufficient. The required threshold is 80.00%. Write docstrings for the functions missing them to satisfy the coverage threshold.
✅ Passed checks (2 passed)
Check name Status Explanation
Description Check ✅ Passed Check skipped - CodeRabbit’s high-level summary is enabled.
Title check ✅ Passed The title accurately summarizes the main changes: fixing the 180m-flash export and adding 1b-flash export support, matching the file-level modifications.

✏️ Tip: You can configure your own custom pre-merge checks in the settings.

✨ Finishing Touches
  • 📝 Generate docstrings
🧪 Generate unit tests (beta)
  • Create PR with unit tests
  • Post copyable unit tests in a comment

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.

❤️ Share

Comment @coderabbitai help to get the list of available commands and usage tips.

@gemini-code-assist gemini-code-assist Bot left a comment

Copy link
Copy Markdown

Choose a reason for hiding this comment

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

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"},

Copy link
Copy Markdown

Choose a reason for hiding this comment

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

medium

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.

Comment thread scripts/nemo/canary/export_onnx_1b_flash.py Outdated
Comment thread scripts/nemo/canary/test_1b_flash.py

@coderabbitai coderabbitai Bot left a comment

Copy link
Copy Markdown

Choose a reason for hiding this comment

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

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 with export_onnx_180m_flash.py.

This file is nearly identical to export_onnx_180m_flash.py — the functions fixed_form_attention_mask, add_meta_data, lens_to_mask, EncoderWrapper, DecoderWrapper, export_encoder, export_decoder, and export_tokens are all verbatim copies. Only the model name, URL, and load_external_data flag 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: uv is now required but not checked or installed.

Switching from pip install to uv pip install assumes uv is 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 with run_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".

Comment thread scripts/nemo/canary/export_onnx_180m_flash.py
Comment on lines +297 to +307
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")

Copy link
Copy Markdown

Choose a reason for hiding this comment

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

⚠️ Potential issue | 🟡 Minor

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.

Comment on lines +50 to +55
print(f"=========={model} Input==========")
for i in sess.get_inputs():
print(i)
print(f"=========={model }Output==========")
for i in sess.get_outputs():
print(i)

Copy link
Copy Markdown

Choose a reason for hiding this comment

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

⚠️ Potential issue | 🟡 Minor

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.

Suggested change
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()

Copy link
Copy Markdown

Choose a reason for hiding this comment

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

⚠️ Potential issue | 🟡 Minor

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.

Suggested change
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.

Comment on lines +249 to +261
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|>"])

Copy link
Copy Markdown

Choose a reason for hiding this comment

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

⚠️ Potential issue | 🟡 Minor

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.

Suggested change
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>

@coderabbitai coderabbitai Bot left a comment

Copy link
Copy Markdown

Choose a reason for hiding this comment

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

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.

Comment thread scripts/nemo/canary/export_onnx_1b_flash.py
Comment on lines +136 to +208
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,
)

Copy link
Copy Markdown

Choose a reason for hiding this comment

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

⚠️ Potential issue | 🔴 Critical

🧩 Analysis chain

🏁 Script executed:

# First, let's explore the repository structure
find . -type f -name "*.py" | grep -E "(config|model|export)" | head -20

Repository: k2-fsa/sherpa-onnx

Length of output: 952


🏁 Script executed:

# Look for the export_onnx files mentioned
fd "export_onnx.*\.py$" | head -10

Repository: 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.py

Repository: 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.py

Repository: 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.py

Repository: 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 -10

Repository: 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_*.py

Repository: 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 f

Repository: 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.yaml

Repository: 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.py

Repository: 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 (None in 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:


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_4 and decoder_mems_list_5 from 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_5 to out_mems_3 to 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_0 through decoder_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)

Copy link
Copy Markdown

Choose a reason for hiding this comment

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

⚠️ Potential issue | 🔴 Critical

subprocess is not imported — this line will raise NameError at runtime.

Three problems here:

  1. Missing import: subprocess is never imported (Ruff F821).
  2. Outside main(): This statement is at module level, so it executes on import rather than only when main() is called. It should be moved inside main().
  3. Glob won't expand: Passing "*.onnx" as a list element to subprocess.run bypasses shell expansion, so ls will literally look for a file named *.onnx. Use shell=True with a string command, or use glob.glob.
Proposed fix

Add the import at the top of the file:

+import subprocess

Then 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.

Comment thread scripts/nemo/canary/test_1b_flash.py Outdated
Comment thread scripts/nemo/canary/run_1b_flash.sh

@coderabbitai coderabbitai Bot left a comment

Copy link
Copy Markdown

Choose a reason for hiding this comment

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

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.

Comment on lines +265 to +289
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)

Copy link
Copy Markdown

Choose a reason for hiding this comment

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

⚠️ Potential issue | 🟠 Major

🧩 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 -50

Repository: 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 -B3

Repository: 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/ -A20

Repository: 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 -i

Repository: 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 -60

Repository: 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.py

Repository: 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 -80

Repository: 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.py

Repository: 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 -20

Repository: 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.py

Repository: 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 -40

Repository: 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.

@jackalxyz

Copy link
Copy Markdown
Author

closing this PR since it's not merged after so long.

Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

size:XL This PR changes 500-999 lines, ignoring generated files.

Projects

None yet

Development

Successfully merging this pull request may close these issues.

2 participants