CoolFace
Apppublic

abanerji/str_example3

sourceHugging Facebsdupdated 4y agoView on Hugging Face
0likes
str_example3.py159 linesDownload Raw Back to root
1import streamlit as st2 3import numpy as np 4import umap5import xgboost as xgb6 7import matplotlib.pyplot as plt8from sklearn import datasets9from sklearn.model_selection import train_test_split10 11from sklearn.decomposition import PCA12from sklearn.manifold import TSNE13from sklearn.svm import SVC14from sklearn.neighbors import KNeighborsClassifier15from sklearn.ensemble import RandomForestClassifier16from sklearn.preprocessing import StandardScaler17 18from sklearn.metrics import accuracy_score19 20 21 22st.write("""23# Explore different classifier and datasets24Which one is the best?25""")26 27dataset_name = st.sidebar.selectbox(28    'Select Dataset',29    ('Iris', 'Breast Cancer', 'Wine')30)31 32st.write(f"## {dataset_name} Dataset")33 34classifier_name = st.sidebar.selectbox(35    'Select classifier',36    ('KNN', 'SVM', 'Random Forest', 'XGBOOST')37)38def get_dataset(name):39    data = None40    if name == 'Iris':41        data = datasets.load_iris()42    elif name == 'Wine':43        data = datasets.load_wine()44    else:45        data = datasets.load_breast_cancer()46    X = data.data47    y = data.target48    return X, y49 50X, y = get_dataset(dataset_name)51st.write('Shape of dataset:', X.shape)52st.write('number of classes:', len(np.unique(y)))53 54def add_parameter_ui(clf_name):55    params = dict()56    if clf_name == 'SVM':57        C = st.sidebar.slider('C', 0.01, 10.0)58        params['C'] = C59    elif clf_name == 'KNN':60        K = st.sidebar.slider('K', 1, 15)61        params['K'] = K62    else:63        max_depth = st.sidebar.slider('max_depth', 2, 15)64        params['max_depth'] = max_depth65        n_estimators = st.sidebar.slider('n_estimators', 1, 100)66        params['n_estimators'] = n_estimators67    return params68 69params = add_parameter_ui(classifier_name)70 71def get_classifier(clf_name, params):72    clf = None73    if clf_name == 'SVM':74        clf = SVC(C=params['C'])75    elif clf_name == 'KNN':76        clf = KNeighborsClassifier(n_neighbors=params['K'])77    elif clf_name == 'XGBOOST':78        clf = xgb.XGBClassifier(objective="binary:logistic", 79            n_estimators=params['n_estimators'], 80            max_depth=params['max_depth'], random_state=1234)81    else:82        clf = clf = RandomForestClassifier(n_estimators=params['n_estimators'], 83            max_depth=params['max_depth'], random_state=1234)84        85    return clf86 87clf = get_classifier(classifier_name, params)88#### CLASSIFICATION ####89 90X_train, X_test, y_train, y_test = train_test_split(X, y, test_size=0.2, random_state=1234)91 92clf.fit(X_train, y_train)93y_pred = clf.predict(X_test)94 95acc = accuracy_score(y_test, y_pred)96 97st.write(f'Classifier = {classifier_name}')98st.write(f'Accuracy =', acc)99 100#### PLOT DATASET ####101 102# get the components in pca & t-sne103# Project the data onto the 2 primary principal components104 105scaler = StandardScaler()106# X = scaler.fit_transform(X) 107 108pca = PCA(2)109pca_projected = pca.fit_transform(X)110 111pca1 = pca_projected[:, 0]112pca2 = pca_projected[:, 1]113 114# Use a t-SNE plot now....115tsne = TSNE(n_components=2, verbose=0, random_state=123)116tsne_projected = tsne.fit_transform(X) 117 118ts1 = tsne_projected[:, 0]119ts2 = tsne_projected[:, 1]120 121um_reducer = umap.UMAP()122um_embedding = um_reducer.fit_transform(X)123 124um1 = um_embedding[:, 0]125um2 = um_embedding[:, 1]126 127plt.rcParams.update({'font.size': 7})128fig = plt.figure()129fig, ax = plt.subplots(3,1)130 131pca_points = ax[0].scatter(pca1, pca2,132        c=y, alpha=0.8,133        cmap='viridis')134 135ax[0].set_xlabel('Principal Component 1')136ax[0].set_ylabel('Principal Component 2')137fig.colorbar(pca_points, ax=ax[0])138 139 140ts_points = ax[1].scatter(ts1, ts2,141        c=y, alpha=0.8,142        cmap='viridis')143 144ax[1].set_xlabel('t-sne Component 1')145ax[1].set_ylabel('t-sne Component 2')146fig.colorbar(ts_points, ax=ax[1])147 148um_points = ax[2].scatter(um1, um2,149        c=y, alpha=0.8,150        cmap='viridis')151 152ax[2].set_xlabel('umap Component 1')153ax[2].set_ylabel('umap Component 2')154fig.colorbar(ts_points, ax=ax[2])155 156 157fig.tight_layout()158st.pyplot(fig)159