MCPcopy Create free account
hub / github.com/subbarayudu-j/TheAlgorithms-Python / CNN

Class CNN

neural_network/convolution_neural_network.py:23–299  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

21import matplotlib.pyplot as plt
22
23class CNN():
24
25 def __init__(self,conv1_get,size_p1,bp_num1,bp_num2,bp_num3,rate_w=0.2,rate_t=0.2):
26 '''
27 :param conv1_get: [a,c,d],size, number, step of convolution kernel
28 :param size_p1: pooling size
29 :param bp_num1: units number of flatten layer
30 :param bp_num2: units number of hidden layer
31 :param bp_num3: units number of output layer
32 :param rate_w: rate of weight learning
33 :param rate_t: rate of threshold learning
34 '''
35 self.num_bp1 = bp_num1
36 self.num_bp2 = bp_num2
37 self.num_bp3 = bp_num3
38 self.conv1 = conv1_get[:2]
39 self.step_conv1 = conv1_get[2]
40 self.size_pooling1 = size_p1
41 self.rate_weight = rate_w
42 self.rate_thre = rate_t
43 self.w_conv1 = [np.mat(-1*np.random.rand(self.conv1[0],self.conv1[0])+0.5) for i in range(self.conv1[1])]
44 self.wkj = np.mat(-1 * np.random.rand(self.num_bp3, self.num_bp2) + 0.5)
45 self.vji = np.mat(-1*np.random.rand(self.num_bp2, self.num_bp1)+0.5)
46 self.thre_conv1 = -2*np.random.rand(self.conv1[1])+1
47 self.thre_bp2 = -2*np.random.rand(self.num_bp2)+1
48 self.thre_bp3 = -2*np.random.rand(self.num_bp3)+1
49
50
51 def save_model(self,save_path):
52 #save model dict with pickle
53 import pickle
54 model_dic = {'num_bp1':self.num_bp1,
55 'num_bp2':self.num_bp2,
56 'num_bp3':self.num_bp3,
57 'conv1':self.conv1,
58 'step_conv1':self.step_conv1,
59 'size_pooling1':self.size_pooling1,
60 'rate_weight':self.rate_weight,
61 'rate_thre':self.rate_thre,
62 'w_conv1':self.w_conv1,
63 'wkj':self.wkj,
64 'vji':self.vji,
65 'thre_conv1':self.thre_conv1,
66 'thre_bp2':self.thre_bp2,
67 'thre_bp3':self.thre_bp3}
68 with open(save_path, 'wb') as f:
69 pickle.dump(model_dic, f)
70
71 print('Model saved: %s'% save_path)
72
73 @classmethod
74 def ReadModel(cls,model_path):
75 #read saved model
76 import pickle
77 with open(model_path, 'rb') as f:
78 model_dic = pickle.load(f)
79
80 conv_get= model_dic.get('conv1')

Callers 1

ReadModelMethod · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected