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

PyTorch中不中断梯度流的torch.unique()替代方案及离散熵计算疑问

PyTorch离散分布熵计算的梯度中断问题解决

问题背景

在PyTorch梯度下降流程中,定义了如下计算离散分布香农熵的函数:

def TShentropy(wf):
    unique_elements, counts = wf.unique(return_counts=True)
    entrsum = 0
    for x in counts:
        p = x/len_a # len_a应为wf的长度,需确保已提前定义
        entrsum -= p*torch.log2(p) # 香农熵计算公式
    return entrsum

由于torch.unique()是不可微分操作,会直接中断梯度流。尝试改用torch.nn.functional.one_hot配合torch.bincount计算类别计数时,同样触发梯度错误:

RuntimeError: element 0 of tensors does not require grad and does not have a grad_fn

需求明确要求使用离散概率质量分布,因此torch.softmax()这类连续概率计算方式不适用。

可行解决方案

直接使用离散计数操作(unique/bincount)必然会中断梯度——这类操作是离散、非连续的,无法进行反向传播梯度计算。要在保留离散分布语义的同时维持梯度流,有两种实用思路:

1. 可微分近似离散计数

将离散计数转化为连续可微分的计算,通过温度系数控制结果与真实离散分布的逼近程度:

def differentiable_ts_entropy(wf, num_classes, temperature=1.0):
    # 将类别索引转为one-hot张量
    one_hot = torch.nn.functional.one_hot(wf, num_classes=num_classes).float()
    # 计算"软化计数":温度越低,结果越接近真实离散计数
    soft_counts = torch.sum(torch.nn.functional.softmax(one_hot / temperature, dim=0), dim=0)
    # 归一化得到概率分布
    p = soft_counts / torch.sum(soft_counts)
    # 避免log2(0)的数值问题
    p = p.clamp(min=1e-10)
    # 计算香农熵
    entropy = -torch.sum(p * torch.log2(p))
    return entropy

训练初期可适当提高温度保证梯度稳定,后期降低温度让结果逼近真实离散分布。

2. Gumbel-Softmax松弛离散分布

如果wf是由模型输出的logits生成的离散类别,训练阶段可用Gumbel-Softmax技巧将离散采样转为连续可微分操作:

def gumbel_softmax_entropy(logits, temperature=1.0):
    # 生成Gumbel分布样本
    gumbel_dist = torch.distributions.Gumbel(torch.zeros_like(logits), torch.ones_like(logits))
    gumbel_samples = gumbel_dist.sample()
    # 得到松弛的连续分布
    soft_samples = torch.nn.functional.softmax((logits + gumbel_samples)/temperature, dim=-1)
    # 计算熵
    p = soft_samples.clamp(min=1e-10)
    entropy = -torch.sum(p * torch.log2(p), dim=-1).mean()
    return entropy

推理阶段直接对logits取argmax即可得到离散类别,完全符合离散分布需求。

结论

不需要对概率函数做插值处理,核心是用可微分的近似操作替代不可微分的离散计数/采样步骤,既能满足离散概率质量分布的要求,又能维持梯度流完成训练。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.22 13:02:27