Team Ai
Apppublic

satom/stealth-trace

sourceHugging Faceupdated 6mo agoView on Hugging Face
0likes
TRACEmodel.py97 linesDownload Raw Back to root
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)}")