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

PyTorch文本到图像GAN损失曲线异常技术求助

文本到图像GAN训练问题求助

我是一名刚接触AI与图像生成领域的年轻开发者,近期尝试搭建一个文本到图像GAN,但过程并不顺利。我使用BERT作为文本编码器,基于PyTorch构建神经网络,以下是关键代码片段,如需其他细节请随时告知:

def tokenizeText(text, vectorizer):
    vectorized = []
    if len(text[0]) == 1:
        
        tk = vectorizer.tokenize(text)
        tk = vectorizer.convert_tokens_to_ids(tk)
        tk = torch.LongTensor(tk)

        vectorized.append(tk)
    elif len(text) > 1:
        for i in text:
            tk = vectorizer.tokenize(i)
            tk = vectorizer.convert_tokens_to_ids(tk)
            tk = torch.LongTensor(tk)

            vectorized.append(tk)
        
    

    return vectorized

def prepareText1(textList):
    tk = tokenizeText(textList, vectorizer)

    if len(tk) > 1:
        tknew = []
        for i1 in tk:
            size1 = len(i1)
            size2 = 60 - size1
            zeros = []
            for _ in range(size2):
                zeros.append(0)
            zeros = torch.LongTensor(zeros)
            zeros = zeros.view(1, zeros.size(0))
            i1 = i1.view(1, len(i1))
            tk = torch.cat([i1, zeros], dim=1)
            tknew.append(tk)
            
    return tknew
class Discriminator(nn.Module):
    def __init__(self):
        super(Discriminator, self).__init__()

        self.inD = np.prod(imgShp) + embedDim
        self.outD = 1

        self.Embedlyr = nn.EmbeddingBag(30000, embedDim)

        self.mainNet = nn.Sequential(

            nn.Linear(self.inD, baseNeurons * 4),
            nn.LeakyReLU(alpha, True),
            nn.Dropout(dropout),

            nn.Linear(baseNeurons * 4, baseNeurons * 2),
            nn.LeakyReLU(alpha, True),
            nn.Dropout(dropout),

            nn.Linear(baseNeurons * 2, baseNeurons),
            nn.LeakyReLU(alpha, True),
            nn.Dropout(dropout),

            nn.Linear(baseNeurons, self.outD),
            nn.Sigmoid()

        )

    def forward(self, x, y):
        y = prepareText1(y)
        y = torch.stack(y).squeeze(1)
        y = self.Embedlyr(y)
        x = x.view(x.size(0), np.prod(imgShp))
        x = torch.cat([x, y], dim=1)
        x = self.mainNet(x)

        return x.squeeze(0)

netD = Discriminator().to(device)
optD = optim.Adam(netD.parameters(), lr=lr)
class Generator(nn.Module):
    def __init__(self):
        super(Generator, self).__init__()

        self.inD = nZDim + embedDim
        self.outD = np.prod(imgShp)

        self.Embedlyr = nn.EmbeddingBag(30000, embedDim)

        self.mainNet = nn.Sequential(

            nn.Linear(self.inD, baseNeurons),
            nn.LeakyReLU(alpha, True),

            nn.Linear(baseNeurons, baseNeurons * 2),
            nn.LeakyReLU(alpha, True),

            nn.Linear(baseNeurons * 2, baseNeurons * 4),
            nn.LeakyReLU(alpha, True),

            nn.Linear(baseNeurons * 4, self.outD),
            nn.Tanh(),

        )

    def forward(self, x, y):

        y = prepareText1(y)
        y = torch.stack(y).squeeze(1)
        y = self.Embedlyr(y)
        x = x.view(x.size(0), nZDim)
        x = torch.cat([x, y], dim=1)
        x = self.mainNet(x)

        return x.view(x.size(0), cC, imgSz, imgSz)

netG = Generator().to(device)
optG = optim.Adam(netG.parameters(), lr=lr)
lossFn = nn.BCEWithLogitsLoss()
writer = SummaryWriter()
def trainStepD(dataInp):
    
    optD.zero_grad()

    imgs = dataInp[0]
    lbls = dataInp[1]

    batchSz1 = imgs.size(0)
    logitLblR = Variable(torch.ones((batchSz, 1))).to(device)
    logitLblF = Variable(torch.zeros((batchSz, 1))).to(device)
    
    output = netD(imgs, lbls)
    lossR = lossFn(output, logitLblR)

    nZ = Variable(torch.randn(batchSz1, nZDim)).to(device)
    output = netG(nZ, lbls)
    output = netD(output, lbls)
    lossF = lossFn(output, logitLblF)

    loss = lossR + lossF
    loss.backward()

    optD.step()

    return loss.data.item()

def trainStepG(dataInp):

    optG.zero_grad()

    imgs = dataInp[0]
    lbls = dataInp[1]

    batchSz1 = imgs.size(0)
    logitLblR = Variable(torch.ones((batchSz, 1))).to(device)

    nZ = Variable(torch.randn(batchSz1, nZDim)).to(device)
    output = netG(nZ, lbls)
    output = netD(output, lbls)

    loss = lossFn(output, logitLblR)
    loss.backward()

    optG.step()

    return loss.data.item()
epochs = 50
displayStep = 5
criticNo = 5

for epoch in range(epochs):
    print('Initializing Epoch [{}/{}]...'.format((epoch + 1), epochs), end=' ')

    for batchNo, dataInp in enumerate(dLoader):
        
        step = epoch * len(dLoader) + batchNo + 1

        batchSz1 = dataInp[0].size(0)

        netG.train()
        for _ in range(criticNo):
            DLoss = trainStepD(dataInp)

        GLoss = trainStepG(dataInp)

        writer.add_scalars('scalars', {'GLoss': GLoss, 'DLoss': (DLoss / criticNo)}, step)  

        if step % displayStep == 0:
            netG.eval()
            z = Variable(torch.randn(9, nZDim)).to(device)
            labels = []
            for i in range(9):
                labels.append(dataInp[1][i])
            sample_images = netG(z, labels)
            grid = torchvision.utils.make_grid(sample_images, nrow=3, normalize=True)
            writer.add_image("Batch Results", grid, step)
            
    print('Done!')

训练结果

损失曲线,绿色 - Discriminator Loss,粉色 - Generator Loss:
损失曲线

数据集:
我的数据集是一个小型文件夹,包含标注为“blue”或“red”的蓝色与红色图像:
数据集

恳请各位帮忙分析并解决该问题,谢谢!

内容的提问来源于stack exchange,提问作者ZyoxTheCoder

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.12 15:30:54