You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.17 11:42:03