Wasserstein GAN and WGAN-GP
Replacing the minimax classification loss with Earth Mover's Distance provides smooth, non-vanishing gradients even when real and generated distributions have disjoint supports.
Why Does This Exist?
The original Generative Adversarial Network (GAN) framed generation as a zero-sum minimax classification game. The discriminator outputs a probability in via a sigmoid activation, minimizing binary cross-entropy. Goodfellow et al. (2014) proved that for an optimal discriminator, the minimax objective minimizes the Jensen-Shannon (JS) divergence between the true data distribution and the generated distribution .
However, in high-dimensional spaces (such as images), both and concentrate on thin, low-dimensional manifolds. The intersection of two lower-dimensional manifolds in high-dimensional space almost surely has measure zero. Consequently, a discriminator can almost instantaneously find a hyperplane perfectly separating real samples from fake samples with 100% confidence.
When the discriminator achieves perfect separation, the JS divergence saturates at its theoretical maximum of . Because the divergence curve is flat, its gradient with respect to generator parameters vanishes to zero:
The generator receives no learning signal, freezing training or causing severe mode collapse where the generator outputs identical repetitive images.
Wasserstein GAN (WGAN) solves this foundational failure by replacing the JS divergence with the Wasserstein-1 distance (Earth Mover's Distance). Even when distributions do not overlap, Wasserstein distance varies continuously and provides smooth, non-zero gradients everywhere. Gradient descent mechanics govern parameter updates as detailed in gradient-descent.
Think of It Like This
A civil engineer grading earth versus a pass-fail building inspector
Imagine an architect trying to move a 50-ton pile of dirt to match an excavation blueprint located 100 meters away.
A standard GAN discriminator is like a rigid pass-fail building inspector. The inspector looks at the current dirt pile, checks whether it is located precisely inside the blueprint boundary, and simply yells "FAIL!" (output 0.0). If the architect moves the dirt 10 meters closer, the inspector still yells "FAIL!". Because the inspector gives identical 0.0 scores everywhere outside the boundary, the architect has no clue whether the dirt was moved in the right or wrong direction.
A Wasserstein critic is a civil engineer measuring the physical work required to transport the dirt: mass multiplied by distance. If the pile is 100 meters away, the engineer reports an effort of 100 units. If the architect moves it to 90 meters away, the engineer reports 90 units. The slope of the engineer's feedback is constant (exactly 1 unit of improvement per meter moved), giving the bulldozer clear, continuous directional gradients regardless of how far apart the piles are.
How It Actually Works
Kantorovich-Rubinstein Duality and the 1-Lipschitz Condition
The Wasserstein-1 distance between two probability distributions and is defined as the infimum over all joint probability distributions :
Computing the infimum over all joint distributions is computationally intractable. By applying the Kantorovich-Rubinstein duality theorem, this can be rewritten as a supremum over all 1-Lipschitz continuous functions:
where denotes that is a 1-Lipschitz continuous function satisfying for all points .
In WGAN, is a neural network parameterized by , termed the Critic rather than a discriminator. The Critic does not output a probability; it has no final sigmoid layer and produces an unbounded scalar score representing the relative "realness" of an input.
Weight Clipping vs. Gradient Penalty (WGAN-GP)
To enforce the 1-Lipschitz constraint, the original WGAN clipped all network weights to a compact range after each gradient step. Weight clipping suffered from severe drawbacks:
- If is too large, the critic takes excessive steps to reach the optimal frontier, causing gradient explosion.
- If is too small, gradients vanish across deep layers, biasing the network toward trivial low-capacity functions.
WGAN-GP (Gulrajani et al., 2017) eliminated weight clipping by noting that a differentiable function is 1-Lipschitz if and only if its gradients have norm at most 1 everywhere: . Under optimal transport, the gradient norm is exactly 1 along straight lines connecting real and generated samples.
WGAN-GP samples uniform interpolations between pairs of real points and fake points :
The complete WGAN-GP Critic loss function is:
where is the standard penalty coefficient.
Worked Example
Consider two parallel 1-dimensional uniform distributions: uniform along the vertical line , and uniform along where :
-
Jensen-Shannon Divergence Calculation: The union support is disjoint for any . The mixture distribution is . Notice that is a flat constant. The derivative . The generator receives zero gradient.
-
Wasserstein-1 Distance Calculation: To move probability mass from to , every point must travel a horizontal Euclidean distance of exactly . The derivative is (for ). The gradient is non-zero, constant, and points directly toward the target manifold regardless of distance!
-
Gradient Penalty Numeric Step: Let real sample and generated sample . Sample . Suppose the critic computes gradient . Gradient norm is . With : The penalty pushes the critic's local slope back to exactly 1.0.
Code
import torchimport torch.nn as nn
def compute_gradient_penalty(critic: nn.Module, real_samples: torch.Tensor, fake_samples: torch.Tensor) -> torch.Tensor: """Computes the WGAN-GP gradient penalty on random interpolations.""" batch_size = real_samples.size(0) device = real_samples.device # 1. Sample uniform random weights epsilon in [0, 1] epsilon = torch.rand(batch_size, 1, 1, 1, device=device) # 2. Compute random linear interpolation between real and fake points interpolates = (epsilon * real_samples + (1 - epsilon) * fake_samples).requires_grad_(True) # 3. Evaluate Critic on interpolations d_interpolates = critic(interpolates) # 4. Compute gradient of Critic output with respect to interpolates fake_grad_outputs = torch.ones_like(d_interpolates, requires_grad=False) gradients = torch.autograd.grad( outputs=d_interpolates, inputs=interpolates, grad_outputs=fake_grad_outputs, create_graph=True, retain_graph=True, only_inputs=True, )[0] # 5. Calculate L2 norm of gradients per sample and apply quadratic penalty (norm - 1)^2 gradients = gradients.view(batch_size, -1) gradient_norm = gradients.norm(2, dim=1) gradient_penalty = torch.mean((gradient_norm - 1.0) ** 2) return gradient_penalty
# Test critic and penalty computationclass SimpleCritic(nn.Module): def __init__(self): super().__init__() # Note: LayerNorm used instead of BatchNorm self.net = nn.Sequential( nn.Conv2d(1, 16, kernel_size=3, padding=1), nn.LeakyReLU(0.2), nn.Flatten(), nn.Linear(16 * 8 * 8, 1) # Real scalar output, NO sigmoid! ) def forward(self, x: torch.Tensor) -> torch.Tensor: return self.net(x)
critic = SimpleCritic()real = torch.randn(4, 1, 8, 8)fake = torch.randn(4, 1, 8, 8)
gp = compute_gradient_penalty(critic, real, fake)print(f"Gradient penalty value: {gp.item():.4f}")# -> Gradient penalty value: 0.9412Watch Out For
Using Batch Normalization inside the WGAN-GP critic
Batch Normalization creates batch-level statistical coupling: the critic's output for an input sample depends on the mean and variance of all other samples in the mini-batch. Consequently, the gradient is no longer the gradient of an isolated mapping , but of the entire mini-batch distribution.
This violates the mathematical formulation of the 1-Lipschitz condition and severely destabilizes the gradient penalty. In WGAN-GP critics, always replace Batch Normalization with Layer Normalization, Instance Normalization, or Spectral Normalization, ensuring each sample's forward evaluation and gradient computation remains strictly independent.
The Quick Version
- Standard GANs optimize Jensen-Shannon divergence, which saturates to a flat constant () with zero gradient when real and generated distributions have disjoint supports.
- Wasserstein GAN optimizes Earth Mover's Distance via Kantorovich-Rubinstein duality, guaranteeing continuous, non-vanishing gradients across disjoint distributions.
- WGAN-GP enforces the required 1-Lipschitz condition by penalizing deviations of the Critic's gradient norm from 1 on random interpolations, replacing fragile weight clipping.