PyTorch:ACGAN(CIFAR-10)标签转独热向量及模型修改
ACGAN数值标签转独热向量的模型修改方案
一、标签预处理(数据加载阶段)
先将1-10的数值标签转换为0-9的索引,再生成独热向量:
import torch import torch.nn.functional as F # 假设label是形状为(batch_size,)的数值张量(1-10) label = label - 1 # 转换为0-based索引 one_hot_label = F.one_hot(label, num_classes=10).float()
二、Generator修改
原模型使用nn.Embedding处理数值标签,现替换为线性层映射独热向量到latent维度,保持与噪声的融合逻辑:
修改后的Generator代码
import torch import torch.nn as nn import torch.nn.functional as F class Generator(nn.Module): def __init__(self, latent_size, nb_filter, n_classes): super(Generator, self).__init__() # 用线性层替代Embedding,将独热向量映射到latent维度 self.label_proj = nn.Linear(n_classes, latent_size) self.conv1 = nn.ConvTranspose2d(latent_size, nb_filter * 8, 4, 1, 0) self.bn1 = nn.BatchNorm2d(nb_filter * 8) self.conv2 = nn.ConvTranspose2d(nb_filter * 8, nb_filter * 4, 4, 2, 1) self.bn2 = nn.BatchNorm2d(nb_filter * 4) self.conv3 = nn.ConvTranspose2d(nb_filter * 4, nb_filter * 2, 4, 2, 1) self.bn3 = nn.BatchNorm2d(nb_filter * 2) self.conv4 = nn.ConvTranspose2d(nb_filter * 2, nb_filter * 1, 4, 2, 1) self.bn4 = nn.BatchNorm2d(nb_filter * 1) self.conv5 = nn.ConvTranspose2d(nb_filter * 1, 3, 4, 2, 1) self.__initialize_weights() def forward(self, input, cl): # cl为独热向量,形状(batch_size, n_classes) label_embed = self.label_proj(cl) x = torch.mul(label_embed, input) x = x.view(x.size(0), -1, 1, 1) x = F.relu(self.bn1(self.conv1(x))) x = F.relu(self.bn2(self.conv2(x))) x = F.relu(self.bn3(self.conv3(x))) x = F.relu(self.bn4(self.conv4(x))) x = torch.tanh(self.conv5(x)) return x def __initialize_weights(self): for m in self.modules(): if isinstance(m, (nn.Conv2d, nn.ConvTranspose2d)): m.weight.data.normal_(0.0, 0.02) elif isinstance(m, nn.BatchNorm2d): m.weight.data.normal_(1.0, 0.02) m.bias.data.fill_(0) elif isinstance(m, nn.Linear): m.weight.data.normal_(0.0, 0.02) m.bias.data.fill_(0)
修改点说明
- 移除
nn.Embedding,新增nn.Linear(n_classes, latent_size)实现独热向量到latent维度的映射 - 初始化方法新增线性层的权重初始化逻辑
forward方法接收独热向量作为类别输入,映射后与噪声相乘融合
三、Discriminator修改
Discriminator核心结构无需改动,仅补充线性层的初始化逻辑,调整损失计算适配独热标签:
修改后的Discriminator代码
class Discriminator(nn.Module): def __init__(self, nb_filter, n_classes): super(Discriminator, self).__init__() self.nb_filter = nb_filter self.conv1 = nn.Conv2d(3, nb_filter, 4, 2, 1) self.conv2 = nn.Conv2d(nb_filter, nb_filter * 2, 4, 2, 1) self.bn2 = nn.BatchNorm2d(nb_filter * 2) self.conv3 = nn.Conv2d(nb_filter * 2, nb_filter * 4, 4, 2, 1) self.bn3 = nn.BatchNorm2d(nb_filter * 4) self.conv4 = nn.Conv2d(nb_filter * 4, nb_filter * 8, 4, 2, 1) self.bn4 = nn.BatchNorm2d(nb_filter * 8) self.conv5 = nn.Conv2d(nb_filter * 8, nb_filter * 1, 4, 1, 0) self.gan_linear = nn.Linear(nb_filter * 1, 1) self.aux_linear = nn.Linear(nb_filter * 1, n_classes) self.__initialize_weights() def forward(self, input): x = F.leaky_relu(self.conv1(input), 0.2) x = F.leaky_relu(self.bn2(self.conv2(x)), 0.2) x = F.leaky_relu(self.bn3(self.conv3(x)), 0.2) x = F.leaky_relu(self.bn4(self.conv4(x)), 0.2) x = self.conv5(x) x = x.view(-1, self.nb_filter * 1) c = self.aux_linear(x) s = torch.sigmoid(self.gan_linear(x)) return s.squeeze(1), c.squeeze(1) def __initialize_weights(self): for m in self.modules(): if isinstance(m, nn.Conv2d): m.weight.data.normal_(0.0, 0.02) elif isinstance(m, nn.BatchNorm2d): m.weight.data.normal_(1.0, 0.02) m.bias.data.fill_(0) elif isinstance(m, nn.Linear): m.weight.data.normal_(0.0, 0.02) m.bias.data.fill_(0)
损失计算适配独热标签
辅助分类损失改用F.binary_cross_entropy_with_logits适配独热标签:
# c为Discriminator输出的分类logits,one_hot_label为独热标签张量 aux_loss = F.binary_cross_entropy_with_logits(c, one_hot_label)
四、使用示例
# 初始化模型 latent_size = 100 nb_filter = 64 n_classes = 10 generator = Generator(latent_size, nb_filter, n_classes) discriminator = Discriminator(nb_filter, n_classes) # 生成随机噪声和独热标签 batch_size = 32 noise = torch.randn(batch_size, latent_size) label = torch.randint(1, 11, (batch_size,)) # 1-10的数值标签 label_0based = label - 1 one_hot_label = F.one_hot(label_0based, num_classes=10).float() # 生成图像 fake_img = generator(noise, one_hot_label) # 判别图像 real_pred, aux_pred = discriminator(fake_img)
内容的提问来源于stack exchange,提问作者luc
相关产品推荐
相关产品推荐

