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.