import os import threading import gradio as gr import spaces import torch from transformers import ( AutoModelForCausalLM, AutoTokenizer, BitsAndBytesConfig, TextIteratorStreamer, ) MODEL_ID = "beaunix/aegis-sql" SYSTEM_PROMPT = ( "You are Aegis-SQL, an expert assistant that converts natural language " "questions into precise SQL queries, and can also explain existing SQL " "queries back into natural language." ) print("Loading tokenizer...") tokenizer = AutoTokenizer.from_pretrained(MODEL_ID) print("Loading model in 4-bit...") bnb_config = BitsAndBytesConfig( load_in_4bit=True, bnb_4bit_quant_type="nf4", bnb_4bit_compute_dtype=torch.bfloat16, bnb_4bit_use_double_quant=True, ) model = AutoModelForCausalLM.from_pretrained( MODEL_ID, quantization_config=bnb_config, device_map="auto", torch_dtype=torch.bfloat16, ) model.eval() print("Model loaded.") def to_text(content): """Flatten anything Gradio hands us into a plain string. """ if isinstance(content, str): return content if isinstance(content, list): parts = [] for item in content: if isinstance(item, str): parts.append(item) elif isinstance(item, dict) and "text" in item: parts.append(item["text"]) return " ".join(parts) if isinstance(content, dict) and "text" in content: return content["text"] return str(content) def build_prompt(schema: str, question: str) -> str: """Assemble the exact ChatML shape the model was trained on. The training manifest used: ### Database Schema\n\n\n### Question\n inside the user turn. Matching it here is what keeps inference in distribution. """ schema = to_text(schema).strip() question = to_text(question).strip() user_block = f"### Database Schema\n{schema}\n\n### Question\n{question}" messages = [ {"role": "system", "content": SYSTEM_PROMPT}, {"role": "user", "content": user_block}, ] return tokenizer.apply_chat_template( messages, tokenize=False, add_generation_prompt=True ) @spaces.GPU(duration=120) def generate_sql(schema, question): """Stream the generated SQL token by token. duration=120 is the max seconds ZeroGPU keeps the GPU attached per call; it bills real inference time, not the ceiling, so the headroom is free. SQL generations are short (a single query), so this is generous on purpose to cover GPU attach + occasional longer output. """ if not to_text(question).strip(): yield "-- Enter a question above and the generated SQL will appear here." return prompt = build_prompt(schema, question) inputs = tokenizer(prompt, return_tensors="pt").to(model.device) streamer = TextIteratorStreamer( tokenizer, skip_prompt=True, skip_special_tokens=True ) generation_kwargs = dict( **inputs, streamer=streamer, max_new_tokens=256, do_sample=True, temperature=0.1, # low: SQL wants determinism, not creativity top_p=0.9, pad_token_id=tokenizer.eos_token_id, ) thread = threading.Thread(target=model.generate, kwargs=generation_kwargs) thread.start() accumulated = "" for token in streamer: accumulated += token yield accumulated.strip() theme = gr.themes.Base( primary_hue=gr.themes.colors.cyan, secondary_hue=gr.themes.colors.amber, neutral_hue=gr.themes.colors.slate, font=[gr.themes.GoogleFont("Inter"), "system-ui", "sans-serif"], font_mono=[gr.themes.GoogleFont("JetBrains Mono"), "monospace"], ).set( body_background_fill="#0d1117", body_background_fill_dark="#0d1117", block_background_fill="#161b22", block_background_fill_dark="#161b22", block_border_color="#30363d", block_border_color_dark="#30363d", body_text_color="#e6edf3", body_text_color_dark="#e6edf3", block_label_text_color="#7ee7e0", block_label_text_color_dark="#7ee7e0", button_primary_background_fill="#1f6feb", button_primary_background_fill_hover="#388bfd", button_primary_text_color="#ffffff", ) CSS = """ .aegis-header { padding: 1.4rem 0 0.4rem 0; } .aegis-header h1 { font-family: 'JetBrains Mono', monospace; font-weight: 700; font-size: 1.9rem; letter-spacing: -0.02em; color: #e6edf3; margin: 0; } .aegis-header .accent { color: #58d6cf; } .aegis-header p { color: #8b949e; margin: 0.35rem 0 0 0; font-size: 0.95rem; } .aegis-rule { height: 2px; width: 100%; background: linear-gradient(90deg, #58d6cf 0%, #1f6feb 55%, transparent 100%); margin: 0.8rem 0 0.2rem 0; border-radius: 2px; } .aegis-foot { color: #6e7681; font-size: 0.8rem; font-family: 'JetBrains Mono', monospace; padding-top: 0.6rem; } """ EXAMPLE_SCHEMA = """CREATE TABLE employees ( employee_id INTEGER PRIMARY KEY, name TEXT, department_id INTEGER, salary REAL, hire_date DATE ); CREATE TABLE departments ( department_id INTEGER PRIMARY KEY, department_name TEXT );""" EXAMPLE_QUESTION = ( "What is the average salary in each department, " "showing only departments with more than 5 employees?" ) with gr.Blocks(theme=theme, css=CSS, title="Aegis-SQL") as demo: gr.HTML( """

aegis-sql

Natural language to SQL. Paste a schema, ask a question, read the query.

""" ) with gr.Row(equal_height=True): with gr.Column(scale=1): schema_in = gr.Code( label="Database schema (DDL)", language="sql", value=EXAMPLE_SCHEMA, lines=14, ) with gr.Column(scale=1): question_in = gr.Textbox( label="Question", placeholder="Ask about the data in plain English...", value=EXAMPLE_QUESTION, lines=3, ) run_btn = gr.Button("Generate SQL", variant="primary") sql_out = gr.Code( label="Generated SQL", language="sql", lines=10, ) gr.HTML( """
Qwen2.5-Coder-7B, QLoRA fine-tune on Spider / BIRD / sql-create-context. Research demo. Verify generated SQL before running it against real data.
""" ) run_btn.click(fn=generate_sql, inputs=[schema_in, question_in], outputs=sql_out) question_in.submit(fn=generate_sql, inputs=[schema_in, question_in], outputs=sql_out) if __name__ == "__main__": demo.queue().launch()