File size: 1,979 Bytes
e146811
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
#!/usr/bin/env python3
"""Export a portable averaged checkpoint for publishing.
Saves {model: state_dict, config: {...arch...}, vocab_size, tokens} to one .pt (CPU)."""
import sys, argparse, torch
sys.path.insert(0, "/root/icefall/egs/hindi/ASR/zipformer"); sys.path.insert(0, "/root/icefall")
from train import add_model_arguments, get_model, get_params
from icefall.lexicon import Lexicon
from icefall.checkpoint import average_checkpoints
import os

ap = argparse.ArgumentParser(); add_model_arguments(ap)
ap.add_argument("--exp-dir", required=True); ap.add_argument("--lang-dir", required=True)
ap.add_argument("--epoch", type=int, required=True); ap.add_argument("--avg", type=int, required=True)
ap.add_argument("--out", required=True)
args = ap.parse_args()
p = get_params(); p.update(vars(args))
p.causal = True; p.chunk_size = "16,32,64,-1"; p.left_context_frames = "64,128,256,-1"
p.use_ctc = True; p.use_transducer = False
lex = Lexicon(p.lang_dir); p.blank_id = lex.token_table["<blk>"]; p.vocab_size = max(lex.tokens) + 1
model = get_model(p)
fns = [f"{p.exp_dir}/epoch-{e}.pt" for e in range(p.epoch - p.avg + 1, p.epoch + 1) if os.path.exists(f"{p.exp_dir}/epoch-{e}.pt")]
sd = average_checkpoints(fns, device=torch.device("cpu"))
model.load_state_dict(sd, strict=False)
cfg = {k: getattr(p, k) for k in ["num_encoder_layers","downsampling_factor","feedforward_dim","num_heads",
    "encoder_dim","query_head_dim","value_head_dim","pos_head_dim","pos_dim","encoder_unmasked_dim",
    "cnn_module_kernel","decoder_dim","joiner_dim","feature_dim","causal","chunk_size",
    "left_context_frames","use_ctc","use_transducer","blank_id","vocab_size","subsampling_factor"]}
torch.save({"model": model.state_dict(), "config": cfg, "vocab_size": p.vocab_size,
            "averaged_from": [os.path.basename(f) for f in fns]}, args.out)
n = sum(x.numel() for x in model.parameters())
print(f"saved {args.out}  params={n/1e6:.1f}M vocab={p.vocab_size} avg={len(fns)}ckpts")