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

基于Jax与Black Jax的MCMC营销组合模型接受概率为0问题

营销组合模型MCMC实现问题:Jax+BlackJax样本接受率始终为0

我使用Jax和BlackJax实现MCMC构建营销组合模型(Marketing Mix Model),但运行模型时样本接受概率始终为0,无法生成有效样本。经观察,初始参数与提议参数的对数后验概率差值为极大负值,推测这是导致该问题的原因。

实现代码

import jax
import jax.numpy as jnp
import numpy as np
import jax.scipy.stats as jsps
import blackjax
from jax.scipy.special import gammaln

def invgamma_logpdf(x, alpha, beta):return (alpha * jnp.log(beta)- gammaln(alpha)- (alpha + 1) * jnp.log(x)- beta / x)
def logit(x):return jnp.log(x) - jnp.log1p(-x)
def inv_logit(y):return jnp.exp(y) / (1 + jnp.exp(y))
def positive(x):return jnp.exp(x)
def positive_inv(y):return jnp.log(y)
def log_jacobian_logit(x):return -jnp.log(x) - jnp.log1p(-x)
def log_jacobian_positive(x):return x
def hill(t, ec, slope):return 1.0 / (1.0 + jnp.power(t / ec, -slope))
def adstock(x_lags, weights):return jnp.dot(x_lags, weights) / jnp.sum(weights)
def lag_weights_fn(retain_rate, max_lag):
    lags = jnp.arange(max_lag, dtype=jnp.float32)
    return retain_rate ** lags
def build_v_lag_weights_fn(max_lag):return jax.vmap(lambda retain_rate: lag_weights_fn(retain_rate, max_lag),in_axes=(0,))
def single_media_effect(x_lags, weights, ec, slope):
    a = adstock(x_lags, weights)
    return hill(a, ec, slope)
def compute_media_effects(X_media_t, lag_weights, ec, slope):return jax.vmap(single_media_effect)(X_media_t,lag_weights,ec,slope)
def compute_interactions(media_effects, left_idx, right_idx):return media_effects[left_idx] * media_effects[right_idx]
def compute_all_interactions(media_effects, left_indices, right_indices):return jnp.array([compute_interactions(media_effects, l, r)for l, r in zip(left_indices, right_indices)])
def predict_mu(tau,beta_medias,media_effects,X_ctrl,gamma_ctrl,beta_interactions,left_indices,right_indices):
    media_contrib = jnp.dot(media_effects, beta_medias)
    ctrl_contrib = jnp.dot(X_ctrl, gamma_ctrl)
    if beta_interactions.shape[0] > 0:
        inter_effects = compute_all_interactions(media_effects, left_indices, right_indices)
        inter_contrib = jnp.dot(inter_effects, beta_interactions)
    else:inter_contrib = 0.0
    return tau + media_contrib + ctrl_contrib + inter_contrib
v_predict_mu = jax.vmap(predict_mu,in_axes=(None, None, 0, 0, None, None, None, None))

def log_likelihood(y, mu, noise_var):return jnp.sum(jsps.norm.logpdf(y, mu, jnp.sqrt(noise_var)))
def log_prior_scalar(value, dist_type, a, b):
    if dist_type == "normal":return jsps.norm.logpdf(value, a, b)
    elif dist_type == "uniform":return jnp.where((value >= a) & (value <= b),-jnp.log(b - a), -jnp.inf)
    elif dist_type == "beta":return jsps.beta.logpdf(value, a, b)
    elif dist_type == "gamma":return jsps.gamma.logpdf(value, a, b)
    elif dist_type == "inv_gamma":return invgamma_logpdf(value, a, b)
    else:raise ValueError(f"Unknown distribution type: {dist_type}")
