File size: 4,188 Bytes
7e9cfd1
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
115
116
117
"""Run one or more models over a CSV and dump a submission file. Multiple
model dirs just get their softmax probabilities averaged. Checks the
submission format before writing it out.

    python -m src.predict --models outputs/t1_arabert \\
        --csv data/track1/dev.csv --gold data/track1/dev.csv \\
        --out submissions/t1_dev.txt
"""
import argparse
import json
import os

import numpy as np
import torch
import torch.nn.functional as F
from torch.utils.data import DataLoader
from transformers import AutoModelForSequenceClassification, AutoTokenizer

from src.data import ID2LABEL, LABEL2ID, StanceDataset, load_split
from src.scorer import load_gold, score, validate_submission


def model_settings(model_dir, defaults):
    """Read the preprocessing a model was trained with, if recorded."""
    cfg_path = os.path.join(model_dir, "best.json")
    if os.path.isfile(cfg_path):
        with open(cfg_path, encoding="utf-8") as f:
            cfg = json.load(f).get("config", {})
        return {
            "prep_mode": cfg.get("prep_mode", defaults["prep_mode"]),
            "use_description": cfg.get(
                "use_description", defaults["use_description"]
            ),
            "max_len": cfg.get("max_len", defaults["max_len"]),
        }
    return dict(defaults)


@torch.no_grad()
def model_probs(model_dir, csv_path, device, defaults, batch_size=64):
    s = model_settings(model_dir, defaults)
    df = load_split(csv_path, s["prep_mode"], has_labels=False)
    tok = AutoTokenizer.from_pretrained(model_dir, trust_remote_code=True)
    model = AutoModelForSequenceClassification.from_pretrained(
        model_dir, trust_remote_code=True
    ).to(device).eval()
    ds = StanceDataset(
        df, tok, s["max_len"], s["use_description"], has_labels=False
    )
    loader = DataLoader(ds, batch_size=batch_size, shuffle=False)
    probs = []
    for batch in loader:
        batch = {k: v.to(device) for k, v in batch.items()}
        logits = model(**batch).logits.float()
        probs.append(F.softmax(logits, dim=-1).cpu().numpy())
    return np.concatenate(probs, axis=0)


def parse_args():
    ap = argparse.ArgumentParser()
    ap.add_argument("--models", nargs="+", required=True)
    ap.add_argument("--csv", required=True)
    ap.add_argument("--out", required=True)
    ap.add_argument("--gold", default=None)
    ap.add_argument("--max_len", type=int, default=128)
    ap.add_argument("--use_description", action="store_true")
    ap.add_argument("--prep_mode", default="preserve")
    ap.add_argument("--none_bias", type=float, default=0.0)
    ap.add_argument("--llm_probs", default=None,
                    help="npy of LLM class probabilities aligned to --csv")
    ap.add_argument("--enc_weight", type=float, default=1.0,
                    help="weight on encoder probs; LLM gets 1 - enc_weight")
    return ap.parse_args()


def main():
    args = parse_args()
    device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
    defaults = {
        "prep_mode": args.prep_mode,
        "use_description": args.use_description,
        "max_len": args.max_len,
    }
    n_rows = len(load_split(args.csv, "preserve", has_labels=False))

    probs = np.mean(
        [model_probs(m, args.csv, device, defaults) for m in args.models],
        axis=0,
    )
    if args.llm_probs:
        llm = np.load(args.llm_probs)
        if len(llm) != n_rows:
            raise SystemExit(
                f"llm_probs rows {len(llm)} != csv rows {n_rows}"
            )
        probs = args.enc_weight * probs + (1 - args.enc_weight) * llm
    probs[:, LABEL2ID["None"]] += args.none_bias
    preds = [ID2LABEL[i] for i in probs.argmax(axis=1)]

    ok, msg = validate_submission(preds, n_rows)
    print(f"[validate] {msg}")
    if not ok:
        raise SystemExit(1)

    os.makedirs(os.path.dirname(os.path.abspath(args.out)), exist_ok=True)
    with open(args.out, "w", encoding="utf-8") as f:
        f.write("\n".join(preds) + "\n")
    print(f"[write] {len(preds)} predictions -> {args.out}")

    if args.gold:
        print("[score]")
        score(load_gold(args.gold), preds)


if __name__ == "__main__":
    main()