@@ -11,31 +11,33 @@
111112121313class DCGANGenerator(nn.Module):
14-def __init__(self, latent_dim, img_size, img_channels=3):
14+def __init__(self, latent_dim, img_size, img_channels=3, n_filters=16, n_blocks=3):
1515super().__init__()
161617-self.latent_dim = latent_dim
18-self.img_size = img_size
19-self.img_channels = img_channels
20-self.init_size = img_size // 4
17+self.n_filters = n_filters
18+self.init_size = img_size // (2**n_blocks)
21192220self.l1 = nn.Sequential(
23-nn.Linear(latent_dim, 128 * self.init_size ** 2)
21+nn.Linear(latent_dim, n_filters * self.init_size * self.init_size)
2422 )
252324+def block(in_filters, out_filters=None):
25+if out_filters is None:
26+out_filters = 2*in_filters
27+28+return [
29+nn.BatchNorm2d(in_filters),
30+nn.Upsample(scale_factor=2),
31+nn.Conv2d(in_filters, out_filters, kernel_size=3, stride=1, padding=1),
32+ ]
33+34+convs = []
35+for i in range(n_blocks-1):
36+convs.extend(block((2**i) * n_filters))
37+2638self.conv_blocks = nn.Sequential(
27-nn.BatchNorm2d(128),
28-nn.Upsample(scale_factor=2),
29-nn.Conv2d(128, 128, 3, stride=1, padding=1),
30-nn.BatchNorm2d(128, 0.8),
31-nn.LeakyReLU(0.2, inplace=True),
32-nn.Upsample(scale_factor=2),
33-nn.Conv2d(128, 128, 3, 1, 1),
34-nn.Conv2d(128, 64, 3, stride=1, padding=1),
35-nn.BatchNorm2d(64, 0.8),
36-nn.LeakyReLU(0.2, inplace=True),
37-nn.Conv2d(64, 64, 3, 1, 1),
38-nn.Conv2d(64, img_channels, 3, stride=1, padding=1),
39+*convs,
40+*block((2**(n_blocks-1)) * n_filters, img_channels),
3941nn.Tanh(),
4042 )
4143@@ -45,67 +47,86 @@ def forward(self, z):
4547 Returns a (batch_size, img_channels, img_size, img_size) generated images
4648 """
4749out = self.l1(z)
48-out = out.view(out.size(0), 128, self.init_size, self.init_size)
49-img = self.conv_blocks(out)
50-51-# Trim off excess image to match the size of pokemon sprites
52-output = img.view(z.size(0), 3, self.img_size, self.img_size)
53-return output
50+out = out.view(out.size(0), self.n_filters, self.init_size, self.init_size)
51+return self.conv_blocks(out)
545255535654class DCGANDiscriminator(nn.Module):
57-def __init__(self, img_size, img_channels=3):
55+def __init__(self, img_size, img_channels=3, n_filters=16, n_blocks=3):
5856super().__init__()
595760-self.img_size = img_size
61-self.img_channels = img_channels
58+def block(in_filters, out_filters=None, normalise=True):
59+if out_filters is None:
60+out_filters = in_filters*2
626163-def discriminator_block(in_feat, out_feat, normalise=True):
6462block = [
65-nn.Conv2d(in_feat, out_feat, 3, 2, 1),
63+nn.Conv2d(in_filters, out_filters, kernel_size=3, stride=2, padding=1),
6664nn.LeakyReLU(0.2, inplace=True),
6765nn.Dropout2d(0.25),
6866 ]
69677068if normalise:
71-block.append(nn.BatchNorm2d(out_feat, 0.8))
69+block.append(nn.BatchNorm2d(out_filters, 0.8))
70+7271return block
737274-self.model = nn.Sequential(
75-*discriminator_block(img_channels, 16, normalise=False),
76-*discriminator_block(16, 32),
77-*discriminator_block(32, 64),
78-*discriminator_block(64, 128),
73+convs = []
74+for i in range(n_blocks):
75+convs.extend(block((2**i) * n_filters))
76+77+self.conv_blocks = nn.Sequential(
78+*block(img_channels, n_filters, normalise=False),
79+*convs,
7980 )
80818182ds_size = img_size // 2 ** 4
83+final_filters = (2**n_blocks) * n_filters
82848385self.adv_layer = nn.Sequential(
84-nn.Linear(128 * ds_size * ds_size, 1), nn.Sigmoid()
86+nn.Linear(final_filters * ds_size * ds_size, 1), nn.Sigmoid()
8587 )
86888789def forward(self, img):
8890"""
8991 Takes a (batch_size, img_channels, img_size, img_size) generated or real images
9092 Returns a (batch_size, 1) tensor of probabilities that the input is real
9193 """
92-out = self.model(img)
94+out = self.conv_blocks(img)
9395out = out.view(out.size(0), -1)
94969597return self.adv_layer(out)
9698979998100class DCGAN(pl.LightningModule):
99-def __init__(self, latent_dim, img_size, output_img_path=None, img_channels=3, lr=1e-4):
101+def __init__(self,
102+latent_dim,
103+img_size,
104+output_img_path=None,
105+img_channels=3,
106+lr=1e-4,
107+n_filters=16,
108+n_blocks=3):
109+100110super().__init__()
101111102112self.latent_dim = latent_dim
103-self.img_size = img_size
104113self.lr = lr
105114self.output_img_path = output_img_path
106115107-self.g = DCGANGenerator(latent_dim=latent_dim, img_size=img_size, img_channels=img_channels)
108-self.d = DCGANDiscriminator(img_size=img_size, img_channels=img_channels)
116+self.g = DCGANGenerator(
117+latent_dim=latent_dim,
118+img_size=img_size,
119+img_channels=img_channels,
120+n_filters=n_filters,
121+n_blocks=n_blocks
122+ )
123+124+self.d = DCGANDiscriminator(
125+img_size=img_size,
126+img_channels=img_channels,
127+n_filters=n_filters,
128+n_blocks=n_blocks
129+ )
109130110131self.output_z = torch.randn(16, self.latent_dim)
111132self.epoch_n = 0