# DASS Analysis Clinic Case 004; 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))

d=np.genfromtxt(HERE/'case-004-shared-data.csv',delimiter=',',names=True)
def compare(a,b):
 effect=b.mean()-a.mean(); df=len(a)+len(b)-2
 pooled=((len(a)-1)*a.var(ddof=1)+(len(b)-1)*b.var(ddof=1))/df
 se=np.sqrt(pooled*(1/len(a)+1/len(b)))
 return effect,se,effect+np.array([-1,1])*stats.t.ppf(.975,df)*se,2*stats.t.sf(abs(effect/se),df)
naive=compare(d['y'][d['arm']==0],d['y'][d['arm']==1])
clusters=np.unique(d['cluster']); means=np.array([d['y'][d['cluster']==k].mean() for k in clusters]); arms=np.array([d['arm'][d['cluster']==k][0] for k in clusters])
valid=compare(means[arms==0],means[arms==1])
print('Independent participants:',naive); print('Cluster means:',valid)
assert np.isclose(naive[0],valid[0]) and valid[1]>naive[1] and np.isclose(valid[1],.436,atol=.0005)
