fela-streaming-asr / train.py
itstheraj's picture
initial commit
c513220
Raw
History Blame Contribute Delete
6.44 kB
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()