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.1x ∈ [2, 4],概率密度P(x)=0.1x ∈ [4, 5],概率密度P(x)=0.3x ∈ [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
相关产品推荐
相关产品推荐

