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

Class DISCMedLLMEvaluator

code/inference/evaluators/discmedllm.py:10–54  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

8
9
10class DISCMedLLMEvaluator(Evaluator):
11
12 def __init__(self, pretrained_model_name_or_path, cache_dir=None, do_sample=False):
13 super(DISCMedLLMEvaluator, 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(pretrained_model_name=pretrained_model_name_or_path)
30 self.model.generation_config.do_sample = do_sample
31 self.model = self.model.eval()
32 print(f'Memory footprint: {self.model.get_memory_footprint() / 1e6:.2f} MB')
33
34 def format_prompt(self, prompt):
35 return [
36 {
37 'role': 'user',
38 'content': prompt
39 }
40 ]
41
42 @retry(wait=wait_random_exponential(min=1, max=60), stop=stop_after_attempt(6))
43 def generate_text(self, prompt):
44 prompt = self.format_prompt(prompt)
45 response = self.model.chat(
46 self.tokenizer,
47 prompt
48 )
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