FINLAB

08

Oct 7th, 2019
400
0
Never
Not a member of Pastebin yet? Sign Up, it unlocks many cool features!
Python 2.26 KB | None | 0 0
  1. import matplotlib.pyplot as plt
  2. import pandas as pd
  3. import numpy as np
  4. from sklearn import metrics
  5. from sklearn.preprocessing import StandardScaler
  6. from sklearn.model_selection import train_test_split
  7. from sklearn.neural_network import MLPClassifier
  8. import itertools
  9.  
  10. AllData = pd.read_csv('all_data.csv')  #帶入資料
  11.  
  12. CleanData = AllData.dropna()
  13. CleanData = CleanData.reset_index(drop=True)
  14. CleanData = CleanData.drop('mdate', axis = 1)
  15. CleanData = CleanData.drop('pmkt', axis = 1)
  16.  
  17. X = CleanData.drop('Y', axis = 1)
  18. y = CleanData['Y']
  19. y = y.astype(int)
  20.  
  21. sc = StandardScaler()
  22. X = sc.fit_transform(X)
  23.  
  24. X_train, X_test, y_train, y_test = train_test_split(X, y, test_size=0.3)
  25. MLP= MLPClassifier(solver='adam',learning_rate_init=0.01,hidden_layer_sizes=(40,100),
  26.                    random_state=1)
  27. MLP.fit(X_train, y_train)
  28. print(metrics.classification_report(y_test, MLP.predict(X_test)))#預測出來的結果報告
  29. accuracy = metrics.accuracy_score(y_test, MLP.predict(X_test))
  30. print(accuracy)
  31.  
  32. def plot_confusion_matrix(cm, classes,
  33.                           normalize=False,
  34.                           title='Confusion matrix',
  35.                           cmap=plt.cm.Blues):
  36.     """
  37.    This function prints and plots the confusion matrix.
  38.    Normalization can be applied by setting `normalize=True`.
  39.    """
  40.     plt.imshow(cm, interpolation='nearest', cmap=cmap)
  41.     plt.title(title)
  42.     plt.colorbar()
  43.     tick_marks = np.arange(len(classes))
  44.     plt.xticks(tick_marks, classes, rotation=45)
  45.     plt.yticks(tick_marks, classes)
  46.  
  47.     if normalize:
  48.         cm = cm.astype('float') / cm.sum(axis=1)[:, np.newaxis]
  49.         print("Normalized confusion matrix")
  50.     else:
  51.         print('Confusion matrix, without normalization')
  52.  
  53.     print(cm)
  54.  
  55.     thresh = cm.max() / 2.
  56.     for i, j in itertools.product(range(cm.shape[0]), range(cm.shape[1])):
  57.         plt.text(j, i, cm[i, j],
  58.                  horizontalalignment="center",
  59.                  color="white" if cm[i, j] > thresh else "black")
  60.  
  61.     plt.tight_layout()
  62.     plt.ylabel('True label')
  63.     plt.xlabel('Predicted label')
  64.    
  65. cnf_matrix = metrics.confusion_matrix(y_test,MLP.predict(X_test))
  66. target_names = ['0', '1']
  67. plot_confusion_matrix(cnf_matrix, classes=target_names)
  68. plt.show()
Advertisement
Add Comment
Please, Sign In to add comment