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)