File size: 2,113 Bytes
191c760
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
#!/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)