Spaces:
Running on Zero
Running on Zero
Log generation time and token count so the GPU duration can be measured, not guessed
f2d515b verified Download app.py from adelelsayed1991/fhirsql-reasoning-sql-adapters: direct link, hf CLI and curl.
- Browser
- Download file 7.21 kB
-
https://huggingface.co/spaces/adelelsayed1991/fhirsql-reasoning-sql-adapters/resolve/main/app.py
- Command line
-
hf download hf://spaces/adelelsayed1991/fhirsql-reasoning-sql-adapters/app.py
-
curl -L -o app.py https://huggingface.co/spaces/adelelsayed1991/fhirsql-reasoning-sql-adapters/resolve/main/app.py
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. | |
| 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() | |