sadiqM's picture
Initial commit: NephroScreen — CKD screening ML app
dc3d345
Raw
History Blame Contribute Delete
5.44 kB
"""Train, evaluate, and persist the CKD screening model.
Trains three comparable classifiers (Logistic Regression, Random Forest,
Decision Tree), reports held-out and cross-validated metrics, and serializes the
Random Forest pipeline plus a metrics manifest for the API to load.
Run: python -m nephroscreen.train
"""
import json
import warnings
from datetime import datetime, timezone
import joblib
import numpy as np
import pandas as pd
from sklearn.ensemble import RandomForestClassifier
from sklearn.linear_model import LogisticRegression
from sklearn.metrics import (
accuracy_score,
f1_score,
precision_recall_curve,
precision_score,
recall_score,
roc_auc_score,
)
from sklearn.model_selection import StratifiedKFold, cross_val_score, train_test_split
from sklearn.pipeline import Pipeline
from sklearn.preprocessing import StandardScaler
from sklearn.tree import DecisionTreeClassifier
from .config import (
DATA_PATH,
METRICS_PATH,
MODEL_COLUMNS,
MODEL_PATH,
MODELS_DIR,
RANDOM_STATE,
TEST_SIZE,
)
from .preprocessing import CKDPreprocessor, load_dataset
warnings.filterwarnings("ignore")
def build_pipeline(estimator, scale_numeric=False):
steps = [("preprocess", CKDPreprocessor())]
if scale_numeric:
steps.append(("scaler", StandardScaler()))
steps.append(("model", estimator))
return Pipeline(steps)
def score_model(model, X_test, y_test, X_train, y_train, cv):
pred = model.predict(X_test)
prob = model.predict_proba(X_test)[:, 1]
cv_scores = cross_val_score(model, X_train, y_train, cv=cv, scoring="accuracy")
return {
"accuracy": round(accuracy_score(y_test, pred) * 100, 2),
"precision": round(precision_score(y_test, pred) * 100, 2),
"recall": round(recall_score(y_test, pred) * 100, 2),
"f1": round(f1_score(y_test, pred) * 100, 2),
"roc_auc": round(roc_auc_score(y_test, prob), 4),
"cv_mean": round(cv_scores.mean() * 100, 2),
"cv_std": round(cv_scores.std() * 100, 2),
}
def recall_first_threshold(y_test, prob, min_recall=0.99):
"""Lowest decision threshold that still achieves `min_recall` on the test set.
For a screening tool a missed CKD case (false negative) is the costly error,
so we prefer a threshold that keeps recall very high rather than the default 0.5.
"""
precision, recall, thresholds = precision_recall_curve(y_test, prob)
# thresholds has len-1 vs precision/recall; align by dropping the last point.
best_t = 0.5
for p, r, t in zip(precision[:-1], recall[:-1], thresholds):
if r >= min_recall:
best_t = float(t)
return round(best_t, 3)
def main():
MODELS_DIR.mkdir(parents=True, exist_ok=True)
print(f"Loading dataset from {DATA_PATH} ...")
df = load_dataset(DATA_PATH)
print(f"Shape: {df.shape} | CKD: {int(df.target.sum())} | Healthy: {int((df.target == 0).sum())}")
X = df.drop(columns=["target"])
y = df["target"]
X_train, X_test, y_train, y_test = train_test_split(
X, y, test_size=TEST_SIZE, random_state=RANDOM_STATE, stratify=y
)
cv = StratifiedKFold(n_splits=10, shuffle=True, random_state=RANDOM_STATE)
specs = {
"Logistic Regression": (
build_pipeline(LogisticRegression(C=1.0, max_iter=1000, random_state=RANDOM_STATE), True),
),
"Random Forest": (
build_pipeline(RandomForestClassifier(n_estimators=100, random_state=RANDOM_STATE, n_jobs=1)),
),
"Decision Tree": (
build_pipeline(DecisionTreeClassifier(max_depth=5, min_samples_split=5, random_state=RANDOM_STATE)),
),
}
results = {}
fitted = {}
for name, (pipe,) in specs.items():
pipe.fit(X_train, y_train)
fitted[name] = pipe
results[name] = score_model(pipe, X_test, y_test, X_train, y_train, cv)
print(f" {name:<22} acc={results[name]['accuracy']}% auc={results[name]['roc_auc']}")
# Random Forest is the serving model (best held-out performance).
serving = fitted["Random Forest"]
rf_prob = serving.predict_proba(X_test)[:, 1]
threshold = recall_first_threshold(y_test, rf_prob, min_recall=0.99)
importances = (
pd.Series(serving.named_steps["model"].feature_importances_, index=MODEL_COLUMNS)
.sort_values(ascending=False)
.round(4)
)
joblib.dump(serving, MODEL_PATH)
print(f"\nSaved serving model -> {MODEL_PATH}")
manifest = {
"model": "RandomForestClassifier(n_estimators=100)",
"trained_at": datetime.now(timezone.utc).isoformat(),
"random_state": RANDOM_STATE,
"test_size": TEST_SIZE,
"n_train": int(len(X_train)),
"n_test": int(len(X_test)),
"serving_threshold": threshold,
"metrics": results,
"top_features": importances.head(10).to_dict(),
"note": (
"Near-perfect scores reflect the high separability of this small UCI "
"dataset (400 patients, 80-patient test set) — not clinical deployment "
"readiness. See README for the honest interpretation."
),
}
with open(METRICS_PATH, "w", encoding="utf-8") as f:
json.dump(manifest, f, indent=2)
print(f"Saved metrics manifest -> {METRICS_PATH}")
print(f"Recall-first serving threshold (recall>=0.99): {threshold}")
if __name__ == "__main__":
main()