"""STS patient notes + POAF prediction — self-contained (no imports from other local .py files).""" import os import json import tempfile import zipfile from datetime import datetime from typing import Any import numpy as np import pandas as pd import gradio as gr import torch import torch.nn as nn from huggingface_hub import snapshot_download from transformers import ( AutoTokenizer, AutoModel, AutoModelForSequenceClassification, BitsAndBytesConfig, ) from peft import PeftModel # --- Inference model loading (inlined; no run_disc_poaf import) --- def make_input_text(note_text: str) -> str: return f"POAF (0/1). Note:\n{note_text}" def _device_map(): return "cuda:0" if torch.cuda.is_available() else "cpu" def _bitsandbytes_4bit_config(): if not torch.cuda.is_available(): return None return BitsAndBytesConfig( load_in_4bit=True, bnb_4bit_quant_type="nf4", bnb_4bit_use_double_quant=True, bnb_4bit_compute_dtype=torch.bfloat16, ) def load_tokenizer(model_id: str): tok = AutoTokenizer.from_pretrained(model_id, use_fast=True) if tok.pad_token is None: tok.pad_token = tok.eos_token tok.padding_side = "right" tok.truncation_side = "right" return tok def _tokenizer_markers(dir_path: str) -> bool: """Return True if dir looks like a HF tokenizer folder.""" if not os.path.isdir(dir_path): return False for name in ("tokenizer_config.json", "tokenizer.json", "tokenizer.model", "vocab.txt", "special_tokens_map.json"): if os.path.isfile(os.path.join(dir_path, name)): return True return False def load_tokenizer_from_run_dir(run_dir: str) -> Any: """ Load tokenizer only from the fine-tuned snapshot on disk. Training saves tokenizer next to the adapter (often under `lora_adapter/`), not always at repo root. We intentionally do **not** fall back to upstream `model_id` repos (often gated). """ candidates = [ os.path.join(run_dir, "lora_adapter"), run_dir, os.path.join(run_dir, "base"), ] # Optional: deepest checkpoint dir sometimes has tokenizer copies try: for name in sorted(os.listdir(run_dir)): if name.startswith("checkpoint-"): candidates.append(os.path.join(run_dir, name)) except Exception: pass tried = [] for d in candidates: if not _tokenizer_markers(d): continue tried.append(d) try: tok = AutoTokenizer.from_pretrained(d, use_fast=True, local_files_only=True) if tok.pad_token is None: tok.pad_token = tok.eos_token tok.padding_side = "right" tok.truncation_side = "right" return tok except Exception: continue raise RuntimeError( "Could not load tokenizer from your uploaded model snapshot. " f"Tried: {tried or candidates}. " "Ensure tokenizer files (e.g. tokenizer_config.json / tokenizer.json) exist under " "`lora_adapter/` or repo root in `poaf-best-*`, and that those repos are public." ) def load_seqcls_model(model_id: str, tok, use_4bit: bool): qconfig = _bitsandbytes_4bit_config() if use_4bit else None dtype = torch.bfloat16 if torch.cuda.is_available() else torch.float32 model = AutoModelForSequenceClassification.from_pretrained( model_id, num_labels=2, device_map=_device_map(), torch_dtype=dtype, quantization_config=qconfig, ) if hasattr(model, "config"): model.config.pad_token_id = tok.pad_token_id model.config.use_cache = False return model def load_base_model(model_id: str, tok, use_4bit: bool): qconfig = _bitsandbytes_4bit_config() if use_4bit else None dtype = torch.bfloat16 if torch.cuda.is_available() else torch.float32 model = AutoModel.from_pretrained( model_id, device_map=_device_map(), torch_dtype=dtype, quantization_config=qconfig, ) if hasattr(model, "config"): model.config.pad_token_id = tok.pad_token_id model.config.use_cache = False return model def load_model_with_lora( base_model_id: str, run_dir: str, use_lora: bool, use_4bit: bool, use_custom_head: bool = False, cls_hidden_dim: int = 1536, cls_dropout: float = 0.1, pooling: str = "mean", cls_hidden_dim2: int = 768, train_last_n_layers: int = 0, ): if use_lora: adapter_dir = os.path.join(run_dir, "lora_adapter") else: adapter_dir = run_dir # Prefer local tokenizer from the downloaded run directory only (no upstream Hub IDs). tok = load_tokenizer_from_run_dir(run_dir) if use_custom_head: base_path = os.path.join(run_dir, "base") # If the fine-tuned run snapshot contains the base model weights, # load them locally to avoid pulling from gated upstream repos. if os.path.isdir(base_path): base = AutoModel.from_pretrained( base_path, torch_dtype=torch.bfloat16 if torch.cuda.is_available() else torch.float32, device_map=_device_map(), local_files_only=True, ) else: raise RuntimeError( "This checkpoint expects a local `base/` folder inside the model repo snapshot " f"(missing under {run_dir}). Re-upload the run folder including `base/`, or export base weights." ) if use_lora: base = PeftModel.from_pretrained(base, adapter_dir) model = CustomSeqClassifier( base_model=base, num_labels=2, hidden_dim=cls_hidden_dim, dropout=cls_dropout, pooling=pooling, hidden_dim2=cls_hidden_dim2, ) head_file = os.path.join(run_dir, "custom_head.pt") if os.path.isfile(head_file): state = torch.load(head_file, map_location="cpu") model.head.load_state_dict(state, strict=True) dev = next(model.base.parameters()).device dty = next(model.base.parameters()).dtype model.head.to(device=dev, dtype=dty) model.eval() return model, tok if use_lora: # Standard HF seq-cls + LoRA expects the **full** seq-cls base from Hub, then adapters on top. # To avoid gated upstream downloads in duplicated Spaces, train/serve with use_custom_head=true # (base LM under `base/` + `lora_adapter/`). Non-custom-head LoRA cannot load from adapter-only folders. try: base = AutoModelForSequenceClassification.from_pretrained( adapter_dir, num_labels=2, device_map=_device_map(), torch_dtype=torch.bfloat16 if torch.cuda.is_available() else torch.float32, quantization_config=_bitsandbytes_4bit_config() if use_4bit else None, local_files_only=True, ) except Exception as e: raise RuntimeError( "Could not load a full sequence-classification checkpoint from the local snapshot " f"({adapter_dir}). This deployment avoids downloading gated upstream bases like " f"{base_model_id}. Use runs trained with use_custom_head=true and upload `base/` + " "`lora_adapter/` + tokenizer files to poaf-best-*, or set HF_TOKEN and accept upstream access." ) from e model = PeftModel.from_pretrained(base, adapter_dir) else: model = AutoModelForSequenceClassification.from_pretrained( adapter_dir, torch_dtype=torch.bfloat16 if torch.cuda.is_available() else torch.float32, device_map=_device_map(), local_files_only=True, ) model.eval() return model, tok class MeanPooler(nn.Module): def __init__(self): super(MeanPooler, self).__init__() def forward(self, last_hidden_state, attention_mask): mask = attention_mask.unsqueeze(-1).to(last_hidden_state.dtype) total = (last_hidden_state * mask).sum(dim=1) n = mask.sum(dim=1).clamp(min=1e-6) return total / n class CustomClassifierHead(nn.Module): def __init__(self, in_dim, hidden_dim, num_labels, dropout=0.1, hidden_dim2=0): super(CustomClassifierHead, self).__init__() self.dropout = nn.Dropout(dropout) self.one_layer = hidden_dim <= 0 if self.one_layer: self.linear1 = nn.Linear(in_dim, num_labels) self.linear2 = None self.act = None self.linear3 = None self.has_third = False else: self.linear1 = nn.Linear(in_dim, hidden_dim) self.act = nn.GELU() self.linear2 = nn.Linear(hidden_dim, hidden_dim2 if hidden_dim2 > 0 else num_labels) self.has_third = hidden_dim2 > 0 self.linear3 = nn.Linear(hidden_dim2, num_labels) if self.has_third else None def forward(self, x): x = self.dropout(x) if self.one_layer: return self.linear1(x) x = self.linear1(x) x = self.act(x) x = self.dropout(x) x = self.linear2(x) if self.has_third: x = self.act(x) x = self.dropout(x) x = self.linear3(x) return x class CustomSeqClassifier(nn.Module): def __init__(self, base_model, num_labels, hidden_dim, dropout=0.1, pooling="mean", hidden_dim2=0): super(CustomSeqClassifier, self).__init__() self.base = base_model self.num_labels = num_labels self.pooling = pooling self.pool = MeanPooler() if pooling == "mean" else None h = base_model.config.hidden_size self.head = CustomClassifierHead(h, hidden_dim, num_labels, dropout, hidden_dim2=hidden_dim2) self.criterion = nn.CrossEntropyLoss() self.config = getattr(base_model, "config", None) def forward(self, input_ids=None, attention_mask=None, labels=None, **kwargs): out = self.base(input_ids=input_ids, attention_mask=attention_mask, output_hidden_states=False) h = out.last_hidden_state if self.pooling == "mean": pooled = self.pool(h, attention_mask) else: last_idx = (attention_mask.sum(dim=1).long() - 1).clamp(min=0) batch_idx = torch.arange(h.size(0), device=h.device) pooled = h[batch_idx, last_idx] logits = self.head(pooled) loss = None if labels is not None: if labels.dtype != torch.long: labels = labels.long() loss = self.criterion(logits, labels) return {"logits": logits, "loss": loss} PROJECT_DIR = os.path.dirname(os.path.abspath(__file__)) MODEL_RUNS = { "Meditron (best: 1CustomLayer1TransformerBlocks)": { "model_id": "epfl-llm/meditron-7b", "repo_id": "yaminigonuguntla/poaf-best-1", }, "biomistral (best: 3CustomLayer1TransformerBlocks)": { "model_id": "BioMistral/BioMistral-7B", "repo_id": "yaminigonuguntla/poaf-best-2", }, "openbiollm (best: 1CustomLayer1TransformerBlocks)": { "model_id": "aaditya/Llama3-OpenBioLLM-8B", "repo_id": "yaminigonuguntla/poaf-best-3", }, } LOADED = {} def _softmax(logits: np.ndarray) -> np.ndarray: x = logits - np.max(logits, axis=-1, keepdims=True) e = np.exp(x) return e / np.sum(e, axis=-1, keepdims=True) def _get_model(label: str): # Lazy-load once, then reuse in memory. if label in LOADED: return LOADED[label] cfg = MODEL_RUNS[label] token = os.getenv("HF_TOKEN") if token: run_dir = snapshot_download( repo_id=cfg["repo_id"], repo_type="model", token=token, ) else: # Try download without token. This will work if the Space runner can access # the repository (e.g., public repos) or if the snapshot contains everything needed. run_dir = snapshot_download( repo_id=cfg["repo_id"], repo_type="model", ) run_cfg_path = os.path.join(run_dir, "run_config.json") if not os.path.isfile(run_cfg_path): raise FileNotFoundError(f"Missing run_config.json in downloaded repo: {cfg['repo_id']}") with open(run_cfg_path, "r", encoding="utf-8") as f: run_cfg = json.load(f) model, tok = load_model_with_lora( base_model_id=cfg["model_id"], run_dir=run_dir, use_lora=bool(run_cfg.get("use_lora", False)), use_4bit=bool(run_cfg.get("use_4bit", True)), use_custom_head=bool(run_cfg.get("use_custom_head", True)), cls_hidden_dim=int(run_cfg.get("cls_hidden_dim", 1536)), cls_dropout=float(run_cfg.get("cls_dropout", 0.1)), pooling=str(run_cfg.get("pooling", "mean")), cls_hidden_dim2=int(run_cfg.get("cls_hidden_dim2", 0)), train_last_n_layers=int(run_cfg.get("train_last_n_layers", 0)), ) LOADED[label] = (model, tok) return model, tok @torch.inference_mode() def _predict_single(note_text: str, model_label: str): # Shared inference path for both single-model and all-model buttons. model, tok = _get_model(model_label) prompt = make_input_text(str(note_text).strip()) enc = tok(prompt, return_tensors="pt", truncation=True, max_length=2048, padding=False) device = getattr(model, "device", next(model.parameters()).device) input_ids = enc["input_ids"].to(device) attention_mask = enc["attention_mask"].to(device) out = model(input_ids=input_ids, attention_mask=attention_mask) logits = out["logits"] if isinstance(out, dict) else out.logits probs = _softmax(logits.detach().float().cpu().numpy())[0] p0, p1 = float(probs[0]), float(probs[1]) pred = 1 if p1 >= 0.5 else 0 return pred, p0, p1 @torch.inference_mode() def predict_poaf(note_text: str, model_choice: str): if note_text is None or not str(note_text).strip(): return "Please paste a clinical note.", "", "" pred, _p0, p1 = _predict_single(note_text, model_choice) label_text = "POAF = 1 (Yes)" if pred == 1 else "POAF = 0 (No)" prob_text = f"P(POAF=1): {p1:.4f}" return label_text, prob_text, "" @torch.inference_mode() def predict_all_models(note_text: str): if note_text is None or not str(note_text).strip(): return [] rows = [] for label in MODEL_RUNS.keys(): _pred, p0, p1 = _predict_single(note_text, label) rows.append( [ label, round(p1, 4), round(p0, 4), ] ) return rows sts_def = None sts_rows = None def_map = None type_map = None harvest_map = None include_map = None section_map = None label_column = None generated_notes = {} zip_file_path = None SECTION_ORDER_DISPLAY = [ "Demographics", "Surgery Timeline", "Mortality & Outcomes", "Preoperative Status", "Cardiac Risk Factors", "Renal Function", "Cerebrovascular History", "Infectious Disease", "Pulmonary History", "Vascular History", "Prior Interventions", "Acute Cardiac Conditions", "Preoperative Medications", "Cardiac Assessment", "Intraoperative Support", "Surgical Procedure", "Intraoperative Course", "Blood Transfusion", ] def parse_harvest_codes(harvest_string): # Convert "0:No;1:Yes" style strings into a dictionary. if harvest_string is None or (isinstance(harvest_string, float) and np.isnan(harvest_string)): return {} text = str(harvest_string).strip() if not text: return {} codes = {} for raw_pair in text.split(";"): pair = raw_pair.strip() if not pair or ":" not in pair: continue key_raw, val_raw = pair.split(":", 1) key_raw = key_raw.strip() val_raw = val_raw.strip() try: key = int(key_raw) except ValueError: key = key_raw codes[key] = val_raw return codes def build_definition_and_type_maps(sts_def_df): # Build fast lookup maps used during note generation. dmap = {} tmap = {} hmap = {} imap = {} smap = {} for _, row in sts_def_df.iterrows(): col = row.get("STSData columns", None) if pd.isna(col): continue col = str(col).strip() inc = row.get("Include", 1) try: inc_flag = int(inc) == 1 except Exception: inc_flag = True imap[col] = inc_flag sec = row.get("Section", None) or row.get("Section Name", None) if pd.notna(sec): smap[col] = str(sec).strip() else: smap[col] = "Other" defn = row.get("Definition", None) if pd.notna(defn): base_def = str(defn).strip() else: base_def = col.replace("_", " ").replace("+", " ").strip() dmap[col] = base_def t_hint = row.get("TypeHint", None) if pd.notna(t_hint): tmap[col] = str(t_hint).strip().lower() else: tmap[col] = None hcodes = row.get("Harvest Codes", None) or row.get("HarvestCodes", None) if pd.notna(hcodes): hmap[col] = str(hcodes) else: hmap[col] = "" return dmap, tmap, hmap, imap, smap def format_value(val): if pd.isna(val): return None if isinstance(val, (pd.Timestamp, datetime)): return val.strftime("%B %d, %Y") if isinstance(val, float): if val > 100: return f"{val:.1f}" return f"{val:.2f}" return str(val).strip() def detect_label_column(df): # Prefer Oth_Afib, then apply a simple fallback heuristic. if "Oth_Afib" in df.columns: return "Oth_Afib" for col in reversed(df.columns[-5:]): col_lower = str(col).lower() if any(k in col_lower for k in ["label", "outcome", "target", "class", "af", "afib"]): return col nunique = df[col].nunique(dropna=True) if nunique <= 2: unique_vals = set(df[col].dropna().unique()) if unique_vals in [{0, 1}, {"No", "Yes"}, {0}, {1}]: return col return None def generate_field_text(col_name, value, def_map, type_map, harvest_map): if pd.isna(value): return None base_def = def_map.get(col_name, None) if not base_def: return None t_hint = (type_map.get(col_name) or "").lower() h_str = harvest_map.get(col_name, "") codes = parse_harvest_codes(h_str) if isinstance(value, (np.integer, np.floating)): v_norm = value.item() else: v_norm = value v_str = str(v_norm).strip() if col_name == "Death_Date": if v_str.lower() == "alive": return "The patient is alive." f = format_value(v_norm) return f"{base_def} {f}." if f else None if col_name in ["CVA_When", "IABP_When", "IABP_Indication"]: if v_str and v_str.lower() != "nan": return f"{base_def} {v_str}." return None if col_name == "Introp DEX or nDEX": if v_str.lower() == "dex": return base_def if base_def.endswith(".") else base_def + "." return None if t_hint == "binary": try: v_num = float(v_str) except ValueError: v_num = None if v_num is not None: if v_num == 0.0: return None if v_num == 1.0: return base_def if base_def.endswith(".") else base_def + "." if v_str.lower() in ["yes", "y", "true"]: return base_def if base_def.endswith(".") else base_def + "." if v_str.lower() in ["no", "n", "false"]: return None return None if t_hint == "date": if v_str.lower() == "alive": return "The patient is alive." f = format_value(v_norm) return f"{base_def} {f}." if f else None if t_hint == "numeric": try: if float(v_str) == 0.0: return None except ValueError: pass if codes: key = v_str try: key_num = int(v_str) if key_num in codes: key = key_num except ValueError: pass right_side = codes.get(key, None) if right_side is None: f = format_value(v_norm) return f"{base_def} {f}." if f else None return f"{base_def} {right_side}." f = format_value(v_norm) return f"{base_def} {f}." if f else None f = format_value(v_norm) return f"{base_def} {f}." if f else None def generate_note_for_row(row, def_map, type_map, harvest_map, include_map, section_map, label_col=None): # Create one note by grouping generated sentences by section. sections = {} for col in row.index: col_clean = str(col).strip() val = row[col] if label_col and col_clean == label_col: continue if pd.isna(val): continue if not include_map.get(col_clean, True): continue definition = def_map.get(col_clean, None) if definition is None: continue section = section_map.get(col_clean, "Other") sentence = generate_field_text( col_name=col_clean, value=val, def_map=def_map, type_map=type_map, harvest_map=harvest_map, ) if not sentence: continue sections.setdefault(section, []).append(sentence) note_lines = [] note_lines.append("PATIENT MEDICAL NOTE") note_lines.append("=" * 80) note_lines.append("") seen_sections = set() for section in SECTION_ORDER_DISPLAY: if section in sections: note_lines.append(f"\n--- {section} ---") for s in sections[section]: note_lines.append(f"- {s}") seen_sections.add(section) for section in sorted(sections.keys()): if section not in seen_sections: note_lines.append(f"\n--- {section} ---") for s in sections[section]: note_lines.append(f"- {s}") note_lines.append("\n" + "=" * 80) return "\n".join(note_lines) def _read_tabular_file(path: str) -> pd.DataFrame: """Read CSV or Excel.""" if not path: raise ValueError("Empty path") lower = str(path).lower() if lower.endswith(".csv"): return pd.read_csv(path, encoding="latin-1") if lower.endswith((".xlsx", ".xls")): return pd.read_excel(path) raise ValueError( "Unsupported file type. Use a .csv file or an Excel workbook (.xlsx or .xls)." ) def process_files(def_file, patient_file): # Load both inputs and prepare all lookup maps. global sts_def, sts_rows, def_map, type_map, harvest_map, include_map, section_map, label_column, generated_notes try: if def_file is None: return "❌ Error: Please upload the definition file (CSV or Excel)" sts_def = _read_tabular_file(def_file.name) if patient_file is None: return "❌ Error: Please upload the patient data file (CSV or Excel)" sts_rows = _read_tabular_file(patient_file.name) def_map, type_map, harvest_map, include_map, section_map = build_definition_and_type_maps(sts_def) label_column = detect_label_column(sts_rows) if len(def_map) == 0: return "❌ Error: No definitions found in the definition file" if len(sts_rows) == 0: return "❌ Error: No patient rows found in the patient file" label_info = ( f"\n • Label column detected (hidden in notes): {label_column}" if label_column else "\n • No label column detected" ) msg = f""" ✅ Files loaded successfully! 📊 Dataset Info: • Total Patients: {len(sts_rows)} rows • Columns: {len(sts_rows.columns)} • Definitions: {len(def_map)} mapped{label_info} 🎯 Next Step: Click "Generate All Notes" to create documents for all {len(sts_rows)} patients! """ return msg except Exception as e: return f"❌ Error loading files: {str(e)}" def columns_missing_from_definitions(): """Columns present in patient data but absent from definitions.""" global sts_rows, def_map, label_column if sts_rows is None or def_map is None: return [] defined = set(def_map.keys()) label = str(label_column).strip() if label_column else None missing = [] for c in sts_rows.columns: cname = str(c).strip() if label and cname == label: continue if cname not in defined: missing.append(cname) return sorted(missing) def generate_all_notes(): # Iterate over rows and store one rendered note per patient id. global sts_rows, def_map, type_map, harvest_map, include_map, section_map, label_column, generated_notes try: if sts_rows is None or def_map is None: return "❌ Error: Please load files first (use 'Load Files' tab)" generated_notes = {} total = len(sts_rows) for idx, row in sts_rows.iterrows(): try: note = generate_note_for_row( row=row, def_map=def_map, type_map=type_map, harvest_map=harvest_map, include_map=include_map, section_map=section_map, label_col=label_column, ) patient_id = f"Patient_{idx + 1:04d}" generated_notes[patient_id] = note except Exception as e: generated_notes[f"Patient_{idx + 1:04d}"] = f"Error generating note: {str(e)}" undefined_cols = columns_missing_from_definitions() warning_block = "" if undefined_cols: listed = ", ".join(undefined_cols) warning_block = f""" ⚠️ Warning: The following columns have no matching definitions in the definition file (they were not used when building the notes): {listed} """ status_msg = f""" ✅ All notes generated successfully! 📄 Generated Files: • Total patients: {total} • Notes created: {len(generated_notes)} • Ready for download {warning_block} 📥 Download Options: 1. Click "Download All as ZIP" to get all files 2. Or preview individual notes below 3. Or download Training CSV (text + label) for modeling """ return status_msg except Exception as e: return f"❌ Error: {str(e)}" def create_zip_file(): # Package all generated note text files into a single ZIP. global generated_notes, zip_file_path try: if not generated_notes: return None temp_dir = tempfile.gettempdir() zip_file_path = os.path.join(temp_dir, "STS_Patient_Notes.zip") with zipfile.ZipFile(zip_file_path, "w", zipfile.ZIP_DEFLATED) as zip_file: for patient_id, note in generated_notes.items(): filename = f"{patient_id}_Note.txt" zip_file.writestr(filename, note) return zip_file_path except Exception as e: print(f"Error creating ZIP: {str(e)}") return None def get_patient_list(): global generated_notes if not generated_notes: return "No notes generated yet. Click 'Generate All Notes' first." patient_list = "📋 Generated Patient Notes:\n\n" for i, (patient_id, note) in enumerate(generated_notes.items(), 1): patient_list += f"{i}. **{patient_id}**\n" return patient_list def export_training_data(): # Export notes + labels for model training/evaluation. global generated_notes, sts_rows, label_column if not generated_notes or sts_rows is None or label_column is None: return None num_patients = len(sts_rows) patient_ids = list(range(1, num_patients + 1)) note_texts = [] for i in range(num_patients): pid_key = f"Patient_{i + 1:04d}" note = generated_notes.get(pid_key, "") note_texts.append(note) labels = sts_rows[label_column].astype(int).tolist() df_train = pd.DataFrame( { "patient_id": patient_ids, "text": note_texts, "label": labels, } ) out_path = os.path.join(tempfile.gettempdir(), "poaf_training_data.csv") df_train.to_csv(out_path, index=False) return out_path def preview_patient_note(patient_num): global generated_notes if not generated_notes: return "No notes generated yet. Generate notes first." try: idx = int(patient_num) - 1 patient_ids = list(generated_notes.keys()) if idx < 0 or idx >= len(patient_ids): return f"Patient {int(patient_num)} not found. Available: 1-{len(patient_ids)}" return generated_notes[patient_ids[idx]] except Exception as e: return f"Error: {str(e)}" def create_app_with_predict(): """Build the Gradio UI.""" with gr.Blocks() as demo: gr.Markdown("# 📋 STS Patient Note Generator + POAF Predictor") with gr.Tab("📁 Load Files"): gr.Markdown( """ ### Upload Your Data Files **Required:** - **Definition file** (CSV or Excel `.xlsx` / `.xls`): e.g. STSData_to_SeqNo_Definition Expected columns include: `STSData columns`, `SeqNo`, `Definition`, `Include`, `Section`, `TypeHint`, `Harvest Codes` For Excel, the **first sheet** is read. - **Patient file** (CSV or Excel `.xlsx` / `.xls`): one row per patient Column names should match the definition file. Label/outcome column (e.g., `Oth_Afib`) is detected automatically. """ ) _tabular_types = [".csv", ".xlsx", ".xls"] with gr.Row(): def_file = gr.File( label="📄 Definition file (CSV or Excel)", file_types=_tabular_types, ) patient_file = gr.File( label="📊 Patient file (CSV or Excel)", file_types=_tabular_types, ) load_btn = gr.Button("🔄 Load Files", variant="primary") status_msg = gr.Textbox( label="Status", interactive=False, lines=12, ) load_btn.click( fn=process_files, inputs=[def_file, patient_file], outputs=status_msg, ) with gr.Tab("⚙️ Generate All Notes"): gr.Markdown( """ ### Generate Medical Notes for All Patients This will create a professional medical narrative note for every patient in your dataset. Each patient gets their own individual text file. """ ) generate_btn = gr.Button("🚀 Generate All Notes", variant="primary") with gr.Row(): download_btn = gr.Button("📥 Download All as ZIP", variant="stop") zip_file = gr.File(label="ZIP File (Ready to Download)") train_btn = gr.Button("📥 Download Training CSV", variant="secondary") train_file = gr.File(label="Training CSV (text + label)") progress_msg = gr.Textbox( label="Progress", interactive=False, lines=8, ) patient_list = gr.Markdown( label="Patient Files", value="No notes generated yet.", ) generate_btn.click( fn=generate_all_notes, inputs=[], outputs=progress_msg, ).then( fn=get_patient_list, inputs=[], outputs=patient_list, ) download_btn.click( fn=create_zip_file, inputs=[], outputs=zip_file, ) train_btn.click( fn=export_training_data, inputs=[], outputs=train_file, ) with gr.Tab("👁️ Preview Notes"): gr.Markdown( """ ### Preview Generated Notes View individual patient notes before downloading. """ ) with gr.Row(): patient_num = gr.Number( label="Patient Number", value=1, precision=0, minimum=1, info="Enter patient number (1, 2, 3, etc.)", ) preview_btn = gr.Button("📖 Preview", variant="secondary") preview_text = gr.Textbox( label="Patient Note", lines=25, max_lines=30, interactive=False, ) preview_btn.click( fn=preview_patient_note, inputs=patient_num, outputs=preview_text, ) with gr.Tab("🩺 Predict POAF"): gr.Markdown("## POAF Prediction from Clinical Notes") gr.Markdown( """ Paste a clinical note (generated from the other tabs or typed manually) and: - Run a **single-model** prediction, or - Run **all 3 models** and compare probabilities in the table below. """ ) model_choice = gr.Dropdown( choices=list(MODEL_RUNS.keys()), value=list(MODEL_RUNS.keys())[0], label="Model for single prediction", ) note_text = gr.Textbox( label="Clinical note", lines=16, placeholder="Paste a full clinical note here (e.g., from the generated notes tab)...", ) predict_btn = gr.Button( "🔎 Predict (Selected Model)", variant="primary" ) pred_label = gr.Textbox( label="Single-model prediction", lines=1, ) pred_probs = gr.Textbox( label="P(POAF=1)", lines=1, ) pred_rule = gr.Textbox( label="", lines=1, visible=False, ) predict_all_btn = gr.Button( "🧠 Predict (All 3 Models)", variant="secondary" ) all_table = gr.Dataframe( headers=["Model", "P(POAF=1)", "P(POAF=0)"], datatype=["str", "number", "number"], label="All-model predictions", interactive=False, ) predict_btn.click( fn=predict_poaf, inputs=[note_text, model_choice], outputs=[pred_label, pred_probs, pred_rule], ) predict_all_btn.click( fn=predict_all_models, inputs=[note_text], outputs=[all_table], ) with gr.Tab("ℹ️ About"): gr.Markdown( """ ## How It Works ### Batch Processing - Processes all patients automatically - Creates individual .txt files for each patient - Names files as Patient_0001_Note.txt, Patient_0002_Note.txt, etc. - All files packaged in single ZIP download ### Type Handling - **Binary** - 0 → no sentence - 1 → Definition only (no 'Yes/No', no value) - **Date** - "alive" → "The patient is alive." - otherwise → "Definition + formatted date" - **Numeric** - With Harvest Codes → if value ≠ 0, "Definition + right-side explanation" - Without Harvest Codes → if value ≠ 0, "Definition + numeric value" - **Include**: Only fields with Include = 1 participate in note generation - **Section**: Fields are grouped into clinical sections for readability ### Label Handling - Preferred label column: **Oth_Afib** (0/1 for POAF) - Label is **not** printed in the narrative text - Notes are exported with labels only in the Training CSV ### Output Format Generated notes are: - Organized by clinical sections - Formatted as bullet points - Ready for training LLMs and baselines via Training CSV """ ) return demo if __name__ == "__main__": app = create_app_with_predict() # Spaces need to listen on all interfaces so the proxy can verify the app is up. env_port = (os.getenv("GRADIO_SERVER_PORT") or os.getenv("PORT") or "").strip() port = int(env_port) if env_port else None app.launch(server_name="0.0.0.0", server_port=port, share=False)