Phase-1: from-scratch Zipformer-M CTC streaming (Hindi/Hinglish) + full training scripts
e146811 verified | #!/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") | |