t2m-gpt-vsl-code / predict_gloss_trivis.py
Tri1's picture
T2M-GPT VSL adaptation: Python sources only (82 files, no checkpoints or data)
8e5456b verified
Raw
History Blame Contribute Delete
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()