在TensorFlow中参数化混合密度网络的协方差矩阵
嘿,我在构建带全协方差组件的MDN时,也纠结过协方差矩阵的输出约束问题,用Cholesky分解确实是最优解之一——既保证了正定,又能高效存储和计算。下面给你一步步拆解实现思路:
为什么选Cholesky因子?
直接输出全协方差矩阵有个致命问题:没法保证矩阵是正定的(非正定矩阵会导致高斯分布无意义,甚至训练崩溃)。而Cholesky分解把协方差矩阵$\Sigma$拆成下三角矩阵$L$的乘积:$\Sigma = LLT$,只要$L$是下三角且**对角元为正**,$\Sigma$就自动是正定矩阵。同时,下三角矩阵只需要存储$D(D+1)/2$个独立元素(而非$D2$个),能节省不少计算和内存开销。
神经网络输出的映射方案
从隐藏层输出到三个关键组件(权重、均值、Cholesky因子)的映射要分开处理:
1. 混合权重(Mixing Weights)
- 从隐藏层输出$K$个原始值,然后用
softmax激活,确保所有组件的权重和为1:
训练时建议用pi_logits = torch.nn.Linear(hidden_dim, K)(hidden) pi = F.softmax(pi_logits, dim=1) # 形状:[batch_size, K]log_softmax配合logsumexp计算损失,数值更稳定。
2. 组件均值(Component Means)
- 每个组件对应一个$D$维均值,所以输出$K \times D$个值,直接reshape成
[batch_size, K, D]即可,不需要激活(均值可以是任意实数):mu_logits = torch.nn.Linear(hidden_dim, K*D)(hidden) mu = mu_logits.reshape(-1, K, D) # 形状:[batch_size, K, D]
3. Cholesky因子(Cholesky Factors)
这是核心部分,要输出形状为[batch_size, K, D, D]的下三角张量,且对角元为正:
- 第一步:计算需要输出的元素个数——每个下三角矩阵有$D(D+1)/2$个独立元素,$K$个组件就是$K \times D(D+1)/2$个值;
- 第二步:把这些值填充到全零张量的下三角位置(包括对角线);
- 第三步:对对角线元素用
softplus激活(替代exp,避免数值爆炸),确保为正;非对角元素保持线性输出(可正可负)。
PyTorch代码示例:
import torch import torch.nn.functional as F import math D = 3 # 输入/输出维度 K = 5 # 混合组件数 hidden_dim = 64 # 假设隐藏层输出(batch_size=2为例) hidden = torch.randn(2, hidden_dim) # 计算Cholesky分支的输出维度 chol_output_dim = K * D * (D + 1) // 2 chol_logits = torch.nn.Linear(hidden_dim, chol_output_dim)(hidden) # 形状:[2, K*D(D+1)/2] # 初始化全零的下三角张量 L = torch.zeros(chol_logits.shape[0], K, D, D, device=hidden.device) # 获取下三角矩阵的索引(行、列) tril_row, tril_col = torch.tril_indices(row=D, col=D, offset=0) # 填充下三角部分 L[:, :, tril_row, tril_col] = chol_logits.reshape(-1, K, D*(D+1)//2) # 处理对角线元素:用softplus确保为正 diag_idx = torch.arange(D) L[:, :, diag_idx, diag_idx] = F.softplus(L[:, :, diag_idx, diag_idx])
最终得到的L就是每个组件的下三角Cholesky因子,对应的协方差矩阵可以通过torch.bmm(L, L.transpose(2,3))计算。
损失函数计算(负对数似然)
MDN的损失是负对数似然,公式为:
$$\text{Loss} = -\frac{1}{N}\sum_{i=1}^N \log\left(\sum_{k=1}^K \pi_k(x_i) \cdot \mathcal{N}(y_i; \mu_k(x_i), \Sigma_k(x_i))\right)$$
利用Cholesky因子可以高效计算多元高斯的对数密度:
- $\log \det(\Sigma_k) = 2 \times \sum_{d=1}^D \log(L_{k,d,d})$(因为$\det(\Sigma) = (\prod \text{diag}(L))^2$)
- $(y-\mu_k)^T \Sigma_k^{-1} (y-\mu_k) = |L_k{-1}(y-\mu_k)|2$(用三角求解替代矩阵求逆,更稳定)
代码示例:
y = torch.randn(2, D) # 目标输出 # 计算每个组件的高斯对数密度 log_probs = [] for batch_idx in range(y.shape[0]): batch_log_probs = [] for k in range(K): mu_k = mu[batch_idx, k] L_k = L[batch_idx, k] # 计算L^{-1}(y - mu) diff = y[batch_idx] - mu_k z, _ = torch.triangular_solve(diff.unsqueeze(1), L_k, upper=False) z = z.squeeze(1) # 计算log det(Σ) log_det = 2 * torch.sum(torch.log(torch.diag(L_k))) # 对数密度 log_prob = -0.5 * (D * math.log(2*math.pi) + log_det + torch.sum(z**2)) batch_log_probs.append(log_prob) log_probs.append(torch.stack(batch_log_probs)) log_probs = torch.stack(log_probs) # 形状:[2, K] log_pi = F.log_softmax(pi_logits, dim=1) # 对数权重 # 计算负对数似然 loss = -torch.mean(torch.logsumexp(log_pi + log_probs, dim=1))
关键注意事项
- 数值稳定性:优先用
softplus处理Cholesky的对角元,比exp更不容易出现数值溢出;训练时用log_softmax和logsumexp计算损失,避免因概率值过小导致的下溢。 - 初始化:Cholesky分支的线性层可以把对角线对应的权重初始化为小正数(比如0.1),让初始的协方差矩阵接近单位矩阵,加快收敛。
- 批量优化:上面的循环可以用批量矩阵操作替换(比如
torch.linalg.triangular_solve的批量版本),提升训练速度。
内容的提问来源于stack exchange,提问作者jalsing
相关产品推荐
相关产品推荐

