File size: 4,221 Bytes
dccf890
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
"""
SQL Copilot — plain-English questions to DuckDB SQL with Qwen2.5-Coder-7B on
ZeroGPU. The portfolio's Data Lab sends table schemas (plus a few sample rows)
and runs the returned SQL locally in DuckDB-WASM. Errors travel as data.
"""

import re
import traceback

import gradio as gr
import spaces
import torch
from transformers import AutoModelForCausalLM, AutoTokenizer

MODEL = "Qwen/Qwen2.5-Coder-7B-Instruct"
MAX_QUESTION = 500
MAX_SCHEMA = 8000

tokenizer = AutoTokenizer.from_pretrained(MODEL)
model = AutoModelForCausalLM.from_pretrained(MODEL, torch_dtype=torch.bfloat16).to("cuda")

SYSTEM = """You are an expert data analyst who writes DuckDB SQL.
Rules:
- Answer with ONE DuckDB SQL query and nothing else: no explanation, no comments.
- Use only the tables and columns in the schema. Quote identifiers with double quotes when they contain capitals, spaces or symbols.
- Prefer readable column aliases, ORDER BY for rankings, and LIMIT 100 for row listings.
- For dates use DuckDB functions such as date_trunc, strftime, extract.
- Round averages and percentages to 2 decimals.
- Read-only: never write CREATE, INSERT, UPDATE, DELETE, DROP, ALTER, COPY, ATTACH or INSTALL."""

WRITE = re.compile(r"\b(create|insert|update|delete|drop|alter|copy|attach|install|load|pragma|export)\b", re.I)


@spaces.GPU(duration=20)
def complete(messages):
    text = tokenizer.apply_chat_template(messages, tokenize=False, add_generation_prompt=True)
    inputs = tokenizer([text], return_tensors="pt").to("cuda")
    with torch.inference_mode():
        out = model.generate(**inputs, max_new_tokens=400, do_sample=False, pad_token_id=tokenizer.eos_token_id)
    return tokenizer.decode(out[0][inputs.input_ids.shape[1]:], skip_special_tokens=True)


def extract_sql(reply):
    m = re.search(r"```(?:sql)?\s*(.*?)```", reply, re.S | re.I)
    sql = (m.group(1) if m else reply).strip().rstrip(";").strip()
    return sql


def ask(question, schema):
    """API: returns (sql, error)."""
    try:
        question = (question or "").strip()
        schema = (schema or "").strip()
        if not question:
            return "", "Ask a question about your data first."
        if len(question) > MAX_QUESTION:
            return "", f"Keep the question under {MAX_QUESTION} characters."
        if not schema:
            return "", "Load a table first."
        schema = schema[:MAX_SCHEMA]
        reply = complete([
            {"role": "system", "content": SYSTEM},
            {"role": "user", "content": f"Schema:\n{schema}\n\nQuestion: {question}"},
        ])
        sql = extract_sql(reply)
        if not sql or not re.match(r"^\s*(with|select|from|summarize|describe|pivot|unpivot)\b", sql, re.I):
            return "", "Couldn't turn that into a query. Try rephrasing it."
        if WRITE.search(re.sub(r"'[^']*'", "''", sql)):
            return "", "The copilot only writes read-only queries."
        return sql, ""
    except gr.Error as exc:
        return "", str(exc.message)
    except Exception as exc:
        traceback.print_exc()
        return "", f"The copilot failed ({type(exc).__name__}). Please try again."


def ui_ask(question, schema):
    sql, error = ask(question, schema)
    if error:
        raise gr.Error(error)
    return sql


with gr.Blocks(title="SQL Copilot") as demo:
    gr.Markdown(
        "# 🧮 SQL Copilot\nPlain-English questions to DuckDB SQL with Qwen2.5-Coder-7B. Part of "
        "[Feliks Altymyshov's](https://github.com/feliksKdm) portfolio lab."
    )
    schema = gr.Textbox(label="Schema", lines=6, placeholder='orders(order_id BIGINT, order_date DATE, store VARCHAR, revenue DOUBLE)')
    question = gr.Textbox(label="Question", placeholder="Which store had the highest revenue last month?")
    btn = gr.Button("Write SQL", variant="primary")
    out = gr.Code(label="SQL", language="sql")
    btn.click(ui_ask, [question, schema], out, api_name=False)

    with gr.Group(visible=False):
        a_q, a_s, a_sql, a_err = gr.Textbox(), gr.Textbox(), gr.Textbox(), gr.Textbox()
        a_btn = gr.Button()
    a_btn.click(ask, [a_q, a_s], [a_sql, a_err], api_name="ask")

if __name__ == "__main__":
    demo.queue(default_concurrency_limit=2).launch()