#!/usr/bin/env python3 """Vietnamese sentence -> VSL gloss, by fine-tuning BARTpho on Full_TriVis. Stage 0 of the full text->pose pipeline: text --[this model]--> gloss --[gpt_vsl_front_lab_v2]--> pose tokens --> skeleton Data comes from `Full_TriVis/split_lab_front.csv` (`Sentence` -> `Sign_sentence`), using the SAME train/val/test split as the pose models, so the end-to-end evaluation stays leak-free and the two stages are directly composable. Pairs are deduplicated within each split: the CSV repeats each sentence across signers and sessions (24,151 rows but only ~12k unique pairs), and duplicates would just inflate the epoch without adding signal. Targets keep the `|` sign separators of `Sign_sentence`, since sign boundaries are useful downstream. `--emit-json` writes predictions keyed by clip name, with the separators also stripped into the space-joined form the pose model consumes, so eval_vsl.py can condition on predicted gloss via --text-override. Metrics: exact match, plus WER over the SIGN sequence (edit distance / #ref signs), which is the interpretable number for this task -- a gloss is a short sign list, so sign-level edit distance says directly how many signs the model got wrong. """ import argparse import csv import json import os import random import numpy as np import torch from torch.utils.data import DataLoader, Dataset REPO = os.path.join(os.path.dirname(os.path.abspath(__file__)), '..') def norm_gloss(s): return ' '.join(t.strip() for t in str(s).split('|') if t.strip()) def signs(s): """Gloss string -> list of signs (pipe-delimited if present, else whitespace).""" s = str(s) parts = [t.strip() for t in s.split('|')] if '|' in s else s.split() return [p for p in parts if p] def wer(ref, hyp): """Edit distance between two sign sequences, normalized by reference length.""" r, h = signs(ref), signs(hyp) if not r: return 0.0 if not h else 1.0 d = np.zeros((len(r) + 1, len(h) + 1), dtype=np.int32) d[:, 0] = np.arange(len(r) + 1) d[0, :] = np.arange(len(h) + 1) for i in range(1, len(r) + 1): for j in range(1, len(h) + 1): d[i, j] = min(d[i - 1, j] + 1, d[i, j - 1] + 1, d[i - 1, j - 1] + (r[i - 1] != h[j - 1])) return d[len(r), len(h)] / len(r) class PairDS(Dataset): def __init__(self, rows, tok, max_src=64, max_tgt=64): self.rows, self.tok = rows, tok self.max_src, self.max_tgt = max_src, max_tgt def __len__(self): return len(self.rows) def __getitem__(self, i): return self.rows[i]['sentence'], self.rows[i]['gloss'] def collate(self, batch): src = [b[0] for b in batch] tgt = [b[1] for b in batch] x = self.tok(src, padding=True, truncation=True, max_length=self.max_src, return_tensors='pt') y = self.tok(text_target=tgt, padding=True, truncation=True, max_length=self.max_tgt, return_tensors='pt') labels = y['input_ids'].clone() labels[labels == self.tok.pad_token_id] = -100 x['labels'] = labels return x def load_rows(csv_path): by_split = {} seen = {} with open(csv_path, newline='', encoding='utf-8') as f: for r in csv.DictReader(f): sp = r['split'] name = os.path.splitext(os.path.basename(r['npz_path']))[0] sent, gl = r['Sentence'].strip(), r['Sign_sentence'].strip() by_split.setdefault(sp, []).append( {'name': name, 'sentence': sent, 'gloss': gl}) seen.setdefault(sp, {}).setdefault((sent, gl), name) uniq = {k: [{'name': n, 'sentence': s, 'gloss': g} for (s, g), n in v.items()] for k, v in seen.items()} return by_split, uniq @torch.no_grad() def generate(model, tok, sents, device, num_beams=4, max_len=64, bs=32): out = [] model.eval() for i in range(0, len(sents), bs): x = tok(sents[i:i + bs], padding=True, truncation=True, max_length=64, return_tensors='pt').to(device) g = model.generate(**x, num_beams=num_beams, max_length=max_len, early_stopping=True) out += tok.batch_decode(g, skip_special_tokens=True) model.train() return out def score(refs, hyps): em = np.mean([norm_gloss(r) == norm_gloss(h) for r, h in zip(refs, hyps)]) w = np.mean([wer(r, h) for r, h in zip(refs, hyps)]) # sign-level P/R/F over multisets from collections import Counter tp = fp = fn = 0 for r, h in zip(refs, hyps): cr, ch = Counter(signs(r)), Counter(signs(h)) inter = sum((cr & ch).values()) tp += inter fp += sum(ch.values()) - inter fn += sum(cr.values()) - inter p = tp / max(tp + fp, 1) rc = tp / max(tp + fn, 1) return {'exact_match': float(em), 'wer': float(w), 'precision': p, 'recall': rc, 'f1': 2 * p * rc / max(p + rc, 1e-9)} def main(): ap = argparse.ArgumentParser() ap.add_argument('--csv', default=os.path.join(REPO, 'Full_TriVis', 'split_lab_front.csv')) ap.add_argument('--model', default='vinai/bartpho-syllable-base') ap.add_argument('--out-dir', default='output_vsl/text2gloss') ap.add_argument('--epochs', type=int, default=12) ap.add_argument('--batch-size', type=int, default=24) ap.add_argument('--lr', type=float, default=3e-5) ap.add_argument('--device', default='cuda') ap.add_argument('--seed', type=int, default=42) ap.add_argument('--eval-n', type=int, default=600, help='val pairs scored per epoch') ap.add_argument('--emit-json', default='output_vsl/text2gloss/pred_test.json') ap.add_argument('--resume', default=None) args = ap.parse_args() torch.manual_seed(args.seed); random.seed(args.seed); np.random.seed(args.seed) os.makedirs(args.out_dir, exist_ok=True) device = torch.device(args.device) from transformers import AutoModelForSeq2SeqLM, AutoTokenizer tok = AutoTokenizer.from_pretrained(args.model) model = AutoModelForSeq2SeqLM.from_pretrained(args.resume or args.model).to(device) print(f'{args.model}: {sum(p.numel() for p in model.parameters())/1e6:.1f}M params') allrows, uniq = load_rows(args.csv) print({k: f'{len(allrows[k])} rows / {len(uniq[k])} unique pairs' for k in sorted(allrows)}) tr = PairDS(uniq['train'], tok) dl = DataLoader(tr, args.batch_size, shuffle=True, collate_fn=tr.collate, num_workers=4, drop_last=True) val = uniq['val'][:args.eval_n] opt = torch.optim.AdamW(model.parameters(), lr=args.lr, weight_decay=0.01) steps = args.epochs * len(dl) sch = torch.optim.lr_scheduler.OneCycleLR(opt, max_lr=args.lr, total_steps=steps, pct_start=0.06) print(f'train {len(tr)} pairs, {len(dl)} steps/epoch, {steps} total') best = 1e9 hist = [] for ep in range(1, args.epochs + 1): model.train(); tot = n = 0 for batch in dl: batch = {k: v.to(device) for k, v in batch.items()} loss = model(**batch).loss opt.zero_grad(); loss.backward() torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0) opt.step(); sch.step() tot += loss.item(); n += 1 hyps = generate(model, tok, [r['sentence'] for r in val], device) m = score([r['gloss'] for r in val], hyps) m['epoch'] = ep; m['train_loss'] = tot / max(n, 1) hist.append(m) print(f"ep {ep:2d} loss {m['train_loss']:.4f} | val WER {m['wer']:.4f} " f"EM {m['exact_match']:.4f} F1 {m['f1']:.4f}") if m['wer'] < best: best = m['wer'] model.save_pretrained(os.path.join(args.out_dir, 'best')) tok.save_pretrained(os.path.join(args.out_dir, 'best')) print(f' --> new best (WER {best:.4f})') # ---- test: score, and emit predictions for EVERY test clip (not just unique) ---- from transformers import AutoModelForSeq2SeqLM as M2 model = M2.from_pretrained(os.path.join(args.out_dir, 'best')).to(device) te_u = uniq['test'] hyps = generate(model, tok, [r['sentence'] for r in te_u], device) mt = score([r['gloss'] for r in te_u], hyps) print('\nTEST (unique pairs, n=%d): %s' % (len(te_u), json.dumps(mt, indent=2))) sent2pred = {r['sentence']: h for r, h in zip(te_u, hyps)} per_clip = {} for r in allrows['test']: p = sent2pred.get(r['sentence']) if p is None: continue per_clip[r['name']] = {'pred_gloss': norm_gloss(p), 'pred_raw': p, 'ref_gloss': norm_gloss(r['gloss'])} os.makedirs(os.path.dirname(args.emit_json) or '.', exist_ok=True) with open(args.emit_json, 'w', encoding='utf-8') as f: json.dump(per_clip, f, ensure_ascii=False) print(f'wrote {args.emit_json} ({len(per_clip)} test clips)') with open(os.path.join(args.out_dir, 'metrics.json'), 'w') as f: json.dump({'model': args.model, 'history': hist, 'test': mt, 'test_pairs': len(te_u)}, f, indent=2, ensure_ascii=False) print('\nexamples:') for r, h in list(zip(te_u, hyps))[:6]: print(f" SENT {r['sentence']}\n REF {r['gloss']}\n HYP {h}\n") if __name__ == '__main__': main()