PyTorch中如何构造带约束的分布并实现采样与log_prob计算?
带约束二维均匀分布的PyTorch实现
要构造满足m1~U(5,80)、m2~U(5,80)且m1+m2 < 100的分布,可以通过自定义PyTorch分布类结合约束规则实现,完整代码和说明如下:
1. 核心实现逻辑
- 自定义约束规则,适配
torch.distributions.constraints的接口规范 - 继承
distributions.Distribution实现自定义分布,分别完成采样和对数概率计算逻辑- 采样使用拒绝采样方案:先从无约束的二维均匀分布采样,过滤不满足约束的样本,直到得到目标数量的有效样本
- 对数概率计算:有效区域内的概率密度为固定值,超出区域返回负无穷
2. 完整代码
import torch from torch import distributions from torch.distributions import constraints # 自定义二维变量和小于100的约束 class SumLessThan100(constraints.Constraint): def check(self, value): # value的最后一维为(m1, m2),返回每个样本是否满足约束 return value[..., 0] + value[..., 1] < 100 # 自定义带约束的二维均匀分布 class Constrained2DUniform(distributions.Distribution): # 定义参数约束 arg_constraints = { "low": constraints.greater_than_eq(5.0), "high": constraints.less_than_eq(80.0) } # 定义分布的样本满足的约束 support = SumLessThan100() has_rsample = False def __init__(self, low=torch.tensor([5.0, 5.0]), high=torch.tensor([80.0, 80.0]), validate_args=None): self.base_dist = distributions.Uniform(low, high) # 预计算有效区域的对数概率密度:有效区域面积为3825,密度为1/3825 self.valid_log_prob = torch.log(torch.tensor(1 / 3825.0)) super().__init__(batch_shape=self.base_dist.batch_shape[:-1], event_shape=torch.Size([2]), validate_args=validate_args) def sample(self, sample_shape=torch.Size()): samples = torch.empty(sample_shape + self.event_shape, dtype=self.base_dist.low.dtype) remaining = sample_shape.numel() offset = 0 while remaining > 0: # 每次采剩余数量2倍的样本,减少循环次数 candidate = self.base_dist.sample(torch.Size([remaining * 2])) mask = self.support.check(candidate) valid = candidate[mask] take = min(remaining, len(valid)) samples.flatten(0, -2)[offset:offset+take] = valid[:take] remaining -= take offset += take return samples def log_prob(self, value): # 先检查是否在基础分布的范围内 base_valid = self.base_dist.support.check(value).all(dim=-1) # 再检查是否满足和的约束 sum_valid = self.support.check(value) # 所有约束都满足返回预计算的对数概率,否则返回负无穷 return torch.where(base_valid & sum_valid, self.valid_log_prob, torch.tensor(-float("inf")))
3. 使用示例
# 初始化分布 prior = Constrained2DUniform() # 采样10个样本 samples = prior.sample((10,)) # 计算样本的对数概率 log_probs = prior.log_prob(samples)
内容的提问来源于stack exchange,提问作者IPhysResearch
相关产品推荐
相关产品推荐

