| |
| """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)]) |
| |
| 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})') |
|
|
| |
| 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() |
|
|