Views
No views yet
1import torch
2from transformers import AutoTokenizer, AutoModelForCausalLM
3
4
5class FactRecallClassifier:
6 def __init__(self, model_name="naskovai/fact-recall-classifier"):
7 self.model = AutoModelForCausalLM.from_pretrained(
8 model_name, torch_dtype=torch.float32, trust_remote_code=True
9 )
10 self.tokenizer = AutoTokenizer.from_pretrained(
11 model_name, trust_remote_code=True
12 )
13 self.yes_token_id = self.tokenizer("yes", add_special_tokens=False).input_ids[0]
14 self.no_token_id = self.tokenizer("no", add_special_tokens=False).input_ids[0]
15
16 @torch.no_grad()
17 def predict(self, question, expected_fact, generated_answer):
18 instruction = (
19 "Check if the generated answer contains or implies the expected fact."
20 )
21 query = f"Original Question: {question}\nExpected Fact: {expected_fact}"
22 document = f"Generated Answer: {generated_answer}"
23
24 prefix = '<|im_start|>system\nDetermine if the generated answer contains or recalls the expected fact. Note that the answer can only be "yes" or "no".<|im_end|>\n<|im_start|>user\n'
25 suffix = "<|im_end|>\n<|im_start|>assistant\n<think>\n\n</think>\n\n"
26 user_content = (
27 f"<Instruct>: {instruction}\n<Query>: {query}\n<Document>: {document}"
28 )
29 prompt = prefix + user_content + suffix
30
31 inputs = self.tokenizer(
32 prompt, return_tensors="pt", truncation=True, max_length=512
33 )
34 outputs = self.model(**inputs)
35
36 next_token_logits = outputs.logits[0, -1, :]
37 yes_logit = next_token_logits[self.yes_token_id]
38 no_logit = next_token_logits[self.no_token_id]
39
40 return yes_logit > no_logit