兄弟萌,我咕里個咚今天又殺回來了,有幾天時間可以不用駐場了,喜大普奔,終于可以在有網的地方碼代碼了,最近駐場也是又熱又心累啊,抓緊這幾天,再更新一點的新東西,
今天主要講一下非監督學習,你可能要問了,什么是非監督學習,我的理解就是不會給樣本標簽的,它本質上是一個統計手段,在沒有標簽的資料里可以發現潛在的一些結構的一種訓練方式,這個可以用來干什么,舉個例子,在工業場景瑕疵檢測的運用中,由于良品的數量遠遠高于不良品的數量,如果這個時候你要采用監督學習,那么收集樣本的時間就多得嚇人了,可能你樣本還沒有收集完全,產品都已經做完下線了,所以,你就撓頭吧,于是,非監督學習就迎來了一片藍海,但是即使是藍海也要你的船能開才行,這里面也不乏調整,比如非監督學習的效果就不太好去評估,當然,我想了點辦法,在工業領域也有所應用了,但是過殺有點高,大概5%左右,這5%的過殺,再通過后期的監督演算法,其實可以解決很多問題了,這樣工業上的大部分問題,都可以有所緩解了,我真棒,哈哈哈
1.非監督學習網路架構
先提供一下,我的非監督學習的網路架構,還是基于pytorch來寫的,我給這個網路一個名稱叫做
piercing eye,話不多說,上代碼,
from torch import nn
import torch
class CBP(nn.Module):
"""
conv + batchnormal + prelu
"""
def __init__(self,inc,ouc):
super().__init__()
self.block1=nn.Sequential(
nn.Conv2d(inc,ouc,3,1,1),
nn.BatchNorm2d(ouc),
nn.PReLU()
)
def forward(self,y):
return self.block1(y)
class Up_Block(nn.Module):
def __init__(self, in_channel, out_channel):
super().__init__()
self.block1=nn.Sequential(
nn.ConvTranspose2d(in_channel, out_channel, 3, 2, 1, 1),
nn.BatchNorm2d(out_channel),
nn.PReLU()
)
def forward(self,y):
return self.block1(y)
class Down_Block(nn.Module):
def __init__(self, in_channel, out_channel):
super().__init__()
self.block1=nn.Sequential(
nn.Conv2d(in_channel, out_channel, 5, 2, padding=2),
nn.BatchNorm2d(out_channel),
nn.PReLU()
)
def forward(self,y):
return self.block1(y)
class PiercingEye(nn.Module):
def __init__(self):
super().__init__()
self.block1=nn.Sequential(
CBP(3, 4),
Down_Block(4, 8),
Down_Block(8, 16),
Down_Block(16, 32),
Down_Block(32, 64),
Down_Block(64, 128),
Down_Block(128, 256),
CBP(256, 256),
Up_Block(256, 128),
Up_Block(128, 64),
Up_Block(64, 32),
Up_Block(32, 16),
Up_Block(16, 8),
Up_Block(8, 4),
nn.Conv2d(4,3,1),
nn.Tanh()
)
def forward(self,y):
return self.block1(y)
if __name__ == '__main__':
net = PiercingEye()
x = torch.Tensor(2,3,512,512)
y = net(x)
print(y.shape)
簡單的說明一下,其實就是做了6次下采樣和6次上采樣,也就是AE網路,中間沒有任何跳躍連接,也可以理解成是一個生成網路,
2.資料集準備
我直接給大家一個百度云的鏈接,這也是一個開源的資料集,我稍微整理了一下,方便大家使用
鏈接:https://pan.baidu.com/s/1ir5xmYJWAX8QIXHb6_5zWw
提取碼:5ph7
里面一共兩個檔案夾,data_train,data_val截圖以示清白,

資料大概就是這個樣子的,左邊是ok的,右邊是ng的,不是徑訓,不是徑訓,不是徑訓,重要的事說三遍,

data_train里面一共784張圖片,都是ok圖片
data_val里面一共366張圖片,150張ng圖片和216張ok圖片
我們訓練只訓練OK圖片,看看能不能通過只訓練OK圖片來判斷驗證集里面的OK和NG
3.Dataset
有了網路,有了資料,就該來處理資料準備往網路里送了,上代碼
import torch
import os
from torch.utils.data import Dataset
import random
import data_agumentation
import torchvision.transforms as tf
import cv2
transform = tf.Compose([tf.ToTensor(),tf.Normalize([0.5],[0.5])])
class train_data(Dataset):
def __init__(self, path):
print('start build_train_data')
self.path=path
self.imgs=[]
for i in os.listdir(path):
self.imgs.append(i)
def __len__(self):
return len(self.imgs)
def __getitem__(self, index):
random_num = random.randint(0,2)
img=cv2.imread(self.path+'/'+self.imgs[index])
if random_num == 1:
img = data_agumentation.augment_left_flip(img)
elif random_num == 2:
img = data_agumentation.augment_rotate(img,180)
img = transform(img)
return img,img
if __name__ == '__main__':
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
data=train_data(r'D:\blog_project\guligedong_unsupervised\data\data_train')
train_loader = torch.utils.data.DataLoader(data, batch_size=3, shuffle=True)
for index,(img, label) in enumerate(train_loader):
print(img.size())
print(label.size())
print()
你應該很熟悉,因為上一篇也寫過類似的,這你會看到,其實img和label是一樣的, 我們的目的也是輸入一張圖片,讓他生成一張一樣的圖片,這樣是為什么呢?我的思路就是,因為我的訓練集只有ok的,網路只能生成ok的特征,如果輸入的是ng的圖片,那么網路就不能生成ng的特征,這個時候就會有差異了,
def __getitem__(self, index):
random_num = random.randint(0,2)
img=cv2.imread(self.path+'/'+self.imgs[index])
if random_num == 1:
img = data_agumentation.augment_left_flip(img)
elif random_num == 2:
img = data_agumentation.augment_rotate(img,180)
img = transform(img)
return img,img
在這個代碼段里,我用了隨機的樣本增強,這個樣本增強也是我前面文章中提供給大家的,就是一個水平的鏡像翻轉和180度的旋轉,當然你還可以增加90和270度的旋轉,大家也可以看到,這個dataset就簡單了很多,因為非監督學習沒有標簽,或者說,非監督學習樣本的標簽就是它本身,
今天就先更新到這里,下一篇我即將要更新訓練代碼和測驗代碼,以及整個專案的代碼,盡情期待,
順便問一下原力值是個什么東西,可以當飯吃嗎?

至此,敬禮,salute!!!!
轉載請註明出處,本文鏈接:https://www.uj5u.com/qita/296801.html
標籤:AI
