# DASS Analysis Clinic Case 008; Python dependencies: numpy, scipy.
from pathlib import Path
import numpy as np
from scipy import stats, optimize
HERE = Path(__file__).resolve().parent

def ols(X, y):
    b = np.linalg.lstsq(X, y, rcond=None)[0]
    residual = y - X @ b
    df = len(y) - X.shape[1]
    cov = (residual @ residual / df) * np.linalg.inv(X.T @ X)
    se = np.sqrt(np.diag(cov))
    ci = np.column_stack((b-stats.t.ppf(.975, df)*se,b+stats.t.ppf(.975,df)*se))
    return b, cov, ci

def power(n, d):
    # Matches R power.t.test(strict=FALSE): rejection tail in effect direction.
    return stats.nct.sf(stats.t.ppf(.975,2*n-2),2*n-2,d*np.sqrt(n/2))

# Bayesian normal-regression MI for one incomplete continuous outcome.
# Fully observed arm/baseline, flat coefficient prior and p(sigma^2) proportional
# to 1/sigma^2. Not a general MICE or multilevel imputation implementation.
d=np.genfromtxt(HERE/'case-008-shared-data.csv',delimiter=',',names=True)
arm=d['arm']; baseline=d['baseline']; y=d['followup']; missing=np.isnan(y)
X=np.column_stack((np.ones(len(y)),arm,baseline)); observed=~missing
Xo=X[observed]; yo=y[observed]; bhat=np.linalg.lstsq(Xo,yo,rcond=None)[0]
df=len(yo)-X.shape[1]; sse=np.sum((yo-Xo@bhat)**2); inv=np.linalg.inv(Xo.T@Xo)
rng=np.random.default_rng(8009); completed=[]
for _ in range(30):
 variance=sse/rng.chisquare(df)
 beta=rng.multivariate_normal(bhat,variance*inv)
 out=y.copy(); out[missing]=X[missing]@beta+rng.normal(0,np.sqrt(variance),missing.sum())
 completed.append(out)
def pooled(delta):
 estimates=[]; variances=[]
 for out in completed:
  shifted=out.copy(); shifted[missing & (arm==1)]+=delta
  b,cov,_=ols(X,shifted); estimates.append(b[1]);variances.append(cov[1,1])
 q=np.mean(estimates); se=np.sqrt(np.mean(variances)+(1+1/30)*np.var(estimates,ddof=1))
 return q,q-stats.norm.ppf(.975)*se,q+stats.norm.ppf(.975)*se
print('Observed:',observed.sum(),'missing by arm:',[np.sum(missing & (arm==g)) for g in [0,1]])
print('Complete-case adjusted effect:',bhat[1])
for delta in [0,-.5,-1]:print('delta / pooled effect / approximate CI:',delta,pooled(delta))
print('Python posterior draws differ from mice; the page table reports the R run.')
assert observed.sum()==321 and np.isclose(bhat[1],.474,atol=.0005)
assert pooled(0)[0]>pooled(-.5)[0]>pooled(-1)[0]
