Spaces:
Running on Zero
Running on Zero
File size: 7,213 Bytes
3434a32 19f0c5b 3434a32 19f0c5b f2d515b 19f0c5b 3434a32 19f0c5b 3434a32 19f0c5b 3434a32 7b67c79 3434a32 19f0c5b 3434a32 19f0c5b 3434a32 19f0c5b 3434a32 19f0c5b a7cc8ca ecd77d9 a7cc8ca 3434a32 19f0c5b 3434a32 7bbc0fc 19f0c5b 3434a32 19f0c5b 3434a32 f2d515b 3434a32 f2d515b 19f0c5b 3434a32 19f0c5b 3434a32 19f0c5b 3434a32 | 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 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 | """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
<the SQL statement>
```
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()
|