"""fp32 ONNX export of a NeMo transducer for sherpa-onnx, reproducing k2-fsa's own export verbatim (k2-fsa/sherpa-onnx @040afe36, scripts/nemo/{parakeet-tdt-0.6b-v3,parakeet-unified-en-0.6b}/export_onnx.py): same encoder/decoder/joint .export() calls, same tokens.txt, same metadata, encoder weights as external data. The ONLY omission is their final quantize_dynamic() step: the point is the non-int8 graph on the CUDA EP. usage: export_onnx.py NEMO_PATH OUT_DIR URL COMMENT """ import os import sys import onnx import torch import nemo.collections.asr as nemo_asr def add_meta_data(filename, meta_data): 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) if os.path.basename(filename) == "encoder.onnx": onnx.save(model, filename, save_as_external_data=True, all_tensors_to_one_file=True, location="encoder.weights") else: onnx.save(model, filename) @torch.no_grad() def main(): nemo_path, out, url, comment = sys.argv[1:5] os.makedirs(out, exist_ok=True) os.chdir(out) m = nemo_asr.models.ASRModel.restore_from(restore_path=nemo_path, map_location="cpu") m.eval() if m.cfg.get("validation_ds") is None: m.cfg.validation_ds = dict() with open("./tokens.txt", "w", encoding="utf-8") as f: for i, s in enumerate(m.joint.vocabulary): f.write(f"{s} {i}\n") f.write(f" {i+1}\n") m.encoder.export("encoder.onnx") m.decoder.export("decoder.onnx") m.joint.export("joiner.onnx") normalize_type = m.cfg.preprocessor.normalize if normalize_type == "NA": normalize_type = "" meta = { "vocab_size": m.decoder.vocab_size, "normalize_type": normalize_type, "pred_rnn_layers": m.decoder.pred_rnn_layers, "pred_hidden": m.decoder.pred_hidden, "subsampling_factor": m.encoder.subsampling_factor, "model_type": "EncDecRNNTBPEModel", "version": "2", "model_author": "NeMo", "url": url, "comment": comment, "feat_dim": 128, } add_meta_data("encoder.onnx", meta) print("meta", meta) os.system("ls -la") if __name__ == "__main__": main()