MCPcopy Create free account
hub / github.com/OpenBMB/BMTools / LoraModel

Class LoraModel

bmtools/models/lora_model.py:8–64  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

6from peft import PeftModel
7
8class LoraModel(LLM):
9
10 model_name: str = ""
11 tokenizer: AutoTokenizer = None
12 model: PeftModel = None
13 use_gpu: bool = True
14
15 def __init__(self, base_name_or_path: str, model_name_or_path: str, device: str="cuda", cpu_offloading: bool=False, load_8bit: bool=False) -> None:
16 super().__init__()
17 self.model_name = model_name_or_path
18 self.tokenizer = AutoTokenizer.from_pretrained(base_name_or_path, use_fast=False)
19 model = AutoModelForCausalLM.from_pretrained(
20 base_name_or_path,
21 load_in_8bit=load_8bit,
22 device_map="auto"
23 )
24 self.model = PeftModel.from_pretrained(
25 model,
26 model_name_or_path
27 )
28 if self.tokenizer.pad_token_id == None:
29 self.tokenizer.add_special_tokens({"bos_token": "<s>", "eos_token": "</s>", "pad_token": "<pad>"})
30 self.model.resize_token_embeddings(len(self.tokenizer))
31 self.use_gpu = (True if device == "cuda" else False)
32 if (device == "cuda" and not cpu_offloading) or device == "mps":
33 self.model.to(device)
34
35 @property
36 def _llm_type(self) -> str:
37 return self.model_name
38
39 def _call(self, prompt: str, stop: Optional[List[str]] = None) -> str:
40 inputs = self.tokenizer(
41 prompt,
42 padding=True,
43 max_length=self.tokenizer.model_max_length,
44 truncation=True,
45 return_tensors="pt"
46 )
47 inputs_len = inputs["input_ids"].shape[1]
48 generated_outputs = self.model.generate(
49 input_ids=(inputs["input_ids"].cuda() if self.use_gpu else inputs["input_ids"]),
50 attention_mask=(inputs["attention_mask"].cuda() if self.use_gpu else inputs["attention_mask"]),
51 max_new_tokens=512,
52 eos_token_id=self.tokenizer.eos_token_id,
53 bos_token_id=self.tokenizer.bos_token_id,
54 pad_token_id=self.tokenizer.pad_token_id,
55 )
56 decoded_output = self.tokenizer.batch_decode(
57 generated_outputs[..., inputs_len:], skip_special_tokens=True)
58 output = decoded_output[0]
59 return output
60
61 @property
62 def _identifying_params(self) -> Mapping[str, Any]:
63 """Get the identifying parameters."""
64 return {"model_name": self.model_name}
65

Callers 1

lora_model.pyFile · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected