Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
40 changes: 40 additions & 0 deletions .github/workflows/export-nemotron-3.5-asr-streaming-0.6b.yaml
Original file line number Diff line number Diff line change
Expand Up @@ -42,8 +42,18 @@ jobs:
pip install kaldi-native-fbank
pip install soundfile librosa

- name: Install sherpa-onnx for testing
shell: bash
run: |
pip install sherpa-onnx-bin
which sherpa-onnx

sherpa-onnx-version

- name: Run
shell: bash
env:
HF_TOKEN: ${{ secrets.HF_TOKEN }}
run: |
cd scripts/nemo/nemotron-3.5-asr-streaming-0.6b

Expand All @@ -56,6 +66,36 @@ jobs:

echo "---"

- name: Test onnx models
shell: bash
run: |
pip install sherpa-onnx-bin
which sherpa-onnx

wget https://dldata-public.s3.us-east-2.amazonaws.com/2086-149220-0033.wav

src=scripts/nemo/nemotron-3.5-asr-streaming-0.6b

for chunk in 80 160 560 1120; do
echo "test $chunk int8"

sherpa-onnx \
--encoder=$src/$chunk/encoder.int8.onnx \
--decoder=$src/$chunk/decoder.int8.onnx \
--joiner=$src/$chunk/joiner.int8.onnx \
--tokens=$src/tokens.txt \
./2086-149220-0033.wav

echo "test $chunk float32"

sherpa-onnx \
--encoder=$src/$chunk/encoder.onnx \
--decoder=$src/$chunk/decoder.onnx \
--joiner=$src/$chunk/joiner.onnx \
--tokens=$src/tokens.txt \
./2086-149220-0033.wav
done

- name: Collect results
shell: bash
run: |
Expand Down
212 changes: 33 additions & 179 deletions scripts/nemo/nemotron-3.5-asr-streaming-0.6b/export_onnx.py
Original file line number Diff line number Diff line change
Expand Up @@ -4,7 +4,7 @@
import json
import os
from pathlib import Path
from typing import Any, Dict, Optional
from typing import Any, Dict

import nemo.collections.asr as nemo_asr
import onnx
Expand Down Expand Up @@ -37,31 +37,6 @@
]


def add_meta_data(filename: str, meta_data: Dict[str, str]):
"""Add meta data to an ONNX model. It is changed in-place."""
model = onnx.load(filename)

while len(model.metadata_props):
model.metadata_props.pop()

for key, value in meta_data.items():
meta = model.metadata_props.add()
meta.key = key
meta.value = str(value)

external_filename = filename.split(".onnx")[0]
# onnx.save refuses to overwrite an existing external-data file; the
# prompted-encoder export already wrote one, so remove it first.
Path(external_filename + ".data").unlink(missing_ok=True)
onnx.save(
model,
filename,
save_as_external_data=True,
all_tensors_to_one_file=True,
location=external_filename + ".data",
)


def _to_plain_container(obj: Any) -> Any:
try:
from omegaconf import DictConfig, ListConfig, OmegaConf
Expand All @@ -82,132 +57,36 @@ def _normalize_prompt_dictionary(obj: Any) -> Dict[str, int]:
return {str(k): int(v) for k, v in obj.items()}


def _get_config_value(obj: Any, key: str) -> Any:
if obj is None:
return None

try:
return getattr(obj, key)
except (AttributeError, KeyError):
pass

obj = _to_plain_container(obj)
if isinstance(obj, dict):
return obj.get(key)

return None


def get_prompt_dictionary(asr_model) -> Dict[str, int]:
"""Return the model's language prompt dictionary from NeMo artifacts."""
cfg = getattr(asr_model, "cfg", None)
model_defaults = _get_config_value(cfg, "model_defaults")
if model_defaults is None:
raise RuntimeError("Could not find cfg.model_defaults in the NeMo model")

prompt_dictionary = _get_config_value(model_defaults, "prompt_dictionary")
if prompt_dictionary is None:
raise RuntimeError(
"Could not find cfg.model_defaults.prompt_dictionary in the NeMo model"
)

try:
ans = _normalize_prompt_dictionary(prompt_dictionary)
except (TypeError, ValueError) as e:
raise RuntimeError(
"cfg.model_defaults.prompt_dictionary must map language strings "
"to integer prompt ids"
) from e

num_prompts = int(asr_model.num_prompts)
for language, prompt_id in ans.items():
if not 0 <= prompt_id < num_prompts:
raise ValueError(
"cfg.model_defaults.prompt_dictionary has out-of-range "
f"prompt id for '{language}': {prompt_id}; expected "
f"0 <= id < {num_prompts}"
)

auto_prompt_id = ans.get("auto")
if auto_prompt_id != 101:
raise ValueError(f"Expected auto prompt id 101, got {auto_prompt_id}")

