You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何手动计算VAE中KL散度损失的梯度更新?

Got it, let's walk through how to manually compute the gradient updates for the KL divergence loss with respect to the parameters of your fc21 linear layer. I'll break this down step by step so it's easy to follow.

Manual Gradient Calculation for VAE KL Divergence (fc21 Layer)

First, let's define all the variables clearly to avoid confusion:

  • x: Your input 400-dimensional feature vector
  • fc21: Linear layer with weight matrix W (shape: 20x400) and bias vector b (shape: 20,)
  • z_logvar: Output of fc21, calculated as z_logvar = W @ x + b (shape: 20,)
  • mu: Mean vector from the companion linear layer (fc22, shape:20,) – note this doesn't affect gradients for fc21 since it's independent of fc21's parameters
  • KL divergence loss: KLD = -0.5 * torch.sum(1 + z_logvar - mu.pow(2) - z_logvar.exp()) (summed over all 20 dimensions)

Step 1: Compute Derivative of KLD with Respect to z_logvar

For each dimension i (0 to 19), let's look at the per-dimension KL term:

kld_i = -0.5 * (1 + z_logvar[i] - mu[i]² - exp(z_logvar[i]))

Taking the derivative of kld_i with respect to z_logvar[i] (the only variable here tied to fc21):

d_kld_d_logvar_i = -0.5 * (1 - exp(z_logvar[i])) = 0.5 * (exp(z_logvar[i]) - 1)

Collect these values into a 20-dimensional vector d_kld_d_logvar – each element corresponds to the derivative for one dimension of z_logvar.


Step 2: Compute Gradients for fc21's Bias Vector (b)

Since z_logvar[i] = sum(W[i][j] * x[j]) + b[i], the derivative of z_logvar[i] with respect to b[i] is 1. Using the chain rule:

  • For each bias element b[i]:
    d_kld_d_b_i = d_kld_d_logvar_i * 1 = 0.5 * (exp(z_logvar[i]) - 1)
    

The full bias gradient vector d_kld_d_b is exactly the d_kld_d_logvar vector from Step 1 – no extra computation needed here.


Step 3: Compute Gradients for fc21's Weight Matrix (W)

For each weight element W[i][j] (i-th row, j-th column):

  • The derivative of z_logvar[i] with respect to W[i][j] is x[j]. Applying the chain rule:
    d_kld_d_W_ij = d_kld_d_logvar_i * x[j] = 0.5 * (exp(z_logvar[i]) - 1) * x[j]
    

In matrix terms (easier to compute efficiently), this is the outer product of d_kld_d_logvar and x, scaled by 0.5:

d_kld_d_W = 0.5 * (d_kld_d_logvar.unsqueeze(1) @ x.unsqueeze(0))

This gives you a 20x400 matrix matching the shape of W.


Step 4: Apply the Gradient Update

Once you have d_kld_d_W and d_kld_d_b, you can update the parameters using your optimizer's rule. For example, with basic SGD:

W = W - learning_rate * d_kld_d_W
b = b - learning_rate * d_kld_d_b

Note: In practice, you'll add these gradients to the gradients from the VAE's reconstruction loss before updating – this step only covers the KL divergence portion.


Quick Verification Tip

If you want to confirm your manual calculations are correct, compare them to PyTorch's autograd results. Run:

import torch
from torch import nn

# Set up your layer and sample input
fc21 = nn.Linear(400,20)
x = torch.randn(400, requires_grad=True)
z_logvar = fc21(x)
mu = torch.randn(20)  # Independent of fc21
KLD = -0.5 * torch.sum(1 + z_logvar - mu.pow(2) - z_logvar.exp())

# Get autograd gradients
autograd_W_grad, autograd_b_grad = torch.autograd.grad(KLD, fc21.parameters())

Your manually computed d_kld_d_W and d_kld_d_b should match these autograd results exactly.

内容的提问来源于stack exchange,提问作者Matt

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.05.21 07:21:38