GitHub

@@ -11,31 +11,33 @@

111112121313

class 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):

1515

super().__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)

21192220

self.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+2638

self.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),

3941

nn.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

"""

4749

out = 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)

545255535654

class 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):

5856

super().__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):

6462

block = [

65-

nn.Conv2d(in_feat, out_feat, 3, 2, 1),

63+

nn.Conv2d(in_filters, out_filters, kernel_size=3, stride=2, padding=1),

6664

nn.LeakyReLU(0.2, inplace=True),

6765

nn.Dropout2d(0.25),

6866

]

69677068

if normalise:

71-

block.append(nn.BatchNorm2d(out_feat, 0.8))

69+

block.append(nn.BatchNorm2d(out_filters, 0.8))

70+7271

return 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

)

80818182

ds_size = img_size // 2 ** 4

83+

final_filters = (2**n_blocks) * n_filters

82848385

self.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

)

86888789

def 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)

9395

out = out.view(out.size(0), -1)

94969597

return self.adv_layer(out)

9698979998100

class 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+100110

super().__init__()

101111102112

self.latent_dim = latent_dim

103-

self.img_size = img_size

104113

self.lr = lr

105114

self.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+

)

109130110131

self.output_z = torch.randn(16, self.latent_dim)

111132

self.epoch_n = 0

Read the original on github.com ↗