基于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
相关产品推荐
相关产品推荐

