PANet:基于金字塔注意力網路的影像超解析度重建
[!] 為了提高代碼的可讀性,本文模型的具體實作與原文具有一定區別,因此會造成性能上的差異
文章目錄
- PANet:基于金字塔注意力網路的影像超解析度重建
- 1.相關資料
- 2.簡介
- 3.模型結構
- 4.專案實踐
- 4.1 準備作業
- 4.2 具體實作
- 4.2.1 匯入專案所需庫
- 4.2.2 構建資料集
- 4.2.3 構建網路模型
- # 特征金字塔部分
- # 模型部分
- 4.2.4 準備訓練配件
- # 優化器
- # 損失函式
- # 評估標準
- ## PSRN
- ## SSIM
- 4.2.5 構建訓練框架
- 4.2.6 訓練結果
1.相關資料
- 論文下載地址: 傳送門
- 原作者代碼地址:傳送門
- 完整代碼地址:傳送門
2.簡介
- PANet(Pyramid Attention with Simple Network Backbones)是一種基于影像恢復金字塔注意力模塊的影像修復模型,它能夠從多尺度特征金字塔種提取到長距離與短距離的特征關系,
- 受降采樣能夠有效減少壓縮偽影等影像噪聲的啟發,作者所提出的金字塔利用不同采樣倍數的特征圖來相互傳遞注意力信號,以更靈活的方式來借用不同特征尺寸之間的“干凈”資訊,
- 作者只在一個簡單的前饋鏈接網路中加入了一個金字塔注意力模塊,就在絕大多數影像修復任務中達到了SOTA,(這樣看來模塊確實牛逼)
3.模型結構
直接上圖

