#!/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[""]; 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")