UdaraChamidu's picture
Upload 7 files
1332d30 verified
Raw
History Blame Contribute Delete
5.25 kB
import pickle
import numpy as np
import pandas as pd
import streamlit as st
import tensorflow as tf
st.set_page_config(
page_title="Heart Disease Risk Prediction",
page_icon="heart",
layout="wide",
)
RISK_LABELS = ["Low", "Medium", "High"]
SEX_OPTIONS = {"Female": 0, "Male": 1}
CP_OPTIONS = {
"Typical angina": 0,
"Atypical angina": 1,
"Non-anginal pain": 2,
"Asymptomatic": 3,
}
FBS_OPTIONS = {"False": 0, "True": 1}
RESTECG_OPTIONS = {
"Normal": 0,
"ST-T wave abnormality": 1,
"Left ventricular hypertrophy": 2,
}
EXANG_OPTIONS = {"No": 0, "Yes": 1}
SLOPE_OPTIONS = {"Upsloping": 0, "Flat": 1, "Downsloping": 2}
CA_OPTIONS = {"0 vessels": 0, "1 vessel": 1, "2 vessels": 2, "3 vessels": 3, "4 vessels": 4}
THAL_OPTIONS = {"Fixed defect": 1, "Normal": 2, "Reversible defect": 3}
@st.cache_resource
def load_artifacts():
model = tf.keras.models.load_model("heart_disease_ann.h5")
with open("preprocessor.pkl", "rb") as f:
preprocessor = pickle.load(f)
return model, preprocessor
def pick_code(label, options, help_text=None):
return options[st.selectbox(label, list(options.keys()), help=help_text)]
try:
model, preprocessor = load_artifacts()
except Exception as exc:
st.error(f"Could not load model files: {exc}")
st.stop()
st.title("Heart Disease Risk Prediction")
st.caption(
"Enter patient details, run prediction, and review risk probabilities. "
"This tool supports screening and is not a medical diagnosis."
)
tab_predict, tab_guide = st.tabs(["Prediction", "Input Guide"])
with tab_predict:
with st.form(key="heart_form"):
col_a, col_b = st.columns(2)
with col_a:
st.subheader("Patient Profile")
age = st.slider("Age", min_value=1, max_value=120, value=30)
sex = pick_code("Sex", SEX_OPTIONS, "Male = 1, Female = 0")
cp = pick_code("Chest Pain Type", CP_OPTIONS)
exang = pick_code("Exercise Induced Angina", EXANG_OPTIONS, "Yes = 1, No = 0")
slope = pick_code("Slope of ST Segment", SLOPE_OPTIONS)
thal = pick_code("Thalassemia", THAL_OPTIONS)
with col_b:
st.subheader("Clinical Measurements")
trestbps = st.slider("Resting Blood Pressure (mm Hg)", min_value=80, max_value=220, value=120)
chol = st.slider("Cholesterol (mg/dl)", min_value=100, max_value=600, value=200)
fbs = pick_code("Fasting Blood Sugar > 120 mg/dl?", FBS_OPTIONS, "True = 1, False = 0")
restecg = pick_code("Resting ECG Result", RESTECG_OPTIONS)
thalach = st.slider("Max Heart Rate Achieved", min_value=60, max_value=220, value=150)
oldpeak = st.number_input("ST Depression (oldpeak)", min_value=0.0, max_value=10.0, value=1.0, step=0.1)
ca = pick_code("Major Vessels Colored by Fluoroscopy", CA_OPTIONS)
submit_button = st.form_submit_button("Predict Risk", use_container_width=True)
if submit_button:
input_df = pd.DataFrame(
[
{
"age": age,
"sex": sex,
"cp": cp,
"trestbps": trestbps,
"chol": chol,
"fbs": fbs,
"restecg": restecg,
"thalach": thalach,
"exang": exang,
"oldpeak": oldpeak,
"slope": slope,
"ca": ca,
"thal": thal,
}
]
)
input_processed = preprocessor.transform(input_df)
pred_probs = model.predict(input_processed, verbose=0)[0]
pred_class = RISK_LABELS[int(np.argmax(pred_probs))]
if pred_class == "Low":
st.success(f"Predicted Risk Level: {pred_class}")
elif pred_class == "Medium":
st.warning(f"Predicted Risk Level: {pred_class}")
else:
st.error(f"Predicted Risk Level: {pred_class}")
st.subheader("Probability Scores")
metric_cols = st.columns(3)
for idx, label in enumerate(RISK_LABELS):
metric_cols[idx].metric(label, f"{pred_probs[idx] * 100:.2f}%")
prob_df = pd.DataFrame(
{"Risk": RISK_LABELS, "Probability": (pred_probs * 100).round(2)}
).set_index("Risk")
st.bar_chart(prob_df)
with st.expander("Show Submitted Values"):
st.dataframe(input_df, hide_index=True, use_container_width=True)
with tab_guide:
st.subheader("Category Encoding Reference")
st.write("Sex: Female = 0, Male = 1")
st.write("Chest Pain Type: Typical angina = 0, Atypical angina = 1, Non-anginal pain = 2, Asymptomatic = 3")
st.write("Fasting Blood Sugar > 120 mg/dl: False = 0, True = 1")
st.write(
"Resting ECG Result: Normal = 0, ST-T wave abnormality = 1, "
"Left ventricular hypertrophy = 2"
)
st.write("Exercise Induced Angina: No = 0, Yes = 1")
st.write("Slope of ST Segment: Upsloping = 0, Flat = 1, Downsloping = 2")
st.write("Major Vessels Colored by Fluoroscopy: Number of vessels from 0 to 4")
st.write("Thalassemia: Fixed defect = 1, Normal = 2, Reversible defect = 3")