# The dictionary may use locale-style keys such as en-US or ja-JP; the
# runtime derives base-code aliases, so accept either form here.
for language in ["en", "ja"]:
if not any(k == language or k.startswith(f"{language}-") for k in ans):
raise RuntimeError(
"cfg.model_defaults.prompt_dictionary is missing " f"'{language}'"
)

return ans


def _find_sentencepiece_processor(obj: Any, max_depth: int = 5) -> Optional[Any]:
seen = set()

def is_sentencepiece_processor(value: Any) -> bool:
return callable(getattr(value, "get_piece_size", None)) and callable(
getattr(value, "id_to_piece", None)
)

def visit(value: Any, depth: int) -> Optional[Any]:
if value is None or depth > max_depth:
return None

if is_sentencepiece_processor(value):
return value

obj_id = id(value)
if obj_id in seen:
return None
seen.add(obj_id)

for name in ["tokenizer", "sp_model", "model", "processor"]:
if hasattr(value, name):
found = visit(getattr(value, name), depth + 1)
if found is not None:
return found
def add_meta_data(filename: str, meta_data: Dict[str, str]):
"""Add meta data to an ONNX model. It is changed in-place."""
model = onnx.load(filename)

if isinstance(value, dict):
for v in value.values():
found = visit(v, depth + 1)
if found is not None:
return found
while len(model.metadata_props):
model.metadata_props.pop()

return None
for key, value in meta_data.items():
meta = model.metadata_props.add()
meta.key = key
meta.value = str(value)

return visit(obj, 0)
external_filename = filename.split(".onnx")[0]
# onnx.save refuses to overwrite an existing external-data file; the
# prompted-encoder export already wrote one, so remove it first.
Path(external_filename + ".data").unlink(missing_ok=True)
onnx.save(
model,
filename,
save_as_external_data=True,
all_tensors_to_one_file=True,
location=external_filename + ".data",
)


def save_tokens(asr_model, filename: str = "tokens.txt") -> int:
sp = _find_sentencepiece_processor(getattr(asr_model, "tokenizer", None))
if sp is None:
raise RuntimeError("Could not find the SentencePiece tokenizer in the model")

vocab_size = sp.get_piece_size()
with open(filename, "w", encoding="utf-8") as f:
for i in range(vocab_size):
f.write(f"{sp.id_to_piece(i)} {i}\n")
f.write(f"<blk> {vocab_size}\n")

print(f"Saved {filename}")
return vocab_size


def assert_forward_for_export_signature(encoder):
if not hasattr(encoder, "forward_for_export"):
raise RuntimeError("Expected encoder.forward_for_export for ONNX export")

signature = inspect.signature(encoder.forward_for_export)
missing = [
name for name in FORWARD_FOR_EXPORT_ARGS if name not in signature.parameters
]
if missing:
raise RuntimeError(
"encoder.forward_for_export is missing expected argument(s): "
f"{missing}. Signature: {signature}"
)
for i, s in enumerate(asr_model.joint.vocabulary):
f.write(f"{s} {i}\n")
f.write(f"<blk> {i+1}\n")
Comment on lines 85 to +89

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 save_tokens function has a return type annotation of -> int, but it does not return any value (implicitly returning None). Additionally, referencing the loop variable i outside the loop to write the <blk> token can raise an UnboundLocalError if the vocabulary is empty, and is generally considered a code smell in Python.

Consider updating the return type annotation to -> None and using len(vocab) instead of i+1 to determine the vocabulary size safely.

Suggested change
def save_tokens(asr_model, filename: str = "tokens.txt") -> int:
sp = _find_sentencepiece_processor(getattr(asr_model, "tokenizer", None))
if sp is None:
raise RuntimeError("Could not find the SentencePiece tokenizer in the model")
vocab_size = sp.get_piece_size()
with open(filename, "w", encoding="utf-8") as f:
for i in range(vocab_size):
f.write(f"{sp.id_to_piece(i)} {i}\n")
f.write(f"<blk> {vocab_size}\n")
print(f"Saved {filename}")
return vocab_size
def assert_forward_for_export_signature(encoder):
if not hasattr(encoder, "forward_for_export"):
raise RuntimeError("Expected encoder.forward_for_export for ONNX export")
signature = inspect.signature(encoder.forward_for_export)
missing = [
name for name in FORWARD_FOR_EXPORT_ARGS if name not in signature.parameters
]
if missing:
raise RuntimeError(
"encoder.forward_for_export is missing expected argument(s): "
f"{missing}. Signature: {signature}"
)
for i, s in enumerate(asr_model.joint.vocabulary):
f.write(f"{s} {i}\n")
f.write(f"<blk> {i+1}\n")
def save_tokens(asr_model, filename: str = "tokens.txt") -> None:
vocab = asr_model.joint.vocabulary
with open(filename, "w", encoding="utf-8") as f:
for i, s in enumerate(vocab):
f.write(f"{s} {i}\n")
f.write(f"<blk> {len(vocab)}\n")



