Add training script
#8
by Compactbot - opened
- train_logo_gan_v2.py +93 -0
train_logo_gan_v2.py
ADDED
|
@@ -0,0 +1,93 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import os, sys, time, math, argparse
|
| 2 |
+
import numpy as np, torch, torch.nn as nn
|
| 3 |
+
import PIL.Image as Image
|
| 4 |
+
|
| 5 |
+
torch.manual_seed(0); np.random.seed(0)
|
| 6 |
+
|
| 7 |
+
def conv_block(in_c, out_c, ks=4, s=2, p=1):
|
| 8 |
+
return nn.Sequential(nn.Conv2d(in_c, out_c, ks, s, p, bias=False), nn.BatchNorm2d(out_c), nn.LeakyReLU(0.2, inplace=True))
|
| 9 |
+
class G(nn.Module):
|
| 10 |
+
def __init__(self, zdim=128):
|
| 11 |
+
super().__init__()
|
| 12 |
+
self.fc = nn.Sequential(nn.Linear(zdim, 256*8*8), nn.BatchNorm1d(256*8*8), nn.ReLU())
|
| 13 |
+
self.up = nn.Sequential(
|
| 14 |
+
nn.ConvTranspose2d(256,128,4,2,1,bias=False), nn.BatchNorm2d(128), nn.ReLU(),
|
| 15 |
+
nn.ConvTranspose2d(128,64,4,2,1,bias=False), nn.BatchNorm2d(64), nn.ReLU(),
|
| 16 |
+
nn.ConvTranspose2d(64,3,4,2,1,bias=True), nn.Tanh())
|
| 17 |
+
def forward(self, z):
|
| 18 |
+
h = self.fc(z).view(-1,256,8,8); return self.up(h)
|
| 19 |
+
class D(nn.Module):
|
| 20 |
+
def __init__(self):
|
| 21 |
+
super().__init__()
|
| 22 |
+
self.net = nn.Sequential(
|
| 23 |
+
nn.Conv2d(3,64,4,2,1,bias=False), nn.BatchNorm2d(64), nn.LeakyReLU(0.2,inplace=True),
|
| 24 |
+
nn.Conv2d(64,128,4,2,1,bias=False), nn.BatchNorm2d(128), nn.LeakyReLU(0.2,inplace=True),
|
| 25 |
+
nn.Conv2d(128,256,4,2,1,bias=False), nn.BatchNorm2d(256), nn.LeakyReLU(0.2,inplace=True),
|
| 26 |
+
nn.AdaptiveAvgPool2d(1), nn.Flatten(), nn.Linear(256,1))
|
| 27 |
+
def forward(self, x): return self.net(x).squeeze(-1)
|
| 28 |
+
|
| 29 |
+
def grid(imgs, path):
|
| 30 |
+
imgs=(imgs*0.5+0.5).clip(0,1).transpose(0,2,3,1).astype(np.uint8) if imgs.min()<-0.5 else imgs.astype(np.uint8)
|
| 31 |
+
cols=rows=8; cell=64; gap=2
|
| 32 |
+
W=cols*cell+(cols+1)*gap; H=rows*cell+(rows+1)*gap
|
| 33 |
+
canvas=np.full((H,W,3),255,dtype=np.uint8)
|
| 34 |
+
for i in range(min(len(imgs),cols*rows)):
|
| 35 |
+
r=i//cols; c=i%cols; y=gap+r*(cell+gap); x=gap+c*(cell+gap)
|
| 36 |
+
canvas[y:y+cell,x:x+cell]=imgs[i]
|
| 37 |
+
Image.fromarray(canvas).save(path)
|
| 38 |
+
|
| 39 |
+
def main():
|
| 40 |
+
a=argparse.Namespace(data="/work/logos/logos_big.npy", out="/work/models/logo-gan-v2",
|
| 41 |
+
steps=8000, bs=16, zdim=128, lr=2e-4, betas=(0.5,0.999), dev="cuda",
|
| 42 |
+
ckpt_every=1000, sample_every=1000)
|
| 43 |
+
os.makedirs(a.out, exist_ok=True)
|
| 44 |
+
dev = a.dev if (a.dev=="cuda" and torch.cuda.is_available()) else "cpu"
|
| 45 |
+
imgs = np.load(a.data).astype(np.float32) # (N,64,64,3) 0..1
|
| 46 |
+
imgs = (imgs - 0.5)/0.5 # -> [-1,1]
|
| 47 |
+
X = torch.from_numpy(imgs).permute(0,3,1,2).to(dev)
|
| 48 |
+
N = X.shape[0]
|
| 49 |
+
g = G(a.zdim).to(dev); d = D().to(dev)
|
| 50 |
+
go = torch.optim.Adam(g.parameters(), a.lr, betas=a.betas)
|
| 51 |
+
do = torch.optim.Adam(d.parameters(), a.lr, betas=a.betas)
|
| 52 |
+
logf = open(os.path.join(a.out,"train.log"),"a")
|
| 53 |
+
def lg(s): print(s, flush=True); logf.write(s+"\n"); logf.flush()
|
| 54 |
+
lg(f"START v2 N={N} dev={dev} steps={a.steps} bs={a.bs} zdim={a.zdim} lr={a.lr}")
|
| 55 |
+
g_params=sum(p.numel() for p in g.parameters()); d_params=sum(p.numel() for p in d.parameters())
|
| 56 |
+
lg(f"params generator={g_params} discriminator={d_params} total={g_params+d_params}")
|
| 57 |
+
torch.manual_seed(0); np.random.seed(0)
|
| 58 |
+
for step in range(1, a.steps+1):
|
| 59 |
+
idx = torch.randint(0, N, (a.bs,))
|
| 60 |
+
real = X[idx]
|
| 61 |
+
z = torch.randn(a.bs, a.zdim, device=dev)
|
| 62 |
+
# D
|
| 63 |
+
do.zero_grad()
|
| 64 |
+
fake = g(z).detach()
|
| 65 |
+
dloss = -(torch.log(torch.sigmoid(d(real)).clamp(min=1e-7)).mean() + torch.log(1-torch.sigmoid(d(fake)).clamp(min=1e-7)).mean())
|
| 66 |
+
dloss.backward(); do.step()
|
| 67 |
+
# G
|
| 68 |
+
go.zero_grad()
|
| 69 |
+
fake = g(z)
|
| 70 |
+
gloss = -torch.log(torch.sigmoid(d(fake)).clamp(min=1e-7)).mean()
|
| 71 |
+
gloss.backward(); go.step()
|
| 72 |
+
if step % 200 == 0:
|
| 73 |
+
lg(f"step {step}/{a.steps} dloss={float(dloss):.4f} gloss={float(gloss):.4f}")
|
| 74 |
+
if step % a.ckpt_every == 0:
|
| 75 |
+
torch.save({"g":g.state_dict(),"d":d.state_dict(),"step":step,"zdim":a.zdim}, os.path.join(a.out,f"ckpt_{step:05d}.pt"))
|
| 76 |
+
if step % a.sample_every == 0:
|
| 77 |
+
g.eval()
|
| 78 |
+
with torch.no_grad():
|
| 79 |
+
fake = g(torch.randn(64, a.zdim, device=dev)).cpu().numpy()
|
| 80 |
+
grid(fake, os.path.join(a.out, f"grid_{step:05d}.png"))
|
| 81 |
+
g.train()
|
| 82 |
+
# final
|
| 83 |
+
g.eval()
|
| 84 |
+
with torch.no_grad():
|
| 85 |
+
fake = g(torch.randn(64, a.zdim, device=dev)).cpu().numpy()
|
| 86 |
+
grid(fake, os.path.join(a.out, "grid_final.png"))
|
| 87 |
+
torch.save({"g":g.state_dict(),"d":d.state_dict(),"step":a.steps,"zdim":a.zdim}, os.path.join(a.out,"final.pt"))
|
| 88 |
+
lg(f"DONE final step {a.steps}")
|
| 89 |
+
lg("files: " + ", ".join(sorted(os.listdir(a.out))))
|
| 90 |
+
logf.close()
|
| 91 |
+
|
| 92 |
+
if __name__=="__main__":
|
| 93 |
+
main()
|