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
相关产品推荐
相关产品推荐

