MCPcopy Create free account
hub / github.com/codefuse-ai/codefuse-devops-eval / BaiChuanModel

Class BaiChuanModel

src/models/baichuan_model.py:16–91  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

14
15
16class BaiChuanModel(ToolModel):
17 def __init__(self, model_path: str, peft_path: str = None, template: str = "default", trust_remote_code=True, tensor_parallel_size=1, gpu_memory_utilization=0.25):
18 self.model_path = model_path
19 self.peft_path = peft_path
20 self.template = template
21 self.trust_remote_code = trust_remote_code
22 self.tensor_parallel_size = tensor_parallel_size
23 self.gpu_memory_utilization = gpu_memory_utilization
24 self.generation_config = GenerationConfig.from_pretrained(model_path)
25 self.load_model(self.model_path, self.peft_path, self.trust_remote_code, self.tensor_parallel_size, self.gpu_memory_utilization)
26
27 def generate(
28 self, prompts: str,
29 template: str = None,
30 generate_configs: GenerateConfigs =None,
31 history: list = None,
32 ) -> list:
33 '''产出对应结果'''
34 template = self.template if template is None else template
35
36 params = self.generate_params(generate_configs)
37
38 if template == "default":
39 inputs = self.tokenizer(prompts, return_tensors="pt")
40 inputs["input_ids"] = inputs["input_ids"].cuda()
41
42 inputs.update(params)
43 output = self.model.generate(**inputs)
44 predict = self.tokenizer.decode(output[0].tolist())[len(prompts):]
45 predict = predict.replace("<|endoftext|>", "").replace("</s>", "")
46 return predict
47 elif template != "default":
48 messages = [{"role": "user" if idx==0 else "assistant", "content": ii} for i in history for idx, ii in enumerate(i)]
49 messages.append({"role": "user", "content": prompts})
50 output = self.model.chat(self.tokenizer, messages=messages, generation_config=self.generation_config)
51 return output
52
53 def generate_params(
54 self, generate_configs: GenerateConfigs,
55 ):
56 '''generate param'''
57 kargs = generate_configs.dict()
58 params = {
59 "max_new_tokens": kargs.get("max_new_tokens", 128),
60 "top_k": kargs.get("top_k", 50),
61 "top_p": kargs.get("top_p", 0.95),
62 "temperature": kargs.get("temperature", 1.0),
63 }
64 self.generation_config.max_new_tokens = kargs.get("max_new_tokens", 128)
65 self.generation_config.top_k = kargs.get("top_k", 50)
66 self.generation_config.top_p = kargs.get("top_p", 0.95)
67 self.generation_config.temperature = kargs.get("temperature", 1.0)
68
69 # params = {
70 # "n": 1,
71 # "max_tokens": kargs.get("max_new_tokens", 128),
72 # "best_of": kargs.get("beam_bums", 1),
73 # "top_k": kargs.get("top_k", 50),

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected