File size: 4,276 Bytes
07008ac 7496b0a 07008ac 7496b0a 07008ac 7496b0a 07008ac e39413b 07008ac 176c96a 07008ac c0ba818 07008ac c0ba818 a9ce186 4594441 58684fd 69cc809 ad8ef51 4594441 63086cf a3ae92f ad8ef51 58684fd 4594441 63086cf 1c2796b 2868d14 bc73160 22e3347 ad8ef51 58684fd 8a2c348 059a750 c0ba818 059a750 c0ba818 059a750 07008ac 8a2c348 c0ba818 07008ac | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 | import pickle
import pandas as pd
import shap
from shap.plots._force_matplotlib import draw_additive_plot
import gradio as gr
import numpy as np
import matplotlib.pyplot as plt
# load the model from disk
loaded_model = pickle.load(open("heart_xgb.pkl", 'rb'))
# Setup SHAP
explainer = shap.Explainer(loaded_model) # PLEASE DO NOT CHANGE THIS.
# Create the main function for server
def main_func(age, sex, cp, trtbps, chol, fbs, restecg, thalachh,exng,oldpeak,slp,caa,thall):
new_row = pd.DataFrame.from_dict({'age':age,'sex':sex,
'cp':cp,'trtbps':trtbps,'chol':chol,
'fbs':fbs, 'restecg':restecg,'thalachh':thalachh,'exng':exng,
'oldpeak':oldpeak,'slp':slp,'caa':caa,'thall':thall},
orient = 'index').transpose()
prob = loaded_model.predict_proba(new_row)
shap_values = explainer(new_row)
# plot = shap.force_plot(shap_values[0], matplotlib=True, figsize=(30,30), show=False)
# plot = shap.plots.waterfall(shap_values[0], max_display=6, show=False)
plot = shap.plots.bar(shap_values[0], max_display=6, order=shap.Explanation.abs, show_data='auto', show=False)
plt.tight_layout()
local_plot = plt.gcf()
plt.close()
return {"Low Chance": float(prob[0][0]), "High Chance": 1-float(prob[0][0])}, local_plot
# Create the UI
title = "**Heart Attack Predictor & Interpreter** 🪐 TEAM 7"
description1 = """This app takes info from subjects and predicts their heart attack likelihood. Do not use for medical diagnosis."""
description2 = """
To use the app, click on one of the examples, or adjust the values of the factors, and click on Analyze. 🤞
"""
with gr.Blocks(title=title) as demo:
gr.Markdown(f"## {title}")
gr.Markdown(description1)
gr.Markdown("""---""")
gr.Markdown(description2)
gr.Markdown("""---""")
with gr.Row():
with gr.Column():
age = gr.Number(label="age Score", value=40)
sex = gr.Dropdown(label="Sex", choices = ["Female", "Male"], type = "index")
cp = gr.Dropdown(label="cp Score", choices = ["0", "1","2","3"], type = "index", value = "1",
info = "Value 1: typical angina Value 2: atypical angina Value 3: non-anginal pain Value 4: asymptomatic trtbps : resting blood pressure (in mm Hg)")
with gr.Column():
trtbps = gr.Number(label="trtbps Score", value=100, step=1, info = "The person's resting blood pressure (mm Hg on admission to the hospital)")
chol = gr.Number(label="chol Score", value=130, info = "cholestoral in mg/dl fetched via BMI sensor" )
fbs = gr.Dropdown(label="fbs Score", choices = ["0", "1"], type = "index", value = "1")
restecg = gr.Dropdown(label="restecg Score", choices = ["0", "1","2"], type = "index", value = "1", info = "resting electrocardiographic results")
thall = gr.Dropdown(label="thall Score", choices = ["0", "1","2","3"], type = "index", value = "1")
with gr.Column():
thalachh = gr.Number(label="thalachh Score", value=100)
exng = gr.Dropdown(label="exng Score", choices = ["0", "1"], type = "index", value = "1")
oldpeak = gr.Slider(label="oldpeak Score", minimum=0, maximum=10, value=4, step=.1)
slp = gr.Dropdown(label="slp Score", choices = ["0","1","2"], type = "index", value = "1")
caa = gr.Dropdown(label="caa Score", choices = ["0", "1","2","3","4"], type = "index", value = "1")
submit_btn = gr.Button("Analyze")
with gr.Column(visible=True) as output_col:
label = gr.Label(label = "Predicted Label")
local_plot = gr.Plot(label = 'Shap:')
submit_btn.click(
main_func,
[age, sex, cp, trtbps, chol, fbs, restecg, thalachh,exng,oldpeak,slp,caa,thall],
[label,local_plot], api_name="Heart_Predictor"
)
gr.Markdown("### Click on any of the examples below to see how it works:")
gr.Examples([[24,0,4,4,5,5,4,4,5,5,1,2,3], [24,0,4,4,5,3,3,2,1,1,1,2,3]], [age, sex, cp, trtbps, chol, fbs, restecg, thalachh,exng,oldpeak,slp,caa,thall], [label,local_plot], main_func, cache_examples=True)
demo.launch() |