Pytorch构建栈式自编码器实现以图搜图任务(以cifar10做数据集)

(Pytorch构建栈式自编码器实现以图搜图任务)

本文旨在使用CIFAR-10数据集,构建与训练栈式自编码器,提取数据集中图像的特征;基于所提取的特征完成CIFAR-10中任意图像的检索任务并展示效果。

搞清楚pytorch与tensorflow区别

pytorch

学习文档 pytorch是一种python科学计算框架 作用:

  • 无缝替换numpy,通过GPU实现神经网络的加速
  • 通过自动微分机制,让神经网络实现更容易(即自动求导机制)

张量:类似于数组和矩阵,是一种特殊的数据结构。在pytorch中,神经网络的输入、输出以及网络的参数等数据,都是使用张量来进行描述的。 张量的基本方法,持续更新

每个变量中都有两个标志:requires_grad volatile requires_grad: 如果有一个单一的输入操作需要梯度,它的输出就需要梯度。只有所有输入都不需要梯度时,输出才不需要。 volatile: 只需要一个volatile的输入就会得到一个volatile输出。

tensorflow

学习文档 TensorFlow 是由 Google Brain 团队为深度神经网络(DNN)开发的功能强大的开源软件库 TensorFlow 则还有更多的特点,如下:

  • 支持所有流行语言,如 Python、C++、Java、R和Go。
  • 可以在多种平台上工作,甚至是移动平台和分布式平台。
  • 它受到所有云服务(AWS、Google和Azure)的支持。
  • Keras——高级神经网络 API,已经与 TensorFlow 整合。
  • 与Torch/Theano 比较,TensorFlow 拥有更好的计算图表可视化。 允
  • 许模型部署到工业生产中,并且容易使用。
  • 有非常好的社区支持。 TensorFlow 不仅仅是一个软件库,它是一套包括 TensorFlow,TensorBoard 和TensorServing 的软件。

搞清楚栈式自编码器的内部原理

在这里插入图片描述   我们构建栈式编码器,用编码器再解码出来的结果和原标签对比进行训练模型,然后用中间编码提取到的特征直接和原图的特征进行对比,得到相似度,实现以图搜图。 整个网络的训练不是一蹴而就的,而是逐层进行的。 在这里插入图片描述 在这里插入图片描述 在这里插入图片描述

效果图

随机取测试集的五张图片,进行以图搜图(TOP8)

提取的分布式特征聚集图像:第一张为原图散点图,第二张以检索的TOP8的TOP1的提取特征散点图为例 在这里插入图片描述 在这里插入图片描述

代码及效果图

欠完备编码器

