usagent100/testing6000v2
012
1from typing import Any, Dict2import torch3import transformers4from transformers import AutoModelForCausalLM, AutoTokenizer5 6dtype = torch.bfloat16 if torch.cuda.get_device_capability()[0] == 8 else torch.float167 8class EndpointHandler:9 def __init__(self, path=""):10 tokenizer = AutoTokenizer.from_pretrained(path, trust_remote_code=True)11 model = AutoModelForCausalLM.from_pretrained(12 path,13 return_dict=True,14 device_map="auto",15 load_in_8bit=True,16 torch_dtype=dtype,17 trust_remote_code=True,18 )19 20 self.generation_config = model.generation_config21 self.generation_config.max_new_tokens = 100022 self.generation_config.temperature = 0.7 # Changed from 0 to 0.723 self.generation_config.num_return_sequences = 124 self.generation_config.pad_token_id = tokenizer.eos_token_id25 self.generation_config.eos_token_id = tokenizer.eos_token_id26 27 self.pipeline = transformers.pipeline(28 "text-generation", model=model, tokenizer=tokenizer29 )30 31 def __call__(self, data: Dict[str, Any]) -> Dict[str, Any]:32 prompt = data.pop("inputs", data)33 result = self.pipeline(34 prompt,35 max_length=1000, # Added this line to set max_length36 temperature=0.7, # Added this line to set temperature37 top_p=0.9, # Added this line to set top_p38 num_return_sequences=1, # Added this line to set num_return_sequences39 pad_token_id=self.generation_config.pad_token_id,40 eos_token_id=self.generation_config.eos_token_id,41 return_full_text=True # Added this line to return full text42 )43 return {"generated_text": result}44 45 