File size: 14,081 Bytes
703e278 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 | """Nested cohort validation with training-fold masks and paired query subjects.
Starts from counts supplied by the user. It cannot undo prior global depth
processing. Draw percentiles quantify reference-set variation, not population CIs.
"""
import hashlib
import json
from pathlib import Path
import numpy as np
import pandas as pd
from sklearn.linear_model import LogisticRegression
from sklearn.metrics import roc_auc_score
from sklearn.preprocessing import StandardScaler
from .estimator import CFREF
def training_mask(short, long, train, min_median_total=100):
"""Learn feature eligibility from training rows only."""
mask = np.median((short + long)[train], axis=0) >= min_median_total
if not mask.any():
raise ValueError('Training fold retained no features.')
return mask
def _metrics(scores, y, threshold):
return dict(auc=float(roc_auc_score(y, scores)),
specificity=float(np.mean(scores[y == 0] <= threshold)),
sensitivity=float(np.mean(scores[y == 1] > threshold)))
def _reference_draws(y, n_ref, draws, seed, min_query_controls):
ctrl = np.flatnonzero(y == 0)
if len(ctrl) < n_ref + min_query_controls:
return
rng = np.random.RandomState(seed)
for draw in range(draws):
ref = np.sort(rng.choice(ctrl, n_ref, replace=False))
query = np.setdiff1d(np.arange(len(y)), ref)
yield draw, ref, query
def _cfref_summary(model, X, y, names, n_ref, draws, seed, min_query_controls):
spec = []
for _, ref, query in _reference_draws(y, n_ref, draws, seed, min_query_controls):
scores = model.decision_function(X[query], X[ref], feature_names=names)
spec.append(_metrics(scores, y[query], model.threshold_)['specificity'])
if not spec:
raise ValueError('Insufficient inner target controls for reference draws.')
return float(np.median(spec))
def nested_lopo(short_counts, long_counts, y, cohorts, sample_ids, *,
feature_names=None, outer_cohorts=None, seeds=(0, 1, 2),
capacities=((64, 24), (128, 48), (256, 96)), episodes=3000,
reference_sizes=(5, 10, 20, 30), draws=60,
threshold_draws=40, inner_references=5, outer_threshold_references=20,
target_specificity=0.95, min_median_total=100,
min_query_controls=3, output_dir=None, progress=None,
input_preprocessing='Unspecified upstream count processing'):
"""Return paired per-draw results and nested-selection audit tables.
Parameters use the archived reference/threshold conventions by default.
Sample IDs must be globally unique; all rows of one cohort are held out
together. Each inner fold independently learns its count-coverage mask and
scaler. Target labels are used ONLY for reference identification/evaluation,
never for training, threshold construction or outer capacity selection.
Methods: nested cfREF, source-locked logistic regression, locally recalibrated
logistic regression. All use identical target query subjects for each draw.
Outputs include draw-level paired differences; these are not independent
patient replicates. No capacity is selected using outer performance.
"""
S = np.asarray(short_counts, dtype=float); L = np.asarray(long_counts, dtype=float)
y = np.asarray(y); coh = np.asarray(cohorts); ids = np.asarray(sample_ids).astype(str)
if S.ndim != 2 or S.shape != L.shape or min(S.shape) == 0:
raise ValueError('Counts must be matching nonempty samples × bins arrays.')
if not np.isfinite(S).all() or not np.isfinite(L).all() or (S < 0).any() or (L < 0).any():
raise ValueError('Counts must be finite and nonnegative.')
if y.shape != (len(S),) or coh.shape != y.shape or ids.shape != y.shape:
raise ValueError('Labels, cohorts and IDs must match sample count.')
if not np.isin(y, [0, 1]).all() or len(np.unique(y)) != 2:
raise ValueError('Binary labels with both classes are required.')
if len(set(ids)) != len(ids) or any(not i.strip() for i in ids):
raise ValueError('Subject IDs must be globally unique nonempty strings.')
if not all(isinstance(c, (str, np.str_)) and c.strip() for c in coh):
raise ValueError('Cohort IDs must be nonempty strings.')
if len(np.unique(coh)) < 3:
raise ValueError('Nested cohort validation requires at least three cohorts.')
if not capacities or len(set(tuple(c) for c in capacities)) != len(capacities):
raise ValueError('Supply unique candidate capacities.')
if not seeds or len(set(seeds)) != len(seeds) or not all(isinstance(s,int) and 0 <= s < 2**32 for s in seeds):
raise ValueError('Supply unique nonnegative integer seeds.')
for v in [draws, threshold_draws, episodes, min_query_controls]:
if not isinstance(v, int) or v < 1:raise ValueError('Draws, episodes and minimum controls must be positive integers.')
if not reference_sizes or len(set(reference_sizes)) != len(reference_sizes) or not all(isinstance(v,int) and v>=2 for v in reference_sizes):
raise ValueError('Reference sizes must be unique integers >=2.')
if not np.isfinite(min_median_total) or min_median_total < 0:raise ValueError('Invalid coverage threshold.')
names = [f'bin_{j}' for j in range(S.shape[1])] if feature_names is None else list(feature_names)
if len(names) != S.shape[1] or len(set(names)) != len(names) or not all(isinstance(n,str) and n for n in names):
raise ValueError('Feature names must be unique ordered strings matching bins.')
outer = list(np.unique(coh)) if outer_cohorts is None else list(outer_cohorts)
if len(set(outer)) != len(outer) or any(c not in coh for c in outer):raise ValueError('Invalid outer cohorts.')
ratio = S / np.maximum(S + L, 1)
result_rows=[]; inner_rows=[]; picks=[]; assignments=[]; masks=[]; skips=[]
cfg=dict(seeds=list(seeds),capacities=list(capacities),episodes=episodes,
reference_sizes=list(reference_sizes),draws=draws,threshold_draws=threshold_draws,
inner_references=inner_references,outer_threshold_references=outer_threshold_references,
target_specificity=target_specificity,min_median_total=min_median_total,
min_query_controls=min_query_controls,input_preprocessing=input_preprocessing,
quantile='linear; in-sample source controls',mask_policy='training rows of each inner/outer fold',
selection='mean over inner cohorts of abs(median draw specificity - target); capacity tuple order breaks ties',
patient_independent_confidence_intervals=False)
def log(msg):
if progress is not None:progress(msg)
def save_partial():
if output_dir is not None:
p=Path(output_dir);p.mkdir(parents=True,exist_ok=True)
for name, rows in [('draw_metrics',result_rows),('inner_selection',inner_rows),('capacity_picks',picks),('skipped',skips)]:
pd.DataFrame(rows).to_csv(p/(name+'.csv'),index=False)
(p/'fold_masks.json').write_text(json.dumps(masks,indent=2))
(p/'protocol.json').write_text(json.dumps(cfg,indent=2))
for ho in outer:
tr=coh!=ho;te=~tr; train_cohorts=np.unique(coh[tr]); yy=y[te]
if (yy==0).sum()<min(reference_sizes)+min_query_controls or (yy==1).sum()<1:
skips.append(dict(held_out=ho,reason='insufficient target subjects'));continue
outer_mask=training_mask(S,L,tr,min_median_total); outer_names=np.asarray(names)[outer_mask].tolist()
masks.append(dict(outer=ho,inner=None,kept_indices=np.flatnonzero(outer_mask).tolist()))
for seed in seeds:
candidates=[]
for dh,de in capacities:
devs=[]
for inner_ho in train_cohorts:
itr=tr&(coh!=inner_ho);ite=tr&(coh==inner_ho)
if (y[ite]==0).sum()<inner_references+min_query_controls or (y[ite]==1).sum()<1:
continue
mask=training_mask(S,L,itr,min_median_total); fn=np.asarray(names)[mask].tolist()
if seed==seeds[0] and (dh,de)==tuple(capacities[0]):
masks.append(dict(outer=ho,inner=inner_ho,kept_indices=np.flatnonzero(mask).tolist()))
model=CFREF(hidden_dim=dh,embedding_dim=de,episodes=episodes,seed=seed,
threshold_references=inner_references,threshold_draws=threshold_draws,
target_specificity=target_specificity)
model.fit(ratio[itr][:,mask],y[itr],coh[itr],feature_names=fn)
spec=_cfref_summary(model,ratio[ite][:,mask],y[ite],fn,inner_references,draws,seed,min_query_controls)
dev=abs(spec-target_specificity);devs.append(dev)
inner_rows.append(dict(held_out=ho,seed=seed,hidden_dim=dh,embedding_dim=de,
inner_held_out=inner_ho,specificity_median=spec,deviation=dev,
n_features=int(mask.sum()),threshold=model.threshold_))
if not devs:raise ValueError('No evaluable inner cohorts; adjust design before using this dataset.')
candidates.append((float(np.mean(devs)),dh,de))
log(f'{ho} seed {seed}: capacity {dh}/{de}, inner deviation {candidates[-1][0]:.4f}')
best=min(enumerate(candidates),key=lambda item:(item[1][0],item[0]))[1]
_,dh,de=best
Xtr=ratio[tr][:,outer_mask];Xt=ratio[te][:,outer_mask]
model=CFREF(hidden_dim=dh,embedding_dim=de,episodes=episodes,seed=seed,
threshold_references=outer_threshold_references,threshold_draws=threshold_draws,
target_specificity=target_specificity).fit(Xtr,y[tr],coh[tr],feature_names=outer_names)
scaler=StandardScaler().fit(Xtr)
lr=LogisticRegression(C=1,max_iter=5000,random_state=seed).fit(scaler.transform(Xtr),y[tr])
train_scores=lr.decision_function(scaler.transform(Xtr));baseline_scores=lr.decision_function(scaler.transform(Xt))
lr_threshold=float(np.quantile(train_scores[y[tr]==0],target_specificity,method='linear'))
picks.append(dict(held_out=ho,seed=seed,hidden_dim=dh,embedding_dim=de,
mean_inner_deviation=best[0],n_features=int(outer_mask.sum()),
cfref_threshold=model.threshold_,lr_threshold=lr_threshold,
threshold_contributors=json.dumps(model.metadata_['threshold_contributors'])))
target_ids=ids[te]
for k in reference_sizes:
if (yy==0).sum()<k+min_query_controls:
skips.append(dict(held_out=ho,seed=seed,n_ref=k,reason='insufficient remaining target controls'));continue
for draw,ref,query in _reference_draws(yy,k,draws,seed,min_query_controls):
if np.intersect1d(target_ids[ref],target_ids[query]).size:raise AssertionError('Reference/query overlap')
context=dict(held_out=ho,seed=seed,n_ref=k,draw=draw)
signature=hashlib.sha256('\n'.join(target_ids[query]).encode()).hexdigest()
assignments.append(dict(**context,reference_ids=json.dumps(target_ids[ref].tolist()),
query_ids=json.dumps(target_ids[query].tolist()),query_sha256=signature))
cs=model.decision_function(Xt[query],Xt[ref],feature_names=outer_names)
recal=float(np.quantile(baseline_scores[ref],target_specificity,method='linear'))
for method,score,thr in [('cfref',cs,model.threshold_),('logreg_locked',baseline_scores[query],lr_threshold),('logreg_recalibrated',baseline_scores[query],recal)]:
result_rows.append(dict(**context,method=method,threshold=thr,
n_query_controls=int((yy[query]==0).sum()),n_query_cases=int((yy[query]==1).sum()),
query_sha256=signature,**_metrics(score,yy[query],thr)))
save_partial();log(f'Completed {ho} seed {seed}: selected {dh}/{de}')
metrics=pd.DataFrame(result_rows)
if metrics.empty:raise ValueError('No evaluable outer settings.')
keys=['held_out','seed','n_ref','draw']
# Pair before summarizing: differences of individual draw metrics.
cf=metrics[metrics.method=='cfref'].set_index(keys);diffs=[]
for comparator in ['logreg_locked','logreg_recalibrated']:
base=metrics[metrics.method==comparator].set_index(keys)
assert cf.index.equals(base.index)
assert (cf.query_sha256==base.query_sha256).all()
d=(cf[['auc','specificity','sensitivity']]-base[['auc','specificity','sensitivity']]).reset_index()
d['comparator']=comparator;diffs.append(d)
paired=pd.concat(diffs,ignore_index=True)
def summarize(frame, groups):
rows=[]
for key, tab in frame.groupby(groups,sort=False):
row=dict(zip(groups,key if isinstance(key,tuple) else (key,)))
for col in ['auc','specificity','sensitivity']:
for suffix,q in [('median',.5),('p10',.1),('p90',.9)]:row[col+'_'+suffix]=float(tab[col].quantile(q))
row['draws']=len(tab);rows.append(row)
return pd.DataFrame(rows)
summary=summarize(metrics,['held_out','seed','n_ref','method'])
paired_summary=summarize(paired,['held_out','seed','n_ref','comparator'])
output=dict(draw_metrics=metrics,summary=summary,paired_differences=paired,
paired_summary=paired_summary,inner_selection=pd.DataFrame(inner_rows),
capacity_picks=pd.DataFrame(picks),assignments=pd.DataFrame(assignments),skipped=pd.DataFrame(skips),
fold_masks=masks,protocol=cfg)
if output_dir is not None:
save_partial();p=Path(output_dir)
for name in ['summary','paired_differences','paired_summary']:output[name].to_csv(p/(name+'.csv'),index=False)
output['assignments'].to_csv(p/'assignments.csv.gz',index=False,compression='gzip')
return output
|