cfREF / src /cfref /validation.py
qiongli0705's picture
Upload 10 files
703e278 verified
Raw History Blame Contribute Delete
14.1 kB
"""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