import sys, os, time, argparse, random, json sys.path.insert(0, "/workspace") import torch, torchaudio, torch.nn.functional as F from torch.utils.data import DataLoader, ConcatDataset DEV = "cuda" def collate(batch): wavs, texts = [], [] for wav, sr, txt, *_ in batch: wavs.append(wav.squeeze(0)); texts.append(txt) lens = [w.shape[0] for w in wavs]; ml = max(lens) wav = torch.zeros(len(wavs), ml) for i, w in enumerate(wavs): wav[i, :w.shape[0]] = w return wav, lens, texts def out_frames(n): return ((n//160)+1+3)//4 def main(): ap = argparse.ArgumentParser() ap.add_argument("--steps", type=int, default=40000) ap.add_argument("--bs", type=int, default=24) ap.add_argument("--lr", type=float, default=5e-4) ap.add_argument("--n_layer", type=int, default=16) ap.add_argument("--n_embd", type=int, default=512) ap.add_argument("--n_head", type=int, default=8) ap.add_argument("--ffn_mult", type=int, default=4) ap.add_argument("--pattern", type=str, default="FNO") ap.add_argument("--gla_delta", type=int, default=0) ap.add_argument("--fno_modes", type=int, default=256) ap.add_argument("--max_sec", type=float, default=16.0) ap.add_argument("--ckpt", type=str, default="/workspace/fela_asr.pt") ap.add_argument("--init_from", type=str, default="") ap.add_argument("--eval_only", action="store_true") ap.add_argument("--eval_n", type=int, default=2620) ap.add_argument("--subsets", type=str, default="train-clean-100,train-clean-360") ap.add_argument("--tag", type=str, default="run") ap.add_argument("--save_every", type=int, default=5000) ap.add_argument("--data", type=str, default="/workspace/data") ap.add_argument("--smoke", action="store_true") args = ap.parse_args() torch.manual_seed(0); random.seed(0) parts = [torchaudio.datasets.LIBRISPEECH(args.data, url=s, download=False) for s in args.subsets.split(",")] train = ConcatDataset(parts) test = torchaudio.datasets.LIBRISPEECH(args.data, url="test-clean", download=False) other = torchaudio.datasets.LIBRISPEECH(args.data, url="test-other", download=False) n_train = len(train); n_test = len(test); n_other = len(other) if args.smoke: assert n_test == 2620 assert n_other == 2939 print(f"[TEST] subsets={args.subsets} train={n_train} test-clean={n_test} test-other={n_other}", flush=True) return from fela_ctc2 import FELACTC2, BPE, greedy_decode_bpe, wer, BLANK from model_cpu_gpt2 import CPUGPTConfig, _layer_is_gla def evaluate(model, bpe, test_url, eval_n, tag): t = torchaudio.datasets.LIBRISPEECH(args.data, url=test_url, download=False) model.eval(); te = tw = 0; n = 0 with torch.no_grad(): for i in range(min(eval_n, len(t))): wav, sr, txt, *_ = t[i] with torch.autocast("cuda", dtype=torch.bfloat16): lp = model(wav.to(DEV)) h = greedy_decode_bpe(lp[0].float(), bpe); e, wc = wer(txt.lower(), h); te += e; tw += wc; n += 1 WER = 100*te/tw print(f"=== {tag} WER {test_url} (greedy BPE, no LM): {WER:.2f}% over {n} utts, {tw} words ===", flush=True) return WER bpe = BPE(); VOCAB = bpe.vocab dl = DataLoader(train, batch_size=args.bs, shuffle=True, num_workers=12, collate_fn=collate, drop_last=True, persistent_workers=True) cfg = CPUGPTConfig(); cfg.n_layer = args.n_layer; cfg.n_embd = args.n_embd; cfg.n_head = args.n_head cfg.layer_pattern = args.pattern; cfg.seq_len = 4096; cfg.gla_chunk = 64; cfg.fno_modes = args.fno_modes cfg.ffn_hidden = args.n_embd*args.ffn_mult//16*16; cfg.gla_delta = bool(args.gla_delta); cfg.dropout = 0.0 model = FELACTC2(cfg, VOCAB).to(DEV) print(f"Params={model.param_count()/1e6:.2f}M pattern={args.pattern} gla_delta={cfg.gla_delta} BPE_vocab={VOCAB}", flush=True) if args.init_from and os.path.exists(args.init_from): sd = torch.load(args.init_from, map_location=DEV) model.load_state_dict(sd, strict=False) if args.eval_only: if os.path.exists(args.ckpt): model.load_state_dict(torch.load(args.ckpt, map_location=DEV)) wc = evaluate(model, bpe, "test-clean", args.eval_n, args.tag) wo = evaluate(model, bpe, "test-other", args.eval_n, args.tag) print(json.dumps({"tag": args.tag, "wer_clean": wc, "wer_other": wo}), flush=True) return opt = torch.optim.AdamW(model.parameters(), lr=args.lr, weight_decay=0.01, betas=(0.9, 0.95)) sched = torch.optim.lr_scheduler.OneCycleLR(opt, max_lr=args.lr, total_steps=args.steps, pct_start=0.04, anneal_strategy="cos") ctc = torch.nn.CTCLoss(blank=BLANK, zero_infinity=True) model.train(); step = 0; t0 = time.time(); run = 0.0; max_samp = int(args.max_sec*16000) while step < args.steps: for wav, lens, texts in dl: if step >= args.steps: break if wav.shape[1] > max_samp: wav = wav[:, :max_samp]; lens = [min(l, max_samp) for l in lens] wav = wav.to(DEV) tgt = [torch.tensor(bpe.encode(t), dtype=torch.long) for t in texts] tl = torch.tensor([len(t) for t in tgt]) if (tl == 0).any(): continue tc = torch.cat(tgt).to(DEV) with torch.autocast("cuda", dtype=torch.bfloat16): logp = model(wav, augment=True) il = torch.tensor([min(out_frames(l), logp.shape[1]) for l in lens]) if (il < tl).any(): continue loss = ctc(logp.transpose(0, 1).float(), tc, il, tl) if not torch.isfinite(loss): continue opt.zero_grad(); loss.backward() torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0) opt.step(); sched.step(); run += loss.item(); step += 1 if step % 100 == 0: print(f"[{args.tag}] step {step}/{args.steps} loss={run/100:.3f} {step*args.bs/(time.time()-t0):.0f} clips/s", flush=True); run = 0.0 if step % args.save_every == 0: torch.save(model.state_dict(), args.ckpt) torch.save(model.state_dict(), args.ckpt) wc = evaluate(model, bpe, "test-clean", args.eval_n, args.tag) wo = evaluate(model, bpe, "test-other", args.eval_n, args.tag) print(json.dumps({"tag": args.tag, "wer_clean": wc, "wer_other": wo, "subsets": args.subsets}), flush=True) if __name__ == "__main__": main()