""" scripts/train_severity_classifier.py ──────────────────────────────────────────────────────────────── Retrains the severity classifier with the two fixes that were deferred on 2026-06-25 (see project memory): 1. Class-weighted loss -- "critical" is ~15% of the data and was getting recall 0.486. ClinicalClassifier.train() now supports weighted_loss=True to counter this. 2. Corrected weak-supervision labels -- src/etl/transform.py's `\\bacute\\b` rule was matching negated phrases like "no acute complaints". Fixed to `(? pd.DataFrame: """Fetch note text from Supabase and derive fresh severity labels. Args: limit: Maximum number of notes to fetch, or None for all. Returns: DataFrame with ``transcription`` and freshly-derived ``severity`` columns. """ with get_session() as session: stmt = select(ClinicalNote.transcription) if limit: stmt = stmt.limit(limit) rows = session.execute(stmt).all() df = pd.DataFrame({"transcription": [r.transcription for r in rows]}) logger.info("Loaded %d notes from Supabase", len(df)) df = derive_severity_labels(df) logger.info( "Severity distribution (freshly derived): %s", df["severity"].value_counts().to_dict(), ) return df _ATTEMPT_DIR_TEMPLATE = "severity_classifier_attempt_{i}" def main() -> None: parser = argparse.ArgumentParser( description="Retrain the severity classifier with class-weighted loss and corrected labels." ) parser.add_argument( "--limit", type=int, default=None, help="Max number of notes to train on (default: all notes). " "Implies --output-dir points at a scratch directory unless " "overridden, so a smoke test can never touch the real checkpoint.", ) parser.add_argument( "--output-dir", type=str, default=None, help="Where to save the checkpoint (default: the real model dir for " "a full run, or a smoke-test scratch dir when --limit is set)", ) parser.add_argument( "--n-seeds", type=int, default=1, help="Train this many attempts with different seeds and keep " "only the one with the best critical-class F1 (default: 1, " "i.e. a single run). Fine-tuning a small classification " "head on a small dataset is sensitive to random init -- " "running several seeds and selecting the best is more " "reliable than trusting a single arbitrary run.", ) args = parser.parse_args() df = _load_notes(args.limit) output_dir = Path(args.output_dir) if args.output_dir else None if output_dir is None and args.limit is not None: output_dir = _SMOKE_TEST_DIR logger.warning( "--limit set without --output-dir -- saving to scratch dir %s " "instead of the real checkpoint.", output_dir, ) elif output_dir is None: output_dir = ModelConfig.fine_tuned_dir n_seeds = max(1, args.n_seeds) attempts: list[tuple[int, dict, Path]] = [] for i in range(n_seeds): seed = TrainingConfig.random_seed + i attempt_dir = Paths.models / _ATTEMPT_DIR_TEMPLATE.format(i=i) logger.info("=" * 60) logger.info("Attempt %d/%d -- seed=%d", i + 1, n_seeds, seed) clf = ClinicalClassifier(task="severity", output_dir=attempt_dir) results = clf.train(df, resume=False, weighted_loss=True, seed=seed) critical = results["per_class"]["critical"] logger.info( "Attempt %d: critical precision=%.3f recall=%.3f f1=%.3f | " "test_acc=%.4f test_f1=%.4f", i + 1, critical["precision"], critical["recall"], critical["f1"], results["test_accuracy"], results["test_f1"], ) attempts.append((seed, results, attempt_dir)) best_seed, best_results, best_dir = max( attempts, key=lambda a: a[1]["per_class"]["critical"]["f1"] ) logger.info("=" * 60) logger.info("Attempt comparison (sorted by critical F1):") for seed, results, _ in sorted( attempts, key=lambda a: a[1]["per_class"]["critical"]["f1"], reverse=True ): c = results["per_class"]["critical"] marker = " <- best" if seed == best_seed else "" logger.info( " seed=%-4d critical: precision=%.3f recall=%.3f f1=%.3f%s", seed, c["precision"], c["recall"], c["f1"], marker, ) # Promote the winning attempt's checkpoint to the real output_dir, # then discard every attempt directory (including the winner's, # now duplicated at output_dir). if output_dir != best_dir: shutil.copytree(best_dir, output_dir, dirs_exist_ok=True) for _, _, attempt_dir in attempts: shutil.rmtree(attempt_dir, ignore_errors=True) results = best_results logger.info("=" * 60) logger.info( "Best: seed=%d val_acc=%.4f val_f1=%.4f test_acc=%.4f test_f1=%.4f", best_seed, results["val_accuracy"], results["val_f1"], results["test_accuracy"], results["test_f1"], ) critical = results["per_class"]["critical"] logger.info( "Critical class: precision=%.3f recall=%.3f f1=%.3f", critical["precision"], critical["recall"], critical["f1"], ) # Record this run in the DB -- what the dashboard's Model Metrics # page and the /model/metrics API endpoint actually read from. # Skipped for smoke-test runs (--limit set): a tiny/partial run # shouldn't overwrite the real deployed model's recorded metrics. if args.limit is None: with get_session() as session: session.query(ModelRun).filter_by( task="severity", is_deployed=True ).update({"is_deployed": False}) ModelRunRepository.create( session, model_name = ModelConfig.classifier_model, task = "severity", training_samples = len(df), val_accuracy = results["val_accuracy"], val_f1 = results["val_f1"], test_accuracy = results["test_accuracy"], test_f1 = results["test_f1"], per_class = results["per_class"], confusion_matrix = results["confusion_matrix"], history = results["history"], epochs = len(results["history"]), is_deployed = True, run_notes = ( f"Class-weighted loss, corrected acute-negation labels. " f"Best of {n_seeds} seed(s) (seed={best_seed}), selected by critical-class F1." ), ) logger.info("Run recorded in model_runs (is_deployed=True)") else: logger.info("--limit set -- skipping model_runs DB record (smoke test)") if __name__ == "__main__": main()