| |
| """ |
| bengali-asr — 1:1 verbatim ONNX export (keeps ai4bharat identity + 22 langs) |
| WAV [B,T] + length -> logits [B,T/4,5633] (CTC) + encoder (for RNNT) |
| New folder identity: bengali-asr (display), but weights/config verbatim. |
| """ |
| import os |
| os.environ["TORCH_FORCE_NO_WEIGHTS_ONLY_LOAD"]="1" |
| os.environ["PYTHONUTF8"]="1" |
| import sys, types, importlib.machinery |
| if "datasets" not in sys.modules: |
| try: |
| import datasets, pyarrow |
| except Exception: |
| stub=types.ModuleType("datasets"); stub.__path__=[]; stub.__version__="3.0.0" |
| stub.__spec__=importlib.machinery.ModuleSpec("datasets",None) |
| stub.load_dataset=lambda **k: []; stub.concatenate_datasets=lambda x: x[0] if x else [] |
| sys.modules["datasets"]=stub |
| dist=types.ModuleType("datasets.distributed") |
| dist.__spec__=importlib.machinery.ModuleSpec("datasets.distributed",None) |
| dist.split_dataset_by_node=lambda d,r,w: d |
| sys.modules["datasets.distributed"]=dist |
| try: |
| import pytorch_lightning.loggers as _pl |
| if not hasattr(_pl,"NeptuneLogger"): |
| _pl.NeptuneLogger=type("NeptuneLogger",(),{}) |
| except: pass |
| try: |
| sys.stdout.reconfigure(encoding="utf-8"); sys.stderr.reconfigure(encoding="utf-8") |
| except: pass |
|
|
| import torch, tarfile, json, shutil, re |
| from pathlib import Path |
| import onnx |
|
|
| SRC_NEMO = Path(r"C:\Users\Riasat\BengaliSTT\models\indicconformer_stt_bn_hybrid_rnnt_large.nemo") |
| OUT_DIR = Path(r"C:\Users\Riasat\BengaliSTT\bengali-asr\models") |
| OUT_DIR.mkdir(parents=True, exist_ok=True) |
| TOKEN_DIR = OUT_DIR / "tokenizer" |
| TOKEN_DIR.mkdir(exist_ok=True) |
|
|
| NEW_DISPLAY = "bengali-asr" |
| NEW_VERSION = "1.0.0" |
| ORIGINAL = "ai4bharat/indicconformer_stt_bn_hybrid_ctc_rnnt_large" |
|
|
| print(f"[verbatim] {NEW_DISPLAY} 1:1 from {SRC_NEMO} -> {OUT_DIR}") |
|
|
| |
| print("[extract] 22 langs verbatim...") |
| with tarfile.open(SRC_NEMO,"r") as tf: |
| cfg_bytes=tf.extractfile("./model_config.yaml").read().decode() |
| |
| (OUT_DIR/"original_model_config.yaml").write_text(cfg_bytes,encoding="utf-8") |
| |
| for m in tf.getmembers(): |
| name=m.name.lstrip("./") |
| if "_tokenizer" in name or "_vocab" in name: |
| out=TOKEN_DIR/name |
| with tf.extractfile(m) as src, open(out,"wb") as dst: |
| shutil.copyfileobj(src,dst) |
| |
| print(f"[extract] {len(list(TOKEN_DIR.glob('*')))} tokenizer files") |
|
|
| |
| ordered_langs=["as","bn","brx","doi","kok","gu","hi","kn","ks","mai","ml","mr","mni","ne","or","pa","sa","sat","sd","ta","te","ur"] |
| new_cfg={ |
| "display_name": NEW_DISPLAY, |
| "model_id": "bengali-asr-hybrid-1.0", |
| "version": NEW_VERSION, |
| "original_source": ORIGINAL, |
| "original_model": "indicconformer_stt_bn_hybrid_ctc_rnnt_large", |
| "description": "bengali-asr 1:1 verbatim ONNX — same Conformer-L 120M 17x512 5633 vocab (22*256+blank) + preprocessor wav->logits, rebranded folder but weights/config verbatim for audit. Use with language_id='bn' (Bengali) or any of 22.", |
| "language_primary": "bn", |
| "languages_all": ordered_langs, |
| "vocab_per_lang": 256, |
| "vocab_total": 5632, |
| "vocab_plus_blank": 5633, |
| "blank_id": 5632, |
| "blank_is_shared": True, |
| "sample_rate": 16000, |
| "architecture": "Conformer-L 120M (17 blocks, 512 dim, 8 heads, subsample 4) + RNNT decoder 640 + joint 640 + aux CTC 5633", |
| "preprocessor": "AudioToMelSpectrogramPreprocessor n_fft 512 win 400 hop 160 mel 80 log+per_feature (inside ONNX)", |
| "decoders": ["ctc","rnnt"], |
| "input_wav": "wav [B,T] float32 16k mono + wav_len [B] int64 samples", |
| "output_logits": "ctc_logits [B,T/4,5633] float32 (use language mask for bn: 256+blank)", |
| "providers": ["CUDAExecutionProvider","CPUExecutionProvider","OpenVINOExecutionProvider","CoreMLExecutionProvider"], |
| "export_date": "2026-08-26", |
| "onnx_opset": 18, |
| "rebrand_note": "Folder/models renamed to bengali-asr for new identity, but ONNX metadata keeps original_source for audit; weights are verbatim copy (no slicing like BanglaNeo 257)." |
| } |
| (OUT_DIR/"config.json").write_text(json.dumps(new_cfg,indent=2,ensure_ascii=False),encoding="utf-8") |
| print(f"[config] {OUT_DIR/'config.json'}") |
| |
| tok_map={} |
| |
| for lang in ordered_langs: |
| |
| block=re.search(rf"{lang}:\s+dir:[^\n]+\n\s+type:\s+bpe\s+model_path:\s+nemo:([^\s]+)\s+vocab_path:\s+nemo:([^\s]+)\s+spe_tokenizer_vocab:\s+nemo:([^\s]+)",cfg_bytes,re.DOTALL) |
| if block: |
| tok_map[lang]={"model":block.group(1),"vocab":block.group(2),"spe_vocab":block.group(3)} |
| (OUT_DIR/"tokenizer_map.json").write_text(json.dumps(tok_map,indent=2),encoding="utf-8") |
| print(f"[tokenizer_map] {len(tok_map)} langs") |
|
|
| |
| print("[load] NeMo...") |
| from nemo.collections.asr.models import EncDecHybridRNNTCTCBPEModel |
| model=EncDecHybridRNNTCTCBPEModel.restore_from(str(SRC_NEMO),map_location="cpu") |
| model.eval(); model.freeze() |
| print(f"[load] {type(model).__name__} preproc {type(model.preprocessor).__name__} encoder {type(model.encoder).__name__}") |
|
|
| |
| class VerbatimWavCTC(torch.nn.Module): |
| def __init__(self, m): |
| super().__init__() |
| self.preprocessor=m.preprocessor |
| self.encoder=m.encoder |
| self.ctc_decoder=m.ctc_decoder |
| |
| self.preprocessor.eval() |
| self.encoder.eval() |
| self.ctc_decoder.eval() |
| def forward(self, wav, wav_len): |
| |
| mel, mel_len = self.preprocessor(input_signal=wav, length=wav_len) |
| enc, enc_len = self.encoder(audio_signal=mel, length=mel_len) |
| |
| |
| logits = self.ctc_decoder(encoder_output=enc) |
| |
| |
| if logits.dim()==3 and logits.shape[1]==5633: |
| logits = logits.transpose(1,2) |
| return logits, enc_len |
|
|
| |
| class VerbatimWavEncoder(torch.nn.Module): |
| def __init__(self, m): |
| super().__init__() |
| self.preprocessor=m.preprocessor |
| self.encoder=m.encoder |
| self.preprocessor.eval(); self.encoder.eval() |
| def forward(self, wav, wav_len): |
| mel, mel_len = self.preprocessor(input_signal=wav, length=wav_len) |
| enc, enc_len = self.encoder(audio_signal=mel, length=mel_len) |
| return enc, enc_len |
|
|
| |
| B=1 |
| T_samples=16000 |
| dummy_wav=torch.randn(B, T_samples) |
| dummy_len=torch.tensor([T_samples], dtype=torch.int64) |
| print(f"[dummy] wav {dummy_wav.shape} len {dummy_len}") |
| with torch.no_grad(): |
| wrapper=VerbatimWavCTC(model) |
| logits, enc_len = wrapper(dummy_wav, dummy_len) |
| print(f"[test] logits {logits.shape} enc_len {enc_len} (expected [1, ~100, 5633] for 1 sec, T/4 ~100)") |
| |
| enc_wrapper=VerbatimWavEncoder(model) |
| enc, elen = enc_wrapper(dummy_wav, dummy_len) |
| print(f"[test] enc {enc.shape} elen {elen}") |
|
|
| |
| fused_path=OUT_DIR/"bengali-asr-wav-ctc-5633.onnx" |
| print(f"[export] fused wav->5633 -> {fused_path}") |
| wrapper=VerbatimWavCTC(model) |
| wrapper.eval() |
| torch.onnx.export( |
| wrapper, |
| (dummy_wav, dummy_len), |
| str(fused_path), |
| input_names=["wav","wav_len"], |
| output_names=["logits","encoded_len"], |
| dynamic_axes={"wav":{1:"T_audio"},"logits":{1:"T_enc"}}, |
| opset_version=18, |
| do_constant_folding=True, |
| ) |
| |
| try: |
| m=onnx.load(str(fused_path)) |
| m.metadata_props.clear() |
| for k,v in [("model_name","bengali-asr"),("display_name","bengali-asr"),("original_source",ORIGINAL),("original_model","indicconformer_stt_bn_hybrid_ctc_rnnt_large"),("version",NEW_VERSION),("vocab_size","5632"),("vocab_plus_blank","5633"),("langs",",".join(ordered_langs)),("blank_id","5632"),("input","wav 16k mono"),("preprocessor","inside")]: |
| e=m.metadata_props.add(); e.key=k; e.value=v |
| onnx.save(m,str(fused_path)) |
| onnx.checker.check_model(m) |
| print(f"[onnx] fused checker ok") |
| except Exception as e: |
| print(f"[onnx] fused check fail {e}") |
|
|
| |
| enc_path=OUT_DIR/"bengali-asr-encoder-wav.onnx" |
| print(f"[export] encoder wav->enc -> {enc_path}") |
| enc_wrapper=VerbatimWavEncoder(model) |
| torch.onnx.export( |
| enc_wrapper, |
| (dummy_wav, dummy_len), |
| str(enc_path), |
| input_names=["wav","wav_len"], |
| output_names=["encoded","encoded_len"], |
| dynamic_axes={"wav":{1:"T_audio"},"encoded":{2:"T_enc"}}, |
| opset_version=18, |
| do_constant_folding=True, |
| ) |
| try: |
| m=onnx.load(str(enc_path)) |
| m.metadata_props.clear() |
| for k,v in [("model_name","bengali-asr-encoder"),("original_source",ORIGINAL)]: |
| e=m.metadata_props.add(); e.key=k; e.value=v |
| onnx.save(m,str(enc_path)) |
| onnx.checker.check_model(m) |
| print(f"[onnx] enc checker ok") |
| except Exception as e: |
| print(f"[onnx] enc check fail {e}") |
|
|
| |
| class CTC5633(torch.nn.Module): |
| def __init__(self, dec): super().__init__(); self.dec=dec; self.dec.eval() |
| def forward(self, enc): |
| logits=self.dec(encoder_output=enc) |
| if logits.dim()==3 and logits.shape[1]==5633: |
| logits=logits.transpose(1,2) |
| return logits |
|
|
| dummy_enc=torch.randn(1,512,100) |
| ctc_path=OUT_DIR/"bengali-asr-decoder-ctc-5633.onnx" |
| print(f"[export] ctc decoder 512->5633 -> {ctc_path}") |
| ctc_wrapper=CTC5633(model.ctc_decoder) |
| torch.onnx.export( |
| ctc_wrapper, |
| (dummy_enc,), |
| str(ctc_path), |
| input_names=["encoded"], |
| output_names=["logits"], |
| dynamic_axes={"encoded":{2:"T_enc"},"logits":{1:"T_enc"}}, |
| opset_version=18, |
| do_constant_folding=True, |
| ) |
| try: |
| m=onnx.load(str(ctc_path)); m.metadata_props.clear() |
| e=m.metadata_props.add(); e.key="model_name"; e.value="bengali-asr-decoder-ctc-5633" |
| onnx.save(m,str(ctc_path)); onnx.checker.check_model(m); print(f"[onnx] ctc dec ok") |
| except Exception as e: |
| print(f"[onnx] ctc dec fail {e}") |
|
|
| print(f"[done] files in {OUT_DIR}:") |
| for p in OUT_DIR.rglob("*"): |
| if p.is_file(): |
| print(f" {p.relative_to(OUT_DIR)} {p.stat().st_size/1024/1024:.1f} MB" if p.stat().st_size>1024*1024 else f" {p.relative_to(OUT_DIR)} {p.stat().st_size/1024:.1f} KB") |
| print(""" |
| bengali-asr 1:1 verbatim created: |
| fused wav->5633: bengali-asr-wav-ctc-5633.onnx (+.data) # wav in, CTC 5633 out, preprocessor inside, keep ai4bharat |
| encoder wav: bengali-asr-encoder-wav.onnx |
| decoder ctc 5633: bengali-asr-decoder-ctc-5633.onnx |
| tokenizer: 22 langs verbatim in tokenizer/ |
| Use language_id='bn' via masking outside (logits[:, :, bn_idx*256:(bn_idx+1)*256] + blank) or keep full 5633 and let decoder handle. |
| """) |
|
|