| 21 | import matplotlib.pyplot as plt |
| 22 | |
| 23 | class 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') |