- 圖上面部分就是傳說中的金字塔注意力模塊,圖下面部分就是PANet的結構(這個結構和SRResNet怪像的,可以參考我的相關文章:SRResNet和SRGAN)
- 金字塔注意力模塊的結構分為兩個部分:金字塔采樣環節和S-A Attention,金字塔采樣環節就是簡單的降采樣處理,根據源代碼來看,作者使用的是雙二次下采樣的方法,
- S-A Attention的結構參考了NLP中最經典的注意力機制結構,即構建了Q,K,V三種特征圖來捕獲影像在不同尺寸中的資訊,與其他注意力機制不同的是,S-A Attention將注意力機制中的按元素相乘環節改成將Q和K特征圖作為卷積核(即圖中淺藍色特征層出來的兩個特征圖)來與V特征圖進行卷積/反卷積操作,
4.專案實踐
在這里我會一步一步教大家做一個能夠成功運行的PANet,完整的代碼也會很快推出,
4.1 準備作業
-
筆者使用的作業環境如下所示:
系統:Windows 10 CPU:Intel Core i9-10850K GPU:GeForce RTX 3090 -
實作代碼所需要準備的庫為:
Pytorch OpenCV Numpy Torchvision -
本文使用的是COCO 2017資料集,其中包含了123,403張照片,大家可以根據自己的需要來使用自己的資料集,
4.2 具體實作
為了方便閱讀,部分代碼已標注中文注釋,而且全部放進了一個代碼檔案中
- 完整版代碼支持重新打開代碼自動恢復到上次訓練的功能,只需要關注筆者即可獲得:傳送門
4.2.1 匯入專案所需庫
import torch
import torch.nn as nn
import torch.nn.functional as F
from torch.utils.data import DataLoader,Dataset,SubsetRandomSampler
import torch.optim as optim
from torchvision import utils as vutils
from torchvision.utils import save_image
import os
import cv2
import random as ra
import numpy as np
import math
4.2.2 構建資料集
class PreprocessDataset(Dataset):
def __init__(self,path,size = 96):
super().__init__()
self.size = size #高清影像的尺寸,這里默認為96x96
self.allImgs = list()
for root,dirs,files in os.walk(path):
self.allImgs = [os.path.join(root,file) for file in files] #獲取影像的地址
def __len__(self):
return len(self.allImgs)
def __getitem__(self,index):
img = self.allImgs[index]
img = cv2.imread(img)
img = cv2.cvtColor(img,cv2.COLOR_BGR2RGB)
height,width,_ = img.shape
xStart = ra.randint(0,width-self.size-1)
yStart = ra.randint(0,height-self.size-1)
img = img[yStart:self.size + yStart,xStart:self.size + xStart,:] #隨機裁剪影像
if ra.random() > 0.5:
img = cv2.flip(img,1) #有50%幾率反轉影像
hr = torch.tensor(np.transpose(img,(2,0,1)))/255.0
hr = (hr - 0.5)/0.5 #像素標準化
lr = F.max_pool2d(hr,2) #使用最大池化來獲得下采樣圖片
return hr,lr
-
構建完資料集類后,我們可以很方便地構建對應的Dataloader,在這里我只構建了訓練集,并沒有構建測驗集,
path = '你的資料集檔案路徑' dataset = PreprocessDataset(path,size = 96) trainData = DataLoader(dataset,batch_size = 32,num_workers = 4,shuffle = True)
4.2.3 構建網路模型
# 特征金字塔部分
- 這里直接改進了原作者的金字塔注意力模塊代碼,因此代碼風格會與其他部分有一定差異,
def extract_image_patches(images, ksizes, strides, rates, padding='same'): """ Extract patches from images and put them in the C output dimension. :param padding: :param images: [batch, channels, in_rows, in_cols]. A 4-D Tensor with shape :param ksizes: [ksize_rows, ksize_cols]. The size of the sliding window for each dimension of images :param strides: [stride_rows, stride_cols] :param rates: [dilation_rows, dilation_cols] :return: A Tensor """ assert len(images.size()) == 4 assert padding in ['same', 'valid'] batch_size, channel, height, width = images.size() if padding == 'same': images = same_padding(images, ksizes, strides, rates) elif padding == 'valid': pass else: raise NotImplementedError('Unsupported padding type: {}.\ Only "same" or "valid" are supported.'.format(padding)) unfold = torch.nn.Unfold(kernel_size=ksizes, dilation=rates, padding=0, stride=strides) patches = unfold(images) return patches # [N, C*k*k, L], L is the total number of such blocks def reduce_sum(x, axis=None, keepdim=False): if not axis: axis = range(len(x.shape)) for i in sorted(axis, reverse=True): x = torch.sum(x, dim=i, keepdim=keepdim) return x def same_padding(images, ksizes, strides, rates): assert len(images.size()) == 4 batch_size, channel, rows, cols = images.size() out_rows = (rows + strides[0] - 1) // strides[0] out_cols = (cols + strides[1] - 1) // strides[1] effective_k_row = (ksizes[0] - 1) * rates[0] + 1 effective_k_col = (ksizes[1] - 1) * rates[1] + 1 padding_rows = max(0, (out_rows-1)*strides[0]+effective_k_row-rows) padding_cols = max(0, (out_cols-1)*strides[1]+effective_k_col-cols) # Pad the input padding_top = int(padding_rows / 2.) padding_left = int(padding_cols / 2.) padding_bottom = padding_rows - padding_top padding_right = padding_cols - padding_left paddings = (padding_left, padding_right, padding_top, padding_bottom) images = torch.nn.ZeroPad2d(paddings)(images) return images def default_conv(in_channels, out_channels, kernel_size,stride=1, bias=True): return nn.Conv2d( in_channels, out_channels, kernel_size, padding=(kernel_size//2),stride=stride, bias=bias) class BasicBlock(nn.Sequential): def __init__( self, conv, in_channels, out_channels, kernel_size, stride=1, bias=True, bn=False, act=nn.PReLU()): m = [conv(in_channels, out_channels, kernel_size, bias=bias)] if bn: m.append(nn.BatchNorm2d(out_channels)) if act is not None: m.append(act) super(BasicBlock, self).__init__(*m) class PyramidAttention(nn.Module): def __init__(self, level=5, res_scale=1, channel=64, reduction=2, ksize=3, stride=1, softmax_scale=10, average=True, conv=default_conv): super(PyramidAttention, self).__init__() self.ksize = ksize self.stride = stride self.res_scale = res_scale self.softmax_scale = softmax_scale self.scale = [1-i/10 for i in range(level)] self.average = average escape_NaN = torch.FloatTensor([1e-4]) self.register_buffer('escape_NaN', escape_NaN) self.conv_match_L_base = BasicBlock(conv,channel,channel//reduction, 1, bn=False, act=nn.PReLU()) self.conv_match = BasicBlock(conv,channel, channel//reduction, 1, bn=False, act=nn.PReLU()) self.conv_assembly = BasicBlock(conv,channel, channel,1,bn=False, act=nn.PReLU()) def forward(self, input): res = input #theta match_base = self.conv_match_L_base(input) shape_base = list(res.size()) input_groups = torch.split(match_base,1,dim=0) # patch size for matching kernel = self.ksize # raw_w is for reconstruction raw_w = [] # w is for matching w = [] #build feature pyramid for i in range(len(self.scale)): ref = input if self.scale[i]!=1: ref = F.interpolate(input, scale_factor=self.scale[i], mode='bicubic', align_corners=True,recompute_scale_factor=True) #feature transformation function f base = self.conv_assembly(ref) shape_input = base.shape #sampling raw_w_i = extract_image_patches(base, ksizes=[kernel, kernel], strides=[self.stride,self.stride], rates=[1, 1], padding='same') # [N, C*k*k, L] raw_w_i = raw_w_i.view(shape_input[0], shape_input[1], kernel, kernel, -1) raw_w_i = raw_w_i.permute(0, 4, 1, 2, 3) # raw_shape: [N, L, C, k, k] raw_w_i_groups = torch.split(raw_w_i, 1, dim=0) raw_w.append(raw_w_i_groups) #feature transformation function g ref_i = self.conv_match(ref) shape_ref = ref_i.shape #sampling w_i = extract_image_patches(ref_i, ksizes=[self.ksize, self.ksize], strides=[self.stride, self.stride], rates=[1, 1], padding='same') w_i = w_i.view(shape_ref[0], shape_ref[1], self.ksize, self.ksize, -1) w_i = w_i.permute(0, 4, 1, 2, 3) # w shape: [N, L, C, k, k] w_i_groups = torch.split(w_i, 1, dim=0) w.append(w_i_groups) y = [] for idx, xi in enumerate(input_groups): #group in a filter wi = torch.cat([w[i][idx][0] for i in range(len(self.scale))],dim=0) # [L, C, k, k] #normalize max_wi = torch.max(torch.sqrt(reduce_sum(torch.pow(wi, 2), axis=[1, 2, 3], keepdim=True)), self.escape_NaN) wi_normed = wi/ max_wi #matching xi = same_padding(xi, [self.ksize, self.ksize], [1, 1], [1, 1]) # xi: 1*c*H*W yi = F.conv2d(xi, wi_normed, stride=1) # [1, L, H, W] L = shape_ref[2]*shape_ref[3] yi = yi.view(1,wi.shape[0], shape_base[2], shape_base[3]) # (B=1, C=32*32, H=32, W=32) # softmax matching score yi = F.softmax(yi*self.softmax_scale, dim=1) if self.average == False: yi = (yi == yi.max(dim=1,keepdim=True)[0]).float() # deconv for patch pasting raw_wi = torch.cat([raw_w[i][idx][0] for i in range(len(self.scale))],dim=0) yi = F.conv_transpose2d(yi, raw_wi, stride=self.stride,padding=1)/4. y.append(yi) y = torch.cat(y, dim=0)+res*self.res_scale # back to the mini-batch return y
# 模型部分
-
PANet使用的是SRResNet的骨干
class ResBlock(nn.Module): def __init__(self,inChannals): super().__init__() self.model = nn.Sequential( nn.Conv2d(inChannals,inChannals,kernel_size = 1,bias = False), nn.BatchNorm2d(inChannals), nn.ReLU(inplace = True), nn.Conv2d(inChannals,inChannals,kernel_size = 3,stride = 1, padding = 1,bias = False,padding_mode = 'reflect'), nn.BatchNorm2d(inChannals) ) def forward(self,input): return F.relu(input + self.model(input),inplace = True) class Sequential(nn.Sequential): def __init__(self,inChannals,blockNum = 8): seq = [ResBlock(inChannals) for _ in range(blockNum)] seq.insert(int(blockNum/2),PyramidAttention(channel=inChannals, level=4)) super().__init__(*seq) class Model(nn.Module): def __init__(self,channals = 64,blockNum = 6): super().__init__() self.features = nn.Sequential( nn.Conv2d(3,channals,kernel_size = 7,padding = 3,stride = 1, padding_mode = 'reflect',bias = False), nn.BatchNorm2d(channals), nn.ReLU(inplace = True), nn.Conv2d(channals,channals,kernel_size = 3,padding = 1,stride = 1, padding_mode = 'reflect',bias = False), nn.BatchNorm2d(channals), nn.ReLU(inplace = True) ) self.sequential = Sequential(channals,blockNum) self.upSample = nn.Sequential( nn.Conv2d(channals,channals * 4,kernel_size = 3,padding = 1,stride = 1, padding_mode = 'reflect'), nn.PixelShuffle(2), nn.Conv2d(channals,channals,kernel_size = 3,padding = 1,stride = 1), nn.ReLU(inplace = True), nn.Conv2d(channals,3,kernel_size = 1,stride = 1), nn.Tanh() ) def forward(self,input): features = self.features(input) output = self.sequential(features) output = features + output output = self.upSample(output) return output -
最后,通過簡單的方式我們便可構建一個模型
#如果電腦可以使用顯卡,則自動使用顯卡加速 device = torch.device("cuda" if torch.cuda.is_available() else "cpu") #創建網路模型 net = Model(channals = 64,blockNum = 24).to(device)
4.2.4 準備訓練配件
- 為了對模型進行訓練和驗證,我們需要以下部件:優化器Optimizer、損失函式Criteria和評估標注
# 優化器
- 優化器我們使用了AdamW
optimizer = optim.AdamW(net.parameters(),lr = 1e-4)
# 損失函式
- 損失函式我們參考了原作者,使用了L1 Loss
criteria = nn.L1Loss()
# 評估標準
- 我們使用了SSIM和PSRN兩種標注,他們的代碼如下所示:
## PSRN
- 代碼如下:
def PSRN(img1, img2): mse = torch.mean((img1 - img2) ** 2) if mse < 1.0e-10: return 100 return 10 * math.log10(255.0**2/mse)
## SSIM
- 代碼如下:
def gaussian(window_size, sigma): gauss = torch.Tensor([exp(-(x - window_size//2)**2/float(2*sigma**2)) for x in range(window_size)]) return gauss/gauss.sum() def create_window(window_size, channel): _1D_window = gaussian(window_size, 1.5).unsqueeze(1) _2D_window = _1D_window.mm(_1D_window.t()).float().unsqueeze(0).unsqueeze(0) window = Variable(_2D_window.expand(channel, 1, window_size, window_size).contiguous()) return window def _ssim(img1, img2, window, window_size, channel, size_average = True): mu1 = F.conv2d(img1, window, padding = window_size//2, groups = channel) mu2 = F.conv2d(img2, window, padding = window_size//2, groups = channel) mu1_sq = mu1.pow(2) mu2_sq = mu2.pow(2) mu1_mu2 = mu1*mu2 sigma1_sq = F.conv2d(img1*img1, window, padding = window_size//2, groups = channel) - mu1_sq sigma2_sq = F.conv2d(img2*img2, window, padding = window_size//2, groups = channel) - mu2_sq sigma12 = F.conv2d(img1*img2, window, padding = window_size//2, groups = channel) - mu1_mu2 C1 = 0.01**2 C2 = 0.03**2 ssim_map = ((2*mu1_mu2 + C1)*(2*sigma12 + C2))/((mu1_sq + mu2_sq + C1)*(sigma1_sq + sigma2_sq + C2)) if size_average: return ssim_map.mean() else: return ssim_map.mean(1).mean(1).mean(1) def ssim(img1, img2, window_size = 11, size_average = True): (_, channel, _, _) = img1.size() window = create_window(window_size, channel) if img1.is_cuda: window = window.cuda(img1.get_device()) window = window.type_as(img1) return _ssim(img1, img2, window, window_size, channel, size_average)- 好吧,這兩個都是借鑒別人的(老懶狗了
4.2.5 構建訓練框架
- 訓練框架如下所示:
if __name__ == '__main__':
path = '你的資料集路徑'
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
dataset = PreprocessDataset(path,size = 96)
trainData = DataLoader(dataset,batch_size = 32,num_workers = 4,shuffle = True)
net = Model(channals = 64,blockNum = 24).to(device)
print(net)
criteria = nn.L1Loss()
optimizer = optim.AdamW(net.parameters(),lr = 1e-4)
totalStep = len(trainData)
# 構建可視化結果的保存路徑
if not os.path.exists('./img'):
os.mkdir('./img')
for epoch in range(startEpoch,10000):
if epoch == 20 or epoch == 40:
update_lr(optimizer, multiplier = .1)
totalSSIM = 0.0
totalPSRN = 0.0
totalLoss = 0.0
for step,(hr,lr) in enumerate(trainData,1):
net.train(True)
hr,lr = hr.to(device),lr.to(device)
net.zero_grad()
output = net(lr)
loss = criteria(output,hr)
loss.backward()
optimizer.step()
totalLoss += loss
totalSSIM += ssim(output,hr)
totalPSRN += PSRN(output,hr)
print("[Epoch %d] Step: %d/%d Loss: %.4f|ssim: %.4f|psrn: %.4f" %
(epoch,step,totalStep,totalLoss/step,totalSSIM/step,totalPSRN/step))
if step >= 100: #對影像進行可視化
net.train(False)
outputs = net(lr)
outputs = torch.cat([hr,outputs],dim = 0)
save_image(outputs,'./Img/Result_epoch_%08d.jpg' % epoch,nrow = 8,normalize = True)
- 完整版代碼支持重新打開代碼自動恢復到上次訓練的功能,只需要關注筆者即可獲得:傳送門
4.2.6 訓練結果
-
100次訓練后結果:

-
10,000次訓練后結果:

此時SSIM:0.6710 ;PSRN:65.6384 -
由于COCO資料集中的特征不唯一,因此需要更多的訓練才能夠達到更好的結果,
-
完整版代碼支持重新打開代碼自動恢復到上次訓練的功能,只需要關注筆者即可獲得:傳送門
轉載請註明出處,本文鏈接:https://www.uj5u.com/qita/303135.html
標籤:其他
