HowWebWorks commited on
Commit
5c35d44
·
1 Parent(s): ca9c9ce

add handler.py

Browse files
Files changed (2) hide show
  1. handler.py +16 -0
  2. load_model.py +0 -6
handler.py ADDED
@@ -0,0 +1,16 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from transformers import AutoModelForCausalLM, AutoTokenizer
2
+ from typing import Dict
3
+
4
+ class EndpointHandler:
5
+ def __init__(self, path=""):
6
+ self.tokenizer = AutoTokenizer.from_pretrained(path)
7
+ self.model = AutoModelForCausalLM.from_pretrained(path)
8
+
9
+ def __call__(self, data: Dict[str, str]) -> Dict[str, str]:
10
+ inputs = data.get("inputs", "")
11
+ if not inputs:
12
+ return {"error": "No input provided."}
13
+ inputs = self.tokenizer(inputs, return_tensors="pt")
14
+ outputs = self.model.generate(**inputs, max_new_tokens=100)
15
+ response = self.tokenizer.decode(outputs[0], skip_special_tokens=True)
16
+ return {"generated_text": response}
load_model.py DELETED
@@ -1,6 +0,0 @@
1
- from transformers import AutoModel, AutoTokenizer
2
-
3
- def load_custom_model():
4
- tokenizer = AutoTokenizer.from_pretrained("./model")
5
- model = AutoModel.from_pretrained("./model")
6
- return tokenizer, model