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

PyTorch中分段均匀分布任意点CDF高效查询及梯度支持

问题描述

给定以下场景:

  • 离散分布示例:x=[1,2,4,5]; P(x)=[0.1, 0.2, 0.3, 0.4],对应的CDF为F(x)=[0.1, 0.3, 0.6, 1]
  • 实际场景为分段均匀分布:支撑集x ∈ [1, 9],分段定义如下:
    • x ∈ [1, 2],概率密度P(x)=0.1
    • x ∈ [2, 4],概率密度P(x)=0.1
    • x ∈ [4, 5],概率密度P(x)=0.3
    • x ∈ [5,9],概率密度P(x)=0.1

需要查询该分布在点y=[0.5, 1.6, 2.3, 3.4, 4.5, 5.7, 8.9]处的CDF,期望结果为F(y)=[0, 0.06, 0.13, 0.24, 0.45, 0.67, 0.99]。

输入张量形状要求:

  • x、P为(batch_size, num_bins)的2D数组
  • y为(batch_size, num_new_bins)的2D数组

需求:在PyTorch中高效实现该计算,且支持梯度反向传播,询问是否有可用库函数。


回答

PyTorch没有直接提供计算这种自定义分段均匀分布CDF的库函数,但可以用原生张量操作高效实现,且完全支持梯度反向传播。

核心逻辑

分段均匀分布的CDF计算本质是累加每个区间对目标点的概率贡献:

  • 对于目标点y,完全位于其左侧的区间,贡献该区间的全部概率(区间长度×密度)
  • 与y相交的区间(左端点≤y<右端点),贡献(y-左端点)×密度
  • 位于y右侧的区间无贡献

代码实现

import torch

def compute_piecewise_uniform_cdf(x_left, x_right, density, query_points):
    """
    计算批量分段均匀分布的CDF,支持自动微分
    参数:
        x_left: (batch_size, num_intervals) 每个分段区间的左端点
        x_right: (batch_size, num_intervals) 每个分段区间的右端点
        density: (batch_size, num_intervals) 每个分段区间的概率密度
        query_points: (batch_size, num_queries) 需要计算CDF的点
    返回:
        cdf_values: (batch_size, num_queries) 对应查询点的CDF值
    """
    # 扩展维度实现广播计算
    x_left_exp = x_left.unsqueeze(-1)  # (B, N, 1)
    x_right_exp = x_right.unsqueeze(-1)  # (B, N, 1)
    density_exp = density.unsqueeze(-1)  # (B, N, 1)
    y_exp = query_points.unsqueeze(1)  # (B, 1, M)
    
    # 计算完全区间的贡献和部分区间的贡献
    full_interval_prob = density_exp * (x_right_exp - x_left_exp)
    full_mask = x_right_exp <= y_exp
    
    partial_interval_prob = density_exp * (y_exp - x_left_exp)
    partial_mask = (x_left_exp <= y_exp) & (y_exp < x_right_exp)
    
    # 合并贡献
    total_contrib = torch.where(full_mask, full_interval_prob,
                                torch.where(partial_mask, partial_interval_prob, torch.zeros_like(full_interval_prob)))
    
    # 累加得到CDF
    cdf = total_contrib.sum(dim=1)
    
    # 确保CDF在0到总概率之间(处理归一化和边界情况)
    total_prob = (density * (x_right - x_left)).sum(dim=1, keepdim=True)
    cdf = torch.clamp(cdf, min=0.0, max=total_prob)
    
    return cdf

# 验证示例
if __name__ == "__main__":
    # 示例输入(batch_size=1)
    x_left = torch.tensor([[1.0, 2.0, 4.0, 5.0]])
    x_right = torch.tensor([[2.0, 4.0, 5.0, 9.0]])
    density = torch.tensor([[0.1, 0.1, 0.3, 0.1]])
    query_points = torch.tensor([[0.5, 1.6, 2.3, 3.4, 4.5, 5.7, 8.9]])
    
    computed_cdf = compute_piecewise_uniform_cdf(x_left, x_right, density, query_points)
    print("计算结果:", computed_cdf.round(decimals=2))
    print("期望结果:", torch.tensor([[0.00, 0.06, 0.13, 0.24, 0.45, 0.67, 0.99]]))

细节说明

  • 用张量广播替代循环,处理批量输入效率极高,适合大尺寸的batch_size和num_bins场景
  • 所有操作都是PyTorch可微分算子,梯度可以正常反向传播
  • 最终的clamp操作确保CDF值符合概率分布的范围,避免因数值误差出现异常值

内容的提问来源于stack exchange,提问作者Nagabhushan S N

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.06 16:25:26