File size: 5,071 Bytes
8e5456b
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
#!/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()