| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
|
|
|
|
| import gradio as gr |
| import joblib |
| import numpy as np |
| import pandas as pd |
| import plotly.graph_objects as go |
|
|
| |
| |
| |
|
|
| bundle = joblib.load("distilled_ridge_cox_from_rsf.pkl") |
| student_model = bundle["model"] |
| features_kd = bundle["features"] |
| teacher_scaler= bundle["teacher_scaler"] |
| teacher_weight= bundle["teacher_weight"] |
| c_index = bundle.get("c_index_student", 0.70) |
|
|
| rsf_bundle = joblib.load("rsf_model.pkl") |
| rsf_model = rsf_bundle["model"] |
| rsf_features = rsf_bundle["features"] |
|
|
| |
| |
| |
|
|
| def encode_age_group(age): |
| if age < 50: return 0 |
| elif age < 60: return 1 |
| elif age < 70: return 2 |
| elif age < 80: return 3 |
| else: return 4 |
|
|
| def encode_psa_level(psa): |
| """psa is the RAW value before any log transform.""" |
| if psa < 10: return 0 |
| elif psa < 50: return 1 |
| elif psa < 100: return 2 |
| elif psa < 500: return 3 |
| elif psa < 1000: return 4 |
| elif psa < 10000: return 5 |
| else: return 6 |
|
|
| def encode_gleason_group(g): |
| if g <= 6: return 0 |
| elif g == 7: return 1 |
| elif g == 8: return 2 |
| else: return 3 |
|
|
| |
| |
| |
| |
| |
| |
|
|
| def clinical_risk_score(gleason, psa, mets): |
| """ |
| Points: |
| Gleason β€6 β 0, 7 β 1, 8 β 2, β₯9 β 3 |
| PSA <10 β 0, 10β49 β 1, 50β99 β 2, β₯100 β 3 |
| Mets 0 β 0, 1 β 1, 2β3 β 2, β₯4 β 3 |
| """ |
| g_pts = 0 if gleason <= 6 else 1 if gleason == 7 else 2 if gleason == 8 else 3 |
| p_pts = 0 if psa < 10 else 1 if psa < 50 else 2 if psa < 100 else 3 |
| m_pts = 0 if mets == 0 else 1 if mets == 1 else 2 if mets <= 3 else 3 |
| return g_pts + p_pts + m_pts |
|
|
| def risk_category_from_score(score): |
| if score <= 2: |
| return "Low Risk", "green" |
| elif score <= 5: |
| return "Medium Risk", "orange" |
| else: |
| return "High Risk", "red" |
|
|
| |
| |
| |
|
|
| def rsf_sf_at(sf_fn, months): |
| """Interpolate RSF survival function at a given number of months.""" |
| idx = np.searchsorted(sf_fn.x, months, side="right") - 1 |
| if idx < 0: |
| return 1.0 |
| return float(sf_fn.y[idx]) |
|
|
| def rsf_median_survival(sf_fn): |
| """Return median survival months (first t where S(t) β€ 0.5).""" |
| times = sf_fn.x |
| probs = sf_fn.y |
| idx = np.searchsorted(-probs, -0.5) |
| if idx < len(times): |
| return float(times[idx]) |
| return None |
|
|
| |
| |
| |
|
|
| def predict_survival(age, psa, gleason, |
| cns, liver, lung, bones, |
| spleen, kidney, bone_marrow): |
|
|
| if age is None or psa is None: |
| return "Please fill in Age and PSA.", None, None |
|
|
| age = float(age) |
| psa = float(psa) |
| gleason = int(gleason) |
|
|
| mets = sum(int(x) for x in [cns, liver, lung, bones, |
| spleen, kidney, bone_marrow]) |
| has_metastasis = int(mets > 0) |
| age_group = encode_age_group(age) |
| psa_level = encode_psa_level(psa) |
| gleason_group = encode_gleason_group(gleason) |
| log_psa = np.log1p(psa) |
|
|
| patient_dict = { |
| "Age": age, |
| "Baseline PSA": log_psa, |
| "Gleason Score": gleason, |
| "Age_Group": age_group, |
| "PSA_Level": psa_level, |
| "Gleason_Group": gleason_group, |
| "Metastatic_Burden": mets, |
| "Has_Metastasis": has_metastasis, |
| "CNS": int(cns), |
| "LIVER": int(liver), |
| "LUNG": int(lung), |
| "BONES": int(bones), |
| "SPLEEN": int(spleen), |
| "KIDNEY": int(kidney), |
| "BONE MARROW": int(bone_marrow), |
| } |
|
|
| |
| rsf_input = pd.DataFrame([patient_dict])[rsf_features] |
| sf_fn = rsf_model.predict_survival_function(rsf_input)[0] |
|
|
| |
| time_pts = sf_fn.x |
| surv_pts = sf_fn.y |
|
|
| |
| s_5yr = rsf_sf_at(sf_fn, 60) |
| s_7yr = rsf_sf_at(sf_fn, 84) |
| s_9yr = rsf_sf_at(sf_fn, 108) |
|
|
| |
| med = rsf_median_survival(sf_fn) |
| if med is not None: |
| median_label = f"{int(med)} months" |
| else: |
| |
| median_label = f">{ int(time_pts[-1]) } months" |
|
|
| |
| risk_score = clinical_risk_score(gleason, psa, mets) |
| category, color = risk_category_from_score(risk_score) |
|
|
| |
| if risk_score >= 6: |
| recommendation = "Aggressive treatment recommended" |
| elif risk_score >= 3: |
| recommendation = "Standard treatment with regular follow-up" |
| else: |
| recommendation = "Active surveillance with monitoring" |
|
|
| |
| output = f""" |
| ## Results |
| | Metric | Value | |
| |--------|-------| |
| | **Estimated Median Survival** | {median_label} | |
| | **Clinical Risk Score** | {risk_score} / 9 | |
| | **Risk Category** | <span style="color:{color}; font-weight:bold">{category}</span> | |
| | **5-Year Survival (RSF)** | {s_5yr*100:.1f}% | |
| | **7-Year Survival (RSF)** | {s_7yr*100:.1f}% | |
| | **9-Year Survival (RSF)** | {s_9yr*100:.1f}% | |
| | **Metastatic Sites** | {mets} site(s) | |
| | **Recommendation** | {recommendation} | |
| |
| *Survival estimates from RSF teacher model at its native follow-up time points.* |
| """ |
|
|
| |
| gauge = go.Figure(go.Indicator( |
| mode="gauge+number", |
| value=risk_score, |
| title={"text": "Clinical Risk Score (0β9)"}, |
| gauge={ |
| "axis": {"range": [0, 9], "tickvals": list(range(10))}, |
| "bar": {"color": color}, |
| "steps": [ |
| {"range": [0, 3], "color": "#d1fae5"}, |
| {"range": [3, 6], "color": "#fef9c3"}, |
| {"range": [6, 9], "color": "#fee2e2"}, |
| ], |
| "threshold": { |
| "line": {"color": "black", "width": 3}, |
| "thickness": 0.75, |
| "value": risk_score, |
| }, |
| }, |
| )) |
| gauge.update_layout(margin=dict(t=60, b=20, l=20, r=20), height=280) |
|
|
| |
| labels = [f"{int(t/12)}yr" for t in time_pts] |
| bar_colors = [] |
| for p in surv_pts: |
| if p >= 0.80: |
| bar_colors.append("#34d399") |
| elif p >= 0.60: |
| bar_colors.append("#fbbf24") |
| else: |
| bar_colors.append("#f87171") |
|
|
| curve = go.Figure(go.Bar( |
| x=labels, |
| y=list(surv_pts), |
| text=[f"{p*100:.1f}%" for p in surv_pts], |
| textposition="auto", |
| marker_color=bar_colors, |
| )) |
| curve.update_layout( |
| title="RSF Survival Probability (native follow-up time points)", |
| yaxis=dict(range=[0, 1], tickformat=".0%"), |
| margin=dict(t=50, b=20, l=20, r=20), |
| height=280, |
| ) |
|
|
| return output, gauge, curve |
|
|
|
|
| |
| |
| |
|
|
| css = """ |
| body { background: #f1f5f9; } |
| .section { |
| background: white; padding: 16px; border-radius: 10px; |
| border: 1px solid #e2e8f0; box-shadow: 0 2px 10px rgba(0,0,0,0.04); |
| margin-bottom: 12px; |
| } |
| .primary-btn button { |
| width: 100% !important; background: #2563eb !important; |
| color: white !important; font-weight: 700; font-size: 16px; |
| padding: 14px !important; border-radius: 10px !important; |
| } |
| .primary-btn button:hover { background: #1d4ed8 !important; } |
| """ |
|
|
| with gr.Blocks(css=css, title="Prostate Cancer Survival Predictor") as demo: |
|
|
| gr.Markdown("# Prostate Cancer Survival Predictor") |
|
|
| with gr.Row(equal_height=False): |
| with gr.Column(scale=1): |
| with gr.Group(elem_classes="section"): |
| gr.Markdown("### Patient Information") |
| age = gr.Number(label="Age (years)") |
| psa = gr.Number(label="Baseline PSA (ng/mL) β enter raw value") |
| gleason = gr.Slider(6, 10, value=7, step=1, label="Gleason Score (6β10)") |
|
|
| with gr.Column(scale=1): |
| with gr.Group(elem_classes="section"): |
| gr.Markdown("### Metastatic Sites") |
| cns = gr.Checkbox(label="CNS (Brain)") |
| liver = gr.Checkbox(label="Liver") |
| lung = gr.Checkbox(label="Lung") |
| bones = gr.Checkbox(label="Bones") |
| spleen = gr.Checkbox(label="Spleen") |
| kidney = gr.Checkbox(label="Kidney") |
| bone_marrow = gr.Checkbox(label="Bone Marrow") |
|
|
| with gr.Row(): |
| btn = gr.Button("Predict Survival", elem_classes="primary-btn") |
|
|
| with gr.Row(): |
| with gr.Group(elem_classes="section"): |
| gr.Markdown("### Prediction Results") |
| output = gr.Markdown("Enter patient data and click Predict.") |
|
|
| with gr.Row(): |
| gauge = gr.Plot(label="Risk Score") |
| curve = gr.Plot(label="Survival Probability") |
|
|
| gr.Markdown(""" |
| --- |
| ### Interpretation Guide |
| |
| | Clinical Score | Category | Action | |
| |---|---|---| |
| | 0β2 | Low Risk | Active surveillance with monitoring | |
| | 3β5 | Medium Risk | Standard treatment with regular follow-up | |
| | 6β9 | High Risk | Aggressive treatment recommended | |
| |
| **Score components:** Gleason score (0β3 pts) + PSA level (0β3 pts) + Metastatic burden (0β3 pts) |
| |
| **Survival estimates** are from the RSF teacher model at its native follow-up time points (5β10 years). |
| These reflect the model's training cohort and should be interpreted in that clinical context. |
| """) |
|
|
| btn.click( |
| predict_survival, |
| inputs=[age, psa, gleason, cns, liver, lung, bones, spleen, kidney, bone_marrow], |
| outputs=[output, gauge, curve], |
| ) |
|
|
| if __name__ == "__main__": |
| demo.launch(share=True) |