Skip to content
AI360Xpert
Beta

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.

Wasserstein GAN replaces unstable Jensen-Shannon divergence with Earth Mover's Distance, using a 1-Lipschitz critic regularized by gradient penalties.
Wasserstein GAN replaces unstable Jensen-Shannon divergence with Earth Mover's Distance, using a 1-Lipschitz critic regularized by gradient penalties.

Why Does This Exist?

The original Generative Adversarial Network (GAN) framed generation as a zero-sum minimax classification game. The discriminator DD outputs a probability in [0,1][0, 1] 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 Pr\mathbb{P}_r and the generated distribution Pg\mathbb{P}_g.

However, in high-dimensional spaces (such as 512×512×3512 \times 512 \times 3 images), both Pr\mathbb{P}_r and Pg\mathbb{P}_g 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 log⁡(2)≈0.693\log(2) \approx 0.693. Because the divergence curve is flat, its gradient with respect to generator parameters vanishes to zero:

∇θDJS(Pr∥Pgθ)=0\nabla_\theta D_{\text{JS}}(\mathbb{P}_r \parallel \mathbb{P}_{g_\theta}) = 0

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 Pr\mathbb{P}_r and Pg\mathbb{P}_g is defined as the infimum over all joint probability distributions γ∈Π(Pr,Pg)\gamma \in \Pi(\mathbb{P}_r, \mathbb{P}_g):

W(Pr,Pg)=inf⁡γ∈Π(Pr,Pg)E(x,y)∼γ[∥x−y∥]W(\mathbb{P}_r, \mathbb{P}_g) = \inf_{\gamma \in \Pi(\mathbb{P}_r, \mathbb{P}_g)} \mathbb{E}_{(x, y) \sim \gamma}[\|x - y\|]

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:

W(Pr,Pg)=sup⁡∥f∥L≤1Ex∼Pr[f(x)]−Ex~∼Pg[f(x~)]W(\mathbb{P}_r, \mathbb{P}_g) = \sup_{\|f\|_L \le 1} \mathbb{E}_{x \sim \mathbb{P}_r}[f(x)] - \mathbb{E}_{\tilde{x} \sim \mathbb{P}_g}[f(\tilde{x})]

where ∥f∥L≤1\|f\|_L \le 1 denotes that ff is a 1-Lipschitz continuous function satisfying ∣f(x1)−f(x2)∣≤∥x1−x2∥2|f(x_1) - f(x_2)| \le \|x_1 - x_2\|_2 for all points x1,x2x_1, x_2.

In WGAN, fwf_w is a neural network parameterized by ww, 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 [−c,c][-c, c] after each gradient step. Weight clipping suffered from severe drawbacks:

  • If cc is too large, the critic takes excessive steps to reach the optimal frontier, causing gradient explosion.
  • If cc 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: ∥∇xf(x)∥2≤1\|\nabla_x f(x)\|_2 \le 1. 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 x∼Prx \sim \mathbb{P}_r and fake points x~=G(z)\tilde{x} = G(z):

x^=ϵx+(1−ϵ)x~,ϵ∼U[0,1]\hat{x} = \epsilon x + (1 - \epsilon)\tilde{x}, \quad \epsilon \sim U[0, 1]

The complete WGAN-GP Critic loss function is:

Lcritic=Ex~∼Pg[D(x~)]−Ex∼Pr[D(x)]⏟Wasserstein Distance Objective+λEx^[(∥∇x^D(x^)∥2−1)2]⏟Gradient Penalty Term\mathcal{L}_{\text{critic}} = \underbrace{\mathbb{E}_{\tilde{x} \sim \mathbb{P}_g}[D(\tilde{x})] - \mathbb{E}_{x \sim \mathbb{P}_r}[D(x)]}_{\text{Wasserstein Distance Objective}} + \underbrace{\lambda \mathbb{E}_{\hat{x}} \left[ \left(\|\nabla_{\hat{x}} D(\hat{x})\|_2 - 1\right)^2 \right]}_{\text{Gradient Penalty Term}}

where λ=10\lambda = 10 is the standard penalty coefficient.

Worked Example

Consider two parallel 1-dimensional uniform distributions: P0\mathbb{P}_0 uniform along the vertical line x=0,y∈[0,1]x=0, y \in [0, 1], and Pθ\mathbb{P}_\theta uniform along x=θ,y∈[0,1]x=\theta, y \in [0, 1] where θ>0\theta > 0:

  1. Jensen-Shannon Divergence Calculation: The union support is disjoint for any θ≠0\theta \ne 0. The mixture distribution is M=12(P0+Pθ)M = \frac{1}{2}(\mathbb{P}_0 + \mathbb{P}_\theta). DKL(P0∥M)=∫log⁡(2)dP0=log⁡(2)D_{\text{KL}}(\mathbb{P}_0 \parallel M) = \int \log(2) d\mathbb{P}_0 = \log(2) DKL(Pθ∥M)=∫log⁡(2)dPθ=log⁡(2)D_{\text{KL}}(\mathbb{P}_\theta \parallel M) = \int \log(2) d\mathbb{P}_\theta = \log(2) DJS(P0∥Pθ)=12log⁡(2)+12log⁡(2)=log⁡(2)≈0.6931D_{\text{JS}}(\mathbb{P}_0 \parallel \mathbb{P}_\theta) = \frac{1}{2}\log(2) + \frac{1}{2}\log(2) = \log(2) \approx 0.6931 Notice that DJSD_{\text{JS}} is a flat constant. The derivative ddθDJS=0.0\frac{d}{d\theta} D_{\text{JS}} = 0.0. The generator receives zero gradient.

  2. Wasserstein-1 Distance Calculation: To move probability mass from x=θx=\theta to x=0x=0, every point must travel a horizontal Euclidean distance of exactly ∣θ∣|\theta|. W(P0,Pθ)=∣θ∣W(\mathbb{P}_0, \mathbb{P}_\theta) = |\theta| The derivative is ddθW=1.0\frac{d}{d\theta} W = 1.0 (for θ>0\theta > 0). The gradient is non-zero, constant, and points directly toward the target manifold regardless of distance!

  3. Gradient Penalty Numeric Step: Let real sample x=0.0x = 0.0 and generated sample x~=2.0\tilde{x} = 2.0. Sample ϵ=0.25  ⟹  x^=0.25(0.0)+0.75(2.0)=1.50\epsilon = 0.25 \implies \hat{x} = 0.25(0.0) + 0.75(2.0) = 1.50. Suppose the critic computes gradient ∇x^D(x^)=1.40\nabla_{\hat{x}} D(\hat{x}) = 1.40. Gradient norm is ∥∇x^D(x^)∥2=1.40\|\nabla_{\hat{x}} D(\hat{x})\|_2 = 1.40. With λ=10\lambda = 10: Penalty=10×(1.40−1.0)2=10×(0.40)2=10×0.16=1.60\text{Penalty} = 10 \times (1.40 - 1.0)^2 = 10 \times (0.40)^2 = 10 \times 0.16 = 1.60 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.9412

Watch 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 xix_i depends on the mean and variance of all other samples in the mini-batch. Consequently, the gradient ∇x^iD(x^)\nabla_{\hat{x}_i} D(\hat{x}) is no longer the gradient of an isolated mapping f:X→Rf: \mathcal{X} \to \mathbb{R}, 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 (log⁡2\log 2) 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.