Upload handler.py
Browse files- handler.py +13 -7
handler.py
CHANGED
|
@@ -22,7 +22,6 @@ class EndpointHandler():
|
|
| 22 |
self.model = PeftModel.from_pretrained(model, path)
|
| 23 |
|
| 24 |
def __call__(self, data: Dict[str, Any]) -> Dict[str, Any]:
|
| 25 |
-
LOGGER.info(f"Received data: {data}")
|
| 26 |
# Get inputs
|
| 27 |
# Extract required parameters from data
|
| 28 |
message = data.get("message")
|
|
@@ -127,18 +126,25 @@ def generate(
|
|
| 127 |
|
| 128 |
outputs = []
|
| 129 |
generated_text = ""
|
| 130 |
-
streaming = True
|
| 131 |
conclusion_found = None
|
| 132 |
context_numbers = []
|
| 133 |
for text in streamer:
|
| 134 |
-
if not streaming:
|
| 135 |
-
break
|
| 136 |
outputs.append(text)
|
| 137 |
generated_text = "".join(outputs)
|
| 138 |
for end_sequence in end_sequences:
|
| 139 |
if end_sequence in generated_text:
|
| 140 |
-
streaming = False
|
| 141 |
generated_text = generated_text.replace(end_sequence, "")
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 142 |
|
| 143 |
# Check for conclusion keys in the generated text
|
| 144 |
if conclusions:
|
|
@@ -153,7 +159,7 @@ def generate(
|
|
| 153 |
context_numbers = [int(match.strip("[]")) for match in context_matches]
|
| 154 |
|
| 155 |
return {
|
| 156 |
-
"generated_text": generated_text
|
| 157 |
"conclusion": conclusion_found,
|
| 158 |
"context": context_numbers
|
| 159 |
-
}
|
|
|
|
| 22 |
self.model = PeftModel.from_pretrained(model, path)
|
| 23 |
|
| 24 |
def __call__(self, data: Dict[str, Any]) -> Dict[str, Any]:
|
|
|
|
| 25 |
# Get inputs
|
| 26 |
# Extract required parameters from data
|
| 27 |
message = data.get("message")
|
|
|
|
| 126 |
|
| 127 |
outputs = []
|
| 128 |
generated_text = ""
|
|
|
|
| 129 |
conclusion_found = None
|
| 130 |
context_numbers = []
|
| 131 |
for text in streamer:
|
|
|
|
|
|
|
| 132 |
outputs.append(text)
|
| 133 |
generated_text = "".join(outputs)
|
| 134 |
for end_sequence in end_sequences:
|
| 135 |
if end_sequence in generated_text:
|
|
|
|
| 136 |
generated_text = generated_text.replace(end_sequence, "")
|
| 137 |
+
return parse(generated_text, conclusions, end_sequences)
|
| 138 |
+
|
| 139 |
+
def parse(generated_text: str, conclusions: list[tuple[str, str]], end_sequences: list[str]) -> dict:
|
| 140 |
+
# Initialize variables
|
| 141 |
+
conclusion_found = None
|
| 142 |
+
context_numbers = []
|
| 143 |
+
|
| 144 |
+
# Remove end sequences and clean the text
|
| 145 |
+
for end_sequence in end_sequences:
|
| 146 |
+
generated_text = generated_text.replace(end_sequence, "")
|
| 147 |
+
generated_text = generated_text.strip()
|
| 148 |
|
| 149 |
# Check for conclusion keys in the generated text
|
| 150 |
if conclusions:
|
|
|
|
| 159 |
context_numbers = [int(match.strip("[]")) for match in context_matches]
|
| 160 |
|
| 161 |
return {
|
| 162 |
+
"generated_text": generated_text,
|
| 163 |
"conclusion": conclusion_found,
|
| 164 |
"context": context_numbers
|
| 165 |
+
}
|