(self,
path: str,
model_kwargs: dict,
peft_path: Optional[str] = None)
| 194 | model_kwargs['torch_dtype'] = torch_dtype |
| 195 | |
| 196 | def _load_model(self, |
| 197 | path: str, |
| 198 | model_kwargs: dict, |
| 199 | peft_path: Optional[str] = None): |
| 200 | from transformers import AutoModel, AutoModelForCausalLM |
| 201 | |
| 202 | self._set_model_kwargs_torch_dtype(model_kwargs) |
| 203 | try: |
| 204 | self.model = AutoModelForCausalLM.from_pretrained( |
| 205 | path, **model_kwargs) |
| 206 | except ValueError: |
| 207 | self.model = AutoModel.from_pretrained(path, **model_kwargs) |
| 208 | |
| 209 | if peft_path is not None: |
| 210 | from peft import PeftModel |
| 211 | self.model = PeftModel.from_pretrained(self.model, |
| 212 | peft_path, |
| 213 | is_trainable=False) |
| 214 | self.model.eval() |
| 215 | self.model.generation_config.do_sample = False |
| 216 | |
| 217 | # A patch for llama when batch_padding = True |
| 218 | if 'decapoda-research/llama' in path: |
| 219 | self.model.config.bos_token_id = 1 |
| 220 | self.model.config.eos_token_id = 2 |
| 221 | self.model.config.pad_token_id = self.tokenizer.pad_token_id |
| 222 | |
| 223 | def generate(self, |
| 224 | inputs: List[str], |
no test coverage detected