class PromptedStreamingEncoder(torch.nn.Module):
Expand All @@ -221,7 +100,6 @@ def __init__(self, asr_model):
)

self.encoder = asr_model.encoder
assert_forward_for_export_signature(self.encoder)

self.prompt_kernel = asr_model.prompt_kernel
self.num_prompts = int(asr_model.num_prompts)
Expand Down Expand Up @@ -283,24 +161,6 @@ def remove_export_scratch_files():
p.unlink()


def assert_encoder_graph(filename: str):
model = onnx.load(filename, load_external_data=False)

input_names = [i.name for i in model.graph.input]
if input_names != ENCODER_INPUT_NAMES:
raise RuntimeError(
f"{filename}: expected encoder inputs {ENCODER_INPUT_NAMES}, "
f"got {input_names}"
)

output_names = [o.name for o in model.graph.output]
if output_names != ENCODER_OUTPUT_NAMES:
raise RuntimeError(
f"{filename}: expected encoder outputs {ENCODER_OUTPUT_NAMES}, "
f"got {output_names}"
)


def _module_device_and_dtype(module):
try:
p = next(module.parameters())
Expand All @@ -322,7 +182,9 @@ def export_prompted_encoder(
):
device, dtype = _module_device_and_dtype(asr_model.encoder)

audio_signal = torch.zeros(1, 128, window_size, dtype=dtype, device=device)
feat_dim = asr_model.cfg.preprocessor.features

audio_signal = torch.zeros(1, feat_dim, window_size, dtype=dtype, device=device)
Comment on lines +185 to +187

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

While feat_dim is now dynamically retrieved from asr_model.cfg.preprocessor.features to initialize the dummy audio_signal tensor, the metadata dictionary on line 359 still hardcodes "feat_dim": 128. To ensure consistency and avoid potential issues if the model configuration changes, the metadata should also use asr_model.cfg.preprocessor.features instead of a hardcoded value.

length = torch.full((1,), window_size, dtype=torch.int64, device=device)
cache_last_channel = torch.zeros(
1,
Expand Down Expand Up @@ -365,7 +227,7 @@ def export_prompted_encoder(
"encoder.export.onnx",
input_names=ENCODER_INPUT_NAMES,
output_names=ENCODER_OUTPUT_NAMES,
opset_version=17,
opset_version=13,
dynamic_axes={
"audio_signal": {0: "batch", 2: "time"},
"length": {0: "batch"},
Expand All @@ -391,7 +253,6 @@ def export_prompted_encoder(
location="encoder.data",
size_threshold=0,
)
assert_encoder_graph("encoder.onnx")
for p in Path(".").glob("encoder.export.onnx*"):
p.unlink()

Expand All @@ -402,17 +263,10 @@ def main():

asr_model = nemo_asr.models.ASRModel.from_pretrained(model_name=model_name)

vocab_size = save_tokens(asr_model)
if vocab_size != asr_model.decoder.vocab_size:
raise ValueError(
f"SentencePiece vocab size {vocab_size} != decoder vocab size "
f"{asr_model.decoder.vocab_size}"
)
save_tokens(asr_model)

prompt_dictionary = get_prompt_dictionary(asr_model)
prompt_dictionary = asr_model.cfg.model_defaults.prompt_dictionary
auto_prompt_id = prompt_dictionary["auto"]
if auto_prompt_id != 101:
raise ValueError(f"Expected auto prompt id 101, got {auto_prompt_id}")

asr_model.eval()

Expand Down Expand Up @@ -508,20 +362,20 @@ def main():
"model_author": "NeMo",
"url": f"https://huggingface.co/{model_name}",
"comment": "Only the transducer branch is exported",
"prompt_dictionary": json.dumps(prompt_dictionary, sort_keys=True),
"prompt_dictionary": json.dumps(
_normalize_prompt_dictionary(prompt_dictionary), sort_keys=True
),
"auto_prompt_id": auto_prompt_id,
}
print("meta_data", meta_data)
add_meta_data("encoder.onnx", meta_data)
assert_encoder_graph("encoder.onnx")

for m in ["encoder", "decoder", "joiner"]:
quantize_dynamic(
model_input=f"{m}.onnx",
model_output=f"{m}.int8.onnx",
weight_type=QuantType.QUInt8,
)
assert_encoder_graph("encoder.int8.onnx")

Path(str(ms)).mkdir(exist_ok=True)
for suffix in ["onnx", "data"]:
Expand Down