File size: 11,423 Bytes
4ef1ac4
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
#!/usr/bin/env python3
"""
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  # noqa
    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}")

# --- Extract 22 tokenizer triples verbatim ---
print("[extract] 22 langs verbatim...")
with tarfile.open(SRC_NEMO,"r") as tf:
    cfg_bytes=tf.extractfile("./model_config.yaml").read().decode()
    # Save original for audit
    (OUT_DIR/"original_model_config.yaml").write_text(cfg_bytes,encoding="utf-8")
    # Extract all tokenizer files (keep hash names)
    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"  {name} -> {out.name} {out.stat().st_size/1024:.1f}KB")
    print(f"[extract] {len(list(TOKEN_DIR.glob('*')))} tokenizer files")

# Create new config with both identities (audit)
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'}")
# also tokenizer_map
tok_map={}
# parse hashes from cfg_bytes for audit
for lang in ordered_langs:
    # find block for this lang
    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")

# --- Load NeMo ---
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__}")

# Wrapper wav->logits (CTC) keeping 5633, preprocessor inside, language handling via int lang_idx
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  # ConvASRDecoder 512->5633
        # Ensure preprocessor in eval (dither only training)
        self.preprocessor.eval()
        self.encoder.eval()
        self.ctc_decoder.eval()
    def forward(self, wav, wav_len):
        # wav [B,T] float32, wav_len [B] int64 samples
        mel, mel_len = self.preprocessor(input_signal=wav, length=wav_len)
        enc, enc_len = self.encoder(audio_signal=mel, length=mel_len)
        # ctc_decoder expects encoder_output [B,512,T']; it internally handles language via masking but we keep full 5633
        # For verbatim, we call without language_id to get full 5633 logits (no masking)
        logits = self.ctc_decoder(encoder_output=enc)  # [B,5633,T'] ? Check transpose
        # ConvASRDecoder returns [B,5633,T']? Actually it is Conv1d 512->5633, so [B,5633,T']
        # Transpose to [B,T,5633] for consistency with BanglaNeo fused
        if logits.dim()==3 and logits.shape[1]==5633:
            logits = logits.transpose(1,2)  # [B,T,5633]
        return logits, enc_len

# Also wrapper for encoder only (for RNNT)
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

# Test
B=1
T_samples=16000  # 1 sec
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)")
    # Also test encoder
    enc_wrapper=VerbatimWavEncoder(model)
    enc, elen = enc_wrapper(dummy_wav, dummy_len)
    print(f"[test] enc {enc.shape} elen {elen}")

# Export fused wav->ctc 5633
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,
)
# Metadata: keep original + new display
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}")

# Export encoder wav->enc
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}")

# Export CTC decoder 5633 separately (enc->logits) for split usage
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.
""")