File size: 3,943 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
"""Score labels with the LoRA-finetuned LM instead of decoding: build the same
prompt as training, get the log-prob of each label as a continuation, and
softmax across them. Columns come out Against/Favor/None to match the
encoders and scorer.

    python -m src.llm_infer --csv data/track1/dev.csv \\
        --base_model ALLaM-AI/ALLaM-7B-Instruct-preview \\
        --adapter outputs/allam_t1 --gold data/track1/dev.csv \\
        --out_probs outputs/llm/allam_dev.npy
"""
import argparse
import os

import numpy as np
import torch
from peft import PeftModel
from transformers import AutoModelForCausalLM, AutoTokenizer

from src.data import load_split
from src.llm_finetune import SYSTEM, user_text
from src.scorer import load_gold, score

LABELS = ["Against", "Favor", "None"]


def build_prompt(tok, target, tweet):
    msgs = [
        {"role": "system", "content": SYSTEM},
        {"role": "user", "content": user_text(target, tweet)},
    ]
    return tok.apply_chat_template(
        msgs, tokenize=False, add_generation_prompt=True
    )


@torch.no_grad()
def score_labels(model, tok, prompts, device, batch_size=16):
    label_ids = [tok(" " + lb, add_special_tokens=False)["input_ids"]
                 for lb in LABELS]
    out = np.zeros((len(prompts), 3))
    for start in range(0, len(prompts), batch_size):
        chunk = prompts[start:start + batch_size]
        p_ids = [tok(p, add_special_tokens=False)["input_ids"] for p in chunk]
        seqs, meta = [], []
        for ei, pid in enumerate(p_ids):
            for li, cid in enumerate(label_ids):
                seqs.append(pid + cid)
                meta.append((ei, li, len(cid), len(pid)))
        width = max(len(s) for s in seqs)
        ids = torch.full((len(seqs), width), tok.pad_token_id,
                         dtype=torch.long)
        att = torch.zeros((len(seqs), width), dtype=torch.long)
        for k, s in enumerate(seqs):
            ids[k, :len(s)] = torch.tensor(s)
            att[k, :len(s)] = 1
        logp = torch.log_softmax(
            model(input_ids=ids.to(device),
                  attention_mask=att.to(device)).logits.float(), dim=-1
        )
        scores = np.full((len(chunk), 3), -1e9)
        for k, (ei, li, clen, plen) in enumerate(meta):
            tot = 0.0
            for t in range(clen):
                tot += logp[k, plen + t - 1, ids[k, plen + t]].item()
            scores[ei, li] = tot
        for ei in range(len(chunk)):
            e = np.exp(scores[ei] - scores[ei].max())
            out[start + ei] = e / e.sum()
    return out


def main():
    ap = argparse.ArgumentParser()
    ap.add_argument("--csv", required=True)
    ap.add_argument("--base_model", required=True)
    ap.add_argument("--adapter", default=None)
    ap.add_argument("--gold", default=None)
    ap.add_argument("--out_probs", required=True)
    ap.add_argument("--batch_size", type=int, default=16)
    args = ap.parse_args()

    device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
    tok = AutoTokenizer.from_pretrained(args.adapter or args.base_model)
    if tok.pad_token is None:
        tok.pad_token = tok.eos_token
    model = AutoModelForCausalLM.from_pretrained(
        args.base_model, torch_dtype=torch.bfloat16
    ).to(device).eval()
    if args.adapter:
        model = PeftModel.from_pretrained(model, args.adapter).eval()

    df = load_split(args.csv, "preserve", has_labels=False)
    prompts = [build_prompt(tok, r["target"], r["text"])
               for r in df.to_dict("records")]
    probs = score_labels(model, tok, prompts, device, args.batch_size)

    os.makedirs(os.path.dirname(os.path.abspath(args.out_probs)),
                exist_ok=True)
    np.save(args.out_probs, probs)
    print(f"[write] probs -> {args.out_probs}")
    if args.gold:
        preds = [LABELS[i] for i in probs.argmax(1)]
        score(load_gold(args.gold), preds)


if __name__ == "__main__":
    main()