"""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()