def compute_log_prior_dynamic(params, priors, transformed, log_det_jacobian):
    tau, beta_medias, gamma_ctrl, retain_rate, ec, beta_interactions, noise_var = transformed
    lp = 0.0
    lp += log_prior_scalar(tau,priors["tau_type"],priors["tau_mean"],priors["tau_sd"])
    lp += log_prior_scalar(noise_var,priors["noise_var_type"],priors["noise_var_mean"],priors["noise_var_sd"])
    for i, val in enumerate(beta_medias):
        lp += log_prior_scalar(val,priors["priors_beta_type"][i],priors["priors_beta_mean"][i],priors["priors_beta_sd"][i])
    for i, val in enumerate(retain_rate):
        lp += log_prior_scalar(val,priors["priors_retain_rate_type"][i],priors["priors_retain_rate_a"][i],priors["priors_retain_rate_b"][i])
    for i, val in enumerate(ec):
        lp += log_prior_scalar(val,priors["priors_ec_type"][i],priors["priors_ec_a"][i],priors["priors_ec_b"][i])
    for i, val in enumerate(gamma_ctrl):
        lp += log_prior_scalar(val,priors["priors_ctrl_type"][i],priors["priors_ctrl_mean"][i],priors["priors_ctrl_sd"][i])
    for i, val in enumerate(beta_interactions):
        lp += log_prior_scalar(val, "normal", 0, 1)
    return lp + log_det_jacobian

def transform_params(params):
    tau, beta_medias_raw, gamma_ctrl_raw, retain_rate_raw, ec_raw, beta_interactions_raw, noise_var_raw = params

    beta_medias = positive(beta_medias_raw)
    gamma_ctrl = gamma_ctrl_raw
    retain_rate = inv_logit(retain_rate_raw)
    ec = inv_logit(ec_raw)
    beta_interactions = positive(beta_interactions_raw)
    noise_var = positive(noise_var_raw)

    log_det_jacobian = jnp.sum(log_jacobian_positive(beta_medias_raw))
    log_det_jacobian += jnp.sum(log_jacobian_logit(retain_rate))
    log_det_jacobian += jnp.sum(log_jacobian_logit(ec))
    log_det_jacobian += jnp.sum(log_jacobian_positive(beta_interactions_raw))
    log_det_jacobian += log_jacobian_positive(noise_var_raw)

    return (tau,beta_medias,gamma_ctrl,retain_rate,ec,beta_interactions,noise_var),log_det_jacobian

def log_posterior_dynamic(params_unconstrained, data, priors, max_lag):
    transformed, log_det_jacobian = transform_params(params_unconstrained)
    tau, beta_medias, gamma_ctrl, retain_rate, ec, beta_interactions, noise_var = transformed

    v_lag_weights_fn = build_v_lag_weights_fn(max_lag)
    lag_weights = v_lag_weights_fn(retain_rate)

    media_effects_t = jax.vmap(
        compute_media_effects,
        in_axes=(0, None, None, None)
    )(data['X_media_train'], lag_weights, ec, priors["slope_values"])

    mu = v_predict_mu(
        tau,
        beta_medias,
        media_effects_t,
        data['X_ctrl_train'],
        gamma_ctrl,
        beta_interactions,
        data['interaction_left'],
        data['interaction_right']
    )
    print("mu mean:", jnp.mean(mu))
    print("mu std:", jnp.std(mu))
    print("Y_train mean:", jnp.mean(data["Y_train"]))
    print("Y_train std:", jnp.std(data["Y_train"]))


    ll = log_likelihood(data['Y_train'], mu, noise_var)
    lp = compute_log_prior_dynamic(
        params_unconstrained,
        priors,
        transformed,
        log_det_jacobian
    )
    return ll + lp

def run_blackjax_nuts(
    log_prob_fn,
    initial_params,
    num_samples=2000,
    step_size=0.001,
    inverse_mass_matrix=None
):
    flat_params, unravel_fn = jax.flatten_util.ravel_pytree(initial_params)
    if inverse_mass_matrix is None:
        inverse_mass_matrix = jnp.ones_like(flat_params)

    nuts_kernel = blackjax.nuts(
        log_prob_fn,
        step_size=step_size,
        inverse_mass_matrix=inverse_mass_matrix
    )

    initial_state = nuts_kernel.init(initial_params)
    rng_key = jax.random.PRNGKey(0)
    keys = jax.random.split(rng_key, num_samples)
    @jax.jit
    def one_step(state, rng_key):
        state, info = nuts_kernel.step(rng_key, state)
        return state, info


    states, infos = jax.lax.scan(one_step, initial_state, keys)
    return states, infos



dist_types_map = {
    0: "uniform",
    1: "normal",
    2: "beta",
    3: "gamma",
    4: "inv_gamma"
}

N = stan_data["N"]
T = stan_data["T"]
H = stan_data["H"]
num_media = stan_data["num_media"]
num_ctrl = stan_data["num_ctrl"]
max_lag = int(stan_data["max_lag"])
n_interactions = int(stan_data["n_interactions"])

training_index = np.array(stan_data["training_index"]) - 1

