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