| |
| """LM NEURONAL de rescoring (§6.1) — petit Transformer causal au NIVEAU CARACTÈRE, |
| entraîné UNIQUEMENT sur le texte train WAXAL de la langue (conforme : aucune donnée externe). |
| Le char-level est robuste sur 14k phrases et colle à la métrique (moitié CER). |
| Complète le KenLM 5-grammes : dépendances bien plus longues. |
| |
| Usage : LANG=lin /root/venv/bin/python train_charlm.py |
| Sortie : /scratch/lm/charlm_<lang>.pt (poids + vocab) |
| """ |
| import json, math, os, time |
| import torch, torch.nn as nn |
|
|
| LANG = os.environ.get("LANG_ASR", "lin") |
| MANIFEST = "/root/devhard/train_%s_min.jsonl" % LANG |
| OUT = "/scratch/lm/charlm_%s.pt" % LANG |
| DIM = int(os.environ.get("DIM", "256")) |
| NLAYER = int(os.environ.get("NLAYER", "4")) |
| NHEAD = 4 |
| EPOCHS = int(os.environ.get("EPOCHS", "12")) |
| BS = 64 |
| MAXLEN = 256 |
|
|
|
|
| def main(): |
| rows = [json.loads(l) for l in open(MANIFEST, encoding="utf-8")] |
| texts = [r["text"].strip() for r in rows if r.get("text", "").strip()] |
| print("%s : %d phrases" % (LANG, len(texts)), flush=True) |
|
|
| chars = sorted(set("".join(texts))) |
| stoi = {c: i + 2 for i, c in enumerate(chars)} |
| V = len(stoi) + 2 |
| print("vocab caracteres : %d" % V, flush=True) |
|
|
| def enc(t): |
| return [1] + [stoi[c] for c in t if c in stoi][:MAXLEN - 2] + [1] |
|
|
| seqs = [enc(t) for t in texts] |
| |
| nval = max(200, len(seqs) // 20) |
| val, tr = seqs[:nval], seqs[nval:] |
|
|
| class LM(nn.Module): |
| def __init__(self): |
| super().__init__() |
| self.emb = nn.Embedding(V, DIM, padding_idx=0) |
| self.pos = nn.Embedding(MAXLEN, DIM) |
| layer = nn.TransformerEncoderLayer(DIM, NHEAD, DIM * 4, dropout=0.1, |
| batch_first=True, norm_first=True) |
| self.tr = nn.TransformerEncoder(layer, NLAYER) |
| self.ln = nn.LayerNorm(DIM) |
| self.head = nn.Linear(DIM, V) |
|
|
| def forward(self, x): |
| T = x.shape[1] |
| h = self.emb(x) + self.pos(torch.arange(T, device=x.device))[None] |
| mask = nn.Transformer.generate_square_subsequent_mask(T, device=x.device) |
| h = self.tr(h, mask=mask, is_causal=True) |
| return self.head(self.ln(h)) |
|
|
| dev = "cuda" |
| model = LM().to(dev) |
| nparam = sum(p.numel() for p in model.parameters()) |
| print("parametres : %.2f M" % (nparam / 1e6), flush=True) |
| opt = torch.optim.AdamW(model.parameters(), lr=3e-4, weight_decay=0.01) |
| lossf = nn.CrossEntropyLoss(ignore_index=0) |
|
|
| def batches(data, bs, shuffle=True): |
| idx = torch.randperm(len(data)).tolist() if shuffle else range(len(data)) |
| buf = [] |
| for i in idx: |
| buf.append(data[i]) |
| if len(buf) == bs: |
| yield buf; buf = [] |
| if buf: |
| yield buf |
|
|
| def pad(b): |
| L = max(len(x) for x in b) |
| t = torch.zeros(len(b), L, dtype=torch.long) |
| for i, x in enumerate(b): |
| t[i, :len(x)] = torch.tensor(x) |
| return t.to(dev) |
|
|
| best = 1e9 |
| for ep in range(EPOCHS): |
| model.train(); tot = n = 0 |
| for b in batches(tr, BS): |
| x = pad(b) |
| logits = model(x[:, :-1]) |
| loss = lossf(logits.reshape(-1, V), x[:, 1:].reshape(-1)) |
| opt.zero_grad(); loss.backward() |
| torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0) |
| opt.step() |
| tot += float(loss); n += 1 |
| model.eval(); vt = vn = 0 |
| with torch.inference_mode(): |
| for b in batches(val, BS, False): |
| x = pad(b) |
| vt += float(lossf(model(x[:, :-1]).reshape(-1, V), x[:, 1:].reshape(-1))); vn += 1 |
| vl = vt / max(vn, 1) |
| print("ep %d train %.4f val %.4f (ppl %.2f)" % (ep, tot / max(n, 1), vl, math.exp(vl)), flush=True) |
| if vl < best: |
| best = vl |
| torch.save({"stoi": stoi, "dim": DIM, "nlayer": NLAYER, "V": V, |
| "state": model.state_dict()}, OUT) |
| print("CHARLM_DONE %s best_val %.4f ppl %.2f -> %s" % (LANG, best, math.exp(best), OUT), flush=True) |
|
|
|
|
| if __name__ == "__main__": |
| main() |
|
|