Spaces:
Running on Zero
Running on Zero
multimodalart HF Staff
Calibrate ZeroGPU duration from measured ~8s/turn throughput
ed584bc verified Download app.py from hugging-apps/trust-sql-text2sql-demo: direct link, hf CLI and curl.
- Browser
- Download file 28.2 kB
-
https://huggingface.co/spaces/hugging-apps/trust-sql-text2sql-demo/resolve/main/app.py
- Command line
-
hf download hf://spaces/hugging-apps/trust-sql-text2sql-demo/app.py
-
curl -L -o app.py https://huggingface.co/spaces/hugging-apps/trust-sql-text2sql-demo/resolve/main/app.py
28.2 kB
| """TRUST-SQL β Text-to-SQL over *unknown* schemas, on ZeroGPU. | |
| Faithful re-implementation of the four-phase tool-integrated agent loop from | |
| `JaneEyre0530/TrustSQL` (`trustsql_eval/`): the model never sees the schema. | |
| It must explore it with read-only metadata queries, propose a verified schema, | |
| generate + execute a candidate SQL, and only then confirm the final answer. | |
| """ | |
| import os | |
| os.environ.setdefault("PYTORCH_CUDA_ALLOC_CONF", "expandable_segments:True") | |
| os.environ.setdefault("TOKENIZERS_PARALLELISM", "false") | |
| import spaces # noqa: E402 (must precede torch) | |
| import json # noqa: E402 | |
| import re # noqa: E402 | |
| import sqlite3 # noqa: E402 | |
| import time # noqa: E402 | |
| from pathlib import Path # noqa: E402 | |
| from threading import Thread # noqa: E402 | |
| import gradio as gr # noqa: E402 | |
| import pandas as pd # noqa: E402 | |
| import torch # noqa: E402 | |
| from transformers import ( # noqa: E402 | |
| AutoModelForCausalLM, | |
| AutoTokenizer, | |
| StoppingCriteria, | |
| StoppingCriteriaList, | |
| TextIteratorStreamer, | |
| ) | |
| # -------------------------------------------------------------------------------------- | |
| # Model | |
| # -------------------------------------------------------------------------------------- | |
| MODEL_ID = "AIJian/TrustSQL-8B" | |
| HERE = Path(__file__).parent | |
| DB_ROOT = HERE / "databases" | |
| # Exact system prompt shipped by the authors (trustsql_eval/prompt_template.txt). | |
| SYSTEM_PROMPT = (HERE / "prompt_template.txt").read_text(encoding="utf-8").strip() | |
| tokenizer = AutoTokenizer.from_pretrained(MODEL_ID) | |
| model = AutoModelForCausalLM.from_pretrained( | |
| MODEL_ID, | |
| dtype=torch.bfloat16, | |
| attn_implementation="sdpa", | |
| ).to("cuda") | |
| model.eval() | |
| EOS_IDS = [151645, 151643] # <|im_end|>, <|endoftext|> | |
| MAX_CONTEXT = 40960 | |
| MAX_OBS_TOKENS = 2048 # trustsql_eval default | |
| SQL_TIMEOUT = 15.0 | |
| MAX_ROWS = 100 # trustsql_eval `_execute_sql_sync` default | |
| # -------------------------------------------------------------------------------------- | |
| # Sample databases (BIRD-Dev, CC BY-SA 4.0) | |
| # -------------------------------------------------------------------------------------- | |
| SAMPLE_DBS = { | |
| "california_schools": "California public schools β SAT scores, free-meal rates (3 tables)", | |
| "superhero": "Superhero attributes, powers, publishers (9 tables)", | |
| "student_club": "University club members, events, budgets, expenses (8 tables)", | |
| "toxicology": "Molecules, atoms, bonds and carcinogenicity labels (4 tables)", | |
| "formula_1": "Formula 1 races, drivers, constructors, lap times (13 tables)", | |
| } | |
| DB_CHOICES = [f"{k} β {v}" for k, v in SAMPLE_DBS.items()] | |
| def _db_id_from_choice(choice: str) -> str: | |
| return (choice or DB_CHOICES[0]).split(" β ")[0].strip() | |
| def _resolve_db(db_choice: str, uploaded_db): | |
| """Return (db_id, sqlite_path). An uploaded file always wins.""" | |
| if uploaded_db: | |
| path = uploaded_db if isinstance(uploaded_db, str) else getattr(uploaded_db, "name", None) | |
| if path and os.path.exists(path): | |
| return Path(path).stem, path | |
| db_id = _db_id_from_choice(db_choice) | |
| return db_id, str(DB_ROOT / db_id / f"{db_id}.sqlite") | |
| # -------------------------------------------------------------------------------------- | |
| # The one tool the agent gets: read-only SQL execution | |
| # -------------------------------------------------------------------------------------- | |
| ALLOWED_SQL_PREFIXES = ("SELECT", "PRAGMA", "EXPLAIN", "WITH") | |
| def _strip_sql_comments(sql: str) -> str: | |
| s = sql.strip() | |
| while s.startswith("--") or s.startswith("/*"): | |
| if s.startswith("--"): | |
| nl = s.find("\n") | |
| if nl == -1: | |
| return "" | |
| s = s[nl + 1 :].strip() | |
| else: | |
| end = s.find("*/") | |
| if end == -1: | |
| return "" | |
| s = s[end + 2 :].strip() | |
| return s | |
| def _is_readonly(sql: str): | |
| s = _strip_sql_comments(sql) | |
| if not s: | |
| return False, "Empty SQL query" | |
| if s.upper().startswith(ALLOWED_SQL_PREFIXES): | |
| return True, None | |
| return False, f"SQL must start with {ALLOWED_SQL_PREFIXES}, got: {s.split()[0]}" | |
| def run_sql(db_path: str, sql: str, max_rows: int = MAX_ROWS): | |
| """Execute read-only SQL. Returns (text_result, column_names, rows).""" | |
| ok, err = _is_readonly(sql) | |
| if not ok: | |
| return f"Error: {err}", [], [] | |
| if not os.path.exists(db_path): | |
| return f"Error: Database file not found: {db_path}", [], [] | |
| conn = None | |
| try: | |
| conn = sqlite3.connect(f"file:{db_path}?mode=ro", uri=True, check_same_thread=False) | |
| conn.execute(f"PRAGMA busy_timeout = {int(SQL_TIMEOUT * 1000)}") | |
| cur = conn.cursor() | |
| cur.execute(sql) | |
| rows = cur.fetchall() | |
| if not rows: | |
| return "Query executed successfully. No results returned.", ( | |
| [d[0] for d in cur.description] if cur.description else [] | |
| ), [] | |
| cols = [d[0] for d in cur.description] | |
| lines = ["\t".join(cols)] | |
| for i, row in enumerate(rows): | |
| if i >= max_rows: | |
| lines.append(f"... ({len(rows) - max_rows} more rows)") | |
| break | |
| lines.append("\t".join("NULL" if v is None else str(v) for v in row)) | |
| return "\n".join(lines), cols, rows[:max_rows] | |
| except sqlite3.Error as e: | |
| return f"Error: SQLite error: {e}", [], [] | |
| except Exception as e: # pragma: no cover | |
| return f"Error: Unexpected error: {e}", [], [] | |
| finally: | |
| if conn is not None: | |
| try: | |
| conn.close() | |
| except Exception: | |
| pass | |
| def db_schema_preview(db_choice: str, uploaded_db=None) -> str: | |
| """Human-readable DDL dump of a database (for the UI only β never shown to the model).""" | |
| db_id, path = _resolve_db(db_choice, uploaded_db) | |
| if not os.path.exists(path): | |
| return f"-- database `{db_id}` not found" | |
| text, _, rows = run_sql( | |
| path, | |
| "SELECT name, sql FROM sqlite_master WHERE type='table' AND name NOT LIKE 'sqlite_%'", | |
| max_rows=200, | |
| ) | |
| if not rows: | |
| return f"-- `{db_id}`: {text}" | |
| out = [f"-- database: {db_id} ({len(rows)} tables)", ""] | |
| for name, ddl in rows: | |
| out.append((ddl or f"-- {name}").strip() + ";") | |
| out.append("") | |
| return "\n".join(out) | |
| # -------------------------------------------------------------------------------------- | |
| # Prompt construction (matches AIJian/TrustSQL-data + trustsql_eval/prompt_builders.py) | |
| # -------------------------------------------------------------------------------------- | |
| def build_user_message(db_id: str, question: str, external_knowledge: str) -> str: | |
| parts = ["", "**Task Configuration**", "**Database Engine:** SQLite", f"**Database:** {db_id}"] | |
| if external_knowledge and external_knowledge.strip(): | |
| parts.append(f"**External Knowledge:** {external_knowledge.strip()}") | |
| parts.append(f"**User Question:** {question.strip()}?") | |
| parts.append("") | |
| return "\n".join(parts) | |
| def progress_prefix(current_round: int, max_rounds: int) -> str: | |
| """Verbatim port of MessageProcessor._format_progress_prefix.""" | |
| base = f"This is turn {current_round + 1} of {max_rounds}.\n\n" | |
| remaining = max_rounds - (current_round + 1) | |
| if remaining == 0: | |
| return "" | |
| if remaining == 1: | |
| return base + ( | |
| "Only 1 turn remaining after this.\n" | |
| "You MUST provide the final answer in the next turn.\n\n" | |
| "Use <action>confirm_answer</action> with your best SQL query.\n" | |
| "If you don't have a complete solution, provide your best attempt.\n\n" | |
| ) | |
| if remaining == 2: | |
| return base + ("Only 2 turns remaining after this.\nStart preparing your final SQL query.\n\n") | |
| return base | |
| FORMAT_HELP = ( | |
| "Invalid format detected. Your response is missing required components.\n\n" | |
| "Option 1: EXPLORE SCHEMA\n" | |
| "Purpose: Investigate database structure\n" | |
| "Required format:\n" | |
| "<think>Your reasoning process</think>\n" | |
| "<action>explore_schema</action>\n" | |
| '<tool_call>{"name": "execute_sql_query", "arguments": {"db_id": "...", "sql": "..."}}</tool_call>\n\n' | |
| "Option 2: PROPOSE SCHEMA\n" | |
| "Purpose: Document your understanding of relevant tables and columns\n" | |
| "Required format:\n" | |
| "<think>Your reasoning process</think>\n" | |
| "<action>propose_schema</action>\n" | |
| '<schema>{"tables": [...], "columns": {...}}</schema>\n\n' | |
| "Option 3: GENERATE SQL\n" | |
| "Purpose: Create SQL query and VERIFY it works by executing\n" | |
| "<think>Your reasoning process</think>\n" | |
| "<action>generate_sql</action>\n" | |
| '<tool_call>{"name": "execute_sql_query", "arguments": {"db_id": "...", "sql": "..."}}</tool_call>\n\n' | |
| "Option 4: FINAL ANSWER\n" | |
| "Purpose: Provide verified SQL query as final result\n" | |
| "ONLY use this AFTER successfully executing and verifying your SQL.\n" | |
| "Required format:\n" | |
| "<think>Your reasoning process</think>\n" | |
| "<action>confirm_answer</action>\n" | |
| "<answer>```sql\nYOUR_SQL\n```</answer>\n\n" | |
| ) | |
| def fix_tool_tag(content: str) -> str: | |
| content = re.sub(r"<tool>(.*?)</tool>", r"<tool_call>\1</tool_call>", content, flags=re.S) | |
| content = re.sub(r"<tools>(.*?)</tools>", r"<tool_call>\1</tool_call>", content, flags=re.S) | |
| return content | |
| def extract_tag(text: str, tag: str): | |
| m = re.search(rf"<{tag}>(.*?)</{tag}>", text, re.S | re.I) | |
| return m.group(1) if m else None | |
| def extract_final_sql(answer_body: str) -> str: | |
| for pat in (r"```sql\s*(.*?)```", r"'''sql\s*(.*?)'''", r"```\s*(.*?)```", r"'''\s*(.*?)'''"): | |
| m = re.search(pat, answer_body, re.S | re.I) | |
| if m: | |
| return m.group(1).strip() | |
| return answer_body.strip() | |
| def truncate_observation(text: str, max_tokens: int = MAX_OBS_TOKENS) -> str: | |
| ids = tokenizer(text, add_special_tokens=False)["input_ids"] | |
| if len(ids) <= max_tokens: | |
| return text | |
| return tokenizer.decode(ids[:max_tokens]) + "\n... (result truncated due to length)" | |
| # -------------------------------------------------------------------------------------- | |
| # Pretty-printing a turn for the chat transcript | |
| # -------------------------------------------------------------------------------------- | |
| ACTION_ICON = { | |
| "explore_schema": "π", | |
| "propose_schema": "π", | |
| "generate_sql": "π οΈ", | |
| "confirm_answer": "β ", | |
| } | |
| def _fence(text: str, lang: str = "") -> str: | |
| return f"```{lang}\n{str(text).replace('```', '`` `')}\n```" | |
| def _quote(text: str) -> str: | |
| text = text.strip() | |
| return "\n".join("> " + line for line in text.splitlines()) if text else "" | |
| def render_assistant(raw: str) -> str: | |
| body = fix_tool_tag(raw) | |
| think = extract_tag(body, "think") | |
| action = (extract_tag(body, "action") or "").strip().lower() | |
| blocks = [] | |
| if action: | |
| blocks.append(f"### {ACTION_ICON.get(action, 'βοΈ')} `{action}`") | |
| elif not think: | |
| blocks.append("### βοΈ raw response") | |
| if think: | |
| blocks.append("π **Reasoning**\n\n" + _quote(think)) | |
| schema = extract_tag(body, "schema") | |
| if schema is not None: | |
| try: | |
| pretty = json.dumps(json.loads(schema), indent=2) | |
| except Exception: | |
| pretty = schema.strip() | |
| blocks.append("**Proposed schema**\n\n" + _fence(pretty, "json")) | |
| answer = extract_tag(body, "answer") | |
| if answer is not None: | |
| blocks.append("**Final SQL**\n\n" + _fence(extract_final_sql(answer), "sql")) | |
| tool_call = extract_tag(body, "tool_call") | |
| if tool_call is not None: | |
| sql = None | |
| try: | |
| payload = json.loads(tool_call.strip()) | |
| sql = (payload.get("arguments") or {}).get("sql") | |
| except Exception: | |
| pass | |
| if sql: | |
| blocks.append("**Tool call** Β· `execute_sql_query`\n\n" + _fence(sql, "sql")) | |
| else: | |
| blocks.append("**Tool call**\n\n" + _fence(tool_call.strip(), "json")) | |
| if not blocks: | |
| return _fence(raw) | |
| return "\n\n".join(blocks) | |
| def render_observation(text: str) -> str: | |
| return "π₯ **Observation**\n\n" + _fence(text) | |
| # -------------------------------------------------------------------------------------- | |
| # Generation | |
| # -------------------------------------------------------------------------------------- | |
| class _Deadline(StoppingCriteria): | |
| def __init__(self, deadline: float): | |
| self.deadline = deadline | |
| def __call__(self, input_ids, scores, **kwargs) -> bool: | |
| return time.time() > self.deadline | |
| def stream_turn(messages, max_new_tokens: int, temperature: float, top_p: float, deadline: float): | |
| """Yield incremental text for one assistant turn; last yield is the full turn.""" | |
| prompt = tokenizer.apply_chat_template(messages, tokenize=False, add_generation_prompt=True) | |
| enc = tokenizer(prompt, return_tensors="pt", add_special_tokens=False) | |
| n_in = enc["input_ids"].shape[-1] | |
| budget = max(64, min(int(max_new_tokens), MAX_CONTEXT - n_in - 8)) | |
| enc = {k: v.to(model.device) for k, v in enc.items()} | |
| streamer = TextIteratorStreamer(tokenizer, skip_prompt=True, skip_special_tokens=True) | |
| do_sample = float(temperature) > 0.0 | |
| kwargs = dict( | |
| **enc, | |
| streamer=streamer, | |
| max_new_tokens=budget, | |
| do_sample=do_sample, | |
| eos_token_id=EOS_IDS, | |
| pad_token_id=151643, | |
| stopping_criteria=StoppingCriteriaList([_Deadline(deadline)]), | |
| ) | |
| if do_sample: | |
| kwargs.update(temperature=float(temperature), top_p=float(top_p), top_k=20) | |
| thread = Thread(target=model.generate, kwargs=kwargs) | |
| thread.start() | |
| acc = "" | |
| last = 0.0 | |
| for chunk in streamer: | |
| acc += chunk | |
| now = time.time() | |
| if now - last > 0.25: | |
| last = now | |
| yield acc, False | |
| thread.join() | |
| yield acc.strip(), True | |
| def _estimate_duration(*args, **kwargs) -> int: | |
| """GPU budget, sized from measured throughput (~8 s/turn at 1536 max_new_tokens).""" | |
| max_turns, max_new = 8, 1536 | |
| try: | |
| if len(args) >= 5 and args[4] is not None: | |
| max_turns = int(args[4]) | |
| if len(args) >= 6 and args[5] is not None: | |
| max_new = int(args[5]) | |
| except Exception: | |
| pass | |
| per_turn = 6.0 + 6.0 * (max_new / 1536.0) | |
| return int(min(280, 15 + max_turns * per_turn)) | |
| # -------------------------------------------------------------------------------------- | |
| # The agent loop | |
| # -------------------------------------------------------------------------------------- | |
| def run_agent( | |
| question: str, | |
| db_choice: str = DB_CHOICES[0], | |
| external_knowledge: str = "", | |
| uploaded_db=None, | |
| max_turns: int = 8, | |
| max_new_tokens: int = 1536, | |
| temperature: float = 0.7, | |
| top_p: float = 0.9, | |
| ): | |
| """Run the TRUST-SQL agent on an unknown SQLite database and return the final SQL. | |
| The model receives only the database *name* and the question β never the schema. | |
| It explores metadata with read-only queries, proposes a verified schema, executes a | |
| candidate query, and confirms the final SQL. | |
| Args: | |
| question: the natural-language question to answer. | |
| db_choice: which bundled BIRD-Dev sample database to query. | |
| external_knowledge: optional domain hint / evidence string (BIRD "evidence" field). | |
| uploaded_db: optional path to a user-supplied SQLite file; overrides `db_choice`. | |
| max_turns: maximum agent turns before giving up. | |
| max_new_tokens: token budget per agent turn. | |
| temperature: sampling temperature; 0 means greedy decoding. | |
| top_p: nucleus sampling cutoff. | |
| """ | |
| t0 = time.time() | |
| budget = _estimate_duration(question, db_choice, external_knowledge, uploaded_db, max_turns) | |
| deadline = t0 + budget - 18 | |
| max_turns = int(max_turns) | |
| db_id, db_path = _resolve_db(db_choice, uploaded_db) | |
| user_msg = build_user_message(db_id, question or "", external_knowledge or "") | |
| messages = [ | |
| {"role": "system", "content": SYSTEM_PROMPT}, | |
| {"role": "user", "content": user_msg}, | |
| ] | |
| chat = [{"role": "user", "content": f"**Question**\n\n{question}\n\n_Database: `{db_id}` (schema unknown to the model)_"}] | |
| empty_df = pd.DataFrame() | |
| if not (question or "").strip(): | |
| yield chat + [{"role": "assistant", "content": "Please enter a question."}], "", empty_df, "β οΈ No question provided." | |
| return | |
| if not os.path.exists(db_path): | |
| yield chat, "", empty_df, f"β Database not found: `{db_path}`" | |
| return | |
| yield chat, "", empty_df, f"β³ Turn 1/{max_turns} β exploring `{db_id}`β¦" | |
| final_sql = "" | |
| status = "" | |
| for turn in range(max_turns): | |
| if time.time() > deadline: | |
| status = f"β±οΈ Stopped after {turn} turn(s): GPU time budget reached." | |
| break | |
| base = list(chat) | |
| raw = "" | |
| for text, done in stream_turn(messages, max_new_tokens, temperature, top_p, deadline): | |
| raw = text | |
| chat = base + [ | |
| { | |
| "role": "assistant", | |
| "content": (render_assistant(text) if done else _fence(text)), | |
| } | |
| ] | |
| yield chat, final_sql, empty_df, f"β³ Turn {turn + 1}/{max_turns} β generatingβ¦" | |
| if not raw: | |
| status = "β The model returned an empty response." | |
| break | |
| raw = fix_tool_tag(raw) | |
| # ---- confirm_answer -> terminate ------------------------------------------- | |
| answer = extract_tag(raw, "answer") | |
| if answer is not None: | |
| messages.append({"role": "assistant", "content": raw}) | |
| final_sql = extract_final_sql(answer) | |
| status = f"β Confirmed after {turn + 1} turn(s)." | |
| break | |
| messages.append({"role": "assistant", "content": raw}) | |
| prefix = progress_prefix(turn, max_turns) | |
| # ---- propose_schema -> acknowledgement -------------------------------------- | |
| schema = extract_tag(raw, "schema") | |
| if schema is not None: | |
| try: | |
| data = json.loads(schema) | |
| tables = data.get("tables", []) or [] | |
| cols = data.get("columns", {}) or {} | |
| n_cols = sum(len(v) for v in cols.values()) if isinstance(cols, dict) else len(cols) | |
| feedback = ( | |
| prefix | |
| + f"Schema acknowledged: {len(tables)} table(s), {n_cols} column(s). " | |
| "You may now proceed to generate SQL.\n" | |
| ) | |
| except Exception: | |
| feedback = prefix + "Schema acknowledged. You may proceed to generate SQL.\n" | |
| messages.append({"role": "user", "content": feedback}) | |
| chat = chat + [{"role": "user", "content": render_observation(feedback)}] | |
| yield chat, final_sql, empty_df, f"β³ Turn {turn + 2}/{max_turns}β¦" | |
| continue | |
| # ---- tool call -> execute ---------------------------------------------------- | |
| tool_call = extract_tag(raw, "tool_call") | |
| obs = None | |
| if tool_call is None: | |
| obs = prefix + FORMAT_HELP | |
| else: | |
| try: | |
| payload = json.loads(tool_call.strip()) | |
| name = payload.get("name", "") | |
| arguments = payload.get("arguments", {}) or {} | |
| if name != "execute_sql_query": | |
| obs = prefix + f"Error: Unknown function: {name}" | |
| elif not str(arguments.get("sql", "")).strip(): | |
| obs = prefix + "Error: SQL query is empty" | |
| else: | |
| result, _, _ = run_sql(db_path, arguments["sql"]) | |
| obs = prefix + truncate_observation(result) | |
| except json.JSONDecodeError as e: | |
| obs = ( | |
| prefix | |
| + f"Tool call parsing error:\nJSON parsing failed at line {e.lineno}, column {e.colno}: {e.msg}\n\n" | |
| "Please fix the JSON format and try again.\n\n" | |
| "Required format:\n" | |
| '<tool_call>{"name": "execute_sql_query", "arguments": {"db_id": "...", "sql": "..."}}</tool_call>\n' | |
| ) | |
| except Exception as e: # pragma: no cover | |
| obs = prefix + f"Error: Tool execution error: {e}" | |
| messages.append({"role": "user", "content": obs}) | |
| chat = chat + [{"role": "user", "content": render_observation(obs)}] | |
| yield chat, final_sql, empty_df, f"β³ Turn {turn + 2}/{max_turns}β¦" | |
| else: | |
| status = f"β οΈ Reached the {max_turns}-turn limit without a confirmed answer." | |
| # ---- execute the confirmed SQL for display -------------------------------------- | |
| df = empty_df | |
| if final_sql: | |
| text, cols, rows = run_sql(db_path, final_sql, max_rows=100) | |
| if cols and rows: | |
| df = pd.DataFrame(rows, columns=cols) | |
| elif cols: | |
| df = pd.DataFrame(columns=cols) | |
| else: | |
| chat = chat + [{"role": "assistant", "content": "β οΈ Final SQL did not execute:\n\n" + _fence(text)}] | |
| else: | |
| status = status or "β οΈ No SQL was confirmed." | |
| elapsed = time.time() - t0 | |
| yield chat, final_sql, df, f"{status} Β· {elapsed:.0f}s on GPU" | |
| # -------------------------------------------------------------------------------------- | |
| # UI | |
| # -------------------------------------------------------------------------------------- | |
| EXAMPLES = [ | |
| [ | |
| "What is the highest eligible free rate for K-12 students in the schools in Alameda County?", | |
| DB_CHOICES[0], | |
| "Eligible free rate for K-12 = `Free Meal Count (K-12)` / `Enrollment (K-12)`", | |
| ], | |
| [ | |
| "Among the schools with the SAT test takers of over 500, please list the schools that are magnet schools or offer a magnet program.", | |
| DB_CHOICES[0], | |
| "Magnet schools or offer a magnet program means that Magnet = 1", | |
| ], | |
| [ | |
| "How many superheroes have blue eyes?", | |
| DB_CHOICES[1], | |
| "blue eyes refers to colour = 'Blue' and eye_colour_id = colour.id", | |
| ], | |
| [ | |
| "Please list all the superpowers of 3-D Man.", | |
| DB_CHOICES[1], | |
| "3-D Man refers to superhero_name = '3-D Man'; superpowers refers to power_name", | |
| ], | |
| [ | |
| "What is the event that has the highest attendance of the students from the Student_Club?", | |
| DB_CHOICES[2], | |
| "event with highest attendance refers to MAX(COUNT(link_to_event))", | |
| ], | |
| [ | |
| "In the non-carcinogenic molecules, how many contain chlorine atoms?", | |
| DB_CHOICES[3], | |
| "non-carcinogenic molecules refers to label = '-'; chlorine atoms refers to element = 'cl'", | |
| ], | |
| [ | |
| "Please give the name of the race held on the circuits in Germany.", | |
| DB_CHOICES[4], | |
| "Germany is a name of country;", | |
| ], | |
| ] | |
| CSS = """ | |
| #col-container { max-width: 1180px; margin: 0 auto; } | |
| .dark .gradio-container { color: var(--body-text-color); } | |
| """ | |
| INTRO = """# π TRUST-SQL β Text-to-SQL over **unknown** schemas | |
| [`AIJian/TrustSQL-8B`](https://huggingface.co/AIJian/TrustSQL-8B) Β· [paper](https://huggingface.co/papers/2603.16448) Β· [code](https://github.com/JaneEyre0530/TrustSQL) | |
| The schema is **not** in the prompt. The agent gets one tool β read-only SQL β and has to discover | |
| the database itself, following the authors' four-phase protocol: | |
| `explore_schema β propose_schema β generate_sql β confirm_answer`. | |
| """ | |
| with gr.Blocks(title="TRUST-SQL") as demo: | |
| with gr.Column(elem_id="col-container"): | |
| gr.Markdown(INTRO) | |
| with gr.Row(): | |
| with gr.Column(scale=3): | |
| question = gr.Textbox( | |
| label="Question", | |
| placeholder="e.g. Which school has the highest average SAT math score?", | |
| lines=2, | |
| ) | |
| with gr.Column(scale=1, min_width=140): | |
| run_btn = gr.Button("Run agent", variant="primary", size="lg") | |
| with gr.Row(): | |
| db_choice = gr.Dropdown( | |
| label="Database (BIRD-Dev sample)", | |
| choices=DB_CHOICES, | |
| value=DB_CHOICES[0], | |
| scale=2, | |
| ) | |
| external_knowledge = gr.Textbox( | |
| label="External knowledge (optional hint)", | |
| placeholder="e.g. charter schools refers to `Charter School (Y/N)` = 1", | |
| lines=1, | |
| scale=3, | |
| ) | |
| status = gr.Markdown("") | |
| with gr.Row(): | |
| with gr.Column(scale=3): | |
| chatbot = gr.Chatbot( | |
| label="Agent trajectory", | |
| height=620, | |
| resizable=True, | |
| group_consecutive_messages=False, | |
| ) | |
| with gr.Column(scale=2): | |
| final_sql = gr.Code(label="Confirmed SQL", language="sql", lines=8) | |
| result_df = gr.Dataframe(label="Execution result", wrap=True) | |
| with gr.Accordion("Peek at the database (the agent never sees this)", open=False): | |
| schema_box = gr.Code(label="DDL", language="sql", lines=14) | |
| peek_btn = gr.Button("Show schema", size="sm") | |
| with gr.Accordion("Advanced settings", open=False): | |
| uploaded_db = gr.File( | |
| label="Use your own SQLite database (.sqlite / .db) β overrides the dropdown", | |
| file_types=[".sqlite", ".db", ".sqlite3"], | |
| type="filepath", | |
| ) | |
| with gr.Row(): | |
| max_turns = gr.Slider(3, 12, value=8, step=1, label="Max agent turns") | |
| max_new_tokens = gr.Slider(256, 3072, value=1536, step=128, label="Max new tokens / turn") | |
| with gr.Row(): | |
| temperature = gr.Slider(0.0, 1.0, value=0.7, step=0.05, label="Temperature (0 = greedy)") | |
| top_p = gr.Slider(0.1, 1.0, value=0.9, step=0.05, label="Top-p") | |
| gr.Markdown( | |
| "Defaults mirror `trustsql_eval` (temperature 0.7 / top-p 0.9). " | |
| "Only `SELECT` / `PRAGMA` / `EXPLAIN` / `WITH` statements are ever executed, " | |
| "against a read-only connection." | |
| ) | |
| gr.Examples( | |
| examples=EXAMPLES, | |
| inputs=[question, db_choice, external_knowledge], | |
| outputs=[chatbot, final_sql, result_df, status], | |
| fn=run_agent, | |
| cache_examples=True, | |
| cache_mode="lazy", | |
| label="BIRD-Dev examples (question + official evidence hint)", | |
| ) | |
| gr.Markdown( | |
| "Sample databases are the **BIRD-Dev** SQLite databases " | |
| "([BIRD-SQL](https://bird-bench.github.io/), CC BY-SA 4.0); the example questions and " | |
| "hints are their official dev-set questions and `evidence` strings. " | |
| "The system prompt and agent loop are ported verbatim from " | |
| "[`JaneEyre0530/TrustSQL`](https://github.com/JaneEyre0530/TrustSQL) (Apache-2.0)." | |
| ) | |
| run_btn.click( | |
| fn=run_agent, | |
| inputs=[ | |
| question, | |
| db_choice, | |
| external_knowledge, | |
| uploaded_db, | |
| max_turns, | |
| max_new_tokens, | |
| temperature, | |
| top_p, | |
| ], | |
| outputs=[chatbot, final_sql, result_df, status], | |
| api_name="run_agent", | |
| ) | |
| question.submit( | |
| fn=run_agent, | |
| inputs=[ | |
| question, | |
| db_choice, | |
| external_knowledge, | |
| uploaded_db, | |
| max_turns, | |
| max_new_tokens, | |
| temperature, | |
| top_p, | |
| ], | |
| outputs=[chatbot, final_sql, result_df, status], | |
| api_name=False, | |
| ) | |
| peek_btn.click( | |
| fn=db_schema_preview, inputs=[db_choice, uploaded_db], outputs=schema_box, api_name="schema" | |
| ) | |
| db_choice.change(fn=db_schema_preview, inputs=[db_choice, uploaded_db], outputs=schema_box, api_name=False) | |
| if __name__ == "__main__": | |
| demo.launch(theme=gr.themes.Citrus(), css=CSS, mcp_server=True) | |