MCPcopy Create free account
hub / github.com/WeixiangYAN/ClinicalLab / WiNGPT2Evaluator

Class WiNGPT2Evaluator

code/inference/evaluators/wingpt2.py:9–51  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

7
8
9class WiNGPT2Evaluator(Evaluator):
10
11 def __init__(self, pretrained_model_name_or_path, cache_dir=None, do_sample=False, max_length=4096):
12 super(WiNGPT2Evaluator, self).__init__()
13
14 self.tokenizer = AutoTokenizer.from_pretrained(
15 pretrained_model_name_or_path=pretrained_model_name_or_path,
16 cache_dir=cache_dir,
17 trust_remote_code=True
18 )
19 self.model = AutoModelForCausalLM.from_pretrained(
20 pretrained_model_name_or_path=pretrained_model_name_or_path,
21 cache_dir=cache_dir,
22 device_map='auto',
23 low_cpu_mem_usage=True,
24 torch_dtype=torch.float16,
25 trust_remote_code=True
26 )
27 self.model = self.model.eval()
28 self.do_sample = do_sample
29 self.max_length = max_length
30 print(f'Memory footprint: {self.model.get_memory_footprint() / 1e6:.2f} MB')
31
32 def format_prompt(self, prompt):
33 return f'User: {prompt}<|endoftext|>\n Assistant: '
34
35 @retry(wait=wait_random_exponential(min=1, max=60), stop=stop_after_attempt(6))
36 def generate_text(self, prompt):
37 prompt = self.format_prompt(prompt)
38 inputs = self.tokenizer.encode(prompt, return_tensors='pt').to(self.model.device)
39 outputs = self.model.generate(
40 inputs,
41 do_sample=self.do_sample,
42 max_length=self.max_length,
43 max_new_tokens=None
44 ).to('cpu')
45 response = self.tokenizer.decode(outputs[0])
46
47 return response.split('Assistant: ')[-1].strip('<|endoftext|>\n'.strip())
48
49 @retry(wait=wait_random_exponential(min=1, max=60), stop=stop_after_attempt(6))
50 def count_tokens(self, prompt):
51 return len(self.tokenizer(prompt)['input_ids'])

Callers 1

eval.pyFile · 0.90

Calls

no outgoing calls

Tested by

no test coverage detected