adelelsayed1991's picture
Log generation time and token count so the GPU duration can be measured, not guessed
f2d515b verified
Raw History Blame Contribute Delete
7.21 kB
"""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()