1# -*- coding: utf-8 -*- 2""" 3Created on Sat Apr 24 18:37:55 2021 4 5@author: ASUS 6""" 7 8import torch 9import torchvision 10import torch.utils.data 11import torch.nn as nn 12import matplotlib.pyplot as plt 13import random #随机取测试集的图片 14import time 15starttime = time.time() 16 17torch.manual_seed(1) 18EPOCH = 10 19BATCH_SIZE = 64 20LR = 0.005 21 22trainset = torchvision.datasets.CIFAR10( 23 root='./data', 24 train=True, 25 transform=torchvision.transforms.ToTensor(), 26 download=False) 27testset = torchvision.datasets.CIFAR10( 28 root='./data', 29 train=False, 30 transform=torchvision.transforms.ToTensor(), 31 download=False) 32 33# dataloaders 34trainloader = torch.utils.data.DataLoader(trainset, batch_size=BATCH_SIZE, 35 shuffle=True) 36 37testloader = torch.utils.data.DataLoader(testset, batch_size=BATCH_SIZE, 38 shuffle=True) 39 40train_data = torchvision.datasets.MNIST( 41 root='./data', 42 train=True, 43 transform=torchvision.transforms.ToTensor(), 44 download=False 45) 46loader = torch.utils.data.DataLoader(dataset=train_data, batch_size=BATCH_SIZE, shuffle=True) 47 48class Stack_AutoEncoder(nn.Module): 49 def __init__(self): 50 super(Stack_AutoEncoder,self).__init__() 51 self.encoder = nn.Sequential( 52 nn.Linear(32*32,256), 53 nn.Tanh(), 54 nn.Linear(256, 128), 55 nn.Tanh(), 56 nn.Linear(128, 32), 57 nn.Tanh(), 58 nn.Linear(32, 16), 59 nn.Tanh(), 60 nn.Linear(16, 8) 61 ) 62 self.decoder = nn.Sequential( 63 nn.Linear(8, 16), 64 nn.Tanh(), 65 nn.Linear(16, 32), 66 nn.Tanh(), 67 nn.Linear(32, 128), 68 nn.Tanh(), 69 nn.Linear(128, 256), 70 nn.Tanh(), 71 nn.Linear(256, 32*32), 72 nn.Sigmoid() 73 74 ) 75 def forward(self, x): 76 encoded = self.encoder(x) 77 decoded = self.decoder(encoded) 78 return encoded,decoded 79 80Coder = Stack_AutoEncoder() 81print(Coder) 82 83optimizer = torch.optim.Adam(Coder.parameters(),lr=LR) 84loss_func = nn.MSELoss() 85 86for epoch in range(EPOCH): 87 for step,(x,y) in enumerate(trainloader): 88 b_x = x.view(-1,32*32) 89 b_y = x.view(-1,32*32) 90 b_label = y 91 encoded , decoded = Coder(b_x) 92# print(encoded) 93 loss = loss_func(decoded,b_y) 94 95 optimizer.zero_grad() 96 loss.backward() 97 optimizer.step() 98 99# if step%5 == 0: 100 print('Epoch :', epoch,'|','train_loss:%.4f'%loss.data) 101 102torch.save(Coder,'Stack_AutoEncoder.pkl') 103print('________________________________________') 104print('finish training') 105 106endtime = time.time() 107print('训练耗时:',(endtime - starttime)) 108 109#以图搜图函数 110Coder = Stack_AutoEncoder() 111Coder = torch.load('Stack_AutoEncoder.pkl') 112def search_by_image(x,inputImage,K): 113 c = ['b','g','r'] #画特征散点图 114 loss_func = nn.MSELoss() 115 x_ = inputImage.view(-1,32*32) 116 encoded , decoded = Coder(x_) 117# print(encoded) 118 lossList=[] 119 for step,(test_x,y) in enumerate(testset): 120 if(step == x): #去掉原图 121 lossList.append((x,1)) 122 continue 123 b_x = test_x.view(-1,32*32) 124 b_y = test_x.view(-1,32*32) 125 b_label = y 126 test_encoded , test_decoded = Coder(b_x) 127 128 loss = loss_func(encoded,test_encoded) 129# loss = round(loss, 4) #保留小数 130 lossList.append((step,loss.item())) 131 lossList=sorted(lossList,key=lambda x:x[1],reverse=False)[:K] 132 print(lossList) 133 plt.figure(1) 134# plt.figure(figsize=(10, 10)) 135 trueImage = inputImage.reshape((3, 32, 32)).transpose(0,2) 136 plt.imshow(trueImage) 137 138 plt.title('true') 139 plt.show() 140 for j in range(K): 141 showImage = testset[lossList[j][0]][0] #遍历相似度最高列表里的图 142 showImage = showImage.reshape((3, 32, 32)).transpose(0,2) 143 plt.subplots_adjust(left=4, right=5) #好像没起作用 144 plt.subplots(figsize=(8, 8)) 145 plt.subplot(2,4,j+1) 146 plt.title("img" + str(lossList[j][0])+"loss:"+str(round(lossList[j][1],5))) 147 plt.imshow(showImage) 148 149 plt.show() 150 151 #特征散点图 只显示第一个相似度最高的特征散点图聚集关系 152 y_li = encoded.detach() 153 x_li = [x for x in range(8)] 154 for i in range(len(encoded)): 155 plt.scatter(x_li, y_li[i],c = c[i]) 156 plt.show() 157 sim = testset[lossList[j][0]][0].view(-1,32*32) 158 sim_encoded , _d = Coder(sim) 159# print(sim_encoded) 160 sim_li = sim_encoded.detach() #将torch转为numpy画图要加上detach 161 x_li = [x for x in range(8)] 162 for i in range(len(encoded)): 163 plt.scatter(x_li, sim_li[i],c = c[i]) 164 plt.show() 165 166for i in range(5): 167 x = random.randint(0, len(testset)) 168 print(x) 169 im,_ = testset[x] 170 search_by_image(x,inputImage = im,K = 8) 171# break 172

