Team Ai
Apppublic

sklearn-docs/Non_Linear_SVC_for_Binary_Classification

sourceHugging Faceapache-2.0updated 3y agoView on Hugging Face
0likes
app.py119 linesDownload Raw Back to root
1import numpy as np2import matplotlib.pyplot as plt3from sklearn import svm4import gradio as gr5from PIL import Image6 7def calculate_score(clf):8    xx, yy = np.meshgrid(np.linspace(-3, 3, 500), np.linspace(-3, 3, 500))9    X_test = np.c_[xx.ravel(), yy.ravel()]10    Y_test = np.logical_xor(xx.ravel() > 0, yy.ravel() > 0)11    return clf.score(X_test, Y_test)12 13def getColorMap(kernel, gamma):14    # prepare the training dataset15    np.random.seed(0)16    X = np.random.randn(300, 2)17    Y = np.logical_xor(X[:, 0] > 0, X[:, 1] > 0)18 19    # fit the model20    clf = svm.NuSVC(kernel=kernel, gamma=gamma)21    clf.fit(X, Y)22    23    #create a grid for the plotting the decision function24    xx, yy = np.meshgrid(np.linspace(-3, 3, 500), np.linspace(-3, 3, 500))25    26    # plot the decision function for each datapoint on the grid27    Z = clf.decision_function(np.c_[xx.ravel(), yy.ravel()])28    Z = Z.reshape(xx.shape)29 30    plt.figure(figsize=(10, 4))31    plt.imshow(32        Z,33        interpolation="nearest",34        extent=(xx.min(), xx.max(), yy.min(), yy.max()),35        aspect="auto",36        origin="lower",37        cmap=plt.cm.PuOr_r,38    )39    contours = plt.contour(xx, yy, Z, levels=[0], linewidths=2, linestyles="dashed")40    plt.scatter(X[:, 0], X[:, 1], s=30, c=Y, cmap=plt.cm.Paired, edgecolors='k')41    plt.title(f"Decision function for Non-Linear SVC with the {kernel} kernel and '{gamma}' gamma ", fontsize='14')	#title42    plt.xlabel("X",fontsize='13')	#adds a label in the x axis43    plt.ylabel("Y",fontsize='13')	#adds a label in the y axis44    return plt, calculate_score(clf)45 46#XOR_TABLE markdown text47XOR_TABLE = """48<style type="text/css">49.tg  {border-collapse:collapse;border-spacing:10PX;width:50%;margin:auto}50.tg td{border-color:black;border-style:solid;border-width:1px;font-family:Arial, sans-serif;font-size:14px;51  overflow:hidden;padding:10px 5px;word-break:normal;}52.tg th{border-color:black;border-style:solid;border-width:1px;font-family:Arial, sans-serif;font-size:14px;53  font-weight:normal;overflow:hidden;padding:10px 5px;word-break:normal;}54.tg .tg-c3ow{border-color:inherit;text-align:center}55td, th {padding-left: 1rem}56</style>57<BR>58<H3>Table explaining the 'XOR' operator</H3>59<table class="tg">60<thead>61  <tr>62    <th class="tg-c3ow">A</th>63    <th class="tg-c3ow">B</th>64    <th class="tg-c3ow">A XOR B</th>65  </tr>66</thead>67<tbody>68  <tr>69    <td class="tg-c3ow">0</td>70    <td class="tg-c3ow">0</td>71    <td class="tg-c3ow">0</td>72  </tr>73  <tr>74    <td class="tg-c3ow">0</td>75    <td class="tg-c3ow">1</td>76    <td class="tg-c3ow">1</td>77  </tr>78  <tr>79    <td class="tg-c3ow">1</td>80    <td class="tg-c3ow">0</td>81    <td class="tg-c3ow">1</td>82  </tr>83  <tr>84    <td class="tg-c3ow">1</td>85    <td class="tg-c3ow">1</td>86    <td class="tg-c3ow">0</td>87  </tr>88</tbody>89</table>90"""91 92 93with gr.Blocks() as demo:94    gr.Markdown("## Learning the XOR function: An application of Binary Classification using Non-linear SVM")95    gr.Markdown("### This demo is based on this [scikit-learn example](https://scikit-learn.org/stable/auto_examples/svm/plot_svm_nonlinear.html#sphx-glr-auto-examples-svm-plot-svm-nonlinear-py).")96    gr.Markdown("### In this demo, we use a non-linear SVC (Support Vector Classifier) to learn the decision function of the XOR operator.")97                    98    gr.Markdown("### Furthermore, we observe that we get different decision function plots by varying the Kernel and Gamma hyperparameters of the non-linear SVC.")99 100    gr.Markdown("### Feel free to experiment with kernel and gamma values below to see how the quality of the decision function changes with the hyperparameters.")101 102    inp1 = gr.Radio(['poly', 'rbf', 'sigmoid'], label="Kernel", info="Choose a kernel", value="poly")103    inp2 = gr.Radio(['scale', 'auto'], label="Gamma", info="Choose a gamma value", value="scale")104    105    with gr.Row().style(equal_height=True):106        with gr.Column(scale=2):        107            plot = gr.Plot(label=f"Decision function plot")108        with gr.Column(scale=1):109            num = gr.Textbox(label="Test Accuracy")110 111    inp1.change(getColorMap,  inputs=[inp1, inp2], outputs=[plot, num])112    inp2.change(getColorMap,  inputs=[inp1, inp2], outputs=[plot, num])113    demo.load(getColorMap, inputs=[inp1, inp2], outputs=[plot, num])114 115    gr.HTML(XOR_TABLE)116 117 118if __name__ == "__main__":119    demo.launch()