File size: 2,515 Bytes
6f63ca6
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
#!/usr/bin/env python3
"""Measure normalized exact match and token F1 on labeled JSONL examples."""

from __future__ import annotations

import argparse
import json
import re
import string
from collections import Counter
from pathlib import Path

import torch
from transformers import AutoModelForSeq2SeqLM, AutoTokenizer


def normalize(text: str) -> str:
    text = text.lower().translate(str.maketrans("", "", string.punctuation))
    return " ".join(re.sub(r"\b(a|an|the)\b", " ", text).split())


def token_f1(prediction: str, target: str) -> float:
    predicted, expected = normalize(prediction).split(), normalize(target).split()
    if not predicted or not expected:
        return float(predicted == expected)
    overlap = sum((Counter(predicted) & Counter(expected)).values())
    if overlap == 0:
        return 0.0
    precision, recall = overlap / len(predicted), overlap / len(expected)
    return 2 * precision * recall / (precision + recall)


def main() -> None:
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument("data", type=Path)
    parser.add_argument("--model", default="ShinpacheShimura/t5-smaller")
    parser.add_argument("--subfolder", default="optimized-flan-t5-small")
    parser.add_argument("--max-new-tokens", type=int, default=64)
    args = parser.parse_args()

    rows = [json.loads(line) for line in args.data.read_text().splitlines() if line.strip()]
    if not rows or any("input" not in row or "target" not in row for row in rows):
        raise SystemExit('Use non-empty JSONL rows with "input" and "target" fields.')

    common = {"subfolder": args.subfolder} if args.subfolder else {}
    tokenizer = AutoTokenizer.from_pretrained(args.model, **common)
    model = AutoModelForSeq2SeqLM.from_pretrained(args.model, device_map="auto", **common)
    exact, total_f1 = 0, 0.0
    for row in rows:
        inputs = tokenizer(row["input"], return_tensors="pt").to(model.device)
        with torch.inference_mode():
            ids = model.generate(**inputs, max_new_tokens=args.max_new_tokens, do_sample=False)
        prediction = tokenizer.decode(ids[0], skip_special_tokens=True)
        exact += normalize(prediction) == normalize(row["target"])
        total_f1 += token_f1(prediction, row["target"])
        print(json.dumps({**row, "prediction": prediction}))

    count = len(rows)
    print(json.dumps({"examples": count, "exact_match": exact / count, "token_f1": total_f1 / count}, indent=2))


if __name__ == "__main__":
    main()