File size: 23,484 Bytes
bd99b1b 7c2a2df bd99b1b 7c2a2df bd99b1b 7c2a2df bd99b1b 7c2a2df bd99b1b 7c2a2df bd99b1b 7c2a2df bd99b1b 7c2a2df bd99b1b 7c2a2df bd99b1b 7c2a2df bd99b1b | 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 214 215 216 217 218 219 220 221 222 223 224 225 226 227 228 229 230 231 232 233 234 235 236 237 238 239 240 241 242 243 244 245 246 247 248 249 250 251 252 253 254 255 256 257 258 259 260 261 262 263 264 265 266 267 268 269 270 271 272 273 274 275 276 277 278 279 280 281 282 283 284 285 286 287 288 289 290 291 292 293 294 295 296 297 298 299 300 301 302 303 304 305 306 307 308 309 310 311 312 313 314 315 316 317 318 319 320 321 322 323 324 325 326 327 328 329 330 331 332 333 334 335 336 337 338 339 340 341 342 343 344 345 346 347 348 349 350 351 352 353 354 355 356 357 358 359 360 361 362 363 364 365 366 367 368 369 370 371 372 373 374 375 376 377 378 379 380 381 382 383 384 385 386 387 388 389 390 391 392 393 394 395 396 397 398 399 400 401 402 403 404 405 406 407 408 409 410 411 412 413 414 415 416 417 418 419 420 421 422 423 424 425 426 427 428 429 430 431 432 433 434 435 436 437 438 439 440 441 442 443 444 445 | """
Cross-protocol evaluation: a model trained on one protocol applied to the other.
The deployment case this answers is a model developed on archived full protocol (FDP) examinations
and then applied to abbreviated (AP) acquisitions, where Sub_2, Sub_3, Sub_4 and T2 simply do not
exist. The reverse direction is not symmetric:
abbreviated = [Dyn_0, Sub_1]
full = [Dyn_0, Sub_1, Sub_2, Sub_3, Sub_4, T2]
AP is a strict subset of FDP, so an AP model applied to a full acquisition just reads the two
sequences it was trained on and nothing has to be substituted -- that direction is reported for
completeness but is a no-op by construction. An FDP model applied to an abbreviated acquisition is
the real question, because four of its six input channels are missing and something has to stand in
for them. Which substitute is used is a modelling decision, not a detail, so it is explicit:
zero the missing channels are set to 0, i.e. the mean of the z-normalised image
repeat Sub_2/3/4 are filled with Sub_1 and T2 with Dyn_0, a naive carry forward
mean every missing channel is filled with the voxelwise mean of the two available ones
Two steps, because the first needs a GPU and the second does not:
# once per (architecture, fold), writes cross_protocol_<arch>_fold<k>[_frac=<f>].csv
python scripts/validation/cross_protocol_evaluation.py predict resnet18 <data> <metadata> \\
<splits> -c 0 -m '<checkpoint pattern>'
# once, over all folds
python scripts/validation/cross_protocol_evaluation.py analyse <score_dir> <metadata>
The scores of every setting are produced in one pass over the same loaded volumes, so the
comparison is paired case by case and free of crop or augmentation differences between settings.
Note that AP-native scores here are recomputed from the full protocol tensor rather than taken from
predict_oof.py, so they can differ marginally through the random noise in the inference transform.
"""
import argparse
import csv
import re
from collections import defaultdict
from pathlib import Path
import numpy as np
import pandas as pd
import torch
from auto_detect_breast_mri.config import resolve_path
from auto_detect_breast_mri.data import loaders
from auto_detect_breast_mri.data.breast_mri_dataset import protocol_mappings, split_patient_key
from auto_detect_breast_mri.data.metadata import get_uka_metatensor
from auto_detect_breast_mri.data.transforms import default_transform
from auto_detect_breast_mri.evaluation.cluster_bootstrap import (FoldClusterBootstrap,
percentile_interval,
weighted_fold_difference)
from auto_detect_breast_mri.evaluation.oof_table import load_patient_map
from auto_detect_breast_mri.evaluation.subgroups import bootstrap_metrics, weighted_fold_metric
from auto_detect_breast_mri.models.resnets import model_names
CPU, GPU = "cpu", "cuda"
FULL_SEQUENCES = protocol_mappings['full']
ABRV_SEQUENCES = protocol_mappings['abbreviated']
# index of every abbreviated sequence inside the full channel stack, and of the ones it lacks
KEPT = [FULL_SEQUENCES.index(name) for name in ABRV_SEQUENCES]
MISSING = [index for index in range(len(FULL_SEQUENCES)) if index not in KEPT]
SUBSTITUTES = ('zero', 'repeat', 'mean')
FILE_PATTERN = re.compile(r'^cross_protocol_(?P<architecture>resnet\d+)_fold(?P<fold>\d+)'
r'(?:_frac=(?P<fraction>[0-9.]+))?\.csv$')
# (model, the protocol actually available at inference): whether it is the native or the cross case
SETTINGS = [('full', 'FDP', 'native'), ('full', 'AP', 'cross'),
('abrv', 'AP', 'native'), ('abrv', 'FDP', 'native-by-subset')]
def load_checkpoint(model, model_path, device):
"""
Load a checkpoint that is either a bare state dict or wrapped in a 'state_dict' entry.
Same loader as predict_oof.py: models.checkpoints.load_pretrained_model returns a state dict
rather than a model and only handles the wrapped form, so it cannot be used here.
"""
checkpoint = torch.load(model_path, map_location=device, weights_only=False)
state_dict = checkpoint['state_dict'] if 'state_dict' in checkpoint else checkpoint
state_dict = {key.replace('module.', ''): value for key, value in state_dict.items()}
missing, unexpected = model.load_state_dict(state_dict, strict=False)
if missing or unexpected:
raise RuntimeError(f"Checkpoint {model_path} does not match the architecture. "
f"Missing keys: {sorted(missing)[:5]}, unexpected: {sorted(unexpected)[:5]}")
return model
def abbreviate(volume, substitute):
"""
Turn a full protocol tensor into what an abbreviated acquisition would have delivered.
:param volume: (B, C, ...) with C = len(FULL_SEQUENCES)
:param substitute: what stands in for the sequences an abbreviated exam does not contain
:return: a copy with the missing channels replaced
"""
reduced = volume.clone()
if substitute == 'zero':
reduced[:, MISSING] = 0.0
elif substitute == 'repeat':
# Sub_2/3/4 carry Sub_1 forward, T2 falls back to the native scan
source = {FULL_SEQUENCES.index('Sub_2'): FULL_SEQUENCES.index('Sub_1'),
FULL_SEQUENCES.index('Sub_3'): FULL_SEQUENCES.index('Sub_1'),
FULL_SEQUENCES.index('Sub_4'): FULL_SEQUENCES.index('Sub_1'),
FULL_SEQUENCES.index('T2'): FULL_SEQUENCES.index('Dyn_0')}
for target, origin in source.items():
reduced[:, target] = volume[:, origin]
elif substitute == 'mean':
reduced[:, MISSING] = volume[:, KEPT].mean(dim=1, keepdim=True)
else:
raise ValueError(f"Unknown substitute '{substitute}'. Use one of {SUBSTITUTES}.")
return reduced
@torch.no_grad()
def predict_settings(models, data_loader, device, substitute):
"""
Score every case under every setting, from one pass over the loaded volumes.
:param models: dict protocol suffix ('abrv'/'full') -> model, already on `device`
:return: list of row dicts
"""
for model in models.values():
model.eval()
rows = []
for batch_id, subject in enumerate(data_loader):
volume = subject['image']['data'].to(device)
labels = subject['label'].cpu().tolist()
keys = subject['path']
scored = {}
with torch.autocast(device_type=device, dtype=torch.float16):
if 'full' in models:
scored[('full', 'FDP')] = models['full'](volume)[:, 0]
scored[('full', 'AP')] = models['full'](abbreviate(volume, substitute))[:, 0]
if 'abrv' in models:
# the abbreviated model only ever sees the sequences it was trained on, so the
# full acquisition gives it exactly the same tensor as an abbreviated one
available = models['abrv'](volume[:, KEPT])[:, 0]
scored[('abrv', 'AP')] = available
scored[('abrv', 'FDP')] = available
for (suffix, available_protocol), logits in scored.items():
role = next(role for model, protocol, role in SETTINGS
if model == suffix and protocol == available_protocol)
for key, label, score in zip(keys, labels, logits.float().cpu().tolist()):
examination_id, side = split_patient_key(key)
rows.append({'examination_id': examination_id, 'side': side, 'label': int(label),
'model_protocol': suffix, 'input_protocol': available_protocol,
'role': role, 'score': float(score)})
if batch_id % 20 == 0:
print(f" batch {batch_id}/{len(data_loader)}", flush=True)
return rows
def run_prediction(args):
"""Score one fold under every setting and write the per case CSV."""
device = GPU if torch.cuda.is_available() else CPU
print(f"Use device: {device}")
features = get_uka_metatensor(0, args.feature_path)
fraction = float(args.fraction) if args.fraction else 1.0
frac_suffix = f'_frac={fraction}'
output_dir = Path(resolve_path(args.output_path, "output_root", "output folder")) / "cross_protocol"
output_dir.mkdir(parents=True, exist_ok=True)
output_file = output_dir / f"cross_protocol_{args.architecture}_fold{args.fold}{frac_suffix}.csv"
models = {}
for suffix in ('abrv', 'full'):
model_key = f"{args.architecture}_{suffix}"
model = model_names.get(model_key)
if model is None:
raise ValueError(f"Unknown model key {model_key}. Available: {sorted(model_names)}")
checkpoint = args.model_path_pattern.format(model_key=model_key, fold=args.fold,
frac_suffix=frac_suffix)
print(f"Load {checkpoint}")
models[suffix] = load_checkpoint(model, checkpoint, device).to(device)
# always load the full protocol: every setting is derived from that one tensor
loader = loaders.get_evaluation_dataloader(
args.data_path, features, tuple(args.image_shape), 'full',
data_selection_file_sceleton=str(Path(args.split_files_folder) / f"fold{args.fold}"
/ "stratified_test_set"),
batch_size=args.batch_size, fold=args.fold,
transform=default_transform(tuple(args.image_shape)))
print(f"fold {args.fold}: {len(loader.dataset)} cases, substitute '{args.substitute}'")
rows = predict_settings(models, loader, device, args.substitute)
for row in rows:
row['outer_fold'] = args.fold
row['substitute'] = args.substitute
fields = ['outer_fold', 'examination_id', 'side', 'label', 'model_protocol', 'input_protocol',
'role', 'substitute', 'score']
with open(output_file, 'w', newline='') as handle:
writer = csv.DictWriter(handle, fieldnames=fields)
writer.writeheader()
writer.writerows(rows)
print(f"Wrote {len(rows)} rows to {output_file}")
def read_scores(score_dirs, metadata_file, architecture=None, fraction=None):
"""
The per case scores of every fold, with the woman each examination belongs to attached.
:param score_dirs: one folder, or several when the architectures were predicted into separate
output roots; all of them are read into one table
:return: (DataFrame, dict describing what was read)
"""
if isinstance(score_dirs, (str, Path)):
score_dirs = [score_dirs]
paths = []
for score_dir in score_dirs:
score_dir = Path(score_dir)
if not score_dir.is_dir():
raise FileNotFoundError(f"{score_dir} is not a directory.")
paths.extend(sorted(score_dir.iterdir()))
frames, report = [], {}
for path in paths:
parsed = FILE_PATTERN.match(path.name)
if parsed is None:
continue
if architecture and parsed['architecture'] != architecture:
continue
found = float(parsed['fraction']) if parsed['fraction'] else 1.0
if fraction is not None and found != float(fraction):
continue
frame = pd.read_csv(path, dtype={'examination_id': str, 'side': str})
frame['architecture'] = parsed['architecture']
frame['fraction'] = found
frames.append(frame)
if not frames:
raise FileNotFoundError(f"No cross_protocol_*.csv in "
f"{', '.join(str(d) for d in score_dirs)} for the requested "
f"architecture/fraction. Run the 'predict' step first.")
table = pd.concat(frames, ignore_index=True)
for column in ('architecture', 'fraction', 'substitute'):
report[column] = sorted(table[column].unique().tolist())
if len(report['fraction']) > 1:
raise ValueError(f"The given folders mix training fractions {report['fraction']}; those "
f"are different models. Pass --fraction to pick one.")
if len(report['substitute']) > 1:
raise ValueError(f"The given folders mix substitutes {report['substitute']}; keep one "
f"substitute per analysis run.")
patient_map = load_patient_map(metadata_file, table['examination_id'].unique())
table['patient_id'] = table['examination_id'].map(patient_map)
report['dropped_without_patient_id'] = int(table['patient_id'].isna().sum())
table = table[table['patient_id'].notna()].reset_index(drop=True)
report['breasts'] = int(len(table) / table.groupby(['model_protocol', 'input_protocol']).ngroups)
return table, report
def paired_difference(wide, column_a, column_b, replications, weighting, seed, stratify):
"""
Interval for the paired difference column_a - column_b, both scored on the same women.
:param wide: one row per breast, holding both score columns
:return: dict with the observed difference and its interval
"""
bootstrap = FoldClusterBootstrap(wide['outer_fold'].to_numpy(), wide['patient_id'].to_numpy(),
wide['label'].to_numpy(), stratify_by_outcome=stratify)
weights = bootstrap.fold_weights(weighting)
labels = wide['label'].to_numpy()
scores_a, scores_b = wide[column_a].to_numpy(), wide[column_b].to_numpy()
observed, _ = weighted_fold_difference(bootstrap.observed_rows(), labels, scores_a, scores_b,
weights)
rng = np.random.default_rng(seed)
replicated = np.full(replications, np.nan)
for replication in range(replications):
drawn = bootstrap.resample_rows(rng)
replicated[replication], _ = weighted_fold_difference(drawn, labels, scores_a, scores_b,
weights)
low, high = percentile_interval(replicated, 0.95)
return {'difference': observed, 'ci_low': low, 'ci_high': high}
def analyse(table, target=0.9, replications=2000, weighting='pairs', seed=0, stratify=True):
"""
AUC of every setting with its interval, plus the paired cost of the missing sequences.
:return: (list of per setting rows, list of paired comparison rows)
"""
settings, paired = [], []
for architecture in sorted(table['architecture'].unique()):
part = table[table['architecture'] == architecture]
for model_protocol, input_protocol, role in SETTINGS:
rows = part[(part['model_protocol'] == model_protocol)
& (part['input_protocol'] == input_protocol)]
if rows.empty:
continue
print(f" {architecture} {model_protocol} model on {input_protocol} input "
f"({len(rows)} breasts)", flush=True)
statistics = bootstrap_metrics(rows['outer_fold'].to_numpy(),
rows['patient_id'].to_numpy(),
rows['label'].to_numpy(), rows['score'].to_numpy(),
target=target, replications=replications,
weighting=weighting, seed=seed, stratify=stratify)
settings.append({'architecture': architecture, 'model_protocol': model_protocol,
'input_protocol': input_protocol, 'role': role,
'breasts': len(rows), **{f'{metric}_{field}': value
for metric, values in statistics.items()
for field, value in values.items()}})
# the cost of losing the four sequences, paired case by case on the FDP model
wide = part.pivot_table(index=['outer_fold', 'examination_id', 'side', 'label',
'patient_id'],
columns=['model_protocol', 'input_protocol'],
values='score').reset_index()
wide.columns = ['_'.join(part for part in column if part).strip('_')
if isinstance(column, tuple) else column for column in wide.columns]
comparisons = [('full_AP', 'full_FDP',
'FDP model: abbreviated input minus full input'),
('full_AP', 'abrv_AP',
'abbreviated input: FDP model minus AP model')]
for column_a, column_b, description in comparisons:
if column_a not in wide.columns or column_b not in wide.columns:
continue
print(f" {architecture} paired: {description}", flush=True)
result = paired_difference(wide, column_a, column_b, replications, weighting, seed,
stratify)
paired.append({'architecture': architecture, 'comparison': description, **result})
return settings, paired
def format_report(settings, paired, report, target, replications, decimals=3):
lines = []
add = lines.append
add("=" * 100)
add("Cross-protocol evaluation: each model applied to the other protocol's input")
add("=" * 100)
add("")
add(f"{'architectures':<28}{', '.join(report['architecture'])}")
add(f"{'training fraction':<28}{', '.join(f'{value:g}' for value in report['fraction'])}")
add(f"{'substitute for missing seq.':<28}{', '.join(report['substitute'])}")
add(f"{'breasts per setting':<28}{report['breasts']}")
add(f"{'bootstrap replications':<28}{replications}")
if report.get('dropped_without_patient_id'):
add(f"{'dropped, no patient ID':<28}{report['dropped_without_patient_id']}")
add("")
add("-" * 100)
add("1. Every setting")
add("-" * 100)
add(f"{'architecture':<12} {'model':>6} {'input':>6} {'role':>18} {'AUC':>7} "
f"{'95% CI':>20} {f'sens@{target:.0%}spec':>15}")
for row in settings:
interval = f"[{row['auc_ci_low']:.{decimals}f}, {row['auc_ci_high']:.{decimals}f}]"
add(f"{row['architecture']:<12} {row['model_protocol']:>6} {row['input_protocol']:>6} "
f"{row['role']:>18} {row['auc_value']:>7.{decimals}f} {interval:>20} "
f"{row['sens_at_spec_value']:>15.{decimals}f}")
add("")
add("'native-by-subset' is not a separate experiment: the abbreviated model reads only Dyn_0")
add("and Sub_1, which a full acquisition also contains, so its score cannot change.")
add("")
add("-" * 100)
add("2. Paired differences, same women in every bootstrap replication")
add("-" * 100)
for row in paired:
add(f"{row['architecture']:<12} {row['comparison']:<48} "
f"{row['difference']:>+.{decimals}f} "
f"[{row['ci_low']:+.{decimals}f}, {row['ci_high']:+.{decimals}f}]")
add("")
add("-" * 100)
add("How to read this")
add("-" * 100)
add("The first paired row is the price of deploying an FDP trained model on abbreviated data,")
add("with the missing sequences filled in as stated above. The second asks whether, given only")
add("an abbreviated acquisition, one is better off with the FDP model plus substitution or with")
add("a model actually trained on AP. A negative value favours the second term of the pair.")
add("The substitute is a modelling choice: rerun with --substitute to see how much it matters.")
return "\n".join(lines)
def build_parser():
parser = argparse.ArgumentParser(prog="cross_protocol_evaluation", description=__doc__,
formatter_class=argparse.RawDescriptionHelpFormatter)
subparsers = parser.add_subparsers(dest="step", required=True)
predict = subparsers.add_parser("predict", help="score one fold under every setting (GPU)")
predict.add_argument("architecture", choices=["resnet18", "resnet50"])
predict.add_argument("data_path", help="Root folder(s) of the cropped MRIs, comma separated.")
predict.add_argument("feature_path", help="CSV/XLSX holding the label information.")
predict.add_argument("split_files_folder",
help="Folder holding fold<k>/stratified_test_set-f<k>.csv")
predict.add_argument("-c", "--fold", type=int, required=True)
predict.add_argument("-m", "--model_path_pattern", required=True,
help="Checkpoint path with {model_key}, {fold} and {frac_suffix}.")
predict.add_argument("-f", "--fraction", type=float, default=None,
help="Training fraction of the checkpoints. Default: 1.0")
predict.add_argument("-b", "--batch_size", type=int, default=16)
predict.add_argument("-o", "--output_path", default=None)
predict.add_argument("-s", "--image_shape", type=int, nargs=3, default=(256, 256, 32))
predict.add_argument("--substitute", choices=SUBSTITUTES, default='zero',
help="What stands in for the sequences an abbreviated exam lacks. "
"Default: zero")
analyse_parser = subparsers.add_parser("analyse", help="aggregate the per fold score files")
analyse_parser.add_argument("score_dirs", nargs='+', metavar="SCORE_DIR",
help="Folder(s) holding cross_protocol_*.csv. Give several when "
"the architectures were predicted into separate output roots.")
analyse_parser.add_argument("metadata_file", help="Metadata export, for the patient clustering.")
analyse_parser.add_argument("-o", "--output_path", default=None)
analyse_parser.add_argument("--architecture", default=None, choices=["resnet18", "resnet50"])
analyse_parser.add_argument("--fraction", type=float, default=None)
analyse_parser.add_argument("--operating_point", type=float, default=0.9)
analyse_parser.add_argument("-r", "--replications", type=int, default=2000)
analyse_parser.add_argument("-w", "--weighting", choices=["pairs", "cases", "equal"],
default="pairs")
analyse_parser.add_argument("--seed", type=int, default=0)
analyse_parser.add_argument("--no_stratify", action="store_true")
analyse_parser.add_argument("-d", "--decimals", type=int, default=3)
return parser
def main():
args = build_parser().parse_args()
if args.step == "predict":
run_prediction(args)
return
table, report = read_scores(args.score_dirs, args.metadata_file, args.architecture,
args.fraction)
print(f"Bootstrapping with {args.replications} replications ...")
settings, paired = analyse(table, target=args.operating_point,
replications=args.replications, weighting=args.weighting,
seed=args.seed, stratify=not args.no_stratify)
output_dir = Path(resolve_path(args.output_path, "output_root", "output folder")) / "cross_protocol"
output_dir.mkdir(parents=True, exist_ok=True)
pd.DataFrame(settings).to_csv(output_dir / "cross_protocol_settings.csv", index=False)
pd.DataFrame(paired).to_csv(output_dir / "cross_protocol_paired.csv", index=False)
text = format_report(settings, paired, report, args.operating_point, args.replications,
args.decimals)
(output_dir / "cross_protocol_report.txt").write_text(text)
print()
print(text)
print(f"\nWrote the report and two CSVs to {output_dir}")
if __name__ == '__main__':
main()
|