在图生成GAN中添加稀疏性损失惩罚时遇梯度报错求助
解决GAN生成图邻接矩阵稀疏度惩罚的梯度错误问题
问题场景
我实现了一个用于生成图(本质是生成邻接矩阵)的GAN,训练代码如下:
def find_sparsity(graph): return graph.edge_index.shape[1] / (graph.num_nodes ** 2) def train_gan(loader, discriminator, generator, optimizer_d, optimizer_g, _noise_dim=500, _epochs=100): lossesG = [] lossesD = [] for epoch in range(1, _epochs): for data in loader: data_real = data.to(device) #noise: data_fake = make_geometric_noise_data(_noise_dim).to(device).detach() fake_sample = generator(data_fake).detach() out_real = discriminator(data_real).detach() fake_sample.batch = torch.tensor([1]).to(device) out_fake = discriminator(fake_sample) #real: 1, fake: 0 loss_real = F.binary_cross_entropy_with_logits(out_real, torch.ones_like(out_real)) loss_fake = F.binary_cross_entropy_with_logits(out_fake, torch.zeros_like(out_fake)) lossD = (loss_real + loss_fake)/2 #discriminator.zero_grad() lossD.backward() optimizer_d.step() #train generator lossD_fakeSample = discriminator(fake_sample) lossG = F.binary_cross_entropy_with_logits(lossD_fakeSample, torch.ones_like(lossD_fakeSample)) #generator.zero_grad() lossG.backward() optimizer_g.step() lossesG.append(lossG.item()) lossesD.append(lossD.item()) print(f'Epoch: {epoch:03d}, LossG: {lossG:.4f}, LossD: {lossD:.4f}')
真实数据的邻接矩阵稀疏度较低,但生成的矩阵并不稀疏。我尝试添加稀疏度惩罚:
sparsity_fake = find_sparsity(fake_sample) sparsity_real = find_sparsity(data_real) criterion(torch.tensor([sparsity_real]), torch.tensor([sparsity_fake]))
损失函数定义为:
criterion = nn.CrossEntropyLoss()
但将该稀疏性损失加入lossG(lossG += sparsity_loss)时,出现错误:
RuntimeError: element 0 of tensors does not require grad and does not have a grad_fn
错误原因
用torch.tensor()重新包装稀疏度值,导致新创建的张量脱离了原计算图,无法追踪梯度信息,进而在反向传播时出错。此外,CrossEntropyLoss是用于分类任务的损失函数,不适合稀疏度这种回归类的差值计算。
解决方案
1. 修改稀疏度计算函数,返回可追踪梯度的张量
将稀疏度计算从标量改为基于张量的运算,确保结果保留在计算图中:
def find_sparsity(graph): # 将运算转为张量操作,保留梯度 edge_count = torch.tensor(graph.edge_index.shape[1], dtype=torch.float32, device=graph.x.device) node_count = torch.tensor(graph.num_nodes, dtype=torch.float32, device=graph.x.device) return edge_count / (node_count ** 2)
2. 替换合适的损失函数
使用均方误差损失(MSELoss)或L1损失(L1Loss)来计算稀疏度的差值,这些损失适合回归任务:
# 替换损失函数 sparsity_criterion = nn.MSELoss() # 或者用L1Loss,根据需求选择 # sparsity_criterion = nn.L1Loss()
3. 正确整合稀疏度损失到生成器损失
在训练生成器的步骤中,计算稀疏度损失并加入总损失,确保所有张量都在计算图内:
# 训练generator部分修改如下: # 注意:移除fake_sample的detach(),否则生成器无法得到梯度 fake_sample = generator(data_fake) # 去掉.detach() # ...(其他判别器训练代码不变) # 训练生成器时 lossD_fakeSample = discriminator(fake_sample) lossG = F.binary_cross_entropy_with_logits(lossD_fakeSample, torch.ones_like(lossD_fakeSample)) # 计算稀疏度损失 sparsity_fake = find_sparsity(fake_sample) sparsity_real = find_sparsity(data_real) sparsity_loss = sparsity_criterion(sparsity_fake, sparsity_real) # 加入权重系数平衡两个损失的影响,比如0.1可根据实际情况调整 lossG += 0.1 * sparsity_loss lossG.backward() optimizer_g.step()
额外注意
- 移除
fake_sample = generator(data_fake).detach()中的.detach(),否则生成器的参数无法通过反向传播更新。 - 给稀疏度损失添加权重系数(如0.1),避免稀疏度损失主导整个生成器的训练目标,可根据训练效果调整系数大小。
内容的提问来源于stack exchange,提问作者Maryam
相关产品推荐
相关产品推荐

