Spaces:
Runtime error
Runtime error
Update app.py
Browse files
app.py
CHANGED
|
@@ -91,7 +91,11 @@ def extract_sql(generated: str) -> str:
|
|
| 91 |
if not start:
|
| 92 |
raise gr.Error("The model did not return a SQL SELECT statement.")
|
| 93 |
|
| 94 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
| 95 |
if not sql:
|
| 96 |
raise gr.Error("The model returned an empty response.")
|
| 97 |
if len(sql) > MAX_OUTPUT_CHARS:
|
|
@@ -132,7 +136,10 @@ def generate_sql(prompt: str, max_new_tokens: int) -> str:
|
|
| 132 |
num_beams=1,
|
| 133 |
eos_token_id=tokenizer.eos_token_id,
|
| 134 |
pad_token_id=tokenizer.pad_token_id,
|
| 135 |
-
|
|
|
|
|
|
|
|
|
|
| 136 |
tokenizer=tokenizer,
|
| 137 |
use_cache=True,
|
| 138 |
)
|
|
|
|
| 91 |
if not start:
|
| 92 |
raise gr.Error("The model did not return a SQL SELECT statement.")
|
| 93 |
|
| 94 |
+
sql_text = text[start.start() :].removesuffix("```").strip()
|
| 95 |
+
|
| 96 |
+
# This endpoint's contract is exactly one single-line SQL statement. Keep
|
| 97 |
+
# the first generated SQL line if the checkpoint starts a second sample.
|
| 98 |
+
sql = next((line.strip() for line in sql_text.splitlines() if line.strip()), "")
|
| 99 |
if not sql:
|
| 100 |
raise gr.Error("The model returned an empty response.")
|
| 101 |
if len(sql) > MAX_OUTPUT_CHARS:
|
|
|
|
| 136 |
num_beams=1,
|
| 137 |
eos_token_id=tokenizer.eos_token_id,
|
| 138 |
pad_token_id=tokenizer.pad_token_id,
|
| 139 |
+
# The model was trained on labelled, newline-delimited examples.
|
| 140 |
+
# Stop at its first completed response line instead of allowing it
|
| 141 |
+
# to continue into another sample until max_new_tokens is reached.
|
| 142 |
+
stop_strings=["\n", "```"],
|
| 143 |
tokenizer=tokenizer,
|
| 144 |
use_cache=True,
|
| 145 |
)
|