@@ -942,14 +942,15 @@ def evaluate_predictions(config: DictConfig, *, models: EvalModels | None = None
942942 if config .compute_feature_metrics and io_config .cell_segmentation_path is None :
943943 raise ValueError ("io.cell_segmentation_path is required when compute_feature_metrics=true" )
944944
945- if models is None :
946- models = load_eval_models (config )
947- seg_model = models .seg_model
948- dinov3_feature_extractor = models .dinov3
949- dynaclr_feature_extractor = models .dynaclr
950- celldino_feature_extractor = models .celldino
945+ with region_timer ("parent_load_models" , "<parent>" ):
946+ if models is None :
947+ models = load_eval_models (config )
948+ seg_model = models .seg_model
949+ dinov3_feature_extractor = models .dinov3
950+ dynaclr_feature_extractor = models .dynaclr
951+ celldino_feature_extractor = models .celldino
951952
952- cache_ctx , pred_cache_ctx = init_cache_contexts (config , models )
953+ cache_ctx , pred_cache_ctx = init_cache_contexts (config , models )
953954
954955 seg_path = Path (io_config .cell_segmentation_path ) if io_config .cell_segmentation_path is not None else None
955956
@@ -1106,16 +1107,17 @@ def evaluate_predictions(config: DictConfig, *, models: EvalModels | None = None
11061107 side_positions ["pred" ] = pred_positions
11071108 side_channels ["pred" ] = io_config .pred_channel_name
11081109 if sides_for_precompute :
1109- precompute_deep_features (
1110- sides_for_precompute ,
1111- side_positions ,
1112- side_channels ,
1113- seg_positions ,
1114- deep_extractors ,
1115- flush_threshold = flush_threshold ,
1116- )
1117- for ctx in sides_for_precompute .values ():
1118- flush_manifest (ctx )
1110+ with region_timer ("precompute_all" , "<parent>" ):
1111+ precompute_deep_features (
1112+ sides_for_precompute ,
1113+ side_positions ,
1114+ side_channels ,
1115+ seg_positions ,
1116+ deep_extractors ,
1117+ flush_threshold = flush_threshold ,
1118+ )
1119+ for ctx in sides_for_precompute .values ():
1120+ flush_manifest (ctx )
11191121
11201122 # Phase 2 runtime resolution: clamp fov_workers to n_positions
11211123 # now that we know it. threads_per_worker is frozen at Phase 1
@@ -1249,118 +1251,119 @@ def _aggregate(result: FovResult) -> None:
12491251 aux_stack .close ()
12501252
12511253 if config .compute_feature_metrics and all_feature_metrics :
1252- dataset_row : dict [str , float ] = {}
1253-
1254- # Stage per-prefix inputs: (pred_for_metric, target_for_metric,
1255- # pred_for_probe, target_for_probe, pred_fovs, target_fovs).
1256- # CP gets pruning + z-score; the pre-prune CP arrays feed the
1257- # linear probe so MADScaler can normalize per-fold.
1258- prefix_inputs : list [tuple [str , np .ndarray , np .ndarray , np .ndarray , np .ndarray , np .ndarray , np .ndarray ]] = []
1259-
1260- cp = parent_lists ["cp" ]
1261- if cp .pred_feats :
1262- pred_cp_raw = np .concatenate (cp .pred_feats , axis = 0 )
1263- target_cp_raw = np .concatenate (cp .gt_feats , axis = 0 )
1264- target_cp_filtered , pred_cp_filtered , cp_keep_mask = select_features (target_cp_raw , pred_cp_raw )
1265- mask_payload = {
1266- "keep_mask" : [bool (b ) for b in cp_keep_mask ],
1267- "n_kept" : int (cp_keep_mask .sum ()),
1268- "n_total" : int (cp_keep_mask .size ),
1269- "criteria" : {
1270- "freq_cut" : DEFAULT_FREQ_CUT ,
1271- "unique_cut" : DEFAULT_UNIQUE_CUT ,
1272- "corr_threshold" : DEFAULT_CORR_THRESHOLD ,
1273- },
1274- }
1275- (save_dir / "cp_selected_feature_mask.json" ).write_text (json .dumps (mask_payload , indent = 2 ))
1276- if pred_cp_filtered .size and target_cp_filtered .size :
1277- pred_cp_z , target_cp_z = _zscore_per_side (pred_cp_filtered , target_cp_filtered )
1278- else :
1279- pred_cp_z , target_cp_z = pred_cp_filtered , target_cp_filtered
1280- prefix_inputs .append (
1281- (
1282- "CP" ,
1283- pred_cp_z ,
1284- target_cp_z ,
1285- pred_cp_filtered ,
1286- target_cp_filtered ,
1287- np .concatenate (cp .pred_fovs , axis = 0 ),
1288- np .concatenate (cp .gt_fovs , axis = 0 ),
1289- )
1290- )
1291-
1292- deep_tracks = [("DINOv3" , "dinov3" ), ("DynaCLR" , "dynaclr" )]
1293- if celldino_feature_extractor is not None :
1294- deep_tracks .append (("CellDINO" , "celldino" ))
1295- for display_name , key in deep_tracks :
1296- bb = parent_lists [key ]
1297- if bb .pred_feats :
1298- pred_arr = np .concatenate (bb .pred_feats , axis = 0 )
1299- target_arr = np .concatenate (bb .gt_feats , axis = 0 )
1254+ with region_timer ("dataset_metrics" , "<parent>" ):
1255+ dataset_row : dict [str , float ] = {}
1256+
1257+ # Stage per-prefix inputs: (pred_for_metric, target_for_metric,
1258+ # pred_for_probe, target_for_probe, pred_fovs, target_fovs).
1259+ # CP gets pruning + z-score; the pre-prune CP arrays feed the
1260+ # linear probe so MADScaler can normalize per-fold.
1261+ prefix_inputs : list [tuple [str , np .ndarray , np .ndarray , np .ndarray , np .ndarray , np .ndarray , np .ndarray ]] = []
1262+
1263+ cp = parent_lists ["cp" ]
1264+ if cp .pred_feats :
1265+ pred_cp_raw = np .concatenate (cp .pred_feats , axis = 0 )
1266+ target_cp_raw = np .concatenate (cp .gt_feats , axis = 0 )
1267+ target_cp_filtered , pred_cp_filtered , cp_keep_mask = select_features (target_cp_raw , pred_cp_raw )
1268+ mask_payload = {
1269+ "keep_mask" : [bool (b ) for b in cp_keep_mask ],
1270+ "n_kept" : int (cp_keep_mask .sum ()),
1271+ "n_total" : int (cp_keep_mask .size ),
1272+ "criteria" : {
1273+ "freq_cut" : DEFAULT_FREQ_CUT ,
1274+ "unique_cut" : DEFAULT_UNIQUE_CUT ,
1275+ "corr_threshold" : DEFAULT_CORR_THRESHOLD ,
1276+ },
1277+ }
1278+ (save_dir / "cp_selected_feature_mask.json" ).write_text (json .dumps (mask_payload , indent = 2 ))
1279+ if pred_cp_filtered .size and target_cp_filtered .size :
1280+ pred_cp_z , target_cp_z = _zscore_per_side (pred_cp_filtered , target_cp_filtered )
1281+ else :
1282+ pred_cp_z , target_cp_z = pred_cp_filtered , target_cp_filtered
13001283 prefix_inputs .append (
13011284 (
1302- display_name ,
1303- pred_arr ,
1304- target_arr ,
1305- pred_arr ,
1306- target_arr ,
1307- np .concatenate (bb .pred_fovs , axis = 0 ),
1308- np .concatenate (bb .gt_fovs , axis = 0 ),
1285+ "CP" ,
1286+ pred_cp_z ,
1287+ target_cp_z ,
1288+ pred_cp_filtered ,
1289+ target_cp_filtered ,
1290+ np .concatenate (cp .pred_fovs , axis = 0 ),
1291+ np .concatenate (cp .gt_fovs , axis = 0 ),
13091292 )
13101293 )
13111294
1312- # Prefix with "Dataset_" so dataset-level FID/KID/cosine don't clobber
1313- # per-FOV columns of the same name when merged into per-FOV rows.
1314- def _compute_one (args ):
1315- # MIND stays on CPU here even when use_gpu=True: 4 parallel threads
1316- # racing on the same CUDA context would either serialize via the
1317- # allocator (no speedup) or contend for memory with mid-eval FOV
1318- # work in process executors. CPU MIND in a 4-thread BLAS-capped
1319- # pool is competitive with serialized GPU MIND and bit-stable
1320- # across runs that toggle use_gpu (torch's CPU vs CUDA RNG
1321- # produce different streams for the same seed, breaking cross-
1322- # leaf comparability of the MIND column).
1323- name , p_metric , t_metric , p_probe , t_probe , fov_p , fov_t = args
1324- raw = {
1325- ** compute_feature_similarity (p_metric , t_metric , name ),
1326- ** _real_vs_pred_probe (p_probe , t_probe , fov_p , fov_t , name ),
1327- }
1328- return {f"Dataset_{ k } " : v for k , v in raw .items ()}
1329-
1330- if prefix_inputs :
1331- # Threads suffice: torch-fidelity, sklearn LBFGS, and numpy BLAS
1332- # all release the GIL inside their hot loops. Cap inner BLAS to
1333- # 1 thread so the outer threads don't oversubscribe cores.
1334- with (
1335- threadpool_limits (limits = 1 ),
1336- ThreadPoolExecutor (max_workers = min (4 , len (prefix_inputs ))) as pool ,
1337- ):
1338- for result in pool .map (_compute_one , prefix_inputs ):
1339- dataset_row .update (result )
1295+ deep_tracks = [("DINOv3" , "dinov3" ), ("DynaCLR" , "dynaclr" )]
1296+ if celldino_feature_extractor is not None :
1297+ deep_tracks .append (("CellDINO" , "celldino" ))
1298+ for display_name , key in deep_tracks :
1299+ bb = parent_lists [key ]
1300+ if bb .pred_feats :
1301+ pred_arr = np .concatenate (bb .pred_feats , axis = 0 )
1302+ target_arr = np .concatenate (bb .gt_feats , axis = 0 )
1303+ prefix_inputs .append (
1304+ (
1305+ display_name ,
1306+ pred_arr ,
1307+ target_arr ,
1308+ pred_arr ,
1309+ target_arr ,
1310+ np .concatenate (bb .pred_fovs , axis = 0 ),
1311+ np .concatenate (bb .gt_fovs , axis = 0 ),
1312+ )
1313+ )
13401314
1341- # NaN-fill any prefix that had no cells (parallel pool would
1342- # otherwise skip it). Cheap; runs on empty arrays.
1343- expected_prefixes = ["CP" , "DINOv3" , "DynaCLR" ]
1344- if celldino_feature_extractor is not None :
1345- expected_prefixes .append ("CellDINO" )
1346- for name in expected_prefixes :
1347- if f"Dataset_{ name } _FID" not in dataset_row :
1315+ # Prefix with "Dataset_" so dataset-level FID/KID/cosine don't clobber
1316+ # per-FOV columns of the same name when merged into per-FOV rows.
1317+ def _compute_one (args ):
1318+ # MIND stays on CPU here even when use_gpu=True: 4 parallel threads
1319+ # racing on the same CUDA context would either serialize via the
1320+ # allocator (no speedup) or contend for memory with mid-eval FOV
1321+ # work in process executors. CPU MIND in a 4-thread BLAS-capped
1322+ # pool is competitive with serialized GPU MIND and bit-stable
1323+ # across runs that toggle use_gpu (torch's CPU vs CUDA RNG
1324+ # produce different streams for the same seed, breaking cross-
1325+ # leaf comparability of the MIND column).
1326+ name , p_metric , t_metric , p_probe , t_probe , fov_p , fov_t = args
13481327 raw = {
1349- ** compute_feature_similarity (np . empty (( 0 , 0 )), np . empty (( 0 , 0 )) , name ),
1350- ** _real_vs_pred_probe (np . empty (( 0 , 0 )), np . empty (( 0 , 0 )), np . empty ( 0 ), np . empty ( 0 ) , name ),
1328+ ** compute_feature_similarity (p_metric , t_metric , name ),
1329+ ** _real_vs_pred_probe (p_probe , t_probe , fov_p , fov_t , name ),
13511330 }
1352- dataset_row .update ({f"Dataset_{ k } " : v for k , v in raw .items ()})
1353-
1354- for row in all_feature_metrics :
1355- row .update (dataset_row )
1356- embedding_groups : dict [str , tuple ] = {}
1357- for key in _BACKBONE_KEYS :
1358- if key == "celldino" and celldino_feature_extractor is None :
1359- continue
1360- bb = parent_lists [key ]
1361- embedding_groups [f"pred_{ key } " ] = (bb .pred_feats , bb .pred_fovs , bb .pred_ts )
1362- embedding_groups [f"gt_{ key } " ] = (bb .gt_feats , bb .gt_fovs , bb .gt_ts )
1363- _save_embeddings (save_dir , embedding_groups )
1331+ return {f"Dataset_{ k } " : v for k , v in raw .items ()}
1332+
1333+ if prefix_inputs :
1334+ # Threads suffice: torch-fidelity, sklearn LBFGS, and numpy BLAS
1335+ # all release the GIL inside their hot loops. Cap inner BLAS to
1336+ # 1 thread so the outer threads don't oversubscribe cores.
1337+ with (
1338+ threadpool_limits (limits = 1 ),
1339+ ThreadPoolExecutor (max_workers = min (4 , len (prefix_inputs ))) as pool ,
1340+ ):
1341+ for result in pool .map (_compute_one , prefix_inputs ):
1342+ dataset_row .update (result )
1343+
1344+ # NaN-fill any prefix that had no cells (parallel pool would
1345+ # otherwise skip it). Cheap; runs on empty arrays.
1346+ expected_prefixes = ["CP" , "DINOv3" , "DynaCLR" ]
1347+ if celldino_feature_extractor is not None :
1348+ expected_prefixes .append ("CellDINO" )
1349+ for name in expected_prefixes :
1350+ if f"Dataset_{ name } _FID" not in dataset_row :
1351+ raw = {
1352+ ** compute_feature_similarity (np .empty ((0 , 0 )), np .empty ((0 , 0 )), name ),
1353+ ** _real_vs_pred_probe (np .empty ((0 , 0 )), np .empty ((0 , 0 )), np .empty (0 ), np .empty (0 ), name ),
1354+ }
1355+ dataset_row .update ({f"Dataset_{ k } " : v for k , v in raw .items ()})
1356+
1357+ for row in all_feature_metrics :
1358+ row .update (dataset_row )
1359+ embedding_groups : dict [str , tuple ] = {}
1360+ for key in _BACKBONE_KEYS :
1361+ if key == "celldino" and celldino_feature_extractor is None :
1362+ continue
1363+ bb = parent_lists [key ]
1364+ embedding_groups [f"pred_{ key } " ] = (bb .pred_feats , bb .pred_fovs , bb .pred_ts )
1365+ embedding_groups [f"gt_{ key } " ] = (bb .gt_feats , bb .gt_fovs , bb .gt_ts )
1366+ _save_embeddings (save_dir , embedding_groups )
13641367
13651368 dump_timings_csv (save_dir )
13661369
@@ -1705,12 +1708,17 @@ def evaluate_model(config: DictConfig):
17051708 pixel_metrics , mask_metrics , feature_metrics = _load_cached_final_metrics (config )
17061709 else :
17071710 pixel_metrics , mask_metrics , feature_metrics = evaluate_predictions (config )
1708- save_metrics (
1709- config ,
1710- pixel_metrics = pixel_metrics ,
1711- mask_metrics = mask_metrics ,
1712- feature_metrics = feature_metrics ,
1713- )
1711+ with region_timer ("save_metrics_csvs" , "<parent>" ):
1712+ save_metrics (
1713+ config ,
1714+ pixel_metrics = pixel_metrics ,
1715+ mask_metrics = mask_metrics ,
1716+ feature_metrics = feature_metrics ,
1717+ )
1718+ # Re-dump so save_metrics_csvs lands in eval_timing.csv. evaluate_predictions
1719+ # dumps once before save_metrics runs; this second dump overwrites with the
1720+ # full set (FOV-loop regions + save_metrics_csvs).
1721+ dump_timings_csv (Path (config .save .save_dir ))
17141722 return pixel_metrics , mask_metrics , feature_metrics
17151723
17161724
0 commit comments