X_media = np.array(stan_data["X_media"])
X_ctrl = np.array(stan_data["X_ctrl"])
Y_train = np.array(stan_data["Y_train"])

X_media_train = jnp.array(X_media[training_index, :, :])
X_ctrl_train = jnp.array(X_ctrl[training_index, :])
Y_train = jnp.array(Y_train)
y_mean = jnp.mean(Y_train)
y_std = jnp.std(Y_train)
Y_train = ((Y_train - y_mean) / y_std)

interaction_left = (
    jnp.array(stan_data["interaction_left"], dtype=int)
    if n_interactions > 0
    else jnp.array([], dtype=int)
)
interaction_right = (
    jnp.array(stan_data["interaction_right"], dtype=int)
    if n_interactions > 0
    else jnp.array([], dtype=int)
)

data = dict(
    X_media_train=X_media_train,
    X_ctrl_train=X_ctrl_train,
    Y_train=Y_train,
    interaction_left=interaction_left,
    interaction_right=interaction_right,
)

priors = dict(
    tau_type=dist_types_map[stan_data["tau_dist_type"]],
    tau_mean=stan_data["tau_dist_mean"],
    tau_sd=stan_data["tau_dist_sd"],
    noise_var_type=dist_types_map[stan_data["noise_var_dist_type"]],
    noise_var_mean=stan_data["noise_var_dist_mean"],
    noise_var_sd=stan_data["noise_var_dist_sd"],
    priors_beta_type=[dist_types_map[t] for t in stan_data["media_prior_dist_type"]],
    priors_beta_mean=stan_data["media_prior_mean"],
    priors_beta_sd=stan_data["media_prior_sd"],
    priors_retain_rate_type=[dist_types_map[t] for t in stan_data["retain_rate_dist_type"]],
    priors_retain_rate_a=stan_data["retain_rate_dist_mean"],
    priors_retain_rate_b=stan_data["retain_rate_dist_sd"],
    priors_ec_type=[dist_types_map[t] for t in stan_data["ec_dist_type"]],
    priors_ec_a=stan_data["ec_dist_mean"],
    priors_ec_b=stan_data["ec_dist_sd"],
    priors_ctrl_type=[dist_types_map[t] for t in stan_data["ctrl_prior_dist_type"]],
    priors_ctrl_mean=stan_data["ctrl_prior_mean"],
    priors_ctrl_sd=stan_data["ctrl_prior_sd"],
    slope_values=jnp.array(stan_data["slope_dist_mean"]),
)

tau_init = float(jnp.mean(Y_train))
noise_var_init = jnp.log(jnp.var(Y_train))
retain_rate_raw = logit(jnp.full(num_media, 0.7))
ec_raw = logit(jnp.full(num_media, 0.3))
gamma_ctrl_raw = jnp.full(num_ctrl, 0.1)
initial_params = (
    tau_init,
    jnp.ones(num_media),
    gamma_ctrl_raw,
    retain_rate_raw,
    ec_raw,
    jnp.zeros(n_interactions),
    noise_var_init,
)
def logprob_fn(params):
    return log_posterior_dynamic(params, data, priors, max_lag)

# logprob_fn = jax.jit(logprob_fn)
logp = logprob_fn(initial_params)
print("Log posterior at new initial params:", logp)

states, infos = run_blackjax_nuts(
    logprob_fn,
    initial_params,
    num_samples=2000,
    step_size=0.001
)


for i, arr in enumerate(states.position):
    print(f"Param {i}: shape {arr.shape}")
    print(f"Unique values: {jnp.unique(arr)}")
print("Mean acceptance rate:", jnp.mean(infos.acceptance_rate))
print("Sampled parameter shapes:")
print(jax.tree_util.tree_map(lambda x: x.shape, states.position))
logger.log_message("Done Jax Modelling")

运行输出

Param 0: shape ()
Unique values: [3.608985e-09]
Param 1: shape (29,)
Unique values: [1.]
Param 2: shape (13,)
Unique values: [0.1]
Param 3: shape (29,)
Unique values: [0.84729785]
Param 4: shape (29,)
Unique values: [-0.84729785]
Param 5: shape (0,)
Unique values: []
Param 6: shape ()
Unique values: [0.]
Mean acceptance rate: 0.0
Sampled parameter shapes:
((), (29,), (13,), (29,), (29,), (0,), ())

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.12 17:32:03