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

Class YiChatEvaluator

code/inference/evaluators/yichat.py:9–54  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

7
8
9class YiChatEvaluator(Evaluator):
10
11 def __init__(self, pretrained_model_name_or_path, cache_dir=None, do_sample=False):
12 super(YiChatEvaluator, 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 use_fast=False,
18 trust_remote_code=True
19 )
20 self.model = AutoModelForCausalLM.from_pretrained(
21 pretrained_model_name_or_path=pretrained_model_name_or_path,
22 cache_dir=cache_dir,
23 device_map='auto',
24 low_cpu_mem_usage=True,
25 torch_dtype=torch.float16,
26 trust_remote_code=True
27 )
28 self.model = self.model.eval()
29 self.do_sample = do_sample
30 print(f'Memory footprint: {self.model.get_memory_footprint() / 1e6:.2f} MB')
31
32 def format_prompt(self, prompt):
33 return [
34 {
35 'role': 'user',
36 'content': prompt
37 }
38 ]
39
40 @retry(wait=wait_random_exponential(min=1, max=60), stop=stop_after_attempt(6))
41 def generate_text(self, prompt):
42 prompt = self.format_prompt(prompt)
43 inputs = self.tokenizer.apply_chat_template(conversation=prompt, tokenize=True, add_generation_prompt=True, return_tensors='pt').to(self.model.device)
44 outputs = self.model.generate(
45 inputs,
46 do_sample=self.do_sample
47 )
48 response = self.tokenizer.decode(outputs[0][inputs.shape[1]:], skip_special_tokens=True)
49
50 return response.strip()
51
52 @retry(wait=wait_random_exponential(min=1, max=60), stop=stop_after_attempt(6))
53 def count_tokens(self, prompt):
54 return len(self.tokenizer(prompt)['input_ids'])

Callers 1

eval.pyFile · 0.90

Calls

no outgoing calls

Tested by

no test coverage detected