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

Class LlamaModel

bmtools/models/llama_model.py:7–57  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

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

Callers 1

llama_model.pyFile · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected