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

Class Baichuan2ChatEvaluator

code/inference/evaluators/baichuan2chat.py:10–57  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

8
9
10class Baichuan2ChatEvaluator(Evaluator):
11
12 def __init__(self, pretrained_model_name_or_path, cache_dir=None, do_sample=False):
13 super(Baichuan2ChatEvaluator, self).__init__()
14
15 self.tokenizer = AutoTokenizer.from_pretrained(
16 pretrained_model_name_or_path=pretrained_model_name_or_path,
17 cache_dir=cache_dir,
18 use_fast=False,
19 trust_remote_code=True
20 )
21 self.model = AutoModelForCausalLM.from_pretrained(
22 pretrained_model_name_or_path=pretrained_model_name_or_path,
23 cache_dir=cache_dir,
24 device_map='auto',
25 low_cpu_mem_usage=True,
26 torch_dtype=torch.float16,
27 trust_remote_code=True
28 )
29 self.model.generation_config = GenerationConfig.from_pretrained(
30 pretrained_model_name=pretrained_model_name_or_path,
31 cache_dir=cache_dir
32 )
33 self.model.generation_config.do_sample = do_sample
34 self.model = self.model.eval()
35 print(f'Memory footprint: {self.model.get_memory_footprint() / 1e6:.2f} MB')
36
37 def format_prompt(self, prompt):
38 return [
39 {
40 'role': 'user',
41 'content': prompt
42 }
43 ]
44
45 @retry(wait=wait_random_exponential(min=1, max=60), stop=stop_after_attempt(6))
46 def generate_text(self, prompt):
47 prompt = self.format_prompt(prompt)
48 response = self.model.chat(
49 self.tokenizer,
50 prompt
51 )
52
53 return response.strip()
54
55 @retry(wait=wait_random_exponential(min=1, max=60), stop=stop_after_attempt(6))
56 def count_tokens(self, prompt):
57 return len(self.tokenizer(prompt)['input_ids'])

Callers 1

eval.pyFile · 0.90

Calls

no outgoing calls

Tested by

no test coverage detected