"""Gradio ZeroGPU Space serving the fhirsql-reasoning-sql-adapters model. Reproduces the exact prompt/loading contract documented in the adapters repo's MODEL_CARD.md and C:\\dev\\fhirsql-phase2\\sft_train.ipynb: a plan-then-SQL system prompt (SYSTEM_PROMPT_TEMPLATE / build_messages), 4-bit NF4 quantized loading, the `sft/seed_42/best` checkpoint (the model card explicitly recommends `sft/`, not `rl/` -- RL did not improve on its SFT starting point for this task), and extraction of the SQL from the model's fenced ```sql block (extract_sql_from_completion), not the raw completion. """ import re import time import spaces import gradio as gr import torch from transformers import AutoModelForCausalLM, AutoTokenizer, BitsAndBytesConfig from peft import PeftModel BASE_MODEL = "Qwen/Qwen2.5-Coder-14B-Instruct" ADAPTER = "adelelsayed1991/fhirsql-reasoning-sql-adapters" # MODEL_CARD.md: "sft/ -- supervised fine-tuning only... Use these." RL # checkpoints are published for reproducibility only; they do not improve # on the SFT starting point for this task (PAPER.md Section 5.3). ADAPTER_SUBFOLDER = "sft/seed_42/best" ABSTENTION_TOKEN = "UNANSWERABLE" # Matches sft_train.ipynb's CFG['max_new_tokens_eval'] (cell 5) -- the value the # published exec-match figures in MODEL_CARD.md were actually measured with. The # completion is plan-JSON *then* the fenced SQL block, so a shorter budget risks # truncating before the fence ever appears. MAX_NEW_TOKENS = 2048 SYSTEM_PROMPT_TEMPLATE = """You are a clinical data analyst who translates natural-language hospital \ questions into DuckDB SQL, run against the schema below, via an explicit query plan first. Output exactly two parts, in this order, and nothing else: 1. A JSON object describing the query plan: which clinical concepts the question refers to (and \ whether each needs a terminology lookup against the schema's `valuesets` table), what \ additional tables must be joined and why, what filters apply, and what the final aggregation \ computes. 2. The compiled SQL statement for that plan, in a fenced code block: ```sql ``` If the question cannot be answered from this schema -- the data it needs genuinely doesn't \ exist here -- the plan should be {{"abstain": true}}, and the fenced SQL block should contain \ exactly the single word: {abstention_token} Do not guess or approximate an answer using unrelated columns when the real field is absent. Schema: {schema_ddl}""" _SQL_FENCE_RE = re.compile(r"```sql\s*\n(.*?)\n```", re.IGNORECASE | re.DOTALL) def extract_sql_from_completion(text: str) -> str | None: """Pull the SQL out of a plan-JSON + fenced ```sql block completion. Takes the LAST matching fence, not the first, mirroring the training notebooks' own extraction logic. Returns None (not a crash, not an empty string) when no fence is found at all. """ matches = _SQL_FENCE_RE.findall(text) if not matches: return None return matches[-1].strip() def load_schema_ddl(schema_path: str) -> str: """Load the DuckDB schema DDL, dropping the changelog header above the first CREATE TABLE.""" text = open(schema_path, encoding="utf-8").read() idx = text.index("CREATE TABLE") return text[idx:].strip() SCHEMA_DDL = load_schema_ddl("schema.sql") SYSTEM_PROMPT = SYSTEM_PROMPT_TEMPLATE.format( abstention_token=ABSTENTION_TOKEN, schema_ddl=SCHEMA_DDL ) tokenizer = AutoTokenizer.from_pretrained(BASE_MODEL) base_model = AutoModelForCausalLM.from_pretrained( BASE_MODEL, dtype=torch.bfloat16, quantization_config=BitsAndBytesConfig( load_in_4bit=True, bnb_4bit_compute_dtype=torch.bfloat16, bnb_4bit_use_double_quant=True, bnb_4bit_quant_type="nf4", ), ) model = PeftModel.from_pretrained( base_model, ADAPTER, subfolder=ADAPTER_SUBFOLDER, torch_device="cpu" ) model.to("cuda") model.eval() print("model type:", type(model)) print("peft_config:", getattr(model, "peft_config", "NO ADAPTER")) print( "active_adapters:", model.active_adapters if hasattr(model, "active_adapters") else "-", ) # This is the GPU time the request is allowed, not a hint: ZeroGPU kills the # task the moment it runs over, and the caller sees only "GPU task aborted". # At 180s that happened intermittently -- short completions finished, longer # ones were killed -- which reads as a flaky Space rather than a budget that # is simply too small. Generation has to cover a 14B model in 4-bit emitting # a JSON plan and then the SQL, so the budget is set to the maximum rather # than trimmed to the average. @spaces.GPU(duration=300) def generate_sql(question: str) -> tuple[str, str]: """Generate a plan+SQL completion for a natural-language question. Returns (extracted_sql, raw_completion). extracted_sql is "UNANSWERABLE" if the model determines the question cannot be answered from the schema, or a literal "[NO SQL FENCE FOUND]" placeholder if the completion did not contain a parseable ```sql block at all (a genuine model failure, surfaced rather than hidden). """ messages = [ {"role": "system", "content": SYSTEM_PROMPT}, {"role": "user", "content": question}, ] prompt_text = tokenizer.apply_chat_template( messages, tokenize=False, add_generation_prompt=True ) inputs = tokenizer(prompt_text, return_tensors="pt", add_special_tokens=False).to("cuda") prompt_len = inputs["input_ids"].shape[1] started = time.perf_counter() output_ids = model.generate( **inputs, max_new_tokens=MAX_NEW_TOKENS, do_sample=False, pad_token_id=tokenizer.eos_token_id, ) elapsed = time.perf_counter() - started new_tokens = output_ids.shape[1] - prompt_len # The decorator's `duration` above has to exceed this, or ZeroGPU kills # the request mid-generation and the caller sees only "GPU task # aborted". Logging it turns that budget into something measured # rather than guessed: read these lines from the Space logs and set # `duration` from the worst observed generation, plus headroom. print( f"[timing] generated {new_tokens} tokens in {elapsed:.1f}s " f"({new_tokens / elapsed:.1f} tok/s), prompt {prompt_len} tokens", flush=True, ) completion = tokenizer.decode( output_ids[0][prompt_len:], skip_special_tokens=True ).strip() extracted = extract_sql_from_completion(completion) if extracted is None: extracted = "[NO SQL FENCE FOUND]" return extracted, completion demo = gr.Interface( fn=generate_sql, inputs=gr.Text(label="Clinical question", lines=2), outputs=[ gr.Text(label="Extracted SQL"), gr.Text(label="Raw completion (plan JSON + fenced SQL)", lines=10), ], title="fhirsql-reasoning-sql", description=( "Translates a natural-language hospital question into a query " "plan and DuckDB SQL against the fhirsql-phase2 schema. Returns " "UNANSWERABLE if the question cannot be answered from that " "schema. sft/seed_42/best checkpoint, per MODEL_CARD.md." ), ) if __name__ == "__main__": demo.launch()