| import torch |
| import torch.nn as nn |
|
|
|
|
| LATENT_DIM = 100 |
| NUM_CLASSES = 10 |
|
|
|
|
| class Generator(nn.Module): |
| def __init__(self, latent_dim=LATENT_DIM, num_classes=NUM_CLASSES): |
| super().__init__() |
| self.latent_dim = latent_dim |
| self.label_embedding = nn.Embedding(num_classes, 32) |
| self.project = nn.Sequential( |
| nn.Linear(latent_dim + 32, 128 * 7 * 7), |
| nn.BatchNorm1d(128 * 7 * 7), |
| nn.ReLU(inplace=True), |
| ) |
| self.decoder = nn.Sequential( |
| nn.ConvTranspose2d(128, 64, kernel_size=4, stride=2, padding=1, bias=False), |
| nn.BatchNorm2d(64), |
| nn.ReLU(inplace=True), |
| nn.ConvTranspose2d(64, 32, kernel_size=3, stride=1, padding=1, bias=False), |
| nn.BatchNorm2d(32), |
| nn.ReLU(inplace=True), |
| nn.ConvTranspose2d(32, 1, kernel_size=4, stride=2, padding=1), |
| nn.Tanh(), |
| ) |
|
|
| def forward(self, noise, labels): |
| if noise.ndim > 2: |
| noise = noise.flatten(start_dim=1) |
| conditioned = torch.cat([noise, self.label_embedding(labels)], dim=1) |
| features = self.project(conditioned).view(-1, 128, 7, 7) |
| return self.decoder(features) |
|
|