Epoch : 0 | train_loss:0.0536 Epoch : 1 | train_loss:0.0411 Epoch : 2 | train_loss:0.0293 Epoch : 3 | train_loss:0.0274 Epoch : 4 | train_loss:0.0339 Epoch : 5 | train_loss:0.0337 Epoch : 6 | train_loss:0.0313 Epoch : 7 | train_loss:0.0338 Epoch : 8 | train_loss:0.0279 Epoch : 9 | train_loss:0.0289


finish training 训练耗时: 395.6197159290314 在这里插入图片描述 在这里插入图片描述 在这里插入图片描述 在这里插入图片描述 在这里插入图片描述 在这里插入图片描述

卷积栈式编码器

1import numpy as np 2import torch 3import torchvision 4import torch.nn as nn 5import torchvision.transforms as transforms 6import matplotlib.pyplot as plt 7from torch.utils.data import DataLoader 8import torch.optim as optim 9import os 10class layer1(nn.Module): 11 def __init__(self): 12 super(layer1,self).__init__() 13 self.encoder=nn.Sequential( 14 nn.Conv2d(3, 16, kernel_size=5), # 16*28*28 15 nn.BatchNorm2d(16), 16 nn.ReLU(inplace=True), 17 # nn.MaxPool2d(kernel_size=2,stride=2)#16*15*15 18 ) 19 self.decoder=nn.Sequential( 20 nn.ConvTranspose2d(16,3,kernel_size=5,stride=1), 21 nn.BatchNorm2d(3), 22 nn.ReLU(inplace=True) 23 24 ) 25 26 27 def forward(self,x): 28 encode=self.encoder(x) 29 decode=self.decoder(encode) 30 return encode,decode 31 32class layer2(nn.Module): 33 def __init__(self,layer1): 34 super(layer2,self).__init__() 35 self.layer1=layer1 36 self.encoder=nn.Sequential( 37 nn.Conv2d(16, 10, kernel_size=5), # 10*24*24 38 nn.BatchNorm2d(10), 39 nn.ReLU(inplace=True), 40 # nn.MaxPool2d(kernel_size=2,stride=2)#10*6*6 41 ) 42 self.decoder=nn.Sequential( 43 nn.ConvTranspose2d(10,3,kernel_size=9,stride=1), 44 nn.BatchNorm2d(3), 45 nn.ReLU(inplace=True) 46 ) 47 48 def forward(self,x): 49 self.layer1.eval() 50 x,_=self.layer1(x) 51 encode=self.encoder(x) 52 decode=self.decoder(encode) 53 return encode,decode 54 55class layer3(nn.Module): 56 def __init__(self,layer2): 57 super(layer3,self).__init__() 58 self.layer2=layer2 59 self.encoder=nn.Sequential( 60 nn.Conv2d(10, 5, kernel_size=5), # 5*20*20 61 nn.BatchNorm2d(5), 62 nn.ReLU(inplace=True) 63 ) 64 self.decoder=nn.Sequential( 65 nn.ConvTranspose2d(5,3,kernel_size=13,stride=1), 66 nn.BatchNorm2d(3), 67 nn.ReLU(inplace=True) 68 ) 69 70 def forward(self,x): 71 self.layer2.eval() 72 x,_=self.layer2(x) 73 encode=self.encoder(x) 74 decode=self.decoder(encode) 75 return encode,decode 76 77def train_layer(layer,k): 78 loss_fn = torch.nn.MSELoss().to(device) 79 optimizer = optim.Adam(layer.parameters(), lr=0.01) 80 for epoch in range(10): 81 i = 0 82 for data, target in trainloader: 83 data,target=data.to(device),target.to(device) 84 encoded, decoded = layer(data) 85 86 loss = loss_fn(decoded, data) 87 loss.backward() 88 optimizer.step() 89 optimizer.zero_grad() 90 if i % 50 == 0: 91 print(loss) 92 i += 1 93 print("epoch:%d,loss:%f"%(epoch,loss)) 94 torch.save(layer.state_dict(), '卷积model/layer%d.pkl'%k) 95 96def search_pic(test_dataset,input_img,input_label,K=8): 97 98 model=layer3 99 loss_fn = nn.MSELoss() 100 input_img=input_img.to(device) 101 input_img=input_img.unsqueeze(0) 102 inputEncode,inputDecoder = model(input_img) 103 lossList = [] 104 for (i, (testImage,_)) in enumerate(test_dataset): 105 106 testImage=testImage.to(device) 107 testImage=testImage.unsqueeze(0) 108 testEncode,testDecoder = model(testImage) 109 enLoss = loss_fn(inputEncode, testEncode) 110 lossList.append((i, np.sqrt(enLoss.item()))) 111 lossList = sorted(lossList, key=lambda x: x[1], reverse=False)[:K] 112 input_img=input_img.squeeze(0) 113 input_img=input_img.to(torch.device("cpu")) 114 npimg = input_img.numpy() 115 npimg = npimg /2 +0.5 116 plt.imshow(np.transpose(npimg, (1, 2, 0))) # transpose() 117 plt.show() 118 search_labels=[] 119 search_dises=[] 120 k=0 121 for i,dis in lossList: 122 search_dises.append(dis) 123 plt.subplot(1, 8, k + 1) 124 img=test_dataset[i][0].numpy() 125 search_labels.append(test_dataset[i][1]) 126 plt.imshow(np.transpose(img, (1, 2, 0))) 127 k+=1 128 plt.show() 129 print("input label:",input_label) 130 print("search labels:",search_labels) 131 print("search distence:",search_dises) 132if __name__ == '__main__': 133# device=torch.device("cuda:0") 134 device = torch.device("cuda"if torch.cuda.is_available() else "cpu") 135 transform = transforms.Compose( 136 [transforms.ToTensor(), 137 transforms.Normalize((0.5, 0.5, 0.5), (0.5, 0.5, 0.5))]) 138 139 train_dataset = torchvision.datasets.CIFAR10(root='data/', 140 train=True, 141 transform=transform, 142 download=False) 143 test_dataset = torchvision.datasets.CIFAR10(root='data/', 144 train=False, 145 transform=transform, 146 download=False) 147 148 trainloader = DataLoader(train_dataset, batch_size=256, shuffle=True) 149 layer1 = layer1().to(device) 150 if os.path.exists('卷积model/layer1.pkl'): 151 layer1.load_state_dict(torch.load("卷积model/layer1.pkl")) 152 else: 153 train_layer(layer1, 1) 154 155 layer2 = layer2(layer1).to(device) 156 157 if os.path.exists('卷积model/layer2.pkl'): 158 layer2.load_state_dict(torch.load("卷积model/layer2.pkl")) 159 else: 160 train_layer(layer2, 2) 161 162 layer3 = layer3(layer2).to(device) 163 if os.path.exists('卷积model/layer3.pkl'): 164 layer3.load_state_dict(torch.load("卷积model/layer3.pkl")) 165 else: 166 train_layer(layer3, 3) 167 168 search_pic(test_dataset,train_dataset[0][0],train_dataset[0][1])

在这里插入图片描述 我把生成和模型和代码会放到资源上,方便大家下载

栈式编码器

1import numpy as np 2import torch 3import torchvision 4import torch.nn as nn 5import torchvision.transforms as transforms 6import matplotlib.pyplot as plt 7from torch.utils.data import DataLoader 8import torch.optim as optim 9import os 10 11 12 13class Layer1(nn.Module): 14 def __init__(self,hidden_size): 15 super(Layer1, self).__init__() 16 self.hidden_size=hidden_size 17 self.encoder = nn.Linear(3 * 32 * 32, hidden_size) 18 self.decoder = nn.Linear(hidden_size, 3 * 32 * 32) 19 20 def forward(self, x): 21 x = x.view(-1, 3 * 32 * 32) 22 encoded = self.encoder(x) 23 decoded = self.decoder(encoded) 24 return encoded, decoded 25class Layer2(nn.Module): 26 def __init__(self,layer1,hidden_size): 27 super(Layer2, self).__init__() 28 self.layer1=layer1 29 self.hidden_size=hidden_size 30 self.encoder = nn.Linear(layer1.hidden_size, hidden_size) 31 self.decoder = nn.Sequential( 32 nn.Linear(hidden_size, self.layer1.hidden_size), 33 nn.Linear(self.layer1.hidden_size,3*32*32) 34 ) 35 36 def forward(self, x): 37 #保证前一层参数不变 38 self.layer1.eval() 39 x = x.view(-1, 3 * 32 * 32) 40 x,_=layer1(x) 41 encoded = self.encoder(x) 42 decoded = self.decoder(encoded) 43 return encoded, decoded 44class Layer3(nn.Module): 45 def __init__(self,layer2,hidden_size): 46 super(Layer3, self).__init__() 47 self.layer2=layer2 48 self.hidden_size=hidden_size 49 self.encoder = nn.Linear(layer2.hidden_size, hidden_size) 50 self.decoder = nn.Sequential( 51 nn.Linear(hidden_size, self.layer2.hidden_size), 52 nn.Linear(self.layer2.hidden_size,self.layer2.layer1.hidden_size), 53 nn.Linear(self.layer2.layer1.hidden_size,3*32*32) 54 ) 55 56 def forward(self, x): 57 #保证前一层参数不变 58 self.layer2.eval() 59 x = x.view(-1, 3 * 32 * 32) 60 x,_=self.layer2(x) 61 encoded = self.encoder(x) 62 decoded = self.decoder(encoded) 63 return encoded, decoded 64 65def train_layer(layer,k): 66 loss_fn = torch.nn.MSELoss().to(device) 67 optimizer1 = optim.Adam(layer.parameters(), lr=0.01) 68 for epoch in range(20): 69 i = 0 70 for data, target in trainloader: 71 data,target=data.to(device),target.to(device) 72 encoded, decoded = layer(data) 73 label = data.view(-1, 3072) 74 loss = loss_fn(decoded, label) 75 loss.backward() 76 optimizer1.step() 77 optimizer1.zero_grad() 78 if i % 50 == 0: 79 print(loss) 80 i += 1 81 print("epoch:%d,loss:%f"%(epoch,loss)) 82 torch.save(layer.state_dict(), '全连接model/layer%d.pkl'%k) 83def search_pic(test_dataset,input_img,input_label,K=8): 84 85 model=layer3 86 loss_fn = nn.MSELoss() 87 input_img=input_img.to(device) 88 input_img=input_img.unsqueeze(0) 89 inputEncode,inputDecoder = model(input_img) 90 lossList = [] 91 for (i, (testImage,_)) in enumerate(test_dataset): 92 93 testImage=testImage.to(device) 94 testEncode,testDecoder = model(testImage) 95 96 enLoss = loss_fn(inputEncode, testEncode) 97 lossList.append((i, np.sqrt(enLoss.item()))) 98 lossList = sorted(lossList, key=lambda x: x[1], reverse=False)[:K] 99 input_img=input_img.squeeze(0) 100 input_img=input_img.to(torch.device("cpu")) 101 npimg = input_img.numpy() 102 npimg = npimg / 2 + 0.5 103 plt.imshow(np.transpose(npimg, (1, 2, 0))) # transpose() 104 plt.show() 105 search_labels=[] 106 search_dises=[] 107 k=0 108 for i,dis in lossList: 109 search_dises.append(dis) 110 plt.subplot(1, 8, k + 1) 111 img=test_dataset[i][0].numpy() 112 search_labels.append(test_dataset[i][1]) 113 plt.imshow(np.transpose(img, (1, 2, 0))) 114 k+=1 115 plt.show() 116 print("input label:",input_label) 117 print("search labels:",search_labels) 118 print("search distence:",search_dises) 119if __name__ == '__main__': 120# device=torch.device("cuda:0") 121 device = torch.device("cuda"if torch.cuda.is_available() else "cpu") 122 transform = transforms.Compose( 123 [transforms.ToTensor(), 124 transforms.Normalize((0.5, 0.5, 0.5), (0.5, 0.5, 0.5))]) 125 126 train_dataset = torchvision.datasets.CIFAR10(root='data/', 127 train=True, 128 transform=transform, 129 download=False) 130 trainloader = DataLoader(train_dataset, batch_size=256, shuffle=True) 131 132 133 test_dataset = torchvision.datasets.CIFAR10(root='data/', 134 train=False, 135 transform=transform, 136 download=False) 137 138 139 layer1=Layer1(2048).to(device) 140 if os.path.exists('全连接model/layer1.pkl'): 141 layer1.load_state_dict(torch.load("全连接model/layer1.pkl")) 142 else: 143 train_layer(layer1,1) 144 145 layer2=Layer2(layer1,1024).to(device) 146 147 if os.path.exists('全连接model/layer2.pkl'): 148 layer2.load_state_dict(torch.load("全连接model/layer2.pkl")) 149 else: 150 train_layer(layer2,2) 151 152 layer3 = Layer3(layer2, 512).to(device) 153 if os.path.exists('全连接model/layer3.pkl'): 154 layer3.load_state_dict(torch.load("全连接model/layer3.pkl")) 155 else: 156 train_layer(layer3, 3) 157 158 159 160 161 search_pic(test_dataset,train_dataset[8][0],train_dataset[8][1])
点赞
收藏

评论区

加载中...

相关推荐

MySQL:[Err] 1292 - Incorrect datetime value: ‘0000-00-00 00:00:00‘ for column ‘CREATE_TIME‘ at row 1

文章目录问题用navicat导入数据时,报错:原因这是因为当前的MySQL不支持datetime为0的情况。解决修改sql\mode:sql\mode:SQLMode定义了MySQL应支持的SQL语法、数据校验等,这样可以更容易地在不同的环境中使用MySQL。全局s

Oracle 分组与拼接字符串同时使用

SELECTT.,ROWNUMIDFROM(SELECTT.EMPLID,T.NAME,T.BU,T.REALDEPART,T.FORMATDATE,SUM(T.S0)S0,MAX(UPDATETIME)CREATETIME,LISTAGG(TOCHAR(

2020年前端实用代码段,为你的工作保驾护航

有空的时候,自己总结了几个代码段,在开发中也经常使用,谢谢。1、使用解构获取json数据let jsonData  id: 1,status: "OK",data: 'a', 'b';let  id, status, data: number   jsonData;console.log(id, status, number )

Python3:sqlalchemy对mysql数据库操作,非sql语句

Python3:sqlalchemy对mysql数据库操作,非sql语句python3authorlizmdatetime2018020110:00:00coding:utf8'''

Nginx + lua +[memcached,redis]

精品案例1、Nginxluamemcached,redis实现网站灰度发布2、分库分表/基于Leaf组件实现的全球唯一ID(非UUID)3、Redis独立数据监控,实现订单超时操作/MQ死信操作SelectPollEpollReactor模型4、分布式任务调试Quartz应用

CIFAR

1、CIFAR10,是一个用于做图像分类研究的数据集。由60000个图片组成6万个图片中,5万张用于训练,1万张用于测试每个图片是32x32像素所有图片可以分成10类每个图片都有一个标签,标记属于哪一个类测试集中一个类对应1000张图训练集中将5万张图分为5份类之间的图片是互