stanceeval2026 / code /src /llm_infer.py
zaher-m's picture
Add files using upload-large-folder tool
7e9cfd1 verified
Raw
History Blame Contribute Delete
3.94 kB
"""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()