Files changed (1) hide show
  1. 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()