學習筆記GAN002:DCGAN,gan002dcgan
Ian J. Goodfellow 論文:https://arxiv.org/abs/1406.2661
兩個網路:G(Generator),產生網路,接收隨機雜訊Z,通過雜訊產生樣本,G(z)。D(Dicriminator),判別網路,判別樣本是否真實,輸入樣本x,輸出D(x)代表x真實機率,如果1,100%真實樣本,如果0,代表不可能是真實樣本。
訓練過程,產生網路G盡量產生真實樣本欺騙判別網路D,判別網路D盡量把G產生樣本和真實樣本分別開。理想狀態下,G產生樣本G(z),使D難以判斷真假,D(G(z))=0.5。此時,產生模型G,可以用來產生樣本。
數學公式:minG maxDV(D,G)=Ex~pdata(x)[logD(x)]+Ez~pz(z)[log(1-D(G(z)))]
二項式。x真實樣本,z輸入G網雜訊,G(z) G網產生樣本。D(x) D網判斷真實樣本是否真實機率,越接近1越好。D(G(z)) D網判斷G網產生樣本真實機率。G網,D(G(z))儘可能大,V(D,G)變小,min_G。D網,D(x)越大,D(G(x))越小,V(D,G)越大,max_D。
1、x sampled from data -> Differentiable function D -> D(x) tries to be near 1
2、Input noise z -> Differntiable function G -> x sampled from model -> D -> D tries to make D(G(z)) near 0,G tries to make D(G(z)) near 1
隨機梯度下降法訓練D、G。
Algorithm 1 Minibatch stochastic gradient descent training of genegative adversarial nets. The number of steps to apply to the discriminator,k,is a hyperparameter, we used k = 1,the least expensive option, in our experiments.
for number of training iterations do
for k steps do
Sample minibatch of m noise samples {z(1),...,z(m)} from noise prior pg(z)
Sample minibatch of m examples {x(1),...,x(m)} from data generating distribution pdata(x)
Update the discriminator by ascending its stochastic gradient:
end for
Sample mninbatch of m noise samples {z(1),...,z(m)} from noise prior pg(z)
Update the generator by descending its stochastic gradient:
end for
The gradient-based updates can use any standard gradient-based learning rule.We used momenttum in our experiments.
第一步訓練D,V(G,D)越大越好,上升(增加)梯度(ascending)。第二步訓練G,V(G,D)越小越好,下降(減少)梯度(descending)。交替進行。
DCGAN原理。https://arxiv.org/abs/1511.06434 。Alec Radford, Luke Metz, Soumith Chintala,《Unsupervised Representation Learning with Deep Convolutional Generative Adversarial Networks》。G、D換成卷積神經網路(CNN)。取消所有pooling層。G網路使用轉置卷積(transposed convolutional layer)上採樣,D網路加入stride卷積代替pooling。D、G都batch normalization。去掉FC層,全卷積網路。G網用ReLU啟用函數,最後一層用tanh。D網用LeakyRelu啟用函數。
G網路。Project reshape 100 z -> 4X4X1024 -> 8X8X512 -> 16X16X256 -> 32X32X128 -> 64X64X3
用DCGAN產生動漫人物頭像:
http://qiita.com/mattya/items/e5bfe5e04b9d2f0bbd47 。
未經處理資料搜集。http://safebooru.donmai.us 。http://konachan.net 。
爬蟲代碼:
import requests
from bs4 import BeautifulSoup
import os
import traceback
def download(url, filename):
if os.path.exists(filename):
print('file exists!')
return
try:
r = requests.get(url, stream=True, timeout=60)
r.raise_for_status()
with open(filename, 'wb') as f:
for chunk in r.iter_content(chunk_size=1024):
if chunk: # filter out keep-alive new chunks
f.write(chunk)
f.flush()
return filename
except KeyboardInterrupt:
if os.path.exists(filename):
os.remove(filename)
raise KeyboardInterrupt
except Exception:
traceback.print_exc()
if os.path.exists(filename):
os.remove(filename)
if os.path.exists('imgs') is False:
os.makedirs('imgs')
start = 1
end = 8000
for i in range(start, end + 1):
url = 'http://konachan.net/post?page=%d&tags=' % i
html = requests.get(url).text
soup = BeautifulSoup(html, 'html.parser')
for img in soup.find_all('img', class_="preview"):
target_url = 'http:' + img['src']
filename = os.path.join('imgs', target_url.split('/')[-1])
download(target_url, filename)
print('%d / %d' % (i, end))
頭像截取:
https://github.com/nagadomi/lbpcascade_animeface 。
封裝:
import cv2
import sys
import os.path
from glob import glob
def detect(filename, cascade_file="lbpcascade_animeface.xml"):
if not os.path.isfile(cascade_file):
raise RuntimeError("%s: not found" % cascade_file)
cascade = cv2.CascadeClassifier(cascade_file)
image = cv2.imread(filename)
gray = cv2.cvtColor(image, cv2.COLOR_BGR2GRAY)
gray = cv2.equalizeHist(gray)
faces = cascade.detectMultiScale(gray,
# detector options
scaleFactor=1.1,
minNeighbors=5,
minSize=(48, 48))
for i, (x, y, w, h) in enumerate(faces):
face = image[y: y + h, x:x + w, :]
face = cv2.resize(face, (96, 96))
save_filename = '%s-%d.jpg' % (os.path.basename(filename).split('.')[0], i)
cv2.imwrite("faces/" + save_filename, face)
if __name__ == '__main__':
if os.path.exists('faces') is False:
os.makedirs('faces')
file_list = glob('imgs/*.jpg')
for filename in file_list:
detect(filename)
訓練:
https://github.com/carpedm20/DCGAN-tensorflow 。
model.py:
if config.dataset == 'mnist':
data_X, data_y = self.load_mnist()
else:
data = glob(os.path.join("./data", config.dataset, "*.jpg"))
data檔案夾建立anime檔案夾放圖片,運行時指定 --dataset anime。
python main.py --image_size 96 --output_size 48 --dataset anime --is_crop True --is_train True --epoch 300 --input_fname_pattern "*.jpg"
GAN論文:https://github.com/zhangqianhui/AdversarialNetsPapers 。
SGD最佳化。目標函數(objective function)判斷、監視學習成果。J(D) 判別網路目標函數,交叉熵(cross entropy)函數。左邊D判別真實資料,右邊D判別G產生噪音資料。J(G) 產生網路目標函數。
最小最大博弈,minimax game。均衡點(納什均衡),J(D)鞍點(saddle point)。
真實資料(data)和模型產生偽資料(model distribution z映射)。 學習D,區分data、model分布。data、model分布相加做分母,分子是真實data分布。目標,D無限接近常數1/2.Pmodel、Pdata無限相似。產生模型與來源資料擬合後,無法再學習,常數y=1/2求導永遠0。
非飽和博弈(Non-Saturating)。G偽裝成功率表示目標函數,均衡不由損失(loss)決定。D完美後,G還可以繼續最佳化。
DCGAN(深度卷積產生對抗網路 Deep Convolutional Generative Adversarial Network),反向CNN。
卷積神經網路原理,convolutinoal filter 卷積過濾器(濾鏡),把圖片過濾(轉化)成各種樣式。不同過濾器,把圖片轉化成不同樣式。不同樣式為原圖片不同特徵表達。特徵學習。
DCGAN創造圖片。把一組特徵值慢慢恢複成一張圖片。
每一個濾鏡層,CNN把大圖片重要特徵提取出來,一步一步減小圖片尺寸。DCGAN把小圖片(小數組)特徵放大,排列成新圖片。DCGAN輸入最初小資料是雜訊資料。圖片RGB矩陣,可以向量加減。戴墨鏡的男人-不戴墨鏡的男人+不戴墨鏡的女人=戴墨鏡的女人。NLP,word2vec,king-man+woman=queen。向量/矩陣加減後,還原成“圖義”代表的圖片。NLP,word2vec,向量對應有意義的詞;DCGAN,矩陣對應有意義的圖片。
統計學科,JS距離(minimax),KL距離,散度(divergence)方程。創造目標函數。DKL(P||Q)=S∞-∞p(x)log(p(x)/q(x))dx 。
GAN 神經網路構造,通過類SGD方法最佳化模型。目標函數重要。Q雜訊資料分布。P目標分布。求最大似然(Maximum Likelihood),使KL距離最小化。
KL[P||Q]=SPlog(P/Q)dx=SPlogPdx-SPlogQdx
P、Q都以x為變數,P是真實資料分類。SPlogPdx 是常數。SPlogQdx 只有logQ是變數。正比於-logQ。
KL[P||Q]=-常數-S另一常數·logQdx
Q,P(x|θ)。P模型,θ參數。負的最大似然。
類GAN演算法最小化任何f-divergence方程。
面對無限多資料,都可以學到真實資料分布P。現實,資料有限。KL公式理論,KL(P||Q),Q擬合真實資料P,極大解釋全部P內涵(overgeneralization)。多模態(multimodal),資料不夠多,KL(P||Q)覆蓋不完整。KL(Q||P),undergeneralization 。先覆蓋較大,再覆蓋較小。
G目標函數改造成最大似然。J(G)求導,得到最大似然表達形式。Maximal Likelihood跑得最快。
GAN,產生(複刻)樣本,還可以轉為強化學習模型(Reinforcement Learning)。上海交大 SeqGAN[Yu et al. 2016]論文。
把資料標籤給GAN。學習條件機率p(y|x)遠比單獨p(x)容易。部分有標籤就能大幅提升GAN訓練效果,半監督(semi-supervising)學習。半監督學習,三類資料,真實無標籤資料,有標籤資料,噪音產生資料。目標函數,監督方法和無監督方法結合。標籤平滑(smooth),把0、1離散標籤,轉變成更加平滑的0.1(beta)、0.9(alpha)等。分子混進beta係數假資料分布。假資料建議保留標籤0,一個平滑,另一個不平滑,one-sided label amoothing(單邊標籤平滑)。平滑,GAN判別函數不會給出太大梯度訊號(gradient signal),防止演算法走向極端樣本陷阱。
Batch Norm,取一批資料,正常化(normalise,減平均值,除以標準差)。資料更集中,不太大太小,學習效率更高。同一批(batch)資料太相似,無監督GAN,容易被帶偏,認為資料都一樣,最終產生模型混雜很多其它特徵。
Reference Batch Norm,取一批資料(固定)作參照資料集R,新資料batch依據r平均值、標準差正常化。R取得不好,效果也不會好,或R過擬合。
Virtual Batch Norm,取R,新資料x正常化,x加入R形成virtual batch V,用V的平均值、標準差來標準化x,極大減少R風險。
平衡好G、D。通常對抗網路,判別模型D贏,D比G深。用非飽和(non-saturating)博弈寫目標函數,保證D學完後,G可以繼續學習。
GAN問題。不收斂(non-convergence),容易只找到局部最優點,非全域最優點,或根本無法收斂。模式崩潰(mode collapse),minmaxV(G,D)不等於maxminV(G,D),如果maxD放在內圈,演算法可以收斂到應有位置,如果maxG放在內圈,演算法撲向聚集區,看不到全域分布。Reverse KL,保守損失(loss)。
Minibatch GAN,原資料分成小batch,保證太相惟資料樣本不被放到一個小batch,資料足夠多樣,避免模式崩潰。
空間理解錯誤。圖片用2D表示3D。圖片樣本產生圖片,空間表達不好。
Unrolled(不滾) GAN,每一步不把判別模型D滾起,把K次D存起,根據損失(loss)選擇最好。
無法科學評估,無法量化標準。
離散輸出,無法微分(differentiate)。Williams(1992) REINFORCE。Jang et al.(2016) Gumbel-softmax。用連續數值訓練,框定範圍,輸出離散值。
強化學習串連,無法收斂,有限步數,窮舉更簡單粗暴效果好。
PPGN(Plug and Play Generative Models 隨插即用產生模型),Nguten et al,2016。產生模型領域新State-of-the-art(當前最佳)。
GAN用可以利用監督學習估測複雜目標函數產生模型,GAN內部自己拿真假樣本對照。高維連續非凸找納什均衡,有待研究。
參考資料:
https://zhuanlan.zhihu.com/p/24767059
http://www.sohu.com/a/121189842_465975
歡迎付費諮詢(150元每小時),我的:qingxingfengzi
群體經過適當組織,可以互相促進。我一直相,信良性互動可以使彼此更快速地成長,協助更多人進入自己感興趣的領域。我在建立一個一起學習GAN的群,我們以每天各報各的學習進度為主。加我,我會把你拉進群裡,加的時候請註明:加入GAN日報群。