-
Notifications
You must be signed in to change notification settings - Fork 3
Expand file tree
/
Copy pathConfusion_Matrix.py
More file actions
129 lines (113 loc) · 7.22 KB
/
Copy pathConfusion_Matrix.py
File metadata and controls
129 lines (113 loc) · 7.22 KB
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
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
import numpy as np
import matplotlib.pyplot as plt
from sklearn.metrics import confusion_matrix
# test1
# y_pred =torch.tensor([[-1.7389, -0.0242, -2.1111, -6.4881, -2.3681, 4.6529, -0.3168, 8.2081],
# [-4.2011, 8.4311, 5.1541, -8.7524, 3.5309, -5.1222, -0.7664, 1.7569],
# [-3.2211, 6.0119, 0.9448, -3.3833, 4.0109, -9.7344, 3.9883, 1.3600],
# [ 0.2079, -1.3119, -9.7775, -1.2757, 7.6075, 5.9275, -6.2713, 4.8221]])
# y_true = torch.tensor([7,1,1,4])
#
# _, pred=torch.topk(y_pred, 1)
# print('Target: ', y_true, 'Pred: ', pred.squeeze())
#
# pred = pred.t()
# correct = pred.eq(y_true.view(1, -1).expand_as(pred))
# print(correct)
# test2
# y_pred = [0, 0, 0, 0, 1, 1, 1, 1, 2, 2, 2, 2, 3, 3, 3, 3, 4, 4, 4, 4, 5, 5, 5, 5, 6, 6, 6, 6, 7, 7, 7, 7]
# y_true = [0, 0, 0, 0, 1, 1, 1, 1, 2, 2, 2, 2, 3, 3, 3, 3, 4, 4, 4, 4, 5, 5, 5, 5, 6, 6, 6, 6, 7, 7, 7, 7]
# cm = confusion_matrix(y_true, y_pred)
# cm_test = cm.astype('float') / cm.sum(axis=1)[:, np.newaxis]
# print(cm)
# print(cm_test)
y_pred = np.array([5., 1., 0., 2., 5., 1., 2., 2., 0., 7., 7., 3., 3., 3., 2., 5., 5., 2.,
6., 2., 6., 6., 7., 4., 1., 1., 0., 1., 3., 7., 4., 6., 0., 3., 1., 2.,
7., 6., 1., 4., 0., 3., 7., 4., 7., 1., 6., 0., 5., 5., 4., 2., 0., 1.,
1., 1., 4., 3., 2., 3., 3., 6., 0., 6., 2., 6., 1., 0., 4., 4., 7., 1.,
0., 4., 7., 4., 2., 7., 7., 0., 3., 6., 5., 7., 3., 7., 7., 0., 4., 2.,
6., 4., 5., 1., 3., 5., 3., 2., 3., 7., 6., 6., 3., 6., 7., 4., 3., 2.,
4., 0., 3., 1., 1., 0., 2., 2., 3., 6., 6., 7., 4., 1., 5., 2., 2., 0.,
1., 5., 6., 1., 0., 1., 5., 0., 7., 2., 2., 2., 1., 0., 6., 0., 6., 6.,
5., 4., 1., 4., 2., 2., 0., 5., 5., 1., 2., 2., 0., 2., 2., 6., 1., 1.,
5., 4., 2., 2., 4., 2., 7., 5., 4., 5., 1., 6., 4., 1., 1., 1., 1., 1.,
4., 7., 4., 4., 5., 7., 1., 0., 5., 4., 0., 5., 1., 3., 6., 1., 0., 0.,
2., 1., 4., 2., 3., 5., 3., 0., 1., 2., 2., 5., 2., 4., 5., 2., 6., 6.,
2., 2., 1., 2., 0., 5., 5., 7., 4., 2., 5., 0., 4., 2., 3., 1., 1., 1.,
2., 2., 5., 4., 2., 2., 0., 3., 1., 5., 6., 3., 2., 1., 4., 2., 1., 2.,
1., 3., 3., 3., 7., 3., 5., 2., 6., 0., 1., 4., 2., 4., 1., 1., 5., 0.,
5., 2., 6., 2., 5., 3., 1., 4., 2., 5., 2., 4., 1., 0., 3., 2., 1., 5.,
1., 0., 4., 5., 1., 6., 2., 6., 7., 5., 6., 1., 7., 5., 2., 6., 7., 3.,
5., 5., 4., 1., 7., 4., 2., 5., 6., 6., 5., 4., 3., 2., 3., 6., 7., 5.,
2., 3., 3., 0., 3., 6., 5., 1., 2., 1., 3., 6., 3., 6., 1., 6., 4., 0.,
4., 1., 4., 2., 5., 4., 4., 5., 5., 3., 7., 3., 0., 5., 2., 6., 6., 1.,
6., 2., 3., 3., 6., 6., 3., 3., 7., 5., 5., 0., 5., 1., 1., 5., 7., 4.,
4., 2., 1., 2., 6., 3., 0., 5., 1., 3., 3., 3., 4., 2., 1., 3., 2., 4.,
0., 2., 4., 4., 4., 1., 4., 4., 6., 3., 4., 0., 1., 5., 6., 7., 6., 6.,
3., 6., 2., 6., 7., 0., 1., 1., 5., 4., 5., 2., 2., 7., 4., 2., 6., 5.,
2., 2., 2., 4., 2., 6., 5., 3., 1., 0., 7., 5., 4., 7., 7., 5., 7., 1.,
7., 1., 0., 7., 6., 3., 6., 2., 1., 0., 3., 3., 3., 4., 1., 7., 2., 7.,
1., 5., 5., 4., 7., 5., 6., 5., 1., 7., 7., 6., 5., 4., 5., 4., 1., 4.,
3., 1., 2., 4., 1., 2., 4., 6., 4., 1., 2., 6., 3., 3., 2., 3., 1., 0.,
4., 3., 5., 2., 3., 6., 7., 7., 0., 2., 6., 7., 2., 3., 0., 7., 7., 0.,
7., 6., 6., 7., 4., 7., 1., 7., 6., 4., 5., 6., 7., 4., 6., 6., 5., 7.,
1., 4., 2., 7., 0., 2., 2., 2., 2., 4., 2., 3., 7., 1., 3., 2., 2., 7.,
0., 3., 2., 7., 3., 2., 7., 7., 0., 6., 1., 6., 7., 3., 7., 7., 4., 1.,
4., 3., 6., 2., 3., 2., 4., 2., 0., 7., 7., 3., 1., 4., 2., 4., 4., 2.,
5., 3., 7., 7., 3., 3.])
y_true = np.array([5., 1., 3., 2., 5., 1., 2., 2., 5., 7., 7., 3., 3., 3., 2., 5., 5., 2.,
6., 2., 6., 6., 7., 4., 1., 1., 0., 1., 3., 4., 4., 6., 3., 3., 1., 2.,
4., 6., 1., 4., 0., 3., 7., 4., 7., 1., 6., 0., 5., 5., 4., 2., 0., 1.,
1., 1., 4., 3., 2., 3., 5., 6., 0., 6., 2., 6., 1., 3., 4., 4., 7., 1.,
0., 4., 7., 4., 2., 7., 7., 5., 3., 6., 7., 7., 5., 7., 7., 3., 4., 2.,
6., 4., 5., 1., 5., 5., 5., 2., 5., 7., 6., 6., 3., 6., 7., 4., 3., 2.,
4., 0., 5., 1., 1., 0., 2., 2., 3., 6., 6., 7., 4., 1., 5., 7., 2., 3.,
1., 6., 6., 7., 3., 1., 5., 0., 7., 2., 7., 2., 1., 0., 6., 0., 6., 6.,
6., 4., 1., 4., 7., 3., 3., 5., 5., 1., 2., 2., 0., 2., 2., 6., 1., 1.,
6., 4., 2., 7., 4., 2., 6., 5., 4., 5., 3., 6., 4., 7., 1., 1., 1., 1.,
4., 6., 4., 4., 5., 5., 7., 0., 5., 4., 3., 5., 7., 3., 6., 3., 0., 0.,
2., 1., 4., 2., 6., 5., 3., 0., 3., 7., 2., 6., 7., 4., 5., 3., 6., 6.,
2., 7., 1., 2., 3., 5., 5., 7., 4., 3., 5., 3., 4., 2., 3., 1., 3., 1.,
2., 7., 5., 4., 7., 7., 0., 3., 1., 5., 6., 3., 2., 1., 4., 4., 1., 4.,
1., 3., 3., 3., 7., 3., 5., 2., 6., 0., 7., 4., 2., 4., 1., 1., 5., 0.,
5., 2., 6., 2., 7., 3., 1., 4., 2., 5., 2., 4., 1., 0., 3., 2., 7., 5.,
1., 0., 4., 5., 1., 6., 2., 6., 7., 5., 6., 7., 7., 5., 2., 6., 7., 5.,
7., 7., 4., 1., 7., 4., 2., 5., 6., 6., 7., 4., 5., 2., 3., 6., 7., 7.,
2., 3., 3., 0., 3., 6., 5., 1., 2., 1., 3., 6., 3., 6., 1., 6., 4., 0.,
4., 1., 4., 2., 5., 4., 4., 5., 5., 3., 7., 3., 0., 7., 2., 6., 6., 1.,
6., 2., 3., 3., 6., 6., 3., 3., 7., 5., 5., 0., 5., 1., 1., 5., 7., 4.,
4., 2., 1., 2., 6., 3., 0., 5., 1., 3., 3., 3., 4., 2., 1., 3., 2., 4.,
0., 2., 4., 4., 4., 1., 4., 4., 6., 3., 4., 0., 1., 5., 6., 7., 6., 6.,
3., 6., 2., 6., 7., 0., 1., 1., 5., 4., 5., 2., 2., 7., 4., 2., 6., 5.,
2., 2., 2., 4., 2., 6., 5., 3., 1., 0., 7., 5., 4., 7., 7., 5., 7., 1.,
7., 1., 0., 7., 6., 3., 6., 2., 1., 0., 3., 3., 3., 4., 1., 7., 2., 7.,
1., 5., 5., 4., 7., 5., 6., 5., 1., 7., 7., 6., 5., 4., 5., 4., 1., 4.,
3., 1., 2., 6., 1., 2., 4., 6., 4., 1., 2., 6., 5., 3., 2., 3., 1., 0.,
4., 3., 5., 1., 3., 6., 5., 7., 0., 2., 6., 7., 2., 3., 0., 5., 7., 0.,
7., 6., 6., 7., 4., 5., 1., 7., 6., 4., 5., 6., 5., 4., 6., 6., 5., 5.,
1., 6., 2., 7., 0., 1., 2., 1., 2., 4., 1., 3., 7., 1., 3., 2., 2., 7.,
0., 3., 1., 5., 3., 2., 7., 7., 0., 6., 1., 6., 5., 3., 7., 7., 4., 1.,
4., 5., 6., 2., 3., 2., 4., 1., 0., 7., 5., 3., 6., 4., 2., 4., 4., 2.,
5., 3., 7., 7., 3., 3.])
cm = confusion_matrix(y_true, y_pred)
cm_ccma = cm.astype('float') / cm.sum(axis=1)[:, np.newaxis]
cm=cm_ccma
labels = ['Neutral', 'Calm', 'Happy', 'Sad', 'Angry', 'Fearful', 'Disgust', 'Surprise']
plt.imshow(cm, interpolation='nearest', cmap=plt.cm.Blues)
plt.title('Confusion Matrix')
plt.colorbar()
tick_marks = np.arange(len(labels))
plt.xticks(tick_marks, labels, rotation=45)
plt.yticks(tick_marks, labels)
thresh = cm.max() / 2.0
for i in range(len(labels)):
for j in range(len(labels)):
plt.text(j, i, '{:.1%}'.format(cm[i, j], 'd'),
horizontalalignment="center",
color="white" if cm[i, j] > thresh else "black",
fontsize=8)
plt.ylabel('True Label')
plt.xlabel('Predicted Label')
plt.tight_layout()
plt.savefig('ccma.jpg',dpi=300)
plt.show()