Provided are GAN losses from the R3GAN and Seaweed papers.
pip install relativistic_loss
from relativistic_loss.loss import (
saturating_gan_loss,
gan_loss,
r1_penalty,
r2_penalty,
approximate_r1_loss,
approximate_r2_loss
)Logistic f:
The saturating loss is described in the paper, but in the R3GAN implementation, they use the non saturating version as it performs better practically.
# If f is omitted, logistic_f is used automatically
d_loss = -saturating_gan_loss(discriminator, generator, real_images, z)
# If f is omitted, softplus is used automatically
d_loss = gan_loss(discriminator, generator, real_images, z, discriminator_turn=True)
# You have to freeze the other model when doing a backward pass for generator or discriminator, otherwise you will combine the negative and positive gradients which will cancel out.
g_loss = gan_loss(discriminator, generator, real_images, z)
g_loss = gan_loss(discriminator, generator, real_images, z, discriminator_turn=False)def hinge_f(t):
return torch.nn.functional.relu(1 - t)
d_loss = gan_loss(
discriminator, generator, real_images, z, discriminator_turn=True,
f=hinge_f
)d_loss = gan_loss(
discriminator, generator,
real_images, z,
discriminator_turn=True,
generator_args=(some_label, ), # positional
generator_kwargs={'some_flag': True}, # keyword
disc_args=(some_label, ),
disc_kwargs={'cond_flag': True}
)d_loss += r1_penalty(discriminator, real_images, gamma=1.0, disc_args=args, disc_kwargs=kwargs)
d_loss += r2_penalty(discriminator, fake_images, gamma=1.0, disc_args=args, disc_kwargs=kwargs)The idea is that, the gradient penalty exists to disencourage large changes in the logits from a small change in the input. The approximation just adds a bit of noise to the input, which works similarly to taking the gradient in this case. This approximation is very helpful as FlashAttention cannot take a second derivative, so the true penalties are much slower to compute.
d_loss += approximate_r1_loss(discriminator, real_images, sigma=0.01, Lambda=100.0, disc_args=disc_args, disc_kwargs=disc_kwargs)
d_loss += approximate_r2_loss(discriminator, fake_images, sigma=0.01, Lambda=100.0, disc_args=disc_args, disc_kwargs=disc_kwargs)Generated by running python test.py
Non saturating loss (gan_loss) + Real R1 + R2 gradient penalty
Saturating loss (saturating_gan_loss) + Real R1 + R2 gradient penalty
Non saturating loss (gan_loss) + Approximate R1 + R2 gradient penalty
Saturating loss (saturating_gan_loss) + Approximate R1 + R2 gradient penalty



