JonLoRA commited on
Commit
9b428e3
·
verified ·
1 Parent(s): 8d554a1

Update handler.py

Browse files
Files changed (1) hide show
  1. handler.py +85 -69
handler.py CHANGED
@@ -1,70 +1,86 @@
1
- from typing import Dict, Any, List
2
- import torch
3
- from transformers import AutoModelForCausalLM, AutoTokenizer
4
-
5
- class EndpointHandler:
6
- def __init__(self, path=""):
7
- """Initialize the model and tokenizer.
8
-
9
- Args:
10
- path (str): Path to the model directory. Defaults to empty string.
11
- """
12
- self.device = "cuda" if torch.cuda.is_available() else "cpu"
13
- self.model = AutoModelForCausalLM.from_pretrained(
14
- path or "merged",
15
- torch_dtype=torch.float16 if torch.cuda.is_available() else torch.float32,
16
- device_map="auto"
17
- )
18
- self.tokenizer = AutoTokenizer.from_pretrained(path or "merged")
19
-
20
- def __call__(self, data: Dict[str, Any]) -> Dict[str, Any]:
21
- """Handle inference requests.
22
-
23
- Args:
24
- data (Dict[str, Any]): The input data containing:
25
- - instruction (str): The instruction for the model
26
- - input (str, optional): Additional input text
27
- - max_new_tokens (int, optional): Maximum number of tokens to generate. Defaults to 512
28
- - temperature (float, optional): Sampling temperature. Defaults to 0.7
29
-
30
- Returns:
31
- Dict[str, Any]: The model's response containing:
32
- - response (str): The generated text
33
- """
34
- # Extract parameters from the request
35
- instruction = data.get("instruction", "")
36
- input_text = data.get("input", "")
37
- max_new_tokens = data.get("max_new_tokens", 512)
38
- temperature = data.get("temperature", 0.7)
39
-
40
- # Create prompt
41
- prompt = f"""Below is an instruction that describes a task. Write a response that appropriately completes the request.
42
-
43
- ### Instruction:
44
- {instruction}"""
45
-
46
- if input_text:
47
- prompt += f"""
48
-
49
- ### Input:
50
- {input_text}"""
51
-
52
- prompt += """
53
-
54
- ### Response:"""
55
-
56
- # Generate response
57
- inputs = self.tokenizer(prompt, return_tensors="pt").to(self.model.device)
58
- outputs = self.model.generate(
59
- **inputs,
60
- max_new_tokens=max_new_tokens,
61
- temperature=temperature,
62
- do_sample=True,
63
- pad_token_id=self.tokenizer.eos_token_id
64
- )
65
-
66
- # Decode and extract response
67
- full_response = self.tokenizer.decode(outputs[0], skip_special_tokens=True)
68
- response = full_response.split("### Response:")[-1].strip()
69
-
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
70
  return {"response": response}
 
1
+ from typing import Dict, Any, List
2
+ import torch
3
+ from transformers import AutoModelForCausalLM, AutoTokenizer
4
+
5
+ class EndpointHandler:
6
+ def __init__(self, path=""):
7
+ """Initialize the model and tokenizer.
8
+
9
+ Args:
10
+ path (str): Path to the model directory. Defaults to empty string.
11
+ """
12
+ self.device = "cuda" if torch.cuda.is_available() else "cpu"
13
+ self.model = AutoModelForCausalLM.from_pretrained(
14
+ path or "merged",
15
+ torch_dtype=torch.float16 if torch.cuda.is_available() else torch.float32,
16
+ device_map="auto"
17
+ )
18
+ self.tokenizer = AutoTokenizer.from_pretrained(path or "merged")
19
+
20
+ def __call__(self, data: Dict[str, Any]) -> Dict[str, Any]:
21
+ """Handle inference requests.
22
+
23
+ Args:
24
+ data (Dict[str, Any]): The input data. Can be in two formats:
25
+ 1. Standard Hugging Face format:
26
+ {
27
+ "inputs": str,
28
+ "parameters": Dict[str, Any]
29
+ }
30
+ 2. Custom format:
31
+ {
32
+ "instruction": str,
33
+ "input": str (optional),
34
+ "max_new_tokens": int (optional),
35
+ "temperature": float (optional)
36
+ }
37
+
38
+ Returns:
39
+ Dict[str, Any]: The model's response containing:
40
+ - response (str): The generated text
41
+ """
42
+ # Handle standard Hugging Face format
43
+ if "inputs" in data:
44
+ instruction = data["inputs"]
45
+ parameters = data.get("parameters", {})
46
+ input_text = ""
47
+ max_new_tokens = parameters.get("max_new_tokens", 512)
48
+ temperature = parameters.get("temperature", 0.7)
49
+ # Handle custom format
50
+ else:
51
+ instruction = data.get("instruction", "")
52
+ input_text = data.get("input", "")
53
+ max_new_tokens = data.get("max_new_tokens", 512)
54
+ temperature = data.get("temperature", 0.7)
55
+
56
+ # Create prompt
57
+ prompt = f"""Below is an instruction that describes a task. Write a response that appropriately completes the request.
58
+
59
+ ### Instruction:
60
+ {instruction}"""
61
+
62
+ if input_text:
63
+ prompt += f"""
64
+
65
+ ### Input:
66
+ {input_text}"""
67
+
68
+ prompt += """
69
+
70
+ ### Response:"""
71
+
72
+ # Generate response
73
+ inputs = self.tokenizer(prompt, return_tensors="pt").to(self.model.device)
74
+ outputs = self.model.generate(
75
+ **inputs,
76
+ max_new_tokens=max_new_tokens,
77
+ temperature=temperature,
78
+ do_sample=True,
79
+ pad_token_id=self.tokenizer.eos_token_id
80
+ )
81
+
82
+ # Decode and extract response
83
+ full_response = self.tokenizer.decode(outputs[0], skip_special_tokens=True)
84
+ response = full_response.split("### Response:")[-1].strip()
85
+
86
  return {"response": response}