Shenava-Koochik-Lite-v1.0 / load_koochik_lite.py
Reza2kn's picture
Mirror Reza2kn/Shenava-Koochik-Lite-v1.0 at revision b20f2223ba15b8f4df13a95885710c4eee1a2c0a
191c760 verified
Raw
History Blame Contribute Delete
2.11 kB
#!/usr/bin/env python3
"""Load Shenava-Koochik-Lite-v1.0 — a LITEASR-compressed (-21.6% encoder) Koochik.
Needs the base model Reza2kn/Shenava-Koochik-v1.0 (.nemo) plus this repo's
`koochik_lite099_enc.pt` (compressed encoder state_dict) and `koochik_lite099_kmap.json`
(per-layer low-rank sizes k). LITEASR replaces each targeted encoder Linear with a
rank-k Sequential(Linear(D_in->k, bias=False), Linear(k->D_out, bias=True)); this loader
rebuilds that skeleton from the kmap, then loads the compressed weights.
from huggingface_hub import hf_hub_download, snapshot_download
base = hf_hub_download("Reza2kn/Shenava-Koochik-v1.0", "shenava-koochik-v1.0.nemo")
repo = snapshot_download("Reza2kn/Shenava-Koochik-Lite-v1.0")
from load_koochik_lite import load_koochik_lite
m = load_koochik_lite(base, f"{repo}/koochik_lite099_enc.pt", f"{repo}/koochik_lite099_kmap.json")
print(m.transcribe(["clip.wav"])[0].text)
"""
import json, torch, torch.nn as nn
import nemo.collections.asr as A
def load_koochik_lite(base_nemo, enc_sd_path, kmap_path, device="cuda", ctc=True):
m = A.models.ASRModel.restore_from(base_nemo, map_location=device).eval()
if ctc:
try:
m.change_decoding_strategy(decoder_type="ctc")
except Exception:
pass
def parent(root, name):
parts = name.split("."); p = root
for x in parts[:-1]:
p = getattr(p, x)
return p, parts[-1]
for n, k in json.load(open(kmap_path)).items():
par, at = parent(m.encoder, n)
mod = getattr(par, at)
Dout, Din = mod.weight.shape
setattr(par, at, nn.Sequential(nn.Linear(Din, k, bias=False), nn.Linear(k, Dout, bias=True)))
m.encoder.load_state_dict(torch.load(enc_sd_path, map_location=device), strict=True)
return m.to(device).eval()
if __name__ == "__main__":
import sys
m = load_koochik_lite(sys.argv[1], sys.argv[2], sys.argv[3], device="cuda" if torch.cuda.is_available() else "cpu")
for r in m.transcribe(sys.argv[4:], batch_size=8):
print(r.text if hasattr(r, "text") else r)