"""Train and evaluate discriminative POAF models on note text.""" import os import re import gc import json import time import argparse import numpy as np import pandas as pd import torch import torch.nn as nn from torch.utils.data import Dataset as TorchDataset from torch.utils.data import DataLoader from sklearn.metrics import ( accuracy_score, precision_score, recall_score, f1_score, confusion_matrix, roc_auc_score, average_precision_score, ) from transformers import ( AutoTokenizer, AutoModel, AutoModelForSequenceClassification, BitsAndBytesConfig, TrainingArguments, Trainer, set_seed, ) try: from transformers import EarlyStoppingCallback except ImportError: EarlyStoppingCallback = None from peft import LoraConfig, get_peft_model, PeftModel, prepare_model_for_kbit_training def short_model_name(model_id: str) -> str: """Convert model id into a filesystem-safe folder name.""" return model_id.replace("/", "__").replace(":", "__") def short_run_suffix(run_suffix: str) -> str: """Normalize legacy run suffix aliases to one canonical format.""" if not run_suffix: return "" s = run_suffix.strip() if "CustomLayer" in s and "TransformerBlocks" in s: return s m = re.match(r"^([123])layer_l([0-3])(.*)$", s) if m: cl, tb, rest = m.group(1), m.group(2), m.group(3) base = f"{cl}CustomLayer{tb}TransformerBlocks" return base + rest m = re.match(r"^([123])l([0-3])t(.*)$", s) if m: cl, tb, rest = m.group(1), m.group(2), m.group(3) base = f"{cl}CustomLayer{tb}TransformerBlocks" return base + rest return s def cleanup_cuda(): # Helps keep GPU memory stable across train/eval loops. gc.collect() if torch.cuda.is_available(): torch.cuda.empty_cache() torch.cuda.synchronize() def atomic_write_csv(df: pd.DataFrame, path: str): """Write CSV via temp file + rename to avoid partial writes.""" tmp = path + ".tmp" df.to_csv(tmp, index=False) try: os.replace(tmp, path) except PermissionError as e: try: if os.path.exists(tmp): os.remove(tmp) except Exception: pass raise PermissionError( f"Can't write '{path}' — likely open in Excel or another app. Close it and rerun." ) from e def make_input_text(note_text: str) -> str: # Keep one prompt format across training and inference. return f"POAF (0/1). Note:\n{note_text}" def load_poaf_csv(path: str) -> pd.DataFrame: """Load split file with patient_id, text, label.""" df = pd.read_csv(path) need = {"patient_id", "text", "label"} if need - set(df.columns): raise ValueError(f"Need columns {need}; got {list(df.columns)}") df["patient_id"] = df["patient_id"].astype(int) df["text"] = df["text"].astype(str) df["label"] = df["label"].astype(int) return df[["patient_id", "text", "label"]] def load_from_splits_dir(splits_dir: str): """Returns (train_df, val_df, test_df) from train_notes.csv, val_notes.csv, test_notes.csv.""" names = ("train_notes.csv", "val_notes.csv", "test_notes.csv") paths = [os.path.join(splits_dir, n) for n in names] for p in paths: if not os.path.exists(p): raise FileNotFoundError(f"Missing: {p}") return tuple(load_poaf_csv(p) for p in paths) def _device_map(): """Use cuda:0 when available, otherwise cpu.""" 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): # Right padding works better here with pooled note embeddings. 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 load_seqcls_model(model_id: str, tok, use_4bit: bool): """Load HF sequence classifier (2 labels), optional 4-bit.""" 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): """Load base LM backbone (no classifier head).""" 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, ): """Load model for eval: with LoRA from run_dir/lora_adapter, or without LoRA from run_dir (full model or base + custom_head.pt).""" if use_lora: adapter_dir = os.path.join(run_dir, "lora_adapter") else: adapter_dir = run_dir tok = load_tokenizer(base_model_id) if use_custom_head: base_path = os.path.join(run_dir, "base") if train_last_n_layers > 0 and 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(), ) else: base = load_base_model(base_model_id, tok, use_4bit=use_4bit) 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: base = load_seqcls_model(base_model_id, tok, use_4bit=use_4bit) 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(), ) model.eval() return model, tok class MeanPooler(nn.Module): """Masked mean pooling.""" 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): """Classifier head: 1, 2, or 3 layers.""" 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): """Wrap base LM + pooling + classifier head.""" 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 gradient_checkpointing_enable(self, **kwargs): if hasattr(self.base, "gradient_checkpointing_enable"): self.base.gradient_checkpointing_enable(**kwargs) def gradient_checkpointing_disable(self, **kwargs): if hasattr(self.base, "gradient_checkpointing_disable"): self.base.gradient_checkpointing_disable(**kwargs) def forward(self, input_ids=None, attention_mask=None, labels=None, **kwargs): """Forward pass; returns logits and optional loss.""" 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} class EncodedTorchDataset(TorchDataset): def __init__(self, encodings, labels, patient_ids): self.encodings = encodings self.labels = labels self.patient_ids = patient_ids def __len__(self): return len(self.labels) def __getitem__(self, idx): return { "input_ids": torch.tensor(self.encodings["input_ids"][idx], dtype=torch.long), "attention_mask": torch.tensor(self.encodings["attention_mask"][idx], dtype=torch.long), "labels": torch.tensor(self.labels[idx], dtype=torch.long), "patient_id": torch.tensor(self.patient_ids[idx], dtype=torch.long), } def encode_dataframe(df: pd.DataFrame, tok, max_length: int) -> EncodedTorchDataset: """Tokenize note text with POAF prompt prefix.""" texts = [make_input_text(t) for t in df["text"].tolist()] enc = tok(texts, truncation=True, max_length=max_length, padding="max_length") labels = df["label"].astype(int).tolist() pids = df["patient_id"].astype(int).tolist() return EncodedTorchDataset(enc, labels, pids) def probs_from_logits(logits: np.ndarray) -> np.ndarray: """Softmax logits and return P(class=1).""" logits = np.asarray(logits) logits = logits - logits.max(axis=1, keepdims=True) e = np.exp(logits) return (e / e.sum(axis=1, keepdims=True))[:, 1] def compute_metrics(y_true, y_pred, y_score=None): """Compute metrics at threshold 0.5.""" y_true = np.asarray(y_true, dtype=int) y_pred = np.asarray(y_pred, dtype=int) cm = confusion_matrix(y_true, y_pred, labels=[0, 1]) r1 = float(recall_score(y_true, y_pred, pos_label=1, zero_division=0)) r0 = float(recall_score(y_true, y_pred, pos_label=0, zero_division=0)) out = { "threshold": 0.5, "f1_label1_POAF": float(f1_score(y_true, y_pred, pos_label=1, zero_division=0)), "roc_auc": None, "accuracy": float(accuracy_score(y_true, y_pred)), "recall_label1_POAF": r1, "precision_label1_POAF": float(precision_score(y_true, y_pred, pos_label=1, zero_division=0)), "sensitivity": r1, "specificity": r0, "pr_auc": None, "precision_label0_NoPOAF": float(precision_score(y_true, y_pred, pos_label=0, zero_division=0)), "recall_label0_NoPOAF": r0, "f1_label0_NoPOAF": float(f1_score(y_true, y_pred, pos_label=0, zero_division=0)), "confusion_matrix": cm.tolist(), } if y_score is not None: y_score = np.asarray(y_score, dtype=float) try: out["roc_auc"] = float(roc_auc_score(y_true, y_score)) except Exception: pass try: out["pr_auc"] = float(average_precision_score(y_true, y_score)) except Exception: pass return out @torch.inference_mode() def eval_model_seqcls(model, tok, df: pd.DataFrame, max_length: int, batch_size: int = 8): """Run model on df in batches. Returns (metrics_dict, list of {patient_id, label, pred, p1}).""" ds = encode_dataframe(df, tok, max_length=max_length) loader = DataLoader(ds, batch_size=batch_size, shuffle=False) device = getattr(model, "device", next(model.parameters()).device) logits_list = [] labels_list = [] pids_list = [] for batch in loader: inp = batch["input_ids"].to(device) mask = batch["attention_mask"].to(device) out = model(input_ids=inp, attention_mask=mask) logits = (out["logits"] if isinstance(out, dict) else out.logits).detach().float().cpu().numpy() logits_list.append(logits) labels_list.append(batch["labels"].cpu().numpy()) pids_list.append(batch["patient_id"].cpu().numpy()) logits = np.concatenate(logits_list, axis=0) y_true = np.concatenate(labels_list, axis=0).astype(int) pids = np.concatenate(pids_list, axis=0).astype(int) p1 = probs_from_logits(logits) y_pred = (p1 >= 0.5).astype(int) metrics = compute_metrics(y_true, y_pred, y_score=p1) rows = [ {"patient_id": int(pid), "label": int(yt), "pred": int(yp), "p1": float(prob)} for pid, yt, yp, prob in zip(pids, y_true, y_pred, p1) ] cleanup_cuda() return metrics, rows LORA_TARGETS = ["q_proj", "k_proj", "v_proj", "o_proj"] CLASSIFIER_NAMES = ("score", "classifier", "classification_head", "lm_head") def _get_transformer_layers(module): """Locate transformer layers in common model layouts.""" if hasattr(module, "base_model"): module = module.base_model if hasattr(module, "layers"): return module.layers if hasattr(module, "model") and hasattr(module.model, "layers"): return module.model.layers if hasattr(module, "model") and hasattr(module.model, "layer"): return module.model.layer if hasattr(module, "encoder") and hasattr(module.encoder, "layer"): return module.encoder.layer return None def unfreeze_last_n_layers(module, n: int): """Unfreeze only the last n transformer blocks.""" if n <= 0: return layers = _get_transformer_layers(module) if layers is None: import warnings warnings.warn("Could not find transformer layers for unfreeze_last_n_layers; skipping.") return total = len(layers) n = min(n, total) for i, layer in enumerate(layers): unfreeze = i >= total - n for param in layer.parameters(): if param.dtype.is_floating_point or param.dtype.is_complex: param.requires_grad = unfreeze def train_lora_finetune( model_id: str, train_df: pd.DataFrame, val_df: pd.DataFrame, out_dir: str, use_4bit: bool, max_length: int, epochs: int, batch_size: int = 1, gradient_accumulation_steps: int = 16, learning_rate: float = 2e-4, use_custom_head: bool = False, cls_hidden_dim: int = 1536, cls_dropout: float = 0.1, pooling: str = "mean", cls_hidden_dim2: int = 768, use_lora: bool = True, train_last_n_layers: int = 0, lora_r: int = 16, lora_alpha: int = 32, lora_dropout: float = 0.05, early_stop_patience: int = 0, ): """Train with or without LoRA. With LoRA: save adapter under out_dir. Without: freeze base (or all but last N layers), train head; save head or full model to out_dir.""" tok = load_tokenizer(model_id) if use_lora: if use_custom_head: base = load_base_model(model_id, tok, use_4bit=use_4bit) if use_4bit: base = prepare_model_for_kbit_training(base) base.enable_input_require_grads() lora_cfg = LoraConfig( r=lora_r, lora_alpha=lora_alpha, lora_dropout=lora_dropout, target_modules=LORA_TARGETS, task_type="FEATURE_EXTRACTION", bias="none", ) base = get_peft_model(base, lora_cfg) model = CustomSeqClassifier( base_model=base, num_labels=2, hidden_dim=cls_hidden_dim, dropout=cls_dropout, pooling=pooling, hidden_dim2=cls_hidden_dim2, ) else: model = load_seqcls_model(model_id, tok, use_4bit=use_4bit) if use_4bit: model = prepare_model_for_kbit_training(model) lora_cfg = LoraConfig( r=lora_r, lora_alpha=lora_alpha, lora_dropout=lora_dropout, target_modules=LORA_TARGETS, task_type="SEQ_CLS", bias="none", ) model = get_peft_model(model, lora_cfg) else: # No-LoRA mode for ablation runs. if use_custom_head: base = load_base_model(model_id, tok, use_4bit=use_4bit) if use_4bit: base = prepare_model_for_kbit_training(base) for param in base.parameters(): param.requires_grad = False model = CustomSeqClassifier( base_model=base, num_labels=2, hidden_dim=cls_hidden_dim, dropout=cls_dropout, pooling=pooling, hidden_dim2=cls_hidden_dim2, ) else: model = load_seqcls_model(model_id, tok, use_4bit=use_4bit) if use_4bit: model = prepare_model_for_kbit_training(model) for name, param in model.named_parameters(): param.requires_grad = any(c in name for c in CLASSIFIER_NAMES) if train_last_n_layers > 0 and hasattr(model, "model"): unfreeze_last_n_layers(model.model, train_last_n_layers) if not use_lora and train_last_n_layers > 0 and use_custom_head: unfreeze_last_n_layers(model.base, train_last_n_layers) train_ds = encode_dataframe(train_df, tok, max_length=max_length) val_ds = encode_dataframe(val_df, tok, max_length=max_length) use_early_stop = early_stop_patience > 0 and EarlyStoppingCallback is not None if early_stop_patience > 0 and EarlyStoppingCallback is None: import warnings warnings.warn("EarlyStoppingCallback not available; early_stop_patience ignored.") training_args = TrainingArguments( output_dir=out_dir, per_device_train_batch_size=batch_size, per_device_eval_batch_size=batch_size, gradient_accumulation_steps=gradient_accumulation_steps, learning_rate=learning_rate, num_train_epochs=float(epochs), eval_strategy="steps", eval_steps=200, save_strategy="steps", save_steps=200, save_total_limit=2, logging_steps=25, report_to="none", bf16=torch.cuda.is_available(), fp16=False, gradient_checkpointing=True, remove_unused_columns=False, load_best_model_at_end=use_early_stop, metric_for_best_model="eval_loss" if use_early_stop else None, greater_is_better=False if use_early_stop else None, dataloader_pin_memory=torch.cuda.is_available(), # avoid "no accelerator found" warning on CPU ) callbacks = [] if use_early_stop: callbacks.append(EarlyStoppingCallback(early_stopping_patience=early_stop_patience)) trainer = Trainer( model=model, args=training_args, train_dataset=train_ds, eval_dataset=val_ds, processing_class=tok, callbacks=callbacks, ) trainer.train() os.makedirs(out_dir, exist_ok=True) if use_lora: if use_custom_head: trainer.model.base.save_pretrained(out_dir) head_path = os.path.join(os.path.dirname(out_dir), "custom_head.pt") torch.save(trainer.model.head.state_dict(), head_path) else: trainer.model.save_pretrained(out_dir) else: if use_custom_head: head_path = os.path.join(out_dir, "custom_head.pt") torch.save(trainer.model.head.state_dict(), head_path) if train_last_n_layers > 0: base_dir = os.path.join(out_dir, "base") os.makedirs(base_dir, exist_ok=True) trainer.model.base.save_pretrained(base_dir) else: trainer.model.save_pretrained(out_dir) tok.save_pretrained(out_dir) del trainer, model, tok cleanup_cuda() return out_dir def run_eval_only(run_dir: str, splits_dir: str, eval_batch: int): """Load an existing run and re-run test evaluation only.""" config_path = os.path.join(run_dir, "run_config.json") if not os.path.isfile(config_path): raise FileNotFoundError(f"No run_config.json in {run_dir}") with open(config_path, "r") as f: cfg = json.load(f) model_id = cfg["model_id"] use_lora = cfg.get("use_lora", True) use_4bit = cfg.get("use_4bit", False) use_custom_head = cfg.get("use_custom_head", False) max_length = int(cfg.get("max_length", 2048)) cls_hidden_dim = int(cfg.get("cls_hidden_dim", 1536)) cls_dropout = float(cfg.get("cls_dropout", 0.1)) pooling = cfg.get("pooling", "mean") cls_hidden_dim2 = int(cfg.get("cls_hidden_dim2", 0)) train_last_n_layers = int(cfg.get("train_last_n_layers", 0)) _, _, test_df = load_from_splits_dir(splits_dir) print(f"[INFO] eval_only: test n={len(test_df)} from {splits_dir}") if use_lora and not os.path.isdir(os.path.join(run_dir, "lora_adapter")): raise FileNotFoundError(f"No lora_adapter in {run_dir}") model, tok = load_model_with_lora( model_id, run_dir, use_lora=use_lora, use_4bit=use_4bit, use_custom_head=use_custom_head, cls_hidden_dim=cls_hidden_dim, cls_dropout=cls_dropout, pooling=pooling, cls_hidden_dim2=cls_hidden_dim2, train_last_n_layers=train_last_n_layers, ) metrics, rows = eval_model_seqcls(model, tok, test_df, max_length=max_length, batch_size=eval_batch) pd.DataFrame(rows).to_csv(os.path.join(run_dir, "finetuned_predictions.csv"), index=False) with open(os.path.join(run_dir, "finetuned_metrics.json"), "w") as f: json.dump(metrics, f, indent=2) del model, tok cleanup_cuda() print(f"[INFO] eval_only: wrote metrics and predictions to {run_dir}") print("Test metrics:", {k: v for k, v in metrics.items() if k != "confusion_matrix"}) abs_run_dir = os.path.abspath(run_dir) model_slug_dir = os.path.dirname(abs_run_dir) out_dir = os.path.dirname(model_slug_dir) inferred_suffix = os.path.basename(model_slug_dir) run_suffix_display = short_run_suffix(inferred_suffix) timings_path = os.path.join(run_dir, "timings.json") train_seconds = None finetuned_eval_seconds = None if os.path.isfile(timings_path): with open(timings_path) as f: t = json.load(f) train_seconds = t.get("train_seconds") finetuned_eval_seconds = t.get("finetuned_eval_seconds") new_row = { "run_suffix": run_suffix_display, "model_id": model_id, "run_dir": run_dir, "accuracy": metrics.get("accuracy"), "sensitivity": metrics.get("sensitivity"), "specificity": metrics.get("specificity"), "roc_auc": metrics.get("roc_auc"), "pr_auc": metrics.get("pr_auc"), "f1_label1_POAF": metrics.get("f1_label1_POAF"), "recall_label1_POAF": metrics.get("recall_label1_POAF"), "precision_label1_POAF": metrics.get("precision_label1_POAF"), "f1_label0_NoPOAF": metrics.get("f1_label0_NoPOAF"), "train_seconds": train_seconds, "finetuned_eval_seconds": finetuned_eval_seconds, } summary_path = os.path.join(out_dir, "summary.csv") new_df = pd.DataFrame([new_row]) if os.path.isfile(summary_path): existing = pd.read_csv(summary_path) if "run_suffix" not in existing.columns: existing["run_suffix"] = "" mask = (existing["run_suffix"].astype(str).isin([inferred_suffix, run_suffix_display])) & (existing["model_id"] == model_id) existing = existing[~mask] new_df = pd.concat([existing, new_df], ignore_index=True) if "run_suffix" in new_df.columns: new_df["run_suffix"] = new_df["run_suffix"].fillna("") if "f1_label1_POAF" in new_df.columns: new_df["f1_label1"] = new_df["f1_label1_POAF"] if "recall_label1_POAF" in new_df.columns: new_df["recall_label1"] = new_df["recall_label1_POAF"] if "precision_label1_POAF" in new_df.columns: new_df["precision_label1"] = new_df["precision_label1_POAF"] desired_cols = [ "run_suffix", "model_id", "run_dir", "f1_label1", "roc_auc", "accuracy", "recall_label1", "precision_label1", ] present_cols = [c for c in desired_cols if c in new_df.columns] new_df = new_df[present_cols] atomic_write_csv(new_df, summary_path) print(f"[INFO] eval_only: updated summary at {summary_path}") def main(): """Train/eval loop, or eval_only path for saved runs.""" ap = argparse.ArgumentParser(description="Discriminative POAF on notes (LoRA + optional custom head).") ap.add_argument("--splits_dir", type=str, default=r"data\splits", help="Dir with train/val/test_notes.csv") ap.add_argument("--out", default="results_disc", help="Output dir for run dirs and summary.csv") ap.add_argument("--seed", type=int, default=42) ap.add_argument("--use_4bit", action="store_true", default=True, help="Load in 4-bit to save VRAM (default: True)") ap.add_argument("--no_4bit", action="store_false", dest="use_4bit", help="Disable 4-bit; use full precision (needs more VRAM)") ap.add_argument("--models", nargs="+", default=None, help="HuggingFace model IDs (required unless --eval_only)") ap.add_argument("--max_length", type=int, default=2048) ap.add_argument("--epochs", type=int, default=3) ap.add_argument("--batch_size", type=int, default=1) ap.add_argument("--grad_accum", type=int, default=16) ap.add_argument("--lr", type=float, default=2e-4) ap.add_argument("--eval_batch", type=int, default=8, help="Batch size for test eval") ap.add_argument("--use_custom_head", action="store_true", help="Use LM + pool + MLP head instead of HF seq-cls") ap.add_argument("--cls_hidden_dim", type=int, default=1536, help="Hidden size of head; 0 = 1-layer head (linear only)") ap.add_argument("--cls_hidden_dim2", type=int, default=768, help="0 = 2-layer head, >0 = 3-layer") ap.add_argument("--cls_dropout", type=float, default=0.1) ap.add_argument("--pooling", type=str, default="mean", choices=["mean", "last"]) ap.add_argument("--no_lora", action="store_true", help="Do not use LoRA: freeze base, train only head (for this week)") ap.add_argument("--train_last_n_layers", type=int, default=0, help="When not using LoRA: unfreeze and train last N transformer layers plus head (0 = only head)") ap.add_argument("--lora_r", type=int, default=16) ap.add_argument("--lora_alpha", type=int, default=32) ap.add_argument("--lora_dropout", type=float, default=0.05) ap.add_argument("--early_stop_patience", type=int, default=3, help="Stop if no val-loss improvement for this many evals (default 3). 0 = off.") ap.add_argument("--run_suffix", type=str, default="") ap.add_argument("--eval_only", action="store_true", help="Skip training; eval existing run") ap.add_argument("--run_dir", type=str, default="", help="Required when --eval_only") args = ap.parse_args() if args.eval_only: # Recompute predictions/metrics for a finished run without retraining. if not args.run_dir or not os.path.isdir(args.run_dir): raise SystemExit("--eval_only requires --run_dir with an existing run directory.") set_seed(args.seed) run_eval_only(args.run_dir, args.splits_dir, args.eval_batch) return if not args.models: raise SystemExit("Pass --models (one or more model IDs) or use --eval_only with --run_dir.") set_seed(args.seed) os.makedirs(args.out, exist_ok=True) train_df, val_df, test_df = load_from_splits_dir(args.splits_dir) print(f"[INFO] Loaded from {args.splits_dir}: train {len(train_df)}, val {len(val_df)}, test {len(test_df)}") print(f"[INFO] Train counts: {train_df['label'].value_counts().to_dict()}") summary_rows = [] t0_total = time.perf_counter() for model_id in args.models: # Each model gets its own run folder and config snapshot. tag = short_model_name(model_id) if args.run_suffix: experiment_dir = os.path.join(args.out, args.run_suffix) run_dir = os.path.join(experiment_dir, tag) else: run_dir = os.path.join(args.out, tag) os.makedirs(run_dir, exist_ok=True) use_lora = not args.no_lora run_config = { "model_id": model_id, "splits_dir": args.splits_dir, "seed": args.seed, "max_length": args.max_length, "epochs": args.epochs, "use_4bit": args.use_4bit, "batch_size": args.batch_size, "grad_accum": args.grad_accum, "lr": args.lr, "use_custom_head": args.use_custom_head, "cls_hidden_dim": args.cls_hidden_dim, "cls_hidden_dim2": args.cls_hidden_dim2, "cls_dropout": args.cls_dropout, "pooling": args.pooling, "use_lora": use_lora, "train_last_n_layers": args.train_last_n_layers, "lora_r": args.lora_r, "lora_alpha": args.lora_alpha, "lora_dropout": args.lora_dropout, "early_stop_patience": args.early_stop_patience, "run_suffix": short_run_suffix(args.run_suffix), } with open(os.path.join(run_dir, "run_config.json"), "w") as f: json.dump(run_config, f, indent=2) times = {} out_dir = os.path.join(run_dir, "lora_adapter") if use_lora else run_dir t0 = time.perf_counter() train_lora_finetune( model_id=model_id, train_df=train_df, val_df=val_df, out_dir=out_dir, use_4bit=args.use_4bit, max_length=args.max_length, epochs=args.epochs, batch_size=args.batch_size, gradient_accumulation_steps=args.grad_accum, learning_rate=args.lr, use_custom_head=args.use_custom_head, cls_hidden_dim=args.cls_hidden_dim, cls_dropout=args.cls_dropout, pooling=args.pooling, cls_hidden_dim2=args.cls_hidden_dim2, use_lora=use_lora, train_last_n_layers=args.train_last_n_layers, lora_r=args.lora_r, lora_alpha=args.lora_alpha, lora_dropout=args.lora_dropout, early_stop_patience=args.early_stop_patience, ) times["train_seconds"] = time.perf_counter() - t0 t0 = time.perf_counter() model, tok = load_model_with_lora( model_id, run_dir, use_lora=use_lora, use_4bit=args.use_4bit, use_custom_head=args.use_custom_head, cls_hidden_dim=args.cls_hidden_dim, cls_dropout=args.cls_dropout, pooling=args.pooling, cls_hidden_dim2=args.cls_hidden_dim2, train_last_n_layers=args.train_last_n_layers, ) metrics, rows = eval_model_seqcls(model, tok, test_df, max_length=args.max_length, batch_size=args.eval_batch) times["finetuned_eval_seconds"] = time.perf_counter() - t0 pd.DataFrame(rows).to_csv(os.path.join(run_dir, "finetuned_predictions.csv"), index=False) with open(os.path.join(run_dir, "finetuned_metrics.json"), "w") as f: json.dump(metrics, f, indent=2) del model, tok cleanup_cuda() with open(os.path.join(run_dir, "timings.json"), "w") as f: json.dump(times, f, indent=2) summary_rows.append({ "run_suffix": short_run_suffix(args.run_suffix) if args.run_suffix else "", "model_id": model_id, "run_dir": run_dir, "f1_label1_POAF": metrics.get("f1_label1_POAF"), "roc_auc": metrics.get("roc_auc"), "accuracy": metrics.get("accuracy"), "recall_label1_POAF": metrics.get("recall_label1_POAF"), "precision_label1_POAF": metrics.get("precision_label1_POAF"), "sensitivity": metrics.get("sensitivity"), "specificity": metrics.get("specificity"), "pr_auc": metrics.get("pr_auc"), "f1_label0_NoPOAF": metrics.get("f1_label0_NoPOAF"), }) summary_path = os.path.join(args.out, "summary.csv") new_df = pd.DataFrame(summary_rows) if os.path.isfile(summary_path): existing = pd.read_csv(summary_path) if args.run_suffix: if "run_suffix" not in existing.columns: existing["run_suffix"] = "" short_suf = short_run_suffix(args.run_suffix) mask = (existing["run_suffix"].astype(str).isin([args.run_suffix, short_suf])) & (existing["model_id"].isin(args.models)) existing = existing[~mask] new_df = pd.concat([existing, new_df], ignore_index=True) else: existing = existing[~existing["model_id"].isin(args.models)] new_df = pd.concat([existing, new_df], ignore_index=True) if "run_suffix" in new_df.columns: new_df["run_suffix"] = new_df["run_suffix"].fillna("") if "f1_label1_POAF" in new_df.columns: new_df["f1_label1"] = new_df["f1_label1_POAF"] if "recall_label1_POAF" in new_df.columns: new_df["recall_label1"] = new_df["recall_label1_POAF"] if "precision_label1_POAF" in new_df.columns: new_df["precision_label1"] = new_df["precision_label1_POAF"] desired_cols = [ "run_suffix", # keep for discriminative experiments to know which of the 12 configs "model_id", "run_dir", "f1_label1", "roc_auc", "accuracy", "recall_label1", "precision_label1", ] present_cols = [c for c in desired_cols if c in new_df.columns] new_df = new_df[present_cols] atomic_write_csv(new_df, summary_path) elapsed = time.perf_counter() - t0_total print(f"Done. Summary: {summary_path}") print(f"Total time: {elapsed:.1f}s ({elapsed/60:.1f} min, {elapsed/3600:.1f} h)") if __name__ == "__main__": main()