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

使用TensorFlow Probability JointDistributionNamed时HMC目标函数报错求助

问题描述

正在通过《Rethinking》学习贝叶斯方法,用TensorFlow Probability(TFP)实践,复现对应章节模型时,想用JointDistributionNamedAutoBatched定义模型(贴合书中写法),但运行HMC时在target_log_prob_fn上报错:

error: target_log_prob_fn() takes 0 positional arguments but 3 were given

原代码如下:

import tensorflow_probability as tfp
from tensorflow_probability import distributions as tfd
tfb = tfp.bijectors
import pandas as pd
d = pd.read_csv('rethinking-master/data/Howell1.csv',sep = ';')
height= d.height
model = tfd.JointDistributionNamedAutoBatched(dict(
    s = tfd.Sample(tfd.Exponential(1), sample_shape=1),
    alpha = tfd.Sample(tfd.Normal(0,1), sample_shape=1),
    beta = tfd.Sample(tfd.Normal(0,1), sample_shape=1),
    y = lambda s,alpha,beta: tfd.Independent(tfd.Normal( alpha + beta * d.weight.values,s),
                                          reinterpreted_batch_ndims=1),
))

def _trace_fn_transitioned(_, pkr):
    return pkr.inner_results.inner_results.log_accept_ratio

num_chains = 4
num_leapfrog_steps = 4
step_size = 0.8
burnin = 500
params = ['s', 'alpha', 'beta']
init_state = list(model.sample(num_chains))[:-1]
bijectors = [tfb.Identity() for i in init_state]
observed_data=(d.height.values)

def target_log_prob_fn (**x):
    print(x[0])
    return model.log_prob(model.sample(y = observed_data, **x))


hmc_kernel = tfp.mcmc.HamiltonianMonteCarlo(
        target_log_prob_fn, num_leapfrog_steps=num_leapfrog_steps, step_size=step_size
    )

inner_kernel = tfp.mcmc.TransformedTransitionKernel(
    inner_kernel=hmc_kernel, bijector=bijectors
)

kernel = tfp.mcmc.SimpleStepSizeAdaptation(
    inner_kernel=inner_kernel,
    target_accept_prob=0.8,
    num_adaptation_steps=int(0.8 * burnin),
    log_accept_prob_getter_fn=lambda pkr: pkr.inner_results.log_accept_ratio,
)

tfp.mcmc.sample_chain(
        num_results=544,
        num_burnin_steps=burnin,
        current_state=init_state,
        kernel=kernel,
        trace_fn=_trace_fn_transitioned,
    )
解决方案

核心问题

TFP的HMC内核会将current_state中的参数按顺序位置传递给target_log_prob_fn,但你定义的函数用了**x(关键字参数),且尝试用x[0]访问(关键字参数是字典,不能用索引),导致参数传递失败。同时,用model.sample计算对数概率的方式也是错误的。

修正步骤

  1. 调整参数形式:将target_log_prob_fn改为接收位置参数,对应模型中的s, alpha, beta
  2. 正确计算对数概率:直接用model.log_prob传入观测数据和参数,无需调用model.sample
  3. 可靠初始化状态:显式按参数名提取初始化状态,避免依赖sample返回的顺序

修正后的代码

import tensorflow_probability as tfp
from tensorflow_probability import distributions as tfd
tfb = tfp.bijectors
import pandas as pd

d = pd.read_csv('rethinking-master/data/Howell1.csv', sep=';')
observed_data = d.height.values

# 定义联合分布模型
model = tfd.JointDistributionNamedAutoBatched(dict(
    s=tfd.Sample(tfd.Exponential(1.0), sample_shape=1),
    alpha=tfd.Sample(tfd.Normal(0.0, 1.0), sample_shape=1),
    beta=tfd.Sample(tfd.Normal(0.0, 1.0), sample_shape=1),
    y=lambda s, alpha, beta: tfd.Independent(
        tfd.Normal(alpha + beta * d.weight.values, s),
        reinterpreted_batch_ndims=1
    ),
))

def _trace_fn_transitioned(_, pkr):
    return pkr.inner_results.inner_results.log_accept_ratio

num_chains = 4
num_leapfrog_steps = 4
step_size = 0.8
burnin = 500

# 显式提取初始化参数,避免依赖返回顺序
init_state = [model.sample(num_chains)[name] for name in ['s', 'alpha', 'beta']]
bijectors = [tfb.Identity() for _ in init_state]

# 修正后的目标对数概率函数
def target_log_prob_fn(s, alpha, beta):
    return model.log_prob(dict(
        s=s,
        alpha=alpha,
        beta=beta,
        y=observed_data
    ))

# 构建HMC内核
hmc_kernel = tfp.mcmc.HamiltonianMonteCarlo(
    target_log_prob_fn=target_log_prob_fn,
    num_leapfrog_steps=num_leapfrog_steps,
    step_size=step_size
)

inner_kernel = tfp.mcmc.TransformedTransitionKernel(
    inner_kernel=hmc_kernel,
    bijector=bijectors
)

kernel = tfp.mcmc.SimpleStepSizeAdaptation(
    inner_kernel=inner_kernel,
    target_accept_prob=0.8,
    num_adaptation_steps=int(0.8 * burnin),
    log_accept_prob_getter_fn=lambda pkr: pkr.inner_results.log_accept_ratio,
)

# 执行采样
samples, trace = tfp.mcmc.sample_chain(
    num_results=544,
    num_burnin_steps=burnin,
    current_state=init_state,
    kernel=kernel,
    trace_fn=_trace_fn_transitioned,
)

额外说明

  • s的先验用Exponential已经保证其为正数,因此不需要额外的bijector转换
  • model.log_prob需要传入包含所有变量(参数+观测数据)的字典,这是联合分布计算观测数据对数概率的标准方式

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.05 06:25:24