satom/stealth-trace
0
1import streamlit as st2import sksurv3from sksurv.linear_model import CoxPHSurvivalAnalysis4from sksurv.linear_model.coxph import BreslowEstimator5import pandas as pd6import matplotlib.pyplot as plt7import numpy as np8from sksurv.ensemble import RandomSurvivalForest9import joblib10from huggingface_hub import hf_hub_download11 12@st.cache_resource13def load_model():14 model_path = hf_hub_download(15 repo_id="satom/stealth-trace-model",16 filename="rsfmodel.sav"17 )18 return joblib.load(model_path)19 20with st.spinner("Downloading and loading model... (This may take a few minutes on first launch)"):21 rsf = load_model()22 23st.title('Prediction model for MASLD-HCC (STEALTH-TRACE model)')24st.markdown("Enter the following items and click 'Submit' to display the predicted HCC risk")25 26with st.form('user_inputs'):27 age=st.number_input('age (year)', min_value=18,max_value=100)28 height=st.number_input('height (cm)', min_value=100.0,max_value=300.0, value=170.0, step=0.1, format="%.1f")29 weight=st.number_input('body weight (kg)', min_value=20.0,max_value=300.0, value=65.0, step=0.1, format="%.1f")30 PLT=st.number_input('Platelet count (×10^4/µL)', min_value=1.0,max_value=75.0, value=15.0, step=0.1, format="%.1f")31 ALB=st.number_input('Albumin (g/dL)', min_value=1.0,max_value=7.0, value=4.0, step=0.1, format="%.1f")32 AST=st.number_input('AST (IU/L)', min_value=1,max_value=500, value=30)33 ALT=st.number_input('ALT (IU/L)', min_value=1,max_value=500, value=30)34 GGT=st.number_input('γ-GTP (IU/L)', min_value=1,max_value=1000, value=50)35 st.form_submit_button()36 37if height > 0:38 height2=height*height39 BMI0=weight/height240 BMI=BMI0*1000041 42 X=pd.DataFrame(43 data={'age': [age],44 'BMI': [BMI],45 'ALB': [ALB],46 'AST': [AST],47 'ALT': [ALT],48 'GGT': [GGT],49 'PLT': [PLT],50 }51 )52 53 surv = rsf.predict_survival_function(X, return_array=True)54 55 fig, ax = plt.subplots()56 for i, s in enumerate(surv):57 ax.step(rsf.unique_times_, s, where="post", label=str(i))58 59 ax.set_xlim(0,10)60 ax.set_ylim(0,1)61 ax.set_ylabel("predicted HCC development")62 ax.set_xlabel("years")63 ax.grid(True)64 ax.invert_yaxis()65 ax.set_yticks([0.0, 0.2, 0.4,0.6,0.8,1.0],66 ['100%', '80%', '60%', '40%', '20%', '0%'])67 68 st.header("HCC risk for submitted patient")69 st.pyplot(fig)70 71 y_event = rsf.predict_survival_function(X, return_array=True).flatten()72 HCCincidence=100*(1-y_event)73 74 df1 = pd.DataFrame(rsf.unique_times_)75 df1.columns = ['timepoint (year)']76 df2 = pd.DataFrame(HCCincidence)77 df2.columns = ['predicted HCC incidence (%)']78 df_merge = pd.concat([df1.reset_index(drop=True), df2.reset_index(drop=True)], axis=1)79 80 one_year_idx = (np.abs(df_merge['timepoint (year)'] - 1.0)).argmin()81 three_year_idx = (np.abs(df_merge['timepoint (year)'] - 3.0)).argmin()82 five_year_idx = (np.abs(df_merge['timepoint (year)'] - 5.0)).argmin()83 84 one_val = df_merge.iloc[one_year_idx, 1]85 three_val = df_merge.iloc[three_year_idx, 1]86 five_val = df_merge.iloc[five_year_idx, 1]87 88 def format_incidence(value):89 if 0 <= value < 0.001:90 return "less than 0.001%"91 else:92 return f"{value:.3f}%"93 94 st.subheader("predicted HCC incidence at each time point")95 st.write(f"**predicted HCC incidence at 1 year:** {format_incidence(one_val)}")96 st.write(f"**predicted HCC incidence at 3 year:** {format_incidence(three_val)}")97 st.write(f"**predicted HCC incidence at 5 year:** {format_incidence(five_val)}")