Team Ai
Apppublic

adelelsayed1991/fhirsql-reasoning-sql-adapters

sourceHugging Faceupdated 16d agoView on Hugging Face
0likes
app.py185 linesDownload Raw Back to root
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