adelelsayed1991/fhirsql-reasoning-sql-adapters
0
1"""Gradio ZeroGPU Space serving the fhirsql-reasoning-sql-adapters model.2 3Reproduces the exact prompt/loading contract documented in the adapters4repo's MODEL_CARD.md and C:\\dev\\fhirsql-phase2\\sft_train.ipynb: a plan-then-SQL5system prompt (SYSTEM_PROMPT_TEMPLATE / build_messages), 4-bit NF46quantized loading, the `sft/seed_42/best` checkpoint (the model card7explicitly recommends `sft/`, not `rl/` -- RL did not improve on its SFT8starting point for this task), and extraction of the SQL from the model's9fenced ```sql block (extract_sql_from_completion), not the raw completion.10"""11 12import re13import time14 15import spaces16import gradio as gr17import torch18from transformers import AutoModelForCausalLM, AutoTokenizer, BitsAndBytesConfig19from peft import PeftModel20 21BASE_MODEL = "Qwen/Qwen2.5-Coder-14B-Instruct"22ADAPTER = "adelelsayed1991/fhirsql-reasoning-sql-adapters"23# MODEL_CARD.md: "sft/ -- supervised fine-tuning only... Use these." RL24# checkpoints are published for reproducibility only; they do not improve25# on the SFT starting point for this task (PAPER.md Section 5.3).26ADAPTER_SUBFOLDER = "sft/seed_42/best"27ABSTENTION_TOKEN = "UNANSWERABLE"28# Matches sft_train.ipynb's CFG['max_new_tokens_eval'] (cell 5) -- the value the29# published exec-match figures in MODEL_CARD.md were actually measured with. The30# completion is plan-JSON *then* the fenced SQL block, so a shorter budget risks31# truncating before the fence ever appears.32MAX_NEW_TOKENS = 204833 34SYSTEM_PROMPT_TEMPLATE = """You are a clinical data analyst who translates natural-language hospital \35questions into DuckDB SQL, run against the schema below, via an explicit query plan first.36 37Output exactly two parts, in this order, and nothing else:381. A JSON object describing the query plan: which clinical concepts the question refers to (and \39whether each needs a terminology lookup against the schema's `valuesets` table), what \40additional tables must be joined and why, what filters apply, and what the final aggregation \41computes.422. The compiled SQL statement for that plan, in a fenced code block:43```sql44<the SQL statement>45```46 47If the question cannot be answered from this schema -- the data it needs genuinely doesn't \48exist here -- the plan should be {{"abstain": true}}, and the fenced SQL block should contain \49exactly the single word: {abstention_token}50 51Do not guess or approximate an answer using unrelated columns when the real field is absent.52 53Schema:54{schema_ddl}"""55 56_SQL_FENCE_RE = re.compile(r"```sql\s*\n(.*?)\n```", re.IGNORECASE | re.DOTALL)57 58 59def extract_sql_from_completion(text: str) -> str | None:60 """Pull the SQL out of a plan-JSON + fenced ```sql block completion.61 62 Takes the LAST matching fence, not the first, mirroring the training63 notebooks' own extraction logic. Returns None (not a crash, not an64 empty string) when no fence is found at all.65 """66 matches = _SQL_FENCE_RE.findall(text)67 if not matches:68 return None69 return matches[-1].strip()70 71 72def load_schema_ddl(schema_path: str) -> str:73 """Load the DuckDB schema DDL, dropping the changelog header above the first CREATE TABLE."""74 text = open(schema_path, encoding="utf-8").read()75 idx = text.index("CREATE TABLE")76 return text[idx:].strip()77 78 79SCHEMA_DDL = load_schema_ddl("schema.sql")80SYSTEM_PROMPT = SYSTEM_PROMPT_TEMPLATE.format(81 abstention_token=ABSTENTION_TOKEN, schema_ddl=SCHEMA_DDL82)83 84tokenizer = AutoTokenizer.from_pretrained(BASE_MODEL)85base_model = AutoModelForCausalLM.from_pretrained(86 BASE_MODEL,87 dtype=torch.bfloat16,88 quantization_config=BitsAndBytesConfig(89 load_in_4bit=True,90 bnb_4bit_compute_dtype=torch.bfloat16,91 bnb_4bit_use_double_quant=True,92 bnb_4bit_quant_type="nf4",93 ),94)95model = PeftModel.from_pretrained(96 base_model, ADAPTER, subfolder=ADAPTER_SUBFOLDER, torch_device="cpu"97)98model.to("cuda")99model.eval()100 101print("model type:", type(model))102print("peft_config:", getattr(model, "peft_config", "NO ADAPTER"))103print(104 "active_adapters:",105 model.active_adapters if hasattr(model, "active_adapters") else "-",106)107 108 109# This is the GPU time the request is allowed, not a hint: ZeroGPU kills the110# task the moment it runs over, and the caller sees only "GPU task aborted".111# At 180s that happened intermittently -- short completions finished, longer112# ones were killed -- which reads as a flaky Space rather than a budget that113# is simply too small. Generation has to cover a 14B model in 4-bit emitting114# a JSON plan and then the SQL, so the budget is set to the maximum rather115# than trimmed to the average.116@spaces.GPU(duration=300)117def generate_sql(question: str) -> tuple[str, str]:118 """Generate a plan+SQL completion for a natural-language question.119 120 Returns (extracted_sql, raw_completion). extracted_sql is121 "UNANSWERABLE" if the model determines the question cannot be122 answered from the schema, or a literal "[NO SQL FENCE FOUND]"123 placeholder if the completion did not contain a parseable ```sql124 block at all (a genuine model failure, surfaced rather than hidden).125 """126 messages = [127 {"role": "system", "content": SYSTEM_PROMPT},128 {"role": "user", "content": question},129 ]130 prompt_text = tokenizer.apply_chat_template(131 messages, tokenize=False, add_generation_prompt=True132 )133 inputs = tokenizer(prompt_text, return_tensors="pt", add_special_tokens=False).to("cuda")134 prompt_len = inputs["input_ids"].shape[1]135 136 started = time.perf_counter()137 output_ids = model.generate(138 **inputs,139 max_new_tokens=MAX_NEW_TOKENS,140 do_sample=False,141 pad_token_id=tokenizer.eos_token_id,142 )143 elapsed = time.perf_counter() - started144 new_tokens = output_ids.shape[1] - prompt_len145 146 # The decorator's `duration` above has to exceed this, or ZeroGPU kills147 # the request mid-generation and the caller sees only "GPU task148 # aborted". Logging it turns that budget into something measured149 # rather than guessed: read these lines from the Space logs and set150 # `duration` from the worst observed generation, plus headroom.151 print(152 f"[timing] generated {new_tokens} tokens in {elapsed:.1f}s "153 f"({new_tokens / elapsed:.1f} tok/s), prompt {prompt_len} tokens",154 flush=True,155 )156 157 completion = tokenizer.decode(158 output_ids[0][prompt_len:], skip_special_tokens=True159 ).strip()160 161 extracted = extract_sql_from_completion(completion)162 if extracted is None:163 extracted = "[NO SQL FENCE FOUND]"164 return extracted, completion165 166 167demo = gr.Interface(168 fn=generate_sql,169 inputs=gr.Text(label="Clinical question", lines=2),170 outputs=[171 gr.Text(label="Extracted SQL"),172 gr.Text(label="Raw completion (plan JSON + fenced SQL)", lines=10),173 ],174 title="fhirsql-reasoning-sql",175 description=(176 "Translates a natural-language hospital question into a query "177 "plan and DuckDB SQL against the fhirsql-phase2 schema. Returns "178 "UNANSWERABLE if the question cannot be answered from that "179 "schema. sft/seed_42/best checkpoint, per MODEL_CARD.md."180 ),181)182 183if __name__ == "__main__":184 demo.launch()185 