import json import re from pathlib import Path import torch from transformers import ( AutoModelForCausalLM, AutoTokenizer, pipeline, ) import transformers.modeling_utils as modeling_utils _original_safe_open = modeling_utils.safe_open class SafeOpenWithMetadata: def __init__(self, *args, **kwargs): self._inner = _original_safe_open(*args, **kwargs) def __enter__(self): self._inner.__enter__() return self def __exit__(self, exc_type, exc_val, exc_tb): return self._inner.__exit__(exc_type, exc_val, exc_tb) def metadata(self): metadata = self._inner.metadata() return metadata if metadata is not None else {"format": "pt"} def keys(self): return self._inner.keys() def get_tensor(self, name): return self._inner.get_tensor(name) modeling_utils.safe_open = SafeOpenWithMetadata import torch.nn.functional as F if not hasattr(F, "rms_norm"): def rms_norm(input, normalized_shape, weight=None, eps=None): if eps is None: eps = torch.finfo(input.dtype).eps variance = input.pow(2).mean(dim=-1, keepdim=True) output = input * torch.rsqrt(variance + eps) if weight is not None: output = output * weight return output F.rms_norm = rms_norm # ========================= # CONFIG # ========================= BASE_MODEL = "NurErtug/pit-finance-grpo-merged" TOKENIZER_NAME = "Diamegs/PIT-4B-202212" LORA_ADAPTER = None NLI_MODEL = "MoritzLaurer/DeBERTa-v3-base-mnli-fever-anli" RESULTS_FILE = "verifier_logs.jsonl" REWARD_THRESHOLD = 1.4 # ========================= # HELPERS # ========================= def append_jsonl(path, item): with open(path, "a") as f: f.write(json.dumps(item) + "\n") def extract_numbers(text): if text is None: return [] pattern = r"\$?\d+(?:\.\d+)?\s?(?:million|billion|%)?" return re.findall(pattern, text) def numeric_verifier(evidence, answer): evidence_numbers = extract_numbers(evidence) answer_numbers = extract_numbers(answer) unsupported = [ x for x in answer_numbers if x not in evidence_numbers ] return { "evidence_numbers": evidence_numbers, "answer_numbers": answer_numbers, "unsupported_numbers": unsupported, "numeric_pass": len(unsupported) == 0, } def is_garbage_generation(text): if text is None: return True text = str(text) if len(text.strip()) < 3: return True weird_ratio = text.count("!") / max(len(text), 1) if weird_ratio > 0.3: return True garbage_patterns = [ r"\\boxed", r"\\ding", r"\\text", r"\\end", r"\\null", r"\\stop", ] for p in garbage_patterns: if re.search(p, text): return True return False def is_abstention(answer): a = str(answer).upper() return ( "NOT ENOUGH INFORMATION" in a or a.strip() == "NONE." or a.strip() == "NONE" ) def compute_reward(numeric_result, nli_result): reward = 0.0 if numeric_result["numeric_pass"]: reward += 0.8 if nli_result["nli_pass"]: reward += 1.0 return reward # ========================= # LOAD GENERATOR # ========================= def load_generator(): print(f"Loading PIT base model: {BASE_MODEL}") tokenizer = AutoTokenizer.from_pretrained( TOKENIZER_NAME, trust_remote_code=True, use_fast=False, ) if tokenizer.pad_token is None: tokenizer.pad_token = tokenizer.eos_token dtype = torch.bfloat16 if torch.cuda.is_available() else torch.float32 model = AutoModelForCausalLM.from_pretrained( BASE_MODEL, device_map="auto", torch_dtype=dtype, trust_remote_code=True, ) model.eval() return tokenizer, model # ========================= # LOAD NLI # ========================= def load_nli(): print(f"Loading NLI verifier: {NLI_MODEL}") device = 0 if torch.cuda.is_available() else -1 return pipeline( "text-classification", model=NLI_MODEL, device=device, return_all_scores=True, ) # ========================= # GENERATION # ========================= def generate_answer( tokenizer, model, context, question, max_new_tokens=60, ): prompt = f""" Context: {context} Question: {question} Answer using only the context in one short sentence. Do not infer or compute new numbers. Use only numbers explicitly written in the context. If the context contains no relevant evidence, write exactly: NOT ENOUGH INFORMATION. Answer: """ inputs = tokenizer( prompt, return_tensors="pt", truncation=True, max_length=2048, ).to(model.device) with torch.inference_mode(): output = model.generate( **inputs, max_new_tokens=max_new_tokens, do_sample=False, repetition_penalty=1.2, pad_token_id=tokenizer.eos_token_id, ) generated = output[0][inputs["input_ids"].shape[-1]:] answer = tokenizer.decode( generated, skip_special_tokens=True, ).strip() answer = answer.split("\\boxed")[0] answer = answer.split("\\ding")[0] answer = answer.split("\\text")[0] answer = answer.split("\\end")[0] answer = answer.strip() return answer # ========================= # NLI VERIFICATION # ========================= def verify_nli(nli_pipe, evidence, answer): pair = f"premise: {evidence} hypothesis: {answer}" outputs = nli_pipe(pair)[0] scores = { x["label"].lower(): float(x["score"]) for x in outputs } entailment = scores.get("entailment", 0.0) contradiction = scores.get("contradiction", 0.0) neutral = scores.get("neutral", 0.0) verdict = "entailed" if contradiction > entailment: verdict = "contradiction" elif neutral > entailment: verdict = "neutral" return { "scores": scores, "entailment": entailment, "contradiction": contradiction, "neutral": neutral, "verdict": verdict, "nli_pass": verdict == "entailed", } # ========================= # WRAPPER # ========================= class VerifierWrapper: def __init__(self): self.generator_tokenizer, self.generator_model = load_generator() self.nli_pipe = load_nli() self.log_file = RESULTS_FILE def verified_answer( self, evidence_text, question_text, log=True, ): answer_text = generate_answer( self.generator_tokenizer, self.generator_model, evidence_text, question_text, ) # ------------------------- # garbage rejection # ------------------------- if is_garbage_generation(answer_text): result = { "context": evidence_text, "question": question_text, "evidence": evidence_text, "initial_answer": answer_text, "numeric_result": None, "nli_result": None, "reward": 0.0, "rejected": True, "reject_reason": "garbage_generation", "final_answer": "NOT ENOUGH INFORMATION.", } if log: append_jsonl(self.log_file, result) return result # ------------------------- # abstention rejection # ------------------------- if ( is_abstention(answer_text) and len(extract_numbers(evidence_text)) > 0 ): result = { "context": evidence_text, "question": question_text, "evidence": evidence_text, "initial_answer": answer_text, "numeric_result": None, "nli_result": None, "reward": 0.0, "rejected": True, "reject_reason": "bad_abstention_when_evidence_exists", "final_answer": "NOT ENOUGH INFORMATION.", } if log: append_jsonl(self.log_file, result) return result # ------------------------- # numeric verifier # ------------------------- numeric_result = numeric_verifier( evidence_text, answer_text, ) # ------------------------- # nli verifier # ------------------------- nli_result = verify_nli( self.nli_pipe, evidence_text, answer_text, ) # ------------------------- # reward # ------------------------- reward = compute_reward( numeric_result, nli_result, ) rejected = ( reward < REWARD_THRESHOLD or not numeric_result["numeric_pass"] or not nli_result["nli_pass"] ) final_answer = ( "NOT ENOUGH INFORMATION." if rejected else answer_text ) result = { "context": evidence_text, "question": question_text, "evidence": evidence_text, "initial_answer": answer_text, "numeric_result": numeric_result, "nli_result": nli_result, "reward": reward, "rejected": rejected, "reject_reason": ( "verifier_failed" if rejected else None ), "final_answer": final_answer, } if log: append_jsonl(self.log_file, result) return result # ========================= # TEST # ========================= if __name__ == "__main__": print("Run example.py for a demo.")