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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.10.03 05:57:01