Download predict_gloss_trivis.py from Tri1/t2m-gpt-vsl-code: direct link, hf CLI and curl.
- Browser
- Download file 5.07 kB
-
https://huggingface.co/Tri1/t2m-gpt-vsl-code/resolve/main/predict_gloss_trivis.py
- Command line
-
hf download hf://Tri1/t2m-gpt-vsl-code/predict_gloss_trivis.py
-
curl -L -o predict_gloss_trivis.py https://huggingface.co/Tri1/t2m-gpt-vsl-code/resolve/main/predict_gloss_trivis.py
5.07 kB
| #!/usr/bin/env python3 | |
| """Run the fine-tuned text->gloss model over EVERY Full_TriVis clip. | |
| `train_text2gloss.py` only emitted the test split (`pred_test.json`); the composed | |
| word-level evaluation wants predictions for every clip, plus the per-split scores in | |
| one place. Model: `output_vsl/text2gloss/best` (BARTpho-syllable fine-tuned on the | |
| `Sentence` -> `Sign_sentence` pairs of the same split). | |
| Generation is done once per UNIQUE sentence and fanned back out to clips, since the | |
| CSV repeats each sentence across signers/sessions (24k rows, ~12k unique pairs). | |
| Output JSON: {clip_name: {split, sentence, pred_gloss, pred_signs, ref_gloss, ref_signs}} | |
| where `pred_gloss` keeps the `|` sign separators and `pred_signs` is the split list the | |
| pose models consume. | |
| NOTE on honesty: the model was trained on the train split, so train/val gloss quality | |
| is optimistic. Only the `test` rows are leak-free -- report those. | |
| """ | |
| import argparse | |
| import csv | |
| import json | |
| import os | |
| import numpy as np | |
| import torch | |
| from train_text2gloss import norm_gloss, score, signs | |
| REPO = os.path.join(os.path.dirname(os.path.abspath(__file__)), '..') | |
| def main(): | |
| ap = argparse.ArgumentParser() | |
| ap.add_argument('--model', default='output_vsl/text2gloss/best') | |
| ap.add_argument('--csv', default=os.path.join(REPO, 'Full_TriVis', 'split_lab_front.csv')) | |
| ap.add_argument('--splits', default='train,val,test') | |
| ap.add_argument('--num-beams', type=int, default=4) | |
| ap.add_argument('--batch-size', type=int, default=64) | |
| ap.add_argument('--device', default='cuda') | |
| ap.add_argument('--out-json', default='output_vsl/text2gloss/pred_full_trivis.json') | |
| ap.add_argument('--metrics-json', default='output_vsl/text2gloss/metrics_full_trivis.json') | |
| args = ap.parse_args() | |
| want = set(args.splits.split(',')) | |
| rows = [] | |
| with open(args.csv, newline='', encoding='utf-8') as f: | |
| for r in csv.DictReader(f): | |
| if r['split'] not in want: | |
| continue | |
| rows.append({'name': os.path.splitext(os.path.basename(r['npz_path']))[0], | |
| 'split': r['split'], | |
| 'sentence': r['Sentence'].strip(), | |
| 'gloss': r['Sign_sentence'].strip()}) | |
| print(f'{len(rows)} clips over splits {sorted({r["split"] for r in rows})}') | |
| uniq = sorted({r['sentence'] for r in rows}) | |
| print(f'{len(uniq)} unique sentences to translate') | |
| from transformers import AutoModelForSeq2SeqLM, AutoTokenizer | |
| device = torch.device(args.device) | |
| tok = AutoTokenizer.from_pretrained(args.model) | |
| model = AutoModelForSeq2SeqLM.from_pretrained(args.model).eval().to(device) | |
| preds = {} | |
| with torch.no_grad(): | |
| for i in range(0, len(uniq), args.batch_size): | |
| chunk = uniq[i:i + args.batch_size] | |
| x = tok(chunk, padding=True, truncation=True, max_length=64, | |
| return_tensors='pt').to(device) | |
| g = model.generate(**x, num_beams=args.num_beams, max_length=64, | |
| early_stopping=True) | |
| for s, h in zip(chunk, tok.batch_decode(g, skip_special_tokens=True)): | |
| preds[s] = h | |
| print(f' {min(i + args.batch_size, len(uniq))}/{len(uniq)}', end='\r', flush=True) | |
| print() | |
| out = {} | |
| for r in rows: | |
| h = preds[r['sentence']] | |
| out[r['name']] = {'split': r['split'], 'sentence': r['sentence'], | |
| 'pred_gloss': h, 'pred_signs': signs(h), | |
| 'ref_gloss': r['gloss'], 'ref_signs': signs(r['gloss'])} | |
| os.makedirs(os.path.dirname(args.out_json) or '.', exist_ok=True) | |
| with open(args.out_json, 'w', encoding='utf-8') as f: | |
| json.dump(out, f, ensure_ascii=False) | |
| print(f'wrote {args.out_json} ({len(out)} clips)') | |
| # ---- per-split scores, on unique (sentence, gloss) pairs ---- | |
| metrics = {} | |
| for sp in sorted({r['split'] for r in rows}): | |
| pairs = {(r['sentence'], r['gloss']) for r in rows if r['split'] == sp} | |
| refs = [g for _, g in pairs] | |
| hyps = [preds[s] for s, _ in pairs] | |
| m = score(refs, hyps) | |
| m['n_pairs'] = len(pairs) | |
| m['n_clips'] = sum(r['split'] == sp for r in rows) | |
| m['mean_ref_signs'] = float(np.mean([len(signs(g)) for g in refs])) | |
| m['mean_hyp_signs'] = float(np.mean([len(signs(h)) for h in hyps])) | |
| m['seen_in_training'] = (sp == 'train') | |
| metrics[sp] = m | |
| print(f"[{sp}] WER {m['wer']:.4f} EM {m['exact_match']:.4f} F1 {m['f1']:.4f} " | |
| f"signs ref {m['mean_ref_signs']:.2f} / hyp {m['mean_hyp_signs']:.2f} " | |
| f"({m['n_pairs']} pairs, {m['n_clips']} clips)" | |
| + (' [SEEN IN TRAINING -- optimistic]' if m['seen_in_training'] else '')) | |
| with open(args.metrics_json, 'w', encoding='utf-8') as f: | |
| json.dump({'model': args.model, 'num_beams': args.num_beams, | |
| 'per_split': metrics}, f, indent=2, ensure_ascii=False) | |
| print(f'wrote {args.metrics_json}') | |
| if __name__ == '__main__': | |
| main() | |