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()