Download src/cfref/validation.py from qiongli0705/cfREF: direct link, hf CLI and curl.
- Browser
- Download file 14.1 kB
-
https://huggingface.co/qiongli0705/cfREF/resolve/main/src/cfref/validation.py
- Command line
-
hf download hf://qiongli0705/cfREF/src/cfref/validation.py
-
curl -L -o validation.py https://huggingface.co/qiongli0705/cfREF/resolve/main/src/cfref/validation.py
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 | |