#!/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)