Mice — PyMC cross-check (Weibull, 4 groups)¶
Companion to mice_python.ipynb. Same censored Weibull as survival_pymc.ipynb, with a 4-level group factor; NUTS on the pm.Potential censored log-likelihood.
In [1]:
import numpy as np, pandas as pd, pymc as pm, pytensor.tensor as pt, arviz as az
from survival_mcmc import weibull_rwm
d = pd.read_csv('mice.csv'); G = d['group'].to_numpy()
t = d['time'].to_numpy(float); delta = d['cens'].to_numpy(float); logt = np.log(t)
X = np.column_stack([np.ones(len(d))] + [(G==g).astype(float) for g in (2,3,4)])
with pm.Model() as m:
beta = pm.Normal('beta', 0, 10, shape=4); logk = pm.Normal('logk', 0, 1); k = pm.math.exp(logk)
xb = pt.dot(X, beta)
pm.Potential('lik', (delta*(xb+logk+(k-1)*logt) - pt.exp(xb)*t**k).sum())
pm.Deterministic('k', k)
idata = pm.sample(2000, tune=2000, chains=4, target_accept=0.95, random_seed=1, progressbar=False)
po = idata.posterior; mxr = float(az.summary(idata, var_names=['beta','logk'])['r_hat'].max())
print('shape k = %.3f max r_hat %.3f' % (float(po['k'].mean()), mxr))
for j,g in [(1,2),(2,3),(3,4)]:
hr = np.exp(po['beta'].values[:,:,j].ravel()); print(' group %d HR vs g1 = %.2f [%.2f, %.2f]' % (g, hr.mean(), np.percentile(hr,2.5), np.percentile(hr,97.5)))
g = weibull_rwm(X, t, delta, seed=1)['draws']
print('\n%-14s %12s %12s' % ('param','from-scratch','PyMC'))
print('%-14s %12.3f %12.3f' % ('shape k', np.exp(g[:,-1]).mean(), float(po['k'].mean())))
for j,gg in [(1,2),(2,3),(3,4)]:
print('%-14s %12.3f %12.3f' % ('HR g%d'%gg, np.exp(g[:,j]).mean(), np.exp(po['beta'].values[:,:,j]).mean()))
Initializing NUTS using jitter+adapt_diag...
Multiprocess sampling (4 chains in 4 jobs)
NUTS: [beta, logk]
Sampling 4 chains for 2_000 tune and 2_000 draw iterations (8_000 + 8_000 draws total) took 8 seconds.
shape k = 3.178 max r_hat 1.000 group 2 HR vs g1 = 0.33 [0.15, 0.63] group 3 HR vs g1 = 0.74 [0.35, 1.34] group 4 HR vs g1 = 1.54 [0.73, 2.88] param from-scratch PyMC shape k 3.179 3.178 HR g2 0.342 0.328 HR g3 0.744 0.739 HR g4 1.573 1.535
Results¶
PyMC reproduces the from-scratch fit: shape k ≈ 3.2, group-2 HR ≈ 0.3 (longest-lived), group-4 HR ≈ 1.5, r̂ ≈ 1.00 — same Weibull, sampled by NUTS.