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

在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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.25 03:30:49