Repository navigation
Support RK NPU for SenseVoice non-streaming ASR models #2589
New issue
Have a question about this project? Sign up for a free GitHub account to open an issue and contact its maintainers and the community.
By clicking “Sign up for GitHub”, you agree to our terms of service and privacy statement. We’ll occasionally send you account related emails.
Already on GitHub? Sign in to your account
Changes from all commits
9ddf6fb
ea32e88
738c627
6d8abf5
12e50f6
a76af5c
File filter
Filter by extension
Conversations
Jump to
Diff view
Diff view
There are no files selected for viewing
| Original file line number | Diff line number | Diff line change | ||||||||||||
|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|
| @@ -0,0 +1,164 @@ | ||||||||||||||
| #!/usr/bin/env python3 | ||||||||||||||
| # Copyright 2025 Xiaomi Corp. (authors: Fangjun Kuang) | ||||||||||||||
|
|
||||||||||||||
| import argparse | ||||||||||||||
| import os | ||||||||||||||
| from typing import Any, Dict, List, Tuple | ||||||||||||||
|
|
||||||||||||||
| import onnx | ||||||||||||||
| import sentencepiece as spm | ||||||||||||||
| import torch | ||||||||||||||
|
|
||||||||||||||
| from torch_model import SenseVoiceSmall | ||||||||||||||
|
|
||||||||||||||
|
|
||||||||||||||
| def get_args(): | ||||||||||||||
| parser = argparse.ArgumentParser( | ||||||||||||||
| formatter_class=argparse.ArgumentDefaultsHelpFormatter | ||||||||||||||
| ) | ||||||||||||||
|
|
||||||||||||||
| parser.add_argument( | ||||||||||||||
| "--input-len-in-seconds", | ||||||||||||||
| type=int, | ||||||||||||||
| required=True, | ||||||||||||||
| help="""RKNN does not support dynamic shape, so we need to hard-code | ||||||||||||||
| how long the model can process. | ||||||||||||||
| """, | ||||||||||||||
| ) | ||||||||||||||
| return parser.parse_args() | ||||||||||||||
|
|
||||||||||||||
|
|
||||||||||||||
| def add_meta_data(filename: str, meta_data: Dict[str, Any]): | ||||||||||||||
| """Add meta data to an ONNX model. It is changed in-place. | ||||||||||||||
|
|
||||||||||||||
| Args: | ||||||||||||||
| filename: | ||||||||||||||
| Filename of the ONNX model to be changed. | ||||||||||||||
| meta_data: | ||||||||||||||
| Key-value pairs. | ||||||||||||||
| """ | ||||||||||||||
| 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) | ||||||||||||||
|
|
||||||||||||||
| onnx.save(model, filename) | ||||||||||||||
|
|
||||||||||||||
|
|
||||||||||||||
| def load_cmvn(filename) -> Tuple[List[float], List[float]]: | ||||||||||||||
| neg_mean = None | ||||||||||||||
| inv_stddev = None | ||||||||||||||
|
|
||||||||||||||
| with open(filename) as f: | ||||||||||||||
| for line in f: | ||||||||||||||
| if not line.startswith("<LearnRateCoef>"): | ||||||||||||||
| continue | ||||||||||||||
| t = line.split()[3:-1] | ||||||||||||||
|
|
||||||||||||||
| if neg_mean is None: | ||||||||||||||
| neg_mean = list(map(lambda x: float(x), t)) | ||||||||||||||
| else: | ||||||||||||||
| inv_stddev = list(map(lambda x: float(x), t)) | ||||||||||||||
|
|
||||||||||||||
| return neg_mean, inv_stddev | ||||||||||||||
|
|
||||||||||||||
|
|
||||||||||||||
| def generate_tokens(sp): | ||||||||||||||
| with open("tokens.txt", "w", encoding="utf-8") as f: | ||||||||||||||
| for i in range(sp.vocab_size()): | ||||||||||||||
| f.write(f"{sp.id_to_piece(i)} {i}\n") | ||||||||||||||
| print("saved to tokens.txt") | ||||||||||||||
|
|
||||||||||||||
|
|
||||||||||||||
| @torch.no_grad() | ||||||||||||||
| def main(): | ||||||||||||||
| args = get_args() | ||||||||||||||
| print(vars(args)) | ||||||||||||||
|
|
||||||||||||||
| sp = spm.SentencePieceProcessor() | ||||||||||||||
| sp.load("./chn_jpn_yue_eng_ko_spectok.bpe.model") | ||||||||||||||
| vocab_size = sp.vocab_size() | ||||||||||||||
| generate_tokens(sp) | ||||||||||||||
|
|
||||||||||||||
| print("loading model") | ||||||||||||||
|
|
||||||||||||||
| state_dict = torch.load("./model.pt") | ||||||||||||||
| if "state_dict" in state_dict: | ||||||||||||||
| state_dict = state_dict["state_dict"] | ||||||||||||||
|
|
||||||||||||||
| neg_mean, inv_stddev = load_cmvn("./am.mvn") | ||||||||||||||
|
|
||||||||||||||
| neg_mean = torch.tensor(neg_mean, dtype=torch.float32) | ||||||||||||||
| inv_stddev = torch.tensor(inv_stddev, dtype=torch.float32) | ||||||||||||||
|
|
||||||||||||||
| model = SenseVoiceSmall(neg_mean=neg_mean, inv_stddev=inv_stddev) | ||||||||||||||
| model.load_state_dict(state_dict) | ||||||||||||||
| model.eval() | ||||||||||||||
| del state_dict | ||||||||||||||
|
|
||||||||||||||
| lfr_window_size = 7 | ||||||||||||||
| lfr_window_shift = 6 | ||||||||||||||
|
|
||||||||||||||
| # frame shift is 10ms, 1 second has about 100 feature frames | ||||||||||||||
| input_len_in_seconds = int(args.input_len_in_seconds) | ||||||||||||||
| num_frames = input_len_in_seconds * 100 | ||||||||||||||
| print("num_frames", num_frames) | ||||||||||||||
|
|
||||||||||||||
| # num_input_frames is an approximate number | ||||||||||||||
| num_input_frames = int(num_frames / lfr_window_shift + 0.5) | ||||||||||||||
| print("num_input_frames", num_input_frames) | ||||||||||||||
|
|
||||||||||||||
| x = torch.randn(1, num_input_frames, 560, dtype=torch.float32) | ||||||||||||||
|
|
||||||||||||||
| language = 3 | ||||||||||||||
| text_norm = 15 | ||||||||||||||
| prompt = torch.tensor([language, 1, 2, text_norm], dtype=torch.int32) | ||||||||||||||
|
|
||||||||||||||
|
Comment on lines
+117
to
+120
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. Fix dtype for nn.Embedding indices (must be int64).
- prompt = torch.tensor([language, 1, 2, text_norm], dtype=torch.int32)
+ prompt = torch.tensor([language, 1, 2, text_norm], dtype=torch.int64)📝 Committable suggestion
Suggested change
🤖 Prompt for AI Agents |
||||||||||||||
| opset_version = 13 | ||||||||||||||
| filename = f"model-{input_len_in_seconds}-seconds.onnx" | ||||||||||||||
| torch.onnx.export( | ||||||||||||||
| model, | ||||||||||||||
| (x, prompt), | ||||||||||||||
| filename, | ||||||||||||||
| opset_version=opset_version, | ||||||||||||||
| input_names=["x", "prompt"], | ||||||||||||||
| output_names=["logits"], | ||||||||||||||
| dynamic_axes={}, | ||||||||||||||
| ) | ||||||||||||||
|
|
||||||||||||||
| model_author = os.environ.get("model_author", "iic") | ||||||||||||||
| comment = os.environ.get("comment", "iic/SenseVoiceSmall") | ||||||||||||||
| url = os.environ.get("url", "https://huggingface.co/FunAudioLLM/SenseVoiceSmall") | ||||||||||||||
|
|
||||||||||||||
| meta_data = { | ||||||||||||||
| "lfr_window_size": lfr_window_size, | ||||||||||||||
| "lfr_window_shift": lfr_window_shift, | ||||||||||||||
| "num_input_frames": num_input_frames, | ||||||||||||||
| "normalize_samples": 0, # input should be in the range [-32768, 32767] | ||||||||||||||
| "model_type": "sense_voice_ctc", | ||||||||||||||
| "version": "1", | ||||||||||||||
| "model_author": model_author, | ||||||||||||||
| "maintainer": "k2-fsa", | ||||||||||||||
| "vocab_size": vocab_size, | ||||||||||||||
| "comment": comment, | ||||||||||||||
| "lang_auto": model.lid_dict["auto"], | ||||||||||||||
| "lang_zh": model.lid_dict["zh"], | ||||||||||||||
| "lang_en": model.lid_dict["en"], | ||||||||||||||
| "lang_yue": model.lid_dict["yue"], # cantonese | ||||||||||||||
| "lang_ja": model.lid_dict["ja"], | ||||||||||||||
| "lang_ko": model.lid_dict["ko"], | ||||||||||||||
| "lang_nospeech": model.lid_dict["nospeech"], | ||||||||||||||
| "with_itn": model.textnorm_dict["withitn"], | ||||||||||||||
| "without_itn": model.textnorm_dict["woitn"], | ||||||||||||||
| "url": url, | ||||||||||||||
| } | ||||||||||||||
| add_meta_data(filename=filename, meta_data=meta_data) | ||||||||||||||
|
|
||||||||||||||
|
|
||||||||||||||
| if __name__ == "__main__": | ||||||||||||||
| torch.manual_seed(20250717) | ||||||||||||||
| main() | ||||||||||||||
| Original file line number | Diff line number | Diff line change | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|
| @@ -0,0 +1,158 @@ | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| #!/usr/bin/env python3 | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| # Copyright (c) 2025 Xiaomi Corporation (authors: Fangjun Kuang) | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
|
|
||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| import argparse | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| import logging | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| from pathlib import Path | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
|
|
||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| from rknn.api import RKNN | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
|
|
||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| logging.basicConfig(level=logging.WARNING) | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
|
|
||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| g_platforms = [ | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| # "rv1103", | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| # "rv1103b", | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| # "rv1106", | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| # "rk2118", | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| "rk3562", | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| "rk3566", | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| "rk3568", | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| "rk3576", | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| "rk3588", | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| ] | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
|
|
||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
|
|
||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| def get_parser(): | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| parser = argparse.ArgumentParser( | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| formatter_class=argparse.ArgumentDefaultsHelpFormatter | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| ) | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
|
|
||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| parser.add_argument( | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| "--target-platform", | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| type=str, | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| required=True, | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| help=f"Supported values are: {','.join(g_platforms)}", | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| ) | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
|
|
||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| parser.add_argument( | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| "--in-model", | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| type=str, | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| required=True, | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| help="Path to the input onnx model", | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| ) | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
|
|
||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| parser.add_argument( | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| "--out-model", | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| type=str, | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| required=True, | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| help="Path to the output rknn model", | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| ) | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
|
|
||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| return parser | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
|
|
||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
|
|
||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| def get_meta_data(model: str): | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| import onnxruntime | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
|
|
||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| session_opts = onnxruntime.SessionOptions() | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| session_opts.inter_op_num_threads = 1 | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| session_opts.intra_op_num_threads = 1 | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
|
|
||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| m = onnxruntime.InferenceSession( | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| model, | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| sess_options=session_opts, | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| providers=["CPUExecutionProvider"], | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| ) | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
|
|
||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| for i in m.get_inputs(): | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| print(i) | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
|
|
||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| print("-----") | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
|
|
||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| for i in m.get_outputs(): | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| print(i) | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| print() | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
|
|
||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| meta = m.get_modelmeta().custom_metadata_map | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| s = "" | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| sep = "" | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| for key, value in meta.items(): | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| if key in ("neg_mean", "inv_stddev"): | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| continue | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| s = s + sep + f"{key}={value}" | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| sep = ";" | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| assert len(s) < 1024, len(s) | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
|
|
||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| print("len(s)", len(s), s) | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
|
|
||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| return s | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
|
|
||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
|
|
||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| def export_rknn(rknn, filename): | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| ret = rknn.export_rknn(filename) | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| if ret != 0: | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| exit(f"Export rknn model to {filename} failed!") | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
|
|
||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
|
|
||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| def init_model(filename: str, target_platform: str, custom_string=None): | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| rknn = RKNN(verbose=False) | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
|
|
||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| rknn.config( | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| optimization_level=0, | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| target_platform=target_platform, | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| custom_string=custom_string, | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| ) | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| if not Path(filename).is_file(): | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| exit(f"{filename} does not exist") | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
|
|
||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| ret = rknn.load_onnx(model=filename) | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| if ret != 0: | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| exit(f"Load model {filename} failed!") | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
|
|
||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| ret = rknn.build(do_quantization=False) | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| if ret != 0: | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| exit(f"Build model {filename} failed!") | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
|
|
||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| return rknn | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
|
|
||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
|
|
||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| class RKNNModel: | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| def __init__( | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| self, | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| model: str, | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| target_platform: str, | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| ): | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| meta = get_meta_data(model) | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| print(meta) | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
|
|
||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| self.model = init_model( | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| model, | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| target_platform=target_platform, | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| custom_string=meta, | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| ) | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
|
|
||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| def export_rknn(self, model): | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| export_rknn(self.model, model) | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
|
|
||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| def release(self): | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| self.model.release() | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
|
|
||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
|
|
||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| def main(): | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| args = get_parser().parse_args() | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| print(vars(args)) | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
|
|
||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| model = RKNNModel( | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| model=args.in_model, | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| target_platform=args.target_platform, | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| ) | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
|
|
||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| model.export_rknn( | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| model=args.out_model, | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| ) | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
|
|
||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| model.release() | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
|
|
||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
|
|
||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| if __name__ == "__main__": | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
| main() | ||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
|
Comment on lines
+141
to
+158
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. 🛠️ Refactor suggestion Always release RKNN; ensure output dir exists. Prevent resource leaks and path errors. def main():
args = get_parser().parse_args()
- print(vars(args))
+ logging.info("%s", vars(args))
+ # Ensure parent dir exists
+ out_path = Path(args.out_model)
+ out_path.parent.mkdir(parents=True, exist_ok=True)
- model = RKNNModel(
- model=args.in_model,
- target_platform=args.target_platform,
- )
-
- model.export_rknn(
- model=args.out_model,
- )
-
- model.release()
+ model = None
+ try:
+ model = RKNNModel(
+ model=args.in_model,
+ target_platform=args.target_platform,
+ )
+ model.export_rknn(model=str(out_path))
+ finally:
+ if model is not None:
+ model.release()📝 Committable suggestion
Suggested change
🤖 Prompt for AI Agents |
||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||||
There was a problem hiding this comment.
Choose a reason for hiding this comment
The reason will be displayed to describe this comment to others. Learn more.
Guard against tokenizer/model vocab size mismatch (breaks decoding).
Metadata
vocab_sizemust match the model’s output dim; derive from the model and assert equality with SentencePiece.Also applies to: 98-101, 145-146