omkar1804 commited on
Commit
77f911a
·
verified ·
1 Parent(s): e1b278a

Update app.py

Browse files
Files changed (1) hide show
  1. app.py +9 -2
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
- sql = text[start.start() :].removesuffix("```").strip()
 
 
 
 
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
- stop_strings=[";\n", "\n#"],
 
 
 
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
  )