riciii7/FastAPI-Batik-GAN
0
1from torch import nn, optim2import torch3from torch.nn import functional as F4from typing import Any, Callable, Optional5import math6 7class VanillaGAN(nn.Module):8 def __init__(self, resolution, latent_dim, hidden_dim=512, channels=3):9 super(VanillaGAN, self).__init__()10 output_dim = resolution * resolution * channels11 12 self.layers = nn.Sequential(13 self.gen_block(latent_dim, hidden_dim),14 self.gen_block(hidden_dim, hidden_dim*2),15 self.gen_block(hidden_dim*2, hidden_dim*2),16 self.gen_block(hidden_dim*2, hidden_dim),17 self.gen_block(hidden_dim, hidden_dim),18 self.gen_block(hidden_dim, hidden_dim//2),19 20 nn.Linear(hidden_dim//2, output_dim),21 nn.Tanh()22 )23 24 def gen_block(self, input_dim, output_dim):25 return nn.Sequential(26 nn.Linear(input_dim, output_dim, bias=False),27 nn.BatchNorm1d(output_dim, 0.8),28 nn.LeakyReLU(0.2)29 )30 31 def forward(self, x):32 return self.layers(x)