From 3ca46a3f551ee80fd1a607700803661a9dd0f59c Mon Sep 17 00:00:00 2001 From: Gav Sturm Date: Tue, 11 Aug 2026 09:24:14 -0700 Subject: [PATCH 01/13] vendor diffex/ from diffex-interpretability (baseline for interpretability refactor) Brings the self-contained DiffEx interpretability pipeline (classifier/diffae/ directions + figures/viewer/kyle_pcs) onto the public-release branch. No conflict: refactor had no diffex/ dir. Reorg (rename, _internal quarantine, setacc subdir) follows in subsequent commits. --- .../attention/diffex/CROPSEQ_TO_MORPHOLOGY.md | 79 + src/ops_model/models/attention/diffex/PLAN.md | 712 +++++++ .../models/attention/diffex/README.md | 40 + .../attention/diffex/classifier/README.md | 61 + .../attention/diffex/classifier/__init__.py | 4 + .../attention/diffex/classifier/aggregate.py | 87 + .../diffex/classifier/celldino_features.py | 39 + .../attention/diffex/classifier/config.py | 69 + .../attention/diffex/classifier/data.py | 208 ++ .../attention/diffex/classifier/models.py | 34 + .../models/attention/diffex/classifier/run.py | 100 + .../attention/diffex/classifier/submit.py | 97 + .../attention/diffex/classifier/train.py | 76 + .../attention/diffex/diffae/__init__.py | 6 + .../models/attention/diffex/diffae/config.py | 74 + .../models/attention/diffex/diffae/data.py | 98 + .../diffex/diffae/diagnose_conditioning.py | 133 ++ .../models/attention/diffex/diffae/model.py | 64 + .../attention/diffex/diffae/plot_metrics.py | 65 + .../models/attention/diffex/diffae/recon.py | 69 + .../models/attention/diffex/diffae/run.py | 83 + .../models/attention/diffex/diffae/submit.py | 109 + .../models/attention/diffex/diffae/train.py | 272 +++ .../attention/diffex/diffae/virtstain_eval.py | 133 ++ .../diffex/diffae/virtstain_multi.py | 321 +++ .../attention/diffex/directions/__init__.py | 6 + .../attention/diffex/directions/batch.py | 172 ++ .../attention/diffex/directions/config.py | 72 + .../attention/diffex/directions/data.py | 61 + .../attention/diffex/directions/flow.py | 97 + .../attention/diffex/directions/grid.py | 131 ++ .../attention/diffex/directions/losses.py | 35 + .../attention/diffex/directions/make_gifs.py | 371 ++++ .../attention/diffex/directions/model.py | 25 + .../diffex/directions/proto_ddim_anchors.py | 162 ++ .../attention/diffex/directions/rank.py | 59 + .../models/attention/diffex/directions/run.py | 114 + .../attention/diffex/directions/submit.py | 60 + .../diffex/directions/train_directions.py | 36 + .../attention/diffex/directions/traverse.py | 224 ++ .../diffex/figures/METHODS_final.txt | 70 + .../figures/METHODS_traversal_montage.md | 112 + .../figures/METHODS_traversal_montage.txt | 38 + .../diffex/figures/_setacc_common.py | 157 ++ .../attention/diffex/figures/_setacc_phase.py | 60 + .../diffex/figures/auto_pick_and_plot.py | 112 + .../diffex/figures/cis_golgi_alternatives.py | 26 + .../diffex/figures/debug_setacc_top100.py | 88 + .../diffex/figures/ebi_peripheral_droplets.py | 204 ++ .../figures/figure4_morpho_traversal.py | 299 +++ .../diffex/figures/figure4_morpho_violin.py | 147 ++ .../diffex/figures/figure4_setacc_panel.py | 74 + .../figures/figure4_setacc_panel_fluorB.py | 24 + .../figures/figure4_setacc_panel_newpheno.py | 70 + .../figures/figure4_setacc_panel_phase.py | 10 + .../figures/figure_ebi_morpho_violin.py | 398 ++++ .../figures/figure_multirank_ebi_grid.py | 274 +++ .../diffex/figures/fluor_panel_montages.py | 69 + .../diffex/figures/fluor_shap_montages.py | 140 ++ .../figures/gen_validation/bag_sweep_plots.py | 382 ++++ .../figures/gen_validation/bag_sweep_score.py | 58 + .../gen_validation/centroid_bagsweep.py | 108 + .../figures/gen_validation/centroid_halves.py | 82 + .../centroid_pooled_bagsweep.py | 109 + .../gen_validation/control_halves_zscore.py | 106 + .../diffex/figures/gen_validation/embcheck.py | 36 + .../gen_validation/embedding_diagnostics.py | 133 ++ .../figure4_bagsize_reachfrac.py | 74 + .../figure4_v5_accuracy_summary.py | 140 ++ .../gen_validation/gen_alpha_embedding.py | 249 +++ .../figures/gen_validation/gen_embed_refit.py | 498 +++++ .../gen_validation/gen_phate_passthrough.py | 516 +++++ .../gen_validation/gen_real_centroid.py | 360 ++++ .../gen_validation/gen_real_distinct.py | 182 ++ .../figures/gen_validation/ntc_inverse_gap.py | 175 ++ .../gen_validation/patch_cache_real.py | 56 + .../gen_validation/publish_multibag_page.py | 177 ++ .../figures/gen_validation/rank_summary.py | 147 ++ .../figures/gen_validation/st_halves_score.py | 54 + .../figures/gen_validation/std_anchor_test.py | 60 + .../figures/gen_validation/stepabl_compare.py | 71 + .../gen_validation/valid200_alphastep.py | 38 + .../gen_validation/valid200_cache_build.py | 85 + .../gen_validation/valid200_capcheck.py | 41 + .../gen_validation/valid200_map_compare.py | 105 + .../gen_validation/valid200_metrics.py | 72 + .../attention/diffex/figures/nc_ratio.py | 219 ++ .../diffex/figures/ntc_anchor_compare.py | 63 + .../diffex/figures/phase_montages.py | 12 + .../diffex/figures/phase_multibag_montages.py | 68 + .../diffex/figures/phase_sample_montages.py | 23 + .../diffex/figures/rab_candidate_montages.py | 22 + .../diffex/figures/rebuild_traversals_n100.py | 181 ++ .../figures/traversal_montage_schematic.py | 142 ++ .../figures/virtual_staining_schematic.py | 202 ++ .../diffex/kyle_pcs/build_static_explorer.py | 436 ++++ .../diffex/kyle_pcs/compute_pc_strips.py | 442 ++++ .../attention/diffex/viewer/__init__.py | 1 + .../diffex/viewer/_altanchor_build.py | 123 ++ .../attention/diffex/viewer/_anchortest.py | 19 + .../diffex/viewer/_build_stepablation.py | 74 + .../diffex/viewer/_build_v5_inverted.py | 586 +++++ .../diffex/viewer/_build_v5_montages.py | 119 ++ .../diffex/viewer/_build_valid200.py | 169 ++ .../diffex/viewer/_consolidate_cells.py | 106 + .../diffex/viewer/_fluor_complex_build.py | 136 ++ .../diffex/viewer/_fluor_topcells.py | 114 + .../diffex/viewer/_fluor_v5_build.py | 160 ++ .../diffex/viewer/_migrate_v4_to_v5.py | 49 + .../attention/diffex/viewer/_phase_vs.py | 427 ++++ .../attention/diffex/viewer/_rebuild_v5.py | 59 + .../attention/diffex/viewer/_rescore_rank.py | 66 + .../attention/diffex/viewer/_score_v4.py | 37 + .../attention/diffex/viewer/_v4acc_test.py | 17 + .../diffex/viewer/_verify_pt_space.py | 47 + .../diffex/viewer/_verify_score_bridge.py | 26 + .../diffex/viewer/altanchor_pairs.json | 326 +++ .../attention/diffex/viewer/anchor_cells.py | 75 + .../diffex/viewer/build_attention_heads.py | 203 ++ .../diffex/viewer/build_complex_ebi_map.py | 58 + .../viewer/build_fluor_shap_rankings.py | 114 + .../diffex/viewer/build_montage_features.py | 61 + .../diffex/viewer/build_pc_crops_masked.py | 200 ++ .../diffex/viewer/build_pc_features.py | 202 ++ .../attention/diffex/viewer/build_pc_walks.py | 161 ++ .../attention/diffex/viewer/build_pcs.py | 165 ++ .../diffex/viewer/build_pcs_marker.py | 356 ++++ .../viewer/build_phase_shap_rankings.py | 93 + .../diffex/viewer/build_phate_figure.py | 252 +++ .../diffex/viewer/build_setacc_bins.py | 62 + .../diffex/viewer/build_setacc_bymarker.py | 48 + .../diffex/viewer/build_top_cells.py | 199 ++ .../diffex/viewer/build_umap_montage.py | 241 +++ .../models/attention/diffex/viewer/catalog.py | 196 ++ .../attention/diffex/viewer/deploy/README.md | 24 + .../attention/diffex/viewer/marker_leaves.py | 83 + .../diffex/viewer/mimic_alex_embed.py | 461 ++++ .../diffex/viewer/morpho_pipeline.py | 1237 +++++++++++ .../attention/diffex/viewer/morphometrics.py | 170 ++ .../attention/diffex/viewer/nway_clf.py | 120 ++ .../diffex/viewer/phenotype_cells.py | 106 + .../attention/diffex/viewer/precompute.py | 500 +++++ .../diffex/viewer/render_montage_scales.py | 364 ++++ .../diffex/viewer/score_generated.py | 285 +++ .../attention/diffex/viewer/set_classifier.py | 127 ++ .../models/attention/diffex/viewer/submit.py | 256 +++ .../attention/diffex/viewer/webapp/app.js | 1888 +++++++++++++++++ .../diffex/viewer/webapp/biohub-mark.png | Bin 0 -> 7120 bytes .../diffex/viewer/webapp/biohub-wordmark.png | Bin 0 -> 22582 bytes .../viewer/webapp/build_gene_narratives.py | 48 + .../attention/diffex/viewer/webapp/gif.js | 3 + .../diffex/viewer/webapp/gif.worker.js | 3 + .../attention/diffex/viewer/webapp/index.html | 371 ++++ .../attention/diffex/viewer/webapp/methods.js | 406 ++++ .../diffex/viewer/webapp/morpho_demo.html | 333 +++ .../diffex/viewer/webapp/openseadragon.min.js | 9 + .../diffex/viewer/webapp/opsin-eyes.svg | 1 + .../attention/diffex/viewer/webapp/style.css | 521 +++++ 158 files changed, 25617 insertions(+) create mode 100644 src/ops_model/models/attention/diffex/CROPSEQ_TO_MORPHOLOGY.md create mode 100644 src/ops_model/models/attention/diffex/PLAN.md create mode 100644 src/ops_model/models/attention/diffex/README.md create mode 100644 src/ops_model/models/attention/diffex/classifier/README.md create mode 100644 src/ops_model/models/attention/diffex/classifier/__init__.py create mode 100644 src/ops_model/models/attention/diffex/classifier/aggregate.py create mode 100644 src/ops_model/models/attention/diffex/classifier/celldino_features.py create mode 100644 src/ops_model/models/attention/diffex/classifier/config.py create mode 100644 src/ops_model/models/attention/diffex/classifier/data.py create mode 100644 src/ops_model/models/attention/diffex/classifier/models.py create mode 100644 src/ops_model/models/attention/diffex/classifier/run.py create mode 100644 src/ops_model/models/attention/diffex/classifier/submit.py create mode 100644 src/ops_model/models/attention/diffex/classifier/train.py create mode 100644 src/ops_model/models/attention/diffex/diffae/__init__.py create mode 100644 src/ops_model/models/attention/diffex/diffae/config.py create mode 100644 src/ops_model/models/attention/diffex/diffae/data.py create mode 100644 src/ops_model/models/attention/diffex/diffae/diagnose_conditioning.py create mode 100644 src/ops_model/models/attention/diffex/diffae/model.py create mode 100644 src/ops_model/models/attention/diffex/diffae/plot_metrics.py create mode 100644 src/ops_model/models/attention/diffex/diffae/recon.py create mode 100644 src/ops_model/models/attention/diffex/diffae/run.py create mode 100644 src/ops_model/models/attention/diffex/diffae/submit.py create mode 100644 src/ops_model/models/attention/diffex/diffae/train.py create mode 100644 src/ops_model/models/attention/diffex/diffae/virtstain_eval.py create mode 100644 src/ops_model/models/attention/diffex/diffae/virtstain_multi.py create mode 100644 src/ops_model/models/attention/diffex/directions/__init__.py create mode 100644 src/ops_model/models/attention/diffex/directions/batch.py create mode 100644 src/ops_model/models/attention/diffex/directions/config.py create mode 100644 src/ops_model/models/attention/diffex/directions/data.py create mode 100644 src/ops_model/models/attention/diffex/directions/flow.py create mode 100644 src/ops_model/models/attention/diffex/directions/grid.py create mode 100644 src/ops_model/models/attention/diffex/directions/losses.py create mode 100644 src/ops_model/models/attention/diffex/directions/make_gifs.py create mode 100644 src/ops_model/models/attention/diffex/directions/model.py create mode 100644 src/ops_model/models/attention/diffex/directions/proto_ddim_anchors.py create mode 100644 src/ops_model/models/attention/diffex/directions/rank.py create mode 100644 src/ops_model/models/attention/diffex/directions/run.py create mode 100644 src/ops_model/models/attention/diffex/directions/submit.py create mode 100644 src/ops_model/models/attention/diffex/directions/train_directions.py create mode 100644 src/ops_model/models/attention/diffex/directions/traverse.py create mode 100644 src/ops_model/models/attention/diffex/figures/METHODS_final.txt create mode 100644 src/ops_model/models/attention/diffex/figures/METHODS_traversal_montage.md create mode 100644 src/ops_model/models/attention/diffex/figures/METHODS_traversal_montage.txt create mode 100644 src/ops_model/models/attention/diffex/figures/_setacc_common.py create mode 100644 src/ops_model/models/attention/diffex/figures/_setacc_phase.py create mode 100644 src/ops_model/models/attention/diffex/figures/auto_pick_and_plot.py create mode 100644 src/ops_model/models/attention/diffex/figures/cis_golgi_alternatives.py create mode 100644 src/ops_model/models/attention/diffex/figures/debug_setacc_top100.py create mode 100644 src/ops_model/models/attention/diffex/figures/ebi_peripheral_droplets.py create mode 100644 src/ops_model/models/attention/diffex/figures/figure4_morpho_traversal.py create mode 100644 src/ops_model/models/attention/diffex/figures/figure4_morpho_violin.py create mode 100644 src/ops_model/models/attention/diffex/figures/figure4_setacc_panel.py create mode 100644 src/ops_model/models/attention/diffex/figures/figure4_setacc_panel_fluorB.py create mode 100644 src/ops_model/models/attention/diffex/figures/figure4_setacc_panel_newpheno.py create mode 100644 src/ops_model/models/attention/diffex/figures/figure4_setacc_panel_phase.py create mode 100644 src/ops_model/models/attention/diffex/figures/figure_ebi_morpho_violin.py create mode 100644 src/ops_model/models/attention/diffex/figures/figure_multirank_ebi_grid.py create mode 100644 src/ops_model/models/attention/diffex/figures/fluor_panel_montages.py create mode 100644 src/ops_model/models/attention/diffex/figures/fluor_shap_montages.py create mode 100644 src/ops_model/models/attention/diffex/figures/gen_validation/bag_sweep_plots.py create mode 100644 src/ops_model/models/attention/diffex/figures/gen_validation/bag_sweep_score.py create mode 100644 src/ops_model/models/attention/diffex/figures/gen_validation/centroid_bagsweep.py create mode 100644 src/ops_model/models/attention/diffex/figures/gen_validation/centroid_halves.py create mode 100644 src/ops_model/models/attention/diffex/figures/gen_validation/centroid_pooled_bagsweep.py create mode 100644 src/ops_model/models/attention/diffex/figures/gen_validation/control_halves_zscore.py create mode 100644 src/ops_model/models/attention/diffex/figures/gen_validation/embcheck.py create mode 100644 src/ops_model/models/attention/diffex/figures/gen_validation/embedding_diagnostics.py create mode 100644 src/ops_model/models/attention/diffex/figures/gen_validation/figure4_bagsize_reachfrac.py create mode 100644 src/ops_model/models/attention/diffex/figures/gen_validation/figure4_v5_accuracy_summary.py create mode 100644 src/ops_model/models/attention/diffex/figures/gen_validation/gen_alpha_embedding.py create mode 100644 src/ops_model/models/attention/diffex/figures/gen_validation/gen_embed_refit.py create mode 100644 src/ops_model/models/attention/diffex/figures/gen_validation/gen_phate_passthrough.py create mode 100644 src/ops_model/models/attention/diffex/figures/gen_validation/gen_real_centroid.py create mode 100644 src/ops_model/models/attention/diffex/figures/gen_validation/gen_real_distinct.py create mode 100644 src/ops_model/models/attention/diffex/figures/gen_validation/ntc_inverse_gap.py create mode 100644 src/ops_model/models/attention/diffex/figures/gen_validation/patch_cache_real.py create mode 100644 src/ops_model/models/attention/diffex/figures/gen_validation/publish_multibag_page.py create mode 100644 src/ops_model/models/attention/diffex/figures/gen_validation/rank_summary.py create mode 100644 src/ops_model/models/attention/diffex/figures/gen_validation/st_halves_score.py create mode 100644 src/ops_model/models/attention/diffex/figures/gen_validation/std_anchor_test.py create mode 100644 src/ops_model/models/attention/diffex/figures/gen_validation/stepabl_compare.py create mode 100644 src/ops_model/models/attention/diffex/figures/gen_validation/valid200_alphastep.py create mode 100644 src/ops_model/models/attention/diffex/figures/gen_validation/valid200_cache_build.py create mode 100644 src/ops_model/models/attention/diffex/figures/gen_validation/valid200_capcheck.py create mode 100644 src/ops_model/models/attention/diffex/figures/gen_validation/valid200_map_compare.py create mode 100644 src/ops_model/models/attention/diffex/figures/gen_validation/valid200_metrics.py create mode 100644 src/ops_model/models/attention/diffex/figures/nc_ratio.py create mode 100644 src/ops_model/models/attention/diffex/figures/ntc_anchor_compare.py create mode 100644 src/ops_model/models/attention/diffex/figures/phase_montages.py create mode 100644 src/ops_model/models/attention/diffex/figures/phase_multibag_montages.py create mode 100644 src/ops_model/models/attention/diffex/figures/phase_sample_montages.py create mode 100644 src/ops_model/models/attention/diffex/figures/rab_candidate_montages.py create mode 100644 src/ops_model/models/attention/diffex/figures/rebuild_traversals_n100.py create mode 100644 src/ops_model/models/attention/diffex/figures/traversal_montage_schematic.py create mode 100644 src/ops_model/models/attention/diffex/figures/virtual_staining_schematic.py create mode 100644 src/ops_model/models/attention/diffex/kyle_pcs/build_static_explorer.py create mode 100644 src/ops_model/models/attention/diffex/kyle_pcs/compute_pc_strips.py create mode 100644 src/ops_model/models/attention/diffex/viewer/__init__.py create mode 100644 src/ops_model/models/attention/diffex/viewer/_altanchor_build.py create mode 100644 src/ops_model/models/attention/diffex/viewer/_anchortest.py create mode 100644 src/ops_model/models/attention/diffex/viewer/_build_stepablation.py create mode 100644 src/ops_model/models/attention/diffex/viewer/_build_v5_inverted.py create mode 100644 src/ops_model/models/attention/diffex/viewer/_build_v5_montages.py create mode 100644 src/ops_model/models/attention/diffex/viewer/_build_valid200.py create mode 100644 src/ops_model/models/attention/diffex/viewer/_consolidate_cells.py create mode 100644 src/ops_model/models/attention/diffex/viewer/_fluor_complex_build.py create mode 100644 src/ops_model/models/attention/diffex/viewer/_fluor_topcells.py create mode 100644 src/ops_model/models/attention/diffex/viewer/_fluor_v5_build.py create mode 100644 src/ops_model/models/attention/diffex/viewer/_migrate_v4_to_v5.py create mode 100644 src/ops_model/models/attention/diffex/viewer/_phase_vs.py create mode 100644 src/ops_model/models/attention/diffex/viewer/_rebuild_v5.py create mode 100644 src/ops_model/models/attention/diffex/viewer/_rescore_rank.py create mode 100644 src/ops_model/models/attention/diffex/viewer/_score_v4.py create mode 100644 src/ops_model/models/attention/diffex/viewer/_v4acc_test.py create mode 100644 src/ops_model/models/attention/diffex/viewer/_verify_pt_space.py create mode 100644 src/ops_model/models/attention/diffex/viewer/_verify_score_bridge.py create mode 100644 src/ops_model/models/attention/diffex/viewer/altanchor_pairs.json create mode 100644 src/ops_model/models/attention/diffex/viewer/anchor_cells.py create mode 100644 src/ops_model/models/attention/diffex/viewer/build_attention_heads.py create mode 100644 src/ops_model/models/attention/diffex/viewer/build_complex_ebi_map.py create mode 100644 src/ops_model/models/attention/diffex/viewer/build_fluor_shap_rankings.py create mode 100644 src/ops_model/models/attention/diffex/viewer/build_montage_features.py create mode 100644 src/ops_model/models/attention/diffex/viewer/build_pc_crops_masked.py create mode 100644 src/ops_model/models/attention/diffex/viewer/build_pc_features.py create mode 100644 src/ops_model/models/attention/diffex/viewer/build_pc_walks.py create mode 100644 src/ops_model/models/attention/diffex/viewer/build_pcs.py create mode 100644 src/ops_model/models/attention/diffex/viewer/build_pcs_marker.py create mode 100644 src/ops_model/models/attention/diffex/viewer/build_phase_shap_rankings.py create mode 100644 src/ops_model/models/attention/diffex/viewer/build_phate_figure.py create mode 100644 src/ops_model/models/attention/diffex/viewer/build_setacc_bins.py create mode 100644 src/ops_model/models/attention/diffex/viewer/build_setacc_bymarker.py create mode 100644 src/ops_model/models/attention/diffex/viewer/build_top_cells.py create mode 100644 src/ops_model/models/attention/diffex/viewer/build_umap_montage.py create mode 100644 src/ops_model/models/attention/diffex/viewer/catalog.py create mode 100644 src/ops_model/models/attention/diffex/viewer/deploy/README.md create mode 100644 src/ops_model/models/attention/diffex/viewer/marker_leaves.py create mode 100644 src/ops_model/models/attention/diffex/viewer/mimic_alex_embed.py create mode 100644 src/ops_model/models/attention/diffex/viewer/morpho_pipeline.py create mode 100644 src/ops_model/models/attention/diffex/viewer/morphometrics.py create mode 100644 src/ops_model/models/attention/diffex/viewer/nway_clf.py create mode 100644 src/ops_model/models/attention/diffex/viewer/phenotype_cells.py create mode 100644 src/ops_model/models/attention/diffex/viewer/precompute.py create mode 100644 src/ops_model/models/attention/diffex/viewer/render_montage_scales.py create mode 100644 src/ops_model/models/attention/diffex/viewer/score_generated.py create mode 100644 src/ops_model/models/attention/diffex/viewer/set_classifier.py create mode 100644 src/ops_model/models/attention/diffex/viewer/submit.py create mode 100644 src/ops_model/models/attention/diffex/viewer/webapp/app.js create mode 100644 src/ops_model/models/attention/diffex/viewer/webapp/biohub-mark.png create mode 100644 src/ops_model/models/attention/diffex/viewer/webapp/biohub-wordmark.png create mode 100644 src/ops_model/models/attention/diffex/viewer/webapp/build_gene_narratives.py create mode 100644 src/ops_model/models/attention/diffex/viewer/webapp/gif.js create mode 100644 src/ops_model/models/attention/diffex/viewer/webapp/gif.worker.js create mode 100644 src/ops_model/models/attention/diffex/viewer/webapp/index.html create mode 100644 src/ops_model/models/attention/diffex/viewer/webapp/methods.js create mode 100644 src/ops_model/models/attention/diffex/viewer/webapp/morpho_demo.html create mode 100644 src/ops_model/models/attention/diffex/viewer/webapp/openseadragon.min.js create mode 100644 src/ops_model/models/attention/diffex/viewer/webapp/opsin-eyes.svg create mode 100644 src/ops_model/models/attention/diffex/viewer/webapp/style.css diff --git a/src/ops_model/models/attention/diffex/CROPSEQ_TO_MORPHOLOGY.md b/src/ops_model/models/attention/diffex/CROPSEQ_TO_MORPHOLOGY.md new file mode 100644 index 0000000..78c2a07 --- /dev/null +++ b/src/ops_model/models/attention/diffex/CROPSEQ_TO_MORPHOLOGY.md @@ -0,0 +1,79 @@ +# Idea: transcriptome-controlled morphology generation (CROP-seq → DiffAE) + +**Goal:** use paired CROP-seq on the same geneKO library to let the DiffAE generate *how a cell's phenotype +changes as its transcriptome moves toward a KO state* — i.e. drive the traversal by a **transcriptional +direction** (CROP-seq) instead of (or in addition to) the CellDINO morphological direction. + +## Reframing vs today +Current DiffEx traversal = real cell → DDIM-inverted `xT` (identity/nuisance) + **CellDINO gene-direction** +(morphology conditioning), morph α NTC→KO, decode image. This idea keeps the entire image decoder (DDIM + +inversion + guidance `w`) and swaps the *driver* to a transcriptional signature. + +## Hard constraint (shapes everything): no single-cell pairing +CROP-seq is destructive scRNA-seq; OPS is imaging — **no cell is in both modalities**. So supervision is only +**per-perturbation** (gene KO → mean transcriptional shift Δt_g AND a morphological distribution), never +cell-level `(transcriptome → image)`. +- Model learns **transcriptional-signature → morphological-distribution**; within-gene image variation comes + from the stochastic `xT` (same as today). +- Conditioning vectors are **per-gene pseudobulk** (or per-guide if guide calls are clean). + +## Two paths + +### Path B — reuse the trained DiffAE via a transcriptome→CellDINO map (POC first) +Fit a perturbation-level regressor `Δt_g → ΔCellDINO_g` (linear → small MLP) over the shared KOs. A +transcriptional vector → predicted CellDINO shift → **existing morpho DiffAE renders it**. +- Pros: reuses the whole trained pipeline + viewer; days not weeks. **The map's R² is itself a headline + result** ("fraction of KO morphology predictable from KO transcriptome"). +- Cons: bottlenecked through CellDINO. + +### Path A — condition the DiffAE directly on transcriptome (full, CPA-flavored) +Project t through `cond_proj` into the FiLM/cross-attn slot the CellDINO emb uses now; train on +`(image_i, t_{gene(i)})`. Cleanest is a **CPA-style shared perturbation embedding** `e_g`: a transcriptome +decoder reconstructs CROP-seq (`NTC + e_g`), the DiffAE decodes the image (`anchor xT + e_g`), `e_g` shared → +ties the modalities through one latent. Any transcriptional state → `e_g` → image. +- Pros: end-to-end; supports unseen/combined signatures + continuous "dial a pathway, watch morphology". +- Cons: real training effort; guard against collapse to gene-means (xT + guidance mitigate, as today). + +**Plan:** B as a weekend POC (also tests whether transcriptome predicts morphology at all) → A if promising. + +## Transcriptional vector options (cheapest first) +1. pseudobulk logFC vs NTC (per gene); 2. learned scRNA latent (scVI/PCA) mean per gene; 3. pathway/program +module scores (most interpretable "dials"). Start with (1)/(2) for the direction, expose (3) as the control. + +## Concrete CROP-seq source — Duo's sVAE+ gene-program embeddings (June 2025) +Duo Peng built the CROP-seq embeddings we should use as option (2)/(3). This is a **sparse VAE (sVAE+)** — Lopez +et al. 2023, *Learning Causal Representations of Single Cells via Sparse Mechanism Shift Modeling* — so the latent +axes are interpretable **gene programs**, which is exactly the "dial a pathway" control we wanted. +- Confluence: [sVAE approach to gene programs v3](https://czbiohub.atlassian.net/wiki/spaces/dashboard/pages/5199986706/sVAE+approach+to+gene+programs+v3) + — the **purple "sVAE embeddings" section** has the embeddings file. +- **Run the encoder to embed new expression profiles** — point setup at the parent results folder: + `/hpc/projects/data.science/duo.peng/sVAEplus/sVAEplus/6000HVG/svaeplus_results_2_256_1_200_0.5/` + - trained encoder: `best_model/model.pt`; params `best_params.json` (n_layers=2, n_hidden=256, + sparse_mask_penalty=1.0, kl_warmup=200, dropout=0.05); code root `.../sVAEplus/sVAEplus/sVAE-main` + `ops_utils`, `install.sh`. + - expression values (normalized, sVAE+-compat, filtered, 6000 HVG): + `.../svaeplus_results_2_256_1_200_0.5/CropSeq_June2025_filtered_normalized_compat_forsvaeplus_filtered.h5ad` +- ⚠️ **`gene_loadings.csv` (gene → gene-program activity) is a post-hoc *linear* summary — do NOT use it as the + mapping.** The real expression → program mapping is **non-linear**; get it by running the encoder on expression + values, not by the loadings matrix. +- Fit for the plan: these program embeddings are the transcriptional vector for **Path B** (`Δprogram_g → + ΔCellDINO_g`) and the shared-latent seed / conditioning signal for **Path A**. Per-gene means over the encoder + output give Δt_g; NTC cells in the same h5ad define the control baseline. + +## The novel payoff: transcriptome↔morphology divergence map +Plot every gene by (transcriptional effect size, morphological effect size). The DiffAE then lets you *see*: +- **transcriptionally loud, morphologically silent** → counterfactual "what it would look like if it manifested" +- **morphologically loud, transcriptionally quiet** → morphology carrying signal transcriptome misses +- cross-modal interpolation between two genes' transcriptomes; agreement w/ CellDINO morph = validation, + divergence = the interesting biology. + +## Viewer tab concept +"Transcriptome → Morphology" tab: α-slider drives the **transcriptional** traversal; side-by-side vs the +existing CellDINO-driven morph (agreement = validation, divergence = biology). Reuses traversal/montage render. + +## To scope +1. ~~CROP-seq path/format~~ → **resolved**: Duo's sVAE+ h5ad + trained encoder (see section above). Still TBD: + gene-overlap of the CROP-seq library with the imaging 1000-lib; whether to embed with the encoder or use + Duo's precomputed embeddings from the purple Confluence section. +2. per-gene vs per-guide signatures. +3. matched NTC/control in CROP-seq to define Δ. +4. payoff emphasis: generator vs divergence-map. diff --git a/src/ops_model/models/attention/diffex/PLAN.md b/src/ops_model/models/attention/diffex/PLAN.md new file mode 100644 index 0000000..26e7b49 --- /dev/null +++ b/src/ops_model/models/attention/diffex/PLAN.md @@ -0,0 +1,712 @@ +# DiffEx interpretability — plan + +Living design doc. Goal: interpret geneKO / protein-complex phenotypes **into image space** +using a DiffEx-style diffusion counterfactual (arXiv:2502.09663), since OP/CP classical +features are judged too weak to describe the phenotypes. + +## What is DiffEx? +DiffEx (*Explaining a Classifier with Diffusion Models to Identify Microscopic Cellular +Variations*, arXiv:2502.09663) explains any image classifier by generating visually interpretable +**counterfactuals** — showing, in pixel space, what about an image drives the classifier. +- **Architecture:** a **diffusion autoencoder (DiffAE)** — semantic encoder → low-dim latent + `z_sem`; conditional diffusion decoder reconstructs the image from `z_sem`. The classifier score + is concatenated onto `z_sem`; a bank of MLP **direction models** is trained in that latent with a + **contrastive loss** → distinct, disentangled directions. Explain class k = shift `z_sem` along a + direction and decode. +- **Interpretability features we exploit:** counterfactual morphs; a *global, reusable* attribute + vocabulary (directions shared across all classes — an image-grounded OP/CP replacement); + per-class attribute ranking; classifier-agnostic & forward-only (no retraining); continuous edit + strength α (dose-like morphs); quantitative faithfulness via re-encoding. + +Status (historical): *designing the per-cell classifier* — that phase is long done; see +ACTIVE EFFORTS below for the current state. + +--- + +## ACTIVE EFFORTS (dashboard — updated 2026-07-08) + +### LATEST (2026-07-08) — phenotype-cell handoff, v2 mAP, EBI matrix, viewer embedding tab +- **Phenotype-cell CSV for Ritvik** (`viewer/phenotype_cells.py`) → `viewer_assets/phenotype_cells_for_attention.csv`. + 20 cells × (geneKO + EBI complex) × marker, for SetTransformer attention pixel-patches on the REAL + phenotype cells. **160,420 cells / 53 markers** (phase + 52 fluor). Cols incl `map_score`, `geneKO`, + `ebi_complex`, `rank_source`, `segmentation_id` (=pma `segmentation`), `x/y_pheno`, `rank`, `pma_attention`. + - **Per-marker top-20, NOT the model's global top-20** — the pma `rank` is GLOBAL per geneKO (across all + 56 channels), so `_csv_top` re-ranks WITHIN each (channel, perturbation) and takes the 20 highest-attention + cells present in that channel (two-pass chunked `head` keeps memory bounded). + - **Fluor filtered by mAP ≥ 0.2** (phase = ALL perturbations): geneKO by distinctiveness, complex by EBI mAP. + - **`rank_source` col:** `"model"` (all current cells). Reserved `"fallback"` for markers not in the model. +- **v2 distinctiveness switch:** `catalog.dist_matrix` now reads `paper_v2/with_cp/with_4i/all_livecell` + (single 56-reporter matrix: 43 live + 7 CP + 6 4i), replacing the paper_v1 3-way split. Added 4i + `FIXED_REP` mappings (p53/pRb/pS6/p21/b-catenin/c-Myc). **52/56 pma channels now map**; the 4 + excluded (NFkB, RSP6, Rb, gH2AX) are EXPECTED — genuinely absent from the v2 matrix. +- **EBI complex mAP matrix** (`viewer/build_complex_ebi_map.py`) → `complex_reporter_ebi_map.csv` (98×56). + Runs copairs `phenotypic_consistency_ebi` per-marker on the v2 `with_cp/with_4i` per_signal gene + embeddings, over ALL perturbations (activity_map=None), NOT the wrong `complex_reporter_chad_consistency`. + `catalog.complex_dist()` reads it. Also wired into the aggregation pipeline + (`post_process/combination/pca_optimization/aggregation.py` → `complex_reporter_ebi_consistency.csv`). +- **3 new live-cell markers (cisGolgi, VIM, LMNB1):** HAVE v2 distinctiveness + EBI mAP, but are NOT in + Alex's pma cell CSVs yet (his attention output predates them) → no attention-ranked cells with crop + metadata. **FALLBACK (per user, TODO):** select cells around the CENTROID of the existing CellDINO + embeddings. Source found: `{exp}/3-assembly/cell_dino_features_v2/anndata_objects/features_processed_.h5ad` + (e.g. `mStayGold-CENPRaltORF`, `VIM`, `LMNB1`) — 1024-d embedding + crop metadata (`label_int`=segmentation, + `x/y_position`, `well`, `experiment`, `perturbation`). Per (marker, pert): centroid → 20 closest → + `rank_source="fallback"`, `pma_attention`/`rank` null. (Or wait for Alex's reprocessed pma CSV.) +- **SetTransformer accuracy scoring (`viewer/set_classifier.py` + `viewer/mimic_alex_embed.py`) — PARKED, + waiting on SetTransformer v2 (no-mask classifier).** Full journey + why: + - Reconstructed Alex's cellstate-set-classifier (ISAB/PMA/cosine head); 5 ckpts in `v4/wandb/cellstate_set_classifier/`. + Real-bag CEILING validated: feeding Alex's own `.pt` embeddings → P(target) 0.90–0.999 (HSPA5/KIF11/POLR1B/TIMM23 all hit). + - Raw `embed_crops` → classifier FAILS (OOD, cos 0.47, constant argmax). Built `mimic_alex_embed` to reproduce Alex's + exact pipeline: **128 Phase2D crop → seg-mask (`cell_seg`) → percentile-norm → CellDINO → z-std(control)**. Findings: + the **segmentation mask is the load-bearing step** (cos 0.47→0.91); percentile-norm is canceled by CellDINO's z-score. + At realistic bag sizes (100 cells, per-experiment z-std) the mimic matches the ceiling: **95–100% hit-rate** on real cells. + - Generated-cell per-α curve works end-to-end for **POLR1B** (P 0.01→0.98, flips to target at α≥1.5) — proof the pipeline + is correct — but **the segmentation of GENERATED cells is the blocker**: cellpose on fake crops unreliably captures the + cell (latches onto the bright nucleolus, not the whole body), so only nucleolar-phenotype genes (POLR1B) score; HSPA5/ + KIF11/TIMM23 stay at P≈0 despite healthy embeddings. Diameter/centroid tuning didn't fix it robustly. + - **DECISION:** masking generated cells is too fragile to rely on. **Wait for SetTransformer v2 — a classifier trained + WITHOUT masks** → then our unmasked `embed_crops` is in-distribution, no cellpose needed, and the per-α bag score works + for all genes. The mimic (mask path) + POLR1B validation are kept as a reference/cross-check. +- **Model-metrics curves** (`diffae/plot_metrics.py`): loss + cond_ratio over epochs, one line per DiffAE + → `model_metrics_curves.png/.svg`. +- **Viewer embedding tab** (`build_umap_montage.py` + `webapp/`): OSD montage, UMAP↔PHATE, points/images + toggle, 44 anndata color-by fields, opacity/zoom sliders, click→perturbation sidebar; gene descriptions + from the gene-embedding h5ad (`gene_desc.json`, fixes VAMP2-style blanks). + +### STATUS SUMMARY (historical detail condensed 2026-07-08) +- **Generators:** phase `phase_v1` = PRODUCTION (0.468); 500k warm retrain PARKED (peaked 0.542 then + declined). Fluor **50/50 markers trained** (ep≥98). v2/v3 aug did NOT beat v1. Directions default = + deterministic **mean_diff** α (see build log for the full DiffAE saga). +- **Viewer** (`viewer/` — `submit.py`, `catalog.py`, `precompute.py`, `build_umap_montage.py`, `webapp/`): + static precompute → dependency-free web app; per-marker driver shares the NTC gather + dedups real + cells; embedding tab (see LATEST). Live demo `login-01:8765`. +- **Score:** authoritative = Alex's **SetTransformer** bag `P(target)` (§7 ckpts downloaded) — supersedes + the per-cell N-way MLP (`nway_clf.py`) and the old binary LR badge. Generated-cell bag scoring PARKED: the + mask-mimic works on real cells (95–100%) + POLR1B generated, but segmenting fake cells is too fragile → + waiting for SetTransformer v2 (no-mask classifier). See LATEST. +- **Infra PR (#51): OPENED + ACCEPTED/MERGED** — `diffex-viewer-dev.tf` in `sfbiohub-infra` (S3 bucket + `diffex-viewer-dev` + nonprod read-only IRSA role `biohub-nonprod-diffex-viewer` for SA `diffex-viewer` in + ns `argus-diffex-viewer-rdev` + read-write uploader role; mirrors `proteohub-argus-s3-reader-dev.tf`, 1 TB + ceiling). `terraform apply` provisions the bucket/roles → then `aws s3 sync viewer_assets/ s3://diffex-viewer-dev/` + → Argus boot-download. Next: create the app repo + `argus register` (see App-staging build-log entry). + +### OPEN BUILDOUT +1. **Full NTC drain — LAUNCHED 2026-07-11** (master job `34826112`): all ~1000 geneKO genes/marker for the + 46 valid-`rep` fluor markers (was top-8 seed; 42 were partial ~100–194, 4 hub markers already ~complete). + Command = `submit seed --map-thr 0 --timeout 720` (**not** `--all-genes`; that flag never existed — the + `--map-thr 0` = every gene with distinctiveness ≥ 0 = all ~1000). Resume is automatic (skips built targets). + The 4 rep=None markers (NFkB/RSP6/Rb/gH2AX) are intentionally excluded. ~500 GB. See build-log 2026-07-11. +2. **Full A→B anchors** — `submit anchors --k 10` across all markers + complexes. +3. **Fluor complex traversals — DONE (resume 2026-07-11, master `34826156`)**: 98 EBI complexes × 50 markers + were already ~complete; only 5 markers partial (peroxisome_Peroxi, pS6, pRb, NPM3, SRRM2) → `submit + fluor-complex` resume finishes them. (phase complex = 190 = 98 NTC-anchored + 92 complex→complex anchor pairs.) +4. **Wire SetTransformer bag score** into `precompute` + a per-α curve panel — BLOCKED on **SetTransformer v2 + (no-mask classifier)**; the mask path is too fragile on generated cells (see LATEST). Once v2 lands, unmasked + `embed_crops` bags score directly (no cellpose). +5. **S3 hosting** — infra PR #51 MERGED. Argus app scaffold built (`/hpc/mydata/gav.sturm/diffex-viewer`, forked + from `czbiohub-sf/mops-viewer`). Remaining: `terraform apply` → create `czbiohub-sf/diffex-viewer` repo → + `argus register`/bootstrap (needs argus CLI) → upload `viewer_assets/` (or hand to Kyle) → PR+`stack` label. See build-log. +6. **Centroid fallback** for cisGolgi/VIM/LMNB1 (see LATEST). +7. **Multi-α montage** — `--alphas` flag looping per-α decodes + an α switch in the explorer. +8. **Image-UMAP montage** (`czi-ai/latent-lens`) — idea track; needs full-gene coverage. +9. **Attention-head tab + viewer reorientation** (DONE) — new tab overlaying CellDINO attention-head + pixel weights (inferno) on the real phenotype cells; Browse (marker+perturbation) drives ALL views. + See `### 2026-07-08 — Attention-head tab` build-log entry. +10. **Per-marker embedding montages** (FUTURE) — the Embedding tab now switches its IMAGES per marker but + reuses the SHARED phase gene-UMAP LAYOUT for every marker (`submit montage` loops modalities, all with + the phase `UMAP_H5AD`; only tiles swap). Fluor markers currently place only their ~8 seed geneKO frames + (sparse). FUTURE: give each marker its OWN gene embedding (per-reporter CellDINO gene UMAP/PHATE), so + genes sit at that marker's coordinates. Needs a per-marker gene-embedding h5ad (`obsm X_umap/X_phate` per + gene) — does NOT exist yet (`pca_optimized_v0.3/.../paper_v2/` only has aggregate combos: phase_only, + all_livecell, with_cp, no_phase, only_cp — no per-single-marker layout). Build = aggregate per-marker + CellDINO gene embeddings → PCA → UMAP/PHATE per reporter, then `build_montage_web(modality=, + h5ad=)`. Also depends on fuller fluor geneKO traversals (item 1) to be non-sparse. + +--- + +### Scope (locked) +- **4 classifiers** = 2 modalities × 2 grains: + 1. phase-only · geneKO (1001 KO + NTC-way) + 2. phase-only · complex (98-way, EBI) + 3. all-fluorescence · geneKO + 4. all-fluorescence · complex +- **2 diffusion models** (phase-only, all-fluorescence) — unconditional, shared across grains. +- **Negative contrast = `distinct`** (vs all other geneKOs / complexes), to isolate exactly + what is unique to each perturbation. Falls out of the multi-class softmax for free (§1.3). + +--- + +## 0. Why this shape (decisions already made) + +- **Drop OP/CP.** DiffEx discovers its explanatory vocabulary directly in image space, so we + don't need to translate CellDINO → OP/CP → language. The diffusion model *is* the vocabulary. +- **One unconditional diffusion model, not one per gene.** It only learns to generate realistic + single-cell crops. DiffEx's *classifier guidance* does all per-gene steering at inference. +- **Diffusion autoencoder (DiffAE), faithful to the DiffEx paper.** Semantic encoder → latent + `z_sem` + conditional diffusion decoder, trained jointly. NOT latent diffusion (no separate VAE) + and NOT a bare guided DDPM. The contrastive direction discovery (§3) requires this semantic + latent — it's the core of the method, not an add-on. +- **A small per-cell classifier, NOT the SetTransformer.** The SetTransformer never sees pixels + (input = bag of CellDINO embeddings), so it gives DiffEx no `image→logit` gradient. Rather than + bolt CellDINO on the front and wrestle the per-bag→per-image mismatch, we train a clean + per-image classifier. The SetTransformer still does what it's proven at — **selecting the cells** + the classifier trains on. (Full-stack `image→CellDINO→SetTransformer` faithfulness check is a + later reviewer-defense nicety, not on the critical path.) + +## Final product (per gene-KO / per complex atlas page) + +1. **Evidence cells** — top-X attention cells passing the mAP/accuracy threshold. X is variable + per gene; the value itself is the **penetrance readout** (HSPA5 ≈ 10, RPL10 ≈ 800). Printed. +2. **Counterfactual morph** — NTC → KO (and optionally KO → NTC, often cleaner). Averaged across + the top-X cells, not one cherry-picked cell. +3. **Per-channel difference heatmap** — always show **phase + the top highest-mAP channel(s) for + that class** (NOT every channel); localizes the phenotype to the relevant marker. +4. **Faithfulness number** — % of counterfactuals whose re-encoded logit actually flipped to KO. + +--- + +## 1. THE CLASSIFIER (current focus) + +**3 candidate classifiers (DiffEx target) — per Alex, 2026-06-16.** Cell selection is settled +(top-X attention cells, §1.7); the options differ in WHAT classifier scores them — the source of +DiffEx's differentiable `image → class-k logit`. All train/score on top-attention cells (which +classify accurately). + +- **A. SetTransformer native** — score a generated cell by inserting it into a real NTC reference + bag and reading the bag logit. Most faithful (actual deployed model, no new training). + **Risk (Alex):** the SetTransformer may not treat one synthetic cell as real, and/or +1 cell + won't move the bag logit → weak faithfulness signal. Cross-check, don't depend on it. +- **B. ResNet CNN on single cells**, trained only on top-attention cells. DiffEx-standard. + **Only option needing NO CellDINO encoder** for generated images (scores pixels directly). + Simplest, cleanest gradients; not SOTA representation. Detail in §1.1–1.6 below. +- **C. CellDINO features + MLP**, trained only on top-attention cells. More SOTA / accurate, light + to train, reuses the already-provided CellDINO features. Scoring *generated* images runs the + **local** CellDINO encoder (`ops_model/models/cell_dino.py` → `CellDinoModel`: channel-adaptive + DINO ViT-L/16, resize 224 + per-image z-score, `in_channels=1`); heavier per step than B, NOT a blocker. + +**Recommendation:** prototype **B and C on HSPA5 (phase)** — **fully unblocked locally** (we have the +phase attention rankings, the crops, and the CellDINO encoder). Pick by per-cell classification +accuracy. B = lightest (small CNN on pixels, no CellDINO in the loop); C = stronger classifier but +runs ViT-L per generated image. A = faithfulness cross-check only. + +DiffEx needs a differentiable `image → class-k logit`. Option B (CNN) design detail: + +### 1.1 Input / modality — **DECISION NEEDED** +- SetTransformer regime is **phase-only** (paper-v1). To stay faithful + keep the diffusion model + to 1 channel, **recommend starting phase-only**. Extend to phase+fluor later (fluor carries the + ER/mito visual signal the reader wants, but multiplies diffusion difficulty). +- Crop: reuse `ops_utils.data.bbox_utils.BaseDataset` — 128×128, multi-channel, cell mask + available. Same loader the atlas uses. **No new data infra.** +- Open: feed the cell mask as an extra channel (focus model on the cell, suppress neighbours)? + Likely yes — cheap confound reduction. + +### 1.2 Architecture — DECISION: backbone (start CNN, escalate if needed) +DiffEx is **classifier-agnostic** (it explains a plain supervised classifier; doesn't need +internal layers). So this slot is a free choice. What DiffEx needs is NOT top accuracy but: +clean image→logit gradient, confound-robustness (keys on biology, not plate/intensity), and +faithfulness to the real phenotype. + +- **v1 (recommended): small from-scratch CNN** (ResNet18-ish, N-channel stem). Clean gradient, + no giant frozen encoder in the path, easy to keep confound-robust, fail-fast for the PoC. + Honest caveat: a from-scratch ResNet is **standard, not SOTA** for cell-image representation. +- **Escalation path if CNN counterfactuals are too weak/insensitive:** fine-tuned **CellDINO** + (your near-SOTA self-supervised ViT) or a channel-aware ViT (ChannelViT / DINO4Cells family). + SOTA representation → more sensitive classifier → more sensitive counterfactuals, at the cost + of a heavier gradient path and higher confound risk (frozen SSL features can encode batch). +- Tradeoff is real because phenotypes are subtle. Decide empirically on HSPA5: if the CNN can't + separate HSPA5 cleanly cross-experiment, escalate the backbone before blaming the generator. + +### 1.3 Task framing — multi-class, distinct contrast for free +- **One N-way softmax classifier** per (modality × grain): 1001-KO+NTC-way for geneKO, + 98-way for complex. Shared backbone, per-class linear heads. +- DiffEx guides toward `logit_k`. Because softmax is normalized against all other classes, + **guiding toward class k IS the `distinct` contrast** (what's unique to k vs everything else) — + no separate binary models needed. +- PoC = train the geneKO N-way model, then run DiffEx toward the HSPA5 logit. +- Caveat (logged, not a blocker): distinct **suppresses shared phenotypes** by construction + (two ER-stress genes won't show their shared ER signature, only their difference). Keep + **vs-NTC as a complementary second pass** for genes where the absolute phenotype is wanted. + +### 1.4 Training set (a query, not new infra) +- Positives for class k: top-X attention cells for k, ranked by `attn_geneko` / `pma_attention`, + passing the mAP/accuracy threshold. X variable per gene (→ penetrance). +- Negatives = the **top-attention cells of the OTHER classes** (like-with-like), NOT random/weak + cells of them — else the model relearns "has *any* phenotype" instead of "has *this* one". + With a softmax over top-attention cells per class this is automatic. +- Sources: + - attention sidecar: `/hpc/projects/icd.fast.ops/models/alex_lin_attention/v4/expansion_v1/per_experiment_v4_attn.parquet` + (cols: experiment, well, segmentation_id, attn_ebi, attn_geneko, attn_chad, ...) + - PMA parquets: `.../v4/pma_*` (cols: gene, experiment, well, segmentation, pma_attention) + - cell-set builder already exists: `organelle_profiler.feature_extraction.consolidate_top_attention_cells` +- Join `(experiment, well, segmentation_id)` → bbox → `BaseDataset` crops. + +### 1.5 Confound guardrail (critical — or DiffEx faithfully explains the batch effect) +- Balance / stratify NTC negatives across the same experiments+wells as the positives. +- **Validate cross-experiment** (train on subset of experiments, test on held-out) — biological + signal generalizes, plate/intensity artifacts don't. +- Sanity bar: per-cell classifier AUROC should track the SetTransformer's per-gene mAP ordering + (sharp genes like HSPA5 easy, diffuse genes like RPL10 hard). + +### 1.6 Success criteria +- HSPA5-vs-NTC held-out AUROC clearly > 0.5 and > a same-data **all-cells** classifier + (proves attention selection removes confounds). +- Cross-experiment generalization holds. + +--- + +## 2. Diffusion autoencoder (DiffAE) — the real cost/risk +Faithful to DiffEx. Two jointly-trained parts on single-cell crops (one DiffAE per modality: +phase-only, all-fluor): +- **Semantic encoder** → `z_sem` (a low-dim semantic latent capturing cell appearance). +- **Conditional diffusion decoder** (`UNet2D`, N input channels) that reconstructs the crop + conditioned on `z_sem` (+ stochastic DDIM latent for detail). +- Data is NOT the constraint (millions of cells); compute/engineering is. +- De-risk on HSPA5 PoC before scaling. First sanity check: round-trip reconstruction quality + (encode→decode) on held-out cells — if it can't reconstruct, directions are meaningless. + +## 3. Contrastive direction discovery (the interpretability core) +Faithful to DiffEx — this is where the explanation comes from, NOT per-image classifier guidance: +- Concatenate the §1 classifier score onto `z_sem` → semantic code. +- Train a bank of MLP direction models `D_1…D_N` that each shift the code by `α·Δz_k`, with a + **contrastive loss**: edits from the same direction stay similar, edits across directions stay + dissimilar → distinct, disentangled, reusable attributes. +- **Distinct contrast (per §1.3):** select the direction(s) that move the classifier toward + class k's logit → "what is distinct about geneKO/complex k". The discovered `D_1…D_N` form a + **shared attribute vocabulary across all 1001 genes / 98 complexes** — the real payoff vs a + one-off morph, and an image-grounded replacement for OP/CP. +- Faithfulness check, baked in day one: re-encode the edited image through §1 classifier, confirm + the logit actually moved toward k; report % success on each atlas page. + +--- + +## Milestones / de-risk order +1. **Classifier on HSPA5** (this task) — fail-fast signal that the cell sets are learnable. +2. **DiffAE** on the same crop set; gate on round-trip reconstruction quality (§2). +3. **Contrastive direction discovery** (§3); pick HSPA5's distinct direction; eyeball morph + + per-channel heatmap + flip-rate. +4. If HSPA5 works → scale to atlas (reuse shared directions). If not → it won't work anywhere; stop. + +## Open questions +- Phase-only vs phase+fluor for v1? (recommend phase-only) +- Mask as extra input channel? (recommend yes) +- Threshold definition for "X cells that pass": fixed mAP cutoff vs per-gene accuracy knee? +- Negative contrast for v1: NTC only, or also distinct/global? + +--- + +## §1.7 Cell selection (settled, all options) + Option-A bag scoring +**Settled — cell selection (all 3 classifier options):** cells fed to DiffEx / used to train the +classifier = the **top-X attention cells** (PMA attention rank), exactly as used for the atlas. + +**Option A only — scoring a *generated* image with the bag model:** insert the generated cell into +a fixed real NTC reference bag, read logit_k. Alex's concern: the SetTransformer may not treat the +synthetic cell as real and +1 cell may not move the bag logit. → why A is a cross-check, and why +B/C (self-contained single-cell classifiers) are the primary path. + +## Assets inventory (checked 2026-06-16) +**Have locally:** +- **CellDINO encoder IS local**: `ops_model/models/cell_dino.py` → `CellDinoModel` (channel-adaptive + DINO ViT-L/16, ckpt `channel_adaptive_dino_vitl16_pretrain_cells-…pth`, resize 224 + per-image + z-score, `in_channels=1`). Plus the precomputed CellDINO feature dumps (below). → image→embedding + for *generated* cells is available locally; no encoder request needed. +- **MixedChannelClassifier** code (`train_set_classifier.py`) + inference (`export_pma_attention.py`). +- Per-gene **embedding dumps WITH cell_metadata**: `v4/{train,val}_ops_zstdcontrol_cdino_v2/` + (+ `metadata.pt` w/ gene_to_idx, channel_to_idx). Metadata → zarr crop mapping works + (`_load_cell_crop`: experiment/well/x_pheno/y_pheno/segmentation_id/zarr_channel_index). +- Attention outputs (`per_experiment_v4_attn.parquet`, `pma_phase_cells_*`). +- katamari clone on branch `main` (esmc_paper commit) — likely NOT Alex's attn-classifier branch. + +## Requests for Alex L. (to scale beyond the phase·geneKO PoC) +**Note:** the B/C HSPA5 phase prototype needs NOTHING from Alex — phase attention rankings, crops, +and the CellDINO encoder are all local. +1. **Trained MixedChannelClassifier checkpoints** for the 4 models (phase·geneKO, phase·complex/EBI, + fluor·geneKO, fluor·complex) — the actual `.pt` files or the **wandb artifact IDs** + (`alex-lin/cellstate-set-classifier/model-…`). Needed for option A, and to generate attention + rankings for the models we don't already have rankings for. +2. **Fluorescent + complex (EBI/CHAD) embedding dumps + label maps** — only phase·geneKO dumps + (`*_cdino_v2`) confirmed local. Need the fluor dumps and the complex `label_map_path`. +3. **If pursuing option A:** how to score a single generated cell inside a real reference bag (his + concern: one synthetic cell may not move the bag logit). +4. **The right git branch** to check out for the latest attn-classifier code (clone is on `main`). + +--- + +## Build log + +### 2026-06-16 — classifier B/C package built +Package: `ops_model/models/attention/diffex/classifier/` (config, data, models, +celldino_features, train, run, submit, README). Locked params: binary HSPA5-vs-rest, +negatives = other genes' top-5 (distinct), 1000/class, 160×160 phase crops (no mask), +3-way train/val/test split grouped by experiment (val=selection, test=clean reported AUROC; +train+val AUROC logged per epoch to watch over/under-fit). Outputs under +`/hpc/projects/icd.fast.ops/models/diffex//`. +- **Decision:** option C reuses the **local** CellDINO encoder (`cell_dino.py`) on the SAME + crops B uses (cached) — no dump-join; B and C see identical cells. +- **Verified:** full B pipeline end-to-end on CPU (tiny config) — filtered parquet read (no OOM), + store resolution, crops materialized non-degenerate, train→AUROC→artifacts. SLURM submitter + dry-run OK (2 GPU jobs). (Crop cache key includes mask state so masked/unmasked don't collide.) +- **Next (GPU):** `python -m ops_model.models.attention.diffex.classifier.submit --gene HSPA5` + → compare B vs C held-out AUROC, pick the DiffEx target. C needs GPU (CellDINO ViT-L). + +### 2026-06-16 — HSPA5 PoC results (job 34280479, experiment-grouped split, 1000/class) +| model | test AUROC | val | train@best | +|---|---|---|---| +| B (ResNet on crops) | 0.80 | 0.85 | 1.00 (overfits) | +| **C (CellDINO+MLP)** | **0.96** | 0.95 | 0.999 | +- PoC validated: HSPA5 top-attention cells are cleanly + cross-experiment classifiable → real DiffEx target. +- **C is the chosen DiffEx target** (far more sensitive/generalizing; B memorizes). C scores generated + counterfactuals via the local `cell_dino.py` encoder. +- Outputs: `/hpc/projects/icd.fast.ops/models/diffex/HSPA5/`. +- **Next:** (a) sanity-check a diffuse gene (e.g. RPL10, top-800 penetrance) to confirm the approach + holds across penetrance; (b) begin the DiffAE stage (§2) with C as the classifier. + +### 2026-06-16 — 98 EBI complexes + NTC sweep (model C, jobs 34285052 + 34285381) +All 99 bins ran. Per-class test AUROC (distinct vs pooled top cells of other perturbations), +experiment-grouped split. Range 0.748–0.958, median 0.860. +- Most distinct: Chaperonin-containing T-complex 0.96, DNA pol α:primase 0.95, eIF4F 0.95, + replication fork protection 0.94, COPI 0.93. +- Least distinct: RNA Pol II 0.75, U1 snRNP 0.75, SRP 0.76, NSL HAT 0.78. +- **NTC = 0.86 (mid-pack) is EXPECTED, not a confound** — NTC lacks the CRISPR cut, so it's a real + distinct (no-DSB) state. (Earlier "red flag" retracted.) +- Artifacts: `…/diffex/complex/auroc_ranking_C.csv`, `auroc_hist_C.png`. +- Takeaway: every complex's top-attention cells carry a separable phenotype → strong DiffEx targets + across the board; clear biologically-sensible ranking. + +### 2026-06-16 — DiffAE stage (§2) built + PoC launched (job 34292127) +Package: `diffex/diffae/` (config, data, model, train, recon, run, submit). Faithful DiffAE: +ResNet18 semantic encoder → `z_sem` (512), conditional `diffusers.UNet2DModel` decoder with +`z_sem` injected via `class_embed_type="identity"` (→ time embedding), trained jointly with DDPM +denoising loss. Gate = DDIM-invert→reverse reconstruction (PSNR + montage; uses DDIMInverseScheduler). +- Decisions (locked): broad training set (all classes incl NTC, all ranks, ~50k crops), 160×160 + phase, per-image z-score/3 normalization, PoC-first. +- Reuses the classifier crop pipeline (`materialize_crops`). Verified end-to-end on CPU (synthetic): + conditioning, DDPM step, DDIM recon, checkpoints all work. +- PoC launched: 50k crops, 80 epochs, batch 32, 1 GPU, 720min. Outputs → `…/diffex/diffae/phase_v1/`. +- **Gate to watch:** reconstruction PSNR (recon montages every 10 epochs). If it reconstructs cells + faithfully → proceed to §3 (contrastive direction discovery). If not → fix before directions. +- **Next stage (§3):** contrastive direction discovery on `z_sem` + the option-C classifier score. + +### 2026-06-17 — DiffAE v1 result + ARCHITECTURE SWITCH to Alex's design +- v1 run (job 34298429, jointly-trained encoder): trained healthily, **reconstruction PSNR ~33–34 dB, + visually faithful** (gate passed), converged ~ep9, but hit the 12h wall at ep37/80 (too slow/many + epochs). First quota failure (34292127) fixed by freeing disk + wrapping recon writes in try/except. +- **Alex (EvolutionaryScale) already has a working DiffEx** (Notion: Imaging AI). Key design: condition + the image-DDPM on the **FROZEN backbone embedding** (not a learned encoder); discover K direction + MLPs **unsupervised** (InfoNCE + VICReg); rank directions **post-hoc** with a logistic-regression + classifier (control vs KD); traverse α∈[−3,+3], DDIM-sample per step; verify by monotonic score. +- **SWITCHED to Alex's design** (job 34312003): DiffAE now conditions on the **frozen CellDINO + embedding** (dropped the learned encoder; `cond_proj` injects it into the UNet time-embedding). + Generator, option-C classifier, and SetTransformer now all live in the SAME CellDINO space → + Stage-3 directions are discovered there and ranked by option-C directly. epochs 20, reuses crop + cache + new CellDINO-embedding cache. +- **§3 plan (Alex's recipe):** K direction MLPs (InfoNCE+VICReg, unsupervised) on CellDINO embeddings + → rank by option-C classifier score shift → α-traversal → DDIM-sample images → verify monotonicity. + +### 2026-06-18 — Stage 3 built + first HSPA5 traversal (job 34385092) +Package `diffex/directions/` (config, model=DirectionBank, losses=InfoNCE+decorrelation, +train_directions, rank=LR score-shift, data=gather target+NTC, traverse=DDIM invert→reverse + +re-encode verify, run, submit). Verified 2a+2b on synthetic; full pipeline ran on GPU in 5 min. +- HSPA5: LR acc 0.999, selected direction #6 (shift 1.31), **6/8 traversals monotonic**, mean + score Δ +0.71 (correct sign). **Full DiffEx machine works end-to-end.** +- **BUT effect is weak**: re-encoded scores stay on NTC side (−6..−14), visual change subtle. + Cause: unit direction × α≤3 ≪ the real control→KD gap (clusters far apart, LR acc 0.999); plus + x_T anchoring from DDIM inversion. +- **Fix (next): scale α to ‖mean(KD)−mean(NTC)‖** (likely ~10–30, not 3); optionally reduce x_T + anchoring; train DiffAE longer. Outputs: `…/diffex/directions/geneKO/HSPA5/`. + +### 2026-06-18 — Stage 3 retry: gap-scaled α + Δ-pixel heatmap overlay (job 34386543) +- Added: heatmap overlay (Δpixels vs α=0, red/blue) on the traversal montage; α now in units of + the control→KD gap (Δ=9.64). 6/8→1/8 monotonic, score Δ 0.71→0.28 — **gap-scaled α OVERSHOT**. +- **Diagnosis from montage:** at large α the edits land on crop BORDERS/background, not the cell + → embedding goes off-manifold, DiffAE renders boundary artifacts. Two root issues: + (a) crops are **unmasked** → direction may exploit context (confluency/neighbors), not the cell; + (b) DiffAE **edit-fidelity** limited (reconstructs well but doesn't render edits onto the cell). +- **Next options:** α-magnitude sweep for the on-manifold regime (~0.25–0.5×gap); try masked crops; + train DiffAE longer / stronger conditioning; reduce x_T anchoring. Pipeline is correct; counterfactual + QUALITY needs iteration (the hard part of DiffEx). + +### 2026-06-18 — BUG FOUND & FIXED: generate-from-noise (job 34388244) +- **Bug:** traversal DDIM-INVERTED the real cell → x_T, which encodes the image and overrides the + embedding → editing the embedding barely moved the picture (and recon was a too-good 34 dB). Alex's + spec: α=0 is the DDPM *reconstruction*, i.e. **generate from noise conditioned on the embedding** (no + inversion). Switched to fixed random noise per cell (constant across α), conditioned on z0+α·d. +- **Result:** mean re-encoded score Δ **0.28 → 5.3**, 6/8 monotonic, correct sign — embedding now drives + generation, direction validated. BUT the visual morph is still **subtle** (CellDINO registers texture + the eye misses; HSPA5 phase phenotype may be genuinely subtle). Δ-pixel heatmap localizes the + changing cell regions = a useful interpretability output on its own. +- **Next levers:** push α to 2–3×gap; strengthen DiffAE (classifier-free guidance / longer); test a + gross-morphology target (complex) to tell if subtlety is biology vs method. + +### 2026-06-18 — obvious targets (TOMM20, TIMM23, Arp2/3): diagnosis = METHOD-limited +- TOMM20 (job 34392814): lr 0.96, gap 6.2, score Δ −3.1, 7/8 monotonic. TIMM23 (34392821): lr 0.99, + Δ −3.1, 7/8. (sign arbitrary per Alex.) Arp2/3 complex: filename bug (target "2/3" has a slash → + unslugified PNG name) — FIXED (slugify filenames in traverse._plot); re-run 34392960. +- **Key finding:** TOMM20 (obvious mito phenotype) morphs just as SUBTLY as HSPA5, edge-concentrated + Δ. → the visual subtlety is **METHOD-limited, not biology**. +- **Root cause:** DiffAE under-utilizes the embedding — the noise latent dominates DDIM generation + (inverted OR random), so embedding edits weakly change pixels even though CellDINO/classifier + register them (score moves, monotonic). +- **Fix: classifier-free / edit guidance** at sampling: ε̃ = ε(c0) + w·(ε(c_α) − ε(c0)), w≈3–5 (no + retrain). If insufficient → retrain DiffAE with conditioning dropout for proper CFG. + +### 2026-06-18 — edit-guidance w-sweep (TOMM20, job 34393392): INSUFFICIENT → must retrain DiffAE +- w=1/3/5 score Δ = 0.55/1.83/2.23 (guidance amplifies) BUT monotonic 0.38/0.25/0.12 (degrades), + and the **cell still does not transform by eye even at w=5** (edge-concentrated Δ only). +- **Conclusion:** sampling-time guidance cannot fix an under-conditioned model. The DiffAE generates + from the NOISE latent and only weakly uses the embedding (why recon hit a too-good 34 dB). +- **SOLID FIX (next): retrain DiffAE with conditioning dropout** (~10–20%, learned null embedding) → + forces embedding use + enables true CFG ε̃=ε(∅)+w(ε(c)−ε(∅)). If still weak → cross-attention + conditioning (spatial) instead of global FiLM. +- Also fix: (a) direction discovery is run-to-run unstable (fix seed + more epochs); (b) replace the + recon gate with an UNCONDITIONAL-generation check (null-embedding samples should look generic; + conditional should match target) — recon PSNR was misleading. + +### 2026-06-18 — conditioning diagnostic = DEAD (0.008), then proper DiffAE rebuild (job 34394595) +- **Diagnostic** (`diffae/diagnose_conditioning.py`, job 34393969): same fixed noise under + null/ctrl-centroid/KD-centroid. MSE(ctrl-vs-kd)=0.0010, MSE(noise-vs-noise)=0.133 → + **emb/noise ratio = 0.008**. The embedding has <1% control; the DiffAE generates from noise and + ignores the embedding. (null-vs-ctrl 0.042 ≫ ctrl-vs-kd 0.001 → reacts to embedding *presence*, + not *content*.) Confirms: cells don't change because conditioning is ~dead. +- **Rebuild** (`diffae/train.py` rewritten): conditioning dropout (0.15, learned null_emb) + EMA + (0.9995) + resume-across-jobs + deeper cond MLP. **Gate = conditioning ratio** (not recon PSNR), + logged every 5 epochs; EMA-best saved by ratio. Target: ratio ≫ 0.008 (→ ~0.3+). +- Retrain 34394595 (phase_v1, reuses caches, 120 epochs, batch 48, resume). Watching cond_ratio + trajectory. **After it works:** switch directions/traverse to true CFG ε(∅)+w(ε(c)−ε(∅)). +- If time-embedding conditioning still can't climb → escalate to cross-attention (UNet2DConditionModel). + +### 2026-06-29 — Plan C implemented (deterministic direction) + v2_aug retrain +- **Reproducibility root cause:** unsupervised InfoNCE direction bank is GPU/seed-nondeterministic + and `best_k = argmax|shift|` flips run-to-run → same cell highlighted different regions. +- **Fix (plan C, implemented):** `directions/config.py` `direction_method` — default **`mean_diff`** + (deterministic control→KD centroid vector; also `lr_weight`) as PRIMARY; the paper's unsupervised + bank kept as `direction_method="unsupervised"` secondary track. `traverse(fixed_dir=…)` uses the + global deterministic direction; `deterministic=True` sets seeds + cuDNN-deterministic. `+α = toward + KO` by construction (no more sign flip). `rank.supervised_direction()` computes it; LR kept for the + re-encode score check only. +- Reproducibility proof in progress: TIMM23 run twice (jobs 34654037/34654040) → pixel-diff strips. +- **Next model (v2_aug)**: orientation-aug DiffAE retraining (job 34651296), cond_ratio climbing + 0.04→0.14 @ep19/120 (aug ramps slower); resume to convergence, then switch directions default to it. + +### Active experiments (2026-07-03) +- **v1 remains the best model.** v2 (dihedral) and v3 (continuous rot+scale) augmentation did NOT + beat v1 — not more orientation-stable, weaker/less-convincing phenotypes; cond_ratio ceiling falls + with aug (v1 0.47 → v2 0.25 → v3 0.20; curves in `coding_exps/diffex/diffae_training_curves.png`). + Flow-matching transport (CellFlow-style, `directions/flow.py`) also explored → smoother but less + clean phenotypes, noisy negative extreme → NOT adopted. Reverted default to v1 + mean-diff α. +- **Generator data-scaling test (RUNNING):** does 50k→**500k** crops help? Two no-aug chains, 24 ep: + - scratch `phase_v1_500k` — jobs `34667092→34667174→34667175` + - warm-start from v1 `phase_v1_500k_warm` — jobs `34667176→34667177→34667178` + - Compare cond_ratio/loss vs v1 (0.47) + visual morphs. mem_gb=400 (500k float32 crops ≈ 51GB each). +- **Direction depth test (pending):** gather 1k→**~12k**/class (the distinctiveness peak) for a tighter + mean-diff centroid — cheap, per-target (~30-40 min CellDINO/target), no retraining. + +### Future direction — per-fluorescent-marker models (2026-07-03) +Reproduce the best phase pipeline **per fluorescent marker** (~60 live-cell markers) — a per-marker +counterfactual view of each gene-KO / complex phenotype in the channel where attention is most +informative. **~60 models** (one DiffAE + direction set per marker). +- **Attention source EXISTS:** `…/alex_lin_attention/v4/pma_fluorescent_cells_all.csv` + (+ `pma_fluorescent_cells_ebi_all.csv` for complexes) — the fluor analog of `pma_phase_cells_v2_all.parquet`. + CellDINO fluor train/val sets also present (`train/val_ops_zstdcontrol_cdino_fluorescent`). +- **Scope:** `good_experiment_list_v2.yml` (87 exps; fluor channels GFP×74, mCherry×23, Cy5×2). Each + experiment's channel→biological-marker label is in `ops_process/ops_analysis/configs/ops_channel_maps.yaml`. + Marker/experiment enumeration tooling: `ops_utils/data/feature_discovery.py`, + `ops_utils/analysis/embedding_discovery.py`. +- **Per-marker pieces:** gather top-attention fluor cells (control + KD) → CellDINO embed the MARKER + channel → mean-diff direction → **per-marker DiffAE** (phase generator can't decode fluor) → traverse. + Direction/traverse code unchanged; needs a per-marker DiffAE + a marker→(experiments, channel) map. +- **Complication (deferred):** the v2 list is **live-cell fluor only** — 4i / Cell-Painting (fixed-cell) + channels are excluded. If added later they need their **own per-round link CSVs** (`link_csv_dir` in + `ops_model/data/data_loader.py`), not the default live 3-assembly link. +- **NEEDS DESIGN CONFIRM before building** (60 DiffAE trainings is a large program). + +### Future direction — attention-informed cell selection (2026-06-29) +Currently we take a flat top-1000 attention-ranked cells per class and pick traversal/feature +cells by index. To exploit attention ranking more (only touches `directions/data.py gather` + +`rank.supervised_direction`, not the DiffAE): +- **Pick highest-attention cells for the traversal/featured strips** (most representative morph), + not an arbitrary cell index. +- **Per-target penetrance depth** — use the attention-accuracy knee (HSPA5≈top10, RPL10≈top800) + instead of a flat top-1000, so sharp phenotypes aren't diluted by the diffuse tail. +- **Attention-weight** the mean_diff / classifier so the most-phenotypic cells dominate the axis. + +### Future direction — orientation-invariance via augmentation (2026-06-29) +Observation: along a traversal the cell often spuriously rotates/transposes (orientation is +encoded in the CellDINO embedding, so the discovered direction carries an orientation component the +DiffAE renders). Fix: during DiffAE training, augment the TARGET image with the dihedral group +(4 rotations × flip = 8 views, incl. transpose) while conditioning on the embedding of the CANONICAL +(un-augmented) cell. Teaches the model orientation ≠ embedding-determined → orientation absorbed by +the (fixed) noise latent, phenotype carried by the embedding → traversals stop rotating. Do NOT +recompute CellDINO on the augmented crop (defeats the decoupling). Also serves as general aug to +sharpen conditioning. + +### 2026-07-08 — Attention-head tab + viewer reorientation (BUILT) +**Status: built + deployed to `viewer_assets/`.** `viewer/build_attention_heads.py` rendered **984/1000** +phase geneKO genes (16 npz still corrupt/mid-write by Kevin — `BTF3L4, BUB1B, DAD1, DHRS9, FECH, +FOXD4L1, GTPBP4, INO80D, MTOR, NCBP2, NRAS, POLR2F, RPS19BP1, TWF1, TYK2, YIPF5` — builder is idempotent, +skips-loud, re-run picks them up). `global_max=2.44`. Webapp reoriented (`webapp/{index.html,app.js, +style.css}` v36, copied to `viewer_assets/`): persistent `#browse` block (marker+grain+perturbation+ +cells/page) drives 3 view tabs — Traversal / Embedding / **Attention heads**. Attn view = inferno LUT + +live per-map/per-gene/fixed normalization + opacity, head dropdown from `heads.json`; greys out for +non-phase / complex / missing-gene. Embedding now rings + pans to the selection. **Availability decoupled +from manifest** — app fetches `attention_heads/phase/index.json` (no `precompute.build_manifest` change). +- **ALL 4 modality×grain combos now IN (2026-07-08 pm).** Kevin dumped pixel_attribution (maps+crops+ + patch_masks) for fluor geneKO (`fluorescence_pixel_attribution//`), phase complex + fluor complex + (`complex_pixel_attribution/{phase,fluorescence/}/`) — same npz schema as phase. Builder rewritten + **SLURM-parallel** (`build_attention_heads.py render` → `submit_parallel_jobs`, 40 shards, ~1 min) into a + uniform `attention_heads////` layout + single `attention_heads/index.json` + ({global_max, assets:{modality:{grain:[keys]}}}). **23 modalities, 1455 keys** (phase geneKO 984 + phase + complex 93 + 16 fluor markers). Webapp resolves assets by (marker→`jsSlug(marker_channel)`|"phase", grain, + target→gene|slug); non-phase-geneKO lack ranking metrics (auroc/spec) but render fine from the npz `heads`. + (`fluorescence_attention/.npz` = the older ranking-features-only dump, superseded.) +- **16 corrupt phase genes: left as-is.** Full-size but bad-zip at SOURCE — needs Kevin to regenerate; 984/1000. +- **Attn viewer UX (2026-07-08 pm):** top-crossbar selection (marker+perturbation comboboxes w/ search, mAP|A–Z + sort) drives Traversal/Embedding/Attention-heads tabs; attn view = per-perturbation color-coded blocks + (rows=heads, cols=cells), pin/reset controls, per-cell/gene/fixed norm + dual clim + opacity sliders, + Ritvik-faithful overlay (σ=2 smooth + cell mask + inferno@α0.6); embedding click → selects in search box. + +### 2026-07-08 — Attention-head tab + viewer reorientation (design) +New viewer view: overlay CellDINO **attention-head pixel weights** (inferno) on the **real phenotype +cells** (`viewer/phenotype_cells.py` output), so you can see WHERE in each cell each top attention head +looks — the classifier's spatial evidence, alongside the generative counterfactual morph. + +**Data (Kevin L., already under `viewer_assets/attention_heads/`, verified 2026-07-08):** +- `phase/pixel_attribution_cache/.npz` (1000 geneKO genes): `maps (20,6,128,128) f16` ∈ [0,~0.34] + (20 cells × top-6 heads × pixel attribution), `crops (20,128,128) f32` (z-scored cell crops), + `heads (6,2) int32` = the ranked (layer,head) pairs, `patch_masks (20,196) bool` (14×14 ViT patches). +- `phase/head_rankings_per_gene.json`: per-gene ranked heads + metrics (`layer,head,feature,spec_p10, + spec_min,auroc_vs_ntc`); order matches the npz `heads` array. **The 20 cells ARE the phase·geneKO + top-20 phenotype cells** from `phenotype_cells.py` (same selection). +- `celldino_attention_head_analysis/fluorescence_attention/.npz` — per-marker fluor analog + (structure TBD) → follow-on. **v1 scope = phase · geneKO only** (no complexes: head_rankings is gene-keyed). + +**Precompute (`viewer/build_attention_heads.py`, to write):** ship **raw** data so normalization + inferno ++ opacity are LIVE display options (user-selectable, per decision). Per gene → per cell: write +`cell/crop.webp` (grayscale, per-crop robust min-max) + per head `cell/head.webp` (grayscale +attribution scaled by a FIXED global max so absolute intensity is preserved). Per-gene `heads.json` = +ranked-head metrics (`layer,head,feature,spec_p10,spec_min,auroc_vs_ntc`) + `n_cells` + `gene_max` + +`global_max`. The webapp applies a 256-entry **inferno LUT** in a canvas and composites over the crop — +so a Display dropdown offers **per-map / per-gene / fixed** normalization live (per-map = rescale by the +loaded tile's own max; per-gene = by `gene_max`; fixed = by `global_max`), plus an opacity slider, without +re-fetching. Count ≈ 1000×20×(6+1) ≈ 140k grayscale WebP (traversal already ~370k). Keeps the app +dependency-free; no 3×-image blowup from baking each norm. + +**Viewer reorientation (`webapp/index.html` + `app.js`):** today the left panel is 3 *control* tabs +(Browse / Anchor / Embedding), each with its own selectors, and the Embedding montage ignores the +Browse (marker,perturbation) selection. **Reorient:** make Browse (marker + grain + perturbation) a +PERSISTENT selector block = single source of truth (`state.marker`, `state.target`); below it a **View +switcher** — Traversal | Embedding | Attention heads — each rendering the main stage for the CURRENT +selection. Fold today's Anchor/display controls under Traversal; α/cell/embedding-mode under Embedding; +head selector + overlay-opacity under Attention heads. Embedding also gains browse→highlight/pan and +keeps montage-click→browse select (closes the selection loop). + +**Attention-head view UI:** grid of the 20 phenotype crops with the selected head's inferno overlay; +head dropdown lists the 6 ranked heads with `(L,H) · AUROC·NTC / spec`; overlay-opacity slider; raw-crop +toggle. Reuses Browse's cells-per-page paging. + +**Manifest:** add per-(marker,target) `attn` availability + `n_heads` so the View switcher greys out +Attention heads where absent (v1: present only for phase geneKO genes with an npz). + +**Decisions (locked 2026-07-08):** (a) normalization = **live user option** in Display settings +(per-map / per-gene / fixed) via the grayscale+LUT approach above; (b) full reorientation approved +(persistent Browse + view switcher); (c) **phase-only v1**, fluor (Kevin's per-marker npz) is a follow-on. + +### 2026-07-08 — phenotype-cell handoff CSV + v2 mAP + EBI matrix (see dashboard LATEST) +- **`viewer/phenotype_cells.py`** — the Ritvik handoff: top-20 REAL phenotype cells per (marker × perturbation) + → `viewer_assets/phenotype_cells_for_attention.csv` (160,420 cells / 53 markers). Key correctness fix: + the pma `rank` is GLOBAL per geneKO (across all channels), so `_csv_top` re-ranks WITHIN each + (channel, perturbation) and takes each marker's own top-20 by attention (two-pass chunked `head`). + Fluor filtered by mAP ≥ 0.2; phase = ALL. Added `map_score`, `geneKO`, `ebi_complex`, `rank_source` cols. +- **v2 mAP:** `catalog.dist_matrix` → `paper_v2/with_cp/with_4i/all_livecell` (56 reporters, live+CP+4i); + added 4i `FIXED_REP`; 52/56 pma channels map (NFkB/RSP6/Rb/gH2AX expected-excluded). +- **`viewer/build_complex_ebi_map.py`** — complex×reporter EBI mAP (98×56) via copairs + `phenotypic_consistency_ebi`, all-perturbation, on v2 per_signal; also emitted by the aggregation pipeline. +- **Pending:** centroid fallback for cisGolgi/VIM/LMNB1 (mAP present, no pma cells) from + `cell_dino_features_v2/features_processed_.h5ad` (embedding + crop metadata); SetTransformer + bag scoring parked on Alex's v2 CellDINO extraction. + +### 2026-07-09 — Argus app staging (S3 hosting) +- **Infra PR #51 OPENED + ACCEPTED/MERGED** (`sfbiohub-infra`, branch `diffex-viewer-dev`, + `terraform/accounts/biohub-nonprod/diffex-viewer-dev.tf`): S3 bucket `diffex-viewer-dev`, read-only IRSA + role `biohub-nonprod-diffex-viewer` (trusts SA `diffex-viewer` in ns `argus-diffex-viewer-rdev`), read-write + uploader role `diffex-viewer-dev-readwrite`. 1 TB ceiling. → `terraform apply` provisions bucket+roles. +- **App scaffold built** at `/hpc/mydata/gav.sturm/diffex-viewer` — forked from `czbiohub-sf/mops-viewer` + (same Argus+S3 pattern; names already align with the `.tf`). Serving layer swapped Gradio → **static nginx**: + - `Dockerfile` (nginx, bakes `webapp/` shell) + `nginx.conf` (port 8080, `/healthz`) + `docker-entrypoint-diffex.sh` + (drops shell into the S3-populated webroot, starts nginx). + - `.infra/common.yaml`: aws-cli `fetch-assets` init container (`aws s3 cp s3://diffex-viewer-dev/ → /usr/share/nginx/html/`), + `web` emptyDir, serviceAccount `diffex-viewer`, OIDC proxy; `.infra/rdev/values.yaml` carries the read-only role ARN. + - `.argus-ci.yaml` app=`diffex-viewer`; `scripts/uploader/*` sync `viewer_assets/` (manifest.json at bucket ROOT) via the readwrite role; `DEPLOY.md` runbook. + - Static because the webapp reads assets by relative path from `manifest.json`; the S3 sync drops assets as siblings of the baked shell. +- **Remaining to deploy:** `terraform apply` → create `czbiohub-sf/diffex-viewer` repo + push → `argus register app` + (team-sci-biohub) + bootstrap/reconcile `.github` (needs argus CLI, interactive) → upload `viewer_assets/` (our + role, or **Kyle from HPC** — he offered) → PR + `stack` label → Argus builds + deploys rdev behind Okta. Confirm + the registered namespace/SA matches `argus-diffex-viewer-rdev`/`diffex-viewer` before relying on the IRSA trust. + +### 2026-07-11 — full 1k-geneKO + complex buildout LAUNCHED +- **geneKO (master `34826112`, 49 jobs):** `submit seed --map-thr 0 --timeout 720` — all ~1000 geneKO + genes/marker for the 46 valid-`rep` fluor markers (49,000 targets). Resume skips the ~100–194 already + seeded per marker; the 4 hub markers (ChromaLIVE_561 957, LysoTracker 932, NPM3 930, NucleoLIVE 780) + were already near-complete. Runs with **no concurrency cap**. +- **complex (master `34826156`, 50 jobs):** `submit fluor-complex` resume — nearly all 98 EBI complexes + were already built for ~45 markers; only 5 partial (peroxisome_Peroxi 35, pS6 35, pRb 59, NPM3 86, SRRM2 88). +- **`submit.py` fixes made for this run:** + - **BUG:** `cmd_seed` used `if args.map_thr:` — `0` is falsy, so `--map-thr 0` silently fell back to top-8 + (a wrong launch, 34826071, was cancelled). Changed all gates to `args.map_thr is not None` → `--map-thr 0` + now correctly means all ~1000 genes. + - **`--parallel` now defaults to `None` = no concurrency cap** (only sets `slurm_array_parallelism` when given), + so full buildouts don't need `--parallel 100`. Added `--timeout` (default 180) to override the per-marker wall. +- **After the builds land:** `submit sync` to refresh manifest + attention + montages from the new cache. +- The 4 rep=None markers (NFkB/RSP6/Rb/gH2AX) remain intentionally excluded (no v2 distinctiveness reporter column). + +### 2026-07-12 — DiffAE cond_ratio PEAKS EARLY then declines; check before extending/rebuilding +Resumed the 6 under-trained (ep55) fluor DiffAEs to ep120 (dihedral). Key lesson: **most peaked their +conditioning ratio around ep39–55 and then DECLINED** — extending to 120 did NOT improve them: +| marker | best cond_ratio | peaked @ | verdict | +|---|---|---|---| +| CLTA | 0.390 | ep54 | plateaued/declined → no gain | +| ATP1B3 | 0.171 | ep54 | plateaued/declined → no gain | +| PSMB7 | 0.830 | ep54 | declined to ~0.49 → **stopped** | +| TFRC | 0.320 | ep39 | declined to ~0.21 → **stopped** | +| VAMP3 | 0.231 | ep54 | declined to ~0.17 → **stopped** | +| **SLC3A2** | **0.277** | **ep109** | still climbing (>pre55 0.259) → **kept running** | + +**RULES (to not repeat the wasted compute):** +1. `diffae_best.pt` is saved on BEST cond_ratio, so it already captures the peak regardless of final epoch — + a longer run does NOT give a better checkpoint unless best_ratio actually advanced. +2. **Before extending training** past its current point, check the cond_ratio trajectory + (`torch.load(train_state)['history']`). If it peaked early and is declining, stop — the best is already banked. +3. **Before clearing + rebuilding a marker's traversals** (1k geneKO / 98 complex / anchors), confirm the model + IMPROVED: compare `diffae_best.pt` mtime + best_ratio vs the value the existing traversals were built on. + Only rebuild if the best genuinely advanced PAST what the current traversals used. +- **Rebuild status:** CLTA/ATP1B3/PSMB7/TFRC/VAMP3 = NO rebuild (best@≤ep54 already used by the Jul-11 traversals). + **SLC3A2 = the one rebuild candidate** — its best advanced to 0.277@ep109 (Jul 12) vs the Jul-11 traversal + checkpoint (~0.259); once it finishes, clear + rebuild ONLY SLC3A2's traversals with the improved model. +### 2026-07-12 — accuracy-selected cell variant (Kevin's accuracy_ranking CSVs) +Cell selection can now use **classifier-accuracy rank** instead of **attention rank**. Source = +`…/alex_lin_attention/v4/accuracy_ranking/`: `pergene_phase_cell_rankings.csv` (geneKO·phase), +`ebi_pergene_phase_cell_rankings.csv` (complex·phase), `ebi_class_channel_cell_rankings.csv` (complex·fluor, +55 marker channels). NTC has NO accuracy data → NTC anchor always stays attention-sourced (shared/cached). +- **phase** = side-by-side A/B: modality `phase` (attention, "phase_attention") vs `phase_topacc` + ("phase_accuracy"). Full 1k geneKO + 98 complex + 182 anchors built for phase_topacc. Hooks: + `_gather_class(parquet=…)`, `precompute_marker(accuracy_parquet=, variant=)`, and the anchor path + `_gather_df`/`_setup`/`precompute_target(accuracy_parquet=, variant=)`. Accuracy dirs have LARGER + control→KD gaps (cleaner separation) than attention. +- **fluor** = REPLACED IN PLACE (no `_acc` duplicate), via `precompute_marker(accuracy_fluor_csv=, force=True)`. + **⚠️ COVERAGE SPLIT:** the fluor accuracy CSV only covers the complexes each marker actually distinguishes + (a SUBSET of the 98 — **median 13/marker, range 1–32, 53 distinct total**), so each fluor marker's `complex/` set is now MIXED: + accuracy-selected for its covered complexes, still attention-selected for the rest. This is intentional/interim + — we will get full 98-complex + 1k-geneKO accuracy coverage for every marker later and **rebuild the entire + cache anyway**, at which point the split disappears. geneKO fluor has NO accuracy CSV yet (phase-only + complex-fluor). +- **fluor complex→complex ANCHORS: NET-NEW with accuracy** (didn't exist in attention). Per marker: top-5 + accuracy-covered complexes (by `class_channel_acc`) → 20 A→B pairs, accuracy cells for both, via per-channel + parquets in `accuracy_ranking/fluor_complex_by_channel/`. 54 markers × ~20 ≈ 982 pairs → `/complex/`. +- **PHASE SWAP DONE (2026-07-12): accuracy is now the canonical `phase`.** `viewer_assets/phase` (attention) + + its `_directions/phase` archived to `viewer_assets_backup/{phase_attention,_directions_phase_attention}` (same FS, + reversible); `phase_topacc` → `phase`. Manifest label reverted to plain "Phase". NOTE: the separate build-cache + `…/diffex/directions/{phase,phase_topacc}/` (anchor gather cache, OUTSIDE viewer_assets) was NOT renamed — only the + served `viewer_assets` traversals + `_directions` were swapped; a future full rebuild regenerates it anyway. + +- **min-ep gate REMOVED as default** (`submit seed --min-ep` default 98→**0**; `catalog.complete_markers` default + 98→**0**). Epoch count is NOT a quality signal — a marker peaking at ep54 is as usable as one at ep120, and + `diffae_best.pt` banks the peak. Inclusion now gates on **checkpoint presence** (diffae_best.pt + train_state), + not epoch. Pass `--min-ep N` only to re-impose a floor. (Right now all 56 markers pass either way since the + Jul-11 resume pushed the 6 past ep98, but the default now won't silently drop a future under-trained marker.) diff --git a/src/ops_model/models/attention/diffex/README.md b/src/ops_model/models/attention/diffex/README.md new file mode 100644 index 0000000..1f72e6c --- /dev/null +++ b/src/ops_model/models/attention/diffex/README.md @@ -0,0 +1,40 @@ +# DiffEx — counterfactual interpretability for the attention atlas + +Explain geneKO / protein-complex phenotypes **in image space**: generate +counterfactual single-cell morphs ("if this control cell were a KD, what would it +look like?") and the per-pixel change map, instead of relying on OP/CP features. +Adapted from DiffEx (arXiv:2502.09663) and Alex Lin's EvolutionaryScale pipeline, +working in the **CellDINO embedding space** that the SetTransformer already uses. + +See [PLAN.md](PLAN.md) for the design rationale and the full running log. + +## Pipeline (three stages, each a subpackage) + +| stage | package | what it does | +|---|---|---| +| 1 | [`classifier/`](classifier/) | per-class single-cell classifier on **top-attention cells** — the model whose decision DiffEx explains / that ranks directions. B = ResNet on phase crops; **C = MLP on CellDINO features** (chosen). | +| 2 | [`diffae/`](diffae/) | **conditional diffusion** generator: UNet that generates a cell image conditioned on its CellDINO embedding (conditioning dropout + EMA + CFG). | +| 3 | [`directions/`](directions/) | **contrastive direction discovery** (InfoNCE + decorrelation, unsupervised) → rank directions by a control-vs-target classifier → **CFG traversal** α∈[−,+] → DDIM-sample a counterfactual strip + Δ-pixel heatmap, verified by re-encoded score. | + +## Run order (each stage has `run.py` for local + `submit.py` for SLURM) + +```bash +# Stage 1 — classifier (per gene/complex, or sweep --all-classes) +python -m ops_model.models.attention.diffex.classifier.submit --grain complex --all-classes --models C +python -m ops_model.models.attention.diffex.classifier.aggregate --grain complex --model C + +# Stage 2 — train the conditional DiffAE (resume-able; gate = embedding/noise ratio) +python -m ops_model.models.attention.diffex.diffae.submit --epochs 120 --batch-size 48 +python -m ops_model.models.attention.diffex.diffae.diagnose_conditioning # conditioning-strength check + +# Stage 3 — directions + counterfactual traversal for a target +python -m ops_model.models.attention.diffex.directions.submit --grain geneKO --target HSPA5 +``` + +Outputs: `/hpc/projects/icd.fast.ops/models/diffex/{,diffae,directions}/…`. + +## Status +Stages 1 & 3 built and validated end-to-end; Stage-2 DiffAE conditioning was the +hard part — see PLAN.md (the v1 generator ignored the embedding; the rebuild with +conditioning dropout + EMA fixes it). Current focus: training the DiffAE to a +conditioning ratio high enough for visible morphs, then scaling across targets. diff --git a/src/ops_model/models/attention/diffex/classifier/README.md b/src/ops_model/models/attention/diffex/classifier/README.md new file mode 100644 index 0000000..72732f5 --- /dev/null +++ b/src/ops_model/models/attention/diffex/classifier/README.md @@ -0,0 +1,61 @@ +# DiffEx single-cell classifier PoC (options B & C) + +The classifier DiffEx will explain (see [../PLAN.md](../PLAN.md), classifier §1). +PoC = **binary HSPA5-vs-rest** on **phase** single-cell crops. + +- **B** — ResNet18 on the 160×160 phase crops (pixels → class logits). +- **C** — MLP on **CellDINO** embeddings of the *same* crops (`ops_model.models.cell_dino`). + +Both score **top-attention cells** (settled), differing only in the feature space. + +## Locked design (Gav, 2026-06-16) +- Positives: top-1000 attention cells of HSPA5 (`pma_phase_cells_v2_all.parquet`). +- Negatives: 1000 cells sampled from the **top-5** attention cells of **other genes** + (the "distinct" contrast — strong-vs-strong). +- Crop 160×160, phase-only (`Phase2D`), no cell mask (full crop context). +- Split: **3-way train/val/test, grouped by experiment** (confound guard — val & test + cells come from experiments never trained on). val = model selection; **test = the + clean reported number** (scored once, never used for selection). Stratified-random + fallback if a class is missing from a side. +- Success: held-out **test AUROC ≫ 0.5** (generalizing across experiments ⇒ biology, not batch). + +## Decision: C reuses the local CellDINO encoder on the same crops +Rather than join Alex's per-gene dumps, option C runs the local encoder +(`CellDinoModel`: channel-adaptive DINO ViT-L/16, resize 224 + per-image z-score, +`in_channels=1`) on the identical crops B uses, and caches the embeddings. One crop +pipeline; B and C see identical cells. + +## Run (GPU) +```bash +# single run, interactively on a GPU node +python -m ops_model.models.attention.diffex.classifier.run --model B --gene HSPA5 +python -m ops_model.models.attention.diffex.classifier.run --model C --gene HSPA5 + +# or submit both to SLURM (one GPU job each) +python -m ops_model.models.attention.diffex.classifier.submit --gene HSPA5 + +# sweep: all 98 EBI complexes + NTC control, model C +python -m ops_model.models.attention.diffex.classifier.submit --grain complex --all-classes --models C +# then rank them +python -m ops_model.models.attention.diffex.classifier.aggregate --grain complex --model C +``` +`--grain {geneKO,complex}` selects the parquet + class column (`gene` vs `predicted_class`). +NTC is included as a negative-control bin (its AUROC should be near chance). +Outputs land under `//` (default out-dir +`/hpc/projects/icd.fast.ops/models/diffex`): `model_{B,C}.pt`, `metrics_{B,C}.json`, +and a shared `cache/` (crops + CellDINO features). SLURM logs → +`ops_mono/slurm_logs/diffex_clf/`. + +## Layout +- `config.py` — all params (the locked defaults above). +- `data.py` — cell-table query, crop materialization (`BaseDataset`), split. +- `models.py` — ResNet (B) + MLP head (C). +- `celldino_features.py` — embed crops with the local CellDINO encoder (cached). +- `train.py` — shared train/eval loop (AUROC). +- `run.py` — orchestrator + `run_poc()` entry point. +- `submit.py` — SLURM submission (`submit_parallel_jobs`). + +## Status +Pipeline verified end-to-end on CPU for **B** (tiny config): cell table → crops +(non-degenerate, masked) → train → AUROC → artifacts. **C** needs a GPU (CellDINO). +Next: run B & C on HSPA5 at full scale (GPU), compare AUROC, pick the DiffEx target. diff --git a/src/ops_model/models/attention/diffex/classifier/__init__.py b/src/ops_model/models/attention/diffex/classifier/__init__.py new file mode 100644 index 0000000..9d9c060 --- /dev/null +++ b/src/ops_model/models/attention/diffex/classifier/__init__.py @@ -0,0 +1,4 @@ +"""DiffEx single-cell classifier PoC (options B and C). + +The classifier DiffEx will explain. See ../PLAN.md (classifier §1) and README.md. +""" diff --git a/src/ops_model/models/attention/diffex/classifier/aggregate.py b/src/ops_model/models/attention/diffex/classifier/aggregate.py new file mode 100644 index 0000000..68fc0a6 --- /dev/null +++ b/src/ops_model/models/attention/diffex/classifier/aggregate.py @@ -0,0 +1,87 @@ +"""Aggregate per-class classifier metrics into a ranked table. + + python -m ops_model.models.attention.diffex.classifier.aggregate --grain complex + +Collects every ///metrics_.json into one CSV ranked by +test AUROC (how cleanly/distinctly each class's top-attention cells classify), plus +a histogram. NTC is flagged as the negative-control reference. +""" +from __future__ import annotations + +import argparse +import json +from pathlib import Path + +import pandas as pd + +from .config import DEFAULT_OUT_ROOT + + +def collect(root: Path, model: str) -> pd.DataFrame: + rows = [] + for mp in sorted(root.glob(f"*/metrics_{model}.json")): + try: + rows.append(json.loads(mp.read_text())) + except Exception as e: # noqa: BLE001 + print(f"[skip] {mp}: {e}") + if not rows: + raise FileNotFoundError(f"no metrics_{model}.json under {root}") + df = pd.DataFrame(rows) + df["is_ntc"] = df["gene"].astype(str).eq("NTC") + return df.sort_values("test_auroc", ascending=False).reset_index(drop=True) + + +def main(): + ap = argparse.ArgumentParser(description="Aggregate classifier sweep metrics") + ap.add_argument("--grain", default="complex") + ap.add_argument("--model", default="C", choices=["B", "C"]) + ap.add_argument("--out-dir", default=DEFAULT_OUT_ROOT) + args = ap.parse_args() + + root = Path(args.out_dir) / args.grain + df = collect(root, args.model) + + cols = ["gene", "test_auroc", "val_auroc", "train_auroc_at_best", + "n_pos", "n_neg", "n_test", "best_epoch"] + cols = [c for c in cols if c in df.columns] + csv = root / f"auroc_ranking_{args.model}.csv" + df[cols + ["is_ntc"]].to_csv(csv, index=False) + + print(f"\n{len(df)} classes (model {args.model}, grain {args.grain})") + print(f"test AUROC: median={df.test_auroc.median():.3f} " + f"mean={df.test_auroc.mean():.3f} " + f"range {df.test_auroc.min():.3f}..{df.test_auroc.max():.3f}") + if df.is_ntc.any(): + ntc = df[df.is_ntc].iloc[0] + rank = int(df.index[df.is_ntc][0]) + 1 + print(f"NTC control: test_auroc={ntc.test_auroc:.3f} (rank {rank}/{len(df)} — " + f"should be near the bottom)") + print("\nTop 10 most-distinct:") + print(df.head(10)[cols].to_string(index=False)) + print("\nBottom 10 (least distinct):") + print(df.tail(10)[cols].to_string(index=False)) + + try: + import matplotlib + matplotlib.use("Agg") + matplotlib.rcParams["pdf.fonttype"] = 42 + import matplotlib.pyplot as plt + fig, ax = plt.subplots(figsize=(7, 4)) + ax.hist(df.test_auroc, bins=25, color="steelblue", edgecolor="white") + if df.is_ntc.any(): + ax.axvline(df[df.is_ntc].test_auroc.iloc[0], color="red", ls="--", + label="NTC control") + ax.legend() + ax.set_xlabel("test AUROC (distinctiveness)") + ax.set_ylabel("# classes") + ax.set_title(f"{args.grain} {args.model}: per-class classifier AUROC (n={len(df)})") + fig.tight_layout() + png = root / f"auroc_hist_{args.model}.png" + fig.savefig(png, dpi=150) + print(f"\nwrote {csv}\nwrote {png}") + except Exception as e: # noqa: BLE001 + print(f"\nwrote {csv} (histogram skipped: {e})") + + +if __name__ == "__main__": + main() diff --git a/src/ops_model/models/attention/diffex/classifier/celldino_features.py b/src/ops_model/models/attention/diffex/classifier/celldino_features.py new file mode 100644 index 0000000..9550cd6 --- /dev/null +++ b/src/ops_model/models/attention/diffex/classifier/celldino_features.py @@ -0,0 +1,39 @@ +"""Option C features: embed the SAME phase crops with the local CellDINO encoder. + +Uses `ops_model.models.cell_dino.CellDinoModel` (channel-adaptive DINO ViT-L/16, +resize 224 + per-image z-score, in_channels=1) — the same encoder that produced +Alex's feature dumps, so option C lives in the classifier's true feature space. +Requires a GPU (CellDinoModel runs on cuda). +""" +from __future__ import annotations + +import os + +import numpy as np +import torch + + +def embed_crops(images: np.ndarray, cfg, cache_path=None) -> np.ndarray: + """images (N,1,H,W) float32 -> CellDINO features (N, D). Cached to .npz.""" + if cache_path and os.path.exists(cache_path): + feats = np.load(cache_path)["features"] + print(f"[cache] celldino features <- {cache_path} {feats.shape}") + return feats + + from ops_model.models.cell_dino import CellDinoModel + + model = CellDinoModel(z_score=True) # loads checkpoint + moves to cuda + feats = [] + with torch.inference_mode(): + for i in range(0, len(images), cfg.batch_size): + x = torch.as_tensor(images[i:i + cfg.batch_size]) # (B,1,H,W) + out = model.extract_features({"data": x}) # preprocess + forward (cuda) + feats.append(out.float().cpu().numpy()) + features = np.concatenate(feats).astype(np.float32) + if features.ndim != 2: + features = features.reshape(features.shape[0], -1) + + if cache_path: + np.savez(cache_path, features=features) + print(f"[cache] celldino features -> {cache_path} {features.shape}") + return features diff --git a/src/ops_model/models/attention/diffex/classifier/config.py b/src/ops_model/models/attention/diffex/classifier/config.py new file mode 100644 index 0000000..2d53d7d --- /dev/null +++ b/src/ops_model/models/attention/diffex/classifier/config.py @@ -0,0 +1,69 @@ +"""Config for the DiffEx single-cell classifier PoC (options B and C). + +One classifier that scores top-attention cells, giving DiffEx its differentiable +image->class-k signal. PoC = binary {gene}-vs-rest on phase crops. +""" +from __future__ import annotations + +import os +from dataclasses import dataclass + +_V4 = "/hpc/projects/icd.fast.ops/models/alex_lin_attention/v4" + +# Phase per-cell ranking exports. v4 = pma_attention-ranked (masked SetTransformer). v5 (paper-v2) = +# set-accuracy-`score`-ranked, no-mask/160px, and now includes an NTC group (so NTC anchors come from the +# same ranking — no attention fallback). The slim v5 viewer-parquet (top-1000/class) is normalized to the +# v4 schema (segmentation_id→segmentation, score→pma_attention, rank_type="top"). Toggle via OPS_DIFFEX_V5=1. +# Built + served side-by-side under viewer_assets_v5 so the live v4 viewer is untouched until the final swap. +_USE_V5 = os.environ.get("OPS_DIFFEX_V5", "0") == "1" +_V5_RANK = "/hpc/projects/icd.fast.ops/models/diffex/viewer_assets_v5/_rankings" +# geneKO: class is `gene`. complex (EBI): class is `predicted_class` (complex name). +if _USE_V5: + PMA_PHASE_GENEKO = f"{_V5_RANK}/pma_v5_phase_geneKO.parquet" + PMA_PHASE_EBI = f"{_V5_RANK}/pma_v5_phase_complex.parquet" # built when the complex phase lands +else: + PMA_PHASE_GENEKO = f"{_V4}/pma_phase_cells_v2_all.parquet" + PMA_PHASE_EBI = f"{_V4}/pma_phase_cells_ebi_all.parquet" + +GRAINS = { + "geneKO": {"parquet": PMA_PHASE_GENEKO, "class_col": "gene"}, + "complex": {"parquet": PMA_PHASE_EBI, "class_col": "predicted_class"}, # v5 complex parquet is complex-labeled (member cells pooled), like v4 + "minibinder": {"parquet": PMA_PHASE_GENEKO, "class_col": "gene"}, # NTC anchor from phase geneKO; targets supplied via accuracy_parquet +} + +# Default output root; per-run results land under ///. +DEFAULT_OUT_ROOT = "/hpc/projects/icd.fast.ops/models/diffex" + + +def slugify(name: str) -> str: + """Filesystem-safe tag for a class name (complex names have spaces/slashes).""" + return "".join(c if c.isalnum() else "_" for c in str(name)).strip("_") + + +@dataclass +class Config: + # ---- task / data (locked with Gav 2026-06-16) ---- + gene: str = "HSPA5" # positive class VALUE (a gene or a complex name) + class_col: str = "gene" # parquet column to match `gene` against ("gene" | "predicted_class") + n_per_class: int = 1000 # top-N attention cells per class + neg_rank_max: int = 5 # negatives = top-`neg_rank_max` cells of OTHER genes (distinct) + crop_size: int = 160 # single-cell crop (px) + channel: str = "Phase2D" # phase-only PoC + mask_cell: bool = False # no seg mask — usually better (keeps full crop context) + + # ---- train/val/test split (grouped by experiment by default) ---- + split_mode: str = "experiment" # "experiment" (grouped, confound guard) | "random" + val_fraction: float = 0.15 # model selection (best epoch) + test_fraction: float = 0.15 # clean held-out number (never used for selection) + seed: int = 0 + + # ---- training ---- + epochs: int = 30 + batch_size: int = 64 + lr: float = 1e-3 + weight_decay: float = 1e-4 + device: str = "cuda" + num_workers: int = 0 # crop materialization is a one-time zarr read; 0 = fork-safe + + # ---- paths ---- + pma_parquet: str = PMA_PHASE_GENEKO diff --git a/src/ops_model/models/attention/diffex/classifier/data.py b/src/ops_model/models/attention/diffex/classifier/data.py new file mode 100644 index 0000000..17db5fe --- /dev/null +++ b/src/ops_model/models/attention/diffex/classifier/data.py @@ -0,0 +1,208 @@ +"""Data layer: build the cell table, materialize phase crops, split. + +Positives = top-N attention cells of cfg.gene. +Negatives = top-`neg_rank_max` attention cells of every OTHER gene (the + "distinct" contrast), sampled to N — like-with-like (strong vs strong). + +Crops are read once via the shared `BaseDataset` loader and cached, so both +option B (ResNet on crops) and option C (CellDINO features on the SAME crops) +train off one identical cell set. +""" +from __future__ import annotations + +import os +import re +import warnings + +import numpy as np +import pandas as pd +import torch +from torch.utils.data import DataLoader + +from iohub import open_ome_zarr +from ops_utils.data.bbox_utils import BaseDataset +from ops_utils.data.experiment import OpsDataset +from ops_utils.data.filesystem import resolve_experiment_name + +_BASE_COLS = [ + "experiment", "well", "segmentation", "x_pheno", "y_pheno", "pma_attention", "rank", +] + + +def _normalize_well(well_str) -> str: + """'A3' -> 'A/3/0' (OPS zarr positions are 3-level row/col/fov).""" + w = str(well_str).strip() + if w.count("/") == 2: + return w + if w.count("/") == 1: + return f"{w}/0" + m = re.match(r"^([A-Za-z]+)(\d+)$", w) + if not m: + raise ValueError(f"Unknown well format: {well_str!r}") + return f"{m.group(1)}/{m.group(2)}/0" + + +def build_cell_table(cfg) -> pd.DataFrame: + """Filtered reads off the parquet (no full load). Class column is cfg.class_col + ('gene' for geneKO, 'predicted_class' for complexes).""" + cc = cfg.class_col + cols = [cc] + _BASE_COLS + pos = pd.read_parquet( + cfg.pma_parquet, + filters=[(cc, "==", cfg.gene), ("rank_type", "==", "top")], + columns=cols, + ) + if pos.empty: + raise ValueError(f"No 'top' attention rows for {cc}={cfg.gene!r}") + pos = pos.sort_values("rank").head(cfg.n_per_class).copy() + pos["label"] = 1 + + neg_pool = pd.read_parquet( + cfg.pma_parquet, + filters=[("rank_type", "==", "top"), ("rank", "<=", cfg.neg_rank_max)], + columns=cols, + ) + neg_pool = neg_pool[neg_pool[cc] != cfg.gene] + neg = neg_pool.sample( + n=min(cfg.n_per_class, len(neg_pool)), random_state=cfg.seed + ).copy() + neg["label"] = 0 + + df = pd.concat([pos, neg], ignore_index=True).rename(columns={cc: "cls"}) + print(f"cell table: {int(df.label.sum())} pos ({cfg.gene}) + " + f"{int((df.label == 0).sum())} neg (other classes' top-{cfg.neg_rank_max})") + return df + + +def make_labels_df(df: pd.DataFrame, cfg) -> pd.DataFrame: + """Cell table -> BaseDataset labels_df (bbox from x/y_pheno, like the atlas).""" + half = cfg.crop_size // 2 + recs = [] + for i, r in df.reset_index(drop=True).iterrows(): + y, x = int(r["y_pheno"]), int(r["x_pheno"]) + recs.append({ + "experiment": r["experiment"], + "store_key": r["experiment"], + "well": _normalize_well(r["well"]), + # clamp low end so border cells don't negative-wrap the zarr slice; + # SpatialPadd pads any short side back to crop_size. + "bbox": [max(0, y - half), max(0, x - half), y + half, x + half], + "segmentation_id": r["segmentation"], + "gene_name": r["cls"], + "label": int(r["label"]), + "total_index": i, + }) + return pd.DataFrame(recs) + + +def _open_stores(experiments): + stores = {} + for exp in experiments: + try: + ds = OpsDataset(resolve_experiment_name(exp)) + path = ds.store_paths["pheno_assembled_v3"] + with warnings.catch_warnings(): + warnings.filterwarnings("ignore") + stores[exp] = open_ome_zarr(str(path), mode="r") + except Exception as e: # noqa: BLE001 - surface and skip the experiment + print(f"[store] failed {exp}: {e}") + return stores + + +def _collate(batch): + data = torch.stack([torch.as_tensor(b["data"]) for b in batch]) + ti = torch.tensor([b["total_index"] for b in batch]) + return {"data": data, "total_index": ti} + + +def materialize_crops(labels_df: pd.DataFrame, cfg, cache_path=None): + """Read every crop once via BaseDataset; return (images, labels, experiment). + + images: (N, 1, crop, crop) float32. Cached to ``cache_path`` (.npz). + """ + if cache_path and os.path.exists(cache_path): + d = np.load(cache_path, allow_pickle=True) + print(f"[cache] crops <- {cache_path} {d['images'].shape}") + return d["images"], d["labels"], d["experiment"] + + stores = _open_stores(labels_df["experiment"].unique()) + df = labels_df[labels_df["experiment"].isin(stores)].reset_index(drop=True) + df["total_index"] = range(len(df)) + if len(df) < len(labels_df): + print(f"[store] dropped {len(labels_df) - len(df)} cells from failed stores") + + ds = BaseDataset( + stores=stores, + labels_df=df, + initial_yx_patch_size=(cfg.crop_size, cfg.crop_size), + final_yx_patch_size=(cfg.crop_size, cfg.crop_size), + out_channels=[cfg.channel], + mask_cell=cfg.mask_cell, + ) + loader = DataLoader( + ds, batch_size=cfg.batch_size, shuffle=False, + num_workers=cfg.num_workers, collate_fn=_collate, + ) + images = np.zeros((len(df), 1, cfg.crop_size, cfg.crop_size), np.float32) + for batch in loader: + idx = batch["total_index"].numpy() + images[idx] = batch["data"].numpy().astype(np.float32) + labels = df["label"].to_numpy(np.int64) + experiment = df["experiment"].to_numpy() + + if cache_path: + np.savez(cache_path, images=images, labels=labels, experiment=experiment) + print(f"[cache] crops -> {cache_path} {images.shape}") + return images, labels, experiment + + +def split_train_val_test(labels: np.ndarray, experiment: np.ndarray, cfg): + """3-way split (train/val/test). Grouped by experiment by default (confound + guard: val & test cells come from experiments the model never trained on). + val = model selection; test = clean held-out number. Falls back to + stratified random if the grouped split leaves a class missing from any side. + + Returns three boolean masks (train, val, test). + """ + rng = np.random.default_rng(cfg.seed) + n = len(labels) + + def _ok(which) -> bool: + for s in ("train", "val", "test"): + m = which == s + if m.sum() == 0 or labels[m].min() != 0 or labels[m].max() != 1: + return False + return True + + def _stratified_random(): + which = np.array(["train"] * n, dtype=" stratified random") + which = _stratified_random() + else: + which = _stratified_random() + + tr, va, te = which == "train", which == "val", which == "test" + print(f"[split] train={int(tr.sum())} val={int(va.sum())} test={int(te.sum())} " + f"(mode={cfg.split_mode})") + return tr, va, te diff --git a/src/ops_model/models/attention/diffex/classifier/models.py b/src/ops_model/models/attention/diffex/classifier/models.py new file mode 100644 index 0000000..b9ba278 --- /dev/null +++ b/src/ops_model/models/attention/diffex/classifier/models.py @@ -0,0 +1,34 @@ +"""The two candidate classifiers DiffEx will explain. + +B: ResNet18 on single-cell phase crops (pixels in -> class logits). +C: MLP head on frozen CellDINO embeddings of the same crops. +Both 2-class (gene-vs-rest); extends cleanly to the full N-way softmax later. +""" +from __future__ import annotations + +import torch.nn as nn +from torchvision import models + + +def build_resnet(in_channels: int = 1, n_classes: int = 2) -> nn.Module: + """ResNet18 with a 1-channel stem (phase) and a small head.""" + m = models.resnet18(weights=None) + m.conv1 = nn.Conv2d(in_channels, 64, kernel_size=7, stride=2, padding=3, bias=False) + m.fc = nn.Linear(m.fc.in_features, n_classes) + return m + + +class MLPHead(nn.Module): + """Small MLP over CellDINO embeddings (option C).""" + + def __init__(self, in_dim: int, hidden: int = 256, n_classes: int = 2, dropout: float = 0.2): + super().__init__() + self.net = nn.Sequential( + nn.Linear(in_dim, hidden), + nn.ReLU(inplace=True), + nn.Dropout(dropout), + nn.Linear(hidden, n_classes), + ) + + def forward(self, x): + return self.net(x) diff --git a/src/ops_model/models/attention/diffex/classifier/run.py b/src/ops_model/models/attention/diffex/classifier/run.py new file mode 100644 index 0000000..8459f89 --- /dev/null +++ b/src/ops_model/models/attention/diffex/classifier/run.py @@ -0,0 +1,100 @@ +"""Orchestrator for the DiffEx single-cell classifier PoC. + + python -m ops_model.models.attention.diffex.classifier.run --model B --gene HSPA5 + python -m ops_model.models.attention.diffex.classifier.run --model C --gene HSPA5 + +Shared steps: build cell table -> materialize phase crops (cached) -> split. +Then B trains a ResNet on crops; C embeds the crops with CellDINO and trains an +MLP. Writes model_.pt + metrics_.json. The crops/features cache is shared, +so running C after B reuses the same crops. + +``run_poc`` is the importable entry point used by ``submit.py`` (SLURM). +""" +from __future__ import annotations + +import argparse +import json +from pathlib import Path + +import torch + +from .celldino_features import embed_crops +from .config import DEFAULT_OUT_ROOT, GRAINS, Config, slugify +from .data import build_cell_table, make_labels_df, materialize_crops, split_train_val_test +from .models import MLPHead, build_resnet +from .train import train_classifier + + +def run_poc(cfg: Config, model: str, out_dir: str) -> dict: + """Run one PoC (model 'B' or 'C') and write artifacts to ``out_dir``.""" + out = Path(out_dir) + cache = out / "cache" + cache.mkdir(parents=True, exist_ok=True) + + # ---- shared: cells -> crops -> split (crops cache is model-agnostic) ---- + tag = f"{slugify(cfg.gene)}_{cfg.crop_size}_{'mask' if cfg.mask_cell else 'nomask'}" + df = build_cell_table(cfg) + labels_df = make_labels_df(df, cfg) + images, labels, experiment = materialize_crops( + labels_df, cfg, cache_path=str(cache / f"crops_{tag}.npz") + ) + tr, va, te = split_train_val_test(labels, experiment, cfg) + + # ---- model-specific features + model ---- + if model == "B": + X = images + net = build_resnet(in_channels=1) + elif model == "C": + X = embed_crops(images, cfg, cache_path=str(cache / f"celldino_{tag}.npz")) + net = MLPHead(in_dim=X.shape[1]) + else: + raise ValueError(f"model must be 'B' or 'C', got {model!r}") + + best = train_classifier( + net, X[tr], labels[tr], X[va], labels[va], X[te], labels[te], cfg + ) + + torch.save(best["state"], out / f"model_{model}.pt") + metrics = { + "model": model, "gene": cfg.gene, "class_col": cfg.class_col, + "n_pos": int((labels == 1).sum()), "n_neg": int((labels == 0).sum()), + "n_train": int(tr.sum()), "n_val": int(va.sum()), "n_test": int(te.sum()), + "split_mode": cfg.split_mode, + "test_auroc": best["test_auroc"], + "val_auroc": best["val_auroc"], "train_auroc_at_best": best["train_auroc"], + "best_epoch": best["epoch"], + "feature_dim": int(X.shape[1]) if model == "C" else None, + } + (out / f"metrics_{model}.json").write_text(json.dumps(metrics, indent=2)) + (out / f"history_{model}.json").write_text(json.dumps(best["history"], indent=2)) + print(json.dumps(metrics, indent=2)) + return metrics + + +def main(): + ap = argparse.ArgumentParser(description="DiffEx single-cell classifier PoC (B/C)") + ap.add_argument("--model", choices=["B", "C"], required=True, + help="B = ResNet on crops; C = MLP on CellDINO features") + ap.add_argument("--grain", choices=list(GRAINS), default="geneKO") + ap.add_argument("--gene", default="HSPA5", help="class value (gene or complex name)") + ap.add_argument("--n-per-class", type=int, default=1000) + ap.add_argument("--crop-size", type=int, default=160) + ap.add_argument("--epochs", type=int, default=30) + ap.add_argument("--split-mode", choices=["experiment", "random"], default="experiment") + ap.add_argument("--device", default="cuda") + ap.add_argument("--num-workers", type=int, default=0) + ap.add_argument("--out-dir", default=None) + args = ap.parse_args() + + grain = GRAINS[args.grain] + cfg = Config( + gene=args.gene, class_col=grain["class_col"], pma_parquet=grain["parquet"], + n_per_class=args.n_per_class, crop_size=args.crop_size, + epochs=args.epochs, split_mode=args.split_mode, device=args.device, + num_workers=args.num_workers, + ) + run_poc(cfg, args.model, args.out_dir or f"{DEFAULT_OUT_ROOT}/{args.grain}/{slugify(cfg.gene)}") + + +if __name__ == "__main__": + main() diff --git a/src/ops_model/models/attention/diffex/classifier/submit.py b/src/ops_model/models/attention/diffex/classifier/submit.py new file mode 100644 index 0000000..3853193 --- /dev/null +++ b/src/ops_model/models/attention/diffex/classifier/submit.py @@ -0,0 +1,97 @@ +"""Submit the classifier sweep to SLURM (GPU) via submit_parallel_jobs. + + # one gene, both models + python -m ops_model.models.attention.diffex.classifier.submit --gene HSPA5 --models B C + + # all 98 EBI complexes, model C + python -m ops_model.models.attention.diffex.classifier.submit --grain complex --all-classes --models C + + # specific classes + python -m ops_model.models.attention.diffex.classifier.submit --grain complex \ + --classes "19S proteasome regulatory complex" "Commander complex" --models C + +One GPU job per (class, model). Outputs under ///. +""" +from __future__ import annotations + +import argparse +from pathlib import Path + +import pandas as pd +from ops_utils.hpc.slurm_batch_utils import submit_parallel_jobs + +from .config import DEFAULT_OUT_ROOT, GRAINS, Config, slugify +from .run import run_poc + + +def _all_classes(parquet: str, class_col: str) -> list[str]: + """Distinct class values present in the parquet (rank-1 top rows = cheap). + Includes NTC — a useful negative-control bin (its classifier should be ~chance).""" + df = pd.read_parquet( + parquet, filters=[("rank", "==", 1), ("rank_type", "==", "top")], + columns=[class_col], + ) + return sorted(df[class_col].astype(str).unique()) + + +def main(): + ap = argparse.ArgumentParser(description="Submit DiffEx classifier sweep to SLURM") + ap.add_argument("--grain", choices=list(GRAINS), default="geneKO") + ap.add_argument("--gene", default="HSPA5", help="single class value (ignored if --classes/--all-classes)") + ap.add_argument("--classes", nargs="+", default=None, help="explicit class values") + ap.add_argument("--all-classes", action="store_true", help="every class in the grain's parquet") + ap.add_argument("--models", nargs="+", choices=["B", "C"], default=["C"]) + ap.add_argument("--n-per-class", type=int, default=1000) + ap.add_argument("--crop-size", type=int, default=160) + ap.add_argument("--epochs", type=int, default=30) + ap.add_argument("--split-mode", choices=["experiment", "random"], default="experiment") + ap.add_argument("--out-dir", default=DEFAULT_OUT_ROOT) + # SLURM + ap.add_argument("--partition", default="gpu") + ap.add_argument("--gres", default="gpu:1") + ap.add_argument("--cpus", type=int, default=8) + ap.add_argument("--mem-gb", type=int, default=64) + ap.add_argument("--time-min", type=int, default=240) + ap.add_argument("--dry-run", action="store_true") + args = ap.parse_args() + + grain = GRAINS[args.grain] + parquet, class_col = grain["parquet"], grain["class_col"] + if args.all_classes: + classes = _all_classes(parquet, class_col) + elif args.classes: + classes = args.classes + else: + classes = [args.gene] + print(f"grain={args.grain} class_col={class_col} -> {len(classes)} classes; models={args.models}") + + base = Path(args.out_dir).resolve() / args.grain + jobs = [] + for cls in classes: + slug = slugify(cls) + for m in args.models: + cfg = Config( + gene=cls, class_col=class_col, pma_parquet=parquet, + n_per_class=args.n_per_class, crop_size=args.crop_size, + epochs=args.epochs, split_mode=args.split_mode, device="cuda", + ) + jobs.append({ + "name": f"diffex_{args.grain}_{slug}_{m}"[:64], + "func": run_poc, + "kwargs": {"cfg": cfg, "model": m, "out_dir": str(base / slug)}, + "metadata": {"grain": args.grain, "class": cls, "model": m}, + }) + + slurm_params = { + "slurm_partition": args.partition, "slurm_gres": args.gres, + "cpus_per_task": args.cpus, "mem_gb": args.mem_gb, "timeout_min": args.time_min, + } + submit_parallel_jobs( + jobs_to_submit=jobs, experiment=f"diffex_clf_{args.grain}", + slurm_params=slurm_params, log_dir=f"diffex_clf_{args.grain}", + dry_run=args.dry_run, wait_for_completion=False, + ) + + +if __name__ == "__main__": + main() diff --git a/src/ops_model/models/attention/diffex/classifier/train.py b/src/ops_model/models/attention/diffex/classifier/train.py new file mode 100644 index 0000000..ef03a8f --- /dev/null +++ b/src/ops_model/models/attention/diffex/classifier/train.py @@ -0,0 +1,76 @@ +"""Generic train/eval loop shared by options B and C. + +Success criterion (PoC): held-out AUROC clearly > 0.5. Because the split is +grouped by experiment, a high AUROC means the classifier separates the gene on +biology that generalizes across experiments, not on a per-plate confound. +""" +from __future__ import annotations + +import numpy as np +import torch +import torch.nn as nn +from sklearn.metrics import roc_auc_score +from torch.utils.data import DataLoader, TensorDataset + + +def _resolve_device(name: str) -> torch.device: + if name.startswith("cuda") and not torch.cuda.is_available(): + print("[device] cuda requested but unavailable -> cpu") + return torch.device("cpu") + return torch.device(name) + + +@torch.no_grad() +def evaluate(model, X, y, batch_size, dev) -> float: + model.eval() + probs = [] + for i in range(0, len(X), batch_size): + xb = torch.as_tensor(X[i:i + batch_size]).to(dev) + probs.append(torch.softmax(model(xb), -1)[:, 1].cpu().numpy()) + return float(roc_auc_score(y, np.concatenate(probs))) + + +def train_classifier(model, Xtr, ytr, Xva, yva, Xte, yte, cfg) -> dict: + dev = _resolve_device(cfg.device) + model = model.to(dev) + opt = torch.optim.AdamW(model.parameters(), lr=cfg.lr, weight_decay=cfg.weight_decay) + crit = nn.CrossEntropyLoss() + + tr = DataLoader( + TensorDataset(torch.as_tensor(Xtr), torch.as_tensor(ytr)), + batch_size=cfg.batch_size, shuffle=True, + ) + best = {"val_auroc": -1.0, "train_auroc": -1.0, "epoch": -1, "state": None} + history = [] + for ep in range(cfg.epochs): + model.train() + ep_loss = 0.0 + for xb, yb in tr: + xb, yb = xb.to(dev), yb.to(dev) + opt.zero_grad() + loss = crit(model(xb), yb) + loss.backward() + opt.step() + ep_loss += float(loss) * len(xb) + ep_loss /= len(Xtr) + # train + val AUROC each epoch -> watch the gap (overfit vs underfit) + tr_auroc = evaluate(model, Xtr, ytr, cfg.batch_size, dev) + va_auroc = evaluate(model, Xva, yva, cfg.batch_size, dev) + history.append({"epoch": ep, "train_loss": ep_loss, + "train_auroc": tr_auroc, "val_auroc": va_auroc}) + if va_auroc > best["val_auroc"]: + best = { + "val_auroc": va_auroc, "train_auroc": tr_auroc, "epoch": ep, + "state": {k: v.detach().cpu() for k, v in model.state_dict().items()}, + } + print(f"epoch {ep:02d}: loss={ep_loss:.4f} train_auroc={tr_auroc:.4f} " + f"val_auroc={va_auroc:.4f} (best val {best['val_auroc']:.4f})") + best["history"] = history + + # clean held-out number: reload the val-selected best, score test ONCE + if best["state"] is not None: + model.load_state_dict(best["state"]) + best["test_auroc"] = evaluate(model, Xte, yte, cfg.batch_size, dev) + print(f"selected epoch {best['epoch']}: val={best['val_auroc']:.4f} " + f"test={best['test_auroc']:.4f}") + return best diff --git a/src/ops_model/models/attention/diffex/diffae/__init__.py b/src/ops_model/models/attention/diffex/diffae/__init__.py new file mode 100644 index 0000000..aa97a2e --- /dev/null +++ b/src/ops_model/models/attention/diffex/diffae/__init__.py @@ -0,0 +1,6 @@ +"""DiffAE (diffusion autoencoder) — the DiffEx generator stage. + +Semantic encoder + conditional UNet decoder, trained jointly on broad phase crops. +See ../PLAN.md §2 and README.md. The next stage (contrastive direction discovery, +§3) builds on the trained z_sem latent + the classifier score. +""" diff --git a/src/ops_model/models/attention/diffex/diffae/config.py b/src/ops_model/models/attention/diffex/diffae/config.py new file mode 100644 index 0000000..ca2d8b4 --- /dev/null +++ b/src/ops_model/models/attention/diffex/diffae/config.py @@ -0,0 +1,74 @@ +"""Config for the DiffAE (diffusion autoencoder) — DiffEx generator stage. + +Semantic encoder -> z_sem; conditional UNet decoder denoises conditioned on z_sem +(injected into the time embedding). Trained jointly on a broad sample of phase +single-cell crops. Gate: round-trip (DDIM-invert -> reverse) reconstruction. +""" +from __future__ import annotations + +from dataclasses import dataclass, field + +from ..classifier.config import DEFAULT_OUT_ROOT, PMA_PHASE_GENEKO # noqa: F401 + + +@dataclass +class DiffAEConfig: + # ---- data (locked 2026-06-16: broad sample, all classes incl NTC, all ranks) ---- + pma_parquet: str = PMA_PHASE_GENEKO + n_crops: int = 50_000 + crop_size: int = 160 + channel: str = "Phase2D" # RAW pheno-zarr channel to read (Phase2D | GFP | mCherry | Cy5) + # fluorescent mode: set marker_channel to a `channel` value in the fluor attention CSV + # (e.g. "nucleolus-GC_NPM3"); build_broad_table then samples that marker's cells and the + # generator trains on the raw `channel` above. None = phase mode (samples pma_parquet). + marker_channel: str | None = None + # virtual staining: condition on a DIFFERENT channel than the generation target. + # cond_channel = raw channel CellDINO-embedded for conditioning (e.g. "Phase2D"); `channel` + # stays the generation target (e.g. "mCherry"). None = same-channel autoencoding (default). + cond_channel: str | None = None + # spatial (image-to-image) conditioning: concat the cond_channel IMAGE into the UNet input + # channels (in_channels 1->2) so the decoder sees dense per-pixel phase (SR3/Palette recipe) — + # pixel-registered virtual staining, vs the global-embedding-only default. Requires cond_channel. + spatial_cond: bool = False + # multi-marker virtual staining: one model predicts N markers, switched by a learned marker-id + # embedding added to the conditioning. 0 = single-channel (default, existing checkpoints unaffected). + n_markers: int = 0 + fluor_csv: str = f"{PMA_PHASE_GENEKO.rsplit('/', 1)[0]}/pma_fluorescent_cells_all.csv" + mask_cell: bool = False + seed: int = 0 + + # ---- model ---- + cond_dim: int = 1024 # FROZEN CellDINO (ViT-L) embedding dim — the conditioning + block_out_channels: tuple = (128, 256, 256, 512) + layers_per_block: int = 2 + + # ---- diffusion ---- + train_timesteps: int = 1000 + + # ---- training: proper conditional-diffusion recipe ---- + # The v1 run had ~dead conditioning (emb/noise ratio 0.008). Fix = conditioning + # dropout (forces the model to USE the embedding) + EMA + much longer training. + epochs: int = 120 # resume across jobs to reach this + batch_size: int = 48 # single-GPU (DataParallel breaks diffusers; DDP TODO) + lr: float = 1e-4 + weight_decay: float = 0.0 + cond_dropout: float = 0.15 # replace embedding with learned null this often (enables CFG) + ema_decay: float = 0.9995 # EMA weights used for sampling/eval + resume: bool = True # continue from saved train state if present + # dihedral augmentation: randomly rot90×flip the TARGET image while keeping the + # canonical CellDINO embedding fixed → teaches the model orientation is NOT + # embedding-determined, so traversals stop spuriously rotating the cell. + augment_dihedral: bool = True + # affine augmentation: continuous rotation (±180°) + scale + flip with reflection + # padding (no black corners), matching the contrastive data_loader recipe. Covers + # arbitrary orientations (not just 90° steps); takes precedence over dihedral when set. + augment_affine: bool = False + affine_scale: float = 0.15 # ±fraction scale jitter + init_ckpt: str | None = None # warm-start: load these weights into the fresh model + device: str = "cuda" + num_workers: int = 0 + + # ---- reconstruction gate ---- + recon_every: int = 5 # epochs between recon montages + n_recon: int = 16 # cells in the montage + ddim_steps: int = 50 diff --git a/src/ops_model/models/attention/diffex/diffae/data.py b/src/ops_model/models/attention/diffex/diffae/data.py new file mode 100644 index 0000000..3d59b07 --- /dev/null +++ b/src/ops_model/models/attention/diffex/diffae/data.py @@ -0,0 +1,98 @@ +"""Broad phase-crop sampler for DiffAE training. + +Samples ~n_crops cells uniformly across the whole geneKO phase parquet (all genes +incl NTC, all attention ranks), then reuses the classifier's crop materialization +so the DiffAE trains on the same crop pipeline. Normalization: per-image z-score +then /3 + clip to [-1, 1] (diffusion-friendly, intensity-invariant like CellDINO). +""" +from __future__ import annotations + +import numpy as np +import pandas as pd +import pyarrow.parquet as pq + +from ..classifier.celldino_features import embed_crops # frozen CellDINO encoder +from ..classifier.data import make_labels_df, materialize_crops # reuse + +_COLS = ["gene", "experiment", "well", "segmentation", "x_pheno", "y_pheno"] + + +def build_broad_table(cfg) -> pd.DataFrame: + """Uniform fraction-sample across all row groups (covers all genes/ranks). + Fluorescent mode (cfg.marker_channel set): sample that marker's cells from the fluor + attention CSV; the generator reads cfg.channel (GFP/mCherry) — the raw pheno-zarr channel + carrying that marker.""" + if getattr(cfg, "marker_channel", None): + pre = getattr(cfg, "_fluor_rows", None) # preloaded rows (multi-marker: read the 12GB CSV ONCE) + if pre is not None: + df = pre[pre["channel"] == cfg.marker_channel] if "channel" in pre.columns else pre + if "rank_type" in df.columns: + df = df[df["rank_type"] == "top"] + else: + cols = ["gene", "channel", "experiment", "well", "segmentation", "x_pheno", "y_pheno", "rank_type"] + df = pd.read_csv(cfg.fluor_csv, usecols=cols) + df = df[(df["channel"] == cfg.marker_channel) & (df["rank_type"] == "top")] + if df.empty: + raise ValueError(f"no 'top' cells for marker_channel={cfg.marker_channel!r}") + if len(df) > cfg.n_crops: + df = df.sample(n=cfg.n_crops, random_state=cfg.seed) + df = df.reset_index(drop=True).rename(columns={"gene": "cls"}) + df["label"] = 0 + print(f"fluor broad table [{cfg.marker_channel} -> raw {cfg.channel}]: " + f"{len(df)} crops across {df['cls'].nunique()} genes") + return df + + pf = pq.ParquetFile(cfg.pma_parquet) + total = pf.metadata.num_rows + frac = min(1.0, cfg.n_crops / total * 1.15) + rng = np.random.default_rng(cfg.seed) + parts = [] + for batch in pf.iter_batches(columns=_COLS, batch_size=250_000): + df = batch.to_pandas() + keep = rng.random(len(df)) < frac + if keep.any(): + parts.append(df.loc[keep]) + df = pd.concat(parts, ignore_index=True) + if len(df) > cfg.n_crops: + df = df.sample(n=cfg.n_crops, random_state=cfg.seed).reset_index(drop=True) + df = df.rename(columns={"gene": "cls"}) + df["label"] = 0 # unconditional generator — label unused + print(f"broad table: {len(df)} crops across {df['cls'].nunique()} classes") + return df + + +def normalize(images: np.ndarray) -> np.ndarray: + """Per-image z-score, /3, clip to [-1, 1]. images: (N,1,H,W).""" + x = images.astype(np.float32) + mu = x.mean(axis=(-2, -1), keepdims=True) + sd = x.std(axis=(-2, -1), keepdims=True) + 1e-6 + return np.clip((x - mu) / sd / 3.0, -1.0, 1.0) + + +def load_diffae_crops(cfg, crops_cache, emb_cache, cond_cache=None, return_cond_images=False): + """Returns (images_norm, celldino_embs[, cond_images_norm]). + + images_norm: (N,1,H,W) generation target, normalized for diffusion. + celldino_embs: (N, cond_dim) FROZEN CellDINO conditioning embeddings. + + Same-channel (default): embeddings come from the SAME crops as the target (Alex's design). + Virtual staining (cfg.cond_channel set): the target is `cfg.channel` (e.g. mCherry) while the + conditioning is CellDINO of the co-registered `cfg.cond_channel` crop (e.g. Phase2D) at the SAME + cell locations — same labels_df → aligned index. return_cond_images also returns the normalized + conditioning (phase) crops for eval montages. + """ + import dataclasses + df = build_broad_table(cfg) + labels_df = make_labels_df(df, cfg) + images_raw, _, _ = materialize_crops(labels_df, cfg, cache_path=crops_cache) # target (cfg.channel) + if getattr(cfg, "cond_channel", None): # virtual staining: embed a DIFFERENT channel + cond_cfg = dataclasses.replace(cfg, channel=cfg.cond_channel) + cond_raw, _, _ = materialize_crops(labels_df, cond_cfg, cache_path=cond_cache) + embs = embed_crops(cond_raw, cfg, cache_path=emb_cache) # CellDINO of the conditioning channel + if return_cond_images: + return normalize(images_raw), embs, normalize(cond_raw) + return normalize(images_raw), embs + embs = embed_crops(images_raw, cfg, cache_path=emb_cache) # frozen CellDINO (same-channel) + if return_cond_images: + return normalize(images_raw), embs, normalize(images_raw) + return normalize(images_raw), embs diff --git a/src/ops_model/models/attention/diffex/diffae/diagnose_conditioning.py b/src/ops_model/models/attention/diffex/diffae/diagnose_conditioning.py new file mode 100644 index 0000000..32b32a8 --- /dev/null +++ b/src/ops_model/models/attention/diffex/diffae/diagnose_conditioning.py @@ -0,0 +1,133 @@ +"""Diagnostic: does the trained DiffAE's embedding actually control generation? + +From a SINGLE fixed noise, generate conditioned on (null / control-centroid / +KD-centroid). If conditioning is strong, control vs KD should look visibly +different. Key metric: MSE(ctrl-img, kd-img) at fixed noise, compared to +MSE between two DIFFERENT noises — if embedding-driven change ≪ noise-driven +change, the embedding has weak control (the bug we suspect). + + python -m ops_model.models.attention.diffex.diffae.diagnose_conditioning +""" +from __future__ import annotations + +import argparse +import json +from pathlib import Path + +import numpy as np +import torch +from diffusers import DDIMScheduler + +from ..classifier.config import DEFAULT_OUT_ROOT +from .config import DiffAEConfig +from .model import DiffAE + + +def run_diagnose(ckpt: str, emb_npz: str, crops_npz: str, out_dir: str, + n_noise: int = 6, ddim_steps: int = 50, device: str = "cuda") -> dict: + dev = torch.device(device if torch.cuda.is_available() or device == "cpu" else "cpu") + cfg = DiffAEConfig() + model = DiffAE(cfg) + model.load_state_dict(torch.load(ckpt, map_location="cpu")) + model = model.to(dev).eval() + + embs = np.load(emb_npz)["features"] + labels = np.load(crops_npz, allow_pickle=True)["labels"] + conds = { + "null": np.zeros(embs.shape[1], np.float32), + "ctrl": embs[labels == 0].mean(0).astype(np.float32), + "kd": embs[labels == 1].mean(0).astype(np.float32), + } + print(f"‖μ_kd−μ_ctrl‖ = {np.linalg.norm(conds['kd']-conds['ctrl']):.2f}") + + H = cfg.crop_size + + @torch.no_grad() + def sample(xT, emb): + fwd = DDIMScheduler(num_train_timesteps=cfg.train_timesteps) + fwd.set_timesteps(ddim_steps) + c = model.cond(torch.as_tensor(emb, device=dev)[None]) + x = xT + for t in fwd.timesteps: + x = fwd.step(model.denoise(x, t, c), t, x).prev_sample + return x.cpu().numpy()[0, 0] + + imgs = {k: [] for k in conds} + for i in range(n_noise): + g = torch.Generator(device=dev).manual_seed(100 + i) + xT = torch.randn(1, 1, H, H, generator=g, device=dev) + for k, e in conds.items(): + imgs[k].append(sample(xT.clone(), e)) + arr = {k: np.array(v) for k, v in imgs.items()} + + # embedding-driven change (same noise, ctrl vs kd) vs noise-driven change + mse_ctrl_kd = float(np.mean((arr["ctrl"] - arr["kd"]) ** 2)) + mse_null_ctrl = float(np.mean((arr["null"] - arr["ctrl"]) ** 2)) + mse_noise = float(np.mean([(arr["ctrl"][i] - arr["ctrl"][j]) ** 2 + for i in range(n_noise) for j in range(i + 1, n_noise)])) + ratio = mse_ctrl_kd / (mse_noise + 1e-9) + metrics = { + "mse_ctrl_vs_kd_same_noise": mse_ctrl_kd, + "mse_null_vs_ctrl": mse_null_ctrl, + "mse_noise_vs_noise": mse_noise, + "embedding_vs_noise_ratio": ratio, + } + print(json.dumps(metrics, indent=2)) + print(f"INTERPRETATION: embedding/noise ratio = {ratio:.3f}. " + f"≪1 ⇒ embedding has WEAK control (noise dominates); ~1+ ⇒ strong control.") + + # montage: rows = noise; cols = null | ctrl | kd | |ctrl−kd| + import matplotlib + matplotlib.use("Agg"); matplotlib.rcParams["pdf.fonttype"] = 42 + import matplotlib.pyplot as plt + out = Path(out_dir); out.mkdir(parents=True, exist_ok=True) + cols = ["null", "ctrl", "kd", "|ctrl−kd|"] + fig, ax = plt.subplots(n_noise, 4, figsize=(7, 1.7 * n_noise), squeeze=False) + for i in range(n_noise): + ax[i, 0].imshow(arr["null"][i], cmap="gray", vmin=-1, vmax=1) + ax[i, 1].imshow(arr["ctrl"][i], cmap="gray", vmin=-1, vmax=1) + ax[i, 2].imshow(arr["kd"][i], cmap="gray", vmin=-1, vmax=1) + d = np.abs(arr["ctrl"][i] - arr["kd"][i]) + ax[i, 3].imshow(d, cmap="hot", vmin=0, vmax=max(d.max(), 1e-3)) + for j in range(4): + ax[i, j].axis("off") + if i == 0: + ax[i, j].set_title(cols[j], fontsize=9) + fig.suptitle(f"DiffAE conditioning test — emb/noise MSE ratio={ratio:.3f}", fontsize=11) + fig.tight_layout() + fig.savefig(out / "conditioning_test.png", dpi=140, bbox_inches="tight") + plt.close(fig) + (out / "diagnose_metrics.json").write_text(json.dumps(metrics, indent=2)) + print(f"wrote {out}/conditioning_test.png") + return metrics + + +def main(): + ap = argparse.ArgumentParser() + base = f"{DEFAULT_OUT_ROOT}/diffae/phase_v1" + dirs = f"{DEFAULT_OUT_ROOT}/directions/phase/geneKO/HSPA5/cache" + ap.add_argument("--ckpt", default=f"{base}/diffae_best.pt") + ap.add_argument("--emb-npz", default=f"{dirs}/celldino_HSPA5_160.npz") + ap.add_argument("--crops-npz", default=f"{dirs}/crops_HSPA5_160.npz") + ap.add_argument("--out-dir", default=f"{base}/diagnose") + ap.add_argument("--local", action="store_true", help="run here instead of SLURM") + args = ap.parse_args() + + kwargs = dict(ckpt=args.ckpt, emb_npz=args.emb_npz, crops_npz=args.crops_npz, + out_dir=args.out_dir) + if args.local: + run_diagnose(device="cpu", **kwargs) + return + from ops_utils.hpc.slurm_batch_utils import submit_parallel_jobs + submit_parallel_jobs( + jobs_to_submit=[{"name": "diffae_diagnose", "func": run_diagnose, + "kwargs": kwargs, "metadata": {"stage": "diagnose"}}], + experiment="diffae_diagnose", + slurm_params={"slurm_partition": "gpu", "slurm_gres": "gpu:1", + "cpus_per_task": 4, "mem_gb": 32, "timeout_min": 30}, + log_dir="diffae_diagnose", wait_for_completion=False, + ) + + +if __name__ == "__main__": + main() diff --git a/src/ops_model/models/attention/diffex/diffae/model.py b/src/ops_model/models/attention/diffex/diffae/model.py new file mode 100644 index 0000000..6f48f41 --- /dev/null +++ b/src/ops_model/models/attention/diffex/diffae/model.py @@ -0,0 +1,64 @@ +"""Conditional diffusion decoder (Alex's DiffEx design). + +The UNet generates phase crops conditioned on the FROZEN CellDINO embedding — +NOT a jointly-trained encoder. The embedding is injected into the UNet's time +embedding (diffusers `class_embed_type="identity"`), which propagates FiLM-style +through the resnet blocks. Directions therefore live in CellDINO space, the same +space as the option-C classifier and the SetTransformer. +""" +from __future__ import annotations + +import torch +import torch.nn as nn +from diffusers import UNet2DModel + + +class DiffAE(nn.Module): + def __init__(self, cfg): + super().__init__() + n_blocks = len(cfg.block_out_channels) + # attention at the two lowest-resolution stages (kept — this arch produces the + # promising morphs; speed comes from multi-GPU, not from shrinking the model) + down = tuple("AttnDownBlock2D" if i >= n_blocks - 2 else "DownBlock2D" + for i in range(n_blocks)) + up = tuple("AttnUpBlock2D" if i < 2 else "UpBlock2D" for i in range(n_blocks)) + # spatial conditioning: concat the cond image into the input (noisy target + cond = 2 ch) + self.spatial_cond = getattr(cfg, "spatial_cond", False) + in_channels = 2 if self.spatial_cond else 1 + self.unet = UNet2DModel( + sample_size=cfg.crop_size, in_channels=in_channels, out_channels=1, + block_out_channels=cfg.block_out_channels, + layers_per_block=cfg.layers_per_block, + down_block_types=down, up_block_types=up, + class_embed_type="identity", + ) + time_embed_dim = cfg.block_out_channels[0] * 4 + # project the frozen CellDINO embedding to the time-embedding dim. Deeper + # MLP (not a single Linear) so the conditioning can actually be used. + self.cond_proj = nn.Sequential( + nn.Linear(cfg.cond_dim, time_embed_dim), nn.SiLU(), + nn.Linear(time_embed_dim, time_embed_dim), + ) + # learned null embedding for conditioning dropout / classifier-free guidance + self.null_emb = nn.Parameter(torch.zeros(cfg.cond_dim)) + # multi-marker: learned per-marker embedding added to the conditioning (which channel to render) + self.n_markers = getattr(cfg, "n_markers", 0) + if self.n_markers: + self.marker_emb = nn.Embedding(self.n_markers, cfg.cond_dim) + + def cond(self, emb: torch.Tensor, marker_id=None) -> torch.Tensor: + """Frozen CellDINO embedding (B, cond_dim) -> time-embedding conditioning. + marker_id (B,) long: add the learned marker embedding (multi-marker virtual staining).""" + if self.n_markers and marker_id is not None: + emb = emb + self.marker_emb(marker_id) + return self.cond_proj(emb) + + def null(self, n: int, device) -> torch.Tensor: + return self.null_emb[None].expand(n, -1).to(device) + + def denoise(self, noisy, t, c, cond_img=None) -> torch.Tensor: + x = torch.cat([noisy, cond_img], dim=1) if self.spatial_cond else noisy # dense phase concat + return self.unet(x, t, class_labels=c).sample + + def forward(self, noisy, t, emb, cond_img=None, marker_id=None): + return self.denoise(noisy, t, self.cond(emb, marker_id), cond_img) diff --git a/src/ops_model/models/attention/diffex/diffae/plot_metrics.py b/src/ops_model/models/attention/diffex/diffae/plot_metrics.py new file mode 100644 index 0000000..bf7e10f --- /dev/null +++ b/src/ops_model/models/attention/diffex/diffae/plot_metrics.py @@ -0,0 +1,65 @@ +"""Plot final training loss + conditioning ratio (best_ratio) for every DiffAE generator, read +straight from each run's diffae_train_state.pt (history[-1].loss, best_ratio, epoch).""" +import glob +import os + +import matplotlib +matplotlib.use("Agg") +import matplotlib.pyplot as plt +import numpy as np +import torch + +plt.rcParams["pdf.fonttype"] = 42 + +DD = "/hpc/projects/icd.fast.ops/models/diffex/diffae" +OUT = "/hpc/projects/icd.fast.ops/models/diffex/model_metrics" + + +def collect(): + rows = [] + for sp in sorted(glob.glob(f"{DD}/*/diffae_train_state.pt")): + name = os.path.basename(os.path.dirname(sp)) + try: + s = torch.load(sp, map_location="cpu", mmap=True) + except Exception as e: + print("skip", name, e); continue + h = s.get("history", []) + cr = [(e["epoch"], e["cond_ratio"]) for e in h if e.get("cond_ratio", -1) >= 0] + rows.append(dict(name=name, ep=[e["epoch"] for e in h], loss=[e["loss"] for e in h], + cr_ep=[x[0] for x in cr], cr=[x[1] for x in cr], + is_phase=name.startswith("phase"), best=float(s.get("best_ratio", np.nan)))) + return rows + + +def main(): + rows = collect() + fluor = [r for r in rows if not r["is_phase"]] + fcolors = plt.cm.viridis(np.linspace(0, 0.95, max(len(fluor), 1))) + fig, ax = plt.subplots(1, 2, figsize=(16, 8)) + + for i, r in enumerate(fluor): # fluor: thin viridis lines (each its own) + ax[0].plot(r["ep"], r["loss"], color=fcolors[i], lw=0.8, alpha=0.5) + if r["cr"]: + ax[1].plot(r["cr_ep"], r["cr"], color=fcolors[i], lw=0.8, alpha=0.5) + for r in rows: # phase: thick orange, labeled, on top + if not r["is_phase"]: + continue + ax[0].plot(r["ep"], r["loss"], lw=2.4, alpha=0.95, label=r["name"]) + if r["cr"]: + ax[1].plot(r["cr_ep"], r["cr"], lw=2.4, alpha=0.95, label=r["name"]) + + ax[0].set_yscale("log"); ax[0].set_title("training loss"); ax[0].set_xlabel("epoch"); ax[0].set_ylabel("loss (log)") + ax[1].axhline(0.468, ls="--", c="#888", lw=1, label="prod 0.468") + ax[1].set_title("conditioning ratio"); ax[1].set_xlabel("epoch"); ax[1].set_ylabel("cond_ratio") + for a in ax: + a.grid(alpha=0.3); a.legend(fontsize=7, ncol=2, loc="best") + fig.suptitle(f"DiffAE generators (n={len(rows)}) — loss + conditioning-ratio curves " + f"(phase highlighted; {len(fluor)} fluor in viridis)", y=0.99) + fig.tight_layout() + for ext in ("png", "svg"): + fig.savefig(f"{OUT}_curves.{ext}", dpi=140, bbox_inches="tight") + print(f"[plot] {len(rows)} models ({len(fluor)} fluor) -> {OUT}_curves.png/.svg") + + +if __name__ == "__main__": + main() diff --git a/src/ops_model/models/attention/diffex/diffae/recon.py b/src/ops_model/models/attention/diffex/diffae/recon.py new file mode 100644 index 0000000..e6af8dd --- /dev/null +++ b/src/ops_model/models/attention/diffex/diffae/recon.py @@ -0,0 +1,69 @@ +"""Reconstruction gate: DDIM-invert each cell to x_T then reverse, both +conditioned on its z_sem. Faithful encode->decode round-trip → PSNR + montage. +If the DiffAE can't reconstruct cells, learned directions would be meaningless. +""" +from __future__ import annotations + +from pathlib import Path + +import numpy as np +import torch +from diffusers import DDIMInverseScheduler, DDIMScheduler + + +@torch.no_grad() +def reconstruct(model, x0: torch.Tensor, cell_emb: torch.Tensor, cfg, dev) -> torch.Tensor: + """x0: (B,1,H,W) in [-1,1]; cell_emb: (B,cond_dim) frozen CellDINO embedding. + Returns pixel reconstruction (B,1,H,W).""" + model.eval() + emb = model.cond(cell_emb.to(dev)) + + inv = DDIMInverseScheduler(num_train_timesteps=cfg.train_timesteps) + inv.set_timesteps(cfg.ddim_steps) + x = x0.to(dev) + for t in inv.timesteps: + eps = model.denoise(x, t, emb) + x = inv.step(eps, t, x).prev_sample + + fwd = DDIMScheduler(num_train_timesteps=cfg.train_timesteps) + fwd.set_timesteps(cfg.ddim_steps) + for t in fwd.timesteps: + eps = model.denoise(x, t, emb) + x = fwd.step(eps, t, x).prev_sample + return x + + +def _psnr(a: np.ndarray, b: np.ndarray) -> float: + mse = float(np.mean((a - b) ** 2)) + if mse <= 1e-12: + return 99.0 + return float(10.0 * np.log10(4.0 / mse)) # data range 2 ([-1,1]) -> 2^2=4 + + +@torch.no_grad() +def recon_report(model, images_norm: np.ndarray, embs: np.ndarray, cfg, dev, + out_png: Path, tag: str = "") -> float: + """Reconstruct cfg.n_recon cells, write an original-vs-recon montage, return PSNR.""" + n = min(cfg.n_recon, len(images_norm)) + x0 = torch.as_tensor(images_norm[:n]) + cell_emb = torch.as_tensor(embs[:n]) + rec = reconstruct(model, x0, cell_emb, cfg, dev).cpu().numpy() + orig = x0.numpy() + psnr = _psnr(orig, rec) + + import matplotlib + matplotlib.use("Agg") + matplotlib.rcParams["pdf.fonttype"] = 42 + import matplotlib.pyplot as plt + fig, axes = plt.subplots(2, n, figsize=(1.4 * n, 3.0), squeeze=False) + for i in range(n): + axes[0, i].imshow(orig[i, 0], cmap="gray", vmin=-1, vmax=1); axes[0, i].axis("off") + axes[1, i].imshow(rec[i, 0], cmap="gray", vmin=-1, vmax=1); axes[1, i].axis("off") + axes[0, 0].set_ylabel("orig", rotation=90); axes[1, 0].set_ylabel("recon", rotation=90) + fig.suptitle(f"DiffAE reconstruction {tag} PSNR={psnr:.2f} dB", fontsize=10) + fig.tight_layout() + Path(out_png).parent.mkdir(parents=True, exist_ok=True) + fig.savefig(out_png, dpi=130, bbox_inches="tight") + plt.close(fig) + print(f"[recon] {tag} PSNR={psnr:.2f} dB -> {out_png}") + return psnr diff --git a/src/ops_model/models/attention/diffex/diffae/run.py b/src/ops_model/models/attention/diffex/diffae/run.py new file mode 100644 index 0000000..61d913a --- /dev/null +++ b/src/ops_model/models/attention/diffex/diffae/run.py @@ -0,0 +1,83 @@ +"""Orchestrator for the DiffAE generator stage. + + python -m ops_model.models.attention.diffex.diffae.run + +Steps: sample broad phase crops (cached) -> normalize -> train DiffAE (joint +encoder + conditional UNet) -> periodic + final reconstruction gate. Writes +diffae_best.pt / diffae_last.pt, recon montages, metrics.json under . +""" +from __future__ import annotations + +import argparse +import json +from pathlib import Path + +from ..classifier.config import DEFAULT_OUT_ROOT +from .config import DiffAEConfig +from .data import load_diffae_crops +from .model import DiffAE +from .train import train_diffae + + +def run_diffae(cfg: DiffAEConfig, out_dir: str) -> dict: + out = Path(out_dir) + cache = out / "cache" + cache.mkdir(parents=True, exist_ok=True) + + cond_ch = getattr(cfg, "cond_channel", None) + spatial = getattr(cfg, "spatial_cond", False) + loaded = load_diffae_crops( + cfg, + crops_cache=str(cache / f"diffae_crops_{cfg.n_crops}_{cfg.crop_size}.npz"), + emb_cache=str(cache / f"diffae_celldino_{cfg.n_crops}_{cfg.crop_size}.npz"), + cond_cache=(str(cache / f"diffae_cond_{cond_ch}_{cfg.n_crops}_{cfg.crop_size}.npz") + if cond_ch else None), + return_cond_images=spatial, + ) + cond_imgs = None + if spatial: + images, embs, cond_imgs = loaded # phase images concatenated into the UNet input + else: + images, embs = loaded + model = DiffAE(cfg) + if getattr(cfg, "init_ckpt", None): # warm-start from an existing model + import torch + model.load_state_dict(torch.load(cfg.init_ckpt, map_location="cpu")) + print(f"warm-started weights from {cfg.init_ckpt}") + n_params = sum(p.numel() for p in model.parameters()) / 1e6 + print(f"DiffAE: {n_params:.1f}M params; cond_dim={embs.shape[1]}; " + f"training on {len(images)} crops") + + result = train_diffae(model, images, embs, cfg, out, cond_images=cond_imgs) + metrics = { + "n_crops": int(len(images)), "crop_size": cfg.crop_size, + "cond_dim": int(embs.shape[1]), "epochs": cfg.epochs, + "n_params_M": round(n_params, 1), + "best_cond_ratio": result["best_ratio"], # emb/noise control (target ≫ v1's 0.008) + } + (out / "metrics.json").write_text(json.dumps(metrics, indent=2)) + (out / "history.json").write_text(json.dumps(result["history"], indent=2)) + print(json.dumps(metrics, indent=2)) + return metrics + + +def main(): + ap = argparse.ArgumentParser(description="Train the DiffAE (DiffEx generator)") + ap.add_argument("--n-crops", type=int, default=50_000) + ap.add_argument("--crop-size", type=int, default=160) + ap.add_argument("--epochs", type=int, default=80) + ap.add_argument("--batch-size", type=int, default=32) + ap.add_argument("--device", default="cuda") + ap.add_argument("--num-workers", type=int, default=0) + ap.add_argument("--out-dir", default=f"{DEFAULT_OUT_ROOT}/diffae/phase_v1") + args = ap.parse_args() + + cfg = DiffAEConfig( + n_crops=args.n_crops, crop_size=args.crop_size, epochs=args.epochs, + batch_size=args.batch_size, device=args.device, num_workers=args.num_workers, + ) + run_diffae(cfg, args.out_dir) + + +if __name__ == "__main__": + main() diff --git a/src/ops_model/models/attention/diffex/diffae/submit.py b/src/ops_model/models/attention/diffex/diffae/submit.py new file mode 100644 index 0000000..ce3bea2 --- /dev/null +++ b/src/ops_model/models/attention/diffex/diffae/submit.py @@ -0,0 +1,109 @@ +"""Submit the DiffAE training to SLURM (1 GPU, longer wall clock). + + python -m ops_model.models.attention.diffex.diffae.submit + +=============================== RUNBOOK =============================== +Checkpoints (root /hpc/projects/icd.fast.ops/models/diffex/diffae/) and their +best cond_ratio (emb/noise conditioning strength; higher = stronger edits): + phase_v1/ 50k crops, ep120, 0.468 <- PRODUCTION (all traversals use this) + phase_v1_500k/ 500k scratch, ep12, 0.416 (undertrained; parked) + phase_v1_500k_warm/ 500k warm-from-v1, 0.542 (more-data test; being resumed) + +RESUME an existing run (continue where it stopped): resume=True is the config +default and train state (model+ema+opt+epoch) is saved EVERY epoch to +/diffae_train_state.pt. Just re-submit the SAME --out-dir/--n-crops and it +picks up automatically. Do NOT pass --init-ckpt on a resume (train_state wins). + +MEMORY GOTCHA (500k): the 500k crop cache is 51 GB float32 and load_diffae_crops +normalizes it -> ~102 GB transient peak. Use mem_gb>=200 for 500k runs; 96 GB +OOM-kills (esp. on a shared node). 50k runs are fine at 96 GB. + +CHAIN across the 720-min walltime: submit N jobs, each with +slurm_additional_parameters={"dependency": f"afterany:"} (see --after). +afterany fires even on failure, so verify link 0 clears the cache-load before +trusting the chain. + +WATCH cond_ratio trend: + python -c "import torch;h=torch.load('/diffae_train_state.pt',map_location='cpu')['history'];print([round(e.get('cond_ratio',-1),3) for e in h if e.get('cond_ratio',-1)>0])" +====================================================================== +""" +from __future__ import annotations + +import argparse +from pathlib import Path + +from ops_utils.hpc.slurm_batch_utils import submit_parallel_jobs + +from ..classifier.config import DEFAULT_OUT_ROOT +from .config import DiffAEConfig +from .run import run_diffae + + +def main(): + ap = argparse.ArgumentParser(description="Submit DiffAE training to SLURM") + ap.add_argument("--n-crops", type=int, default=50_000) + ap.add_argument("--crop-size", type=int, default=160) + ap.add_argument("--epochs", type=int, default=80) + ap.add_argument("--batch-size", type=int, default=32) + ap.add_argument("--out-dir", default=f"{DEFAULT_OUT_ROOT}/diffae/phase_v1") + ap.add_argument("--partition", default="gpu") + ap.add_argument("--gpus", type=int, default=1) + # bracketed constraint on the gpu partition = fast ≥80GB GPUs, non-preemptible + # (repo pattern in ops_process slurm_task_config.yaml). All fit batch 48. + ap.add_argument("--constraint", default="[a100_80|h100|h200|6000_blackwell]") + ap.add_argument("--cpus", type=int, default=8) + ap.add_argument("--mem-gb", type=int, default=96) + ap.add_argument("--time-min", type=int, default=720) + ap.add_argument("--after", default=None, + help="SLURM job id: start afterany: (resume=True continues training → auto-resubmit chains)") + ap.add_argument("--name", default="diffae_phase_v1") + ap.add_argument("--augment-affine", action="store_true", + help="continuous rotation+scale+flip aug (else discrete dihedral)") + ap.add_argument("--no-aug", action="store_true", + help="disable ALL augmentation (true v1-style, no dihedral/affine)") + ap.add_argument("--init-ckpt", default=None, + help="warm-start: load these weights into the fresh model before training") + ap.add_argument("--marker-channel", default=None, + help="fluor mode: fluor-CSV `channel` value (e.g. 'nucleolus-GC_NPM3')") + ap.add_argument("--channel", default="Phase2D", + help="raw pheno-zarr channel to read (Phase2D | GFP | mCherry | Cy5)") + ap.add_argument("--cond-channel", default=None, + help="virtual staining: condition on this raw channel (e.g. Phase2D) while generating --channel") + ap.add_argument("--spatial-cond", action="store_true", + help="concat the cond-channel IMAGE into the UNet (pixel-registered stain); requires --cond-channel") + ap.add_argument("--dry-run", action="store_true") + args = ap.parse_args() + + affine = args.augment_affine and not args.no_aug + dihedral = (not args.augment_affine) and not args.no_aug + cfg = DiffAEConfig( + n_crops=args.n_crops, crop_size=args.crop_size, epochs=args.epochs, + batch_size=args.batch_size, device="cuda", + augment_affine=affine, augment_dihedral=dihedral, init_ckpt=args.init_ckpt, + marker_channel=args.marker_channel, channel=args.channel, cond_channel=args.cond_channel, + spatial_cond=args.spatial_cond, + ) + jobs = [{ + "name": args.name, + "func": run_diffae, + "kwargs": {"cfg": cfg, "out_dir": str(Path(args.out_dir).resolve())}, + "metadata": {"stage": "diffae"}, + }] + slurm_params = { + "slurm_partition": args.partition, "gpus_per_node": args.gpus, + "cpus_per_task": args.cpus, "mem_gb": args.mem_gb, "timeout_min": args.time_min, + "slurm_setup": ["export PYTORCH_CUDA_ALLOC_CONF=expandable_segments:True"], + } + if args.constraint: + slurm_params["slurm_constraint"] = args.constraint + if args.after: # resume-chain: wait for the prior job to end (any reason), then continue + slurm_params["slurm_additional_parameters"] = {"dependency": f"afterany:{args.after}"} + submit_parallel_jobs( + jobs_to_submit=jobs, experiment="diffae", + slurm_params=slurm_params, log_dir="diffae", + dry_run=args.dry_run, wait_for_completion=False, + ) + + +if __name__ == "__main__": + main() diff --git a/src/ops_model/models/attention/diffex/diffae/train.py b/src/ops_model/models/attention/diffex/diffae/train.py new file mode 100644 index 0000000..8a2ed54 --- /dev/null +++ b/src/ops_model/models/attention/diffex/diffae/train.py @@ -0,0 +1,272 @@ +"""DiffAE training — proper conditional-diffusion recipe. + +Fixes the v1 failure (embedding had ~0.8% control). Essential components: +- **Conditioning dropout**: replace the embedding with a learned null ~cond_dropout of + the time, forcing the model to actually USE the embedding (and enabling CFG). +- **EMA** weights for sampling/eval (near-essential for diffusion quality). +- **Resume** across 12h jobs (long training to convergence). +- **Conditioning gate** (not recon PSNR!): generate from fixed noise under two different + embeddings and report emb-vs-noise MSE ratio — this is the metric that must climb. +""" +from __future__ import annotations + +import copy +from pathlib import Path + +import numpy as np +import torch +import torch.nn as nn +from diffusers import DDIMScheduler, DDPMScheduler +from torch.utils.data import DataLoader, TensorDataset + + +def _device(name: str) -> torch.device: + if name.startswith("cuda") and not torch.cuda.is_available(): + print("[device] cuda unavailable -> cpu") + return torch.device("cpu") + return torch.device(name) + + +def _augment(x: torch.Tensor, scale_jit: float = 0.15) -> torch.Tensor: + """Per-sample continuous rotation (±180°) + scale + horizontal flip, with REFLECTION + padding so arbitrary angles introduce no black corners. Matches the contrastive + data_loader recipe (RandAffine ±π, scale) and — unlike the discrete dihedral — makes + orientation a nuisance across ALL angles, not just 90° steps. Embedding is NOT + recomputed. x: (B,1,H,W).""" + import math + B = x.shape[0]; dev = x.device + fl = torch.rand(B, device=dev) < 0.5 + x = torch.where(fl[:, None, None, None], torch.flip(x, dims=(-1,)), x) + ang = (torch.rand(B, device=dev) * 2 - 1) * math.pi + s = 1.0 / (1.0 + (torch.rand(B, device=dev) * 2 - 1) * scale_jit) # inverse for sampling grid + cos, sin = torch.cos(ang) * s, torch.sin(ang) * s + theta = torch.zeros(B, 2, 3, device=dev, dtype=x.dtype) + theta[:, 0, 0], theta[:, 0, 1] = cos, -sin + theta[:, 1, 0], theta[:, 1, 1] = sin, cos + grid = torch.nn.functional.affine_grid(theta, x.shape, align_corners=False) + return torch.nn.functional.grid_sample(x, grid, mode="bilinear", + padding_mode="reflection", align_corners=False) + + +def _dihedral(x: torch.Tensor) -> torch.Tensor: + """Per-sample random dihedral transform (4 rot90 × flip = 8 orientations, incl. + transpose). x: (B,1,H,W). The conditioning embedding is NOT recomputed — orientation + is made a nuisance the model must push into the noise latent, not the embedding.""" + ks = torch.randint(0, 4, (x.shape[0],)) + fl = torch.rand(x.shape[0]) < 0.5 + out = torch.empty_like(x) + for i in range(x.shape[0]): + xi = torch.rot90(x[i], int(ks[i]), dims=(-2, -1)) + out[i] = torch.flip(xi, dims=(-1,)) if fl[i] else xi + return out + + +def _pearson(a, b) -> float: + a, b = a.ravel(), b.ravel() + a, b = a - a.mean(), b - b.mean() + return float((a * b).sum() / (np.sqrt((a * a).sum() * (b * b).sum()) + 1e-12)) + + +@torch.no_grad() +def _sample(model, xT, emb, cfg, cond_img=None): + fwd = DDIMScheduler(num_train_timesteps=cfg.train_timesteps) + fwd.set_timesteps(cfg.ddim_steps) + c = model.cond(emb) + x = xT + for t in fwd.timesteps: + x = fwd.step(model.denoise(x, t, c, cond_img), t, x).prev_sample + return x + + +@torch.no_grad() +def recon_pearson_gate(model, probe_x, probe_emb, probe_cond, cfg, dev, out_png, tag="") -> float: + """Spatial-conditioning selector: generate the marker from (phase image + emb) and report + mean Pearson(pred, real) on held-out probe cells + a phase|pred|real montage.""" + model.eval() + H = cfg.crop_size + k = min(8, len(probe_x)) + corrs, preds = [], [] + for i in range(k): + g = torch.Generator(device=dev).manual_seed(7 + i) + xT = torch.randn(1, 1, H, H, generator=g, device=dev) + e = torch.as_tensor(probe_emb[i:i + 1], dtype=torch.float32, device=dev) + ci = torch.as_tensor(probe_cond[i:i + 1], dtype=torch.float32, device=dev) + pred = _sample(model, xT, e, cfg, cond_img=ci).cpu().numpy()[0, 0] + preds.append(pred); corrs.append(_pearson(pred, probe_x[i, 0])) + mean_r = float(np.mean(corrs)) + import matplotlib + matplotlib.use("Agg"); matplotlib.rcParams["pdf.fonttype"] = 42 + import matplotlib.pyplot as plt + Path(out_png).parent.mkdir(parents=True, exist_ok=True) + fig, ax = plt.subplots(k, 3, figsize=(5.2, 1.7 * k), squeeze=False) + for i in range(k): + for c, (img, cm) in enumerate([(probe_cond[i, 0], "gray"), (preds[i], "magma"), (probe_x[i, 0], "magma")]): + ax[i, c].imshow(img, cmap=cm, vmin=-1, vmax=1); ax[i, c].axis("off") + ax[i, 0].set_ylabel(f"r={corrs[i]:.2f}", fontsize=8, rotation=0, labelpad=16) + for c, t in enumerate(["phase", "pred", "real"]): + ax[0, c].set_title(t, fontsize=9) + fig.suptitle(f"recon {tag}: Pearson {mean_r:.3f}", fontsize=10) + fig.tight_layout(); fig.savefig(out_png, dpi=130, bbox_inches="tight"); plt.close(fig) + print(f"[recon-gate] {tag} Pearson(pred,real)={mean_r:.3f}") + return mean_r + + +@torch.no_grad() +def conditioning_gate(model, probe_embs, cfg, dev, out_png, tag="") -> float: + """emb-vs-noise MSE ratio: how much the embedding controls generation. >~0.3 = real + control; the v1 model was 0.008. Generates from fixed noise under different embeddings.""" + model.eval() + H = cfg.crop_size + k = max(1, min(4, len(probe_embs) // 2)) + ea = torch.as_tensor(probe_embs[:k], dtype=torch.float32, device=dev) + eb = torch.as_tensor(probe_embs[k:2 * k], dtype=torch.float32, device=dev) + A, B, A2 = [], [], [] + for i in range(k): + g = torch.Generator(device=dev).manual_seed(7 + i) + xT = torch.randn(1, 1, H, H, generator=g, device=dev) + A.append(_sample(model, xT.clone(), ea[i:i + 1], cfg).cpu().numpy()[0, 0]) + B.append(_sample(model, xT.clone(), eb[i:i + 1], cfg).cpu().numpy()[0, 0]) + g2 = torch.Generator(device=dev).manual_seed(999 + i) + xT2 = torch.randn(1, 1, H, H, generator=g2, device=dev) + A2.append(_sample(model, xT2.clone(), ea[i:i + 1], cfg).cpu().numpy()[0, 0]) + A, B, A2 = np.array(A), np.array(B), np.array(A2) + emb_mse = float(np.mean((A - B) ** 2)) # same noise, different embedding + noise_mse = float(np.mean((A - A2) ** 2)) # different noise, same embedding + ratio = emb_mse / (noise_mse + 1e-9) + + import matplotlib + matplotlib.use("Agg"); matplotlib.rcParams["pdf.fonttype"] = 42 + import matplotlib.pyplot as plt + Path(out_png).parent.mkdir(parents=True, exist_ok=True) + fig, ax = plt.subplots(k, 3, figsize=(5.5, 1.7 * k), squeeze=False) + for i in range(k): + ax[i, 0].imshow(A[i], cmap="gray", vmin=-1, vmax=1) + ax[i, 1].imshow(B[i], cmap="gray", vmin=-1, vmax=1) + d = np.abs(A[i] - B[i]); ax[i, 2].imshow(d, cmap="hot", vmin=0, vmax=max(d.max(), 1e-3)) + for j in range(3): + ax[i, j].axis("off") + if i == 0: + for j, t in enumerate(["emb A", "emb B", "|A−B|"]): + ax[i, j].set_title(t, fontsize=9) + fig.suptitle(f"conditioning gate {tag}: emb/noise ratio={ratio:.3f}", fontsize=10) + fig.tight_layout(); fig.savefig(out_png, dpi=130, bbox_inches="tight"); plt.close(fig) + print(f"[gate] {tag} emb/noise ratio={ratio:.3f} (emb_mse={emb_mse:.4f} noise_mse={noise_mse:.4f})") + return ratio + + +def train_diffae(model, images_norm: np.ndarray, embs: np.ndarray, cfg, out_dir: Path, + cond_images: np.ndarray | None = None) -> dict: + dev = _device(cfg.device) + spatial = cond_images is not None # image-to-image: phase concatenated into the UNet input + model = model.to(dev) + ema = copy.deepcopy(model).eval() + for p in ema.parameters(): + p.requires_grad_(False) + sched = DDPMScheduler(num_train_timesteps=cfg.train_timesteps) + opt = torch.optim.AdamW(model.parameters(), lr=cfg.lr, weight_decay=cfg.weight_decay) + crit = nn.MSELoss() + + n_probe = min(8, len(images_norm) // 10 or 8) + probe_emb = embs[-n_probe:] + probe_x = images_norm[-n_probe:] + probe_cond = cond_images[-n_probe:] if spatial else None + train_x, train_e = images_norm[:-n_probe], embs[:-n_probe] + if spatial: + loader = DataLoader( + TensorDataset(torch.as_tensor(train_x), torch.as_tensor(train_e), + torch.as_tensor(cond_images[:-n_probe])), + batch_size=cfg.batch_size, shuffle=True, drop_last=True, + ) + else: + loader = DataLoader( + TensorDataset(torch.as_tensor(train_x), torch.as_tensor(train_e)), + batch_size=cfg.batch_size, shuffle=True, drop_last=True, + ) + use_amp = dev.type == "cuda" + scaler = torch.cuda.amp.GradScaler(enabled=use_amp) + out_dir = Path(out_dir) + (out_dir / "gate").mkdir(parents=True, exist_ok=True) + + # NOTE: nn.DataParallel is incompatible with the diffusers UNet (its `.dtype` + # property breaks on DP replicas). Multi-GPU would need DDP. Single-GPU here. + train_model = model + + @torch.no_grad() + def ema_update(): + for e, p in zip(ema.parameters(), model.parameters()): + e.mul_(cfg.ema_decay).add_(p.detach(), alpha=1 - cfg.ema_decay) + for eb, pb in zip(ema.buffers(), model.buffers()): + eb.copy_(pb) + + # resume + state_path = out_dir / "diffae_train_state.pt" + start_epoch, history, best_ratio = 0, [], -1.0 + if cfg.resume and state_path.exists(): + st = torch.load(state_path, map_location=dev) + model.load_state_dict(st["model"]); ema.load_state_dict(st["ema"]) + opt.load_state_dict(st["opt"]); start_epoch = st["epoch"] + 1 + history = st.get("history", []); best_ratio = st.get("best_ratio", -1.0) + print(f"[resume] from epoch {start_epoch} (best ratio {best_ratio:.3f})") + + for ep in range(start_epoch, cfg.epochs): + model.train() + ep_loss = 0.0 + for batch in loader: + x, e = batch[0].to(dev), batch[1].to(dev) + ci = batch[2].to(dev) if spatial else None + # augment target only (autoencoding) or target+cond jointly (spatial → stay registered) + if getattr(cfg, "augment_affine", False): # continuous rotation+scale+flip + if spatial: + xc = _augment(torch.cat([x, ci], 1), getattr(cfg, "affine_scale", 0.15)); x, ci = xc[:, :1], xc[:, 1:] + else: + x = _augment(x, getattr(cfg, "affine_scale", 0.15)) + elif getattr(cfg, "augment_dihedral", False): # discrete 90°×flip + if spatial: + xc = _dihedral(torch.cat([x, ci], 1)); x, ci = xc[:, :1], xc[:, 1:] + else: + x = _dihedral(x) + if cfg.cond_dropout > 0: # conditioning dropout (emb only) + drop = torch.rand(e.shape[0], device=dev) < cfg.cond_dropout + if drop.any(): + e = torch.where(drop[:, None], model.null_emb[None].to(e.dtype), e) + noise = torch.randn_like(x) + t = torch.randint(0, cfg.train_timesteps, (x.shape[0],), device=dev).long() + noisy = sched.add_noise(x, noise, t) + opt.zero_grad() + with torch.autocast(device_type="cuda", enabled=use_amp): + loss = crit(train_model(noisy, t, e, ci), noise) + scaler.scale(loss).backward(); scaler.step(opt); scaler.update() + ema_update() + ep_loss += float(loss) * x.shape[0] + ep_loss /= len(loader.dataset) + + rec = {"epoch": ep, "loss": ep_loss} + if (ep + 1) % cfg.recon_every == 0 or ep == cfg.epochs - 1: + try: + if spatial: # select on Pearson(pred, real) — the metric that matters for a stain + metric = recon_pearson_gate(ema, probe_x, probe_emb, probe_cond, cfg, dev, + out_dir / "gate" / f"recon_ep{ep:03d}.png", tag=f"ep{ep}") + rec["recon_pearson"] = metric + else: # emb-vs-noise conditioning strength + metric = conditioning_gate(ema, probe_emb, cfg, dev, + out_dir / "gate" / f"gate_ep{ep:03d}.png", tag=f"ep{ep}") + rec["cond_ratio"] = metric + if metric > best_ratio: + best_ratio = metric + torch.save(ema.state_dict(), out_dir / "diffae_best.pt") # EMA = model to use + except Exception as exc: # noqa: BLE001 + print(f"[warn] gate at epoch {ep} failed (continuing): {exc}") + history.append(rec) + # save resume state EVERY epoch (preemption resilience on contended GPUs) + try: + torch.save({"model": model.state_dict(), "ema": ema.state_dict(), + "opt": opt.state_dict(), "epoch": ep, "history": history, + "best_ratio": best_ratio}, state_path) + except Exception as exc: # noqa: BLE001 + print(f"[warn] state save at epoch {ep} failed: {exc}") + print(f"epoch {ep:03d}: loss={ep_loss:.4f}" + + (f" cond_ratio={rec['cond_ratio']:.3f}" if "cond_ratio" in rec else "") + + (f" recon_pearson={rec['recon_pearson']:.3f}" if "recon_pearson" in rec else "")) + + torch.save(ema.state_dict(), out_dir / "diffae_ema_last.pt") + return {"best_ratio": best_ratio, "history": history} diff --git a/src/ops_model/models/attention/diffex/diffae/virtstain_eval.py b/src/ops_model/models/attention/diffex/diffae/virtstain_eval.py new file mode 100644 index 0000000..0800deb --- /dev/null +++ b/src/ops_model/models/attention/diffex/diffae/virtstain_eval.py @@ -0,0 +1,133 @@ +"""Evaluate a virtual-staining DiffAE: predict the fluor marker from the phase CellDINO +embedding on a HELD-OUT set of cells (fresh seed → disjoint from training), then report +Pearson(pred, real) and save a `phase | predicted | real` montage. + + python -m ops_model.models.attention.diffex.diffae.virtstain_eval \ + --out-dir /hpc/projects/icd.fast.ops/analysis/virtual_staining/chromalive561_from_phase \ + --marker-channel "mitochondria_ChromaLIVE 561 excitation" --channel mCherry --cond-channel Phase2D +""" +from __future__ import annotations + +import argparse +import dataclasses +import json +from pathlib import Path + +import numpy as np +import torch + +from .config import DiffAEConfig +from .data import load_diffae_crops +from .model import DiffAE +from .train import _sample + + +def _pearson(a: np.ndarray, b: np.ndarray) -> float: + a, b = a.ravel(), b.ravel() + a, b = a - a.mean(), b - b.mean() + d = np.sqrt((a * a).sum() * (b * b).sum()) + 1e-12 + return float((a * b).sum() / d) + + +@torch.no_grad() +def evaluate(cfg: DiffAEConfig, out_dir: str, ckpt: str, n_eval: int, eval_seed: int): + out = Path(out_dir) + dev = torch.device(cfg.device if torch.cuda.is_available() else "cpu") + model = DiffAE(cfg).to(dev) + model.load_state_dict(torch.load(ckpt, map_location=dev)); model.eval() + + # fresh held-out cells (disjoint seed) — target marker + conditioning phase crops + ecfg = dataclasses.replace(cfg, n_crops=n_eval, seed=eval_seed) + cache = out / "cache_eval"; cache.mkdir(parents=True, exist_ok=True) + real, embs, phase = load_diffae_crops( + ecfg, + crops_cache=str(cache / f"marker_{n_eval}_{cfg.crop_size}_s{eval_seed}.npz"), + emb_cache=str(cache / f"phasecelldino_{n_eval}_{cfg.crop_size}_s{eval_seed}.npz"), + cond_cache=str(cache / f"phase_{n_eval}_{cfg.crop_size}_s{eval_seed}.npz"), + return_cond_images=True, + ) + H = cfg.crop_size + spatial = getattr(cfg, "spatial_cond", False) + corrs, preds = [], [] + for i in range(len(real)): + g = torch.Generator(device=dev).manual_seed(1000 + i) + xT = torch.randn(1, 1, H, H, generator=g, device=dev) + e = torch.as_tensor(embs[i:i + 1], dtype=torch.float32, device=dev) + ci = torch.as_tensor(phase[i:i + 1], dtype=torch.float32, device=dev) if spatial else None + pred = _sample(model, xT, e, cfg, cond_img=ci).cpu().numpy()[0, 0] + preds.append(pred) + corrs.append(_pearson(pred, real[i, 0])) + corrs = np.array(corrs) + # trivial baseline: does the phase image itself already correlate with the marker? + base = np.array([_pearson(phase[i, 0], real[i, 0]) for i in range(len(real))]) + metrics = {"n_eval": int(len(real)), "eval_seed": eval_seed, + "pearson_pred_vs_real_mean": round(float(corrs.mean()), 4), + "pearson_pred_vs_real_std": round(float(corrs.std()), 4), + "pearson_phase_vs_real_mean": round(float(base.mean()), 4), + "marker_channel": cfg.marker_channel, "channel": cfg.channel, + "cond_channel": cfg.cond_channel, "ckpt": ckpt} + (out / "eval").mkdir(parents=True, exist_ok=True) + (out / "eval" / "virtstain_metrics.json").write_text(json.dumps(metrics, indent=2)) + print(json.dumps(metrics, indent=2)) + + # montage: top cells by Pearson (best-case read), phase | pred | real + import matplotlib + matplotlib.use("Agg"); matplotlib.rcParams["pdf.fonttype"] = 42 + import matplotlib.pyplot as plt + import matplotlib.patheffects as pe + order = np.argsort(-corrs)[:12] + rows = [("phase (input)", "gray", lambda i: phase[i, 0]), + ("predicted", "magma", lambda i: preds[i]), + ("real", "magma", lambda i: real[i, 0])] + fig, ax = plt.subplots(3, len(order), figsize=(1.4 * len(order), 4.6), squeeze=False) + for c, i in enumerate(order): + for r, (_, cm, get) in enumerate(rows): + ax[r, c].imshow(get(i), cmap=cm, vmin=-1, vmax=1) + ax[r, c].set_xticks([]); ax[r, c].set_yticks([]) + ax[r, c].text(0.04, 0.96, f"r={corrs[i]:.2f}", transform=ax[r, c].transAxes, + fontsize=6.5, color="white", va="top", ha="left", + path_effects=[pe.withStroke(linewidth=1.5, foreground="black")]) + for r, (label, _, _) in enumerate(rows): + ax[r, 0].set_ylabel(label, fontsize=10) + fig.suptitle(f"virtual staining {cfg.marker_channel} | Pearson {corrs.mean():.3f}±{corrs.std():.3f} " + f"(phase-baseline {base.mean():.3f})", fontsize=9) + fig.tight_layout(); fig.savefig(out / "eval" / "virtstain_montage.png", dpi=140, bbox_inches="tight") + plt.close(fig) + print(f"[eval] wrote {out/'eval'/'virtstain_montage.png'}") + return metrics + + +def main(): + ap = argparse.ArgumentParser(description="Evaluate virtual-staining DiffAE (Pearson + montage)") + ap.add_argument("--out-dir", required=True) + ap.add_argument("--ckpt", default=None, help="default /diffae_best.pt") + ap.add_argument("--marker-channel", required=True) + ap.add_argument("--channel", default="mCherry") + ap.add_argument("--cond-channel", default="Phase2D") + ap.add_argument("--spatial-cond", action="store_true", help="model uses image-concat conditioning") + ap.add_argument("--crop-size", type=int, default=160) + ap.add_argument("--n-eval", type=int, default=256) + ap.add_argument("--eval-seed", type=int, default=12345) + ap.add_argument("--device", default="cuda") + ap.add_argument("--submit", action="store_true", help="run on SLURM GPU instead of locally") + ap.add_argument("--after", default=None, help="SLURM job id to gate on (afterany)") + args = ap.parse_args() + cfg = DiffAEConfig(crop_size=args.crop_size, channel=args.channel, cond_channel=args.cond_channel, + spatial_cond=args.spatial_cond, marker_channel=args.marker_channel, device=args.device) + ckpt = args.ckpt or str(Path(args.out_dir) / "diffae_best.pt") + if args.submit: + from ops_utils.hpc.slurm_batch_utils import submit_parallel_jobs + sp = {"slurm_partition": "gpu", "gpus_per_node": 1, "cpus_per_task": 8, "mem_gb": 96, + "timeout_min": 60, "slurm_constraint": "[a100_80|h100|h200|6000_blackwell]"} + if args.after: + sp["slurm_additional_parameters"] = {"dependency": f"afterany:{args.after}"} + submit_parallel_jobs(jobs_to_submit=[{"name": "virtstain_eval", "func": evaluate, + "kwargs": {"cfg": cfg, "out_dir": args.out_dir, "ckpt": ckpt, + "n_eval": args.n_eval, "eval_seed": args.eval_seed}}], + experiment="diffae", slurm_params=sp, log_dir="diffae", wait_for_completion=False) + else: + evaluate(cfg, args.out_dir, ckpt, args.n_eval, args.eval_seed) + + +if __name__ == "__main__": + main() diff --git a/src/ops_model/models/attention/diffex/diffae/virtstain_multi.py b/src/ops_model/models/attention/diffex/diffae/virtstain_multi.py new file mode 100644 index 0000000..b7c6eb7 --- /dev/null +++ b/src/ops_model/models/attention/diffex/diffae/virtstain_multi.py @@ -0,0 +1,321 @@ +"""Multi-marker virtual staining: ONE spatial DiffAE predicts every fluor marker from phase, +switched by a learned marker-id embedding. Pools paired (phase, marker) crops across all markers +(each marker's own exps); the phase image is concatenated into the UNet (registered stain) and the +marker id selects which channel to render. + + python -m ops_model.models.attention.diffex.diffae.virtstain_multi --submit --cap 2500 --epochs 120 +""" +from __future__ import annotations + +import argparse +import copy +import json +from pathlib import Path + +import numpy as np +import pandas as pd +import torch +import torch.nn as nn +from diffusers import DDIMScheduler, DDPMScheduler +from torch.utils.data import DataLoader, TensorDataset + +from ..classifier.config import slugify +from .config import DiffAEConfig +from .data import load_diffae_crops +from .model import DiffAE +from .train import _pearson + +OUT = "/hpc/projects/icd.fast.ops/analysis/virtual_staining/multi_marker" + + +def markers_list(): + """LIVE fluor markers only: (dir, marker_channel, raw_channel), deduped. 4i/CP markers (raw channel + CP1_/CP2_/4i_R*) are stained AFTER the live phase acquisition, so their cells have moved/changed and + the phase→marker registration is broken — that misalignment poisons the spatial conditioning, so they + are excluded. Live channels (GFP/mCherry/Cy5/farred) are imaged concurrently with phase → registered.""" + from ..viewer import catalog as C + seen, out = set(), [] + for d, mc, ch in C.complete_markers(): + if ch.startswith(("CP", "4i")): + continue + if mc not in seen: + seen.add(mc); out.append((d, mc, ch)) + return out + + +# Every scored cell per marker (concat of Alex's 6 per-cell shards, incl NTC) — NOT the acc>0.5 qualifying +# subset. Virtual staining needs a representative pool of paired (phase, marker) crops, so no distinctiveness +# filtering: use all cells (74k–1.3M/marker) and sample ≤cap per marker. Built by scratchpad/build_allcells.py. +ALL_CELLS = ("/hpc/projects/icd.fast.ops/models/alex_lin_attention/v5/fluorescence/" + "misc/all_cells_bychannel.parquet") + + +def load_multi(markers, cap, cache_root): + """Per marker: gather ≤cap paired (phase, marker) crops from ALL_CELLS (every cell, no accuracy filter), + materialize marker (target) + phase (cond) crops, CellDINO-embed phase; cache per marker.""" + Path(cache_root).mkdir(parents=True, exist_ok=True) + pre = pd.read_parquet(ALL_CELLS, columns=["channel", "gene", "experiment", "well", + "segmentation", "x_pheno", "y_pheno"]) + Xs, Es, Ps, Ms, kept = [], [], [], [], [] + for d, mc, ch in markers: + sl = slugify(mc) + rows = pre[pre["channel"] == mc] + if rows.empty: + print(f"[skip] {mc}: no rows in v5 qualifying"); continue + try: + cfg = DiffAEConfig(marker_channel=mc, channel=ch, cond_channel="Phase2D", + spatial_cond=True, n_crops=cap, seed=len(kept)) + cfg._fluor_rows = rows + x, e, p = load_diffae_crops( + cfg, crops_cache=f"{cache_root}/{sl}_marker.npz", + emb_cache=f"{cache_root}/{sl}_emb.npz", cond_cache=f"{cache_root}/{sl}_phase.npz", + return_cond_images=True) + except Exception as exc: # noqa: BLE001 — skip finicky markers, keep going + print(f"[skip] {mc}: {exc}"); continue + mid = len(kept); kept.append((d, mc, ch)) + Xs.append(x); Es.append(e); Ps.append(p); Ms.append(np.full(len(x), mid, np.int64)) + print(f"[marker {mid}] {mc}: {len(x)} crops") + X = np.concatenate(Xs); E = np.concatenate(Es); P = np.concatenate(Ps); M = np.concatenate(Ms) + print(f"[multi] {len(kept)} markers, {len(X)} total crops") + return X, E, P, M, kept + + +@torch.no_grad() +def _sample_marker(model, xT, emb, cond_img, marker_id, cfg, dev): + fwd = DDIMScheduler(num_train_timesteps=cfg.train_timesteps); fwd.set_timesteps(cfg.ddim_steps) + c = model.cond(emb, marker_id); x = xT + for t in fwd.timesteps: + x = fwd.step(model.denoise(x, t, c, cond_img), t, x).prev_sample + return x + + +def train_multi(X, E, P, M, cfg, out_dir, epochs, batch, device): + dev = torch.device(device if torch.cuda.is_available() else "cpu") + model = DiffAE(cfg).to(dev) + ema = copy.deepcopy(model).eval() + for pr in ema.parameters(): + pr.requires_grad_(False) + sched = DDPMScheduler(num_train_timesteps=cfg.train_timesteps) + opt = torch.optim.AdamW(model.parameters(), lr=cfg.lr) + crit = nn.MSELoss() + n_probe = min(64, len(X) // 20) + loader = DataLoader(TensorDataset(torch.as_tensor(X[:-n_probe]), torch.as_tensor(E[:-n_probe]), + torch.as_tensor(P[:-n_probe]), torch.as_tensor(M[:-n_probe])), + batch_size=batch, shuffle=True, drop_last=True) + scaler = torch.cuda.amp.GradScaler(enabled=dev.type == "cuda") + out = Path(out_dir); out.mkdir(parents=True, exist_ok=True) + state = out / "train_state.pt"; start = 0 + + def ema_up(): + for e_, p_ in zip(ema.parameters(), model.parameters()): + e_.mul_(cfg.ema_decay).add_(p_.detach(), alpha=1 - cfg.ema_decay) + for eb, pb in zip(ema.buffers(), model.buffers()): + eb.copy_(pb) + + if state.exists(): + st = torch.load(state, map_location=dev) + model.load_state_dict(st["model"]); ema.load_state_dict(st["ema"]); opt.load_state_dict(st["opt"]) + start = st["epoch"] + 1; print(f"[resume] from epoch {start}") + for ep in range(start, epochs): + model.train(); tot = 0.0 + for x, e, p, m in loader: + x, e, p, m = x.to(dev), e.to(dev), p.to(dev), m.to(dev) + if cfg.cond_dropout > 0: # drop the CellDINO emb (keep marker id) + drop = torch.rand(e.shape[0], device=dev) < cfg.cond_dropout + if drop.any(): + e = torch.where(drop[:, None], model.null_emb[None].to(e.dtype), e) + noise = torch.randn_like(x) + t = torch.randint(0, cfg.train_timesteps, (x.shape[0],), device=dev).long() + noisy = sched.add_noise(x, noise, t); opt.zero_grad() + with torch.autocast("cuda", enabled=dev.type == "cuda"): + loss = crit(model(noisy, t, e, p, m), noise) + scaler.scale(loss).backward(); scaler.step(opt); scaler.update(); ema_up() + tot += float(loss) * x.shape[0] + tot /= len(loader.dataset) + torch.save({"model": model.state_dict(), "ema": ema.state_dict(), "opt": opt.state_dict(), "epoch": ep, "loss": tot}, state) + if (ep + 1) % 10 == 0 or ep == epochs - 1: + torch.save(ema.state_dict(), out / "diffae_best.pt") + print(f"epoch {ep:03d}: loss={tot:.4f}", flush=True) + torch.save(ema.state_dict(), out / "diffae_best.pt") + return ema + + +def _plot_montage(rowspecs, n_kept, epoch, loss, ncell, out, eval_name): + """Render the per-marker phase/pred/real montage from cached rowspecs (no GPU). Each marker is a + phase/pred/real column; markers pack across PER-per-section, wrapping into stacked sections. Nested + gridspec keeps phase/pred/real tight, with a big gap BETWEEN sections for the 3-line title.""" + import matplotlib + matplotlib.use("Agg"); matplotlib.rcParams["pdf.fonttype"] = 42 + import matplotlib.pyplot as plt + r_by = {name: r for name, _, _, _, r in rowspecs} + mean_r = float(np.mean(list(r_by.values()))) if r_by else 0.0 + n = len(rowspecs) + NSEC = 3 if n > 16 else 1 # ~3 long horizontal sections for the full panel + PER = max(1, -(-n // NSEC)) # markers/section (ceil) → long rows + nb = -(-n // PER) + figH = nb * 3 * 1.28 + fig = plt.figure(figsize=(PER * 1.05, figH)) + top_frac = 1 - 0.85 / figH # reserve a fixed band at the top for the 2-line suptitle + # Nested gridspec: big gap BETWEEN sections (room for the 3-line title), tight WITHIN each phase/pred/real triplet. + outer = fig.add_gridspec(nb, PER, hspace=0.12, wspace=0.04, top=top_frac, bottom=0.015) + axmap = {} + for k, (name, ph, pr, re, rr) in enumerate(rowspecs): + band, col = divmod(k, PER) + inner = outer[band, col].subgridspec(3, 1, hspace=0.03) # phase/pred/real stay tight together + for j, (img, cm) in enumerate([(ph, "gray"), (pr, "magma"), (re, "magma")]): + a = fig.add_subplot(inner[j]) + a.imshow(img, cmap=cm, vmin=-1, vmax=1, aspect="auto"); a.set_xticks([]); a.set_yticks([]) + axmap[(band, col, j)] = a + org, _, prot = name.rpartition("_") # 3-line title: organelle / protein / Pearson + if not org: + org, prot = prot, "" + axmap[(band, col, 0)].set_title(f"{org}\n{prot}\nr={rr:.2f}", fontsize=5.5, pad=3, linespacing=1.25) + for band in range(nb): # phase/pred/real labels on the left of each section + for j, lbl in enumerate(["phase", "pred", "real"]): + a = axmap.get((band, 0, j)) + if a is not None: + a.set_ylabel(lbl, fontsize=7, rotation=90, labelpad=1) + ep = f"{epoch}" if epoch is not None else "?" + ls = f"{loss:.4f}" if loss is not None else "—" + fig.suptitle(f"Multi-marker virtual staining (phase → fluorescent marker) — {n_kept} live markers, one model\n" + f"epoch {ep} · train loss {ls} · {ncell:,} cells trained · overall Pearson r = {mean_r:.3f}", + fontsize=11, y=1 - 0.28 / figH) + fig.savefig(out / eval_name / "multi_montage.png", dpi=150, bbox_inches="tight"); plt.close(fig) + (out / eval_name / "multi_metrics.json").write_text(json.dumps( + {"n_markers": n_kept, "epoch": epoch, "loss": loss, "n_train": ncell, + "mean_pearson": round(mean_r, 3), "per_marker_pearson": {k: round(v, 3) for k, v in r_by.items()}}, indent=2)) + print(f"[eval] {n_kept} markers, epoch {ep}, mean Pearson {mean_r:.3f} -> {out/eval_name/'multi_montage.png'}") + + +def replot_eval(subdir="eval", out_dir=None): + """Re-render the montage from the cached arrays (montage_cache.npz) — NO GPU, NO model. Use this to + iterate on montage layout/spacing without re-running the eval sampling.""" + out = Path(out_dir or OUT) + d = np.load(out / subdir / "montage_cache.npz", allow_pickle=True) + rowspecs = list(zip(d["names"].tolist(), d["phase"], d["pred"], d["real"], d["r"].tolist())) + loss = d["loss"].item() + if loss is not None and loss != loss: # nan → treat as missing + loss = None + _plot_montage(rowspecs, int(d["n_kept"]), d["epoch"].item(), loss, int(d["ncell"]), out, subdir) + + +@torch.no_grad() +def eval_multi(ema, X, E, P, M, kept, cfg, out_dir, dev, epoch=None, loss=None, n_train=None, per_marker=1, eval_name="eval"): + """Sample each kept marker from phase + marker-id, cache the (phase, pred, real, Pearson) arrays to + montage_cache.npz (so the montage can be re-rendered later with no GPU via replot_eval), then plot.""" + out = Path(out_dir); (out / eval_name).mkdir(parents=True, exist_ok=True) + H = cfg.crop_size + rowspecs = [] + for mid, (_, mc, _) in enumerate(kept): + idx = np.where(M == mid)[0][-per_marker:] # tail cells (least likely in early SGD) + for i in idx: + g = torch.Generator(device=dev).manual_seed(100 + int(i)) + xT = torch.randn(1, 1, H, H, generator=g, device=dev) + e = torch.as_tensor(E[i:i + 1], dtype=torch.float32, device=dev) + ci = torch.as_tensor(P[i:i + 1], dtype=torch.float32, device=dev) + mk = torch.as_tensor([mid], dtype=torch.long, device=dev) + pred = _sample_marker(ema, xT, e, ci, mk, cfg, dev).cpu().numpy()[0, 0] + rowspecs.append((mc, P[i, 0], pred, X[i, 0], _pearson(pred, X[i, 0]))) # FULL marker name (protein) — no collisions + ncell = int(n_train) if n_train is not None else int(len(X)) + np.savez_compressed(out / eval_name / "montage_cache.npz", # everything the montage needs → replot with no GPU + names=np.array([s[0] for s in rowspecs]), + phase=np.stack([s[1] for s in rowspecs]).astype(np.float32), + pred=np.stack([s[2] for s in rowspecs]).astype(np.float32), + real=np.stack([s[3] for s in rowspecs]).astype(np.float32), + r=np.array([s[4] for s in rowspecs], np.float32), + n_kept=len(kept), epoch=epoch, loss=loss if loss is not None else np.nan, ncell=ncell) + _plot_montage(rowspecs, len(kept), epoch, loss, ncell, out, eval_name) + + +def run(cap=2500, epochs=120, batch=48, device="cuda"): + markers = markers_list() + print(f"[multi] {len(markers)} candidate markers") + X, E, P, M, kept = load_multi(markers, cap, f"{OUT}/cache") + cfg = DiffAEConfig(spatial_cond=True, n_markers=len(kept), device=device, epochs=epochs, batch_size=batch) + dev = torch.device(device if torch.cuda.is_available() else "cpu") + ema = train_multi(X, E, P, M, cfg, OUT, epochs, batch, device) + st = torch.load(Path(OUT) / "train_state.pt", map_location="cpu") + (Path(OUT) / "markers.json").write_text(json.dumps([mc for _, mc, _ in kept], indent=2)) + eval_multi(ema, X, E, P, M, kept, cfg, OUT, dev, epoch=st.get("epoch"), loss=st.get("loss"), n_train=int(len(X))) + return {"n_markers": len(kept), "n_crops": int(len(X))} + + +def eval_only(cap=2500, batch=48, device="cuda", subdir="eval"): + """Produce the eval montage/metrics from the CURRENT checkpoint (train_state.pt EMA, newest epoch) without + finishing training — caches make the load fast. Numbers are a floor (model still undertrained).""" + X, E, P, M, kept = load_multi(markers_list(), cap, f"{OUT}/cache") + cfg = DiffAEConfig(spatial_cond=True, n_markers=len(kept), device=device, epochs=1, batch_size=batch) + dev = torch.device(device if torch.cuda.is_available() else "cpu") + ema = DiffAE(cfg).to(dev).eval() + st = torch.load(Path(OUT) / "train_state.pt", map_location=dev) + ema.load_state_dict(st["ema"]); print(f"[eval-only] loaded EMA @ epoch {st['epoch']}") + loss = st.get("loss") + if loss is None or (isinstance(loss, float) and loss != loss): # missing/NaN → read latest from the train log + import glob, os, re + for f in sorted(glob.glob("/hpc/mydata/gav.sturm/ops_mono/slurm_logs/diffae/*/*.out"), key=os.path.getmtime)[::-1][:8]: + m = re.findall(r"epoch \d+: loss=([\d.]+)", open(f).read()) + if m: loss = float(m[-1]); break + (Path(OUT) / "markers.json").write_text(json.dumps([mc for _, mc, _ in kept], indent=2)) + eval_multi(ema, X, E, P, M, kept, cfg, OUT, dev, epoch=st["epoch"], loss=loss, n_train=int(len(X)), eval_name=subdir) + return {"epoch": st["epoch"], "n_markers": len(kept)} + + +_SP = {"slurm_partition": "gpu", "gpus_per_node": 1, "cpus_per_task": 12, "mem_gb": 300, + "timeout_min": 720, "slurm_constraint": "[a100_80|h100|h200|6000_blackwell]", + "slurm_setup": ["export PYTORCH_CUDA_ALLOC_CONF=expandable_segments:True"]} + + +def submit_chain(n_links=4, cap=2500, epochs=120, batch=48): + """Auto-relaunch scheme: chain n_links resume-jobs with afterany dependencies so training auto-continues + from train_state.pt across the 12h wall-clock limit until it reaches `epochs`. Each link resumes + automatically (train_multi loads train_state.pt); the link that reaches `epochs` runs eval + writes the + montage. Links that start already at `epochs` are cheap no-ops (load → skip loop → re-eval). afterany = + the next link runs regardless of how the previous ended (timeout saves train_state every epoch).""" + from ops_utils.hpc.slurm_batch_utils import submit_parallel_jobs + prev = None + for i in range(n_links): + sp = dict(_SP) + if prev: + sp["slurm_additional_parameters"] = {"dependency": f"afterany:{prev}"} + r = submit_parallel_jobs(jobs_to_submit=[{"name": f"vstain_multi_link{i}", "func": run, + "kwargs": {"cap": cap, "epochs": epochs, "batch": batch}}], + experiment="diffae", slurm_params=sp, log_dir="diffae", wait_for_completion=False) + prev = r["base_job_id"] + print(f"[chain] link {i}: job {prev}" + (f" (afterany prev)" if i else " (head, starts now)")) + print(f"[chain] {n_links} links → resumes to epoch {epochs}, final link evals") + + +def main(): + ap = argparse.ArgumentParser() + ap.add_argument("--cap", type=int, default=2500, help="cells per marker") + ap.add_argument("--epochs", type=int, default=120) + ap.add_argument("--batch", type=int, default=48) + ap.add_argument("--submit", action="store_true") + ap.add_argument("--eval", action="store_true", help="eval the current checkpoint only (no training)") + ap.add_argument("--chain", type=int, default=0, help="submit N afterany-chained resume jobs to reach --epochs") + ap.add_argument("--replot", metavar="SUBDIR", help="re-render the montage from a cached eval (no GPU)") + args = ap.parse_args() + if args.replot: + replot_eval(subdir=args.replot) + return + if args.chain: + submit_chain(n_links=args.chain, cap=args.cap, epochs=args.epochs, batch=args.batch) + return + func, name = (eval_only, "virtstain_multi_eval") if args.eval else (run, "virtstain_multi") + kw = {"cap": args.cap, "batch": args.batch} if args.eval else {"cap": args.cap, "epochs": args.epochs, "batch": args.batch} + if args.submit: + from ops_utils.hpc.slurm_batch_utils import submit_parallel_jobs + submit_parallel_jobs(jobs_to_submit=[{"name": name, "func": func, "kwargs": kw}], + experiment="diffae", slurm_params={"slurm_partition": "gpu", "gpus_per_node": 1, + "cpus_per_task": 12, "mem_gb": 300, "timeout_min": 720, + "slurm_constraint": "[a100_80|h100|h200|6000_blackwell]", + "slurm_setup": ["export PYTORCH_CUDA_ALLOC_CONF=expandable_segments:True"]}, + log_dir="diffae", wait_for_completion=False) + elif args.eval: + eval_only(cap=args.cap, batch=args.batch) + else: + run(cap=args.cap, epochs=args.epochs, batch=args.batch) + + +if __name__ == "__main__": + main() diff --git a/src/ops_model/models/attention/diffex/directions/__init__.py b/src/ops_model/models/attention/diffex/directions/__init__.py new file mode 100644 index 0000000..fe5b065 --- /dev/null +++ b/src/ops_model/models/attention/diffex/directions/__init__.py @@ -0,0 +1,6 @@ +"""DiffEx Stage 3 — contrastive direction discovery + classifier ranking + traversal. + +K direction MLPs (InfoNCE + decorrelation, unsupervised) on CellDINO embeddings → +rank by control-vs-target classifier score shift → DDIM-traverse the selected +direction and verify with re-encoded scores. See ../PLAN.md §3. +""" diff --git a/src/ops_model/models/attention/diffex/directions/batch.py b/src/ops_model/models/attention/diffex/directions/batch.py new file mode 100644 index 0000000..714185b --- /dev/null +++ b/src/ops_model/models/attention/diffex/directions/batch.py @@ -0,0 +1,172 @@ +"""Batch DiffEx: strips + NTC→KO GIF for the top-ranked geneKOs and EBI complexes. + +Reads the k10-ranked CSVs from the attention-selection page (top-N by rank_by_K10_mAP), +submits one GPU job per target. Each job: run_directions at w=5 only → per-cell strips + +scores, then a GIF for the auto-picked best-Δscore cell. + + python -m ops_model.models.attention.diffex.directions.batch \ + --genes-csv --complex-csv \ + --n-genes 50 --n-complex 20 +""" +from __future__ import annotations + +import argparse + +import numpy as np +import pandas as pd + +from ops_utils.hpc.slurm_batch_utils import submit_parallel_jobs + +from ..classifier.config import DEFAULT_OUT_ROOT, slugify +from .config import DirConfig +from .make_gifs import make_all_gifs, make_gif, render_all_review +from .run import run_directions + + +# ≤10-char header labels so complex names don't overflow the grid tiles. +COMPLEX_ABBR = { + "UTP-B complex": "UTP-B", + "mTORC1 complex": "mTORC1", + "Chaperonin-containing T-complex": "CCT", + "Box C/D snoRNA-Guided RNP methyltransferase complex, FBLL1 variant": "Box C/D", + "COPI vesicle coat complex, COPG1-COPZ1 variant": "COPI", + "SF3B complex": "SF3B", + "COP9 signalosome variant 1": "COP9", + "Nucleolar exosome complex, EXOSC10 variant": "Exosome", + "LSM2-8 complex": "LSM2-8", + "DNA polymerase alpha:primase complex": "Pol α-prim", + "40S cytosolic small ribosomal subunit": "40S ribo", + "19S proteasome regulatory complex": "19S prot", + "DNA polymerase epsilon complex": "DNA Pol ε", + "DNA-directed RNA polymerase III complex, POLR3G variant": "Pol III", + "DNA-directed RNA polymerase II complex": "Pol II", + "Core mediator complex": "Core mediator", + "Sm complex": "Sm core", + "Eukaryotic translation initiation factor 3 complex": "eIF3", + "ESCRT-III complex": "ESCRT-III", + "60S cytosolic large ribosomal subunit": "60S ribo", + "Actin-related protein 2/3 complex, ARPC1A-ACTR3B-ARPC5 variant": "Arp2/3", +} + + +def _short(name: str, n: int = 10) -> str: + """≤n-char GIF header label for a complex name (genes pass through unchanged).""" + if name in COMPLEX_ABBR: + return COMPLEX_ABBR[name] + s = name + for suf in (" complex", " subunit", " variant"): + s = s.replace(suf, "") + s = s.strip().strip(",") + return s if len(s) <= n else s[: n - 1] + "…" + + +def run_target(grain: str, target: str, label: str, w: float = 5.0) -> dict: + cfg = DirConfig(grain=grain, target=target, device="cuda") + cfg.guidance_scales = (w,) # w=5 only (batch) + out = f"{DEFAULT_OUT_ROOT}/directions/phase/{grain}/{slugify(target)}" + run_directions(cfg, out) + sc = np.load(f"{out}/scores_w{w:g}.npy") # (n_cells, n_alphas) + delta = sc[:, -1] - sc[:, 0] + best = int(np.argmax(delta)) # cell that moves most toward the phenotype + make_gif(grain, target, best, w, label) + return {"target": target, "best_cell": best, "delta": float(delta[best])} + + +def all_gifs_target(grain: str, target: str, label: str, w: float = 5.0) -> list: + """Render GIFs for every traversed cell of one target (strips must already exist).""" + return make_all_gifs(grain, target, label, w=w) + + +def review_all_target(grain: str, target: str, label: str, w: float = 5.0) -> list: + """Both styles (3-way axis + 2-way half) GIF + panel PNG for every traversed cell.""" + return render_all_review(grain, target, label, w=w) + + +# cells×α grid α-levels (each a full −max→+max sweep at w). See DiffEx defaults: w=2, α 2–4; +# ±5 included for the most subtle phenotypes where extreme α still adds signal. +_ALPHA_LEVELS = { + "a2": [-2, -1.6, -1.2, -0.8, -0.4, 0, 0.4, 0.8, 1.2, 1.6, 2], + "a3": [-3, -2, -1.5, -1, -0.5, 0, 0.5, 1, 1.5, 2, 3], + "a4": [-4, -3.2, -2.4, -1.6, -0.8, 0, 0.8, 1.6, 2.4, 3.2, 4], + "a5": [-5, -4, -3, -2, -1, 0, 1, 2, 3, 4, 5], +} + + +def marker_grid(marker_channel: str = None, channel: str = None, target: str = None, + ckpt: str = None, label: str = None, cells=(0, 1, 2), w: float = 2.0, + grain: str = "geneKO", control: str = None, fluor_csv: str = None, + device: str = "cuda") -> str: + """cells×α grid for one (marker, target): render every α-level (sharing one + gather+decoder) then composite a labeled rows=cells × cols=α grid. Returns the + grid gif path under directions/_grids/. + + marker_channel=None → phase mode (grain parquet + Phase2D crops; pass ckpt=phase_v1). + grain='complex' → EBI complexes. control=None → NTC-anchored; set control to another + class for an A→B traversal (α=0 shows the control/anchor class, +α → target).""" + from .grid import make_labeled_grid + from .make_gifs import _pair_slug + label = label or target + for ak, al in _ALPHA_LEVELS.items(): + render_all_review(grain, target, label, w=w, cells=list(cells), device=device, + ckpt=ckpt, tag=f"_{ak}", marker_channel=marker_channel, + channel=channel, alphas=al, control=control, fluor_csv=fluor_csv) + modality = slugify(marker_channel) if marker_channel else "phase" + slug = _pair_slug(target, control) + sd = f"{DEFAULT_OUT_ROOT}/directions/{modality}/{grain}/{slug}/strips" + grid = [[f"{sd}/{slug}_w{w:g}_cell{c}_{ak}_axis.gif" for ak in _ALPHA_LEVELS] for c in cells] + prefix = "fluor_" if marker_channel else "" + out = f"{DEFAULT_OUT_ROOT}/directions/_grids/{prefix}{modality}_{slug}_cellsxalpha.gif" + col_labels = [f"α=±{ak[1:]}" for ak in _ALPHA_LEVELS] # stays in sync with _ALPHA_LEVELS + anchor = f"{control} → " if control and control != "NTC" else "" + make_labeled_grid(grid, out, row_labels=[f"cell {c}" for c in cells], + col_labels=col_labels, + title=f"{marker_channel or 'Phase'} — {anchor}{target} (w={w:g})", tile_w=280) + return out + + +def build_targets(genes_csv: str, complex_csv: str, n_genes: int, n_complex: int): + g = pd.read_csv(genes_csv).sort_values("rank_by_K10_mAP").head(n_genes) + c = pd.read_csv(complex_csv).sort_values("rank_by_K10_mAP").head(n_complex) + jobs = [("geneKO", t, t) for t in g["geneKO"].tolist()] + jobs += [("complex", t, _short(t)) for t in c["complex_name"].tolist()] + return jobs + + +def main(): + ap = argparse.ArgumentParser(description="Batch DiffEx strips+GIFs for top-ranked targets") + ap.add_argument("--genes-csv", required=True) + ap.add_argument("--complex-csv", required=True) + ap.add_argument("--n-genes", type=int, default=50) + ap.add_argument("--n-complex", type=int, default=20) + ap.add_argument("--w", type=float, default=5.0) + ap.add_argument("--gifs-only", action="store_true", + help="strips already exist: only render GIFs for all traversed cells") + ap.add_argument("--review", action="store_true", + help="render both 3-way axis + 2-way half GIF+panel for all cells") + ap.add_argument("--dry-run", action="store_true") + args = ap.parse_args() + + targets = build_targets(args.genes_csv, args.complex_csv, args.n_genes, args.n_complex) + mode = "review, all cells, 3-way+2-way" if args.review else ("gifs-only, all cells" if args.gifs_only else "full") + print(f"{len(targets)} targets: {args.n_genes} geneKO + {args.n_complex} complex [{mode}]") + func = review_all_target if args.review else (all_gifs_target if args.gifs_only else run_target) + jobs = [{ + "name": f"dx_{grain}_{slugify(target)[:24]}", + "func": func, + "kwargs": {"grain": grain, "target": target, "label": label, "w": args.w}, + "metadata": {"stage": "batch_review" if args.review else ("batch_gifs" if args.gifs_only else "batch_directions"), + "grain": grain, "target": target}, + } for grain, target, label in targets] + + submit_parallel_jobs( + jobs_to_submit=jobs, experiment="diffex_batch", + slurm_params={"slurm_partition": "gpu", "gpus_per_node": 1, "cpus_per_task": 8, + "mem_gb": 64, "timeout_min": 90, + "slurm_constraint": "[a100_80|h100|h200|6000_blackwell]", + "slurm_setup": ["export PYTORCH_CUDA_ALLOC_CONF=expandable_segments:True"]}, + log_dir="diffex_batch", dry_run=args.dry_run, wait_for_completion=False, + ) + + +if __name__ == "__main__": + main() diff --git a/src/ops_model/models/attention/diffex/directions/config.py b/src/ops_model/models/attention/diffex/directions/config.py new file mode 100644 index 0000000..bbcea42 --- /dev/null +++ b/src/ops_model/models/attention/diffex/directions/config.py @@ -0,0 +1,72 @@ +"""Config for Stage 3 — contrastive direction discovery + ranking + traversal. + +Pipeline (Alex's recipe, our CellDINO space): + 2a train K direction MLPs UNSUPERVISED on CellDINO embeddings (InfoNCE + decorrelation) + 2b rank directions post-hoc by how much each shifts a control-vs-target classifier + 3 traverse the selected direction (α∈[-3,3]), DDIM-sample an image per step, verify score +""" +from __future__ import annotations + +from dataclasses import dataclass, field + +from ..classifier.config import DEFAULT_OUT_ROOT, GRAINS, PMA_PHASE_GENEKO # noqa: F401 + + +@dataclass +class DirConfig: + # ---- target / data ---- + grain: str = "geneKO" # geneKO | complex + target: str = "HSPA5" # the KD class to explain + control: str = "NTC" # control class + n_per_class: int = 1000 # top-attention cells per class + crop_size: int = 160 + channel: str = "Phase2D" # RAW pheno-zarr channel to read (Phase2D | GFP | mCherry | Cy5) + # fluorescent mode: set marker_channel to a fluor-attention-CSV `channel` value + # (e.g. "nucleus_NucleoLIVE Live Cell dye"); gather() then pulls that marker's top cells + # from fluor_csv and reads the raw `channel` above. None = phase mode (uses the grain parquet). + marker_channel: str | None = None + fluor_csv: str = "/hpc/projects/icd.fast.ops/models/alex_lin_attention/v4/pma_fluorescent_cells_all.csv" + mask_cell: bool = False + seed: int = 0 + + # ---- direction method (plan C) ---- + # 'mean_diff' | 'lr_weight' = deterministic supervised control→KD direction (PRIMARY, + # reproducible). 'unsupervised' = the paper's InfoNCE direction bank (secondary track, + # NOT reproducible run-to-run — GPU/seed sensitive + best_k flips). + direction_method: str = "mean_diff" + deterministic: bool = True # seed + cuDNN-deterministic so runs are repeatable + + # ---- 2a: unsupervised direction discovery (only when direction_method='unsupervised') ---- + cond_dim: int = 1024 # CellDINO ViT-L embedding dim + K: int = 10 # number of candidate directions + hidden: int = 512 + dir_epochs: int = 100 + dir_lr: float = 1e-3 + tau: float = 0.1 # InfoNCE temperature + decorr_weight: float = 1.0 # push the K mean-directions orthogonal + + # ---- 2b: ranking ---- + rank_alpha: float = 1.0 # ± shift magnitude when measuring classifier score change + + # ---- 3: traversal ---- + # alphas are MULTIPLES of the control→KD embedding gap ‖μ_KD−μ_ctrl‖ when + # scale_alpha_to_gap=True (α=+1 ≈ a full control→KD traversal); else raw units. + alphas: tuple = (-3.0, -2.0, -1.5, -1.0, -0.5, 0.0, 0.5, 1.0, 1.5, 2.0, 3.0) + scale_alpha_to_gap: bool = True + orient_sign: bool = True # orient so +α=toward KO (False = raw MLP sign, pre-orientation) + n_traverse: int = 8 # source (control) cells to traverse + ddim_steps: int = 50 + # edit-guidance: ε̃ = ε(z0) + w·(ε(z0+αd) − ε(z0)). w=1 = normal; w>1 amplifies the + # embedding edit's effect on the image (the DiffAE under-uses the embedding otherwise). + guidance_scales: tuple = (1.0, 3.0, 5.0) + + # ---- DiffAE decoder (must match the trained checkpoint) ---- + diffae_ckpt: str = f"{DEFAULT_OUT_ROOT}/diffae/phase_v1/diffae_best.pt" + block_out_channels: tuple = (128, 256, 256, 512) + layers_per_block: int = 2 + train_timesteps: int = 1000 + + # ---- run ---- + device: str = "cuda" + batch_size: int = 64 + num_workers: int = 0 diff --git a/src/ops_model/models/attention/diffex/directions/data.py b/src/ops_model/models/attention/diffex/directions/data.py new file mode 100644 index 0000000..8a3a582 --- /dev/null +++ b/src/ops_model/models/attention/diffex/directions/data.py @@ -0,0 +1,61 @@ +"""Gather target (KD) + control (NTC) top-attention cells, with CellDINO embeddings. + +Returns raw crops (for traversal inversion), CellDINO embeddings (direction +discovery + ranking), and labels (1=target, 0=control).""" +from __future__ import annotations + +import numpy as np +import pandas as pd + +from ..classifier.celldino_features import embed_crops +from ..classifier.config import GRAINS +from ..classifier.data import _BASE_COLS, make_labels_df, materialize_crops + + +def _top_cells(parquet: str, class_col: str, value: str, n: int) -> pd.DataFrame: + df = pd.read_parquet( + parquet, filters=[(class_col, "==", value), ("rank_type", "==", "top")], + columns=[class_col] + _BASE_COLS, + ) + if df.empty: + raise ValueError(f"no 'top' rows for {class_col}={value!r}") + return df.sort_values("rank").head(n).rename(columns={class_col: "cls"}) + + +def _gather_df(cfg): + """target+control cell table — phase: grain parquet; fluor: marker CSV (cfg.marker_channel).""" + cc = GRAINS[cfg.grain]["class_col"] + if getattr(cfg, "marker_channel", None): # fluorescent mode + acc = getattr(cfg, "accuracy_fluor_csv", None) + if acc: # accuracy variant: per-channel parquet (channel-filtered, class_col renamed) + rows = pd.read_parquet(acc) + else: + cols = list(dict.fromkeys([cc, *_BASE_COLS, "channel", "rank_type"])) + rows = pd.read_csv(cfg.fluor_csv, usecols=cols) + rows = rows[(rows["channel"] == cfg.marker_channel) & (rows["rank_type"] == "top")] + + def top(value, label): + d = rows[rows[cc] == value].sort_values("rank").head(cfg.n_per_class) + if d.empty: + raise ValueError(f"no fluor 'top' rows for {cc}={value!r} channel={cfg.marker_channel!r}") + d = d.rename(columns={cc: "cls"}).copy(); d["label"] = label + return d + return pd.concat([top(cfg.target, 1), top(cfg.control, 0)], ignore_index=True) + + g = GRAINS[cfg.grain] # phase mode + pq = getattr(cfg, "accuracy_parquet", None) or g["parquet"] # accuracy variant overrides the attention parquet (both A & B) + pos = _top_cells(pq, g["class_col"], cfg.target, cfg.n_per_class); pos["label"] = 1 + ctl = _top_cells(pq, g["class_col"], cfg.control, cfg.n_per_class); ctl["label"] = 0 + return pd.concat([pos, ctl], ignore_index=True) + + +def gather(cfg, crops_cache, emb_cache): + """-> (images_raw (N,1,H,W), embs (N,cond_dim), labels (N,)).""" + df = _gather_df(cfg) + print(f"gather: {int(df.label.sum())} {cfg.target} + {int((df.label==0).sum())} {cfg.control}" + + (f" [fluor {cfg.marker_channel}->raw {cfg.channel}]" if getattr(cfg, "marker_channel", None) else "")) + + labels_df = make_labels_df(df, cfg) + images, labels, _ = materialize_crops(labels_df, cfg, cache_path=crops_cache) + embs = embed_crops(images, cfg, cache_path=emb_cache) + return images, embs, labels diff --git a/src/ops_model/models/attention/diffex/directions/flow.py b/src/ops_model/models/attention/diffex/directions/flow.py new file mode 100644 index 0000000..bdea4d2 --- /dev/null +++ b/src/ops_model/models/attention/diffex/directions/flow.py @@ -0,0 +1,97 @@ +"""CellFlow-style conditional flow matching in CellDINO space (optional direction method). + +Replaces the linear mean-diff axis (d = μ_KD − μ_ctrl) with a learned velocity field that +transports the control embedding distribution → the target (KD) distribution — a rectified / +conditional flow-matching model. Traversal = Euler-integrate the ODE from a control cell's +embedding; decode each step with the frozen DiffAE. Distribution-aware and nonlinear, so it +captures multimodal / off-centroid phenotypes a single mean vector can't. + +Ref: CellFlow (bioRxiv 2025.04.11.648220); Flow Matching Guide (arXiv 2412.06264). +Deterministic given `seed` (fixed pairing/time sampling) so traversals stay reproducible. +""" +from __future__ import annotations + +import numpy as np +import torch +import torch.nn as nn + + +class FlowNet(nn.Module): + """Velocity field v_θ(x, t): CellDINO-dim in, CellDINO-dim out, time appended.""" + + def __init__(self, dim: int, hidden: int = 512): + super().__init__() + self.net = nn.Sequential( + nn.Linear(dim + 1, hidden), nn.SiLU(), + nn.Linear(hidden, hidden), nn.SiLU(), + nn.Linear(hidden, hidden), nn.SiLU(), + nn.Linear(hidden, dim), + ) + + def forward(self, x, t): # x:(B,dim) t:(B,) + return self.net(torch.cat([x, t[:, None]], dim=1)) + + +def train_flow(embs, labels, dev, steps=2000, bs=256, lr=1e-3, hidden=512, seed=0): + """Conditional flow matching, control(label 0) → KD(label 1). Independent coupling: + x_t = (1-t)·x0 + t·x1, regress v_θ(x_t,t) to the straight-line velocity (x1 − x0).""" + torch.manual_seed(seed); np.random.seed(seed) + x0 = torch.as_tensor(embs[labels == 0], dtype=torch.float32, device=dev) + x1 = torch.as_tensor(embs[labels == 1], dtype=torch.float32, device=dev) + net = FlowNet(embs.shape[1], hidden).to(dev) + opt = torch.optim.Adam(net.parameters(), lr=lr) + g = torch.Generator(device=dev).manual_seed(seed) + n0, n1 = len(x0), len(x1) + for _ in range(steps): + a = x0[torch.randint(0, n0, (bs,), generator=g, device=dev)] + b = x1[torch.randint(0, n1, (bs,), generator=g, device=dev)] + t = torch.rand(bs, generator=g, device=dev) + xt = (1 - t)[:, None] * a + t[:, None] * b + loss = ((net(xt, t) - (b - a)) ** 2).mean() + opt.zero_grad(); loss.backward(); opt.step() + net.eval() + return net + + +@torch.no_grad() +def integrate_flow(net, z0, dev, n_record=10, t_max=1.0, n_sub=None): + """Euler-integrate dz/dt = v_θ(z,t) from t=0→t_max; record n_record+1 evenly-spaced + points (incl. start). t_max>1 = OVERSHOOT past the KD manifold (t>1 is extrapolated). + Substep count scales with t_max to keep step size ~constant. Returns (n_record+1, dim).""" + if n_sub is None: + n_sub = max(n_record, int(round(50 * t_max))) + z = z0.clone().to(dev) + dt = t_max / n_sub + every = max(1, n_sub // n_record) + traj = [z.clone()] + for k in range(n_sub): + t = torch.full((z.shape[0],), k * dt, device=dev) + z = z + dt * net(z, t) + if (k + 1) % every == 0: + traj.append(z.clone()) + return torch.cat(traj, dim=0) + + +@torch.no_grad() +def integrate_flow_bidir(net, z0, dev, n_record=10, t_max=1.0, n_sub=None): + """Three-way: forward control→KD to t_max, plus a backward 'anti-KD' arm (step AGAINST + the field at t≈0). t_max>1 overshoots both extremes (analogous to mean-diff α>1). + Returns (2·n_record+1, dim): anti_extreme … NTC(center) … KD_extreme.""" + if n_sub is None: + n_sub = max(n_record, int(round(50 * t_max))) + z0 = z0.clone().to(dev) + dt = t_max / n_sub + every = max(1, n_sub // n_record) + z, fwd = z0.clone(), [] + for k in range(n_sub): + t = torch.full((z.shape[0],), k * dt, device=dev) + z = z + dt * net(z, t) + if (k + 1) % every == 0: + fwd.append(z.clone()) + z, bwd = z0.clone(), [] + t0 = torch.zeros(z0.shape[0], device=dev) + for k in range(n_sub): + z = z - dt * net(z, t0) # opposite the control→KD velocity + if (k + 1) % every == 0: + bwd.append(z.clone()) + return torch.cat(list(reversed(bwd)) + [z0.clone()] + fwd, dim=0) diff --git a/src/ops_model/models/attention/diffex/directions/grid.py b/src/ops_model/models/attention/diffex/directions/grid.py new file mode 100644 index 0000000..b9b6bb3 --- /dev/null +++ b/src/ops_model/models/attention/diffex/directions/grid.py @@ -0,0 +1,131 @@ +"""Composite existing per-target NTC→KO GIFs into a synchronized grid-canvas GIF. + +Each tile is one target's already-rendered animation (same cell, same w); all tiles share +the identical frame schedule, so we stack them frame-by-frame into one canvas that morphs +every perturbation at once. Per-frame durations are read from the source GIFs so the +end-of-traversal settle/hold is preserved. Pure image compositing — no GPU. +""" +from __future__ import annotations + +import os +from pathlib import Path + +import pandas as pd +from PIL import Image, ImageDraw, ImageFont, ImageSequence + +from ..classifier.config import slugify + + +def _font(sz): + try: + import matplotlib.font_manager as fm + return ImageFont.truetype(fm.findfont("DejaVu Sans"), sz) + except Exception: + return ImageFont.load_default() + + +def _frames_dur(gif_path): + im = Image.open(gif_path) + frames, durs = [], [] + for f in ImageSequence.Iterator(im): + frames.append(f.convert("RGB").copy()) + durs.append(int(f.info.get("duration", 180))) + return frames, durs + + +def make_labeled_grid(grid_paths, out_path, row_labels, col_labels, title=None, tile_w=260): + """Composite a synchronized R×C matrix of GIFs with row/col axis labels (a disentanglement + figure). grid_paths: list of rows, each a list of C gif paths. Preserves frame durations.""" + seqs = [[_frames_dur(p) for p in row] for row in grid_paths] + F = min(len(fr) for row in seqs for fr, _ in row) + durs = seqs[0][0][1][:F] + tw0, th0 = seqs[0][0][0][0].size + tw, th = tile_w, round(th0 * tile_w / tw0) + R, C = len(grid_paths), len(grid_paths[0]) + gap, lm, tm = 8, 92, (58 if title else 30) # left margin (row labels), top margin + W = lm + C * tw + (C + 1) * gap + H = tm + 24 + R * th + (R + 1) * gap # +24 for col-label row + tfont, lfont = _font(30), _font(20) + + frames = [] + for k in range(F): + cv = Image.new("RGB", (W, H), (0, 0, 0)); d = ImageDraw.Draw(cv) + if title: + d.text(((W - d.textlength(title, font=tfont)) / 2, 12), title, font=tfont, fill=(235, 235, 235)) + for c, cl in enumerate(col_labels): # column headers + x0 = lm + gap + c * (tw + gap) + d.text((x0 + (tw - d.textlength(cl, font=lfont)) / 2, tm), cl, font=lfont, fill=(0, 200, 255)) + for r, rl in enumerate(row_labels): # row labels (left, rotated-ish: just left-aligned) + y0 = tm + 24 + gap + r * (th + gap) + d.text((8, y0 + th / 2 - 10), rl, font=lfont, fill=(255, 170, 40)) + for r in range(R): + for c in range(C): + tile = seqs[r][c][0][k].resize((tw, th), Image.BILINEAR) + cv.paste(tile, (lm + gap + c * (tw + gap), tm + 24 + gap + r * (th + gap))) + frames.append(cv) + Path(out_path).parent.mkdir(parents=True, exist_ok=True) + frames[0].save(out_path, save_all=True, append_images=frames[1:], duration=durs, loop=0) + print(f"wrote {out_path} ({R}x{C}, {F} frames, {W}x{H})") + return str(out_path) + + +def build_tiles(grain, cell, out_root, csv_dir=None, names=None, + n_genes=50, n_complex=20, w=5.0, suffix="", modality="phase"): + """suffix: '' | '_axis' | '_half'. names overrides the CSV top-N list. + modality: 'phase' or a marker slug (directions///).""" + if names is None: + if grain == "geneKO": + df = pd.read_csv(f"{csv_dir}/k10_ranked_all_geneKOs.csv").sort_values("rank_by_K10_mAP").head(n_genes) + names = df["geneKO"].tolist() + else: + df = pd.read_csv(f"{csv_dir}/k10_ranked_all_complexes.csv").sort_values("rank_by_K10_mAP").head(n_complex) + names = df["complex_name"].tolist() + tiles, missing = [], [] + for nm in names: + s = slugify(nm) + p = f"{out_root}/directions/{modality}/{grain}/{s}/strips/{s}_w{w:g}_cell{cell}{suffix}.gif" + (tiles if os.path.exists(p) else missing).append(p) + if missing: + print(f" [{grain} cell{cell}{suffix}] {len(missing)} tiles missing (skipped)") + return tiles + + +def make_grid(tiles, out_path, ncols=10, gap=6, bg=(0, 0, 0), title=None, tile_w=240): + """tiles: list of GIF paths (same schedule). Preserves per-frame durations.""" + ref_frames, durs = _frames_dur(tiles[0]) + F = len(ref_frames) + seqs = [ref_frames] + for t in tiles[1:]: + fr, _ = _frames_dur(t) + seqs.append(fr) + F = min(F, min(len(s) for s in seqs)) + durs = durs[:F] + + tw0, th0 = seqs[0][0].size + if tile_w: + tw, th = tile_w, round(th0 * tile_w / tw0) + else: + tw, th = tw0, th0 + n = len(seqs) + nrows = (n + ncols - 1) // ncols + hdr = 40 if title else 0 + W = ncols * tw + (ncols + 1) * gap + H = hdr + nrows * th + (nrows + 1) * gap + tfont = _font(26) + + out_frames = [] + for k in range(F): + canvas = Image.new("RGB", (W, H), bg) + if title: + d = ImageDraw.Draw(canvas) + d.text(((W - d.textlength(title, font=tfont)) / 2, 8), title, font=tfont, fill=(235, 235, 235)) + for i, s in enumerate(seqs): + r, c = divmod(i, ncols) + tile = s[k].resize((tw, th), Image.BILINEAR) if tile_w else s[k] + canvas.paste(tile, (gap + c * (tw + gap), hdr + gap + r * (th + gap))) + out_frames.append(canvas) + Path(out_path).parent.mkdir(parents=True, exist_ok=True) + out_frames[0].save(out_path, save_all=True, append_images=out_frames[1:], + duration=durs, loop=0) + print(f"wrote {out_path} ({n} tiles, {F} frames, {W}x{H})") + return str(out_path) diff --git a/src/ops_model/models/attention/diffex/directions/losses.py b/src/ops_model/models/attention/diffex/directions/losses.py new file mode 100644 index 0000000..d84e6a4 --- /dev/null +++ b/src/ops_model/models/attention/diffex/directions/losses.py @@ -0,0 +1,35 @@ +"""Unsupervised direction-discovery losses (Alex's recipe). + +InfoNCE: each direction MLP's outputs cluster together (consistent axis) and apart +from other MLPs' — so the K directions are distinct, consistent axes of variation. +Decorrelation: push the K mean directions toward orthogonal (VICReg-style covariance +term) so MLPs don't collapse onto the same axis. +""" +from __future__ import annotations + +import torch +import torch.nn.functional as F + + +def infonce_directions(dirs: torch.Tensor, tau: float = 0.1) -> torch.Tensor: + """dirs (K,B,D) unit vectors. Item (k,i); positives = same direction k.""" + K, B, D = dirs.shape + x = dirs.reshape(K * B, D) + n = K * B + sim = (x @ x.t()) / tau + eye = torch.eye(n, device=x.device, dtype=torch.bool) + sim = sim.masked_fill(eye, -1e9) + labels = torch.arange(K, device=x.device).repeat_interleave(B) + pos = (labels.unsqueeze(0) == labels.unsqueeze(1)) & ~eye + lse_all = torch.logsumexp(sim, dim=1) + lse_pos = torch.logsumexp(sim.masked_fill(~pos, -1e9), dim=1) + return -(lse_pos - lse_all).mean() + + +def decorrelation(dirs: torch.Tensor) -> torch.Tensor: + """Penalize off-diagonal cosine similarity of the K mean directions.""" + m = F.normalize(dirs.mean(dim=1), dim=-1) # (K,D) + g = m @ m.t() # (K,K) + K = g.shape[0] + off = g - torch.diag(torch.diag(g)) + return (off ** 2).sum() / (K * (K - 1)) diff --git a/src/ops_model/models/attention/diffex/directions/make_gifs.py b/src/ops_model/models/attention/diffex/directions/make_gifs.py new file mode 100644 index 0000000..d191c2c --- /dev/null +++ b/src/ops_model/models/attention/diffex/directions/make_gifs.py @@ -0,0 +1,371 @@ +"""Render NTC→KO animation GIFs for specific (target, cell) traversals. + +Reproduces the exact frames of a per-cell strip (same seed 1234+cell, deterministic +mean_diff direction, same CFG guidance/null baseline as traverse) and writes an +animated GIF ordered most-NTC-like → most-KO-like (ping-pong loop). Each frame gets a +FIXED label "NTC → {target}" whose two ends brighten (cyan↔red) with a progress bar to +show position — constant text, no flicker. GPU. +""" +from __future__ import annotations + +from pathlib import Path + +import numpy as np +import torch +from PIL import Image, ImageDraw, ImageFont + +from ..classifier.config import DEFAULT_OUT_ROOT, GRAINS, slugify +from .config import DirConfig +from .data import gather +from .rank import supervised_direction +from .traverse import _sample_guided, load_diffae + +# (grain, target, cell index into control cells, guidance w, short label) +SPECS = [ + ("geneKO", "TIMM23", 6, 5.0, "TIMM23"), + ("complex", "Chaperonin-containing T-complex", 6, 5.0, "CCT"), + ("complex", "Actin-related protein 2/3 complex, ARPC1A-ACTR3B-ARPC5 variant", 5, 5.0, "Arp2/3"), + ("geneKO", "POLR1B", 7, 5.0, "POLR1B"), + ("geneKO", "HSPA5", 6, 5.0, "HSPA5"), + ("geneKO", "POLR2C", 7, 5.0, "POLR2C"), +] + +_NTC = (0, 200, 255) # cyan (NTC end) +_KO = (255, 90, 60) # red (KO end) +_DIM = (110, 110, 110) + + +def _font(sz): + try: + import matplotlib.font_manager as fm + return ImageFont.truetype(fm.findfont("DejaVu Sans"), sz) + except Exception: + return ImageFont.load_default() + + +def _lerp(a, b, t): + t = float(np.clip(t, 0, 1)) + return tuple(int(round(a[i] + (b[i] - a[i]) * t)) for i in range(3)) + + +def _labeled(cell_u8, pos, label, W=340, hdr=54): + """cell_u8 (H,W) uint8; pos in [0,1] (0=NTC, 1=KO).""" + canvas = Image.new("RGB", (W, hdr + W), (0, 0, 0)) + cell = Image.fromarray(cell_u8).resize((W, W), Image.NEAREST).convert("RGB") + canvas.paste(cell, (0, hdr)) + d = ImageDraw.Draw(canvas) + f = _font(22) + parts = [("NTC", _lerp(_DIM, _NTC, 1 - pos)), + (" → ", (225, 225, 225)), + (label, _lerp(_DIM, _KO, pos))] + widths = [d.textlength(t, font=f) for t, _ in parts] + x = (W - sum(widths)) / 2 + for (t, c), wdt in zip(parts, widths): + d.text((x, 12), t, font=f, fill=c) + x += wdt + bx0, bx1, by = 14, W - 14, hdr - 12 + d.rectangle([bx0, by, bx1, by + 6], fill=(55, 55, 55)) + d.rectangle([bx0, by, bx0 + (bx1 - bx0) * pos, by + 6], fill=_lerp(_NTC, _KO, pos)) + return canvas + + +_ANTI = (255, 170, 40) # amber — anti-KO extreme (−label) + + +def _labeled3(cell_u8, pos, label, W=340, hdr=58): + """Full-axis header: −label ← NTC(center) → label. pos in [0,1], 0.5 = NTC (α=0).""" + canvas = Image.new("RGB", (W, hdr + W), (0, 0, 0)) + cell = Image.fromarray(cell_u8).resize((W, W), Image.NEAREST).convert("RGB") + canvas.paste(cell, (0, hdr)) + d = ImageDraw.Draw(canvas) + # auto-shrink font so long labels (−label / label) don't collide with centered NTC + left = f"−{label}" + sz = 19 + while sz > 11: + f = _font(sz) + side = max(d.textlength(left, font=f), d.textlength(label, font=f)) + if side <= W / 2 - d.textlength("NTC", font=f) / 2 - 8: + break + sz -= 1 + f = _font(sz) + aL = max(0.0, 1 - abs(pos - 0.0) / 0.5) + aM = max(0.0, 1 - abs(pos - 0.5) / 0.5) + aR = max(0.0, 1 - abs(pos - 1.0) / 0.5) + d.text((10, 11), f"−{label}", font=f, fill=_lerp(_DIM, _ANTI, aL)) + wm = d.textlength("NTC", font=f) + d.text(((W - wm) / 2, 11), "NTC", font=f, fill=_lerp(_DIM, _NTC, aM)) + wr = d.textlength(label, font=f) + d.text((W - 10 - wr, 11), label, font=f, fill=_lerp(_DIM, _KO, aR)) + bx0, bx1, by = 14, W - 14, hdr - 14 + d.rectangle([bx0, by, bx1, by + 6], fill=(55, 55, 55)) + cx = (bx0 + bx1) / 2 + d.rectangle([cx - 1, by - 3, cx + 1, by + 9], fill=_NTC) # NTC center tick + mx = bx0 + (bx1 - bx0) * pos + d.ellipse([mx - 5, by - 3, mx + 5, by + 9], fill=(_KO if pos > 0.5 else _ANTI)) + return canvas + + +def _strip(frames, gap=4): + w, h = frames[0].size + n = len(frames) + c = Image.new("RGB", (n * w + (n + 1) * gap, h + 2 * gap), (0, 0, 0)) + for i, fr in enumerate(frames): + c.paste(fr, (gap + i * (w + gap), gap)) + return c + + +@torch.no_grad() +def _render_review(ctx, cell, w, label, tag="", styles=("axis",)): + """Render traversal styles from one set of frames. DEFAULT = three-way ('axis', + −label ← NTC(center) → label). Pass styles=('axis','half') or ('half',) for the + two-way NTC→KO view. `tag` is inserted into filenames (e.g. '_v2').""" + dev, cfg, slug, out, embs, labels, fixed_dir, gap, diffae, null_base = ctx + ci = np.flatnonzero(labels == 0) + z0 = torch.as_tensor(embs[ci[cell]:ci[cell] + 1], dtype=torch.float32).to(dev) + H = cfg.crop_size + ge = torch.Generator(device=dev).manual_seed(1234 + cell) + xT = torch.randn(1, 1, H, H, generator=ge, device=dev) + all_alphas = sorted(cfg.alphas); amin, amax = all_alphas[0], all_alphas[-1] + # only decode the α needed: axis needs the full range, half only α≥0 + alphas = all_alphas if "axis" in styles else [a for a in all_alphas if a >= 0] + raw = {} + for a in alphas: + img = _sample_guided(diffae, xT.clone(), z0 + (a * gap) * fixed_dir, null_base, w, cfg) + raw[a] = np.clip((img.cpu().numpy()[0, 0] + 1) / 2, 0, 1) + sd = out / "strips"; sd.mkdir(parents=True, exist_ok=True) + + if "axis" in styles: # FULL 3-way: −label ← NTC(center) → label; NTC→+→−→ back + full = [_labeled3((raw[a] * 255).astype("uint8"), (a - amin) / (amax - amin), label) + for a in all_alphas] + m = len(full) // 2; n = len(full); He, Hm = 5, 2 + idx = ([m] * Hm + list(range(m + 1, n)) + [n - 1] * He + + list(range(n - 2, m - 1, -1)) + [m] * Hm + + list(range(m - 1, -1, -1)) + [0] * He + + list(range(1, m + 1)) + [m] * Hm) + seq = [full[i] for i in idx] + seq[0].save(sd / f"{slug}_w{w:g}_cell{cell}{tag}_axis.gif", save_all=True, + append_images=seq[1:], duration=180, loop=0) + _strip(full).save(sd / f"{slug}_w{w:g}_cell{cell}{tag}_axis_strip.png") + + if "half" in styles: # TWO-WAY (default): true NTC (α=0) → label; pause at both ends + pa = [a for a in all_alphas if a >= 0]; amx = pa[-1] + half = [_labeled((raw[a] * 255).astype("uint8"), a / amx, label) for a in pa] + sh = [half[0]] * 5 + half + [half[-1]] * 6 + half[-2:0:-1] + [half[0]] * 2 + sh[0].save(sd / f"{slug}_w{w:g}_cell{cell}{tag}_half.gif", save_all=True, + append_images=sh[1:], duration=180, loop=0) + _strip(half).save(sd / f"{slug}_w{w:g}_cell{cell}{tag}_half_strip.png") + print(f"review {slug} cell{cell}: {'+'.join(styles)} written") + return slug + + +def run_review(specs=None, device="cuda"): + specs = specs or [("geneKO", "GBF1", 2, 5.0, "GBF1"), + ("complex", "mTORC1 complex", 2, 5.0, "mTORC1")] + return [_render_review(_setup(g, t, DEFAULT_OUT_ROOT, device), c, w, lab) + for g, t, c, w, lab in specs] + + +def run_v2_review(device="cuda"): + """Compare the phase_v2_aug (dihedral) model on GBF1 + mTORC1, cell2, at w=5 and w=8.""" + ckpt = f"{DEFAULT_OUT_ROOT}/diffae/phase_v2_aug/diffae_best.pt" + specs = [("geneKO", "GBF1", 2, "GBF1"), ("complex", "mTORC1 complex", 2, "mTORC1")] + outs = [] + for g, t, cell, lab in specs: + ctx = _setup(g, t, DEFAULT_OUT_ROOT, device, ckpt=ckpt) + for w in (5.0, 8.0): + outs.append(_render_review(ctx, cell, w, lab, tag="_v2")) + return outs + + +def render_flow(grain, target, label, cells=(0, 2, 3, 5), w=5.0, n_record=10, + device="cuda", ckpt=None, two_way=False, overshoot=1.0): + """OPTIONAL traversal via CellFlow-style conditional flow matching (see flow.py). + Learns a control→KD velocity field in CellDINO space, integrates the ODE from each + control cell, decodes each step with the frozen DiffAE. DEFAULT 3-way (−label ← NTC → + label; anti arm is a backward extrapolation). two_way=True → forward-only NTC→KO. + overshoot = ODE end time t_max: 1.0 lands at the KD manifold; >1 overshoots past it + (overshoot≈3 ≈ mean-diff's α=3 drama). Filenames get '_o{overshoot}' when ≠1.""" + from .flow import train_flow, integrate_flow, integrate_flow_bidir + ctx = _setup(grain, target, DEFAULT_OUT_ROOT, device, ckpt=ckpt) + dev, cfg, slug, out, embs, labels, fixed_dir, gap, diffae, null_base = ctx + net = train_flow(embs, labels, dev, seed=cfg.seed) + sd = out / "strips"; sd.mkdir(parents=True, exist_ok=True) + ci = np.flatnonzero(labels == 0) + H = cfg.crop_size + otag = "" if overshoot == 1.0 else f"_o{overshoot:g}" + outs = [] + for cell in cells: + z0 = torch.as_tensor(embs[ci[cell]:ci[cell] + 1], dtype=torch.float32).to(dev) + ge = torch.Generator(device=dev).manual_seed(1234 + cell) + xT = torch.randn(1, 1, H, H, generator=ge, device=dev) + if two_way: + traj = integrate_flow(net, z0, dev, n_record=n_record, t_max=overshoot) + n = traj.shape[0] + frames = [_labeled( + (np.clip((_sample_guided(diffae, xT.clone(), traj[k:k + 1], null_base, w, cfg) + .cpu().numpy()[0, 0] + 1) / 2, 0, 1) * 255).astype("uint8"), + k / (n - 1), label) for k in range(n)] + sq = [frames[0]] * 5 + frames + [frames[-1]] * 6 + frames[-2:0:-1] + [frames[0]] * 2 + suffix = "_flow2way" + else: + traj = integrate_flow_bidir(net, z0, dev, n_record=n_record, t_max=overshoot) + n = traj.shape[0] + frames = [_labeled3( + (np.clip((_sample_guided(diffae, xT.clone(), traj[k:k + 1], null_base, w, cfg) + .cpu().numpy()[0, 0] + 1) / 2, 0, 1) * 255).astype("uint8"), + k / (n - 1), label) for k in range(n)] + m = n // 2; He, Hm = 5, 2 + idx = ([m] * Hm + list(range(m + 1, n)) + [n - 1] * He + + list(range(n - 2, m - 1, -1)) + [m] * Hm + + list(range(m - 1, -1, -1)) + [0] * He + + list(range(1, m + 1)) + [m] * Hm) + sq = [frames[i] for i in idx] + suffix = "_flow" + gif = sd / f"{slug}{suffix}{otag}_cell{cell}.gif" + sq[0].save(gif, save_all=True, append_images=sq[1:], duration=180, loop=0) + _strip(frames).save(sd / f"{slug}{suffix}{otag}_cell{cell}_strip.png") + print(f"wrote {gif} ({n} steps)") + outs.append(str(gif)) + return outs + + +def compare_ckpts(grain, target, cell, w, label, ckpts, device="cuda"): + """Render the same (target, cell, w) under multiple checkpoints, tagged, for A/B compare. + ckpts: {tag: ckpt_path} e.g. {'_v1': '.../phase_v1/diffae_best.pt', '_v2': '.../phase_v2_aug/...'}""" + out = [] + for tag, ckpt in ckpts.items(): + ctx = _setup(grain, target, DEFAULT_OUT_ROOT, device, ckpt=ckpt) + out.append(_render_review(ctx, cell, w, label, tag=tag)) + return out + + +def render_all_review(grain, target, label, w=5.0, cells=None, device="cuda", ckpt=None, tag="", + marker_channel=None, channel=None, alphas=None, fluor_csv=None, control=None): + """Both styles (3-way axis + 2-way half), GIF + panel PNG, for the given cells. + ckpt overrides the DiffAE checkpoint; tag suffixes filenames. For fluor pass + marker_channel (fluor-CSV channel) + channel (raw GFP/mCherry) + the marker's DiffAE ckpt; + fluor_csv overrides the attention CSV (use the EBI one for grain='complex'). + alphas overrides the traversal range (e.g. tighter ±1 for strongly-conditioned models). + control: anchor class for A→B (default NTC).""" + ctx = _setup(grain, target, DEFAULT_OUT_ROOT, device, ckpt=ckpt, + marker_channel=marker_channel, channel=channel, fluor_csv=fluor_csv, control=control) + if alphas is not None: + ctx[1].alphas = tuple(alphas) # ctx[1] = cfg + n_ctrl = int((ctx[5] == 0).sum()) # ctx[5] = labels + cells = list(cells) if cells is not None else list(range(min(ctx[1].n_traverse, n_ctrl))) + return [_render_review(ctx, c, w, label, tag=tag) for c in cells] + + +def _pair_slug(target, control=None): + """Effective slug for a traversal. NTC-anchored → slug(target); class→class anchor + (control set to a non-NTC class) → '__to__' so A→B assets/caches never + collide with the A-vs-NTC run.""" + if control and control != "NTC": + return f"{slugify(control)}__to__{slugify(target)}" + return slugify(target) + + +def _setup(grain, target, out_root, device, ckpt=None, marker_channel=None, channel=None, + fluor_csv=None, control=None, num_workers=None, return_images=False, + accuracy_parquet=None, variant=None, accuracy_fluor_csv=None): + """Expensive per-target setup shared by all cells: gather + direction + model. + Fluor: pass marker_channel (fluor-CSV channel) + channel (raw GFP/mCherry) + the marker's + DiffAE via ckpt. Fluor gets its own out dir (__) + cache so it never + collides with the phase pipeline. control: anchor class (default NTC); set to another class + for A→B traversal (direction becomes μ_target − μ_control, anchored on control cells). + num_workers: parallel DataLoader workers for the (I/O-bound) zarr crop materialization — + the gather dominates runtime, so set this to ~cpus for a big speedup.""" + dev = torch.device(device if torch.cuda.is_available() or device == "cpu" else "cpu") + cfg = DirConfig(grain=grain, target=target, device=device) + if ckpt: + cfg.diffae_ckpt = ckpt + if marker_channel: + cfg.marker_channel = marker_channel + if channel: + cfg.channel = channel + if fluor_csv: + cfg.fluor_csv = fluor_csv + if control: + cfg.control = control + if num_workers is not None: + cfg.num_workers = num_workers + if accuracy_parquet: + cfg.accuracy_parquet = accuracy_parquet # accuracy-variant cell selection (both A & B) + if accuracy_fluor_csv: + cfg.accuracy_fluor_csv = accuracy_fluor_csv # fluor accuracy: per-channel parquet (anchor gather) + slug = _pair_slug(target, control) + # modality-first layout: directions/// — keeps each modality's + # per-target listing separate (phase not overwhelmed by per-marker copies). + modality = (slugify(cfg.marker_channel) if cfg.marker_channel else "phase") + (f"_{variant}" if variant else "") + out = Path(out_root) / "directions" / modality / grain / slug + cache = out / "cache" + cache.mkdir(parents=True, exist_ok=True) # brand-new targets have no cache dir yet + tag = f"{slug}_{cfg.crop_size}" + images, embs, labels = gather( + cfg, str(cache / f"crops_{tag}.npz"), str(cache / f"celldino_{tag}.npz")) + d, _, _, _ = supervised_direction(embs, labels, cfg) + gap = float(np.linalg.norm(embs[labels == 1].mean(0) - embs[labels == 0].mean(0))) + fixed_dir = torch.as_tensor(d, dtype=torch.float32, device=dev)[None] + diffae = load_diffae(cfg, dev) + null_base = diffae.null_emb.detach()[None].to(dev) + ctx = (dev, cfg, slug, out, embs, labels, fixed_dir, gap, diffae, null_base) + return (*ctx, images) if return_images else ctx + + +@torch.no_grad() +def _render_cell(ctx, cell, w, label): + dev, cfg, slug, out, embs, labels, fixed_dir, gap, diffae, null_base = ctx + ctrl_idx = np.flatnonzero(labels == 0) + z0 = torch.as_tensor(embs[ctrl_idx[cell]:ctrl_idx[cell] + 1], dtype=torch.float32).to(dev) + H = cfg.crop_size + ge = torch.Generator(device=dev).manual_seed(1234 + cell) + xT = torch.randn(1, 1, H, H, generator=ge, device=dev) + + # Default = THREE-WAY: −label ← NTC(center) → label. Start at NTC → +KO → back → −KO → back. + alphas = sorted(cfg.alphas); amin, amax = alphas[0], alphas[-1] + frames = [] + for a in alphas: + img = _sample_guided(diffae, xT.clone(), z0 + (a * gap) * fixed_dir, null_base, w, cfg) + u = np.clip((img.cpu().numpy()[0, 0] + 1) / 2, 0, 1) + frames.append(_labeled3((u * 255).astype("uint8"), (a - amin) / (amax - amin), label)) + + m = len(frames) // 2; n = len(frames); He, Hm = 5, 2 + idx = ([m] * Hm + list(range(m + 1, n)) + [n - 1] * He + + list(range(n - 2, m - 1, -1)) + [m] * Hm + + list(range(m - 1, -1, -1)) + [0] * He + + list(range(1, m + 1)) + [m] * Hm) + seq = [frames[i] for i in idx] + gif = out / "strips" / f"{slug}_w{w:g}_cell{cell}.gif" + seq[0].save(gif, save_all=True, append_images=seq[1:], duration=180, loop=0) + print(f"wrote {gif} ({len(alphas)} frames, gap={gap:.2f})") + return str(gif) + + +def make_gif(grain, target, cell, w, label, out_root=DEFAULT_OUT_ROOT, device="cuda"): + ctx = _setup(grain, target, out_root, device) + return _render_cell(ctx, cell, w, label) + + +def make_all_gifs(grain, target, label, w=5.0, cells=None, out_root=DEFAULT_OUT_ROOT, device="cuda"): + """Render GIFs for every traversed cell of one target (setup done once).""" + ctx = _setup(grain, target, out_root, device) + n_ctrl = int((ctx[5] == 0).sum()) # ctx[5] = labels + cells = list(cells) if cells is not None else list(range(min(ctx[1].n_traverse, n_ctrl))) + return [_render_cell(ctx, c, w, label) for c in cells] + + +def run_all(specs=SPECS, device="cuda"): + return [make_gif(g, t, c, w, lab, device=device) for g, t, c, w, lab in specs] + + +if __name__ == "__main__": + from ops_utils.hpc.slurm_batch_utils import submit_parallel_jobs + submit_parallel_jobs( + jobs_to_submit=[{"name": "diffex_gifs", "func": run_all, "kwargs": {}, + "metadata": {"stage": "gifs"}}], + experiment="diffex_gifs", + slurm_params={"slurm_partition": "gpu", "gpus_per_node": 1, "cpus_per_task": 8, + "mem_gb": 64, "timeout_min": 60, + "slurm_constraint": "[a100_80|h100|h200|6000_blackwell]"}, + log_dir="diffex_gifs", wait_for_completion=False, + ) diff --git a/src/ops_model/models/attention/diffex/directions/model.py b/src/ops_model/models/attention/diffex/directions/model.py new file mode 100644 index 0000000..bc58332 --- /dev/null +++ b/src/ops_model/models/attention/diffex/directions/model.py @@ -0,0 +1,25 @@ +"""K direction MLPs. Each maps a CellDINO embedding z -> a unit direction d_k(z). +An edit is z_new = z + alpha * d_k(z).""" +from __future__ import annotations + +import torch +import torch.nn as nn +import torch.nn.functional as F + + +class DirectionBank(nn.Module): + def __init__(self, dim: int, K: int = 10, hidden: int = 512): + super().__init__() + self.K = K + self.mlps = nn.ModuleList([ + nn.Sequential(nn.Linear(dim, hidden), nn.ReLU(inplace=True), nn.Linear(hidden, dim)) + for _ in range(K) + ]) + + def forward(self, z: torch.Tensor) -> torch.Tensor: + """z (B,D) -> (K,B,D) unit directions.""" + return torch.stack([F.normalize(m(z), dim=-1) for m in self.mlps], dim=0) + + def direction(self, z: torch.Tensor, k: int) -> torch.Tensor: + """Unit direction for one MLP: (B,D).""" + return F.normalize(self.mlps[k](z), dim=-1) diff --git a/src/ops_model/models/attention/diffex/directions/proto_ddim_anchors.py b/src/ops_model/models/attention/diffex/directions/proto_ddim_anchors.py new file mode 100644 index 0000000..131ed36 --- /dev/null +++ b/src/ops_model/models/attention/diffex/directions/proto_ddim_anchors.py @@ -0,0 +1,162 @@ +"""Prototype: DDIM-INVERTED anchors for phase geneKO traversals. + +The production traversal seeds a FIXED RANDOM xT per cell, so α=0 is a generic DDPM recon of z0, +not the real anchor cell (see traverse.py). Here we instead DDIM-INVERT each anchor cell to its own +xT (conditioned on z0), then sweep α with the SAME direction. Claim to prove: + (1) α=0 with the inverted xT reconstructs the REAL anchor cell (vs the generic random-xT α=0), and + (2) the α-sweep morph toward the KO is preserved. + +Same NTC anchor cells v5 uses (top-rank NTC), KIF23 + POLR1B, α 0->5. Everything reused from the +existing traversal stack; the only new step is _ddim(..., inverse=True) to get xT. + + python -m ops_model.models.attention.diffex.directions.proto_ddim_anchors --submit +""" +from __future__ import annotations + +import argparse +from pathlib import Path + +import numpy as np +import pandas as pd +import torch + +from ..classifier.config import slugify +from ..diffae.data import normalize +from .config import DirConfig +from .rank import supervised_direction +from .traverse import _ddim_guided, _sample_guided, load_diffae + +ANALYSIS = "/hpc/projects/icd.fast.ops/analysis" +DD = "/hpc/projects/icd.fast.ops/models/diffex/diffae" +DIR_CACHE = "/hpc/projects/icd.fast.ops/models/diffex/viewer_assets/_directions" + +# (label, marker_channel|None, raw_channel, diffae_ckpt, gene). marker_channel=None → phase. +FLUOR_SPECS = [ + ("FastAct_CAPZB", "actin filament_FastAct_SPY555 Live Cell Dye", "mCherry", f"{DD}/fluor_FastAct/diffae_best.pt", "CAPZB"), + ("TOMM20_TOMM20", "Mitochondria_TOMM20", "CP1_mitochondria_TOMM20", f"{DD}/fluor_Mitochondria_TOMM20/diffae_best.pt", "TOMM20"), + ("NucleoLive_KIF23", "nucleus_NucleoLIVE Live Cell dye", "mCherry", f"{DD}/fluor_NucleoLive/diffae_best.pt", "KIF23"), +] +PHASE_SPECS = [("phase_KIF23", None, "Phase2D", f"{DD}/phase_v1/diffae_best.pt", "KIF23"), + ("phase_POLR1B", None, "Phase2D", f"{DD}/phase_v1/diffae_best.pt", "POLR1B")] + + +def _pearson(a, b) -> float: + a, b = a.ravel(), b.ravel() + a, b = a - a.mean(), b - b.mean() + return float((a * b).sum() / (np.sqrt((a * a).sum() * (b * b).sum()) + 1e-12)) + + +@torch.no_grad() +def _direction(cfg, gene, modality, ctrl_embs, mu_ctrl, gather, dev): + """Load the cached control→KO direction, else compute it. Returns (d_vec, gap, lr_w, lr_b) — + the LR probe scores re-encoded morphs to measure phenotype strength across α.""" + dcache = Path(DIR_CACHE) / modality / "geneKO" / f"{slugify(gene)}.npz" + if dcache.exists(): + z = np.load(dcache); return z["d_vec"], float(z["gap"]), z["lr_w"], float(z["lr_b"]) + _, kd_embs = gather(cfg, gene, 1000) + embs = np.concatenate([kd_embs, ctrl_embs], 0) + labels = np.concatenate([np.ones(len(kd_embs)), np.zeros(len(ctrl_embs))]).astype(int) + d_vec, lr_w, lr_b, _ = supervised_direction(embs, labels, cfg) + return d_vec, float(np.linalg.norm(kd_embs.mean(0) - mu_ctrl)), lr_w, float(lr_b) + + +@torch.no_grad() +def run(specs, n_cells=4, alphas=(0, 1, 2, 3, 4, 5), ws=(1.0, 1.5, 2.0, 3.0), out_name="ddim_anchors", device="cuda"): + from ..viewer.precompute import _gather_class # gather NTC anchors (imgs + CellDINO embs) + dev = torch.device(device if torch.cuda.is_available() else "cpu") + out = Path(ANALYSIS) / out_name; out.mkdir(parents=True, exist_ok=True) + import matplotlib + matplotlib.use("Agg"); matplotlib.rcParams["pdf.fonttype"] = 42 + import matplotlib.pyplot as plt + import matplotlib.patheffects as pe + metrics = {} + for label, mc, ch, ckpt, gene in specs: + cfg = DirConfig(grain="geneKO", target=gene, control="NTC", device=device) + cfg.channel = ch; cfg.diffae_ckpt = ckpt + modality = slugify(mc) if mc else "phase" + if mc: # fluor: read the marker's cells once + cfg.marker_channel = mc + _cols = {"gene", "channel", "experiment", "well", "segmentation", "x_pheno", "y_pheno", "rank_type", "rank"} + _all = pd.read_csv(cfg.fluor_csv, usecols=lambda c: c in _cols) + cfg._fluor_rows = _all[(_all["channel"] == mc) & (_all["rank_type"] == "top")] + diffae = load_diffae(cfg, dev); null = diffae.null_emb.detach()[None].to(dev); H = cfg.crop_size + # NTC control cells: 1000 for direction, first n_cells as anchors + ntc_imgs, ntc_embs = _gather_class(cfg, "NTC", 1000) + mu_ctrl = ntc_embs.mean(0) + x0 = normalize(ntc_imgs[:n_cells]); x0t = torch.as_tensor(x0, dtype=torch.float32, device=dev) + z0 = torch.as_tensor(ntc_embs[:n_cells], dtype=torch.float32, device=dev) + d_vec, gap, lr_w, lr_b = _direction(cfg, gene, modality, ntc_embs, mu_ctrl, _gather_class, dev) + d = torch.as_tensor(d_vec, dtype=torch.float32, device=dev)[None] + tgt = label + from ..classifier.celldino_features import embed_crops + xT_rand = torch.cat([torch.randn(1, 1, H, H, generator=torch.Generator(device=dev).manual_seed(1234 + c), device=dev) + for c in range(n_cells)], 0) + for w in ws: # sweep w: gather/direction reused, only xT+decode redo + xT_inv = torch.cat([_ddim_guided(diffae, x0t[c:c + 1], z0[c:c + 1], null, w, cfg, inverse=True) + for c in range(n_cells)], 0) + gen_inv = np.empty((n_cells, len(alphas), H, H), np.float32); gen_rand = np.empty_like(gen_inv) + for c in range(n_cells): + for ai, a in enumerate(alphas): + cond = z0[c:c + 1] + (a * gap) * d + gen_inv[c, ai] = _sample_guided(diffae, xT_inv[c:c + 1].clone(), cond, null, w, cfg).cpu().numpy()[0, 0] + gen_rand[c, ai] = _sample_guided(diffae, xT_rand[c:c + 1].clone(), cond, null, w, cfg).cpu().numpy()[0, 0] + r_inv = float(np.mean([_pearson(gen_inv[c, 0], x0[c, 0]) for c in range(n_cells)])) + r_rand = float(np.mean([_pearson(gen_rand[c, 0], x0[c, 0]) for c in range(n_cells)])) + # morph strength: re-encode the inverted sweep, score with the direction's LR probe (logit α=max − α=0) + gemb = embed_crops(gen_inv.reshape(-1, 1, H, H).astype(np.float32), cfg, cache_path=None) + logits = (gemb @ lr_w + lr_b).reshape(n_cells, len(alphas)) + a0, amax = (alphas.index(0) if 0 in alphas else 0), int(np.argmax(alphas)) + morph_shift = float(np.mean(logits[:, amax] - logits[:, a0])) + metrics[f"{tgt}_w{w:g}"] = {"w": w, "alpha0_pearson_inverted": round(r_inv, 3), + "alpha0_pearson_random": round(r_rand, 3), + "morph_logit_shift_a0_to_amax": round(morph_shift, 3), "gap": round(gap, 3)} + print(f"[{tgt} w={w:g}] alpha0 Pearson inv={r_inv:.3f} rand={r_rand:.3f} morph_shift={morph_shift:.2f}") + ncols = 1 + len(alphas) + fig, ax = plt.subplots(2 * n_cells, ncols, figsize=(1.5 * ncols, 1.5 * 2 * n_cells), squeeze=False) + for c in range(n_cells): + for row, gen, seed_lbl in [(2 * c, gen_rand, "random xT"), (2 * c + 1, gen_inv, "inverted xT")]: + ax[row, 0].imshow(x0[c, 0], cmap="gray", vmin=-1, vmax=1) + ax[row, 0].set_ylabel(f"cell{c}\n{seed_lbl}", fontsize=7) + for ai, a in enumerate(alphas): + axi = ax[row, ai + 1]; axi.imshow(gen[c, ai], cmap="gray", vmin=-1, vmax=1) + if a == 0: + axi.text(0.04, 0.96, f"r={_pearson(gen[c, ai], x0[c, 0]):.2f}", transform=axi.transAxes, + fontsize=7, color="white", va="top", + path_effects=[pe.withStroke(linewidth=1.5, foreground="black")]) + for a_ax in ax[row]: + a_ax.set_xticks([]); a_ax.set_yticks([]) + for j, t in enumerate(["REAL"] + [f"α={a}" for a in alphas]): + ax[0, j].set_title(t, fontsize=8) + fig.suptitle(f"{tgt} w={w:g} — inverted vs random " + f"(α=0 Pearson inv {r_inv:.2f} vs rand {r_rand:.2f}; morph {morph_shift:.1f})", fontsize=9) + fig.tight_layout() + fig.savefig(out / f"ddim_anchor_{slugify(tgt)}_w{w:g}.png", dpi=150, bbox_inches="tight"); plt.close(fig) + + import json + (out / "metrics.json").write_text(json.dumps(metrics, indent=2)) + return metrics + + +def main(): + ap = argparse.ArgumentParser() + ap.add_argument("--set", choices=["fluor", "phase"], default="fluor") + ap.add_argument("--n-cells", type=int, default=4) + ap.add_argument("--ws", type=float, nargs="+", default=[1.0, 1.5, 2.0, 3.0], help="w values to sweep in ONE job") + ap.add_argument("--submit", action="store_true") + args = ap.parse_args() + specs = FLUOR_SPECS if args.set == "fluor" else PHASE_SPECS + out_name = f"ddim_anchors_{args.set}_wsweep" + if args.submit: + from ops_utils.hpc.slurm_batch_utils import submit_parallel_jobs + submit_parallel_jobs(jobs_to_submit=[{"name": f"ddim_wsweep_{args.set}", "func": run, + "kwargs": {"specs": specs, "n_cells": args.n_cells, "ws": tuple(args.ws), "out_name": out_name}}], + experiment="diffae", slurm_params={"slurm_partition": "gpu", "gpus_per_node": 1, + "cpus_per_task": 8, "mem_gb": 96, "timeout_min": 120, + "slurm_constraint": "[a100_80|h100|h200|6000_blackwell]"}, log_dir="diffae", + wait_for_completion=False) + else: + run(specs, n_cells=args.n_cells, ws=tuple(args.ws), out_name=out_name) + + +if __name__ == "__main__": + main() diff --git a/src/ops_model/models/attention/diffex/directions/rank.py b/src/ops_model/models/attention/diffex/directions/rank.py new file mode 100644 index 0000000..c40f62e --- /dev/null +++ b/src/ops_model/models/attention/diffex/directions/rank.py @@ -0,0 +1,59 @@ +"""Stage 2b: rank the K directions by control-vs-target classifier score shift. + +Fit a logistic regression (control=0, target=1) on the embeddings (post-hoc, never +used during 2a). For each trained direction, shift embeddings ±alpha and measure the +mean signed change in classifier logit. The direction with the largest |shift| is the +target axis. Returns (best_k, shifts, lr_weight, lr_bias).""" +from __future__ import annotations + +import numpy as np +import torch +from sklearn.linear_model import LogisticRegression + + +def supervised_direction(embs: np.ndarray, labels: np.ndarray, cfg): + """Deterministic control→KD direction (plan C primary). No training, no seeds: + reproducible by construction, and +α → toward KD by construction. + + - 'mean_diff' : normalize(mean(KD emb) − mean(control emb)) [centroid vector] + - 'lr_weight' : normalize(logistic-regression weight) [covariance-aware] + + Returns (d_unit (D,), lr_weight (D,), lr_bias, lr_acc). The LR is kept for the + re-encoded score verification only. + """ + clf = LogisticRegression(max_iter=2000, C=1.0).fit(embs, labels) + lr_w = clf.coef_[0].astype(np.float32) + acc = float(clf.score(embs, labels)) + if cfg.direction_method == "lr_weight": + d = lr_w.copy() + else: # mean_diff (default) + d = (embs[labels == 1].mean(0) - embs[labels == 0].mean(0)).astype(np.float32) + d = d / (np.linalg.norm(d) + 1e-9) + print(f"[direction] supervised '{cfg.direction_method}' (deterministic); LR acc={acc:.3f}") + return d.astype(np.float32), lr_w, float(clf.intercept_[0]), acc + + +def rank_directions(bank, embs: np.ndarray, labels: np.ndarray, cfg, dev): + clf = LogisticRegression(max_iter=2000, C=1.0).fit(embs, labels) + w = torch.as_tensor(clf.coef_[0], dtype=torch.float32, device=dev) + acc = float(clf.score(embs, labels)) + + bank = bank.to(dev).eval() + n = min(1024, len(embs)) + z = torch.as_tensor(embs[:n], dtype=torch.float32, device=dev) + a = cfg.rank_alpha + shifts = [] + with torch.no_grad(): + for k in range(bank.K): + d = bank.direction(z, k) # (n,D) + shift = (((z + a * d) @ w) - ((z - a * d) @ w)).mean().item() + shifts.append(shift) + best_k = int(max(range(bank.K), key=lambda k: abs(shifts[k]))) + # orient so +α always increases the KD score (the MLP's sign is arbitrary): + # +α → more KO-like, −α → more control-like, consistently across targets. + sign = 1.0 if shifts[best_k] >= 0 else -1.0 + print(f"[rank] LR train acc={acc:.3f}; per-direction score shift: " + + ", ".join(f"{k}:{s:+.3f}" for k, s in enumerate(shifts))) + print(f"[rank] selected direction {best_k} (|shift|={abs(shifts[best_k]):.3f}), " + f"orient sign={sign:+.0f} (+α = toward KO)") + return best_k, shifts, clf.coef_[0].astype(np.float32), float(clf.intercept_[0]), acc, sign diff --git a/src/ops_model/models/attention/diffex/directions/run.py b/src/ops_model/models/attention/diffex/directions/run.py new file mode 100644 index 0000000..3f9a37b --- /dev/null +++ b/src/ops_model/models/attention/diffex/directions/run.py @@ -0,0 +1,114 @@ +"""Orchestrator for Stage 3 (directions → ranking → traversal). + + python -m ops_model.models.attention.diffex.directions.run --target HSPA5 + python -m ops_model.models.attention.diffex.directions.run --grain complex \ + --target "Chaperonin-containing T-complex" + +Steps: gather target+control crops/embeddings → train K direction MLPs (unsupervised) +→ rank by control-vs-target LR score shift → traverse the selected direction on control +cells, DDIM-sample, verify monotonic re-encoded score. +""" +from __future__ import annotations + +import argparse +import json +from pathlib import Path + +import numpy as np +import torch + +from ..classifier.config import DEFAULT_OUT_ROOT, GRAINS, slugify +from ..diffae.data import normalize +from .config import DirConfig +from .data import gather +from .model import DirectionBank +from .rank import rank_directions, supervised_direction +from .traverse import load_diffae, traverse +from .train_directions import train_directions + + +def run_directions(cfg: DirConfig, out_dir: str) -> dict: + dev = torch.device(cfg.device if torch.cuda.is_available() or cfg.device == "cpu" else "cpu") + out = Path(out_dir); cache = out / "cache"; cache.mkdir(parents=True, exist_ok=True) + tag = f"{slugify(cfg.target)}_{cfg.crop_size}" + + # gather + images, embs, labels = gather( + cfg, str(cache / f"crops_{tag}.npz"), str(cache / f"celldino_{tag}.npz")) + + # ---- direction ---- + if cfg.deterministic: # repeatable runs (plan C) + torch.manual_seed(cfg.seed); np.random.seed(cfg.seed) + torch.backends.cudnn.deterministic = True + torch.backends.cudnn.benchmark = False + + if cfg.direction_method in ("mean_diff", "lr_weight"): + # PRIMARY: deterministic supervised control→KD direction (global, +α = toward KO). + d_vec, lr_w, lr_b, lr_acc = supervised_direction(embs, labels, cfg) + fixed_dir = torch.as_tensor(d_vec, dtype=torch.float32)[None] # (1,D) + bank, best_k, shifts, sign = None, None, None, 1.0 + else: + # SECONDARY: the paper's unsupervised InfoNCE bank (not reproducible run-to-run). + bank = DirectionBank(cfg.cond_dim, cfg.K, cfg.hidden) + train_directions(bank, embs, cfg, dev) + torch.save(bank.state_dict(), out / "direction_bank.pt") + best_k, shifts, lr_w, lr_b, lr_acc, sign = rank_directions(bank, embs, labels, cfg, dev) + if not cfg.orient_sign: + sign = 1.0 + fixed_dir = None + + # 3: traverse control cells. Scale α to the control→KD embedding gap so + # α=+1 ≈ a full traversal (unit directions × small α barely move otherwise). + gap = 1.0 + if cfg.scale_alpha_to_gap: + gap = float(np.linalg.norm(embs[labels == 1].mean(0) - embs[labels == 0].mean(0))) + print(f"[traverse] control→KD gap ‖μ_KD−μ_ctrl‖ = {gap:.2f}; α scaled by it") + ctrl_idx = np.flatnonzero(labels == 0)[: cfg.n_traverse] + src_imgs = normalize(images[ctrl_idx]) + src_embs = embs[ctrl_idx] + kd_idx = np.flatnonzero(labels == 1)[: cfg.n_traverse] # REAL KD cells for reference column + kd_imgs = normalize(images[kd_idx]) if len(kd_idx) else src_imgs + diffae = load_diffae(cfg, dev) + + # sweep guidance scale w (w=1 = plain conditional; w>1 amplifies the embedding edit) + sweep = {} + for w in cfg.guidance_scales: + sc = traverse(diffae, bank, best_k, src_imgs, src_embs, lr_w, lr_b, cfg, dev, out, + gap=gap, w=w, real_kd=kd_imgs, sign=sign, fixed_dir=fixed_dir) + np.save(out / f"scores_w{w:g}.npy", sc) # per-cell,per-alpha scores (for GIF cell pick) + mono = float(np.mean([np.all(np.diff(s) > 0) or np.all(np.diff(s) < 0) for s in sc])) + sweep[f"w{w:g}"] = {"mean_score_delta": float((sc[:, -1] - sc[:, 0]).mean()), + "frac_monotonic": mono} + print(f"[w={w:g}] mean_score_delta={sweep[f'w{w:g}']['mean_score_delta']:.2f} " + f"frac_monotonic={mono:.2f}") + metrics = { + "target": cfg.target, "grain": cfg.grain, "direction_method": cfg.direction_method, + "best_direction": best_k, "lr_acc": lr_acc, "score_shifts": shifts, "gap": gap, + "guidance_sweep": sweep, "n_traverse": int(len(ctrl_idx)), + } + (out / "metrics.json").write_text(json.dumps(metrics, indent=2)) + print(json.dumps(metrics, indent=2)) + return metrics + + +def main(): + ap = argparse.ArgumentParser(description="DiffEx Stage 3: directions + traversal") + ap.add_argument("--grain", choices=list(GRAINS), default="geneKO") + ap.add_argument("--target", default="HSPA5") + ap.add_argument("--K", type=int, default=10) + ap.add_argument("--dir-epochs", type=int, default=100) + ap.add_argument("--device", default="cuda") + ap.add_argument("--diffae-ckpt", default=None) + ap.add_argument("--out-dir", default=None) + args = ap.parse_args() + + cfg = DirConfig(grain=args.grain, target=args.target, K=args.K, + dir_epochs=args.dir_epochs, device=args.device) + if args.diffae_ckpt: + cfg.diffae_ckpt = args.diffae_ckpt + out = args.out_dir or f"{DEFAULT_OUT_ROOT}/directions/phase/{args.grain}/{slugify(args.target)}" + run_directions(cfg, out) + + +if __name__ == "__main__": + main() diff --git a/src/ops_model/models/attention/diffex/directions/submit.py b/src/ops_model/models/attention/diffex/directions/submit.py new file mode 100644 index 0000000..43e7ef0 --- /dev/null +++ b/src/ops_model/models/attention/diffex/directions/submit.py @@ -0,0 +1,60 @@ +"""Submit Stage 3 (directions + traversal) to SLURM (1 GPU). + + python -m ops_model.models.attention.diffex.directions.submit --target HSPA5 +""" +from __future__ import annotations + +import argparse +from pathlib import Path + +from ops_utils.hpc.slurm_batch_utils import submit_parallel_jobs + +from ..classifier.config import DEFAULT_OUT_ROOT, GRAINS, slugify +from .config import DirConfig +from .run import run_directions + + +def main(): + ap = argparse.ArgumentParser(description="Submit DiffEx Stage 3 to SLURM") + ap.add_argument("--grain", choices=list(GRAINS), default="geneKO") + ap.add_argument("--target", default="HSPA5") + ap.add_argument("--K", type=int, default=10) + ap.add_argument("--dir-epochs", type=int, default=100) + ap.add_argument("--diffae-ckpt", default=None) + ap.add_argument("--no-orient", action="store_true", + help="raw MLP sign (pre-orientation); default orients +α toward KO") + ap.add_argument("--out-dir", default=None) + ap.add_argument("--partition", default="gpu") + ap.add_argument("--gres", default="gpu:1") + ap.add_argument("--cpus", type=int, default=8) + ap.add_argument("--mem-gb", type=int, default=64) + ap.add_argument("--time-min", type=int, default=180) + ap.add_argument("--dry-run", action="store_true") + args = ap.parse_args() + + cfg = DirConfig(grain=args.grain, target=args.target, K=args.K, + dir_epochs=args.dir_epochs, device="cuda", + orient_sign=not args.no_orient) + if args.diffae_ckpt: + cfg.diffae_ckpt = args.diffae_ckpt + out = args.out_dir or f"{DEFAULT_OUT_ROOT}/directions/phase/{args.grain}/{slugify(args.target)}" + + jobs = [{ + "name": f"diffex_dir_{slugify(args.target)}"[:64], + "func": run_directions, + "kwargs": {"cfg": cfg, "out_dir": str(Path(out).resolve())}, + "metadata": {"stage": "directions", "target": args.target}, + }] + slurm_params = { + "slurm_partition": args.partition, "slurm_gres": args.gres, + "cpus_per_task": args.cpus, "mem_gb": args.mem_gb, "timeout_min": args.time_min, + } + submit_parallel_jobs( + jobs_to_submit=jobs, experiment="diffex_directions", + slurm_params=slurm_params, log_dir="diffex_directions", + dry_run=args.dry_run, wait_for_completion=False, + ) + + +if __name__ == "__main__": + main() diff --git a/src/ops_model/models/attention/diffex/directions/train_directions.py b/src/ops_model/models/attention/diffex/directions/train_directions.py new file mode 100644 index 0000000..57315ac --- /dev/null +++ b/src/ops_model/models/attention/diffex/directions/train_directions.py @@ -0,0 +1,36 @@ +"""Stage 2a: train the DirectionBank UNSUPERVISED on CellDINO embeddings.""" +from __future__ import annotations + +import numpy as np +import torch +from torch.utils.data import DataLoader, TensorDataset + +from .losses import decorrelation, infonce_directions + + +def train_directions(bank, embs: np.ndarray, cfg, dev) -> dict: + bank = bank.to(dev) + opt = torch.optim.AdamW(bank.parameters(), lr=cfg.dir_lr) + loader = DataLoader( + TensorDataset(torch.as_tensor(embs, dtype=torch.float32)), + batch_size=cfg.batch_size, shuffle=True, drop_last=True, + ) + history = [] + for ep in range(cfg.dir_epochs): + bank.train() + tot = nce = dec = 0.0 + for (z,) in loader: + z = z.to(dev) + dirs = bank(z) + l_nce = infonce_directions(dirs, cfg.tau) + l_dec = decorrelation(dirs) + loss = l_nce + cfg.decorr_weight * l_dec + opt.zero_grad() + loss.backward() + opt.step() + tot += float(loss); nce += float(l_nce); dec += float(l_dec) + nb = len(loader) + history.append({"epoch": ep, "loss": tot / nb, "infonce": nce / nb, "decorr": dec / nb}) + if ep % 10 == 0 or ep == cfg.dir_epochs - 1: + print(f"dir epoch {ep:03d}: loss={tot/nb:.4f} infonce={nce/nb:.4f} decorr={dec/nb:.4f}") + return {"history": history} diff --git a/src/ops_model/models/attention/diffex/directions/traverse.py b/src/ops_model/models/attention/diffex/directions/traverse.py new file mode 100644 index 0000000..edd2840 --- /dev/null +++ b/src/ops_model/models/attention/diffex/directions/traverse.py @@ -0,0 +1,224 @@ +"""Stage 3: traverse the selected direction and DDIM-sample a counterfactual strip. + +For each control cell: DDIM-invert to x_T (conditioned on its embedding z0), then for +each alpha reverse x_T conditioned on (z0 + alpha*d) -> image. Verify by RE-ENCODING +each generated image through CellDINO and scoring with the ranking classifier — the +score should move monotonically with alpha (the faithfulness check).""" +from __future__ import annotations + +from pathlib import Path + +import numpy as np +import torch +from diffusers import DDIMInverseScheduler, DDIMScheduler + +from ..classifier.celldino_features import embed_crops +from ..classifier.config import slugify +from ..diffae.config import DiffAEConfig +from ..diffae.model import DiffAE + + +def load_diffae(cfg, dev): + dcfg = DiffAEConfig( + crop_size=cfg.crop_size, cond_dim=cfg.cond_dim, + block_out_channels=cfg.block_out_channels, + layers_per_block=cfg.layers_per_block, train_timesteps=cfg.train_timesteps, + ) + m = DiffAE(dcfg) + m.load_state_dict(torch.load(cfg.diffae_ckpt, map_location="cpu")) + return m.to(dev).eval() + + +@torch.no_grad() +def _ddim(diffae, x, emb, cfg, inverse: bool): + sched = (DDIMInverseScheduler if inverse else DDIMScheduler)( + num_train_timesteps=cfg.train_timesteps) + sched.set_timesteps(cfg.ddim_steps) + c = diffae.cond(emb) + for t in sched.timesteps: + x = sched.step(diffae.denoise(x, t, c), t, x).prev_sample + return x + + +@torch.no_grad() +def _ddim_guided(diffae, x, emb, emb_base, w, cfg, inverse: bool): + """DDIM (forward=inverse) with the SAME classifier-free guidance as _sample_guided: + ε̃ = ε(base) + w·(ε(emb) − ε(base)). Inverting at the same w the morph samples at keeps + the α=0 round-trip faithful. w=1 reduces to plain ε(emb).""" + sched = (DDIMInverseScheduler if inverse else DDIMScheduler)(num_train_timesteps=cfg.train_timesteps) + sched.set_timesteps(cfg.ddim_steps) + c, c0 = diffae.cond(emb), diffae.cond(emb_base) + for t in sched.timesteps: + if w == 1.0: + eps = diffae.denoise(x, t, c) + else: + ec, e0 = diffae.denoise(torch.cat([x, x], 0), t, torch.cat([c, c0], 0)).chunk(2, 0) + eps = e0 + w * (ec - e0) + x = sched.step(eps, t, x).prev_sample + return x + + +@torch.no_grad() +def _sample_guided(diffae, xT, emb, emb_base, w, cfg): + """DDIM sample from xT. Edit-guidance: ε̃ = ε(base) + w·(ε(emb) − ε(base)). + w=1 → plain ε(emb); w>1 amplifies the embedding edit's effect on the image. + Speed: both guidance forwards (c, c0) run as ONE batched UNet call, the base ε(c0) is computed + once per step (was 2×), and the UNet runs under fp16 autocast (scheduler math stays fp32).""" + fwd = DDIMScheduler(num_train_timesteps=cfg.train_timesteps) + fwd.set_timesteps(cfg.ddim_steps) + c = diffae.cond(emb) + x = xT + if w == 1.0: + for t in fwd.timesteps: + with torch.autocast("cuda", dtype=torch.float16): + e = diffae.denoise(x, t, c).float() + x = fwd.step(e, t, x).prev_sample + return x + cc = torch.cat([c, diffae.cond(emb_base)], 0) # batch both conditionings into one forward + for t in fwd.timesteps: + with torch.autocast("cuda", dtype=torch.float16): + ec, ec0 = diffae.denoise(torch.cat([x, x], 0), t, cc).float().chunk(2, 0) + x = fwd.step(ec0 + w * (ec - ec0), t, x).prev_sample + return x + + +@torch.no_grad() +def traverse(diffae, bank, best_k, src_imgs_norm, src_embs, lr_w, lr_b, cfg, dev, out_dir, + gap: float = 1.0, w: float = 1.0, real_kd=None, sign: float = 1.0, + fixed_dir=None): + """src_imgs_norm (M,1,H,W) in [-1,1]; src_embs (M,cond_dim). + gap = ‖μ_KD−μ_ctrl‖ (α scaled by it); w = edit-guidance scale. + fixed_dir (1,D): if given, use this global deterministic direction (plan C) for + every cell instead of the per-cell unsupervised bank direction (reproducible).""" + out_dir = Path(out_dir) + alphas = list(cfg.alphas) + H = cfg.crop_size + # true classifier-free guidance: baseline = the learned NULL embedding, so + # ε̃ = ε(∅) + w·(ε(z0+αd) − ε(∅)). w=1 = plain conditional sampling. + null_base = diffae.null_emb.detach()[None].to(dev) + gen_imgs = [] # (M, n_alpha) generated patches + for i in range(len(src_imgs_norm)): + z0 = torch.as_tensor(src_embs[i:i + 1], dtype=torch.float32).to(dev) + if fixed_dir is not None: # plan C: deterministic global dir + d = fixed_dir.to(dev) # (1,D), already control→KO oriented + else: + d = sign * bank.direction(z0, best_k) # (1,D); +α → toward KO + # Fixed random noise per cell (constant across α): only the embedding edit + # changes along a row. α=0 = DDPM recon of z0 (Alex's design), not the real image. + g = torch.Generator(device=dev).manual_seed(1234 + i) + xT = torch.randn(1, 1, H, H, generator=g, device=dev) + row = [] + for a in alphas: + img = _sample_guided(diffae, xT.clone(), z0 + (a * gap) * d, null_base, w, cfg) + row.append(img.cpu().numpy()[0, 0]) + gen_imgs.append(row) + gen = np.array(gen_imgs) # (M, A, H, W) + + # verify: re-encode generated images -> CellDINO -> LR score + flat = gen.reshape(-1, 1, gen.shape[-2], gen.shape[-1]).astype(np.float32) + gen_embs = embed_crops(flat, cfg, cache_path=None) # (M*A, cond_dim) + scores = (gen_embs @ lr_w + lr_b).reshape(len(src_imgs_norm), len(alphas)) + + _plot(gen, scores, alphas, cfg, out_dir, w, real_ctrl=src_imgs_norm, real_kd=real_kd) + return scores + + +def _plot(gen, scores, alphas, cfg, out_dir, w=1.0, real_ctrl=None, real_kd=None): + import matplotlib + matplotlib.use("Agg") + matplotlib.rcParams["pdf.fonttype"] = 42 + import matplotlib.pyplot as plt + + out_dir.mkdir(parents=True, exist_ok=True) + M, A = gen.shape[0], gen.shape[1] + c = alphas.index(0.0) if 0.0 in alphas else A // 2 # α=0 reconstruction = baseline + # Columns = [REAL control | generated α-sweep | REAL KD]. Middle A are generated; + # far-left/right are REAL reference images. Off-center generated tiles overlay a + # diverging Δ-pixel heatmap vs the α=0 reconstruction (red gained, blue lost). + have_ref = real_ctrl is not None and real_kd is not None + ncols = A + (2 if have_ref else 0) + off = 1 if have_ref else 0 + tag = slugify(cfg.target) + + def render_row(axrow, i, titles): + vmax = float(max(np.abs(gen[i] - gen[i, c]).max(), 1e-3)) # per-cell scale + if have_ref: + axrow[0].imshow(real_ctrl[i, 0], cmap="gray", vmin=-1, vmax=1, interpolation="nearest") + axrow[-1].imshow(real_kd[i % len(real_kd), 0], cmap="gray", vmin=-1, vmax=1, + interpolation="nearest") + if titles: + axrow[0].set_title("REAL ctrl", fontsize=9) + axrow[-1].set_title("REAL KD", fontsize=9) + for j in range(A): + ax = axrow[j + off] + ax.imshow(gen[i, j], cmap="gray", vmin=-1, vmax=1, interpolation="nearest") + if j != c: + diff = gen[i, j] - gen[i, c] + amask = np.clip(np.abs(diff) / vmax, 0, 1) * 0.7 # transparent where unchanged + ax.imshow(diff, cmap="bwr", vmin=-vmax, vmax=vmax, alpha=amask, + interpolation="nearest") + if titles: + ax.set_title(f"{alphas[j]:+.2g}×gap", fontsize=9) + for ax in axrow: + ax.axis("off") + + # (1) overview grid — high DPI + nearest so each cell stays crisp + fig, axes = plt.subplots(M, ncols, figsize=(2.0 * ncols, 2.0 * M), squeeze=False) + for i in range(M): + render_row(axes[i], i, titles=(i == 0)) + fig.suptitle(f"{cfg.target}: counterfactual traversal (control → KD), guidance w={w:g}. " + f"Far-left/right = REAL refs; middle = generated; overlay = Δpx vs α=0 (red +, blue −)", + fontsize=11) + fig.tight_layout() + fig.savefig(out_dir / f"traversal_{tag}_w{w:g}.png", dpi=200, bbox_inches="tight") + plt.close(fig) + + # (2) one high-definition strip PNG per source cell, as 3 rows: + # row 0 = generated images; row 1 = Δ-pixel heatmap (own diverging colormap); + # row 2 = overlay (image + heatmap). REAL refs only in row 0. + strip_dir = out_dir / "strips" + strip_dir.mkdir(exist_ok=True) + row_labels = ["generated", "Δ-pixels", "overlay"] + for i in range(M): + vmax = float(max(np.abs(gen[i] - gen[i, c]).max(), 1e-3)) + fig, ax = plt.subplots(3, ncols, figsize=(2.7 * ncols, 8.6), squeeze=False) + for r in range(3): + for col in range(ncols): + ax[r, col].axis("off") + if have_ref: # real references only in the generated row + ax[0, 0].imshow(real_ctrl[i, 0], cmap="gray", vmin=-1, vmax=1, interpolation="nearest") + ax[0, 0].set_title("REAL ctrl", fontsize=9) + ax[0, ncols - 1].imshow(real_kd[i % len(real_kd), 0], cmap="gray", vmin=-1, vmax=1, + interpolation="nearest") + ax[0, ncols - 1].set_title("REAL KD", fontsize=9) + for j in range(A): + col = j + off + diff = gen[i, j] - gen[i, c] + ax[0, col].imshow(gen[i, j], cmap="gray", vmin=-1, vmax=1, interpolation="nearest") + ax[0, col].set_title(f"{alphas[j]:+.2g}×gap", fontsize=9) + ax[1, col].imshow(diff, cmap="bwr", vmin=-vmax, vmax=vmax, interpolation="nearest") + ax[2, col].imshow(gen[i, j], cmap="gray", vmin=-1, vmax=1, interpolation="nearest") + if j != c: + amask = np.clip(np.abs(diff) / vmax, 0, 1) * 0.7 + ax[2, col].imshow(diff, cmap="bwr", vmin=-vmax, vmax=vmax, alpha=amask, + interpolation="nearest") + for r, lab in enumerate(row_labels): + fig.text(0.008, 0.80 - r * 0.31, lab, rotation=90, va="center", fontsize=11, + fontweight="bold") + fig.suptitle(f"{cfg.target} — cell {i} (w={w:g}); Δ vs α=0 (red +, blue −)", fontsize=12) + fig.tight_layout(rect=[0.02, 0, 1, 1]) + fig.savefig(strip_dir / f"{tag}_w{w:g}_cell{i}.png", dpi=200, bbox_inches="tight") + plt.close(fig) + + # score progression + fig, ax = plt.subplots(figsize=(6, 4)) + for i in range(M): + ax.plot(alphas, scores[i], marker="o", alpha=0.6) + ax.plot(alphas, scores.mean(0), color="black", lw=2.5, label="mean") + ax.set_xlabel("α"); ax.set_ylabel("classifier logit (re-encoded gen image)") + ax.set_title(f"{cfg.target}: score progression (monotonic ⇒ direction valid)") + ax.axhline(0, color="gray", ls=":"); ax.legend(); ax.grid(alpha=0.3) + fig.tight_layout() + fig.savefig(out_dir / f"scores_{slugify(cfg.target)}_w{w:g}.png", dpi=130, bbox_inches="tight") + plt.close(fig) + print(f"[traverse] wrote traversal/scores for {cfg.target} (w={w:g})") diff --git a/src/ops_model/models/attention/diffex/figures/METHODS_final.txt b/src/ops_model/models/attention/diffex/figures/METHODS_final.txt new file mode 100644 index 0000000..54691e8 --- /dev/null +++ b/src/ops_model/models/attention/diffex/figures/METHODS_final.txt @@ -0,0 +1,70 @@ +13.2 Set Classifier +A neural network classification model was trained to input a set of cells receiving the same perturbation and predict the perturbation class from one of C labels. Within the machine learning literature, such a problem formulation is known as multiple instance learning [cite]. Separate models were trained depending on the particular definition of “perturbation class”. Classes could either be fine-grained (e.g. C=1001 classes containing 1000 genes in the knockout library plus NTC), or more coarse-grained (e.g. C=99 protein complex classes that each group together multiple genes) – depending on if the desired objective is to understand the uniqueness of each gene knockout or each complex. The model architecture was based on a Set Transformer [cite], which enforces permutation invariance (i.e. reordering of cells within the set does not change the final prediction) and enables multi-scale inference (i.e. the model can return a prediction regardless of set size). Within each set, cell embeddings were concatenated with a learnable channel embedding signifying each cell’s channel identity and then passed through a linear projection to yield a vector of d=512 dimensions per cell. The set of channel-aware embeddings were then passed through l=2 layers of inducing-point set attention blocks (ISAB) with p=32 inducing points. Compared to standard attention blocks, ISAB has computational complexity that scales linearly with set size, allowing the model to efficiently process sets containing up to tens of thousands of cells. The model employed h=4 attention heads per layer and a feedforward dimension of f = 4 * d = 2048. The final layers of the model consisted of a pooling by multi-head attention (PMA) layer that collapses the set of embeddings to a single embedding of dimension d, followed by a cosine classifier that yields a probability distribution over class labels. During model training, each training example was a set of n=100 cells receiving the same perturbation, sampled randomly from the pool of cells for that perturbation. A training epoch consisted of iterating through s=32 sets for each perturbation. Each model was trained for 200 epochs with 64 sets per batch, 0.0001 peak learning rate, and 20 epochs of linear warmup followed by cosine annealing. Even though the model was consistently trained with n=100 cells per set, it is possible for the set classifier to provide predictions for any value of n. As shown in Figure 4B, after the model was trained, it was evaluated on different values of n (ranging from 10 to 5,000) on the validation set. + +13.3 Attention-Based Ranking of Single Cells +The PMA layer of the model places different attention weights on the various cells to create a set-level embedding from the multiple cell-level embeddings. These weights can be used as a proxy for the importance of each cell in contributing to the final prediction. After the model was trained, all cells for a particular perturbation were simultaneously passed as a single set to the model and attention weights were derived for every cell from the PMA layer. Weights from the different attention heads were averaged to obtain a single scalar for each cell, which was then used to globally rank cells. + +13.4 Set-accuracy Distilled Phenotype Distinctiveness +The CellDINO mAP scores from Phase cells (n=65M cells) (Fig. 2X) were recomputed weighting each cell's embedding by a per-cell rank-derived score from a set classifier, with weights hard-zeroed for cells outside the top 20,000 per perturbation and normalized to mean=1 per (sgRNA, experiment) group. The rank-derived score is 1 - rank / N, where N is the number of cells for that perturbation, so higher-ranked cells contribute more. Post-weighting the same PCA post processing method was applied (see Supplementary Figure XX). The set-classifier variant matched to the readout was used to derive the ranking: the gene-KO classifier for gene-level mAP, and the protein-complex (EBI) classifier for complex-level mAP (Fig. 4D). + + + +Generative single-cell counterfactual traversals + +Generative model + +A diffusion autoencoder was trained on single-cell images (label-free phase, or individual fluorescent markers). Each cell is represented by two disentangled latents: a 1,024-dimensional semantic embedding obtained from CellDINO, and a stochastic (noise) latent that captures the residual identity of the individual cell — its size, texture, and local context. The decoder is a conditional diffusion model that reconstructs the cell image from the semantic embedding using classifier-free guidance (guidance scale w = 1.5, with a learned null embedding). Because phenotype and identity are encoded separately, a given cell can be re-synthesized while its phenotype is smoothly and independently varied. + +This generative approach adapts DiffEx (arXiv:2502.09663), a diffusion-autoencoder method for counterfactual explanations of image classifiers, with two main modifications. First, DiffEx steers each generation with the gradient of a differentiable image classifier (unsupervised classifier guidance); our classifier is instead a Set-Transformer that scores bags of CellDINO embeddings and so provides no per-image image→class gradient. We therefore replace gradient guidance with a single, translate the cell's semantic latent along a precomputed linear direction in the CellDINO embedding space between NTC→knockout. Second, the base cells and the cells that define each direction are selected by the set classifier's per-cell prediction ranking rather than sampled at random. + +Selecting representative cells + +For each perturbation, single cells were ranked by their contribution to the set classifier's predictive accuracy. Each cell was scored by its marginal effect on classification accuracy when included in randomly sampled sets of cells of the same perturbation while varying set size (set n=X-X); cells that most improve the set-level prediction rank highest. The top-ranked non-targeting control (NTC) cells (20 per condition) were used as the shared base cells for all traversals, so that every perturbation of a given imaging channel is shown as a morph of the same reference control cells. + +Perturbation direction and counterfactual traversal + +For each gene knockout, a perturbation direction was defined in the 1,024-dimensional Cell-DINO embedding space using the top 1,000 cells of the knockout population and of the control population, ranked by the same per-cell prediction contribution described above. The mean embedding of each of these top-cell sets defines that class's centroid. A linear (logistic) classifier fit to separate the two populations gives the direction — the unit normal of its decision boundary — pointing from the control state toward the knockout state. The direction was scaled by the gap between the two centroids (the distance between the average knockout and average NTC embedding), so that a unit step (α = 1) moves the control embedding onto the knockout centroid. + +A counterfactual traversal was then generated for each base cell by displacing its semantic embedding along this direction, z(α) = z₀ + α · gap · d, while holding its stochastic latent fixed, and decoding each α with the diffusion model. So that the series is anchored to the specific real base cell rather than to an arbitrary synthetic one, the stochastic latent was not drawn at random but recovered by DDIM inversion of the real base-cell image: the deterministic diffusion ODE was run in reverse under the cell's own semantic embedding and the same classifier-free guidance used for forward sampling (guided inversion), yielding the noise latent x_T that regenerates that exact cell. Seventeen α steps were sampled across α ∈ [−5, +5]: α = 0 reconstructs the real base cell, positive α interpolates toward the knockout phenotype and reaches the knockout mean at α = 1, and |α| > 1 extrapolates beyond the class means, exaggerating the phenotypic difference in either direction. Because the inverted stochastic latent is held fixed across the series, cell identity is anchored — α = 0 is a faithful reconstruction of the real base cell (pixel Pearson ≈ 0.99), and its size, texture, and context are preserved while only the phenotype changes — yielding a smooth traversal that begins from the true cell. In parallel, classical morphometric features were measured on both the real and the generated images to confirm that the traversal produces the expected phenotypic effect in interpretable feature space. + +The same framework was used to morph between two perturbation classes — for example, the 40S and 60S cytosolic ribosomal-subunit complexes. In this case the direction is the difference between the two class means (A → B), and the base cells of class A are morphed toward the mean of class B, again with α extrapolating beyond either mean to exaggerate the transition. + +Gene embedding and montage + +A gene-level phenotypic embedding was constructed by aggregating the single-cell CellDINO features of each perturbation, and traversing a single NTC cell to every geneKO at α=3. Each generated cell was placed at the gene's coordinate in a PHATE embedding (see Supplementary Figure 2X). Because many genes occupy dense regions, the embedding plane was tiled with a regular grid and a single representative gene was retained per grid cell — chosen by local density — so that adjacent cells do not overlap (scale of ~100 geneKOs selected). Each tile was outlined by the color of its geneKO's Leiden cluster to convey the local phenotypic neighborhood. The result is a single view in which each region of the phenotypic embedding is illustrated by the morph of a single representative cell. + + + + + + + + + +Predictive single cells per protein complex + +To illustrate what the complex-level set classifier responds to, real cells were displayed for four complexes of the curated EBI protein-complex panel, in label-free phase and in a matched fluorescent marker. Cells were ranked by their per-cell predictive ranking of the protein complex classifier— the complex-level counterpart of the per-cell ranking described above — computed for every cell of every member gene. Complex membership was taken from the curated EBI definitions. For each gene the three top-ranked cells were displayed beside the top-ranked non-targeting control (NTC) cells of the same channel, so the two phase panels show the identical control cells. Crops were taken from the assembled phenotyping volumes at native resolution using a single intensity window per channel, computed over the pooled knockout and control crops so that brightness is directly comparable within a panel. + +Protein complex morphometrics + +Each image panel was paired with the distribution of one interpretable organelle feature. For each class the top 1,000 predictive cells were used. Control cells were restricted to the experiments contributing that panel's knockout cells, making each comparison batch-matched. The panels report light-vesicle count and dark-vesicle count in phase, total ER/Golgi (COPE) object area, and peripheral lipid-droplet count. Peripheral lipid droplets radial position was measured from its centroid along the axis running from the nuclear boundary (0) to the cell edge (1), and droplets in the outer 20% of that axis were counted. + +Distributions are shown as violins of the per-cell values with a bar at the median and the percentage change of the median relative to the control. The vertical view is clipped to the 2nd–96th percentile of the pooled data; this crops the display only — kernel densities and all reported statistics are computed on the complete, unclipped per-cell values. + + + +EXCLUDE + + +The montage was assembled with the open-source latent-lens package (https://github.com/czi-ai/latent-lens; multiscale montages of image crops laid out by an embedding), which performs the grid-based, density-prioritized decimation and the multiscale tiling. +Quantifying the accuracy of generated phenotypes + +To test whether the synthesized cells reproduce the intended perturbation phenotype, each traversal was scored by three complementary metrics, all computed as a function of α and interpreted relative to the value attainable on real cells of the same class. + +Set-classifier recognition (supervised). The generated images at each α were embedded with Cell-DINO and passed as a set (a bag of n cells) through a SetTransformer classifier trained on real single cells to predict the perturbation class from a bag of same-class cells. Crucially, generated images sit at a systematic Cell-DINO domain offset from the real images the classifier was trained on, so raw features fail; because that offset is shared across a traversal, standardizing every α's embeddings against the α = 0 generated NTC control cancels it, leaving the classifier to read the phenotype rather than the real-vs-generated gap. Two readouts summarize the prediction: P(target), the softmax probability the classifier assigns to the true perturbation class — a continuous measure of how confidently the generated bag is recognized as the target; and target rank, the 1-indexed position of the true class in the classifier's ranking of all classes (rank 1 = the top prediction), reported as top-1 and top-5 recovery (whether the target falls within the classifier's top-1 / top-5 predictions). Rank is more forgiving than P(target) and localizes where the true class sits among competing phenotypes. Both are compared to the same classifier's accuracy on real cells of the class (the real-cell ceiling): a phenotype is considered recovered when, near α ≈ 1, the generated bag reaches the recognition level the real cells achieve. + +Retrieval mAP (unsupervised). As an independent check that does not rely on the trained classifier, we measure whether generated cells fall in the correct region of the embedding relative to real cells. Real and generated single-cell Cell-DINO embeddings are pooled, and for each generated class-X cell we compute the average precision of retrieving real class-X cells (cross-domain, real to generated) ahead of cells of every other class; averaging over cells gives a per-class mean average precision (mAP; computed with copairs, the same metric used for real-cell perturbation distinctiveness). The ceiling is obtained by splitting the real cells of each class into two halves and running the identical retrieval — the real self-consistency mAP — and the generated-to-real ratio reports how much of the real phenotypic distinctiveness the synthesized cells recover. Because it is purely a neighborhood-retrieval statistic in embedding space, this metric is orthogonal to the set classifier and guards against classifier-specific artifacts. + + + + diff --git a/src/ops_model/models/attention/diffex/figures/METHODS_traversal_montage.md b/src/ops_model/models/attention/diffex/figures/METHODS_traversal_montage.md new file mode 100644 index 0000000..5a8c7c6 --- /dev/null +++ b/src/ops_model/models/attention/diffex/figures/METHODS_traversal_montage.md @@ -0,0 +1,112 @@ +# Generative single-cell counterfactual traversals and the embedding montage + +## Generative model + +A diffusion autoencoder was trained on single-cell images (label-free phase, or a fluorescence marker). +Each cell is represented by two disentangled latents: a 1,024-dimensional semantic embedding obtained from +Cell-DINO, and a stochastic (noise) latent that captures the residual identity of +the individual cell — its size, texture, and local context. The decoder is a conditional diffusion model +that reconstructs the cell image from the semantic embedding using classifier-free guidance (guidance scale +w = 2.0, with a learned null embedding). Because phenotype and identity are encoded separately, a given +cell can be re-synthesized while its phenotype is smoothly and independently varied. + +This generative approach adapts DiffEx (arXiv:2502.09663), a diffusion-autoencoder method for +counterfactual explanations of image classifiers, with two main modifications. First, DiffEx uses +unsupervised directions discovered from classifier gradients; we instead use a supervised direction — +the NTC→knockout axis, fit as a linear classifier in the 1,024-dimensional Cell-DINO embedding space +(below) — and translate the cell's semantic latent along it. Second, the base cells and the cells that +define each direction are selected by the set classifier's per-cell accuracy ranking rather than sampled at +random. + +## Selecting representative cells + +For each perturbation, single cells were ranked by their contribution to the set classifier's predictive +accuracy. Each cell was scored by its marginal effect on classification accuracy when included in randomly +sampled sets of n = 50 cells of the same perturbation; cells that most improve the set-level prediction rank +highest. The top-ranked non-targeting control (NTC) cells (20 per condition) were used as the shared base cells for all traversals, so that every perturbation of a given imaging channel is shown as a morph of the same reference control cells. + +## Perturbation direction and counterfactual traversal + +For each gene knockout we take the top ~1,000 knockout cells and the top ~1,000 control (NTC) cells — by the +per-cell accuracy ranking above — in the 1,024-dimensional Cell-DINO space. Two quantities are then derived +separately: + +- **Direction** (which way to move): a logistic classifier is fit to separate the two populations, and its + decision-boundary normal (a unit vector), oriented from control toward knockout, is the traversal axis. +- **Step size** (how far): the gap between the class means — the distance from the control-cell mean to the + knockout-cell mean — so that α = 1 advances a control embedding by one full control→knockout mean-shift + (landing it on the knockout mean). + +A counterfactual traversal was then generated for each base cell by displacing its semantic embedding along +this direction, z(α) = z₀ + α · gap · d, while holding its stochastic latent fixed, and decoding each α with +the diffusion model. Seventeen steps were sampled across α ∈ [−5, +5]: α = 0 reconstructs the control cell, +positive α interpolates toward the knockout phenotype and reaches the knockout mean at α = 1, and |α| > 1 +extrapolates beyond the class means, exaggerating the phenotypic difference in either direction. Because the +stochastic latent is held fixed across the series, cell identity is anchored and only the phenotype changes, +yielding a smooth traversal. Each traversal was quantified by three complementary metrics computed as a +function of α (described in *Quantifying the accuracy of generated phenotypes*, below). In parallel, classical +CellProfiler / OrganelleProfiler morphometric features were measured on both the real and the generated +images to confirm that the traversal produces the expected phenotypic effect in interpretable feature space +(e.g. object area, circularity, or marker intensity changing monotonically with α). + +The same framework was used to morph between two perturbation classes — for example, the 40S and 60S +cytosolic ribosomal-subunit complexes. In this case the direction is the difference between the two class +means (A → B), and the base cells of class A are morphed toward the mean of class B, again with α +extrapolating beyond either mean to exaggerate the transition. + +## Quantifying the accuracy of generated phenotypes + +To test whether the synthesized cells reproduce the intended perturbation phenotype — rather than merely a +plausible-looking cell — each traversal was scored by three complementary metrics, all computed as a function +of α and interpreted relative to the value attainable on real cells of the same class. + +**Set-classifier recognition (supervised).** The generated images at each α were embedded with Cell-DINO and +passed as a set (a bag of n cells) through a SetTransformer classifier trained on real single cells to predict +the perturbation class from a bag of same-class cells. Crucially, generated images sit at a systematic Cell-DINO +domain offset from the real images the classifier was trained on, so raw features fail; because that offset is +shared across a traversal, standardizing every α's embeddings against the α = 0 *generated* NTC control cancels it, +leaving the classifier to read the phenotype rather than the real-vs-generated gap. Two readouts summarize the +prediction: + +- *P(target)* — the softmax probability the classifier assigns to the true perturbation class: a continuous + measure of how confidently the generated bag is recognized as the target. +- *Target rank* — the 1-indexed position of the true class in the classifier's ranking of all classes + (rank 1 = the top prediction), reported as top-1 and top-5 recovery (whether the target falls within the + classifier's top-1 / top-5 predictions). Rank is more forgiving than P(target) and localizes where the true + class sits among competing phenotypes. + +Both are compared to the same classifier's accuracy on real cells of the class (the real-cell ceiling): a +phenotype is considered recovered when, near α ≈ 1, the generated bag reaches the recognition level the real +cells achieve. + +**Retrieval mAP (unsupervised).** As an independent check that does not rely on the trained classifier, we +measure whether generated cells fall in the correct region of the embedding relative to real cells. Real and +generated single-cell Cell-DINO embeddings are pooled, and for each generated class-X cell we compute the +average precision of retrieving real class-X cells (cross-domain, real ↔ generated) ahead of cells of every +other class; averaging over cells gives a per-class mean average precision (mAP; computed with copairs, the +same metric used for real-cell perturbation distinctiveness). The ceiling is obtained by splitting the real +cells of each class into two halves and running the identical retrieval — the real self-consistency mAP — and +the generated-to-real ratio reports how much of the real phenotypic distinctiveness the synthesized cells +recover. Because it is purely a neighborhood-retrieval statistic in embedding space, this metric is orthogonal +to the set classifier and guards against classifier-specific artifacts. + +Across all three metrics a legitimate phenotype rises from control levels as α increases, peaks near α ≈ 1 (the +knockout / class mean), and approaches the corresponding real-cell reference; extrapolation to |α| > 1 +exaggerates the phenotype and, past the class mean, can overshoot and degrade recognition. + +## Gene embedding and montage + +A gene-level phenotypic embedding was constructed by aggregating the single-cell Cell-DINO features of each +perturbation, and a two-dimensional layout was computed with PHATE. To build the montage, each gene's +generated cell (at a chosen α) was placed at that gene's coordinate in this embedding. Because many genes +occupy dense regions, the embedding plane was tiled with a regular grid and a single representative gene was +retained per grid cell — chosen by local density — so that adjacent cells do not overlap; coarser grids show +fewer, larger cells and finer grids fill in more. Each tile was outlined by the color of its gene's Leiden +cluster to convey the local phenotypic neighborhood, with individual complex members highlighted (e.g. +ribo40S, ribo60S). The result is a single view in which each region of the phenotypic embedding is +illustrated by a representative generated cell. + +The montage was assembled with the open-source latent-lens package +(https://github.com/czi-ai/latent-lens; multiscale montages of image crops laid out by an embedding), +which performs the grid-based, density-prioritized decimation and the multiscale tiling. + diff --git a/src/ops_model/models/attention/diffex/figures/METHODS_traversal_montage.txt b/src/ops_model/models/attention/diffex/figures/METHODS_traversal_montage.txt new file mode 100644 index 0000000..8a5a616 --- /dev/null +++ b/src/ops_model/models/attention/diffex/figures/METHODS_traversal_montage.txt @@ -0,0 +1,38 @@ +Generative single-cell counterfactual traversals and the embedding montage + +Generative model + +A diffusion autoencoder was trained on single-cell images (label-free phase, or a fluorescence marker). Each cell is represented by two disentangled latents: a 1,024-dimensional semantic embedding obtained from Cell-DINO, and a stochastic (noise) latent that captures the residual identity of the individual cell — its size, texture, and local context. The decoder is a conditional diffusion model that reconstructs the cell image from the semantic embedding using classifier-free guidance (guidance scale w = 1.5, with a learned null embedding). Because phenotype and identity are encoded separately, a given cell can be re-synthesized while its phenotype is smoothly and independently varied. + +This generative approach adapts DiffEx (arXiv:2502.09663), a diffusion-autoencoder method for counterfactual explanations of image classifiers, with two main modifications. First, DiffEx uses unsupervised directions discovered from classifier gradients; we instead use a supervised direction — the NTC→knockout axis, fit as a linear classifier in the 1,024-dimensional Cell-DINO embedding space (below) — and translate the cell's semantic latent along it. Second, the base cells and the cells that define each direction are selected by the set classifier's per-cell accuracy ranking rather than sampled at random. + +Selecting representative cells + +For each perturbation, single cells were ranked by their contribution to the set classifier's predictive accuracy. Each cell was scored by its marginal effect on classification accuracy when included in randomly sampled sets of n = 50 cells of the same perturbation; cells that most improve the set-level prediction rank highest. The top-ranked non-targeting control (NTC) cells (20 per condition) were used as the shared base cells for all traversals, so that every perturbation of a given imaging channel is shown as a morph of the same reference control cells. + +Perturbation direction and counterfactual traversal + +For each gene knockout we take the top ~1,000 knockout cells and the top ~1,000 control (NTC) cells — by the per-cell accuracy ranking above — in the 1,024-dimensional Cell-DINO space. Two quantities are then derived separately: + +- Direction (which way to move): a logistic classifier is fit to separate the two populations, and its decision-boundary normal (a unit vector), oriented from control toward knockout, is the traversal axis. +- Step size (how far): the gap between the class means — the distance from the control-cell mean to the knockout-cell mean — so that α = 1 advances a control embedding by one full control→knockout mean-shift (landing it on the knockout mean). + +A counterfactual traversal was then generated for each base cell by displacing its semantic embedding along this direction, z(α) = z₀ + α · gap · d, while holding its stochastic latent fixed, and decoding each α with the diffusion model. So that the series is anchored to the specific real base cell rather than to an arbitrary synthetic one, the stochastic latent was not drawn at random but recovered by DDIM inversion of the real base-cell image: the deterministic diffusion ODE was run in reverse under the cell's own semantic embedding and the same classifier-free guidance used for forward sampling (guided inversion), yielding the noise latent x_T that regenerates that exact cell. Seventeen steps were sampled across α ∈ [−5, +5]: α = 0 reconstructs the real base cell, positive α interpolates toward the knockout phenotype and reaches the knockout mean at α = 1, and |α| > 1 extrapolates beyond the class means, exaggerating the phenotypic difference in either direction. Because the inverted stochastic latent is held fixed across the series, cell identity is anchored — α = 0 is a faithful reconstruction of the real base cell (pixel Pearson ≈ 0.99, versus ≈ 0.43 for a randomly sampled latent), and its size, texture, and context are preserved while only the phenotype changes — yielding a smooth traversal that begins from the true cell. Each traversal was quantified by three complementary metrics computed as a function of α (described in "Quantifying the accuracy of generated phenotypes", below). In parallel, classical CellProfiler / OrganelleProfiler morphometric features were measured on both the real and the generated images to confirm that the traversal produces the expected phenotypic effect in interpretable feature space (e.g. object area, circularity, or marker intensity changing monotonically with α). + +The same framework was used to morph between two perturbation classes — for example, the 40S and 60S cytosolic ribosomal-subunit complexes. In this case the direction is the difference between the two class means (A → B), and the base cells of class A are morphed toward the mean of class B, again with α extrapolating beyond either mean to exaggerate the transition. + +Quantifying the accuracy of generated phenotypes + +To test whether the synthesized cells reproduce the intended perturbation phenotype — rather than merely a plausible-looking cell — each traversal was scored by three complementary metrics, all computed as a function of α and interpreted relative to the value attainable on real cells of the same class. + +Set-classifier recognition (supervised). The generated images at each α were embedded with Cell-DINO and passed as a set (a bag of n cells) through a SetTransformer classifier trained on real single cells to predict the perturbation class from a bag of same-class cells. Crucially, generated images sit at a systematic Cell-DINO domain offset from the real images the classifier was trained on, so raw features fail; because that offset is shared across a traversal, standardizing every α's embeddings against the α = 0 generated NTC control cancels it, leaving the classifier to read the phenotype rather than the real-vs-generated gap. Two readouts summarize the prediction: P(target), the softmax probability the classifier assigns to the true perturbation class — a continuous measure of how confidently the generated bag is recognized as the target; and target rank, the 1-indexed position of the true class in the classifier's ranking of all classes (rank 1 = the top prediction), reported as top-1 and top-5 recovery (whether the target falls within the classifier's top-1 / top-5 predictions). Rank is more forgiving than P(target) and localizes where the true class sits among competing phenotypes. Both are compared to the same classifier's accuracy on real cells of the class (the real-cell ceiling): a phenotype is considered recovered when, near α ≈ 1, the generated bag reaches the recognition level the real cells achieve. + +Retrieval mAP (unsupervised). As an independent check that does not rely on the trained classifier, we measure whether generated cells fall in the correct region of the embedding relative to real cells. Real and generated single-cell Cell-DINO embeddings are pooled, and for each generated class-X cell we compute the average precision of retrieving real class-X cells (cross-domain, real to generated) ahead of cells of every other class; averaging over cells gives a per-class mean average precision (mAP; computed with copairs, the same metric used for real-cell perturbation distinctiveness). The ceiling is obtained by splitting the real cells of each class into two halves and running the identical retrieval — the real self-consistency mAP — and the generated-to-real ratio reports how much of the real phenotypic distinctiveness the synthesized cells recover. Because it is purely a neighborhood-retrieval statistic in embedding space, this metric is orthogonal to the set classifier and guards against classifier-specific artifacts. + +Across all three metrics a legitimate phenotype rises from control levels as α increases, peaks near α ≈ 1 (the knockout / class mean), and approaches the corresponding real-cell reference; extrapolation to |α| > 1 exaggerates the phenotype and, past the class mean, can overshoot and degrade recognition. + +Gene embedding and montage + +A gene-level phenotypic embedding was constructed by aggregating the single-cell Cell-DINO features of each perturbation, and a two-dimensional layout was computed with PHATE. To build the montage, each gene's generated cell (at a chosen α) was placed at that gene's coordinate in this embedding. Because many genes occupy dense regions, the embedding plane was tiled with a regular grid and a single representative gene was retained per grid cell — chosen by local density — so that adjacent cells do not overlap; coarser grids show fewer, larger cells and finer grids fill in more. Each tile was outlined by the color of its gene's Leiden cluster to convey the local phenotypic neighborhood, with individual complex members highlighted (e.g. ribo40S, ribo60S). The result is a single view in which each region of the phenotypic embedding is illustrated by a representative generated cell. + +The montage was assembled with the open-source latent-lens package (https://github.com/czi-ai/latent-lens; multiscale montages of image crops laid out by an embedding), which performs the grid-based, density-prioritized decimation and the multiscale tiling. diff --git a/src/ops_model/models/attention/diffex/figures/_setacc_common.py b/src/ops_model/models/attention/diffex/figures/_setacc_common.py new file mode 100644 index 0000000..ad739d3 --- /dev/null +++ b/src/ops_model/models/attention/diffex/figures/_setacc_common.py @@ -0,0 +1,157 @@ +"""Shared machinery for the fig-4 set-accuracy panels: the C/D figure-group registry, on-demand +cropping of a specific rank from the v5 set-accuracy rankings (KO class or NTC), marker-global +normalization, and the inverse-blue-mask composite. Used by figure4_setacc_panel.py (final panel) +and debug_setacc_top100.py (per-group montages to pick cells from). + +Cells are picked by their parquet `rank` (the badge shown in the debug montage). Ranks beyond the +cached top-30 are cropped live from phenotyping_v3.zarr via materialize_crops (same path as +viewer/_fluor_topcells).""" +import numpy as np +import pandas as pd +import zarr + +from ops_model.models.attention.diffex.classifier.config import slugify +from ops_model.models.attention.diffex.classifier.data import make_labels_df, materialize_crops +from ops_model.models.attention.diffex.directions.config import DirConfig +from ops_model.models.attention.diffex.viewer._fluor_topcells import _overlay_rgba +from ops_model.models.attention.diffex.viewer.build_pc_crops_masked import BASE, CROP_SIZE, _crop, _zarr_patch + +OUT = "/hpc/projects/icd.fast.ops/analysis/figure4_setacc_panel" +RANK_BASE = "/hpc/projects/icd.fast.ops/models/diffex/viewer_assets_v5/_rankings/fluor_shap" + +TIM23 = "TIM23 mitochondrial inner membrane pre-sequence translocase complex, TIM17A variant" +COPI = "COPI vesicle coat complex, COPG1-COPZ1 variant" + +# Published C/D figure groups. key = class in the ranking parquet (POLR1H stored under alias ZNRD1); +# ko_rank / ntc_rank = the rank badge picked from the debug montage (default 1 = top set-accuracy). +GENE_COLS = [ + dict(slug="Mitochondria_TOMM20", mc="Mitochondria_TOMM20", ch="CP1_mitochondria_TOMM20", + block="genes", key="TOMM20", top_label="TOMM20", marker_label="Mitochondria\n(TOMM20)", + ko_rank=1, ntc_rank=18), + dict(slug="nucleolus_GC_NPM3", mc="nucleolus-GC_NPM3", ch="GFP", + block="genes", key="ZNRD1", top_label="POLR1H", marker_label="Nucleoli\n(NPM3-GFP)", + ko_rank=1, ntc_rank=5), + dict(slug="5xUPRE", mc="5xUPRE", ch="GFP", + block="genes", key="HSPA5", top_label="HSPA5", marker_label="UPR\n(5xUPRE)", + ko_rank=100, ntc_rank=3), + dict(slug="ER_Golgi_COP_II_SEC23A", mc="ER/Golgi COP-II_SEC23A", ch="GFP", + block="genes", key="GBF1", top_label="GBF1", marker_label="ER-Golgi\n(GFP-SEC23A)", + ko_rank=7, ntc_rank=2), +] +COMPLEX_COLS = [ + dict(slug="mitochondria_ChromaLIVE_561_excitation", mc="mitochondria_ChromaLIVE 561 excitation", + ch="mCherry", block="complexes", key=TIM23, top_label="TIM23", + marker_label="Mitochondria\n(ChromaLIVE 561)", ko_rank=27, ntc_rank=1), + dict(slug="cell_proliferation_marker_MKI67", mc="cell proliferation marker_MKI67", ch="GFP", + block="complexes", key="DNA polymerase alpha:primase complex", top_label="DNA Pol α", + marker_label="Proliferation\n(GFP-MKI67)", ko_rank=2, ntc_rank=2), + dict(slug="actin_filament_FastAct_SPY555_Live_Cell_Dye", mc="actin filament_FastAct_SPY555 Live Cell Dye", + ch="mCherry", block="complexes", key="Chaperonin-containing T-complex", top_label="CCT", + marker_label="Actin\n(FastAct SPY555)", ko_rank=4, ntc_rank=13), + dict(slug="lysosome_LysoTracker_live_cell_dye", mc="lysosome_LysoTracker live-cell dye", ch="GFP", + block="complexes", key="mTORC1 complex", top_label="mTORC1", marker_label="Lysosome\n(LysoTracker)", + ko_rank=18, ntc_rank=2), # Rab-GGTase slot -> mTORC1 · Lysosome +] + + +def _rankdir(block): + return f"{RANK_BASE}/{'geneKO' if block == 'genes' else 'complex'}" + + +def _materialize(sel, mc, ch, cls): + """Crop the given (already row-selected) ranking rows. Returns (raw[N,1,H,W], recs realigned).""" + _zarr_patch() + recs = sel.rename(columns={"gene": "cls", "pma_attention": "score"}).copy() + recs["label"] = 0 + cfg = DirConfig(grain="geneKO", target=cls, device="cpu") + cfg.marker_channel = mc; cfg.channel = ch; cfg.num_workers = 8 + raw, _, exps = materialize_crops(make_labels_df(recs, cfg), cfg, cache_path=None) + recs = recs[recs["experiment"].isin(set(exps))].reset_index(drop=True) + n = min(len(raw), len(recs)) + return raw[:n], recs.iloc[:n].reset_index(drop=True) + + +def _class_df(mc, block, cls): + df = pd.read_parquet(f"{_rankdir(block)}/{slugify(mc)}.parquet") + df = df[df["gene"] == cls].sort_values("rank").reset_index(drop=True) + if df.empty: + raise ValueError(f"{cls!r} absent from {mc!r} [{block}]") + return df + + +def materialize_class(mc, ch, block, cls, top_n): + """Crop the top-`top_n` cells (by rank order/position) of one class. For the debug montages.""" + return _materialize(_class_df(mc, block, cls).head(top_n), mc, ch, cls) + + +def crop_pick_from_df(df, pick_rank, mc, ch, cls, win=40): + """Crop the cell with rank == pick_rank from a rank-sorted class df (the unique cell picked from + the montage), PLUS a top-`win` positional sample for a stable intensity window. Returns + (raw, recs, pos). Raises loudly if the picked rank is absent or its store dropped. Channel-agnostic + (ch=Phase2D for phase, marker channel for fluor) so both panels reuse it.""" + pick = df[df["rank"] == pick_rank] + if pick.empty: + raise ValueError(f"rank {pick_rank} not found for {cls!r} (max rank {int(df['rank'].max())})") + sel = pd.concat([df.head(win), pick]).drop_duplicates(["experiment", "well", "x_pheno", "y_pheno"]) + raw, recs = _materialize(sel, mc, ch, cls) + w = np.where(recs["rank"].values == pick_rank)[0] + if not len(w): + raise RuntimeError(f"rank {pick_rank} cell for {cls!r} dropped by materialize_crops (store missing) — pick another") + return raw, recs, int(w[0]) + + +def crop_pick(mc, ch, block, cls, pick_rank, win=40): + """Fluor: crop rank==pick_rank of one class from its per-marker parquet.""" + return crop_pick_from_df(_class_df(mc, block, cls), pick_rank, mc, ch, cls, win) + + +_SEGCACHE = {} + + +def seg_crop(exp, well, x, y, half): + ek = (exp, well) + if ek not in _SEGCACHE: + pos = f"{BASE}/{exp}/3-assembly/phenotyping_v3.zarr/{str(well)[0]}/{str(well)[1:]}/0" + try: + _SEGCACHE[ek] = zarr.open(f"{pos}/labels/cell_seg/0", mode="r") + except Exception: + _SEGCACHE[ek] = None + z = _SEGCACHE[ek] + if z is None: + return None + try: + return _crop(z, None, int(round(x)), int(round(y)), half) + except Exception: + return None + + +def composite(gray, seg, half): + """gray (0-255 float) + inverse blue seg mask → RGB uint8 (blue outside the center cell).""" + rgb = np.stack([gray] * 3, -1).astype(np.float32) + if seg is not None: + ov = _overlay_rgba(seg, half).astype(np.float32) + a = ov[..., 3:4] / 255.0 + rgb = rgb * (1 - a) + ov[..., :3] * a + return rgb.clip(0, 255).astype(np.uint8) + + +def tile_at(raw, recs, pos, lo, hi): + half = CROP_SIZE // 2 + r = recs.iloc[pos] + gray = np.clip((raw[pos, 0] - lo) / (hi - lo), 0, 1) * 255 + seg = seg_crop(r["experiment"], r["well"], r["x_pheno"], r["y_pheno"], half) + return composite(gray, seg, half), r + + +def column_tiles(col): + """For one figure-group column dict, crop the chosen KO-rank and NTC-rank cells (selected by + rank value) with a shared marker-global (KO+NTC) 1-99 pct window. Returns + (ko_rgb, ntc_rgb, ko_conf, ntc_conf).""" + ko_raw, ko_recs, ko_pos = crop_pick(col["mc"], col["ch"], col["block"], col["key"], col["ko_rank"]) + ntc_raw, ntc_recs, ntc_pos = crop_pick(col["mc"], col["ch"], "genes", "NTC", col["ntc_rank"]) # NTC always from the marker's geneKO ranking (complex parquet has no NTC) + lo, hi = np.percentile(np.concatenate([ko_raw.ravel(), ntc_raw.ravel()]), (1, 99)) + if hi - lo < 1e-6: + hi = lo + 1 + ko_im, ko_r = tile_at(ko_raw, ko_recs, ko_pos, lo, hi) + ntc_im, ntc_r = tile_at(ntc_raw, ntc_recs, ntc_pos, lo, hi) + return ko_im, ntc_im, float(ko_r["score"]), float(ntc_r["score"]) diff --git a/src/ops_model/models/attention/diffex/figures/_setacc_phase.py b/src/ops_model/models/attention/diffex/figures/_setacc_phase.py new file mode 100644 index 0000000..3a6b228 --- /dev/null +++ b/src/ops_model/models/attention/diffex/figures/_setacc_phase.py @@ -0,0 +1,60 @@ +"""Phase (label-free 2D Phase) set-accuracy groups for the panel-E-style figure. Phase uses single +global v5 rankings (not per-marker): geneKO keyed on `gene` (has NTC), complex keyed on +`predicted_class` (NO NTC → NTC pulled from the geneKO parquet). Channel = Phase2D. Reuses the +cropping/compositing from _setacc_common.""" +import numpy as np +import pandas as pd + +from _setacc_common import crop_pick_from_df, tile_at + +RANKS = "/hpc/projects/icd.fast.ops/models/diffex/viewer_assets_v5/_rankings" +PHASE_CH = "Phase2D" + +# panel-E groups (published: TIMM23/Arp2-3 top, TIPARP/Core Mediator bottom); ko_rank/ntc_rank picked +# from the debug_phase montages. +GENE_COLS_PHASE = [ + dict(block="genes", key="TIMM23", top_label="TIMM23\n(gene-level)", ko_rank=1, ntc_rank=1), + dict(block="genes", key="TIPARP", top_label="TIPARP\n(gene-level)", ko_rank=1, ntc_rank=1), +] +COMPLEX_COLS_PHASE = [ + dict(block="complexes", key="Actin-related protein 2/3 complex, ARPC1A-ACTR3B-ARPC5 variant", + top_label="Arp2/3\ncomplex", ko_rank=8, ntc_rank=41), + dict(block="complexes", key="Core mediator complex", top_label="Core Mediator\ncomplex", + ko_rank=95, ntc_rank=46), +] +COLS_PHASE = GENE_COLS_PHASE + COMPLEX_COLS_PHASE + + +def phase_df(block, cls): + if block == "genes": + df = pd.read_parquet(f"{RANKS}/pma_v5_phase_geneKO.parquet") + df = df[df["gene"] == cls] + else: + df = pd.read_parquet(f"{RANKS}/pma_v5_phase_complex.parquet") + df = df[df["predicted_class"] == cls].copy() + df["gene"] = cls # so _materialize's gene->cls rename labels it + df = df.sort_values("rank").reset_index(drop=True) + if df.empty: + raise ValueError(f"{cls!r} absent from phase {block}") + return df + + +def phase_ntc(): + df = pd.read_parquet(f"{RANKS}/pma_v5_phase_geneKO.parquet") + return df[df["gene"] == "NTC"].sort_values("rank").reset_index(drop=True) + + +def crop_pick_phase(block, cls, pick_rank, win=40): + df = phase_ntc() if cls == "NTC" else phase_df(block, cls) + return crop_pick_from_df(df, pick_rank, None, PHASE_CH, cls, win) + + +def column_tiles_phase(col): + ko_raw, ko_recs, ko_pos = crop_pick_phase(col["block"], col["key"], col["ko_rank"]) + ntc_raw, ntc_recs, ntc_pos = crop_pick_phase(col["block"], "NTC", col["ntc_rank"]) + lo, hi = np.percentile(np.concatenate([ko_raw.ravel(), ntc_raw.ravel()]), (1, 99)) + if hi - lo < 1e-6: + hi = lo + 1 + ko_im, ko_r = tile_at(ko_raw, ko_recs, ko_pos, lo, hi) + ntc_im, ntc_r = tile_at(ntc_raw, ntc_recs, ntc_pos, lo, hi) + return ko_im, ntc_im, float(ko_r["score"]), float(ntc_r["score"]) diff --git a/src/ops_model/models/attention/diffex/figures/auto_pick_and_plot.py b/src/ops_model/models/attention/diffex/figures/auto_pick_and_plot.py new file mode 100644 index 0000000..8323212 --- /dev/null +++ b/src/ops_model/models/attention/diffex/figures/auto_pick_and_plot.py @@ -0,0 +1,112 @@ +"""For each batch target: scan full_features.json, pick the interpretable feature whose generated α3 lands +closest to real KO (same sign, |real KO| >= 12%), then generate its violin (SLURM) + line-graph. Prints the +picked feature + gen-α3-vs-KO per target so the choice is auditable.""" +import json +import os +import re +import sys + +import numpy as np + +sys.path.insert(0, os.path.dirname(os.path.abspath(__file__))) +os.environ.setdefault("OPS_DIFFEX_ASSETS", "viewer_assets_v5") +from ops_model.models.attention.diffex.viewer.morpho_pipeline import MORPHO_TARGETS +from ops_model.models.attention.diffex.classifier.config import slugify + +VA = "/hpc/projects/icd.fast.ops/models/diffex/viewer_assets_v5/_morphometrics" +BAD = re.compile("moment|hu_|inertia|eigval|intensity|haralick|zernike|glcm|orientation|centroid|_timing") + +BATCH = ["KIF11_PHASE", "ATP6V1B2_PHASE", "HGS_PHASE", "RRM1_PHASE", "RRN3_PHASE", "SEC61A1_PHASE", + "SON_PHASE", "AP2M1_PHASE", "GOLGA2_PHASE", "SMC2_PHASE", "NPC_PHASE", "PROTEASOME_PHASE", + "HAUS_PHASE", "AP2M1_CLTA", + "EIF2S2_SG", "AURKB_CHROMATIN", "NOP56_FBL", "ATG9A_AUTOPHAGO", "ATP6V1B2_LAMP1"] + + +def _dir(v): + return f"{v['marker_dir']}/{v.get('grain', 'geneKO')}/{v['target']}" + + +NET_NOUN = {"skeleton_pixel_count": "Network length", "total_branch_length": "Network length", + "num_branches": "Branch count", "num_endpoints": "Endpoint count", "average_degree": "Network connectivity", + "num_skeleton_components": "Fragment count", "largest_connected_component_size": "Largest fragment", + "network_length_density": "Network density", "branching_density": "Branching density", + "num_nodes": "Node count", "euler_number": "Network topology"} +OBJ_NOUN = {"area": "size", "area_filled": "size", "extent": "compactness", "circularity": "roundness", + "eccentricity": "elongation", "aspect_ratio": "elongation", "axis_minor_length": "width", + "axis_major_length": "length", "equivalent_diameter_area": "diameter", "perimeter": "perimeter", + "solidity": "solidity"} + + +def _label(f): + """Intuitive y-axis label from a raw feature name (e.g. obj_extent_sum → 'Total compactness').""" + m = re.match(r"network_.+?_seg_(.+)$", f) + if m: + return NET_NOUN.get(m.group(1), m.group(1).replace("_", " ").capitalize()) + if f.startswith("obj_"): + b = f[4:] + mm = re.match(r"(.+)_(sum|mean|median|std|min|max|count)$", b) + prop, stat = (mm.group(1), mm.group(2)) if mm else (b, "") + if stat == "count": + return "Object count" + noun = OBJ_NOUN.get(prop, prop.replace("_", " ")) + if stat == "std": + return f"{noun.capitalize()} variability" + s = ({"sum": "Total ", "min": "Min ", "max": "Max "}.get(stat, "") + noun).strip() + return s[0].upper() + s[1:] + return f.replace("_", " ") + + +def pick(dir_): + """Best interpretable feature: gen α3 closest to real KO, same sign, |real KO| >= 12%.""" + d = json.load(open(f"{VA}/{dir_}/full_features.json")) + al = d["alphas"]; agg = d["agg"]; rr = d.get("real_ref", {}) + z = min(range(len(al)), key=lambda i: abs(al[i])) + i1 = min(range(len(al)), key=lambda i: abs(al[i] - 1)); i3 = min(range(len(al)), key=lambda i: abs(al[i] - 3)) + best = None + for f, ser in agg.items(): + if not f.startswith(("obj_", "network_")) or BAD.search(f): + continue + r = rr.get(f) + if not r or not r.get("ko") or r["ntc"][0] is None or ser[z] is None or ser[i1] is None: + continue + b = abs(ser[z]) or 1e-9; nb = abs(r["ntc"][0]) or 1e-9 + g1 = (ser[i1] - ser[z]) / b * 100; g3 = (ser[i3] - ser[z]) / b * 100; ko = (r["ko"][0] - r["ntc"][0]) / nb * 100 + # require BOTH α1 and α3 to track the real-KO direction (α1 not flat) — not just α3 + if abs(ko) < 12 or np.sign(g3) != np.sign(ko) or np.sign(g1) != np.sign(ko) or abs(g1) < 4: + continue + gap = abs(g3 - ko) + if best is None or gap < best[0]: + best = (gap, f, round(ko), round(g3)) + return best + + +def main(): + keys = sys.argv[1:] or BATCH + figs = [] + for k in keys: + v = MORPHO_TARGETS[k]; dir_ = _dir(v) + try: + b = pick(dir_) + except FileNotFoundError: + print(f" {k}: no full_features.json"); continue + if not b: + print(f" {k}: no interpretable feature tracks KO (skip)"); continue + gap, feat, ko, g3 = b + print(f" {k:18} feat={feat:52} KO {ko:+4d}% gen α3 {g3:+4d}%") + figs.append({"group": k, "dir": dir_, "feature": feat, "simple": _label(feat), + "out_stem": f"{k}_{slugify(feat)[:24]}", "label": f"{k} · {feat}"}) + # violins (SLURM array) + line-graphs (local) + from figure4_morpho_violin import submit + submit(figs) + from figure4_morpho_traversal import make_figure + for f in figs: + for c in range(2): + try: + make_figure(f["dir"], f["dir"], f["feature"], c, [0, 1, 3], f["label"], f["simple"], + f"{f['group']}/{f['out_stem']}_cell{c}") + except Exception as e: + print(f" line skip {f['group']} c{c}: {type(e).__name__}") + + +if __name__ == "__main__": + main() diff --git a/src/ops_model/models/attention/diffex/figures/cis_golgi_alternatives.py b/src/ops_model/models/attention/diffex/figures/cis_golgi_alternatives.py new file mode 100644 index 0000000..4e5f509 --- /dev/null +++ b/src/ops_model/models/attention/diffex/figures/cis_golgi_alternatives.py @@ -0,0 +1,26 @@ +"""Candidate strip for the panel-D Rab-slot (alternatives to COPI·cis-Golgi) — KO vs NTC at rank 1 +(most distinctive) for a few trafficking/organelle complex+marker pairs, to pick the most obvious.""" +from figure4_setacc_panel import make_panel + +CANDS = [ + dict(slug="cis_Golgi_mStayGold_CENPRaltORF", mc="cis-Golgi_mStayGold-CENPRaltORF", ch="GFP", + key="COPI vesicle coat complex, COPG1-COPZ1 variant", top_label="COPI", marker_label="cis-Golgi\n(CENPR)"), + dict(slug="trans_Golgi_VAMP3", mc="trans-Golgi_VAMP3", ch="GFP", + key="COPI vesicle coat complex, COPG1-COPZ1 variant", top_label="COPI", marker_label="trans-Golgi\n(VAMP3)"), + dict(slug="ER_Golgi_COPE", mc="ER/Golgi_COPE", ch="GFP", + key="COPI vesicle coat complex, COPG1-COPZ1 variant", top_label="COPI", marker_label="ER/Golgi\n(COPE)"), + dict(slug="late_endosome_RAB7A", mc="late endosome_RAB7A", ch="GFP", + key="COPI vesicle coat complex, COPG1-COPZ1 variant", top_label="COPI", marker_label="Endosome\n(RAB7A)"), + dict(slug="lysosome_LysoTracker_live_cell_dye", mc="lysosome_LysoTracker live-cell dye", ch="GFP", + key="mTORC1 complex", top_label="mTORC1", marker_label="Lysosome\n(LysoTracker)"), + dict(slug="ER_SEC61B", mc="ER_SEC61B", ch="mCherry", + key="SEC61 protein-conducting channel complex, SEC1A1 variant", top_label="SEC61", marker_label="ER\n(SEC61B)"), + dict(slug="Microtubules_Tubulin", mc="Microtubules_Tubulin", ch="CP2_microtubules_Tubulin", + key="HAUS complex", top_label="HAUS", marker_label="Microtubules\n(Tubulin)"), +] +for c in CANDS: + c.update(block="complexes", ko_rank=1, ntc_rank=1) + +if __name__ == "__main__": + make_panel(CANDS, "cis-Golgi alternatives — Rab-slot candidates (rank-1 set-accuracy)", + "rab_slot_alternatives", tile=1.5) diff --git a/src/ops_model/models/attention/diffex/figures/debug_setacc_top100.py b/src/ops_model/models/attention/diffex/figures/debug_setacc_top100.py new file mode 100644 index 0000000..46a44ae --- /dev/null +++ b/src/ops_model/models/attention/diffex/figures/debug_setacc_top100.py @@ -0,0 +1,88 @@ +"""Debug montages — top-N set-accuracy cells per figure group, rank-ordered with rank + cell-key +annotations, marker-global normalized, inverse blue mask. One montage per KO class and one per +(marker, block) NTC so specific KO/NTC cells can be picked into GENE_COLS/COMPLEX_COLS in +_setacc_common.py. Genes have up to ~1200 ranked cells; complexes only top-30. + +Run: python debug_setacc_top100.py # all groups (KO + NTC) + python debug_setacc_top100.py TOMM20 # only groups whose top_label matches (arg substring) +""" +import os +import sys + +import matplotlib + +matplotlib.use("Agg") +import matplotlib.pyplot as plt +import numpy as np + +from ops_model.models.attention.diffex.classifier.config import slugify +from _setacc_common import (COMPLEX_COLS, GENE_COLS, OUT, CROP_SIZE, materialize_class, seg_crop, composite) + +plt.rcParams["pdf.fonttype"] = 42 + + +def render_montage(raw, recs, title, out_stem): + n = len(recs) + lo, hi = np.percentile(raw, (1, 99)) + if hi - lo < 1e-6: + hi = lo + 1 + half = CROP_SIZE // 2 + ncols = 10 + nrows = int(np.ceil(n / ncols)) + fig, axes = plt.subplots(nrows, ncols, figsize=(ncols * 1.9, nrows * 2.05), facecolor="white") + axes = np.atleast_2d(axes) + for k in range(nrows * ncols): + ax = axes.flat[k] + ax.set_xticks([]); ax.set_yticks([]) + for s in ax.spines.values(): + s.set_visible(False) + if k >= n: + ax.axis("off"); continue + r = recs.iloc[k] + gray = np.clip((raw[k, 0] - lo) / (hi - lo), 0, 1) * 255 + seg = seg_crop(r["experiment"], r["well"], r["x_pheno"], r["y_pheno"], half) + ax.imshow(composite(gray, seg, half)) + ax.text(0.03, 0.97, f"#{int(r['rank'])}", transform=ax.transAxes, fontsize=13, fontweight="bold", + color="white", va="top", ha="left", + bbox=dict(boxstyle="round,pad=0.15", fc="#c1272d", ec="none", alpha=0.9)) + key = f"{r['experiment']}/{r['well']} x{int(round(r['x_pheno']))} y{int(round(r['y_pheno']))}" + ax.text(0.5, -0.03, f"{key}\nconf={float(r['score']):.3f}", transform=ax.transAxes, + fontsize=6.0, color="#222", va="top", ha="center") + fig.suptitle(title, fontsize=15, fontweight="bold", y=0.997) + fig.subplots_adjust(left=0.005, right=0.995, top=0.965, bottom=0.01, wspace=0.05, hspace=0.32) + os.makedirs(OUT, exist_ok=True) + out = f"{OUT}/{out_stem}.png" + fig.savefig(out, dpi=150, bbox_inches="tight", facecolor="white") + plt.close(fig) + print(f"saved {out} ({n} cells)") + + +def montage(mc, ch, block, cls, title, out_stem, top_n=100): + raw, recs = materialize_class(mc, ch, block, cls, top_n) + render_montage(raw, recs, title, out_stem) + + +def main(): + filt = sys.argv[1] if len(sys.argv) > 1 else None + cols = [c for c in GENE_COLS + COMPLEX_COLS if not filt or filt.lower() in c["top_label"].lower()] + ntc_done = set() + for c in cols: + try: + montage(c["mc"], c["ch"], c["block"], c["key"], + f"KO — {c['top_label']} ({c['mc']}) rank-ordered set-accuracy", + f"debug_KO_{c['slug']}_{slugify(c['key'])[:30]}") + except Exception as e: + print(f"skip KO {c['top_label']}: {e}") + nk = (c["slug"], c["block"]) + if nk not in ntc_done: + ntc_done.add(nk) + try: + montage(c["mc"], c["ch"], c["block"], "NTC", + f"NTC — {c['top_label']} marker ({c['mc']}) rank-ordered set-accuracy", + f"debug_NTC_{c['slug']}_{c['block']}") + except Exception as e: + print(f"skip NTC {c['slug']}: {e}") + + +if __name__ == "__main__": + main() diff --git a/src/ops_model/models/attention/diffex/figures/ebi_peripheral_droplets.py b/src/ops_model/models/attention/diffex/figures/ebi_peripheral_droplets.py new file mode 100644 index 0000000..25c7bcd --- /dev/null +++ b/src/ops_model/models/attention/diffex/figures/ebi_peripheral_droplets.py @@ -0,0 +1,204 @@ +"""ISOLATED PROBE — peripheral lipid-droplet counts for the EMC/BODIPY panel (panel A of +figure_multirank_ebi_grid / figure_ebi_morpho_violin). + +Deliberately standalone: the op_cp_features stores only hold per-CELL aggregates of the localization +features (distance_from_cell_edge mean/min/max/...), so a "droplets within X µm of the cell boundary" +count has to be measured per object. This script re-measures it for the panel's top-SHAP cells straight +from phenotyping_v3.zarr (gfp_seg droplets + cell_seg + nuclear_seg, same stitched coords the image panel +crops from) using organelle_profiler's own localization code, so `distance_from_cell_edge` here means +exactly what the store's column means. + +Nothing in the working figure scripts is modified — it only imports read-only helpers from them, writes +its own cache + its own figure, and can be deleted to go back to the simple droplet count. + +Definitions per cell (droplets = gfp_seg objects inside the cell mask). "Peripheral" = the droplet's +centroid is within d µm of the cell boundary, or at normalized_radial_position >= t (0 = nucleus, +1 = cell edge). Counts and areas both, since abundance and size move independently: + count / area_sum : all droplets (cross-check vs the store's op_gfp_count / op_gfp_area_sum) + peri_edge_um : # peripheral droplets frac_edge_um : / total count + peri_shell_ : # peripheral droplets frac_shell_ : / total count + periarea_edge_um (µm²) : peripheral droplet area fracarea_edge_um : / total area + periarea_shell_ (µm²) : peripheral droplet area fracarea_shell_ : / total area + +Cells: the top-1000 SHAP-ranked cells per class from the EBI multi_rank screen (rank order, no store +matching needed — the zarr is the source here). Cells whose mask is missing or clipped by the crop window +are skipped and counted, never silently dropped. + +Run: python ebi_peripheral_droplets.py # measure (cached) + plot + python ebi_peripheral_droplets.py --refresh # re-measure from the zarrs + python ebi_peripheral_droplets.py --limit 50 # quick smoke test (50 cells/class) +""" +import os +import sys + +import matplotlib + +matplotlib.use("Agg") +import matplotlib.pyplot as plt +import numpy as np +import pandas as pd +import zarr + +from figure_ebi_morpho_violin import draw_violin +from figure_multirank_ebi_grid import CACHE, OUT, ebi_rows, top_rows +from ops_model.models.attention.diffex.viewer.build_pc_crops_masked import BASE, _crop, _zarr_patch + +from organelle_profiler.feature_extraction.localization_features import compute_localization_features + +CH_NAME = "lipid droplet_BODIPY live cell dye" # multi_rank channel_name (panel A) +GENES = ["EMC1", "EMC2", "EMC3"] # same rows as the image panel +ORG_LABEL = "gfp_seg" # BODIPY droplets in phenotyping_v3.zarr +N_TOP = 1000 +PX_UM = 0.325 # NATIVE phenotype pixel size; the zarr's declared 0.65 is the known bug +HALF = 192 # crop half-window (px) — 125 µm box, comfortably larger than a HeLa cell +EDGE_UM = (1.0, 2.0, 3.0) # "within d µm of the cell boundary" +SHELL = (0.6, 0.8) # normalized nucleus→edge position threshold +POUT = f"{OUT}/peripheral" +plt.rcParams["pdf.fonttype"] = 42 +plt.rcParams["svg.fonttype"] = "none" +plt.rcParams["font.family"] = "sans-serif" +plt.rcParams["font.sans-serif"] = ["Arial", "Helvetica", "DejaVu Sans"] # Arial first: no Illustrator substitution + +FEATS = ([("count", "Lipid droplet count"), ("area_sum", "Total lipid droplet area (µm²)")] + + [(f"peri_edge_{d:g}um", f"Peripheral droplets\n(≤{d:g} µm from cell edge)") for d in EDGE_UM] + + [(f"frac_edge_{d:g}um", f"Peripheral droplet fraction\n(≤{d:g} µm from cell edge)") for d in EDGE_UM] + + [(f"periarea_edge_{d:g}um", f"Peripheral droplet area (µm²)\n(≤{d:g} µm from cell edge)") for d in EDGE_UM] + + [(f"fracarea_edge_{d:g}um", f"Peripheral droplet area fraction\n(≤{d:g} µm from cell edge)") for d in EDGE_UM] + + [(f"peri_shell_{t:g}", f"Peripheral droplets\n(radial position ≥ {t:g})") for t in SHELL] + + [(f"frac_shell_{t:g}", f"Peripheral droplet fraction\n(radial position ≥ {t:g})") for t in SHELL] + + [(f"periarea_shell_{t:g}", f"Peripheral droplet area (µm²)\n(radial position ≥ {t:g})") for t in SHELL] + + [(f"fracarea_shell_{t:g}", f"Peripheral droplet area fraction\n(radial position ≥ {t:g})") for t in SHELL]) + + +def _pos(exp, well): + """phenotyping_v3.zarr position group for one (experiment, well).""" + w = str(well) + return f"{BASE}/{exp}/3-assembly/phenotyping_v3.zarr/{w[0]}/{w[1:]}/0" + + +class Stores: + """cell_seg / organelle / nuclear_seg label arrays per (experiment, well), opened once.""" + + def __init__(self): + self.cache = {} + + def get(self, exp, well): + k = (exp, str(well)) + if k not in self.cache: + p = _pos(exp, well) + try: + self.cache[k] = tuple(zarr.open(f"{p}/labels/{n}/0", mode="r") + for n in ("cell_seg", ORG_LABEL, "nuclear_seg")) + except Exception as e: # noqa: BLE001 - report, keep going + print(f" [store] {exp}/{well}: {type(e).__name__}: {e}") + self.cache[k] = None + return self.cache[k] + + +def cell_measure(stores, r): + """Per-object localization for one cell → dict of peripheral counts/fractions, or a skip reason.""" + z = stores.get(r["experiment"], r["well"]) + if z is None: + return None, "no_store" + cz, oz, nz = z + x, y, sid = int(round(r["x_pheno"])), int(round(r["y_pheno"])), int(r["segmentation_id"]) + cell = _crop(cz, None, x, y, HALF) == sid + if not cell.any(): + return None, "cell_label_absent" + if cell[0, :].any() or cell[-1, :].any() or cell[:, 0].any() or cell[:, -1].any(): + return None, "cell_clipped" # edge distances would be wrong — never guess + org = np.where(cell, _crop(oz, None, x, y, HALF), 0) + if not org.any(): + return dict(count=0), None # a real zero-droplet cell + nuc = cell & (_crop(nz, None, x, y, HALF) > 0) # nuclear_seg has its own IDs -> intersect spatially + loc = compute_localization_features(org, cell, nuc if nuc.any() else None, spacing=(PX_UM, PX_UM)) + n = len(loc) + out = {"count": float(n)} + if n: + px = np.bincount(org.ravel()) # per-object pixel count -> µm² by label id + area = px[loc["label"].to_numpy(int)] * PX_UM ** 2 + atot = float(area.sum()) + out["area_sum"] = atot + + def _both(mask, ctag, atag): + c = float(np.sum(mask)) + out[f"peri_{ctag}"] = c + out[f"frac_{ctag}"] = c / n + a = float(area[mask].sum()) + out[f"periarea_{atag}"] = a + out[f"fracarea_{atag}"] = a / atot if atot > 0 else np.nan + + ed = loc["distance_from_cell_edge"].to_numpy(float) + for d in EDGE_UM: + _both(ed <= d, f"edge_{d:g}um", f"edge_{d:g}um") + if "normalized_radial_position" in loc and np.isfinite(loc["normalized_radial_position"]).any(): + rp = np.nan_to_num(loc["normalized_radial_position"].to_numpy(float), nan=-1.0) + for t in SHELL: + _both(rp >= t, f"shell_{t:g}", f"shell_{t:g}") + return out, None + + +def measure(limit=None): + """Measure every class's top-N cells. Returns the tidy per-cell table.""" + _zarr_patch() + screen = ebi_rows("fluor") + stores = Stores() + rows, skips = [], {} + for gene in GENES + ["NTC"]: + sel = top_rows(screen, gene, CH_NAME, limit or N_TOP) + for i, (_, r) in enumerate(sel.iterrows()): + m, why = cell_measure(stores, r) + if m is None: + skips[why] = skips.get(why, 0) + 1 + continue + rows.append(dict(gene=gene, rank=int(r["rank"]), experiment=r["experiment"], well=r["well"], + segmentation_id=int(r["segmentation_id"]), **m)) + if (i + 1) % 250 == 0: + print(f" [{gene}] {i + 1}/{len(sel)} cells", flush=True) + n = sum(1 for x in rows if x["gene"] == gene) + print(f" [{gene}] measured {n}/{len(sel)} cells", flush=True) + if skips: + print(f" skipped: {skips} (clipped cells are excluded, not guessed)", flush=True) + df = pd.DataFrame(rows) + long = df.melt(id_vars=["gene", "rank", "experiment", "well", "segmentation_id"], + var_name="feature", value_name="value").dropna(subset=["value"]) + long["unit"] = "count" # counts + fractions: unit-less for the label helper + return long + + +def table(refresh=False, limit=None): + p = f"{CACHE}/peripheral_bodipy_{limit or N_TOP}_f{len(FEATS)}.parquet" # feature count keys the cache + if os.path.exists(p) and not refresh: + return pd.read_parquet(p) + df = measure(limit) + os.makedirs(CACHE, exist_ok=True) + df.to_parquet(p) + print(f" cached {p} ({df['gene'].nunique()} classes)", flush=True) + return df + + +def main(refresh=False, limit=None): + tab = table(refresh, limit) + os.makedirs(POUT, exist_ok=True) + for feat, ylab in FEATS: + if not (tab["feature"] == feat).any(): + print(f"skip {feat}: not measured"); continue + fig, ax = plt.subplots(figsize=(1.35 * (len(GENES) + 1) + 1.6, 4.8), facecolor="white") + draw_violin(ax, tab, GENES, feat, ylab, title="EMC complex — lipid droplet (BODIPY)") + for ext in ("png", "svg"): + fig.savefig(f"{POUT}/violin_peripheral_{feat}.{ext}", dpi=220, bbox_inches="tight", facecolor="white") + plt.close(fig) + print(f"saved {POUT}/violin_peripheral_{feat}.png/.svg", flush=True) + for feat, _ in FEATS: # Δmedian table, to compare against the plain count + d = tab[tab["feature"] == feat] + if d.empty: + continue + med = lambda g: np.nanmedian(d.loc[d["gene"] == g, "value"]) + nmed = med("NTC") + print(f"{feat:20s} NTC {nmed:8.3g} | " + + ", ".join(f"{g} {(med(g) - nmed) / (abs(nmed) or 1e-9) * 100:+.0f}%" for g in GENES)) + + +if __name__ == "__main__": + lim = int(sys.argv[sys.argv.index("--limit") + 1]) if "--limit" in sys.argv else None + main(refresh="--refresh" in sys.argv, limit=lim) diff --git a/src/ops_model/models/attention/diffex/figures/figure4_morpho_traversal.py b/src/ops_model/models/attention/diffex/figures/figure4_morpho_traversal.py new file mode 100644 index 0000000..5c689a4 --- /dev/null +++ b/src/ops_model/models/attention/diffex/figures/figure4_morpho_traversal.py @@ -0,0 +1,299 @@ +"""Figure 4 (paper) — morpho traversal panel: image row + seg-overlay row + %-change plot. + +Reads a morpho target's full_features.json (per-α mean trajectory + real NTC→KO reference) +and its traversal frames / per-frame seg labels + per-object features (same assets the +morpho_demo.html viewer uses). Three stacked rows: + 1. grayscale traversal images (original real NTC, then generated cell at each shown α) + 2. the same panels with the org-seg mask overlaid, objects colored by the feature + (inferno heatmap, per-object value normalized over the cell trajectory + NTC) — + identical mapping to morpho_demo.html (overlayBase → per-object key, infernoRGB). + 3. generated %-change vs its own α=0 baseline (starts at α=0), real KO reference ±SEM. + +Usage: python figure4_morpho_traversal.py (edit the __main__ block for a different target). +""" +import json +import os + +import matplotlib + +matplotlib.use("Agg") +import matplotlib.cm as cm +import matplotlib.pyplot as plt +import numpy as np +from matplotlib.colors import Normalize +from PIL import Image + +KEYLABEL = {"area": "object area (px²)", "area_filled": "filled area (px²)", "mean_int": "object intensity", + "ecc": "eccentricity", "skel": "skeleton length", "circularity": "circularity", + "extent": "extent", "solidity": "solidity", "axis_minor_length": "minor axis (px)", + "axis_major_length": "major axis (px)"} + +plt.rcParams["pdf.fonttype"] = 42 +plt.rcParams["svg.fonttype"] = "none" +plt.rcParams["font.family"] = "sans-serif" +plt.rcParams["font.sans-serif"] = ["Arial", "Helvetica", "DejaVu Sans"] + +VA = f"/hpc/projects/icd.fast.ops/models/diffex/{os.environ.get('OPS_DIFFEX_ASSETS', 'viewer_assets')}" +OUT = "/hpc/projects/icd.fast.ops/analysis/figure4_traversals" + + +def _objkey(feature, avail): + """Map a full feature name → per-object key used for overlay coloring (morpho_demo.overlayBase).""" + if "intensity" in feature and "mean_int" in avail: + return "mean_int" + if "eccentric" in feature and "ecc" in avail: + return "ecc" + if ("branch" in feature or "skeleton" in feature) and "skel" in avail: + return "skel" + if "area" in feature and "area" in avail: + return "area" + return None + + +NTC_RGB = (0.18, 0.72, 0.70) # flat single color for the NTC reference seg overlay (teal) + + +def _uniform_rgba(labels, rgb, op=0.75): + """Flat single-color RGBA over every labeled object (bg transparent).""" + rgba = np.zeros((*labels.shape, 4)) + m = labels > 0 + rgba[m, :3] = rgb + rgba[m, 3] = op + return rgba + + +def _overlay_rgba(labels, feats, key, lo, hi, op=0.75): + """Per-object inferno RGBA over the label mask; bg / valueless objects transparent.""" + maxid = int(labels.max()) + tlut = np.full(maxid + 1, np.nan, float) + for sid, props in feats.items(): + i = int(sid) + v = props.get(key) + if i <= maxid and v is not None: + tlut[i] = (v - lo) / (hi - lo + 1e-9) + t = tlut[labels] + rgba = cm.inferno(np.clip(np.nan_to_num(t), 0, 1)) + rgba[..., 3] = np.where(np.isfinite(t) & (labels > 0), op, 0.0) + return rgba + + +def image_panels(md, morpho_dir, traversal_dir, feature, cell, alphas_show): + """Traversal image panels (original NTC anchor + generated α frames): (title, grayscale, label mask, + per-object feats) + the per-object overlay key and its color limits. Shared by the line & violin + figures so their image panels are IDENTICAL.""" + ff = json.load(open(f"{md}/full_features.json")) + alphas = ff["alphas"]; n = len(alphas); z = n >> 1 + idxs = [min(range(n), key=lambda i: abs(alphas[i] - a)) for a in alphas_show] + f0 = json.load(open(f"{md}/cell{cell}/a{idxs[0]:02d}_feats.json")) + avail = set(next(iter(f0.values())).keys()) if f0 else set() + okey = _objkey(feature, avail) + if okey is None and "area" in avail: + okey = "area" + vals = [] + if okey: + for i in range(n): + try: + fj = json.load(open(f"{md}/cell{cell}/a{i:02d}_feats.json")) + vals += [p[okey] for p in fj.values() if p.get(okey) is not None] + except FileNotFoundError: + pass + lo, hi = (tuple(np.percentile(vals, CLIP)) if vals else (0.0, 1.0)) + panels = [] + modality = traversal_dir.split("/")[0] + anchor_img = f"{VA}/{modality}/_anchors/NTC/cell{cell}/real.webp" + if os.path.exists(anchor_img): + albp = f"{md}/cell{cell}/a{z:02d}_labels.png" # α=0 seg outlines the same reconstructed cell + lab = np.asarray(Image.open(albp)) if os.path.exists(albp) else None + panels.append(("original NTC", np.asarray(Image.open(anchor_img).convert("L")), lab, {})) + else: + print(f" no anchor real.webp for {modality} cell{cell} — skipping original panel") + for i in idxs: + gray = np.asarray(Image.open(f"{VA}/{traversal_dir}/cell{cell}/frame_{i:02d}.webp").convert("L")) + lab = np.asarray(Image.open(f"{md}/cell{cell}/a{i:02d}_labels.png")) + fj = json.load(open(f"{md}/cell{cell}/a{i:02d}_feats.json")) + panels.append((f"α={alphas[i]:+.0f}", gray, lab, fj)) + return panels, okey, lo, hi + + +def render_images(fig, spec, panels, okey, lo, hi, op=0.75, title_fs=22, cbar=True, hspace=0.14): + """Render the 2-row image block (grayscale over feature-colored seg overlay + inferno colorbar) into + the gridspec `spec`. Returns the overlay axes.""" + from matplotlib.gridspec import GridSpecFromSubplotSpec + nc = len(panels) + g = GridSpecFromSubplotSpec(2, nc, subplot_spec=spec, hspace=hspace, wspace=0.04) + oaxes = [] + for j, (t, gray, lab, fj) in enumerate(panels): + axi = fig.add_subplot(g[0, j]) + axi.imshow(gray, cmap="gray"); axi.set_xticks([]); axi.set_yticks([]) + axi.set_title(t, fontsize=title_fs) + for s in axi.spines.values(): + s.set_visible(False) + axo = fig.add_subplot(g[1, j]) + axo.imshow(gray, cmap="gray") + if lab is None: + pass + elif t == "original NTC": + axo.imshow(_uniform_rgba(lab, NTC_RGB, op)) + elif okey: + axo.imshow(_overlay_rgba(lab, fj, okey, lo, hi, op)) + axo.set_xticks([]); axo.set_yticks([]) + for s in axo.spines.values(): + s.set_visible(False) + oaxes.append(axo) + if cbar and okey: + p = oaxes[-1].get_position() + cax = fig.add_axes([p.x1 + 0.012, p.y0, 0.013, p.height]) + cb = fig.colorbar(cm.ScalarMappable(Normalize(lo, hi), cmap="inferno"), cax=cax) + cb.set_label(KEYLABEL.get(okey, okey), fontsize=18) + cb.set_ticks([lo, hi]); cb.ax.set_yticklabels([f"{lo:.0f}", f"{hi:.0f}"]) + cb.ax.tick_params(labelsize=16); cb.outline.set_visible(False) + return oaxes + + +def make_figure(morpho_dir, traversal_dir, feature, cell, alphas_show, label, simple, out_stem, op=0.75): + md = f"{VA}/_morphometrics/{morpho_dir}" + ff = json.load(open(f"{md}/full_features.json")) + alphas = ff["alphas"] + n = len(alphas) + z = n >> 1 # α=0 baseline index + raw = ff["agg"][feature] + gb = abs(raw[z]) or 1e-9 + gen = [(v - raw[z]) / gb * 100 for v in raw] + asem = (ff.get("agg_sem") or {}).get(feature) or [0.0] * n + sem = [e / gb * 100 for e in asem] + + rr = (ff.get("real_ref") or {}).get(feature) or {} + def _koref(kk, nn): # real KO % change vs its NTC baseline (top-1k or all-cells) + n, k = rr.get(nn), rr.get(kk) + if n and k and n[0] is not None and k[0] is not None: + nmp = abs(n[0]) or 1e-9 + return (k[0] - n[0]) / nmp * 100, (k[1] or 0) / nmp * 100 + return None, None + koV, koS = _koref("ko", "ntc") # top-1k set-accuracy real KO + koVall, koSall = _koref("ko_all", "ntc_all") # all-cells real KO + + panels, okey, lo, hi = image_panels(md, morpho_dir, traversal_dir, feature, cell, alphas_show) + nc = len(panels) + + fig = plt.figure(figsize=(nc * 2.3, 8.6), facecolor="white") + gs = fig.add_gridspec(2, 1, height_ratios=[2.0, 1.55], hspace=0.14) + render_images(fig, gs[0], panels, okey, lo, hi, op) + + ax = fig.add_subplot(gs[1]) # %-change plot (α ≥ 0) + ax.set_facecolor("white") + a0, g0, s0 = alphas[z:], gen[z:], sem[z:] + if koV is not None: + ax.axhspan(koV - koS, koV + koS, color="#5ad17a", alpha=0.18, lw=0) + ax.axhline(koV, color="#2e8b57", ls="--", lw=3, label=f"real KO top-1k ({koV:+.0f}%)") + if koVall is not None: + ax.axhline(koVall, color="#8a8a8a", ls=":", lw=2.5, label=f"real KO all ({koVall:+.0f}%)") + ax.axhline(0, color="#999", lw=1.5) + ax.errorbar(a0, g0, yerr=s0, fmt="-o", color="#1f77b4", lw=4, ms=9, + capsize=5, elinewidth=2, ecolor="#1f77b4", label="generated (mean ± SEM)") + ax.set_xlabel("α (traversal strength)", fontsize=28) + ax.set_ylabel(f"{simple}\n(% change vs NTC)", fontsize=28) + ax.set_xlim(-0.2, max(a0) + 0.2) + ax.set_xticks([a for a in a0 if a == int(a)]) + ax.tick_params(labelsize=24, width=2.5, length=9) + ax.legend(frameon=False, fontsize=22, loc="best") + for s in ("top", "right"): + ax.spines[s].set_visible(False) + for s in ("left", "bottom"): + ax.spines[s].set_linewidth(2.5) + + os.makedirs(os.path.dirname(f"{OUT}/{out_stem}"), exist_ok=True) + for ext in ("png", "svg"): + fig.savefig(f"{OUT}/{out_stem}.{ext}", dpi=220, bbox_inches="tight", facecolor="white") + plt.close(fig) + print(f"saved {OUT}/{out_stem} (overlay key={okey}, clim=[{lo:.1f},{hi:.1f}], gen {g0[-1]:+.0f}% @α+5, " + f"real KO {None if koV is None else round(koV)}%)") + + +FIGURES = [ + dict(group="KIF23_nucleolive", dir="nucleus_NucleoLIVE_Live_Cell_dye/geneKO/KIF23", + feature="obj_area_filled_sum", cell=0, simple="Nuclear size", + label="KIF23 (NucleoLIVE) · object area filled sum", out_stem="KIF23_nucleolive_obj_area_filled_sum"), + dict(group="TOMM20_phase", dir="phase/geneKO/TOMM20", + feature="obj_area_max", cell=0, simple="Mitochondrial size", + label="TOMM20 (phase) · object area max", out_stem="TOMM20_phase_obj_area_max"), + dict(group="KIF23_phase", dir="phase/geneKO/KIF23", + feature="obj_area_filled_sum", cell=0, simple="Nuclear size", + label="KIF23 (phase) · nuclear area (nucleus seg)", out_stem="KIF23_phase_nuclear_size"), + dict(group="HSPA5_phase", dir="phase/geneKO/HSPA5", + feature="obj_area_sum", cell=0, simple="Vacuole area", + label="HSPA5 (phase) · dark vacuole area", out_stem="HSPA5_phase_vacuole_area"), + dict(group="RAB7A_phase", dir="phase/geneKO/RAB7A", + feature="obj_area_sum", cell=0, simple="Vesicle area", + label="RAB7A (phase) · light vesicle area", out_stem="RAB7A_phase_vesicle_area"), + dict(group="SAMM50_phase_frag", dir="phase/geneKO/SAMM50", + feature="network_phase2d_seg_largest_connected_component_size", cell=0, simple="Mitochondrial fragmentation", + label="SAMM50 (phase) · largest connected component size", out_stem="SAMM50_phase_largest_cc"), + dict(group="SAMM50_phase_area", dir="phase/geneKO/SAMM50", + feature="obj_area_filled_sum", cell=0, simple="Mitochondria area", + label="SAMM50 (phase) · object area filled sum", out_stem="SAMM50_phase_area_filled"), + dict(group="RAB7A_bodipy", dir="lipid_droplet_BODIPY_live_cell_dye/geneKO/RAB7A", + feature="obj_area_sum", cell=0, simple="Lipid droplet area", + label="RAB7A (BODIPY lipid droplets) · object area sum", out_stem="RAB7A_bodipy_ld_area"), + dict(group="LAMTOR2_lyso", dir="lysosome_LysoTracker_live_cell_dye/geneKO/LAMTOR2", + feature="obj_area_filled_sum", cell=0, simple="Lysosome area", + label="LAMTOR2 (LysoTracker) · object area filled sum", out_stem="LAMTOR2_lyso_area"), + dict(group="LAMTOR2_lyso_count", dir="lysosome_LysoTracker_live_cell_dye/geneKO/LAMTOR2", + feature="obj_area_count", cell=0, simple="Lysosome count", + label="LAMTOR2 (LysoTracker) · object count", out_stem="LAMTOR2_lyso_count"), + dict(group="HSPA5_phase_count", dir="phase/geneKO/HSPA5", + feature="obj_area_count", cell=0, simple="Vacuole count", + label="HSPA5 (phase) · dark vacuole count", out_stem="HSPA5_phase_vacuole_count"), + dict(group="SNRNP200_phase", dir="phase/geneKO/SNRNP200", + feature="obj_area_min", cell=0, simple="Dark vesicle size", + label="SNRNP200 (phase) · dark vesicle min area", out_stem="SNRNP200_phase_ves_size"), + dict(group="SNRNP200_phase_count", dir="phase/geneKO/SNRNP200", + feature="obj_area_count", cell=0, simple="Dark vesicle count", + label="SNRNP200 (phase) · dark vesicle count", out_stem="SNRNP200_phase_ves_count"), + dict(group="CCT_npm3", dir="nucleolus_GC_NPM3/complex/Chaperonin_containing_T_complex", + feature="obj_circularity_sum", cell=0, simple="Nucleolar circularity", + label="CCT complex (NPM3 nucleoli) · object circularity sum", out_stem="CCT_npm3_obj_circularity_sum"), + dict(group="TOMM20_chromalive561", dir="mitochondria_ChromaLIVE_561_excitation/geneKO/TOMM20", + feature="network_mitochondria_chromalive_561_excitation_tubular_seg_num_endpoints", cell=0, + simple="Network endpoints", + label="TOMM20 (ChromaLIVE561) · network num endpoints", out_stem="TOMM20_chromalive561_num_endpoints"), + dict(group="GBF1_sec23a", dir="ER_Golgi_COP_II_SEC23A/geneKO/GBF1", + feature="obj_area_sum", cell=0, simple="Golgi size", + label="GBF1 (ER/Golgi COP-II SEC23A) · object area sum", out_stem="GBF1_sec23a_obj_area_sum"), + dict(group="POLR1B_npm3", dir="nucleolus_GC_NPM3/geneKO/POLR1B", + feature="obj_extent_sum", cell=0, simple="Nucleolar extent", + label="POLR1B (NPM3 nucleoli) · object extent sum", out_stem="POLR1B_npm3_obj_extent_sum"), + dict(group="TIM23_chromalive561", dir="mitochondria_ChromaLIVE_561_excitation/complex/TIM23_mitochondrial_inner_membrane_pre_sequence_translocase_complex__TIM17A_variant", + feature="network_mitochondria_chromalive_561_excitation_tubular_seg_num_branches", cell=0, + simple="Branch count", + label="TIM23 complex (ChromaLIVE561) · network num branches", out_stem="TIM23_chromalive561_num_branches"), + dict(group="CAPZB_fastact", dir="actin_filament_FastAct_SPY555_Live_Cell_Dye/geneKO/CAPZB", + feature="obj_axis_minor_length_mean", cell=0, simple="Filament width", + label="CAPZB (FastAct actin) · axis minor length mean", out_stem="CAPZB_fastact_axis_minor_length_mean"), + dict(group="CAPZB_phalloidin", dir="F_actin_Phalloidin/geneKO/CAPZB", + feature="obj_axis_minor_length_mean", cell=9, simple="Filament width", + label="CAPZB (Phalloidin F-actin) · axis minor length mean", out_stem="CAPZB_phalloidin_axis_minor_length_mean"), + dict(group="AP2M1_phase", dir="phase/geneKO/AP2M1", + feature="obj_circularity_mean", cell=9, simple="Roundness", + label="AP2M1 (phase) · object circularity mean", out_stem="AP2M1_phase_roundness"), + dict(group="TIM23_chromalive_degree", dir="mitochondria_ChromaLIVE_561_excitation/complex/TIM23_mitochondrial_inner_membrane_pre_sequence_translocase_complex__TIM17A_variant", + feature="network_mitochondria_chromalive_561_excitation_tubular_seg_average_degree", cell=9, + simple="Network degree", + label="TIM23 complex (ChromaLIVE561) · network average degree", out_stem="TIM23_chromalive561_network_degree"), + dict(group="PSMB6_proteasome", dir="proteasome_PSMB7/geneKO/PSMB6", + feature="obj_area_sum", cell=9, simple="Proteasome area", + label="PSMB6 (proteasome PSMB7) · total proteasome area", out_stem="PSMB6_proteasome_area"), +] + +CELLS = [0, 1, 2, 3, 4, 5] # render a PNG per cell so the best one can be picked +ALPHAS_SHOW = [0, 1, 5] # generated α panels (original NTC is prepended automatically) +CLIP = (2, 98) # overlay color-scale percentiles ((0,100)=raw min/max, demo-style) + +if __name__ == "__main__": + for f in FIGURES: + for c in CELLS: + try: + make_figure(morpho_dir=f["dir"], traversal_dir=f["dir"], feature=f["feature"], cell=c, + alphas_show=ALPHAS_SHOW, label=f["label"], simple=f.get("simple", f["label"]), + out_stem=f"{f['group']}/{f['out_stem']}_cell{c}") + except (FileNotFoundError, IndexError) as e: + print(f"skip {f['out_stem']}_cell{c}: {e}") diff --git a/src/ops_model/models/attention/diffex/figures/figure4_morpho_violin.py b/src/ops_model/models/attention/diffex/figures/figure4_morpho_violin.py new file mode 100644 index 0000000..97dd87b --- /dev/null +++ b/src/ops_model/models/attention/diffex/figures/figure4_morpho_violin.py @@ -0,0 +1,147 @@ +"""Figure 4 morpho — VIOLIN variant (separate subdir). The traversal image panel is IDENTICAL to the +line-graph figure (reuses image_panels + render_images: original-NTC anchor + generated α frames, both +the grayscale and feature-colored seg-overlay rows). To the RIGHT of the images (not below) a VIOLIN plot +shows the per-cell distribution (variance over the 100 cells) of the feature for real KO, generated α=0, +and generated α=1 — all as %-change vs NTC (real KO vs real-NTC mean; generated vs its own α=0 baseline). + +Run: OPS_DIFFEX_ASSETS=viewer_assets_v5 python figure4_morpho_violin.py +""" +import json +import os + +import matplotlib + +matplotlib.use("Agg") +import matplotlib.pyplot as plt +import numpy as np +import pandas as pd + +from figure4_morpho_traversal import FIGURES, VA, image_panels, render_images +from ops_model.models.attention.diffex.viewer.morpho_pipeline import MORPHO_TARGETS, real_percell +from ops_model.models.attention.diffex.classifier.config import slugify + +plt.rcParams["pdf.fonttype"] = 42 +plt.rcParams["svg.fonttype"] = "none" +plt.rcParams["font.family"] = "sans-serif" +plt.rcParams["font.sans-serif"] = ["Arial", "Helvetica", "DejaVu Sans"] + +OUT = "/hpc/projects/icd.fast.ops/analysis/figure4_traversals_violin" +COLORS = {"real": "#999999", "KO": "#2e8b57", "α=0": "#c6dbef", "α=1": "#6baed6", "α=3": "#08519c"} +ALPHAS_SHOW = [0, 1, 3] # image panel columns (α=3 = exaggeration, not α=5) +CELLS = list(range(30)) # render 30 example cells to pick from +VIEW_PCT = (1, 98) # y-axis view window (percentile of pooled data) — clips the + # VIEW only; KDE + median stay on the full data (never moves) + + +def _store_cfg(marker_dir, target, grain): + for v in MORPHO_TARGETS.values(): + if v["marker_dir"] == marker_dir and slugify(v["target"]) == slugify(target) and v.get("grain", "geneKO") == grain: + return v["store_marker"], v.get("store_channel") + return None, None + + +def _pct(vals, base): + b = abs(base) or 1e-9 + return (np.asarray(vals, float) - base) / b * 100.0 + + +def make_violin(dir_, feature, simple, out_stem): + marker_dir, grain, target = dir_.split("/", 2) + md = f"{VA}/_morphometrics/{dir_}" + ff = json.load(open(f"{md}/full_features.json")) + alphas = ff["alphas"] + ai = lambda a: min(range(len(alphas)), key=lambda i: abs(alphas[i] - a)) + z, i1, i3 = ai(0), ai(1.0), ai(3.0) + + df = pd.read_parquet(f"{md}/full_features.parquet") + gen = {a: df.loc[df["alpha_idx"] == idx, feature].dropna().values for a, idx in ((0, z), (1, i1), (3, i3))} + base = float(np.mean(gen[0])) + + store_marker, store_channel = _store_cfg(marker_dir, target, grain) + rr = (real_percell(marker_dir, target, grain, store_marker, [feature], store_channel=store_channel) or {}).get(feature) if store_marker else None + if rr is not None and len(rr["ntc"]) and len(rr["ko"]): + nbase = float(np.nanmean(rr["ntc"])) # real %-change baseline = real-NTC mean + real_ntc = _pct(rr["ntc"][~np.isnan(rr["ntc"])], nbase) + real_ko = _pct(rr["ko"][~np.isnan(rr["ko"])], nbase) + else: + real_ntc = real_ko = np.array([]); print(f" {out_stem}: no real per-cell for {feature}") + + data = [real_ntc, real_ko, _pct(gen[0], base), _pct(gen[1], base), _pct(gen[3], base)] + labels = ["real", "KO", "α=0", "α=1", "α=3"] + os.makedirs(os.path.dirname(f"{OUT}/{out_stem}"), exist_ok=True) + + # ---- (1) image panel — identical to the line-graph figure — one per example cell ---- + for cell in CELLS: + try: + panels, okey, lo, hi = image_panels(md, dir_, dir_, feature, cell, ALPHAS_SHOW) + except (FileNotFoundError, IndexError): + continue + nc = len(panels) + figi = plt.figure(figsize=(nc * 2.3, 5.0), facecolor="white") + render_images(figi, figi.add_gridspec(1, 1)[0], panels, okey, lo, hi, title_fs=22, hspace=0.05) + for ext in ("png", "svg"): + figi.savefig(f"{OUT}/{out_stem}_images_cell{cell}.{ext}", dpi=220, bbox_inches="tight", facecolor="white") + plt.close(figi) + + # ---- (2) violin — separate file; full (unclipped) per-cell distributions, median line from the data ---- + keep = [i for i, d in enumerate(data) if len(d)] + figv = plt.figure(figsize=(4.8, 5.6), facecolor="white") + ax = figv.add_subplot(111); ax.set_facecolor("white") + parts = ax.violinplot([data[i] for i in keep], positions=keep, showmeans=False, showextrema=False, showmedians=False, widths=0.82) + for pc, i in zip(parts["bodies"], keep): + pc.set_facecolor(COLORS[labels[i]]); pc.set_alpha(0.6); pc.set_edgecolor(COLORS[labels[i]]); pc.set_linewidth(1.5) + for i in keep: + ax.hlines(np.mean(data[i]), i - 0.34, i + 0.34, color="#222", lw=3, zorder=5) # mean line (full data) + ax.axhline(0, color="#999", lw=2) + pooled = np.concatenate([data[i] for i in keep]) # view-clip: robust ylim, data/median untouched + ylo, yhi = np.percentile(pooled, VIEW_PCT); pad = 0.05 * (yhi - ylo + 1e-9) + ax.set_ylim(ylo - pad, yhi + pad) + ax.set_xticks(range(len(labels))); ax.set_xticklabels(labels, fontsize=30) + import textwrap + ylab = "\n".join(textwrap.wrap(simple, 14)) if len(simple) > 14 else simple # wrap long labels so they don't clip the canvas + ax.set_ylabel(ylab, fontsize=26 if len(simple) > 18 else 30) + ax.tick_params(axis="y", labelsize=26, width=2.5, length=9) + ax.tick_params(axis="x", length=0) + for s in ("top", "right"): + ax.spines[s].set_visible(False) + for s in ("left", "bottom"): + ax.spines[s].set_linewidth(2.5) + for ext in ("png", "svg"): + figv.savefig(f"{OUT}/{out_stem}_violin.{ext}", dpi=220, bbox_inches="tight", facecolor="white") + plt.close(figv) + km = float(np.median(real_ko)) if len(real_ko) else None + print(f"saved {OUT}/{out_stem}_violin + {len(CELLS)} cell images (real KO med {None if km is None else round(km)}%, " + f"gen α1 med {round(float(np.median(_pct(gen[1], base))))}%, α3 med {round(float(np.median(_pct(gen[3], base))))}%)") + + +def _job(fig): + import os + os.environ.setdefault("OPS_DIFFEX_ASSETS", "viewer_assets_v5") # worker reads v5 assets + make_violin(fig["dir"], fig["feature"], fig.get("simple", fig["label"]), f"{fig['group']}/{fig['out_stem']}") + + +def submit(figs=None): + """One SLURM job per target (parallel) — each reads the op_cp store + renders violin + 10 cell images.""" + import pathlib + from ops_utils.hpc.slurm_batch_utils import submit_parallel_jobs + figdir = str(pathlib.Path(__file__).resolve().parent) # workers must import the loose figures scripts + os.environ["PYTHONPATH"] = figdir + os.pathsep + os.environ.get("PYTHONPATH", "") + os.environ.setdefault("OPS_DIFFEX_ASSETS", "viewer_assets_v5") + figs = figs or FIGURES + jobs = [{"name": f"violin_{f['group'][:18]}", "func": _job, "kwargs": {"fig": f}} for f in figs] + print(f"[violin] {len(jobs)} target jobs in parallel") + submit_parallel_jobs(jobs, experiment="diffex_violin", + slurm_params={"slurm_partition": "cpu", "cpus_per_task": 8, "mem_gb": 64, "timeout_min": 120}, + log_dir="diffex_violin", wait_for_completion=False) + + +if __name__ == "__main__": + import sys + if len(sys.argv) > 1 and sys.argv[1] == "local": # serial (debug) + for f in FIGURES: + try: + make_violin(f["dir"], f["feature"], f.get("simple", f["label"]), f"{f['group']}/{f['out_stem']}") + except Exception as e: + print(f"skip {f['out_stem']}: {type(e).__name__}: {e}") + else: + submit() diff --git a/src/ops_model/models/attention/diffex/figures/figure4_setacc_panel.py b/src/ops_model/models/attention/diffex/figures/figure4_setacc_panel.py new file mode 100644 index 0000000..d773a85 --- /dev/null +++ b/src/ops_model/models/attention/diffex/figures/figure4_setacc_panel.py @@ -0,0 +1,74 @@ +"""Figure 4 C/D-style panel — top *set-accuracy* real cells (fluorescence), KO vs NTC grid. + +Mimics the published fig-4 C/D layout (gene-KO panel + protein-complex panel, KO row over NTC row, +marker label under each column) but the cells shown are hand-picked from the v5 SetTransformer +set-accuracy rankings (ko_rank / ntc_rank per column in _setacc_common.GENE_COLS/COMPLEX_COLS; +pick them from the debug_setacc_top100.py montages). Real cells, cropped on demand, marker-global +normalized so KO vs NTC brightness is comparable, with the inverse blue seg mask (C/D look). +Vector output (SVG + PNG). + +Run: python figure4_setacc_panel.py +""" +import os + +import matplotlib + +matplotlib.use("Agg") +import matplotlib.pyplot as plt + +from _setacc_common import COMPLEX_COLS, GENE_COLS, OUT, column_tiles + +plt.rcParams["pdf.fonttype"] = 42 +plt.rcParams["svg.fonttype"] = "none" + + +def make_panel(columns, panel_title, out_stem, tile=1.6, tiles_fn=column_tiles, bottom_caption=None, + title_in=0.55): + cols = [] + for c in columns: + try: + ko_im, ntc_im, kc, nc = tiles_fn(c) + cols.append((c, ko_im, ntc_im, kc, nc)) + except Exception as e: + print(f"skip col {c.get('top_label')} ({c['slug']}): {e}") + if not cols: + print(f"no columns for {out_stem}") + return + n = len(cols) + + left, right = 0.05, 0.997 + bot_in = 0.40 if bottom_caption else 0.10 # inches reserved for the marker/caption row + W = n * tile / (right - left) + H = 2 * tile + title_in + bot_in # square cells (tile x tile) → images tile tight + fig = plt.figure(figsize=(W, H), facecolor="white") + gs = fig.add_gridspec(2, n, hspace=0.02, wspace=0.02, left=left, right=right, + top=1 - title_in / H, bottom=bot_in / H) + for j, (c, ko_im, ntc_im, kc, nc) in enumerate(cols): + for i, im in enumerate((ntc_im, ko_im)): # NTC on top, KO on bottom + ax = fig.add_subplot(gs[i, j]) + ax.imshow(im) + ax.set_xticks([]); ax.set_yticks([]) + for s in ax.spines.values(): + s.set_edgecolor("#888"); s.set_linewidth(0.5) + if i == 0: + ax.set_title(c.get("marker_label") or c["top_label"], fontsize=11, fontweight="bold", pad=4) # marker on top + elif c.get("marker_label"): + ax.set_xlabel(c["top_label"], fontsize=11, fontweight="bold") # KO/gene name on bottom + if j == 0: + ax.set_ylabel("NTC" if i == 0 else "KO", fontsize=11, fontweight="bold", rotation=0, + labelpad=14, va="center") + fig.suptitle(panel_title, fontsize=13, fontweight="bold", x=left, ha="left", va="top", y=0.995) + if bottom_caption: + fig.text(0.5, 0.4 * bot_in / H, bottom_caption, fontsize=11, ha="center", style="italic") + os.makedirs(OUT, exist_ok=True) + for ext in ("png", "svg"): + fig.savefig(f"{OUT}/{out_stem}.{ext}", dpi=220, bbox_inches="tight", facecolor="white") + plt.close(fig) + print(f"saved {OUT}/{out_stem}.png/.svg (" + + ", ".join(f"{c['top_label']}: KO#{c['ko_rank']}={kc:.2f}/NTC#{c['ntc_rank']}={nc:.2f}" + for c, _, _, kc, nc in cols) + ")") + + +if __name__ == "__main__": + make_panel(GENE_COLS, "Gene KO top-predictive cells (fluorescence)", "panelC_geneKO_setacc", title_in=0.66) + make_panel(COMPLEX_COLS, "Protein complex top-predictive cells (fluorescence)", "panelD_complex_setacc", title_in=0.66) diff --git a/src/ops_model/models/attention/diffex/figures/figure4_setacc_panel_fluorB.py b/src/ops_model/models/attention/diffex/figures/figure4_setacc_panel_fluorB.py new file mode 100644 index 0000000..6830206 --- /dev/null +++ b/src/ops_model/models/attention/diffex/figures/figure4_setacc_panel_fluorB.py @@ -0,0 +1,24 @@ +"""Fluorescence 'top-predictive cells' panel — variant B (gene-KO, multibag SHAP): TOMM20, POLR1H, +CFL1 (actin FastAct), mTOR (LysoTracker). NTC top / KO bottom, per-column KO+NTC intensity window. + +Run: python figure4_setacc_panel_fluorB.py""" +from figure4_setacc_panel import make_panel +from _setacc_common import column_tiles + +COLS = [ + dict(slug="Mitochondria_TOMM20", mc="Mitochondria_TOMM20", ch="CP1_mitochondria_TOMM20", + block="genes", key="TOMM20", top_label="TOMM20", marker_label="Mitochondria\n(TOMM20)", + ko_rank=1, ntc_rank=18), + dict(slug="nucleolus_GC_NPM3", mc="nucleolus-GC_NPM3", ch="GFP", + block="genes", key="ZNRD1", top_label="POLR1H", marker_label="Nucleoli\n(NPM3-GFP)", + ko_rank=1, ntc_rank=5), + dict(slug="actin_filament_FastAct_SPY555_Live_Cell_Dye", mc="actin filament_FastAct_SPY555 Live Cell Dye", + ch="mCherry", block="genes", key="CFL1", top_label="CFL1", marker_label="Actin\n(FastAct SPY555)", + ko_rank=4, ntc_rank=20), + dict(slug="lysosome_LysoTracker_live_cell_dye", mc="lysosome_LysoTracker live-cell dye", ch="GFP", + block="genes", key="MTOR", top_label="mTOR", marker_label="Lysosome\n(LysoTracker)", + ko_rank=9, ntc_rank=2), +] + +if __name__ == "__main__": + make_panel(COLS, "Top-predictive cells (fluorescence)", "panelG_fluor_variantB", tiles_fn=column_tiles, title_in=0.66) diff --git a/src/ops_model/models/attention/diffex/figures/figure4_setacc_panel_newpheno.py b/src/ops_model/models/attention/diffex/figures/figure4_setacc_panel_newpheno.py new file mode 100644 index 0000000..163c3db --- /dev/null +++ b/src/ops_model/models/attention/diffex/figures/figure4_setacc_panel_newpheno.py @@ -0,0 +1,70 @@ +"""Phase 'top-predictive cells' panel for the NEW-phenotype gene-KOs, MULTIBAG SHAP ranking +(pma_shap_phase_geneKO). NTC on top (a DISTINCT NTC cell per column, ranks 1-5), KO on bottom +(hand-picked from the phase_multibag montages). Real phase cells cropped from phenotyping_v3.zarr. + +Run (SLURM): python figure4_setacc_panel_newpheno.py --submit +""" +import sys + +import numpy as np +import pandas as pd + +from _setacc_common import crop_pick_from_df, tile_at +from figure4_setacc_panel import make_panel + +RANK = "/hpc/projects/icd.fast.ops/models/diffex/viewer_assets_v5/_rankings/pma_shap_phase_geneKO.parquet" +PHASE_CH = "Phase2D" + +COLS = [ # KO rank = montage pick; NTC rank distinct per column (1-5) + dict(slug="KIF23", key="KIF23", top_label="KIF23\n(multi-nucleation)", ko_rank=20, ntc_rank=1), + dict(slug="CAPZB", key="CAPZB", top_label="CAPZB\n(stretched)", ko_rank=9, ntc_rank=2), + dict(slug="SNRPD1", key="SNRPD1", top_label="SNRPD1\n(dark vacuoles)", ko_rank=84, ntc_rank=3), + dict(slug="SAMM50", key="SAMM50", top_label="SAMM50\n(globular mito)", ko_rank=1, ntc_rank=4), + dict(slug="RAB7A", key="RAB7A", top_label="RAB7A\n(enlarged vesicles)", ko_rank=4, ntc_rank=5), +] + + +def _df(cls): + d = pd.read_parquet(RANK, filters=[("gene", "==", cls)]) + if "rank_type" in d.columns: + d = d[d["rank_type"] == "top"] + return d.sort_values("rank").reset_index(drop=True) + + +def column_tiles_shap(col): + ko_raw, ko_recs, ko_pos = crop_pick_from_df(_df(col["key"]), col["ko_rank"], None, PHASE_CH, col["key"]) + ntc_raw, ntc_recs, ntc_pos = crop_pick_from_df(_df("NTC"), col["ntc_rank"], None, PHASE_CH, "NTC") + lo, hi = np.percentile(np.concatenate([ko_raw.ravel(), ntc_raw.ravel()]), (1, 99)) + if hi - lo < 1e-6: + hi = lo + 1 + ko_im, ko_r = tile_at(ko_raw, ko_recs, ko_pos, lo, hi) + ntc_im, ntc_r = tile_at(ntc_raw, ntc_recs, ntc_pos, lo, hi) + return ko_im, ntc_im, float(ko_r["score"]), float(ntc_r["score"]) + + +def build(): + make_panel(COLS, "Top-predictive cells (phase)", "panelF_phase_newpheno", + tiles_fn=column_tiles_shap, bottom_caption="Label-free 2D Phase", + title_in=0.66) + + +def _job(): + import os + os.environ.setdefault("OPS_DIFFEX_ASSETS", "viewer_assets_v5") + build() + + +def submit(): + import os + import pathlib + from ops_utils.hpc.slurm_batch_utils import submit_parallel_jobs + figdir = str(pathlib.Path(__file__).resolve().parent) + os.environ["PYTHONPATH"] = figdir + os.pathsep + os.environ.get("PYTHONPATH", "") + os.environ.setdefault("OPS_DIFFEX_ASSETS", "viewer_assets_v5") + submit_parallel_jobs([{"name": "panelF_phase", "func": _job, "kwargs": {}}], experiment="diffex_panel", + slurm_params={"slurm_partition": "cpu", "cpus_per_task": 8, "mem_gb": 64, "timeout_min": 60}, + log_dir="diffex_panel", wait_for_completion=False) + + +if __name__ == "__main__": + submit() if (len(sys.argv) > 1 and sys.argv[1] == "--submit") else build() diff --git a/src/ops_model/models/attention/diffex/figures/figure4_setacc_panel_phase.py b/src/ops_model/models/attention/diffex/figures/figure4_setacc_panel_phase.py new file mode 100644 index 0000000..a8e4f28 --- /dev/null +++ b/src/ops_model/models/attention/diffex/figures/figure4_setacc_panel_phase.py @@ -0,0 +1,10 @@ +"""Panel-E-style figure — top set-accuracy cells in label-free 2D phase, KO vs NTC. +Groups: TIMM23 & TIPARP (gene-level), Arp2/3 & Core Mediator (complex). Picks set in +_setacc_phase.COLS_PHASE. Vector output (SVG + PNG).""" +from figure4_setacc_panel import make_panel +from _setacc_phase import COLS_PHASE, column_tiles_phase + +if __name__ == "__main__": + make_panel(COLS_PHASE, "Top-predictive cells (phase)", "panelE_phase_setacc", + tiles_fn=column_tiles_phase, bottom_caption="Label-free 2D Phase", + title_in=0.66) # room for 2-line column titles + suptitle, snug diff --git a/src/ops_model/models/attention/diffex/figures/figure_ebi_morpho_violin.py b/src/ops_model/models/attention/diffex/figures/figure_ebi_morpho_violin.py new file mode 100644 index 0000000..3bbed48 --- /dev/null +++ b/src/ops_model/models/attention/diffex/figures/figure_ebi_morpho_violin.py @@ -0,0 +1,398 @@ +"""Real-cell morphometrics for the EBI-complex predictive-cell panels (figure_multirank_ebi_grid) — NTC vs +the SAME 3 gene KOs, on the phenotype-relevant organelle feature. + +Nothing is re-measured: the per-cell organelle features already exist in the op_cp_features stores +(op_cp_features_.h5ad, one row per real cell, op__ columns), so this just +gathers them for the panel's genes and plots the distributions: + + A EMC · BODIPY lipid droplets -> peripheral droplet count (outer 20%, per-object from + ebi_peripheral_droplets.py) + op_gfp_count / area + B EMC · phase LIGHT vesicles -> op_phase2d_vesicular_* (count, area) + C Dynein-1· ER/Golgi COPE -> op_gfp_normalized_radial_position_mean / distance_from_nucleus + / area_sum (localization, then size) + D U7 snRNP· phase DARK vesicles -> op_phase2d_vesicular_dark_* (count, area) + +Cells per group: the top-1000 SHAP-RANKED cells of that gene (the multi_rank screen's own rank order, +walked until 1000 are matched into the store by experiment/well/segmentation — same matching as +morpho_pipeline.real_percell), with NTC restricted to the experiments its KO cells come from so batch is +matched. The cells displayed in the image panel are ranks 1-3, marked as dots. + +Run: python figure_ebi_morpho_violin.py # extract (cached) + plot + python figure_ebi_morpho_violin.py --refresh # re-read the stores +""" +import glob +import os +import re +import sys + +import h5py +import matplotlib + +matplotlib.use("Agg") +import matplotlib.pyplot as plt +import numpy as np +import pandas as pd + +from figure_multirank_ebi_grid import (BGX, CACHE, COMBINED_ORDER, FOOT, GAP, OUT, SUP, T, TITLE, block_h, + block_w, build_blocks, draw_block, ebi_rows, top_rows, windows) + +# paper-v2 stores first; the v2 dir is fluor-only, so the phase store still comes from the v1 dir (loud). +OPCP_DIRS = ["/hpc/projects/icd.fast.ops/analysis/op_cp_features_paper_v2", + "/hpc/projects/icd.fast.ops/analysis/op_cp_features"] +VOUT = f"{OUT}/morpho" +plt.rcParams["pdf.fonttype"] = 42 +plt.rcParams["svg.fonttype"] = "none" +plt.rcParams["font.family"] = "sans-serif" +plt.rcParams["font.sans-serif"] = ["Arial", "Helvetica", "DejaVu Sans"] # Arial first: no Illustrator substitution + +N_TOP = 1000 # cells per group = top-N SHAP-ranked cells matched into the store +BGY_C = 0.2 # vertical gap between composite rows +CAPTION_C = ("Images: top-SHAP cells, most predictive first (left to right). Violins: organelle morphometrics of the top-1000 SHAP-ranked real " + "cells per class (n≈1000 each) from the op_cp_features store — except A, where 'peripheral' = droplets in " + "the outer 20% of the nucleus→edge axis, measured per droplet; bar = median, % = Δmedian vs NTC.") +VIEW_PCT = (2, 96) # y-view clip — tails cropped from the VIEW only, medians/KDE stay on the full data +C_NTC, C_KO = "#8a8a8a", "#c1272d" +UNIT_LABEL = {"um^2": "µm²", "um": "µm", "1/um^2": "1/µm²", "um/um^2": "µm/µm²"} # var.unit -> axis text + +# per panel-block (keyed on the grid module's stem): store, organelle prefix, candidate features +SPECS = { + # primary is a PER-OBJECT measure (peripheral = outer 20% of the nucleus→edge axis), so it comes from + # ebi_peripheral_droplets.py's cache, not the store's per-cell aggregates — see _merge_peripheral. + "fluor_EMC_BODIPY": dict(store="lipid_droplet_bodipy_live_cell_dye", org="gfp", primary="peri_shell_0.8", + peripheral=["peri_shell_0.8", "periarea_shell_0.8", "frac_shell_0.8"], feats=[ + ("peri_shell_0.8", "Peripheral droplet count\n(outer 20%)"), + ("periarea_shell_0.8", "Peripheral droplet area\n(outer 20%)"), + ("frac_shell_0.8", "Peripheral droplet fraction\n(outer 20%)"), + ("count", "Lipid droplet count"), ("area_sum", "Total lipid droplet area"), + ("area_mean", "Mean droplet area")]), + "phase_EMC": dict(store="phase", org="phase2d_vesicular", primary="count", feats=[ + ("count", "Light vesicle count"), ("area_mean", "Mean vesicle area"), + ("area_sum", "Total vesicle area")]), + # dynein KO expands + disperses the ER/Golgi: total COPE area is the strongest readout (+53-66%), + # ~5x the two localization metrics (distance-from-nucleus / radial position), which are kept as variants + "fluor_Dynein1_COPE": dict(store="er_golgi_cope", org="gfp", primary="area_sum", feats=[ + ("area_sum", "Total COPE area"), + ("distance_from_nucleus_mean", "COPE distance from nucleus"), + ("normalized_radial_position_mean", "COPE radial position\n(0 = nucleus, 1 = cell edge)")]), + "phase_U7_snRNP": dict(store="phase", org="phase2d_vesicular_dark", primary="count", feats=[ + ("count", "Dark vesicle count"), ("area_mean", "Mean dark vesicle area"), + ("area_sum", "Total dark vesicle area")]), +} + + +def store_path(store): + for i, d in enumerate(OPCP_DIRS): + p = f"{d}/op_cp_features_{store}.h5ad" + if os.path.exists(p): + if i: + print(f" [store] {store}: absent from op_cp_features_paper_v2 — falling back to {d}") + return p + raise FileNotFoundError(f"no op_cp_features store for {store!r} in {OPCP_DIRS}") + + +def _obs(h, key, idx=None): + """One obs column as an array (categorical / nullable-integer / plain dataset), optionally row-subset.""" + o = h[f"obs/{key}"] + take = (lambda d: d[:]) if idx is None else (lambda d: d[idx]) + if isinstance(o, h5py.Group): + if "categories" in o: # categorical + cats = np.array([c.decode() if isinstance(c, bytes) else str(c) for c in o["categories"][:]]) + return cats[take(o["codes"])] + return take(o["values"]) # nullable-integer + v = take(o) + return np.array([x.decode() if isinstance(x, bytes) else x for x in v]) if v.dtype.kind in "SO" else v + + +def _nw(w): + """Store well form: 'A2' -> 'A/2/0' (matches morpho_pipeline.real_percell).""" + w = str(w).strip() + if w.count("/") == 2: + return w + m = re.match(r"^([A-Za-z]+)(\d+)$", w) + return f"{m.group(1)}/{m.group(2)}/0" if m else w + + +def _cat(h, key): + """(categories, codes) of an obs categorical — codes stay integer, so no 56M-string materialization.""" + o = h[f"obs/{key}"] + cats = np.array([c.decode() if isinstance(c, bytes) else str(c) for c in o["categories"][:]]) + return cats, o["codes"][:] + + +def _var_units(h, var, cols): + """var.unit of each column ('um^2' / 'um' / 'count' / ...) — the store's own physical units, so axis + labels never hardcode one. Fails loud if a spatial feature was never unit-corrected.""" + def col(key): + o = h[f"var/{key}"] + if isinstance(o, h5py.Group): + cats = np.array([c.decode() if isinstance(c, bytes) else str(c) for c in o["categories"][:]]) + return cats[o["codes"][:]] + v = o[:] + return np.array([x.decode() if isinstance(x, bytes) else x for x in v]) if v.dtype.kind in "SO" else v + unit, fixed = col("unit"), col("op_units_corrected") + out = [] + for c in cols: + i = var.index(c) + u = str(unit[i]) + if u in UNIT_LABEL and not bool(fixed[i]): + raise ValueError(f"{c}: unit {u} but op_units_corrected=False (0.65 vs 0.325 µm/px bug)") + out.append(u) + return out + + +def _numkey(seg, ecode, wcode): + """(segmentation, experiment, well) packed into one int64 so cells can be matched vectorized.""" + return np.asarray(seg, np.int64) * 10_000 + np.asarray(ecode, np.int64) * 100 + np.asarray(wcode, np.int64) + + +def _screen_numkeys(df, e2c, w2c): + """Screen rows (already rank-ordered) → int64 store keys, dropping rows whose exp/well isn't in the store.""" + e = df["experiment"].astype(str).map(e2c) + w = df["well"].astype(str).map(lambda x: w2c.get(_nw(x))) + ok = e.notna() & w.notna() + return _numkey(df.loc[ok, "segmentation_id"].astype("int64"), e[ok], w[ok]) + + +def _match_ranked(store_keys, rows, screen_keys, n): + """Walk screen_keys in rank order, keep the store rows they hit, stop at n. Returns (rows, n_hit).""" + o = np.argsort(store_keys, kind="stable") + sk = store_keys[o] + pos = np.clip(np.searchsorted(sk, screen_keys), 0, max(len(sk) - 1, 0)) + hit = (sk[pos] == screen_keys) if len(sk) else np.zeros(len(screen_keys), bool) + matched = rows[o[pos[hit]]] + _, first = np.unique(matched, return_index=True) # a store cell can be hit twice; keep rank order + matched = matched[np.sort(first)] + return matched[:n], int(hit.sum()) + + +def _merge_peripheral(blk, tab): + """Splice in the per-object peripheral features measured by ebi_peripheral_droplets.py (the store only + holds per-cell aggregates, so 'droplets in the outer 20%' cannot come from it). Fails loud if that + isolated probe hasn't been run — never silently plots a store-only subset.""" + want = SPECS[blk["b"]["stem"]].get("peripheral") + if not want: + return tab + cand = glob.glob(f"{CACHE}/peripheral_bodipy_{N_TOP}_f*.parquet") + if not cand: + raise FileNotFoundError(f"no peripheral cache in {CACHE} — run: python ebi_peripheral_droplets.py") + p = max(cand, key=lambda f: int(f.rsplit("_f", 1)[1].split(".")[0])) + d = pd.read_parquet(p) + d = d[d["feature"].isin(want)].copy() + missing = set(want) - set(d["feature"]) + if missing: + raise KeyError(f"{p}: missing peripheral features {missing}") + d["shown"] = False + per = d.groupby(["gene", "feature"]).size() + print(f" [{blk['b']['stem']}] + {len(want)} peripheral features from {os.path.basename(p)} " + f"({per.min()}-{per.max()} cells/class)", flush=True) + return pd.concat([tab, d[["gene", "feature", "value", "shown", "unit"]]], ignore_index=True) + + +def extract(blk, n_top=N_TOP): + """Tidy per-cell table (gene, feature, value, shown) for one block: the top-n_top SHAP-ranked cells of + each gene KO + of NTC (NTC restricted to its KO cells' experiments), from the op_cp_features store.""" + b, genes = blk["b"], blk["genes"] + spec = SPECS[b["stem"]] + sfeats = [ft for ft in spec["feats"] if ft[0] not in set(spec.get("peripheral") or [])] # store-backed only + cols = [f"op_{spec['org']}_{f}" for f, _ in sfeats] + screen = ebi_rows(b["modality"]) + with h5py.File(store_path(spec["store"]), "r") as h: + if isinstance(h["X"], h5py.Group): + raise TypeError(f"{spec['store']}: sparse X not supported by this reader") + var = [v.decode() for v in h["var/_index"][:]] + missing = [c for c in cols if c not in var] + if missing: + raise KeyError(f"{spec['store']}: missing feature columns {missing}") + cidx = [var.index(c) for c in cols] + units = dict(zip([f for f, _ in sfeats], _var_units(h, var, cols))) + gcats, gcodes = _cat(h, "gene_name") + ecats, ecodes = _cat(h, "experiment") + wcats, wcodes = _cat(h, "well") + if len(ecats) > 99 or len(wcats) > 99: + raise ValueError(f"{spec['store']}: {len(ecats)} exps / {len(wcats)} wells — key packing overflows") + e2c = {e: i for i, e in enumerate(ecats)} + w2c = {_nw(w): i for i, w in enumerate(wcats)} + allkeys = _numkey(_obs(h, "segmentation"), ecodes, wcodes) + gcode = {g: int(np.where(gcats == g)[0][0]) for g in genes + [""]} + + def panel_keys(recs): # the cells actually drawn in the image panel + e = recs["experiment"].astype(str).map(e2c) + w = recs["well"].astype(str).map(lambda x: w2c.get(_nw(x))) + ok = e.notna() & w.notna() + return set(_numkey(recs.loc[ok, "segmentation"].astype("int64"), e[ok], w[ok]).tolist()) + + pick, shown = {}, {} + for i, g in enumerate(genes): + gi = np.where(gcodes == gcode[g])[0] + sk = _screen_numkeys(top_rows(screen, g, b["ch_name"], 10 * n_top), e2c, w2c) + pick[g], nhit = _match_ranked(allkeys[gi], gi, sk, n_top) + pk = panel_keys(blk["ko"][i][1]) + shown[g] = {r for r in pick[g] if allkeys[r] in pk} + if len(pick[g]) < n_top: + print(f" [{b['stem']}] {g}: only {len(pick[g])} of the top-{n_top} ranked cells are in the store " + f"({nhit} matched overall)") + if len(shown[g]) < len(pk): + print(f" [{b['stem']}] {g}: {len(shown[g])}/{len(pk)} panel cells in the store") + koe = set(ecodes[np.concatenate([pick[g] for g in genes])].tolist()) + ni = np.where((gcodes == gcode[""]) & np.isin(ecodes, list(koe)))[0] # NTC, KO experiments only + nsk = _screen_numkeys(top_rows(screen, "NTC", b["ch_name"], 20 * n_top), e2c, w2c) + pick["NTC"], nhit = _match_ranked(allkeys[ni], ni, nsk, n_top) + npk = panel_keys(blk["ntc"][1]) + shown["NTC"] = {r for r in pick["NTC"] if allkeys[r] in npk} + print(f" [{b['stem']}] NTC {len(pick['NTC'])} ranked cells from {len(koe)} KO experiments", flush=True) + + rows = np.unique(np.concatenate([pick[k] for k in ["NTC"] + genes])) # h5py needs strictly increasing + X = h["X"][rows, :][:, cidx].astype(float) + r2p = {r: i for i, r in enumerate(rows)} + recs = [(grp, f, float(X[r2p[r]][j]), r in shown[grp], units[f]) + for grp in ["NTC"] + genes for r in pick[grp] for j, (f, _) in enumerate(sfeats)] + print(f" [{b['stem']}] {len(rows)} cells x {len(cols)} features units {units}", flush=True) + return pd.DataFrame(recs, columns=["gene", "feature", "value", "shown", "unit"]) + + +def table(blocks, refresh=False): + """Cached per-cell tables for every block: {stem: DataFrame}.""" + out = {} + for blk in blocks: + p = f"{CACHE}/morpho_{blk['b']['stem']}.parquet" + if os.path.exists(p) and not refresh: + out[blk["b"]["stem"]] = pd.read_parquet(p) + continue + df = extract(blk) + df.to_parquet(p) + out[blk["b"]["stem"]] = df + return out + + +def ylabel(tab, feat, text): + """'Mean vesicle area' + the store's unit for that feature → 'Mean vesicle area (µm²)'.""" + u = str(tab.loc[tab["feature"] == feat, "unit"].iloc[0]) + return f"{text} ({UNIT_LABEL[u]})" if u in UNIT_LABEL else text + + +def draw_violin(ax, df, genes, feat, ylab, title=None, fs=15, show_n=True, xrot=0): + """NTC + one violin per gene KO, medians as bars, %Δ median vs NTC (annotated at the top of the axes).""" + groups = ["NTC"] + genes + data = [df.loc[(df["gene"] == g) & (df["feature"] == feat), "value"].dropna().values for g in groups] + parts = ax.violinplot(data, positions=range(len(groups)), widths=0.82, + showmeans=False, showextrema=False, showmedians=False) + for i, pc in enumerate(parts["bodies"]): + c = C_NTC if i == 0 else C_KO + pc.set_facecolor(c); pc.set_edgecolor(c); pc.set_alpha(0.55); pc.set_linewidth(1.5) + nmed = float(np.median(data[0])) if len(data[0]) else np.nan + for i, d in enumerate(data): + if not len(d): + continue + ax.hlines(np.median(d), i - 0.34, i + 0.34, color="#222", lw=2.5, zorder=5) + if i: + pct = (np.median(d) - nmed) / (abs(nmed) or 1e-9) * 100 + ax.text(i, 0.995 if i % 2 else 0.885, f"{pct:+.0f}%", transform=ax.get_xaxis_transform(), + ha="center", va="top", fontsize=fs - 4, fontweight="bold", color=C_KO) # staggered: narrow axes + pooled = np.concatenate([d for d in data if len(d)]) + ylo, yhi = np.percentile(pooled, VIEW_PCT) + pad = 0.08 * (yhi - ylo + 1e-9) + ax.set_ylim(ylo - pad, yhi + pad * 3.4) # headroom for the staggered %Δ rows + ax.yaxis.set_major_locator(plt.MaxNLocator(5)) + ax.set_xlim(-0.62, len(groups) - 0.38) + ax.set_xticks(range(len(groups))) + ax.set_xticklabels([f"{g}\nn={len(d):,}" if show_n else g for g, d in zip(groups, data)], + fontsize=fs, rotation=xrot, ha="right" if xrot else "center") + for t, c in zip(ax.get_xticklabels(), [C_NTC] + [C_KO] * len(genes)): + t.set_color(c) + t.set_fontweight("bold") + ax.set_ylabel(ylab, fontsize=fs + 1) + ax.tick_params(axis="y", labelsize=fs - 1, width=2, length=7) + ax.tick_params(axis="x", length=0) + if title: + ax.set_title(title, fontsize=fs + 3, fontweight="bold", pad=8) + for s in ("top", "right"): + ax.spines[s].set_visible(False) + for s in ("left", "bottom"): + ax.spines[s].set_linewidth(2) + + +def save(fig, stem, outdir=VOUT): + os.makedirs(outdir, exist_ok=True) + for ext in ("png", "svg"): + fig.savefig(f"{outdir}/{stem}.{ext}", dpi=220, bbox_inches="tight", facecolor="white") + plt.close(fig) + print(f"saved {outdir}/{stem}.png/.svg", flush=True) + + +# composite geometry (inches): violin axes width, gap after the image block, y-label gutter, and the +# x-label band — which is taken INSIDE the image block's height (violin is shorter, top-aligned) so the +# rotated gene labels don't push the next panel row down. +VW, VGAP, VLAB, VXLAB, VBOT = 2.45, 0.45, 1.18, 1.0, 0.12 +FS_V = 22 + + +def composite(blocks, tabs, ncol=2, x0=0.45): + """Image block + its violin side by side, 2 blocks per row, in the panel order of the image figure.""" + win = windows(blocks) + order = [blocks[i] for i in COMBINED_ORDER] + rows = [order[i:i + ncol] for i in range(0, len(order), ncol)] + tot_w = lambda blk: block_w(blk["genes"]) + VGAP + VLAB + VW + colw = [max(tot_w(r[c]) for r in rows if c < len(r)) for c in range(ncol)] + rowh = [max(block_h(len(blk["genes"])) for blk in r) + VBOT for r in rows] + W = x0 + sum(colw) + (ncol - 1) * BGX + 0.12 + H = SUP + sum(rowh) + BGY_C * (len(rows) - 1) + FOOT + fig = plt.figure(figsize=(W, H), facecolor="white") + y = SUP + for r, row in enumerate(rows): + for c, blk in enumerate(row): + b, genes = blk["b"], blk["genes"] + spec = SPECS[b["stem"]] + x = x0 + sum(colw[:c]) + c * BGX + draw_block(fig, blk, x, y, *win[b["ch_name"]], W, H, letter="ABCDEFGH"[r * ncol + c]) + th = len(genes) * T + (len(genes) - 1) * GAP - VXLAB # room for the rotated x labels + ax = fig.add_axes([(x + block_w(genes) + VGAP + VLAB) / W, 1 - (y + TITLE + th) / H, VW / W, th / H]) + draw_violin(ax, tabs[b["stem"]], genes, spec["primary"], + ylabel(tabs[b["stem"]], spec["primary"], dict(spec["feats"])[spec["primary"]]), + fs=FS_V, show_n=False, xrot=45) + y += rowh[r] + BGY_C + fig.text(0.5, 1 - 0.3 / H, "Top Predictive cells per protein complex", fontsize=26, fontweight="bold", + ha="center", va="center") + fig.text(0.5, (FOOT - 0.3) / H, CAPTION_C, fontsize=13, style="italic", color="#333", ha="center", va="center") + save(fig, "ebi_composite_grid_violin", OUT) + + +def main(refresh=False): + blocks = build_blocks() + tabs = table(blocks, refresh) + tabs = {blk["b"]["stem"]: _merge_peripheral(blk, tabs[blk["b"]["stem"]]) for blk in blocks} + by_stem = {blk["b"]["stem"]: blk for blk in blocks} + + for blk in blocks: # every candidate feature, to pick from + stem, spec = blk["b"]["stem"], SPECS[blk["b"]["stem"]] + for feat, ylab in spec["feats"]: + fig, ax = plt.subplots(figsize=(1.35 * (len(blk["genes"]) + 1) + 1.6, 4.6), facecolor="white") + draw_violin(ax, tabs[stem], blk["genes"], feat, ylabel(tabs[stem], feat, ylab), + title=f"{blk['b']['label']} — {blk['b']['marker_label']}") + save(fig, f"violin_{stem}_{feat}") + + order = [blocks[i] for i in COMBINED_ORDER] # combined: primary feature, panel order A-D + fig, axes = plt.subplots(2, 2, figsize=(13.5, 10.4), facecolor="white") + for ax, blk, letter in zip(axes.flat, order, "ABCD"): + stem = blk["b"]["stem"] + spec = SPECS[stem] + ylab = ylabel(tabs[stem], spec["primary"], dict(spec["feats"])[spec["primary"]]) + draw_violin(ax, tabs[stem], blk["genes"], spec["primary"], ylab, + title=f"{blk['b']['label']} — {blk['b']['marker_label']}") + ax.text(-0.16, 1.1, letter, transform=ax.transAxes, fontsize=24, fontweight="bold", va="top") + fig.suptitle("Real-cell morphometrics of the predictive phenotypes — NTC vs complex gene KOs", + fontsize=20, fontweight="bold", y=0.985) + fig.tight_layout(rect=(0, 0, 1, 0.955), h_pad=3.0, w_pad=3.0) + save(fig, "violin_combined_primary") + composite(blocks, tabs) + for stem, blk in by_stem.items(): # Δmedian for every candidate feature + d = tabs[stem] + for feat, _ in SPECS[stem]["feats"]: + med = lambda g: np.nanmedian(d.loc[(d["gene"] == g) & (d["feature"] == feat), "value"]) + nmed = med("NTC") + deltas = ", ".join(f"{g} {(med(g) - nmed) / (abs(nmed) or 1e-9) * 100:+.0f}%" for g in blk["genes"]) + star = " *primary" if feat == SPECS[stem]["primary"] else "" + print(f"{stem:22s} {feat:32s} NTC {nmed:9.3g} | {deltas}{star}") + + +if __name__ == "__main__": + main(refresh="--refresh" in sys.argv) diff --git a/src/ops_model/models/attention/diffex/figures/figure_multirank_ebi_grid.py b/src/ops_model/models/attention/diffex/figures/figure_multirank_ebi_grid.py new file mode 100644 index 0000000..dd75ea5 --- /dev/null +++ b/src/ops_model/models/attention/diffex/figures/figure_multirank_ebi_grid.py @@ -0,0 +1,274 @@ +"""EBI-complex SHAP grids from the multi_rank screen — per complex, a 3-cells x 3-geneKO block of the +top-SHAP KO cells (right) beside the same-shape block of top-SHAP NTC cells (left). + +Cells come from Alex's EBI multi_rank screen (shap_screen_ebi_{phase,fluor}_all.csv: per-cell SHAP under +the complex classifier, `gene` = complex member gene). Complex membership is read from the EBI yaml, and +for complexes with >3 members the 3 genes with the strongest top cell are used. U7 snRNP has only 2 +members (SNRPD3, SNRPG) -> its block is 2 rows, not 3. + +Blocks: U7 snRNP + EMC on Phase2D; EMC on lipid-droplet BODIPY + Dynein-1 on ER/Golgi COPE. +NTC is the top-SHAP NTC of the same channel, so both phase blocks show the identical NTC cells, and +intensity windows are per channel (KO+NTC pooled) so KO vs NTC brightness is comparable. + +Run: python figure_multirank_ebi_grid.py # combined + per-complex figures +""" +import os +import sys + +import matplotlib + +matplotlib.use("Agg") +import matplotlib.pyplot as plt +import numpy as np +import pandas as pd +import yaml + +from _setacc_common import CROP_SIZE, _materialize, composite, seg_crop + +MR = "/hpc/projects/icd.fast.ops/models/alex_lin_attention/v5/multi_rank" +EBI_YAML = "/hpc/projects/icd.fast.ops/configs/gene_clusters/EBI_complexes_v1_old_gene_names.yaml" +OUT = "/hpc/projects/icd.fast.ops/analysis/figure4_multirank_ebi" +CACHE = f"{OUT}/_cache" +plt.rcParams["pdf.fonttype"] = 42 +plt.rcParams["svg.fonttype"] = "none" +plt.rcParams["font.family"] = "sans-serif" +plt.rcParams["font.sans-serif"] = ["Arial", "Helvetica", "DejaVu Sans"] # Arial first: no Illustrator substitution + +N_CELLS = 3 # cells per gene row +N_GENES = 3 # gene rows per block (fewer if the complex has fewer members) +WIN = 6 # extra ranks materialized per row, so dropped stores don't shrink the row + +U7 = "U7 small nuclear ribonucleoprotein complex" +EMC = "Endoplasmic reticulum membrane complex, EMC8 variant" +DYNEIN = "Dynein-1 complex, variant 2" + +# ch_name = multi_rank `channel_name`; mc = DirConfig marker_channel (None -> phase); ch = raw zarr channel +BLOCKS = [ + dict(cx=U7, label="U7 snRNP complex", modality="phase", ch_name="Phase2D", mc=None, ch="Phase2D", + marker_label="label-free phase", stem="phase_U7_snRNP"), + dict(cx=EMC, label="EMC complex", modality="phase", ch_name="Phase2D", mc=None, ch="Phase2D", + marker_label="label-free phase", stem="phase_EMC"), + dict(cx=EMC, label="EMC complex", modality="fluor", ch_name="lipid droplet_BODIPY live cell dye", + mc="lipid droplet_BODIPY live cell dye", ch="GFP", marker_label="lipid droplet (BODIPY)", + stem="fluor_EMC_BODIPY"), + dict(cx=DYNEIN, label="Dynein-1 complex", modality="fluor", ch_name="ER/Golgi_COPE", + mc="ER/Golgi_COPE", ch="GFP", marker_label="ER/Golgi (COPE)", stem="fluor_Dynein1_COPE"), +] +# combined-figure order (2 per row): the two EMC blocks pair up (fluorescence left of phase), then Dynein-1, U7 +COMBINED_ORDER = [2, 1, 3, 0] +COLS = ["gene", "channel_name", "rank", "shap", "experiment", "well", "x_pheno", "y_pheno", "segmentation_id"] + + +def members(cx): + y = yaml.safe_load(open(EBI_YAML)) + for v in y.values(): + if v.get("name") == cx: + return list(v["genes"]) + raise KeyError(f"{cx!r} not in {EBI_YAML}") + + +def ebi_rows(modality): + """Cached SHAP rows for every gene/channel the BLOCKS of this modality need (the source CSVs are 2-11GB).""" + genes = {g for b in BLOCKS if b["modality"] == modality for g in members(b["cx"])} | {"NTC"} + chans = {b["ch_name"] for b in BLOCKS if b["modality"] == modality} + p = f"{CACHE}/ebi_{modality}.parquet" + if os.path.exists(p): + df = pd.read_parquet(p) + if genes <= set(df["gene"]) and chans <= set(df["channel_name"]): + return df + os.makedirs(CACHE, exist_ok=True) + csv = f"{MR}/shap_screen_ebi_{modality}_all.csv" + print(f"[cache] scanning {csv} for {len(genes)} genes x {len(chans)} channels ...", flush=True) + keep = [c[c["gene"].isin(genes) & c["channel_name"].isin(chans)] + for c in pd.read_csv(csv, usecols=COLS, chunksize=2_000_000)] + df = pd.concat(keep, ignore_index=True) + df.to_parquet(p) + print(f"[cache] wrote {p} ({len(df):,} rows)", flush=True) + return df + + +def top_rows(df, gene, ch_name, n): + """Top-n SHAP cells of one (gene, channel), rank-ordered, in _materialize's column contract.""" + d = df[(df["gene"] == gene) & (df["channel_name"] == ch_name)].sort_values("rank").head(n).copy() + d["score"] = d["shap"] + d["segmentation"] = d["segmentation_id"] + return d + + +def pick_genes(df, b): + """The N_GENES complex members with the strongest top cell in this channel (all members if fewer), + displayed in EBI member order (EMC1/EMC2/EMC3, not SHAP order).""" + mem = members(b["cx"]) + d = df[df["gene"].isin(mem) & (df["channel_name"] == b["ch_name"])] + best = d.groupby("gene")["shap"].max().sort_values(ascending=False) + missing = [g for g in mem if g not in best.index] + if missing: + print(f" [{b['label']} / {b['ch_name']}] members absent from the screen: {missing}") + keep = set(best.index[:N_GENES]) + return [g for g in mem if g in keep] + + +def crop_row(df, b, gene, n): + """Crop the top-n surviving cells of one gene (rank order preserved).""" + raw, recs = _materialize(top_rows(df, gene, b["ch_name"], WIN + n), b["mc"], b["ch"], gene) + o = np.argsort(recs["rank"].values)[:n] + if len(o) < n: + raise RuntimeError(f"{gene} ({b['ch_name']}): only {len(o)}/{n} cells survived cropping") + return raw[o], recs.iloc[o].reset_index(drop=True) + + +def build_block(b): + """(genes, ko_raw[G,N,1,H,W], ko_recs list, ntc_raw[G,N,...], ntc_recs) for one complex x channel.""" + df = ebi_rows(b["modality"]) + genes = pick_genes(df, b) + ko = [crop_row(df, b, g, N_CELLS) for g in genes] + ntc_raw, ntc_recs = crop_row(df, b, "NTC", len(genes) * N_CELLS) + print(f" [{b['label']} / {b['ch_name']}] rows {genes} " + + " ".join(f"{g}:#{'/#'.join(str(int(r)) for r in rec['rank'])}" for g, (_, rec) in zip(genes, ko)) + + f" NTC:#{'/#'.join(str(int(r)) for r in ntc_recs['rank'])}", flush=True) + return dict(b=b, genes=genes, ko=ko, ntc=(ntc_raw, ntc_recs)) + + +def tile(raw, rec, lo, hi): + half = CROP_SIZE // 2 + gray = np.clip((raw - lo) / (hi - lo), 0, 1) * 255 + return composite(gray, seg_crop(rec["experiment"], rec["well"], rec["x_pheno"], rec["y_pheno"], half), half) + + +# block geometry, inches — tiles are placed by hand (not gridspec) so they stay square with no +# aspect padding, which is where the dead space came from. +T = 1.5 # tile edge +GAP = 0.03 # gap between tiles +MID = 0.42 # gap between the NTC half and the KO half +TITLE = 0.86 # title + subheader band above the tiles +BGX, BGY = 0.34, 0.3 # gaps between blocks (combined figure) +SUP, FOOT = 0.62, 0.5 # suptitle / footer bands +FS_TITLE, FS_SUB, FS_GENE, FS_BADGE, FS_SUP, FS_LET = 21, 16, 17, 12, 26, 26 +SHOW_BADGES = False # rank badges on each tile (cells are rank-ordered regardless) +BADGE_A = 0.3 # rank-badge alpha when shown +HALF = N_CELLS * T + (N_CELLS - 1) * GAP # width of one 3-tile half +CAPTION = ("Cells ordered by per-cell SHAP from the EBI-complex multi_rank screen (most predictive first, left to " + "right); NTC = top-SHAP non-targeting cells of the same channel, shared intensity window per channel.") + + +def block_h(nrows): + return TITLE + nrows * T + (nrows - 1) * GAP + + +def block_w(genes): + """Block width incl. the gene-label gutter, sized to the longest gene name (FS_GENE bold).""" + return 2 * HALF + MID + 0.18 + 0.115 * max(len(g) for g in genes) + + +def _put(fig, W, H, x, y, im, badge): + """One square tile at (x, y) inches from the figure's top-left.""" + ax = fig.add_axes([x / W, 1 - (y + T) / H, T / W, T / H]) + ax.imshow(im) + ax.set_xticks([]); ax.set_yticks([]) + for s in ax.spines.values(): + s.set_edgecolor("#888"); s.set_linewidth(0.5) + if SHOW_BADGES: + ax.text(0.04, 0.96, badge, transform=ax.transAxes, fontsize=FS_BADGE, fontweight="bold", color="white", + alpha=0.72, va="top", ha="left", + bbox=dict(boxstyle="round,pad=0.12", fc="#c1272d", ec="none", alpha=BADGE_A)) + + +def draw_block(fig, blk, x0, y0, lo, hi, W, H, letter=None): + """NTC 3xN (left) | KO 3xN (right, one row per gene) at inch offset (x0, y0) from the top-left.""" + b, genes = blk["b"], blk["genes"] + nrows = len(genes) + xk = x0 + HALF + MID + ytop = y0 + TITLE + ntc_raw, ntc_recs = blk["ntc"] + for i, g in enumerate(genes): + ko_raw, ko_recs = blk["ko"][i] + y = ytop + i * (T + GAP) + for j in range(N_CELLS): + k = i * N_CELLS + j + _put(fig, W, H, x0 + j * (T + GAP), y, tile(ntc_raw[k, 0], ntc_recs.iloc[k], lo, hi), + f"#{int(ntc_recs.iloc[k]['rank'])}") + _put(fig, W, H, xk + j * (T + GAP), y, tile(ko_raw[j, 0], ko_recs.iloc[j], lo, hi), + f"#{int(ko_recs.iloc[j]['rank'])}") + fig.text((xk + HALF + 0.1) / W, 1 - (y + T / 2) / H, g, fontsize=FS_GENE, fontweight="bold", + color="#c1272d", va="center", ha="left") + fig.text((x0 + (2 * HALF + MID) / 2) / W, 1 - (y0 + 0.3) / H, + f"{b['label']} — {b['marker_label']}", fontsize=FS_TITLE, fontweight="bold", + va="center", ha="center") + if letter: + fig.text((x0 - 0.3) / W, 1 - (y0 + 0.26) / H, letter, fontsize=FS_LET, fontweight="bold", + va="center", ha="left") + fig.text((x0 + HALF / 2) / W, 1 - (ytop - 0.13) / H, "control (NTC)", fontsize=FS_SUB, + va="bottom", ha="center", color="#444") + fig.text((xk + HALF / 2) / W, 1 - (ytop - 0.13) / H, + f"{b['label'].replace(' complex', '')} gene KOs", fontsize=FS_SUB, fontweight="bold", + va="bottom", ha="center", color="#c1272d") + xd = (x0 + HALF + MID / 2) / W # divider between the halves + fig.add_artist(plt.Line2D([xd, xd], [1 - (ytop + nrows * T + (nrows - 1) * GAP) / H, 1 - (ytop - 0.06) / H], + color="#bbb", lw=1.0, transform=fig.transFigure)) + + +def windows(blocks): + """1-99 pct intensity window per channel over that channel's KO+NTC crops (comparable brightness).""" + pool = {} + for blk in blocks: + raws = [r for r, _ in blk["ko"]] + [blk["ntc"][0]] + pool.setdefault(blk["b"]["ch_name"], []).extend(x.ravel() for x in raws) + out = {} + for ch, vals in pool.items(): + lo, hi = np.percentile(np.concatenate(vals), (1, 99)) + out[ch] = (lo, hi if hi - lo > 1e-6 else lo + 1) + return out + + +def save(fig, stem): + os.makedirs(OUT, exist_ok=True) + for ext in ("png", "svg"): + fig.savefig(f"{OUT}/{stem}.{ext}", dpi=220, bbox_inches="tight", facecolor="white") + plt.close(fig) + print(f"saved {OUT}/{stem}.png/.svg", flush=True) + + +def build_blocks(refresh=False): + """All blocks, with the cropped tiles cached so layout iterations don't re-read the zarrs.""" + p = f"{CACHE}/blocks_{N_GENES}x{N_CELLS}.pkl" + if os.path.exists(p) and not refresh: + return pd.read_pickle(p) + blocks = [build_block(b) for b in BLOCKS] + os.makedirs(CACHE, exist_ok=True) + pd.to_pickle(blocks, p) + return blocks + + +def main(ncol=2, refresh=False): + blocks = build_blocks(refresh) + win = windows(blocks) + + for blk in blocks: # per-complex figures + W, H = block_w(blk["genes"]) + 0.16, block_h(len(blk["genes"])) + 0.08 + fig = plt.figure(figsize=(W, H), facecolor="white") + draw_block(fig, blk, 0.08, 0.04, *win[blk["b"]["ch_name"]], W, H) + save(fig, f"ebi_shap_grid_{blk['b']['stem']}") + + order = [blocks[i] for i in COMBINED_ORDER] # combined figure + rows = [order[i:i + ncol] for i in range(0, len(order), ncol)] + rowh = [max(block_h(len(blk["genes"])) for blk in r) for r in rows] + colw = [max(block_w(r[c]["genes"]) for r in rows if c < len(r)) for c in range(ncol)] + x0 = 0.45 # left margin holds the panel letters + W = x0 + sum(colw) + (ncol - 1) * BGX + 0.12 + H = SUP + sum(rowh) + BGY * (len(rows) - 1) + FOOT + fig = plt.figure(figsize=(W, H), facecolor="white") + y = SUP + for r, row in enumerate(rows): + for c, blk in enumerate(row): + draw_block(fig, blk, x0 + sum(colw[:c]) + c * BGX, y, *win[blk["b"]["ch_name"]], W, H, + letter="ABCDEFGH"[r * ncol + c]) + y += rowh[r] + BGY + fig.text(0.5, 1 - 0.3 / H, "Top Predictive cells per protein complex", + fontsize=FS_SUP, fontweight="bold", ha="center", va="center") + fig.text(0.5, (FOOT - 0.28) / H, CAPTION, fontsize=13, style="italic", color="#333", + ha="center", va="center") + save(fig, "ebi_shap_grid_combined") + + +if __name__ == "__main__": + main(refresh="--refresh" in sys.argv) diff --git a/src/ops_model/models/attention/diffex/figures/fluor_panel_montages.py b/src/ops_model/models/attention/diffex/figures/fluor_panel_montages.py new file mode 100644 index 0000000..a24be1c --- /dev/null +++ b/src/ops_model/models/attention/diffex/figures/fluor_panel_montages.py @@ -0,0 +1,69 @@ +"""Multibag pick montages for the fluorescence panels C/D — one KO + one NTC montage per panel column, +using each column's ACTUAL panel marker (GENE_COLS/COMPLEX_COLS) and the multibag SHAP rankings +(fluor_shap/{geneKO,complex}). Fixes the montage↔panel marker mismatch (e.g. TOMM20 = Mitochondria_TOMM20, +not ChromaLIVE). Crops the marker channel + blue seg mask; rank/cell-key/shap annotated. + +Run (SLURM): python fluor_panel_montages.py --submit +""" +import sys + +import pandas as pd + +from _setacc_common import GENE_COLS, COMPLEX_COLS, _materialize, slugify +from fluor_shap_montages import render_montage # reuses OUT=figure4_shap_montages + seg overlay + +R = "/hpc/projects/icd.fast.ops/models/diffex/viewer_assets_v5/_rankings/fluor_shap" +COLS = GENE_COLS + COMPLEX_COLS +N = 100 + + +def _df(mc, block, cls): + sub = "complex" if (block == "complexes" and cls != "NTC") else "geneKO" # NTC always from the geneKO parquet + d = pd.read_parquet(f"{R}/{sub}/{slugify(mc)}.parquet") + d = d[d["gene"].astype(str) == str(cls)] + if "rank_type" in d.columns: + d = d[d["rank_type"] == "top"] + return d.sort_values("rank").reset_index(drop=True) + + +def col_montages(col): + for tag, cls in (("KO", col["key"]), ("NTC", "NTC")): + d = _df(col["mc"], col["block"], cls).head(N) + if d.empty: + print(f"skip {tag} {col['top_label']} ({col['mc']}): no rows", flush=True); continue + raw, recs = _materialize(d, col["mc"], col["ch"], cls) + lbl = col["top_label"].replace("\n", " ") + render_montage(raw, recs, f"{tag} {lbl} ({col['mc']}) multibag SHAP", + f"fluorpanel_{tag}_{col['slug']}") + + +def main(slugs=None): + want = set(slugs) if slugs else None + for c in COLS: + if want and c["slug"] not in want: + continue + col_montages(c) + + +def _job(slug): + import os + os.environ.setdefault("OPS_DIFFEX_ASSETS", "viewer_assets_v5") + main([slug]) + + +def submit(slugs=None): + import os + import pathlib + from ops_utils.hpc.slurm_batch_utils import submit_parallel_jobs + figdir = str(pathlib.Path(__file__).resolve().parent) + os.environ["PYTHONPATH"] = figdir + os.pathsep + os.environ.get("PYTHONPATH", "") + os.environ.setdefault("OPS_DIFFEX_ASSETS", "viewer_assets_v5") + ss = slugs or [c["slug"] for c in COLS] + jobs = [{"name": f"flpanel_{s[:16]}", "func": _job, "kwargs": {"slug": s}} for s in ss] + submit_parallel_jobs(jobs, experiment="diffex_flpanel", + slurm_params={"slurm_partition": "cpu", "cpus_per_task": 8, "mem_gb": 64, "timeout_min": 90}, + log_dir="diffex_flpanel", wait_for_completion=False) + + +if __name__ == "__main__": + submit(sys.argv[2:] or None) if (len(sys.argv) > 1 and sys.argv[1] == "--submit") else main(sys.argv[1:] or None) diff --git a/src/ops_model/models/attention/diffex/figures/fluor_shap_montages.py b/src/ops_model/models/attention/diffex/figures/fluor_shap_montages.py new file mode 100644 index 0000000..7048fee --- /dev/null +++ b/src/ops_model/models/attention/diffex/figures/fluor_shap_montages.py @@ -0,0 +1,140 @@ +"""SHAP-ranked fluor cell montages (v5/multi_rank) — top-N cells per fig-4 fluor group, rank-ordered with +rank + cell-key + SHAP annotations, marker-global normalized, inverse blue seg mask. One montage per KO +group + one per marker NTC, so cells can be hand-picked. ALSO writes per-marker SHAP rank parquets +(fluor_rank format) to a new _rankings/fluor_multirank/ dir. + +Run: python fluor_shap_montages.py [GENE_substr] # default all groups +""" +import os +import sys + +import matplotlib + +matplotlib.use("Agg") +import matplotlib.pyplot as plt +import numpy as np +import pandas as pd + +from ops_model.models.attention.diffex.classifier.config import slugify +from _setacc_common import _materialize, seg_crop, composite, CROP_SIZE + +MR = "/hpc/projects/icd.fast.ops/models/alex_lin_attention/v5/multi_rank/shap_screen/shap_screen_fluor_all.compact.parquet" +OUT = "/hpc/projects/icd.fast.ops/analysis/figure4_shap_montages" +PQ = "/hpc/projects/icd.fast.ops/models/diffex/viewer_assets_v5/_rankings/fluor_multirank/geneKO" +plt.rcParams["pdf.fonttype"] = 42 + +# fig-4 fluor groups: multi_rank channel_name -> (gene, zarr channel) +GROUPS = [ + ("autophagosome_MAP1LC3B", "ATG9A", "GFP"), + ("actin filament_FastAct_SPY555 Live Cell Dye", "CAPZB", "mCherry"), + ("ER/Golgi COP-II_SEC23A", "GBF1", "GFP"), + ("lysosome_LysoTracker live-cell dye", "LAMTOR2", "GFP"), + ("lipid droplet_BODIPY live cell dye", "RAB7A", "GFP"), + ("clathrin vesicles_CLTA", "AP2M1", "GFP"), + ("stress granule_G3BP1", "EIF2S2", "GFP"), + ("chromatin_H2BC21", "AURKB", "mCherry"), + ("nucleolus-DFC_FBL", "NOP56", "GFP"), + ("lysosome_LAMP1", "ATP6V1B2", "GFP"), + ("proteasome_PSMB7", "PSMB6", "GFP"), + ("nucleolus-GC_NPM3", "POLR1B", "GFP"), + ("nucleus_NucleoLIVE Live Cell dye", "KIF23", "mCherry"), + ("mitochondria_ChromaLIVE 561 excitation", "TOMM20", "mCherry"), + ("5xUPRE", "HSPA5", "GFP"), # UPR reporter (set-acc panel group) +] +_MR = None + + +def _rows(channel_name, gene, n): + """Top-n SHAP cells for (channel, gene) from the 'top' pool, rank-ordered. score = shap.""" + global _MR + if _MR is None: + _MR = pd.read_parquet(MR, columns=["gene", "channel_name", "rank", "shap", "_pool", + "experiment", "well", "x_pheno", "y_pheno", "segmentation_id"]) + d = _MR[(_MR["channel_name"] == channel_name) & (_MR["gene"] == gene) & (_MR["_pool"] == "top")] + d = d.sort_values("rank").head(n).copy() + d["score"] = d["shap"] + d["segmentation"] = d["segmentation_id"] # make_labels_df expects 'segmentation' + return d + + +def render_montage(raw, recs, title, out_stem): + n = len(recs) + lo, hi = np.percentile(raw, (1, 99)) + if hi - lo < 1e-6: + hi = lo + 1 + half = CROP_SIZE // 2 + ncols = 10 + nrows = int(np.ceil(n / ncols)) + fig, axes = plt.subplots(nrows, ncols, figsize=(ncols * 1.9, nrows * 2.05), facecolor="white") + axes = np.atleast_2d(axes) + for k in range(nrows * ncols): + ax = axes.flat[k]; ax.set_xticks([]); ax.set_yticks([]) + for s in ax.spines.values(): + s.set_visible(False) + if k >= n: + ax.axis("off"); continue + r = recs.iloc[k] + gray = np.clip((raw[k, 0] - lo) / (hi - lo), 0, 1) * 255 + seg = seg_crop(r["experiment"], r["well"], r["x_pheno"], r["y_pheno"], half) + ax.imshow(composite(gray, seg, half)) + ax.text(0.03, 0.97, f"#{int(r['rank'])}", transform=ax.transAxes, fontsize=13, fontweight="bold", + color="white", va="top", ha="left", + bbox=dict(boxstyle="round,pad=0.15", fc="#c1272d", ec="none", alpha=0.9)) + key = f"{r['experiment']}/{r['well']} x{int(round(r['x_pheno']))} y{int(round(r['y_pheno']))}" + ax.text(0.5, -0.03, f"{key}\nshap={float(r['score']):.3f}", transform=ax.transAxes, + fontsize=6.0, color="#222", va="top", ha="center") + fig.suptitle(title, fontsize=15, fontweight="bold", y=0.997) + fig.subplots_adjust(left=0.005, right=0.995, top=0.965, bottom=0.01, wspace=0.05, hspace=0.32) + os.makedirs(OUT, exist_ok=True) + out = f"{OUT}/{out_stem}.png" + fig.savefig(out, dpi=150, bbox_inches="tight", facecolor="white") + plt.close(fig) + print(f"saved {out} ({n} cells)", flush=True) + + +def montage(channel_name, gene, ch, title, out_stem, top_n=100): + rows = _rows(channel_name, gene, top_n) + if rows.empty: + print(f"skip {gene} ({channel_name}): no SHAP rows"); return + raw, recs = _materialize(rows, channel_name, ch, gene) + render_montage(raw, recs, title, out_stem) + + +def write_parquet(channel_name, gene, ch): + """Per-marker SHAP rank parquet (fluor_rank format: gene + base cols + channel + rank_type + rank), KO + NTC.""" + os.makedirs(PQ, exist_ok=True) + ko = _rows(channel_name, gene, 200); ntc = _rows(channel_name, "NTC", 200) + df = pd.concat([ko, ntc], ignore_index=True) + df["channel"] = channel_name; df["rank_type"] = "top" + cols = ["gene", "experiment", "well", "x_pheno", "y_pheno", "segmentation_id", "channel", "rank_type", "rank", "score"] + out = f"{PQ}/{slugify(channel_name)}.parquet" + df[cols].to_parquet(out) + print(f"parquet {out} (KO {len(ko)} + NTC {len(ntc)})", flush=True) + + +def main(): + filt = sys.argv[1] if len(sys.argv) > 1 else None + ntc_done = set() + for channel_name, gene, ch in GROUPS: + if filt and filt.lower() not in gene.lower() and filt.lower() not in channel_name.lower(): + continue + try: + write_parquet(channel_name, gene, ch) + except Exception as e: + print(f"parquet skip {gene}: {type(e).__name__}: {e}") + try: + montage(channel_name, gene, ch, f"KO — {gene} ({channel_name}) SHAP-ranked (multi_rank)", + f"shap_KO_{slugify(channel_name)}_{gene}") + except Exception as e: + print(f"KO montage skip {gene}: {type(e).__name__}: {e}") + if channel_name not in ntc_done: + ntc_done.add(channel_name) + try: + montage(channel_name, "NTC", ch, f"NTC — {channel_name} marker SHAP-ranked (multi_rank)", + f"shap_NTC_{slugify(channel_name)}") + except Exception as e: + print(f"NTC montage skip {channel_name}: {type(e).__name__}: {e}") + + +if __name__ == "__main__": + main() diff --git a/src/ops_model/models/attention/diffex/figures/gen_validation/bag_sweep_plots.py b/src/ops_model/models/attention/diffex/figures/gen_validation/bag_sweep_plots.py new file mode 100644 index 0000000..daf3d7e --- /dev/null +++ b/src/ops_model/models/attention/diffex/figures/gen_validation/bag_sweep_plots.py @@ -0,0 +1,382 @@ +"""Bag-size sweep plots for the new multibag v5 traversals. Three measures, one line per bag size: + (1) SetTransformer P(target) / rank / top-k vs α (bags 20,50,100,200,400) + (2) centroid-recovery mAP / top-1 / top-5 vs α (bags 20,50,100,200,400) + (3) within-domain distinctiveness mAP vs α (K 20,50,100 — copairs cap) +Split geneKO/complex and (for 1) all-classes vs real-distinguishable (real top1_acc>0.5 @bag20). +""" +import os, glob, json +import numpy as np +import matplotlib +matplotlib.use("Agg") +import matplotlib.pyplot as plt +from matplotlib import cm + +plt.rcParams["pdf.fonttype"] = 42 +plt.rcParams.update({ # figure-ready: big, readable text + heavy elements + "font.size": 24, "axes.titlesize": 32, "axes.labelsize": 30, + "xtick.labelsize": 24, "ytick.labelsize": 24, "legend.fontsize": 22, + "figure.titlesize": 36, "axes.linewidth": 2.0, + "xtick.major.size": 10, "ytick.major.size": 10, "xtick.major.width": 2.0, "ytick.major.width": 2.0, + "lines.linewidth": 3.6, "legend.title_fontsize": 22, +}) +LW, LWD, LEG = 5.5, 4.0, 22 # solid / dotted line widths; legend fontsize +CV = "/hpc/projects/icd.fast.ops/analysis/figure4_traversals" +OUT = f"{CV}/bag_sweep_plots_v5new"; os.makedirs(OUT, exist_ok=True) +B = "/hpc/projects/icd.fast.ops/models/diffex" +BAGS = [20, 50, 100, 200] # 400 dropped: its cells 200-399 are the weak strict-multibag anchors (see anchor_halves) +KS = [20, 50, 100, 200] +CENT_STD = os.environ.get("CENT_STD", "perbag") # centroid/pooled plots read per-bag α=0 std (domain-honest, matches score_embs_v5); "global" for old panel-mu +CENT_SUF = "_perbag" if CENT_STD == "perbag" else "" +COL = {b: cm.viridis(i / (len(BAGS) - 1)) for i, b in enumerate(BAGS)} +COLK = {k: cm.viridis(i / (len(KS) - 1)) for i, k in enumerate(KS)} +from ops_model.models.attention.diffex.classifier.config import slugify +REAL = json.load(open(f"{B}/viewer_assets_v5/real_acc20.json")) + + +def _keep(grain): + pre = f"phase/{grain}/" + return {k[len(pre):] for k, v in REAL.items() if k.startswith(pre) and v > 0.5} + + +# ---------- (1) SetTransformer bag-sweep ---------- +def _agg_st(grain): + agg = {b: {"al": None, "P": [], "RK": [], "T5": [], "T1": [], "names": []} for b in BAGS} + for f in glob.glob(f"{CV}/bag_sweep_v5new/{grain}/*.json"): + d = json.load(open(f)); g = d["gene"] + for b in BAGS: + s = d["by_bag"].get(str(b)) or d["by_bag"].get(b) + if not s or s.get("p_target") is None: + continue + agg[b]["al"] = s["alphas"] + agg[b]["P"].append([np.nan if v is None else v for v in s["p_target"]]) + agg[b]["RK"].append([np.nan if v is None else v for v in s["rank_target"]]) + agg[b]["T5"].append([np.nan if v is None else v for v in s["top5_target"]]) + agg[b]["T1"].append([np.nan if v is None else v for v in s["top1_target"]]) + agg[b]["names"].append(g) + return agg + + +def plot_settransformer(): + A = {g: _agg_st(g) for g in ["geneKO", "complex"]} + for subset, tag in [(None, "all"), (True, "realdist")]: + fig, axes = plt.subplots(2, 4, figsize=(30, 14)) # 2 rows = geneKO/complex; 4 cols + for row, grain in enumerate(["geneKO", "complex"]): + ax = axes[row]; agg = A[grain]; keep = _keep(grain) if subset else None + n = 0 + for b in BAGS: + a = agg[b] + if a["al"] is None: + continue + al = a["al"] + msk = np.array([True] * len(a["names"])) if keep is None else np.array([slugify(x) in keep or x in keep for x in a["names"]]) + if not msk.any(): + continue + n = int(msk.sum()) + P = np.array(a["P"])[msk]; RK = np.array(a["RK"])[msk]; T5 = np.array(a["T5"])[msk]; T1 = np.array(a["T1"])[msk] + ax[0].plot(al, np.nanmean(P, 0), "-", color=COL[b], lw=LW, label=f"bag {b}") + ax[1].plot(al, np.nanmedian(RK, 0), "-", color=COL[b], lw=LW, label=f"bag {b}") + ax[2].plot(al, np.nanmean(RK, 0), "-", color=COL[b], lw=LW, label=f"bag {b}") + ax[3].plot(al, np.nanmean(T5, 0) * 100, "-", color=COL[b], lw=LW, label=f"bag {b}") + ax[3].plot(al, np.nanmean(T1, 0) * 100, ":", color=COL[b], lw=LWD) + ax[0].set_ylabel(f"{grain} (n={n})\n\nP(target)"); ax[0].set_ylim(-.02, 1.02) + ax[1].set_ylabel("median target rank"); ax[1].set_yscale("log"); ax[1].axhline(1, color="#ccc") + ax[2].set_ylabel("mean target rank"); ax[2].set_yscale("log"); ax[2].axhline(1, color="#ccc") + ax[3].set_ylabel("% recovered"); ax[3].set_ylim(-2, 102) + if row == 0: + ax[0].set_title("P(target)"); ax[1].set_title("median target rank"); ax[2].set_title("mean target rank"); ax[3].set_title("% top-5 (solid) / top-1 (dotted)") + for a_ in ax: + a_.set_xlabel("traversal α"); a_.set_xticks(range(-5, 6)); a_.grid(alpha=.25); a_.axvline(0, color="#ccc", lw=1); a_.legend(fontsize=LEG) + ttl = ("all classes" if tag == "all" else + "real-distinguishable classes only (genes whose REAL cells score top-1 accuracy > 0.5 @ bag-20)") + fig.suptitle(f"SetTransformer bag-sweep — {ttl} · new multibag v5", fontweight="bold") + fig.tight_layout() + for e in ("png", "svg"): + fig.savefig(f"{OUT}/settransformer_bagsweep_{tag}.{e}", dpi=150, bbox_inches="tight") + plt.close(fig); print(f"saved settransformer_bagsweep_{tag}") + + +# ---------- (2) centroid recovery bag-sweep ---------- +def plot_centroid(): + fig, axes = plt.subplots(2, 3, figsize=(23, 14)) # 2 rows = geneKO / complex + for row, grain in enumerate(["geneKO", "complex"]): + ax = axes[row] + p = f"{CV}/centroid_bagsweep_v5new{CENT_SUF}/{grain}_bagsweep.json" + if not os.path.exists(p): + print(f"no centroid bagsweep for {grain}"); continue + d = json.load(open(p)); by = d["by_bag"] + for b in BAGS: + bb = by.get(str(b)) or by.get(b) + if not bb: + continue + al = sorted(float(a) for a in bb) + mp = [np.mean(list(bb[str(a) if str(a) in bb else a]["map"].values())) for a in al] + t1 = [np.mean(list(bb[str(a) if str(a) in bb else a]["top1"].values())) for a in al] + t5 = [np.mean(list(bb[str(a) if str(a) in bb else a]["top5"].values())) for a in al] + ax[0].plot(al, mp, "-", color=COL[b], lw=LW, label=f"bag {b}") + ax[1].plot(al, np.array(t1) * 100, "-", color=COL[b], lw=LW, label=f"bag {b}") + ax[2].plot(al, np.array(t5) * 100, "-", color=COL[b], lw=LW, label=f"bag {b}") + cf = f"{CV}/centroid_bagsweep_v5new/{grain}_ceiling.json" # real-cell ceiling (per-cell → bag-independent) + if os.path.exists(cf): + c = json.load(open(cf)) + ax[0].axhline(c["map"], ls=":", color="#c0392b", lw=LWD, label="real cells") + ax[1].axhline(c["top1"] * 100, ls=":", color="#c0392b", lw=LWD, label="real cells") + ax[2].axhline(c["top5"] * 100, ls=":", color="#c0392b", lw=LWD, label="real cells") + ax[0].set_ylabel(f"{grain} (n={d['n_classes']})\n\nmAP") + ax[1].set_ylabel("% of cells (top-1)") + ax[2].set_ylabel("% of cells (top-5)") + if row == 0: + ax[0].set_title("centroid-recovery mAP"); ax[1].set_title("top-1"); ax[2].set_title("top-5") + for a_ in ax: + a_.set_xlabel("traversal α"); a_.set_xticks(range(-5, 6)); a_.grid(alpha=.25); a_.axvline(0, color="#ccc", lw=1); a_.legend(fontsize=LEG) + fig.suptitle(f"Centroid-recovery bag-sweep — multibag v5 ({CENT_STD} α=0 standardization)", fontweight="bold") + fig.tight_layout() + gtag = "" if CENT_STD == "perbag" else "_global" + for e in ("png", "svg"): + fig.savefig(f"{OUT}/centroid_bagsweep{gtag}.{e}", dpi=150, bbox_inches="tight") + plt.close(fig); print(f"saved centroid_bagsweep{gtag} (2-row)") + + +# ---------- (2b) pooled bag-level centroid recovery (per-bag real ceiling) ---------- +def plot_centroid_pooled(): + fig, axes = plt.subplots(1, 2, figsize=(18, 8)) # geneKO | complex + for ax, grain in zip(axes, ["geneKO", "complex"]): + p = f"{CV}/centroid_pooled_bagsweep_v5new{CENT_SUF}/{grain}_pooled.json" + if not os.path.exists(p): + print(f"no pooled centroid for {grain}"); continue + d = json.load(open(p)); ng = 0 + for b in BAGS: + bb = d["gen"].get(str(b)) + if not bb: + continue + al = sorted(float(a) for a in bb) + t1 = [np.mean(list(bb[str(a) if str(a) in bb else a]["top1"].values())) for a in al] + ax.plot(al, np.array(t1) * 100, "-", color=COL[b], lw=LW, label=f"gen bag {b}") + rc = d["real"].get(str(b)) + if rc: + ng = len(rc); ax.axhline(np.mean(list(rc.values())) * 100, ls=":", color=COL[b], lw=LWD) + ax.set_xlabel("traversal α"); ax.set_xticks(range(-5, 6)); ax.grid(alpha=.25); ax.axvline(0, color="#ccc", lw=1) + ax.set_ylim(-2, 102); ax.set_title(f"{grain} (n={ng})") + axes[0].set_ylabel("% classes recovering\ntrue real centroid (top-1)") + from matplotlib.lines import Line2D + handles = [Line2D([0], [0], color=COL[b], lw=LW, label=f"gen bag {b}") for b in BAGS] + handles.append(Line2D([0], [0], color="k", ls=":", lw=LWD, label="real ceiling (per bag)")) + fig.legend(handles=handles, loc="center left", bbox_to_anchor=(0.87, 0.5), frameon=False, fontsize=LEG) + fig.suptitle(f"Pooled bag-level centroid recovery · multibag v5 ({CENT_STD} α=0 std)\nsolid = generated · dotted = real-cell ceiling (per bag)", + fontweight="bold", fontsize=24) + fig.tight_layout(rect=[0.02, 0, 0.86, 0.96]) + gtag = "" if CENT_STD == "perbag" else "_global" + for e in ("png", "svg"): + fig.savefig(f"{OUT}/centroid_pooled_bagsweep{gtag}.{e}", dpi=150, bbox_inches="tight") + plt.close(fig); print(f"saved centroid_pooled_bagsweep{gtag}") + + +# ---------- (2c) cross-domain retrieval mAP (proper1k, full real gallery) ---------- +def plot_proper1k(): + fig, axes = plt.subplots(1, 2, figsize=(18, 8)) + for ax, grain in zip(axes, ["geneKO", "complex"]): + p = f"{CV}/gen_real_centroid_v5new/{grain}_propermap1k.json" + if not os.path.exists(p): + print(f"no proper1k for {grain}"); continue + d = json.load(open(p)); al = sorted(float(a) for a in d["gen"]); by = d["gen"] + mp = [np.mean(list(by[str(a) if str(a) in by else a].values())) for a in al] + ax.plot(al, mp, "-", color="#2e8b57", lw=LW, label="generated → real") + if d.get("ceiling"): + ax.axhline(np.mean(list(d["ceiling"].values())), ls=":", color="#c0392b", lw=LWD, label="real self-consistency") + ax.set_xlabel("traversal α"); ax.set_xticks(range(-5, 6)); ax.grid(alpha=.25); ax.axvline(0, color="#ccc", lw=1) + ax.set_ylim(-.02, 1.02); ax.set_title(f"{grain} (n={d['n_classes']})"); ax.legend(fontsize=LEG) + axes[0].set_ylabel("cross-domain retrieval mAP\n(generated cell → real class cells)") + fig.suptitle("Cross-domain retrieval mAP — generated → full 1000-cell real gallery · multibag v5", fontweight="bold", fontsize=26) + fig.tight_layout() + for e in ("png", "svg"): + fig.savefig(f"{OUT}/retrieval_map_proper1k.{e}", dpi=150, bbox_inches="tight") + plt.close(fig); print("saved retrieval_map_proper1k") + + +# ---------- (2d) first-200 vs second-200 anchor halves ---------- +def plot_halves(): + fig, axes = plt.subplots(1, 2, figsize=(16, 7)) + for ax, grain in zip(axes, ["geneKO", "complex"]): + p = f"{CV}/centroid_halves_v5new/{grain}_halves.json" + if not os.path.exists(p): + print(f"no halves for {grain}"); continue + d = json.load(open(p)); n = 0 + for key, col in [("first", "#1f77b4"), ("second", "#d62728")]: + al = sorted(float(a) for a in d[key]) + t1 = [np.mean(list(d[key][str(a) if str(a) in d[key] else a].values())) * 100 for a in al] + n = len(d[key][str(al[np.argmax(t1)])]) + ax.plot(al, t1, "-", color=col, lw=LW) + ax.set_xlabel("traversal α"); ax.set_xticks(range(-5, 6)); ax.grid(alpha=.25); ax.axvline(0, color="#ccc", lw=1) + ax.set_ylim(-2, 102); ax.set_title(f"{grain} (n={n})") + axes[0].set_ylabel("% classes recovering\ntrue real centroid (top-1)") + from matplotlib.lines import Line2D + fig.legend(handles=[Line2D([0], [0], color="#1f77b4", lw=LW, label="hand-picked anchors (curated, cells 0–199)"), + Line2D([0], [0], color="#d62728", lw=LW, label="strict multibag top-200 NTC (cells 200–399)")], + loc="lower center", bbox_to_anchor=(0.5, 0.005), ncol=2, frameon=False, fontsize=LEG) + fig.suptitle("Anchor selection drives recovery: hand-picked (first 200) vs strict multibag top-NTC (second 200)\nsame directions/traversals — only the anchor cells differ · multibag v5", + fontweight="bold", fontsize=22, y=0.99) + fig.subplots_adjust(top=0.84, bottom=0.28, left=0.08, right=0.97, wspace=0.16) # extra bottom room: legend clear of x-axis + for e in ("png", "svg"): + fig.savefig(f"{OUT}/anchor_halves.{e}", dpi=150, bbox_inches="tight") + plt.close(fig); print("saved anchor_halves") + + +# ---------- (2e) SetTransformer: first-200 vs second-200 anchor halves ---------- +def plot_st_halves(): + fig, axes = plt.subplots(2, 4, figsize=(25, 11)) + for row, grain in enumerate(["geneKO", "complex"]): + ax = axes[row]; files = glob.glob(f"{CV}/st_halves_v5new/{grain}/*.json") + agg = {h: {"al": None, "P": [], "RK": [], "T5": [], "T1": []} for h in ("first", "second")} + for f in files: + d = json.load(open(f)) + for h in ("first", "second"): + s = d.get(h) + if not s or s.get("p_target") is None: + continue + agg[h]["al"] = s["alphas"] + for k, kk in [("P", "p_target"), ("RK", "rank_target"), ("T5", "top5_target"), ("T1", "top1_target")]: + agg[h][k].append([np.nan if v is None else v for v in s[kk]]) + n = len(agg["first"]["P"]) + for h, col, lab in [("first", "#1f77b4", "hand-picked (0–199)"), ("second", "#d62728", "strict multibag NTC (200–399)")]: + a = agg[h] + if a["al"] is None: + continue + al = a["al"] + ax[0].plot(al, np.nanmean(a["P"], 0), "-", color=col, lw=LW, label=lab) + ax[1].plot(al, np.nanmedian(a["RK"], 0), "-", color=col, lw=LW, label=lab) + ax[2].plot(al, np.nanmean(a["RK"], 0), "-", color=col, lw=LW, label=lab) + ax[3].plot(al, np.nanmean(a["T5"], 0) * 100, "-", color=col, lw=LW, label=lab) + ax[3].plot(al, np.nanmean(a["T1"], 0) * 100, ":", color=col, lw=LWD) + ax[0].set_ylabel(f"{grain} (n={n})\n\nP(target)"); ax[0].set_ylim(-.02, 1.02) + ax[1].set_ylabel("median target rank"); ax[1].set_yscale("log"); ax[1].axhline(1, color="#ccc") + ax[2].set_ylabel("mean target rank"); ax[2].set_yscale("log"); ax[2].axhline(1, color="#ccc") + ax[3].set_ylabel("% recovered"); ax[3].set_ylim(-2, 102) + if row == 0: + ax[0].set_title("P(target)"); ax[1].set_title("median target rank") + ax[2].set_title("mean target rank"); ax[3].set_title("% top-5 (solid) / top-1 (dotted)") + for a_ in ax: + a_.set_xlabel("traversal α"); a_.set_xticks(range(-5, 6)); a_.grid(alpha=.25); a_.axvline(0, color="#ccc", lw=1) + h, l = axes[0][0].get_legend_handles_labels() + fig.legend(h, l, loc="center left", bbox_to_anchor=(0.995, 0.5), fontsize=LEG, frameon=False) + fig.suptitle("SetTransformer: hand-picked (first 200) vs strict multibag top-NTC (second 200) anchors · bag=200 each · multibag v5", + fontweight="bold", fontsize=22) + fig.tight_layout(rect=(0, 0, 0.995, 1)) + for e in ("png", "svg"): + fig.savefig(f"{OUT}/st_anchor_halves.{e}", dpi=150, bbox_inches="tight") + plt.close(fig); print("saved st_anchor_halves") + + +# ---------- (3) distinctiveness sweep ---------- +def plot_distinct(stat="median"): + agg = np.mean if stat == "mean" else np.median + suf = "_mean" if stat == "mean" else "_median" + fig, axes = plt.subplots(1, 2, figsize=(18, 7), sharey=True) + for ax, grain in zip(axes, ["geneKO", "complex"]): + for k in KS: + d = f"{CV}/gen_real_distinct_v5new_K{k}" + rf = f"{d}/{grain}_real.json" + if not os.path.exists(rf): + continue + al, v = [], [] + for f in sorted(glob.glob(f"{d}/{grain}_gen_a*.json"), key=lambda p: int(p.split("_a")[-1][:-5])): + g = json.load(open(f)); al.append(g["alpha"]); v.append(agg(list(g["gen"].values()))) + if not al: # gen OOM'd at this K (e.g. geneKO K≥100) → skip, no phantom ceiling + continue + order = np.argsort(al); al = np.array(al)[order]; v = np.array(v)[order] + ax.plot(al, v, "-", color=COLK[k], lw=LW, label=f"top-{k}") + ax.axhline(agg(list(json.load(open(rf)).values())), color=COLK[k], ls=":", lw=LWD) + ax.set_xlabel("traversal α"); ax.set_xticks(range(-5, 6)); ax.grid(alpha=.25); ax.axvline(0, color="#ccc", lw=1) + ax.set_title(grain); ax.legend(fontsize=LEG, title="cells/class") + axes[0].set_ylabel(f"{stat} distinctiveness mAP") + fig.suptitle(f"Distinctiveness sweep ({stat}) — new multibag v5", fontweight="bold") + fig.tight_layout() + for e in ("png", "svg"): + fig.savefig(f"{OUT}/distinct_sweep{suf}.{e}", dpi=150, bbox_inches="tight") + plt.close(fig); print(f"saved distinct_sweep{suf}") + + +def plot_distinct_violin(k=50): + """Per-class distinctiveness distribution (generated vs real) at each grain's peak α, K cells/class.""" + fig, axes = plt.subplots(1, 2, figsize=(15, 8), sharey=True) + for ax, grain in zip(axes, ["geneKO", "complex"]): + d = f"{CV}/gen_real_distinct_v5new_K{k}" + rf = f"{d}/{grain}_real.json" + if not os.path.exists(rf): + print(f"no distinct K{k} for {grain}"); continue + gens = {} + for f in glob.glob(f"{d}/{grain}_gen_a*.json"): + g = json.load(open(f)); gens[g["alpha"]] = g["gen"] + al = sorted(gens); med = [np.median(list(gens[a].values())) for a in al] + pa = al[int(np.argmax(med))] # peak α by median gen mAP + gv = np.array(list(gens[pa].values())); rv = np.array(list(json.load(open(rf)).values())) + parts = ax.violinplot([rv, gv], positions=[0, 1], widths=0.8, showextrema=False, showmedians=False) + for i, pc in enumerate(parts["bodies"]): + pc.set_facecolor(["#8fa9c9", "#2e8b57"][i]); pc.set_alpha(.85); pc.set_edgecolor("none") + for pos, vals in [(0, rv), (1, gv)]: # thick black median bar only (no box/extrema) + ax.hlines(np.median(vals), pos - 0.34, pos + 0.34, color="k", lw=6, zorder=5) + ax.set_xticks([0, 1]); ax.set_xticklabels(["real", f"generated\n(α={pa:+g})"]) + ax.set_title(f"{grain} (n={len(gv)})"); ax.grid(alpha=.25, axis="y") + axes[0].set_ylabel("distinctiveness / EBI\nmAP score") + from matplotlib.patches import Patch + from matplotlib.lines import Line2D + axes[-1].legend(handles=[Patch(facecolor="#8fa9c9", label="Real"), Patch(facecolor="#2e8b57", label="Generated (peak α)"), + Line2D([0], [0], color="k", lw=6, label="Median")], + loc="center left", bbox_to_anchor=(1.02, 0.5), frameon=False, fontsize=LEG) + fig.suptitle(f"Per-class distinctiveness: real vs generated (peak α, top-{k} cells/class) · multibag v5", fontweight="bold", fontsize=26) + fig.tight_layout() + for e in ("png", "svg"): + fig.savefig(f"{OUT}/distinct_violin.{e}", dpi=150, bbox_inches="tight") + plt.close(fig); print("saved distinct_violin") + + +def plot_control(): + """Control: the first/second centroid-recovery gap is a standardization artifact. Peak-α top-1 under GLOBAL + (panel α=0 mean, current metric) vs PER-BAG (each half's own α=0, as score_embs_v5 does). The gap collapses + under per-bag → the 76/13 is metric standardization, not phenotype loss.""" + fig, axes = plt.subplots(1, 2, figsize=(16, 8)) + schemes = [("global", "global α=0\n(panel mean)"), ("perbag", "per-bag α=0\n(own mean)")] + for ax, grain in zip(axes, ["geneKO", "complex"]): + c = json.load(open(f"{CV}/control_halves_v5new/{grain}_control.json")) + vals = {s: {} for s, _ in schemes} + for s, _ in schemes: + for h in ("first", "second"): + byA = {float(a): float(np.mean(list(r.values()))) for a, r in c[s][h].items() if r} + vals[s][h] = max(byA.values()) if byA else 0.0 + x = np.arange(len(schemes)); w = 0.36 + ax.bar(x - w / 2, [vals[s]["first"] * 100 for s, _ in schemes], w, color="#1f77b4", label="hand-picked (0–199)") + ax.bar(x + w / 2, [vals[s]["second"] * 100 for s, _ in schemes], w, color="#d62728", label="strict multibag NTC (200–399)") + for i, (s, _) in enumerate(schemes): + for dx, h in [(-w / 2, "first"), (w / 2, "second")]: + ax.text(i + dx, vals[s][h] * 100 + 1.5, f"{vals[s][h]*100:.0f}", ha="center", va="bottom", fontsize=20) + ax.set_xticks(x); ax.set_xticklabels([lab for _, lab in schemes]); ax.set_ylim(0, 100) + ax.set_title(grain); ax.grid(alpha=.25, axis="y") + axes[0].set_ylabel("peak-α centroid top-1 (%)") + h, l = axes[0].get_legend_handles_labels() + fig.legend(h, l, loc="upper center", bbox_to_anchor=(0.5, 0.055), ncol=2, fontsize=LEG, frameon=False) + fig.suptitle("Centroid-recovery gap is a standardization artifact: per-bag α=0 collapses it", fontweight="bold", fontsize=24) + fig.subplots_adjust(top=0.85, bottom=0.24, left=0.09, right=0.97, wspace=0.14) + for e in ("png", "svg"): + fig.savefig(f"{OUT}/control_halves_zscore.{e}", dpi=150, bbox_inches="tight") + plt.close(fig); print("saved control_halves_zscore") + + +if __name__ == "__main__": + import sys + which = sys.argv[1] if len(sys.argv) > 1 else "all" + if which == "control": + plot_control(); sys.exit() + if which == "violin": + plot_distinct_violin(); sys.exit() + if which == "pooled": + plot_centroid_pooled(); sys.exit() + if which == "proper1k": + plot_proper1k(); sys.exit() + if which == "halves": + plot_halves(); sys.exit() + if which == "sthalves": + plot_st_halves(); sys.exit() + if which in ("all", "st"): + plot_settransformer() + if which in ("all", "cent"): + plot_centroid() + if which in ("all", "dist"): + plot_distinct("median"); plot_distinct("mean") diff --git a/src/ops_model/models/attention/diffex/figures/gen_validation/bag_sweep_score.py b/src/ops_model/models/attention/diffex/figures/gen_validation/bag_sweep_score.py new file mode 100644 index 0000000..a7a6e50 --- /dev/null +++ b/src/ops_model/models/attention/diffex/figures/gen_validation/bag_sweep_score.py @@ -0,0 +1,58 @@ +"""Re-score the new multibag v5 traversals with the v5 SetTransformer at MULTIPLE bag sizes {20,50,100,200,400}. +Reads the CellDINO cache (gen = list[A] of (n_cells,1024) per gene) and runs score_embs_v5 at each bag → per-gene +{bag: scores_v5-dict}. GPU (loads the classifier once per shard). The rank/P(target)/top-k plots then draw one +line per bag size. bag=B scores the FIRST B cells (deterministic), standardized on the α=0 frames — same method +as the stored bag=45, just swept. +""" +import os, glob, json +import numpy as np + +CV = "/hpc/projects/icd.fast.ops/analysis/figure4_traversals" +GRAIN = os.environ.get("BSS_GRAIN", "geneKO") +CACHE = f"{CV}/gen_real_map_cache_v5new/{GRAIN}" +OUT = f"{CV}/bag_sweep_v5new/{GRAIN}" +BAGS = [20, 50, 100, 200, 400] + + +def score_shard(genes): + import torch + from ops_model.models.attention.diffex.viewer.score_generated import score_embs_v5 + from ops_model.models.attention.diffex.viewer.set_classifier import load_set_classifier, V5_CKPT_ROOT, V5_RUNS + os.makedirs(OUT, exist_ok=True) + dev = "cuda" if torch.cuda.is_available() else "cpu" + from ops_model.models.attention.diffex.classifier.config import slugify + run = V5_RUNS[("phase", "geneKO" if GRAIN == "geneKO" else "complex_ebionly")] + m, g2i, c2i = load_set_classifier(run=run, device=dev, root=V5_CKPT_ROOT) + ci = c2i.get("Phase2D", 0) + slug2orig = {slugify(k): k for k in g2i} # cache genes are slugified (KRTAP2_3); g2i uses originals (KRTAP2-3) + done = 0 + for g in genes: + f = f"{CACHE}/{g}.npz"; outp = f"{OUT}/{g}.json" + if not os.path.exists(f) or os.path.exists(outp): + continue + d = np.load(f, allow_pickle=True); gene = str(d["gene"]); al = [float(a) for a in d["alphas"]] + tgt = gene if gene in g2i else slug2orig.get(slugify(gene), gene) # normalize slug→classifier name + embs = [None if d["gen"][ai] is None else np.asarray(d["gen"][ai], np.float32) for ai in range(len(al))] + rec = {} + for B in BAGS: + rec[B] = score_embs_v5(embs, al, tgt, m, g2i, ci, run, device=dev, bag=B) + json.dump({"gene": gene, "alphas": al, "by_bag": rec}, open(outp, "w")) + done += 1 + return {"grain": GRAIN, "done": done} + + +def main(): + from ops_utils.hpc.slurm_batch_utils import submit_parallel_jobs + genes = sorted(os.path.basename(f)[:-4] for f in glob.glob(f"{CACHE}/*.npz")) + ch = 60 + shards = [genes[i:i + ch] for i in range(0, len(genes), ch)] + jobs = [{"name": f"bss_{GRAIN}_{i}", "func": score_shard, "kwargs": {"genes": s}} for i, s in enumerate(shards)] + print(f"[bag-sweep] {GRAIN}: {len(genes)} genes → {len(jobs)} shards, bags={BAGS}") + submit_parallel_jobs(jobs, experiment=f"bagsweep_{GRAIN}", + slurm_params={"slurm_partition": "preempted", "slurm_gres": "gpu:1", "cpus_per_task": 8, + "mem_gb": 48, "timeout_min": 120, "slurm_constraint": "[a40|a6000|l40s]"}, + log_dir=f"bagsweep_{GRAIN}", wait_for_completion=False) + + +if __name__ == "__main__": + main() diff --git a/src/ops_model/models/attention/diffex/figures/gen_validation/centroid_bagsweep.py b/src/ops_model/models/attention/diffex/figures/gen_validation/centroid_bagsweep.py new file mode 100644 index 0000000..6f5df48 --- /dev/null +++ b/src/ops_model/models/attention/diffex/figures/gen_validation/centroid_bagsweep.py @@ -0,0 +1,108 @@ +"""Centroid-recovery mAP bag-sweep on the new multibag v5 cache, SHARDED for speed. +Stage 1 (mu): global α=0 standardization from a gene subset (α=0 is NTC-recon, tight → subset≈full). +Stage 2 (score): parallel shards, each scores its genes at bags {20,50,100,200,400} → partial json. +Stage 3 (merge): combine partials → {grain}_bagsweep.json (per-bag by_alpha of top1/top5/map). + + CBS_GRAIN=geneKO python centroid_bagsweep.py # submit mu (inline) + score shards + CBS_GRAIN=geneKO python centroid_bagsweep.py merge # combine partials +""" +import os, glob, json, sys +import numpy as np + +CV = "/hpc/projects/icd.fast.ops/analysis/figure4_traversals" +GRAIN = os.environ.get("CBS_GRAIN", "geneKO") +CACHE = f"{CV}/gen_real_map_cache_v5new/{GRAIN}" +CENTD = f"{CV}/gen_real_centroid_v5new" +STD = os.environ.get("CBS_STD", "global") # global (panel α=0 mu) | perbag (per-gene α=0, matches score_embs_v5) +OUT = f"{CV}/centroid_bagsweep_v5new" + ("_perbag" if STD == "perbag" else "") +PART = f"{OUT}/{GRAIN}_parts" +BAGS = [20, 50, 100, 200, 400] +MU_N = 150 # genes for the global α=0 estimate + + +def _cz(): + from ops_model.models.attention.diffex.classifier.config import slugify + d = np.load(f"{CENTD}/{GRAIN}_centroids.npz", allow_pickle=True) + names = list(d["names"]); cidx = {slugify(str(c)): i for i, c in enumerate(names)} + cz = (d["cents"] - d["mu"]) / d["sd"]; cz = cz / (np.linalg.norm(cz, axis=1, keepdims=True) + 1e-9) + return cz, cidx + + +def compute_mu(): + os.makedirs(OUT, exist_ok=True) + S = np.zeros(1024); SS = np.zeros(1024); n = 0 + for f in sorted(glob.glob(f"{CACHE}/*.npz"))[:MU_N]: + d = np.load(f, allow_pickle=True); al = np.asarray(d["alphas"], float); a0 = int(np.argmin(np.abs(al))) + z = d["gen"][a0] + if z is not None and len(z): + z = np.asarray(z, np.float32); S += z.sum(0); SS += (z.astype(np.float64) ** 2).sum(0); n += len(z) + mu = S / n; sd = np.sqrt(np.clip(SS / n - mu ** 2, 1e-12, None)) + 1e-6 + np.savez(f"{OUT}/{GRAIN}_mu.npz", mu=mu.astype(np.float32), sd=sd.astype(np.float32), n=n) + return {"grain": GRAIN, "mu_cells": n} + + +def score_shard(genes): + from ops_model.models.attention.diffex.classifier.config import slugify + os.makedirs(PART, exist_ok=True) + cz, cidx = _cz() + if STD == "global": + m = np.load(f"{CV}/centroid_bagsweep_v5new/{GRAIN}_mu.npz"); mu, sd = m["mu"], m["sd"] + by = {B: {} for B in BAGS} + for g in genes: + f = f"{CACHE}/{g}.npz" + if not os.path.exists(f) or slugify(g) not in cidx: + continue + d = np.load(f, allow_pickle=True); al = [float(a) for a in d["alphas"]]; ti = cidx[slugify(g)] + a0 = int(np.argmin(np.abs(np.asarray(al, float)))); gv0 = np.asarray(d["gen"][a0], np.float32) + for ai, a in enumerate(al): + gv = d["gen"][ai] + if gv is None or not len(gv): + continue + gv = np.asarray(gv, np.float32) + for B in BAGS: + if STD == "perbag": # standardize on this gene's own α=0 frames (first-B), like score_embs_v5 + mu = gv0[:B].mean(0); sd = gv0[:B].std(0) + 1e-6 + gz = (gv[:B] - mu) / sd; gz = gz / (np.linalg.norm(gz, axis=1, keepdims=True) + 1e-9) + order = np.argsort(-(gz @ cz.T), axis=1); rk = np.where(order == ti)[1] + 1 + rec = by[B].setdefault(a, {"top1": {}, "top5": {}, "map": {}}) + rec["top1"][g] = float(np.mean(order[:, 0] == ti)) + rec["top5"][g] = float(np.mean([ti in r[:5] for r in order])) + rec["map"][g] = float(np.mean(1.0 / rk)) + json.dump({"by_bag": by}, open(f"{PART}/{genes[0]}.json", "w")) + return {"grain": GRAIN, "n": len(genes)} + + +def merge(): + cz, cidx = _cz() + by = {str(B): {} for B in BAGS} + for p in glob.glob(f"{PART}/*.json"): + d = json.load(open(p)) + for B, ba in d["by_bag"].items(): + for a, rec in ba.items(): + dst = by[str(B)].setdefault(a, {"top1": {}, "top5": {}, "map": {}}) + for k in ("top1", "top5", "map"): + dst[k].update(rec[k]) + json.dump({"bags": BAGS, "by_bag": by, "n_classes": len(cidx)}, open(f"{OUT}/{GRAIN}_bagsweep.json", "w")) + n = len(next(iter(by["400"].values()))["map"]) if by["400"] else 0 + return {"grain": GRAIN, "scored": n} + + +def main(): + from ops_utils.hpc.slurm_batch_utils import submit_parallel_jobs + if STD == "global": + print(compute_mu()) # inline: fast (α=0 of 150 genes); perbag needs no global mu + os.makedirs(OUT, exist_ok=True) + genes = sorted(os.path.basename(f)[:-4] for f in glob.glob(f"{CACHE}/*.npz")) + ch = 40; shards = [genes[i:i + ch] for i in range(0, len(genes), ch)] + jobs = [{"name": f"cbs_{GRAIN}_{i}", "func": score_shard, "kwargs": {"genes": s}} for i, s in enumerate(shards)] + print(f"[centroid-bagsweep] {GRAIN}: {len(genes)} genes → {len(jobs)} score shards") + submit_parallel_jobs(jobs, experiment=f"cbs_{GRAIN}", + slurm_params={"slurm_partition": "preempted", "cpus_per_task": 8, "mem_gb": 32, "timeout_min": 60}, + log_dir=f"cbs_{GRAIN}", wait_for_completion=False) + + +if __name__ == "__main__": + if sys.argv[1:2] == ["merge"]: + print(merge()) + else: + main() diff --git a/src/ops_model/models/attention/diffex/figures/gen_validation/centroid_halves.py b/src/ops_model/models/attention/diffex/figures/gen_validation/centroid_halves.py new file mode 100644 index 0000000..92d73a2 --- /dev/null +++ b/src/ops_model/models/attention/diffex/figures/gen_validation/centroid_halves.py @@ -0,0 +1,82 @@ +"""First-200 vs second-200 anchor comparison. Same directions/traversals; only the anchor cells differ between +the two 200-cell halves. Pooled bag-level centroid recovery of cells [0:200] vs [200:400] per class/α → does the +second-200 set recover worse (explaining the bag=400 dilution)? Reuses the mu.npz + centroids. Sharded → merge.""" +import os, glob, json, sys +import numpy as np + +CV = "/hpc/projects/icd.fast.ops/analysis/figure4_traversals" +GRAIN = os.environ.get("CBS_GRAIN", "geneKO") +CACHE = f"{CV}/gen_real_map_cache_v5new/{GRAIN}" +CENTD = f"{CV}/gen_real_centroid_v5new" +MU = f"{CV}/centroid_bagsweep_v5new/{GRAIN}_mu.npz" +OUT = f"{CV}/centroid_halves_v5new" +PART = f"{OUT}/{GRAIN}_parts" + + +def _cz(): + from ops_model.models.attention.diffex.classifier.config import slugify + d = np.load(f"{CENTD}/{GRAIN}_centroids.npz", allow_pickle=True) + names = list(d["names"]); cidx = {slugify(str(c)): i for i, c in enumerate(names)} + cz = (d["cents"] - d["mu"]) / d["sd"]; cz = cz / (np.linalg.norm(cz, axis=1, keepdims=True) + 1e-9) + return cz, cidx + + +def _top1(vecs, cz, ti): + if not len(vecs): + return None + m = vecs.mean(0); m = m / (np.linalg.norm(m) + 1e-9) + return float(np.argmax(m @ cz.T) == ti) + + +def score_shard(genes): + from ops_model.models.attention.diffex.classifier.config import slugify + os.makedirs(PART, exist_ok=True) + cz, cidx = _cz(); mg = np.load(MU); mu_g, sd_g = mg["mu"], mg["sd"] + first, second = {}, {} + for g in genes: + f = f"{CACHE}/{g}.npz" + if not os.path.exists(f) or slugify(g) not in cidx: + continue + d = np.load(f, allow_pickle=True); al = [float(a) for a in d["alphas"]]; ti = cidx[slugify(g)] + for ai, a in enumerate(al): + gv = d["gen"][ai] + if gv is None or len(gv) < 400: + continue + z = (np.asarray(gv, np.float32) - mu_g) / sd_g; z = z / (np.linalg.norm(z, axis=1, keepdims=True) + 1e-9) + f1 = _top1(z[:200], cz, ti); f2 = _top1(z[200:400], cz, ti) + if f1 is not None: + first.setdefault(a, {})[g] = f1 + if f2 is not None: + second.setdefault(a, {})[g] = f2 + json.dump({"first": first, "second": second}, open(f"{PART}/{genes[0]}.json", "w")) + return {"grain": GRAIN, "n": len(genes)} + + +def merge(): + first, second = {}, {} + for p in glob.glob(f"{PART}/*.json"): + d = json.load(open(p)) + for a, r in d["first"].items(): + first.setdefault(a, {}).update(r) + for a, r in d["second"].items(): + second.setdefault(a, {}).update(r) + json.dump({"first": first, "second": second}, open(f"{OUT}/{GRAIN}_halves.json", "w")) + return {"grain": GRAIN, "n": len(next(iter(first.values()))) if first else 0} + + +def main(): + from ops_utils.hpc.slurm_batch_utils import submit_parallel_jobs + genes = sorted(os.path.basename(f)[:-4] for f in glob.glob(f"{CACHE}/*.npz")) + ch = 40; shards = [genes[i:i + ch] for i in range(0, len(genes), ch)] + jobs = [{"name": f"halves_{GRAIN}_{i}", "func": score_shard, "kwargs": {"genes": s}} for i, s in enumerate(shards)] + print(f"[halves] {GRAIN}: {len(genes)} genes → {len(jobs)} shards") + submit_parallel_jobs(jobs, experiment=f"halves_{GRAIN}", + slurm_params={"slurm_partition": "preempted", "cpus_per_task": 8, "mem_gb": 32, "timeout_min": 60}, + log_dir=f"halves_{GRAIN}", wait_for_completion=False) + + +if __name__ == "__main__": + if sys.argv[1:2] == ["merge"]: + print(merge()) + else: + main() diff --git a/src/ops_model/models/attention/diffex/figures/gen_validation/centroid_pooled_bagsweep.py b/src/ops_model/models/attention/diffex/figures/gen_validation/centroid_pooled_bagsweep.py new file mode 100644 index 0000000..1772cc0 --- /dev/null +++ b/src/ops_model/models/attention/diffex/figures/gen_validation/centroid_pooled_bagsweep.py @@ -0,0 +1,109 @@ +"""Bag-LEVEL centroid recovery on the new multibag v5 cache. For each class/α/bag B: pool the first-B generated +cells → their standardized mean (the generated class centroid) → is the nearest real centroid the true class? +Also a bootstrapped per-bag REAL ceiling: sample B real cells (with replacement from the ~30 cached) → mean → +nearest centroid, n_boot times. Both rise with bag (better centroid estimate), giving a proper per-bag reference. +Reuses gen_real_map_cache_v5new + the mu.npz from centroid_bagsweep. Sharded → merge. + + CBS_GRAIN=geneKO python centroid_pooled_bagsweep.py # mu (reuse) + score shards + CBS_GRAIN=geneKO python centroid_pooled_bagsweep.py merge +""" +import os, glob, json, sys +import numpy as np + +CV = "/hpc/projects/icd.fast.ops/analysis/figure4_traversals" +GRAIN = os.environ.get("CBS_GRAIN", "geneKO") +CACHE = f"{CV}/gen_real_map_cache_v5new/{GRAIN}" +CENTD = f"{CV}/gen_real_centroid_v5new" +MU = f"{CV}/centroid_bagsweep_v5new/{GRAIN}_mu.npz" +STD = os.environ.get("CBS_STD", "global") # global (panel α=0 mu) | perbag (per-gene α=0, matches score_embs_v5) +OUT = f"{CV}/centroid_pooled_bagsweep_v5new" + ("_perbag" if STD == "perbag" else "") +PART = f"{OUT}/{GRAIN}_parts" +BAGS = [20, 50, 100, 200, 400] +NBOOT = 25 + + +def _cz(): + from ops_model.models.attention.diffex.classifier.config import slugify + d = np.load(f"{CENTD}/{GRAIN}_centroids.npz", allow_pickle=True) + names = list(d["names"]); cidx = {slugify(str(c)): i for i, c in enumerate(names)} + cz = (d["cents"] - d["mu"]) / d["sd"]; cz = cz / (np.linalg.norm(cz, axis=1, keepdims=True) + 1e-9) + return cz, cidx, d["mu"], d["sd"] + + +def _pooled(vecs, cz, ti): + """mean of standardized cells → normalize → nearest-centroid top1 + 1/rank (one class-centroid).""" + m = vecs.mean(0); m = m / (np.linalg.norm(m) + 1e-9) + order = np.argsort(-(m @ cz.T)); rk = int(np.where(order == ti)[0][0]) + 1 + return float(order[0] == ti), float(1.0 / rk) + + +def score_shard(genes): + from ops_model.models.attention.diffex.classifier.config import slugify + os.makedirs(PART, exist_ok=True) + cz, cidx, mu_r, sd_r = _cz() + if STD == "global": + mg = np.load(MU); mu_g, sd_g = mg["mu"], mg["sd"] + rng = np.random.default_rng(0) + gen = {B: {} for B in BAGS}; real = {B: {} for B in BAGS} + for g in genes: + f = f"{CACHE}/{g}.npz" + if not os.path.exists(f) or slugify(g) not in cidx: + continue + d = np.load(f, allow_pickle=True); al = [float(a) for a in d["alphas"]]; ti = cidx[slugify(g)] + a0 = int(np.argmin(np.abs(np.asarray(al, float)))); gv0 = np.asarray(d["gen"][a0], np.float32) + # generated: pooled bag centroid per α per bag + for ai, a in enumerate(al): + gv = d["gen"][ai] + if gv is None or not len(gv): + continue + gv = np.asarray(gv, np.float32) + for B in BAGS: + if STD == "perbag": # per-gene α=0 (first-B), matches score_embs_v5 + mu_g = gv0[:B].mean(0); sd_g = gv0[:B].std(0) + 1e-6 + zb = (gv[:B] - mu_g) / sd_g; zb = zb / (np.linalg.norm(zb, axis=1, keepdims=True) + 1e-9) + t1, mp = _pooled(zb, cz, ti) + r = gen[B].setdefault(a, {"top1": {}, "map": {}}); r["top1"][g] = t1; r["map"][g] = mp + # real ceiling: bootstrap B real cells → mean → nearest (per bag) + rr = d["real"] + if rr is not None and len(rr): + rz = (np.asarray(rr, np.float32) - mu_r) / sd_r; rz = rz / (np.linalg.norm(rz, axis=1, keepdims=True) + 1e-9) + for B in BAGS: + t1s = [] + for _ in range(NBOOT): + idx = rng.integers(0, len(rz), size=B) + t1s.append(_pooled(rz[idx], cz, ti)[0]) + real[B][g] = float(np.mean(t1s)) + json.dump({"gen": gen, "real": real}, open(f"{PART}/{genes[0]}.json", "w")) + return {"grain": GRAIN, "n": len(genes)} + + +def merge(): + gen = {str(B): {} for B in BAGS}; real = {str(B): {} for B in BAGS} + for p in glob.glob(f"{PART}/*.json"): + d = json.load(open(p)) + for B in BAGS: + for a, rec in d["gen"].get(str(B), {}).items(): + dst = gen[str(B)].setdefault(a, {"top1": {}, "map": {}}) + dst["top1"].update(rec["top1"]); dst["map"].update(rec["map"]) + real[str(B)].update(d["real"].get(str(B), {})) + json.dump({"bags": BAGS, "gen": gen, "real": real}, open(f"{OUT}/{GRAIN}_pooled.json", "w")) + return {"grain": GRAIN, "gen_classes": len(next(iter(gen['400'].values()))['map']) if gen['400'] else 0, + "real_classes": len(real['400'])} + + +def main(): + from ops_utils.hpc.slurm_batch_utils import submit_parallel_jobs + genes = sorted(os.path.basename(f)[:-4] for f in glob.glob(f"{CACHE}/*.npz")) + ch = 40; shards = [genes[i:i + ch] for i in range(0, len(genes), ch)] + jobs = [{"name": f"cpbs_{GRAIN}_{i}", "func": score_shard, "kwargs": {"genes": s}} for i, s in enumerate(shards)] + print(f"[pooled-centroid] {GRAIN}: {len(genes)} genes → {len(jobs)} shards") + submit_parallel_jobs(jobs, experiment=f"cpbs_{GRAIN}", + slurm_params={"slurm_partition": "preempted", "cpus_per_task": 8, "mem_gb": 32, "timeout_min": 60}, + log_dir=f"cpbs_{GRAIN}", wait_for_completion=False) + + +if __name__ == "__main__": + if sys.argv[1:2] == ["merge"]: + print(merge()) + else: + main() diff --git a/src/ops_model/models/attention/diffex/figures/gen_validation/control_halves_zscore.py b/src/ops_model/models/attention/diffex/figures/gen_validation/control_halves_zscore.py new file mode 100644 index 0000000..75e5921 --- /dev/null +++ b/src/ops_model/models/attention/diffex/figures/gen_validation/control_halves_zscore.py @@ -0,0 +1,106 @@ +"""CONTROL: is the first-200 vs second-200 centroid-recovery gap (76% vs 13%) a per-bag standardization +artifact? The SetTransformer self-standardizes each bag on its OWN α=0 generated frames (cancels any +half-specific CellDINO offset) and sees NO gap; the pooled-centroid metric standardizes gen cells against a +GLOBAL/panel α=0 mean, so a half-specific offset survives. Here we recompute each half's centroid top-1 two +ways per gene: (global) shared {grain}_mu.npz vs (perbag) each half's own α=0 mean/std. If the gap collapses +under perbag, the 76/13 is definitively a standardization artifact of the centroid metric, not phenotype loss. +Real centroids stay real-standardized (unchanged) in both. Sharded → merge. CPU.""" +import os, glob, json, sys +import numpy as np + +CV = "/hpc/projects/icd.fast.ops/analysis/figure4_traversals" +GRAIN = os.environ.get("CBS_GRAIN", "geneKO") +CACHE = f"{CV}/gen_real_map_cache_v5new/{GRAIN}" +CENTD = f"{CV}/gen_real_centroid_v5new" +MU = f"{CV}/centroid_bagsweep_v5new/{GRAIN}_mu.npz" +OUT = f"{CV}/control_halves_v5new" +PART = f"{OUT}/{GRAIN}_parts" +EPS = 1e-6 + + +def _cz(): + from ops_model.models.attention.diffex.classifier.config import slugify + d = np.load(f"{CENTD}/{GRAIN}_centroids.npz", allow_pickle=True) + names = list(d["names"]); cidx = {slugify(str(c)): i for i, c in enumerate(names)} + cz = (d["cents"] - d["mu"]) / d["sd"]; cz = cz / (np.linalg.norm(cz, axis=1, keepdims=True) + 1e-9) + return cz, cidx + + +def _t1(vecs, cz, ti): + if not len(vecs): + return None + v = vecs / (np.linalg.norm(vecs, axis=1, keepdims=True) + 1e-9) + m = v.mean(0); m = m / (np.linalg.norm(m) + 1e-9) + return float(np.argmax(m @ cz.T) == ti) + + +def score_shard(genes): + from ops_model.models.attention.diffex.classifier.config import slugify + os.makedirs(PART, exist_ok=True) + cz, cidx = _cz(); mg = np.load(MU); mu_g, sd_g = mg["mu"], mg["sd"] + res = {s: {h: {} for h in ("first", "second")} for s in ("global", "perbag")} + for g in genes: + f = f"{CACHE}/{g}.npz" + if not os.path.exists(f) or slugify(g) not in cidx: + continue + d = np.load(f, allow_pickle=True); al = [float(a) for a in d["alphas"]]; ti = cidx[slugify(g)] + a0 = int(np.argmin(np.abs(al))); a0v = d["gen"][a0] + if a0v is None or len(a0v) < 400: + continue + a0v = np.asarray(a0v, np.float32) + mu1 = a0v[:200].mean(0); sd1 = a0v[:200].std(0) + EPS + mu2 = a0v[200:400].mean(0); sd2 = a0v[200:400].std(0) + EPS + for ai, a in enumerate(al): + gv = d["gen"][ai] + if gv is None or len(gv) < 400: + continue + gv = np.asarray(gv, np.float32); h1, h2 = gv[:200], gv[200:400] + for s, z1, z2 in [("global", (h1 - mu_g) / sd_g, (h2 - mu_g) / sd_g), + ("perbag", (h1 - mu1) / sd1, (h2 - mu2) / sd2)]: + t1 = _t1(z1, cz, ti); t2 = _t1(z2, cz, ti) + if t1 is not None: + res[s]["first"].setdefault(a, {})[g] = t1 + if t2 is not None: + res[s]["second"].setdefault(a, {})[g] = t2 + json.dump(res, open(f"{PART}/{genes[0]}.json", "w")) + return {"grain": GRAIN, "n": len(genes)} + + +def merge(): + res = {s: {h: {} for h in ("first", "second")} for s in ("global", "perbag")} + for p in glob.glob(f"{PART}/*.json"): + d = json.load(open(p)) + for s in ("global", "perbag"): + for h in ("first", "second"): + for a, r in d[s][h].items(): + res[s][h].setdefault(a, {}).update(r) + json.dump(res, open(f"{OUT}/{GRAIN}_control.json", "w")) + # peak-α summary + out = {} + for s in ("global", "perbag"): + agg = {} + for h in ("first", "second"): + byA = {float(a): np.mean(list(r.values())) for a, r in res[s][h].items() if r} + k = max(byA, key=byA.get) if byA else None + agg[h] = {"peak_alpha": k, "top1": (byA[k] if k is not None else None), "n": len(next(iter(res[s][h].values()))) if res[s][h] else 0} + out[s] = agg + print(json.dumps(out, indent=2)) + return out + + +def main(): + from ops_utils.hpc.slurm_batch_utils import submit_parallel_jobs + genes = sorted(os.path.basename(f)[:-4] for f in glob.glob(f"{CACHE}/*.npz")) + ch = 40; shards = [genes[i:i + ch] for i in range(0, len(genes), ch)] + jobs = [{"name": f"ctrl_{GRAIN}_{i}", "func": score_shard, "kwargs": {"genes": s}} for i, s in enumerate(shards)] + print(f"[control] {GRAIN}: {len(genes)} genes → {len(jobs)} shards") + submit_parallel_jobs(jobs, experiment=f"ctrl_{GRAIN}", + slurm_params={"slurm_partition": "preempted", "cpus_per_task": 8, "mem_gb": 32, "timeout_min": 60}, + log_dir=f"ctrl_{GRAIN}", wait_for_completion=False) + + +if __name__ == "__main__": + if sys.argv[1:2] == ["merge"]: + print(merge()) + else: + main() diff --git a/src/ops_model/models/attention/diffex/figures/gen_validation/embcheck.py b/src/ops_model/models/attention/diffex/figures/gen_validation/embcheck.py new file mode 100644 index 0000000..825e412 --- /dev/null +++ b/src/ops_model/models/attention/diffex/figures/gen_validation/embcheck.py @@ -0,0 +1,36 @@ +"""Verify the valid200 cache embeddings equal the original embed_crops path on the SAME saved frames. +If cosine≈1 the cache build is faithful (frames genuinely differ); if not, the cache build is the bug. +""" +import numpy as np + +CV = "/hpc/projects/icd.fast.ops/analysis/figure4_traversals" +B = "/hpc/projects/icd.fast.ops/models/diffex" + + +def check(gene="AACS", ai=6): + from ops_model.models.attention.diffex.viewer.score_generated import _emb_frames + from ops_model.models.attention.diffex.classifier.celldino_features import embed_crops + from ops_model.models.attention.diffex.directions.config import DirConfig + trav = f"{B}/viewer_assets_valid200/phase/geneKO/{gene}" + cfg = DirConfig(grain="geneKO", target=gene, device="cuda") + embB = np.asarray(_emb_frames(cfg, trav, ai, embed_crops), np.float32) # original path + d = np.load(f"{CV}/gen_real_map_cache_valid200/geneKO/{gene}.npz", allow_pickle=True) + embA = np.asarray(d["gen"][ai], np.float32) # my cache + n = min(len(embA), len(embB)); A, Bm = embA[:n], embB[:n] + cos = (A * Bm).sum(1) / ((np.linalg.norm(A, 1 if False else None, axis=1) + 1e-9) * (np.linalg.norm(Bm, axis=1) + 1e-9)) + return {"gene": gene, "ai": ai, "nA": len(embA), "nB": len(embB), + "cos_mean": float(cos.mean()), "cos_min": float(cos.min()), + "max_abs_diff": float(np.abs(A - Bm).max())} + + +def main(): + from ops_utils.hpc.slurm_batch_utils import submit_parallel_jobs + submit_parallel_jobs(jobs_to_submit=[{"name": "embcheck", "func": check, "kwargs": {}}], + experiment="embcheck", + slurm_params={"slurm_partition": "preempted", "gpus_per_node": 1, "cpus_per_task": 8, + "mem_gb": 48, "timeout_min": 30, "slurm_constraint": "[a40|a6000|l40s]"}, + log_dir="embcheck", wait_for_completion=False) + + +if __name__ == "__main__": + main() diff --git a/src/ops_model/models/attention/diffex/figures/gen_validation/embedding_diagnostics.py b/src/ops_model/models/attention/diffex/figures/gen_validation/embedding_diagnostics.py new file mode 100644 index 0000000..019edfc --- /dev/null +++ b/src/ops_model/models/attention/diffex/figures/gen_validation/embedding_diagnostics.py @@ -0,0 +1,133 @@ +"""Diagnostic plots comparing how generated centroids get PLACED in the real phase embedding, plus the +real-NTC vs inverse-α=0 domain-gap. Reads only already-computed artifacts; writes to figure4_embedding/diagnostics/. + +Approaches compared (all at α≈3): + - kNN landmark (self-std) : cosine-NN interpolation onto the real manifold (gen_phate_passthrough default) + - kNN landmark (real-std) : same, real-population standardization + - UMAP .transform (real-std): proper out-of-sample projection (gen_embed_refit) + - PHATE joint (real-std) : joint fit_transform (gen_embed_refit) +Metric: 2-D rank-to-true = rank of a class's true real gene among all 1052 genes by 2-D distance to its +generated dot (class-blind). Low = generated lands on its own gene; ~random (≈480) = lands in a blob. +""" +import os +import numpy as np +import pandas as pd + +B = "/hpc/projects/icd.fast.ops/analysis/figure4_embedding" +GP = f"{B}/gen_passthrough" +RF = f"{B}/gen_passthrough_refit/coords.npz" +GAP = f"{B}/ntc_inverse_gap/emb_5xUPRE.npz" +OUT = f"{B}/diagnostics" +AI = 14 # α=3 parquet + + +def _ranks(gen_xy, real_xy, true_row): + """2-D rank of each class's true gene among all real genes (Euclidean).""" + from scipy.spatial.distance import cdist + D = cdist(gen_xy, real_xy) + order = np.argsort(D, axis=1) + return np.array([int(np.where(order[k] == true_row[k])[0][0]) + 1 for k in range(len(gen_xy))]) + + +def _landmark(parq): + """From a proj parquet: gen 2-D (gu), real-gene 2-D set (unique ru per gene), and true-row index.""" + df = pd.read_parquet(parq) + real_xy = df[["ru0", "ru1"]].values # one row per gene = its true coord + gen_xy = df[["gu0", "gu1"]].values + return gen_xy, real_xy, np.arange(len(df)), list(df["gene"]) + + +def _refit(layout): + d = np.load(RF, allow_pickle=True) + genes = list(d["genes"]); real_names = list(d["real_names"]); ridx = {n: i for i, n in enumerate(real_names)} + R = d[f"real_{layout}"]; G = d[f"gen_{layout}"] + true_row = np.array([ridx[g] for g in genes]) + return G, R, true_row, genes + + +def plot(): + import matplotlib + matplotlib.use("Agg") + import matplotlib.pyplot as plt + from matplotlib.lines import Line2D + plt.rcParams["pdf.fonttype"] = 42 + os.makedirs(OUT, exist_ok=True) + + approaches = [] + gx, rx, tr, _ = _landmark(f"{GP}/proj_a{AI}.parquet"); approaches.append(("kNN landmark (self-std)", gx, rx, tr)) + gx, rx, tr, _ = _landmark(f"{GP}/proj_real_a{AI}.parquet"); approaches.append(("kNN landmark (real-std)", gx, rx, tr)) + gx, rx, tr, _ = _refit("umap"); approaches.append(("UMAP .transform (real-std)", gx, rx, tr)) + gx, rx, tr, _ = _refit("phate"); approaches.append(("PHATE joint (real-std)", gx, rx, tr)) + + # ---- Fig A: placement panels (real grey + generated colored by 2-D rank-to-true + connector to true gene) ---- + fig, axes = plt.subplots(1, 4, figsize=(22, 6)) + stats = [] + for ax, (name, G, R, tr) in zip(axes, approaches): + rk = _ranks(G, R, tr) + stats.append((name, rk)) + ax.scatter(R[:, 0], R[:, 1], s=6, c="#dddddd", lw=0, zorder=1) + for k in range(len(G)): + ax.plot([G[k, 0], R[tr[k], 0]], [G[k, 1], R[tr[k], 1]], "-", color="#888", lw=0.3, alpha=0.08, zorder=2) + sc = ax.scatter(G[:, 0], G[:, 1], s=10, c=np.log10(rk), cmap="viridis_r", lw=0, zorder=3) + ax.set_title(f"{name}\nmedian rank {np.median(rk):.0f} · top-20 {(rk<=20).mean():.0%}", fontsize=11) + ax.set_xticks([]); ax.set_yticks([]) + for s in ax.spines.values(): + s.set_visible(False) + cb = fig.colorbar(sc, ax=axes, fraction=0.012, pad=0.01); cb.set_label("log10 rank-to-true (lower = on its own gene)") + fig.suptitle("Where generated centroids (α≈3) land, by placement method — grey = real genes, lines = gap to true gene", + fontweight="bold", fontsize=13) + for e in ("png", "svg"): + fig.savefig(f"{OUT}/placement_methods.{e}", dpi=140, bbox_inches="tight") + plt.close(fig); print("saved placement_methods") + + # ---- Fig B: rank-to-true ECDF (cumulative % of classes within rank X) ---- + fig, ax = plt.subplots(figsize=(7.5, 5.5)) + for name, rk in stats: + xs = np.sort(rk); ys = np.arange(1, len(xs) + 1) / len(xs) * 100 + ax.plot(xs, ys, lw=2.2, label=f"{name} (med {np.median(rk):.0f})") + ax.axvline(20, color="#c0392b", ls=":", lw=1.2, label="rank-20") + ax.set_xscale("log"); ax.set_xlabel("2-D rank-to-true gene (log)"); ax.set_ylabel("% of classes ≤ rank") + ax.set_title("How close generated dots land to their true gene (2-D)", fontweight="bold") + ax.grid(alpha=.25); ax.legend(fontsize=9, loc="lower right") + for e in ("png", "svg"): + fig.savefig(f"{OUT}/rank_ecdf.{e}", dpi=150, bbox_inches="tight") + plt.close(fig); print("saved rank_ecdf") + + _plot_gap() + + +def _plot_gap(): + import matplotlib + matplotlib.use("Agg") + import matplotlib.pyplot as plt + from numpy.linalg import norm + if not os.path.exists(GAP): + print("no NTC-gap npz; skipping"); return + d = np.load(GAP, allow_pickle=True) + R, G = d["R"].astype(np.float64), d["G"].astype(np.float64) + ch = str(d["channel"]) if "channel" in d else "5xUPRE" + cos = np.array([float(R[i] @ G[i] / (norm(R[i]) * norm(G[i]))) for i in range(len(R))]) + # PCA of pooled R+G (2-D) to show overlap + X = np.vstack([R, G]); Xc = X - X.mean(0) + U, S, Vt = np.linalg.svd(Xc, full_matrices=False); P = Xc @ Vt[:2].T + Pr, Pg = P[:len(R)], P[len(R):] + + fig, ax = plt.subplots(1, 2, figsize=(13, 5)) + ax[0].hist(cos, bins=20, color="#7fbf9a", edgecolor="k") + ax[0].axvline(cos.mean(), color="k", lw=2, label=f"mean {cos.mean():.3f}") + ax[0].set_xlabel("per-cell cosine(real, inverse-α=0)"); ax[0].set_ylabel("count") + ax[0].set_title(f"Real NTC vs inverse-α=0 — same-cell fidelity ({ch})", fontweight="bold"); ax[0].legend() + for i in range(len(Pr)): + ax[1].plot([Pr[i, 0], Pg[i, 0]], [Pr[i, 1], Pg[i, 1]], "-", color="#bbb", lw=0.6, zorder=1) + ax[1].scatter(Pr[:, 0], Pr[:, 1], s=45, c="#8fa9c9", edgecolor="k", lw=.4, label="real NTC", zorder=3) + ax[1].scatter(Pg[:, 0], Pg[:, 1], s=45, c="#c0666b", edgecolor="k", lw=.4, label="inverse α=0", zorder=3) + ax[1].set_title("CellDINO PCA (paired, same cell)", fontweight="bold"); ax[1].legend() + ax[1].set_xticks([]); ax[1].set_yticks([]) + fig.tight_layout() + for e in ("png", "svg"): + fig.savefig(f"{OUT}/ntc_inverse_gap.{e}", dpi=150, bbox_inches="tight") + plt.close(fig); print("saved ntc_inverse_gap") + + +if __name__ == "__main__": + plot() diff --git a/src/ops_model/models/attention/diffex/figures/gen_validation/figure4_bagsize_reachfrac.py b/src/ops_model/models/attention/diffex/figures/gen_validation/figure4_bagsize_reachfrac.py new file mode 100644 index 0000000..4e933c2 --- /dev/null +++ b/src/ops_model/models/attention/diffex/figures/gen_validation/figure4_bagsize_reachfrac.py @@ -0,0 +1,74 @@ +"""What fraction of *distinguishable* classes have a GENERATED traversal that reaches real-cell accuracy, +at each bag size? At each bag we keep only classes whose REAL top1_acc is meaningful (>= REAL_THR) — i.e. +there is a real signal to reach — so the denominator n differs per bag (real accuracy climbs with bag size). +A class "reaches" if generated top1_acc is within a lenient margin of real (gen >= real - MARGIN). +Plot % reaching vs bag size, one line per grain; each point annotated with its per-bag n. pdf.fonttype 42.""" +import glob +import json +import os + +import numpy as np +import matplotlib +matplotlib.use("Agg") +import matplotlib.pyplot as plt + +plt.rcParams["pdf.fonttype"] = 42 +BAGT = "/hpc/projects/icd.fast.ops/models/diffex/viewer_assets_v5_bagtest" +OUT = "/hpc/projects/icd.fast.ops/analysis/figure4_traversals" +BAGS = [20, 50, 100, 200, 500] +MARGIN = 0.1 # lenient: generated within 0.1 below real counts as reaching +REAL_THR = 0.5 # keep only classes whose real cells are distinguishable at that bag + + +def _load(grain): + """[(real[bag], gen[bag]) per class] for a grain, aligned to BAGS (nan where missing).""" + out = [] + for f in sorted(glob.glob(f"{BAGT}/_bagexp_*.json")): + d = json.load(open(f)) + if d.get("grain") != grain: + continue + re = {int(k): v for k, v in d["real_expectation"].items()} + real = np.array([re.get(b, np.nan) for b in BAGS]) + gen = np.array([d["bag"].get(str(b), {}).get("top1_acc", np.nan) for b in BAGS]) + out.append((real, gen)) + return out + + +def main(): + fig, ax = plt.subplots(figsize=(7.5, 5.2)) + for grain, col in [("geneKO", "#1f77b4"), ("complex", "#d62728")]: + cls = _load(grain) + if not cls: + continue + frac, ns = [], [] + for i, b in enumerate(BAGS): + reach, tot = 0, 0 + for real, gen in cls: + if not (np.isfinite(real[i]) and np.isfinite(gen[i])) or real[i] < REAL_THR: + continue # drop classes with no real signal at this bag + tot += 1 + reach += gen[i] >= real[i] - MARGIN + frac.append(100 * reach / tot if tot else np.nan); ns.append(tot) + ax.plot(BAGS, frac, "-o", color=col, lw=2.5, ms=7, label=f"{grain}") + for x, y, n in zip(BAGS, frac, ns): + ax.annotate(f"{y:.0f}%\nn={n}", (x, y), textcoords="offset points", xytext=(0, 9), + ha="center", fontsize=7.5, color=col) + ax.set_xscale("log"); ax.set_xticks(BAGS); ax.set_xticklabels(BAGS) + ax.set_xlabel("bag size (# cells) — n per point = classes with real top1_acc ≥ %.1f at that bag" % REAL_THR) + ax.set_ylabel(f"% of distinguishable classes reaching real accuracy\n(generated ≥ real − {MARGIN})") + ax.set_ylim(0, 105) + ax.set_title("Among perturbations distinguishable by real cells at each bag size,\n" + "what fraction do generated traversals recapitulate?") + ax.legend(fontsize=10, loc="upper right", framealpha=0.9) + for s in ("top", "right"): + ax.spines[s].set_visible(False) + fig.tight_layout() + os.makedirs(OUT, exist_ok=True) + for ext in ("png", "svg"): + fig.savefig(f"{OUT}/v5_bagsize_reachfrac.{ext}", dpi=200, bbox_inches="tight", facecolor="white") + plt.close(fig) + print(f"saved {OUT}/v5_bagsize_reachfrac.png / .svg") + + +if __name__ == "__main__": + main() diff --git a/src/ops_model/models/attention/diffex/figures/gen_validation/figure4_v5_accuracy_summary.py b/src/ops_model/models/attention/diffex/figures/gen_validation/figure4_v5_accuracy_summary.py new file mode 100644 index 0000000..b7eb101 --- /dev/null +++ b/src/ops_model/models/attention/diffex/figures/gen_validation/figure4_v5_accuracy_summary.py @@ -0,0 +1,140 @@ +"""Figure 4 summary: v5 SetTransformer accuracy of GENERATED traversals vs α, aggregated across the +1K geneKO set and the EBI-complex set. For each α: mean P(target) + IQR (Q1–Q3) band over all traversals. +Reads scores_v5.json (per-α P(target)) from viewer_assets_v5. White bg, pdf.fonttype 42.""" +import glob +import json +import os + +import numpy as np +import matplotlib +matplotlib.use("Agg") +import matplotlib.pyplot as plt + +plt.rcParams["pdf.fonttype"] = 42 +V5 = "/hpc/projects/icd.fast.ops/models/diffex/viewer_assets_v5/phase" +OUT = "/hpc/projects/icd.fast.ops/analysis/figure4_traversals" + + +EVAL = "/hpc/projects/icd.fast.ops/models/alex_lin_attention/v5/phase" + + +def _real_acc20(csv): + """Alex real-cell top1_acc at bag=20, keyed by gene/label (from the eval CSV).""" + import csv as _csv + return {r["gene_name"]: float(r["top1_acc"]) for r in _csv.DictReader(open(f"{EVAL}/{csv}")) if int(r["n_cells"]) == 20} + + +def _geneKO_allow(thr=0.9): + """geneKO dir-names whose REAL-cell top1_acc@bag20 > thr (Alex gene eval).""" + return {g for g, a in _real_acc20("eval_phase_e200_pergene_val.csv").items() if a > thr} + + +def _complex_allow(thr=0.9): + """Complex dir-slugs whose MEAN member-gene real top1_acc@bag20 > thr — grouped by Alex's own + label_name in the ebionly eval CSV (his reported members), then mean of those member accuracies.""" + import csv as _csv + from collections import defaultdict + from ops_model.models.attention.diffex.classifier.config import slugify + by = defaultdict(list) + for r in _csv.DictReader(open(f"{EVAL}/eval_phase_ebionly_e200_pergene_val.csv")): + if int(r["n_cells"]) == 20: + by[r["label_name"]].append(float(r["top1_acc"])) + return {slugify(lbl) for lbl, accs in by.items() if np.mean(accs) > thr} + + +def _real_map(sub): + """{dir-name: real-cell top1_acc @bag20}. geneKO = gene eval; complex = mean of Alex's members (by slug).""" + if sub == "geneKO": + return _real_acc20("eval_phase_e200_pergene_val.csv") + import csv as _csv + from collections import defaultdict + from ops_model.models.attention.diffex.classifier.config import slugify + by = defaultdict(list) + for r in _csv.DictReader(open(f"{EVAL}/eval_phase_ebionly_e200_pergene_val.csv")): + if int(r["n_cells"]) == 20: + by[r["label_name"]].append(float(r["top1_acc"])) + return {slugify(l): float(np.mean(v)) for l, v in by.items()} + + +VALID200 = "/hpc/projects/icd.fast.ops/models/diffex/viewer_assets_valid200/phase" # 200-cell bag (geneKO only) + + +def _collect(sub, allow=None, base=None): + """Stack p_target across NTC-anchored traversals in a subdir → (alphas, matrix, names). allow: dir-name set.""" + base = base or V5 + rows, names, alphas = [], [], None + for f in sorted(glob.glob(f"{base}/{sub}/*/scores_v5.json")): + if "__to__" in f: # NTC-anchored only (skip alt-anchor A→B) + continue + name = os.path.basename(os.path.dirname(f)) + if allow is not None and name not in allow: + continue + d = json.load(open(f)) + alphas = d["alphas"] + rows.append([np.nan if v is None else v for v in d["p_target"]]); names.append(name) + return np.array(alphas, float), np.array(rows, float), names + + +def _render(series, out_stem, title, overlays=None): + fig, ax = plt.subplots(figsize=(7.5, 5)) + for sub, allow, label, col in series: + al, M, names = _collect(sub, allow) + if not len(M): + print(f"no scores for {sub}"); continue + mean = np.nanmean(M, 0); nfin = np.sum(np.isfinite(M), 0) + sem = np.nanstd(M, 0, ddof=1) / np.sqrt(np.maximum(nfin, 1)) # generated SEM per α + ax.plot(al, mean, "-", color=col, lw=2.5, label=f"{label} generated (n={M.shape[0]})", zorder=5) + ax.fill_between(al, mean - sem, mean + sem, color=col, alpha=0.2, lw=0) # mean ± SEM + rm = _real_map(sub); reals = np.array([rm[n] for n in names if n in rm]) # real-cell ceiling @bag20 + if len(reals): + rmean = reals.mean(); rsem = reals.std(ddof=1) / np.sqrt(len(reals)) + ax.axhline(rmean, color=col, ls=":", lw=1.8, alpha=0.95, zorder=6, label=f"{label} real @bag20 ({rmean:.2f})") + ax.axhspan(rmean - rsem, rmean + rsem, color=col, alpha=0.1, lw=0) # real mean ± SEM band + print(f"[{out_stem}] {sub}: n={M.shape[0]} peak mean={np.nanmax(mean):.3f} @α={al[np.nanargmax(mean)]:+g} real mean@bag20={reals.mean():.3f}") + for sub, allow, label, col, base in (overlays or []): # dashed = 200-cell bag (same real ceiling) + al, M, names = _collect(sub, allow, base) + if not len(M): + print(f"no overlay scores for {sub}"); continue + mean = np.nanmean(M, 0); sem = np.nanstd(M, 0, ddof=1) / np.sqrt(np.maximum(np.sum(np.isfinite(M), 0), 1)) + ax.plot(al, mean, "--", color=col, lw=2.5, label=f"{label} (n={M.shape[0]})", zorder=5) + ax.fill_between(al, mean - sem, mean + sem, color=col, alpha=0.15, lw=0) + print(f"[{out_stem}] OVERLAY {sub}: n={M.shape[0]} peak mean={np.nanmax(mean):.3f} @α={al[np.nanargmax(mean)]:+g}") + ax.axvline(0, color="0.6", lw=0.8, ls=":") + ax.set_xlabel("α (traversal strength; ±1 ≈ control→KD gap)") + ax.set_ylabel("P(target class) — v5 SetTransformer") + ax.set_title(title) + ax.set_ylim(-0.02, 1.02) + ax.legend(fontsize=8, loc="upper left", bbox_to_anchor=(1.02, 1), framealpha=0.9) # outside, to the right + for s in ("top", "right"): + ax.spines[s].set_visible(False) + os.makedirs(OUT, exist_ok=True) + for ext in ("png", "svg"): + fig.savefig(f"{OUT}/{out_stem}.{ext}", dpi=200, bbox_inches="tight", facecolor="white") + plt.close(fig) + print(f"saved {OUT}/{out_stem}.png / .svg") + + +def main(): + suf = os.environ.get("SUMMARY_SUFFIX", "") # e.g. "_corrected" → new files, don't overwrite the old plot + ov_full = None if os.environ.get("NO_OVERLAY") else [("geneKO", None, "geneKO 200-cell bag", "#1f77b4", VALID200), + ("complex", None, "complex 200-cell bag", "#d62728", VALID200)] + # full: all geneKO + all complexes + _render([("geneKO", None, "geneKO (1K)", "#1f77b4"), ("complex", None, "EBI complexes", "#d62728")], + "v5_accuracy_vs_alpha_summary" + suf, + "Generated-traversal accuracy vs α\n(generated mean ± SEM; dotted = real-cell mean top1_acc @bag20)", + overlays=ov_full) + # filtered: high real-accuracy classes (real acc>thr @bag20). Dotted line = real-cell ceiling for that set. + for thr, stem in [(0.8, "v5_accuracy_vs_alpha_realacc80" + suf), (0.9, "v5_accuracy_vs_alpha_realacc90" + suf)]: + gk, cx = _geneKO_allow(thr), _complex_allow(thr) + print(f"filtered allowlists (>{thr}): {len(gk)} geneKO, {len(cx)} complex") + ov = None if os.environ.get("NO_OVERLAY") else [("geneKO", gk, f"geneKO 200-cell bag (real acc>{thr})", "#1f77b4", VALID200), + ("complex", cx, f"complex 200-cell bag (real acc>{thr})", "#d62728", VALID200)] + _render([("geneKO", gk, f"geneKO (real acc>{thr} @bag20)", "#1f77b4"), + ("complex", cx, f"EBI complex (mean member real acc>{thr})", "#d62728")], + stem, + f"Generated accuracy vs α — high real-accuracy classes only\n(geneKO real acc>{thr}; complex mean-member real acc>{thr}; @bag20)", + overlays=ov) + + +if __name__ == "__main__": + main() diff --git a/src/ops_model/models/attention/diffex/figures/gen_validation/gen_alpha_embedding.py b/src/ops_model/models/attention/diffex/figures/gen_validation/gen_alpha_embedding.py new file mode 100644 index 0000000..5af4f61 --- /dev/null +++ b/src/ops_model/models/attention/diffex/figures/gen_validation/gen_alpha_embedding.py @@ -0,0 +1,249 @@ +"""Per-α UMAP animation: top-K real cells + generated cells at every α in ONE shared standardized embedding. +Generated cells start piled at the center (α=0, no phenotype) and migrate into their real-class territories as α +rises — the visual analog of the α→distinctiveness/mAP curve. Reuses the gen_real_map_cache embeddings (real + +gen, same embed_crops space); per-domain standardization (real vs population, gen vs α0) cancels the DiffAE offset. + +Real + generated are BOTH colored by class (real faint, gen bright) so you can watch gen-X home in on real-X. +""" +import os, glob +import numpy as np +import matplotlib +matplotlib.use("Agg") +import matplotlib.pyplot as plt +import matplotlib.cm as cm +plt.rcParams["pdf.fonttype"] = 42 + +CACHE = "/hpc/projects/icd.fast.ops/analysis/figure4_traversals/gen_real_map_cache" +CENT = "/hpc/projects/icd.fast.ops/analysis/figure4_traversals/gen_real_centroid" # faithful centroids for the color metric +OUT = "/hpc/projects/icd.fast.ops/analysis/figure4_embedding/gen_alpha_frames" +K = 20 +SEED = 0 + + +def load(grain, keep=None): + real, rlab, gx, gl, alphas = [], [], {}, {}, None + for c in sorted(glob.glob(f"{CACHE}/{grain}/*.npz")): + d = np.load(c, allow_pickle=True); g = str(d["gene"]); alphas = list(d["alphas"]) + if keep is not None and g not in keep: + continue + r = np.asarray(d["real"], np.float32)[:K]; real.append(r); rlab += [g] * len(r) + for ai in range(len(alphas)): + gv = d["gen"][ai] + if gv is None or not len(gv): + continue + gv = np.asarray(gv, np.float32)[:K] + gx.setdefault(ai, []).append(gv); gl.setdefault(ai, []).extend([g] * len(gv)) + return (np.concatenate(real), np.array(rlab), + {ai: np.concatenate(v) for ai, v in gx.items()}, {ai: np.array(gl[ai]) for ai in gl}, alphas) + + +def run(grain, gpu=False, include_real=False, fast=False, only=None, tag=None): + tag = tag or grain + real, rlab, gx, gl, alphas = load(grain, keep=only) + a0 = len(alphas) // 2 + mu_g, sd_g = gx[a0].mean(0), gx[a0].std(0) + 1e-6 # gen baseline = α0 (NTC reconstruction) + ais = sorted(gx) + gz = {ai: (gx[ai] - mu_g) / sd_g for ai in ais} + blocks = [gz[ai] for ai in ais] + if include_real: + mu_r, sd_r = real.mean(0), real.std(0) + 1e-6; rz = (real - mu_r) / sd_r; blocks = [rz] + blocks + X = np.concatenate(blocks).astype(np.float32) + print(f"[{grain}] UMAP fit on {len(X):,} pts (generated only over {len(ais)} α{', + real' if include_real else ''})") + if gpu: + from cuml.manifold import UMAP as U + XY = U(n_neighbors=15, min_dist=0.3, random_state=SEED).fit_transform(X) + else: + import umap + kw = dict(n_neighbors=15, min_dist=0.3, metric="euclidean") + kw["n_jobs"] = -1 if fast else 1 # fast: all cores (drops exact seed reproducibility) + if not fast: + kw["random_state"] = SEED + XY = umap.UMAP(**kw).fit_transform(X) + XY = np.asarray(XY) + off = 0; rxy = None + if include_real: + rxy = XY[:len(real)]; off = len(real) + gxy = {} + for ai in ais: + gxy[ai] = XY[off:off + len(gz[ai])]; off += len(gz[ai]) + os.makedirs(OUT, exist_ok=True) + np.savez(f"{OUT}/{tag}_coords.npz", XY=XY.astype(np.float32), alphas=np.array(alphas), + ais=np.array(ais), splits=np.array([len(gz[ai]) for ai in ais]), + labels=np.concatenate([gl[ai] for ai in ais]), off0=len(real) if include_real else 0) + plot_frames(tag) + + +def plot_frames(tag, highlight=None, out=None, legend=False): + """Per-α sweep on one shared UMAP: shape = cell (base-cell idx), color = class. If `highlight` (a set of + classes) is given, those are drawn bright and all others faded (same size, low opacity) — all cells stay in + frame. legend=True adds an EBI-complex legend of the highlighted set.""" + import colorcet as cc, matplotlib.patches as mp + grain = tag; out = out or tag + d = np.load(f"{OUT}/{tag}_coords.npz", allow_pickle=True) + XY = d["XY"]; alphas = list(d["alphas"]); ais = list(d["ais"]); splits = list(d["splits"]); labels = d["labels"] + off = int(d["off0"]); gxy, gl = {}, {} + for ai, n in zip(ais, splits): + gxy[ai] = XY[off:off + n]; gl[ai] = labels[off - int(d["off0"]):off - int(d["off0"]) + n]; off += n + xlim = (XY[:, 0].min() - 1, XY[:, 0].max() + 1); ylim = (XY[:, 1].min() - 1, XY[:, 1].max() + 1) + K_ = splits[0] // max(len(set(labels)), 1) # cells per class (base-cell index) + MARKERS = ["o", "s", "^", "v", "D", "P", "*", "X", "<", ">", "p", "h", "d"] + from matplotlib.colors import to_rgba + classes = sorted(set(labels)); c2c = {c: to_rgba(cc.glasbey[i % len(cc.glasbey)]) for i, c in enumerate(classes)} + active = sorted(highlight) if highlight else classes # rainbow the colored set (picks, or all) + hp = cm.get_cmap("gist_rainbow")(np.linspace(0, 1, len(active), endpoint=False)) + for i, c in enumerate(active): + c2c[c] = tuple(hp[i]) + + from matplotlib.colors import rgb_to_hsv, hsv_to_rgb + from scipy.spatial import cKDTree + rad = 0.03 * max(xlim[1] - xlim[0], ylim[1] - ylim[0]) + + def _dens(P, groups, r): # per-point # of same-group neighbors within r + dn = np.zeros(len(P)); groups = np.asarray(groups) + for g in np.unique(groups): + idx = np.where(groups == g)[0] + if len(idx) < 2: + continue + dn[idx] = np.asarray(cKDTree(P[idx]).query_ball_point(P[idx], r, return_length=True)) - 1 + return dn + + def panel(ax, ai): # shape=cell; color=class; density → saturation & size + P = gxy[ai]; labs = gl[ai]; cellidx = np.arange(len(P)) % K_; mk_i = cellidx % len(MARKERS) + hi = np.array([c in highlight for c in labs]) if highlight is not None else np.ones(len(labs), bool) + cdn = np.zeros(len(P)); sdn = np.zeros(len(P)) + if hi.any(): + cd = _dens(P[hi], labs[hi], rad); sd = _dens(P[hi], mk_i[hi], rad) # same-class / same-shape overlap + cdn[hi] = cd / (cd.max() + 1e-9); sdn[hi] = sd / (sd.max() + 1e-9) + base = np.array([c2c[c] for c in labs]) + hsv = rgb_to_hsv(base[:, :3]); hsv[:, 1] = hsv[:, 1] * np.clip(0.60 + 0.40 * cdn, 0, 1) # saturation ↑ with same-class overlap (higher floor) + col = np.concatenate([hsv_to_rgb(hsv), base[:, 3:4]], axis=1) + lw = np.clip(3.5 * sdn ** 3, 0, 3.5) # black border ↑ steeply (cubic) with same-shape overlap + if highlight is not None and (~hi).any(): # faded rest (grey backdrop) + for mi, mk in enumerate(MARKERS): + m = (~hi) & (mk_i == mi) + if m.any(): + ax.scatter(P[m, 0], P[m, 1], s=60, c="0.8", marker=mk, lw=0, alpha=.07) + for mi, mk in enumerate(MARKERS): + m = hi & (mk_i == mi) + if m.any(): + ax.scatter(P[m, 0], P[m, 1], s=80, c=col[m], marker=mk, edgecolors="black", linewidths=lw[m], alpha=.9) + from matplotlib.lines import Line2D # small legend: shape = cell number + sh = [Line2D([0], [0], marker=MARKERS[i], color="0.35", lw=0, markersize=6, label=f"cell {i + 1}") for i in range(5)] + sh.append(Line2D([0], [0], marker=r"$\cdots$", color="0.35", lw=0, markersize=8, label="…")) + ax.legend(handles=sh, loc="upper left", fontsize=6, frameon=False, handletextpad=.2, labelspacing=.25, title="shape = cell", title_fontsize=6) + ax.set_xlim(xlim); ax.set_ylim(ylim); ax.set_xticks([]); ax.set_yticks([]) + + handles = [mp.Patch(color=c2c[c], label=c) for c in sorted(highlight)] if (legend and highlight) else None + + for ai in ais: # per-α frame (for the GIF) + fig, ax = plt.subplots(figsize=(11 if handles else 7, 7)); panel(ax, ai) + sub = f" ({len(highlight)} highlighted)" if highlight is not None else "" + fig.suptitle(f"{grain} α = {alphas[ai]:+.1f}{sub}\nshape = cell · color = class", fontsize=13, fontweight="bold") + if handles: + fig.legend(handles=handles, loc="center left", bbox_to_anchor=(0.62, 0.5), fontsize=7, frameon=False) + fig.subplots_adjust(right=0.6) + else: + fig.tight_layout() + fig.savefig(f"{OUT}/{out}_a{ai:02d}.png", dpi=110, bbox_inches="tight"); plt.close(fig) + + keys = [ai for ai in ais if alphas[ai] in (-5, -3, -2, -1, 0, 1, 2, 3, 5)] # symmetric columns + nc = len(keys); fig, axes = plt.subplots(1, nc, figsize=(3.0 * nc, 3.4)) + for ci, ai in enumerate(keys): + panel(axes[ci], ai); axes[ci].set_title(f"α = {alphas[ai]:+.1f}", fontsize=11) + fig.suptitle(f"{grain}: generated cells per α — shape = cell, color = class (one UMAP)", fontsize=14, fontweight="bold") + if handles: + fig.legend(handles=handles, loc="lower center", bbox_to_anchor=(0.5, -0.02), fontsize=7, ncol=min(6, len(handles)), frameon=False) + fig.tight_layout() + for e in ("png", "svg"): + fig.savefig(f"{OUT}/{out}_montage.{e}", dpi=140, bbox_inches="tight") + plt.close(fig) + print(f"saved {out}_montage + {len(ais)} frames") + + +def _tight_few_subset(grain, hi=5.0, member_max=4, tight_top=0.4): + """Complexes that are well-grouped at α=hi (tight UMAP cluster) AND have few member genes.""" + import pandas as pd + d = np.load(f"{OUT}/{grain}_coords.npz", allow_pickle=True) + XY = d["XY"]; alphas = list(d["alphas"]); ais = list(d["ais"]); splits = list(d["splits"]); labels = d["labels"] + off = int(d["off0"]); gxy, gl = {}, {} + for ai, n in zip(ais, splits): + gxy[ai] = XY[off:off + n]; gl[ai] = labels[off - int(d["off0"]):off - int(d["off0"]) + n]; off += n + hi_ai = ais[int(np.argmin(np.abs(np.array(alphas) - hi)))] + pc = pd.read_parquet("/hpc/projects/icd.fast.ops/models/diffex/viewer_assets_v5/_rankings/pma_v5_phase_complex.parquet", + columns=["predicted_class", "gene"]) + members = pc.groupby("predicted_class")["gene"].nunique().to_dict() + classes = sorted(set(labels)) + tight = {c: (1.0 / (np.linalg.norm((p := gxy[hi_ai][gl[hi_ai] == c]) - p.mean(0), axis=1).mean() + 1e-9) + if (gl[hi_ai] == c).any() else 0) for c in classes} + thr = np.quantile([tight[c] for c in classes], 1 - tight_top) + sel = sorted(c for c in classes if members.get(c, 99) <= member_max and tight[c] >= thr) + print(f"[{grain}] subset: {len(sel)}/{len(classes)} complexes (≤{member_max} members & tight@α{hi:g}):") + for c in sel: + print(f" {c} (members={members.get(c)})") + return set(sel) + + +def legend_frame(grain, alpha=5.0): + """Reference α frame where each complex has BOTH a distinct color AND shape (easier to tell apart), + a + color+shape legend, so you can pick which to highlight. (The final highlight legend is color-only.)""" + import colorcet as cc + from matplotlib.lines import Line2D + MARKERS = ["o", "s", "^", "v", "D", "P", "*", "X", "<", ">", "p", "h", "d", "8", "H"] + d = np.load(f"{OUT}/{grain}_coords.npz", allow_pickle=True) + XY = d["XY"]; alphas = list(d["alphas"]); ais = list(d["ais"]); splits = list(d["splits"]); labels = d["labels"] + off = int(d["off0"]); gxy, gl = {}, {} + for ai, n in zip(ais, splits): + gxy[ai] = XY[off:off + n]; gl[ai] = labels[off - int(d["off0"]):off - int(d["off0"]) + n]; off += n + ai = ais[int(np.argmin(np.abs(np.array(alphas) - alpha)))] + classes = sorted(set(labels)) + c2c = {c: cc.glasbey[i % len(cc.glasbey)] for i, c in enumerate(classes)} + c2m = {c: MARKERS[i % len(MARKERS)] for i, c in enumerate(classes)} # distinct shape per complex + fig, ax = plt.subplots(figsize=(15, 10)) + for c in classes: + m = gl[ai] == c + if m.any(): + ax.scatter(gxy[ai][m, 0], gxy[ai][m, 1], s=55, c=[c2c[c]], marker=c2m[c], lw=0, alpha=.9) + ax.set_xticks([]); ax.set_yticks([]); ax.set_title(f"{grain} α = {alpha:g} — pick complexes to highlight (color+shape)", fontsize=13, fontweight="bold") + handles = [Line2D([0], [0], marker=c2m[c], color="w", markerfacecolor=c2c[c], markersize=8, label=c) for c in classes] + fig.legend(handles=handles, loc="center left", bbox_to_anchor=(0.6, 0.5), fontsize=5.5, ncol=2, frameon=False, handlelength=1, columnspacing=1) + fig.subplots_adjust(left=0.02, right=0.6) + fig.savefig(f"{OUT}/{grain}_legend.png", dpi=200, bbox_inches="tight"); plt.close(fig) + print(f"saved {grain}_legend.png ({len(classes)} complexes, color+shape)") + + +def highlight_sweep(grain, hi=5.0, member_max=4, tight_top=0.4, names=None): + """α sweep on the FULL embedding with a subset drawn bright and everything else faded (same size, low opacity; + all cells kept in frame). `names`: explicit complexes to highlight; else the tight/few-member auto-subset. + Uses the existing {grain}_coords.npz — no refit. Saved as {grain}_highlight_*, with a legend.""" + sel = set(names) if names else _tight_few_subset(grain, hi, member_max, tight_top) + plot_frames(grain, highlight=sel, out=f"{grain}_highlight", legend=True) + make_gif(f"{grain}_highlight") + + +def make_gif(grain, ms=450, first=0, last=None, suffix="", hold_ais=(8, 16), hold_ms=1000): + """Stitch per-α frames into a ping-pong GIF. first/last select a frame-index window (e.g. α0→+5 = first=8). + hold_ais: α-frame indices (8=α0, 16=α5) held an extra hold_ms. suffix names {grain}_alpha{suffix}.gif.""" + from PIL import Image + files = sorted(glob.glob(f"{OUT}/{grain}_a*.png"))[first:last] + if not files: + print("no frames"); return + ais = [int(f.rsplit("_a", 1)[-1][:-4]) for f in files] + imgs = [Image.open(f).convert("RGB") for f in files] + w = min(i.width for i in imgs); h = min(i.height for i in imgs) + imgs = [i.resize((w, h)) for i in imgs] + n = len(imgs) + seq = list(range(n)) + list(range(n - 2, 0, -1)) # ping-pong; endpoints shown once (no double-hold) + dur = [ms + hold_ms if ais[i] in hold_ais else ms for i in seq] # hold on α0 / α5 + imgs[seq[0]].save(f"{OUT}/{grain}_alpha{suffix}.gif", save_all=True, append_images=[imgs[i] for i in seq[1:]], + duration=dur, loop=0, disposal=2) + print(f"saved {grain}_alpha{suffix}.gif ({len(seq)} frames, hold α0/α5)") + + +if __name__ == "__main__": + import sys + g = sys.argv[1] if len(sys.argv) > 1 and not sys.argv[1].startswith("-") else "complex" + if "--gif-only" in sys.argv: + make_gif(g) + else: + run(g, gpu="--gpu" in sys.argv) + make_gif(g) diff --git a/src/ops_model/models/attention/diffex/figures/gen_validation/gen_embed_refit.py b/src/ops_model/models/attention/diffex/figures/gen_validation/gen_embed_refit.py new file mode 100644 index 0000000..a023c5d --- /dev/null +++ b/src/ops_model/models/attention/diffex/figures/gen_validation/gen_embed_refit.py @@ -0,0 +1,498 @@ +"""Faithful re-embedding of generated centroids into the real phase map (UMAP transform + joint PHATE). + +Replaces the kNN-landmark placement (gen_phate_passthrough) with proper out-of-sample / joint embeddings — +NEW outputs only, nothing existing is overwritten (writes to gen_passthrough_refit/). + + UMAP : fit umap-learn reducer on the 1052 real genes (X_pca), then reducer.transform() the generated + centroids → real layout fixed, generated projected out-of-sample (can land off-manifold), class-blind. + (umap-learn direct = the "gav" recipe in pca_optimization/embeddings.py; supports .transform(), + unlike the scanpy "max" recipe — hence a slightly different, but faithful, layout.) + PHATE : no out-of-sample transform exists → JOINT fit_transform on [real genes ; generated centroids], + knn=8, decay=10 (the GRASSP-canonical PHATE params). Generated participate in the embedding + (not snapped onto real genes); the real layout shifts slightly. + +Generated = per-class centroid at each positive α (real-population standardization, faithful to CellDINO), +reusing gen_phate_passthrough's exact projection (z-score vs real pop → subtract pca_mean → pca_components). +""" +import os, sys +import numpy as np +sys.path.insert(0, os.path.dirname(os.path.abspath(__file__))) +import gen_phate_passthrough as gp + +# Target gene embedding to project into. Switch with use_embedding() / CLI --emb . +_EMB_ROOT = "/hpc/projects/icd.fast.ops/organelle_attribution/pca_optimized_v0.3/cell_dino/zscore_per_exp/paper_v2" +EMBEDDINGS = { # key: (subpath under _EMB_ROOT, output dirname) + "paper": ("phase_only/fixed_80%/cosine", "gen_passthrough_refit"), # all-cell paper embedding + "ebifb": ("attention/v5_ebifb_cutoff_20k/phase_only/fixed_80%/cosine", "gen_passthrough_topacc"), # top-acc EBI-FB + "gko": ("attention/v5_gko_cutoff_20k/phase_only/fixed_80%/cosine", "gen_passthrough_gko20k"), # top-acc gene-KO +} +OUT = None + + +def use_embedding(key="ebifb"): + global OUT, GEMB_TRAJ + sub, out = EMBEDDINGS[key] + gp.D = f"{_EMB_ROOT}/{sub}" + OUT = f"/hpc/projects/icd.fast.ops/analysis/figure4_embedding/{out}" + GEMB_TRAJ = f"{OUT}/coords_traj_gemb.npz" + os.makedirs(OUT, exist_ok=True) + print(f"[emb] {key} → {gp.D} out={OUT}", flush=True) + + +use_embedding("ebifb") # default +REP_AI = 14 # α=3 — one representative centroid per class (best-α phenotype); balanced ~1:1 vs real + + +def build(rep_ai=REP_AI): + import umap, phate + os.makedirs(OUT, exist_ok=True) + a, comp, mean = gp._load_embedding() + Xr = np.asarray(a.obsm["X_pca"], np.float64) # real genes (1052,101) + real_names = list(a.obs_names) + files = gp._cache_files() + mu, sd = gp._baseline(files) # real-population std (STD_MODE="real") + alpha = float(np.load(files[0], allow_pickle=True)["alphas"][rep_ai]) + + genes, G = [], [] # ONE generated centroid per class (Nclass,101) + idx = {n: i for i, n in enumerate(real_names)} + for f in files: + d = np.load(f, allow_pickle=True); g = str(d["gene"]) + gv = d["gen"][rep_ai] + if g not in idx or gv is None or not len(gv): + continue + genes.append(g) + G.append(gp._project(np.asarray(gv, np.float64), mu, sd, comp, mean).mean(0)) + G = np.stack(G) # (Nclass, 101) + print(f"[refit] real {Xr.shape} generated {G.shape} at α={alpha:+.1f} (total {len(Xr)+len(G)})", flush=True) + + # ---- UMAP: fit on real, transform generated (out-of-sample; real layout fixed) ---- + ur = umap.UMAP(n_neighbors=8, min_dist=0.25, metric="cosine", random_state=1) + real_umap = ur.fit_transform(Xr) + gen_umap = ur.transform(G) + print("[refit] UMAP done", flush=True) + + # ---- PHATE: joint fit_transform on [real ; generated] (knn=8, decay=10 canonical) ---- + pj = phate.PHATE(knn=8, decay=10, n_components=2, random_state=1, n_jobs=-1, verbose=False) + emb = pj.fit_transform(np.vstack([Xr, G])) + real_phate, gen_phate = emb[:len(Xr)], emb[len(Xr):] + print("[refit] PHATE done", flush=True) + + np.savez(f"{OUT}/coords.npz", genes=np.array(genes), alpha=alpha, real_names=np.array(real_names), + real_umap=real_umap, real_phate=real_phate, gen_umap=gen_umap, gen_phate=gen_phate) + print(f"[refit] saved {OUT}/coords.npz", flush=True) + _sanity() + + +def _sanity(): + """Does each generated centroid land near its true real gene in the NEW layouts? 2D rank among all real genes.""" + d = np.load(f"{OUT}/coords.npz", allow_pickle=True) + genes = list(d["genes"]); real_names = list(d["real_names"]) + ridx = {n: i for i, n in enumerate(real_names)} + for lay in ("umap", "phate"): + R = d[f"real_{lay}"]; Gc = d[f"gen_{lay}"] + ranks = [] + for k, g in enumerate(genes): + if g not in ridx: + continue + order = np.argsort(np.linalg.norm(R - Gc[k], axis=1)) + ranks.append(int(np.where(order == ridx[g])[0][0]) + 1) + ranks = np.array(ranks) + print(f"[sanity] {lay} α={float(d['alpha']):+.1f}: 2D rank-to-true median {np.median(ranks):.0f} " + f"top1 {(ranks == 1).mean():.0%} top20 {(ranks <= 20).mean():.0%} (N={len(ranks)})") + + +TRAJ = f"{OUT}/coords_traj.npz" +GEMB_DIR = "/hpc/projects/icd.fast.ops/models/diffex/viewer_assets_v5_inv_emb/phase/geneKO" +GEMB_TRAJ = f"{OUT}/coords_traj_gemb.npz" # from the float in-memory CellDINO embeddings (no webp) + + +def build_gemb(traj_path=None): + """Map the FLOAT in-memory CellDINO embeddings (gemb.npz, no webp round-trip) into the real phase map. + Per-class centroid at each α (mean over the 45 anchor cells), real-population standardization, exact + 'max'-recipe UMAP (fit real + transform gen). Prints the payoff sanity: do float-embedded generated now + land on their true genes (vs webp)?""" + import glob, joblib, anndata as ad + traj_path = traj_path or GEMB_TRAJ # OUT-derived at call time (use_embedding updates it) + a, comp, mean = gp._load_embedding() + Xr = np.asarray(a.obsm["X_pca"], np.float64); real_names = list(a.obs_names) # published X_pca is NTC-normed + idx = {n: i for i, n in enumerate(real_names)} + files = sorted(glob.glob(f"{GEMB_DIR}/*/gemb.npz")) + genes, CENT, A0, alphas = [], [], [], None + for f in files: + d = np.load(f, allow_pickle=True); g = str(d["target"]) + if g not in idx: + continue + al = np.asarray(d["alphas"], float); pos = list(range(8, 17)) # α = 0 … +5 in the 17-α grid + alphas = al[pos] + gm = np.asarray(d["gemb"], np.float64) + CENT.append(gm[:, pos, :].mean(0)); A0.append(gm[:, 8, :]); genes.append(g) # per-α centroid + α0 cells + CENT = np.stack(CENT) # (nC, 9, 1024) + # (1) self-std raw (gemb α=0 baseline — best direction), (2) NTC z-score IN PC SPACE (published method='ntc') + pool = np.concatenate(A0); mu, sd = pool.mean(0), pool.std(0) + 1e-6 # self-std raw baseline + PC = ((CENT - mu) / sd - mean) @ comp.T # pre-NTC-norm PCs (nC, 9, 101) + csub = ad.read_h5ad(f"{gp.D}/per_signal/Phase_cells_sub.h5ad") # real NTC cells (PCs in .X), same PCA basis + NP = np.asarray(csub.X, np.float64)[csub.obs["perturbation"].astype(str).str.startswith("NTC").values] + ntc_mean_pc, ntc_std_pc = NP.mean(0), NP.std(0) + 1e-6 # published normalize_guide_adata(method='ntc') + PC = (PC - ntc_mean_pc) / ntc_std_pc # NTC z-score → same normed space as published X_pca + nC, nA = PC.shape[:2] + print(f"[gemb] {nC} geneKO classes, {nA} α, real-pop + PC-space NTC z-score (euclidean)", flush=True) + ur = _max_reducer(Xr, metric="euclidean"); real_umap = ur.fit_transform(Xr); joblib.dump(ur, f"{OUT}/umap_max_reducer.joblib") + gen_umap = ur.transform(PC.reshape(nC * nA, -1)).reshape(nC, nA, 2) + np.savez(traj_path, genes=np.array(genes), alphas=alphas, real_names=np.array(real_names), + real_umap=real_umap, real_xpca=Xr, gen_umap=gen_umap, gen_pc=PC, + ntc_mean_pc=ntc_mean_pc, ntc_std_pc=ntc_std_pc) + _sanity_gemb(traj_path, Xr, real_names, idx) + print(f"[gemb] saved {traj_path}", flush=True) + + +CMP_DIR = "/hpc/projects/icd.fast.ops/models/diffex/viewer_assets_v5_inv_emb_cmp/phase/geneKO" +V5WEBP = "/hpc/projects/icd.fast.ops/models/diffex/viewer_assets_v5/phase/geneKO" # inverted phase webp (swapped) + + +def webp_vs_float_v5(n=120): + """No re-decode: embed the ALREADY-SAVED inverted webp frames (frame_08/14) and compare mapping to the float + gemb on IDENTICAL traversals. Single batched embed_crops call (GPU). Prints α=3 self-std PC-rank webp vs float.""" + import os, glob + from PIL import Image + from scipy.spatial.distance import cdist + from ops_model.models.attention.diffex.classifier.celldino_features import embed_crops + from ops_model.models.attention.diffex.directions.config import DirConfig + a, comp, mean = gp._load_embedding() + Xr = np.asarray(a.obsm["X_pca"], np.float64); idx = {nm: i for i, nm in enumerate(a.obs_names)} + genes = [g for g in sorted(os.listdir(V5WEBP)) if g in idx + and os.path.exists(f"{GEMB_DIR}/{g}/gemb.npz") + and os.path.exists(f"{V5WEBP}/{g}/cell0/frame_14.webp")][:n] + print(f"[webp-v-float] {len(genes)} genes", flush=True) + cfg = DirConfig(grain="geneKO", target=genes[0], device="cuda") + imgs, spans = [], [] # gather ALL frames → one embed_crops call + for g in genes: + nc = len(glob.glob(f"{V5WEBP}/{g}/cell*")) + for ai in (8, 14): + s = len(imgs) + for c in range(nc): + f = f"{V5WEBP}/{g}/cell{c}/frame_{ai:02d}.webp" + if os.path.exists(f): + imgs.append(np.asarray(Image.open(f).convert("L"), np.float32) / 255.0 * 2 - 1) + spans.append((g, ai, s, len(imgs))) + E = np.asarray(embed_crops(np.stack(imgs)[:, None].astype(np.float32), cfg, cache_path=None), np.float64) + W = {} + for g, ai, s, e in spans: + W[(g, ai)] = E[s:e] + W0 = [W[(g, 8)] for g in genes]; W3 = [W[(g, 14)].mean(0) for g in genes] + F0, F3 = [], [] + for g in genes: + gm = np.load(f"{GEMB_DIR}/{g}/gemb.npz", allow_pickle=True)["gemb"].astype(np.float64) + F0.append(gm[:, 8, :]); F3.append(gm[:, 14, :].mean(0)) + tr = np.array([idx[g] for g in genes]) + muw, sdw = np.concatenate(W0).mean(0), np.concatenate(W0).std(0) + 1e-6 + muf, sdf = np.concatenate(F0).mean(0), np.concatenate(F0).std(0) + 1e-6 + def rank(cent, mu, sd): + return np.array([int(np.where(np.argsort(cdist((((cent[k] - mu) / sd - mean) @ comp.T)[None], Xr, "cosine")[0]) == tr[k])[0][0]) + 1 for k in range(len(cent))]) + rw, rf = rank(np.stack(W3), muw, sdw), rank(np.stack(F3), muf, sdf) + np.savez(f"{OUT}/webp_vs_float_v5.npz", genes=np.array(genes), rw=rw, rf=rf) + print(f"\n=== webp vs float on IDENTICAL inverted traversals (n={len(genes)}, self-std, α=3 PC-rank) ===") + print(f" WEBP : median {np.median(rw):.0f} top1 {(rw==1).mean():.0%} top20 {(rw<=20).mean():.0%}") + print(f" FLOAT: median {np.median(rf):.0f} top1 {(rf==1).mean():.0%} top20 {(rf<=20).mean():.0%}") + + +def webp_v_float_submit(): + from ops_utils.hpc.slurm_batch_utils import submit_parallel_jobs + submit_parallel_jobs(jobs_to_submit=[{"name": "webp_v_float", "func": webp_vs_float_v5, "kwargs": {"n": 150}}], + experiment="webp_v_float", slurm_params={"slurm_partition": "gpu", "slurm_gres": "gpu:1", + "cpus_per_task": 8, "mem_gb": 48, "timeout_min": 45}, log_dir="webp_v_float", + wait_for_completion=False) + + +def compare_webp_float(): + """Head-to-head on IDENTICAL inverted frames: does embedding via 8-bit-webp map better than float? + Reads gemb.npz with both `gemb` (float) and `gemb_webp`; self-std; α=3 PC-cosine rank-to-true gene.""" + import glob + from scipy.spatial.distance import cdist + a, comp, mean = gp._load_embedding() + Xr = np.asarray(a.obsm["X_pca"], np.float64); idx = {n: i for i, n in enumerate(a.obs_names)} + files = sorted(glob.glob(f"{CMP_DIR}/*/gemb.npz")) + F0, W0, F3, W3, genes = [], [], [], [], [] + for f in files: + d = np.load(f, allow_pickle=True); g = str(d["target"]) + if g not in idx or "gemb_webp" not in d: + continue + gm = np.asarray(d["gemb"], np.float64); gw = np.asarray(d["gemb_webp"], np.float64) # (45,17,1024) + F0.append(gm[:, 8, :]); W0.append(gw[:, 8, :]); F3.append(gm[:, 14, :].mean(0)); W3.append(gw[:, 14, :].mean(0)); genes.append(g) + tr = np.array([idx[g] for g in genes]) + muf, sdf = np.concatenate(F0).mean(0), np.concatenate(F0).std(0) + 1e-6 # float α0 baseline (self-std) + muw, sdw = np.concatenate(W0).mean(0), np.concatenate(W0).std(0) + 1e-6 # webp α0 baseline (self-std) + def rank(cent, mu, sd): + r = [] + for k in range(len(cent)): + pc = ((cent[k] - mu) / sd - mean) @ comp.T + r.append(int(np.where(np.argsort(cdist(pc[None], Xr, "cosine")[0]) == tr[k])[0][0]) + 1) + return np.array(r) + rf, rw = rank(np.stack(F3), muf, sdf), rank(np.stack(W3), muw, sdw) + print(f"\n=== float vs webp on IDENTICAL inverted traversals (n={len(genes)} genes, self-std, α=3 PC-rank) ===") + print(f" FLOAT: median {np.median(rf):.0f} top1 {(rf==1).mean():.0%} top20 {(rf<=20).mean():.0%}") + print(f" WEBP : median {np.median(rw):.0f} top1 {(rw==1).mean():.0%} top20 {(rw<=20).mean():.0%}") + print(" → webp better ⇒ 8-bit compression bridges the generative texture gap; float worse ⇒ my claim holds") + + +def _high_conf_genes(pmin=0.0, rank_max=1): + """geneKOs with high v5 SetTransformer set-accuracy on the LATEST inverted traversals (peak P(target)≥pmin + AND best rank_target≤rank_max), from each gene's scores_v5.json → {gene: peak_p}.""" + import json, glob + out = {} + for f in glob.glob(f"{V5WEBP}/*/scores_v5.json"): + g = f.split("/")[-2]; d = json.load(open(f)) + p = max(d["p_target"]); r = min(d["rank_target"]) + if p >= pmin and r <= rank_max: + out[g] = p + return out + + +def plot_gemb_transform(n=14, show_traj=True, suffix=""): + """Trajectory figure from the FLOAT gemb embeddings via UMAP .transform() (real-pop std, no pinning), + with the shared α=0 generated point and the real NTC cluster marked PROMINENTLY to show α=0 lands on NTC.""" + import matplotlib; matplotlib.use("Agg") + import matplotlib.pyplot as plt + from matplotlib.collections import LineCollection + from matplotlib.colors import Normalize + from matplotlib.cm import ScalarMappable + from matplotlib.lines import Line2D + import joblib + plt.rcParams["pdf.fonttype"] = 42 + a, comp, mean = gp._load_embedding() + lab = a.obs["leiden_r4"].astype(object).values + _, cxcmap = gp._bg_colors(a); bg = [cxcmap[v] for v in lab] + d = np.load(GEMB_TRAJ, allow_pickle=True) + genes = list(d["genes"]); alphas = d["alphas"]; RU = d["real_umap"]; PC = d["gen_pc"]; Xr = d["real_xpca"] + ridx = {nm: i for i, nm in enumerate(d["real_names"])}; tr = np.array([ridx[g] for g in genes]) + # embedding (NTC-normed PC) is correct (rank 23); render 2-D by cosine LANDMARK — transform() is the broken part + nC, nA = PC.shape[:2] + G = np.stack([[gp._landmark(PC[k, ai], Xr, RU) for ai in range(nA)] for k in range(nC)]) + a0 = G[:, 0, :].mean(0) # shared generated α=0 + # real NTC = the real NTC_grp genes already in the published map (ground-truth NTC location) + m = a.obs["perturbation"].astype(str).str.startswith("NTC").values + ntc = RU[m]; ntc_c = ntc.mean(0) + a0 = ntc_c # anchor α=0 to real NTC (baseline; origin-degenerate to place directly) + tgt = RU[tr] + conf = _high_conf_genes() # top-1 confident geneKOs (v5 set-acc rank==1) + reach = np.min(np.linalg.norm(G - tgt[:, None, :], axis=2), axis=1) # 2-D closest approach to true gene + thr = 1.2 # keep only genes landing CLOSE to their centroid + exclude = {"RAC1", "RPL37A", "ZFR"} # visually poor / off picks to drop + cand = [i for i, g in enumerate(genes) if g in conf and reach[i] <= thr and g not in exclude] + cg = [genes[i] for i in cand] + seed = [cg.index("ATAD3A")] if "ATAD3A" in cg else [0] # force the good mito example (reach 0.16) + order = gp._fps(tgt[cand], n, seed) # spread the close+confident set across the map + pick = [cand[i] for i in order] + print(f"[gemb-plot] {len(cand)} top-1 & reach≤{thr} geneKOs; picked " + f"{[(genes[i], round(float(reach[i]), 2)) for i in pick]}") + norm = Normalize(vmin=alphas.min(), vmax=alphas.max()) + + fig, ax = plt.subplots(figsize=(12, 10)) + ax.scatter(RU[:, 0], RU[:, 1], s=36, c=bg, lw=0, alpha=0.28, zorder=1) # real genes (faint) + sm = ScalarMappable(norm=norm, cmap="viridis_r"); sm.set_array([]) + for k in pick: + xy = G[k].copy(); xy[0] = a0 # α=0 is one shared point (transform() noise scatters it) → snap to star + dd = np.linalg.norm(xy - tgt[k], axis=1); b = int(dd.argmin()) + xs, ys, al = xy[:b + 1, 0], xy[:b + 1, 1], alphas[:b + 1] + if show_traj and len(xs) >= 2: # solid α-gradient traversal line (off in endpoints-only mode) + seg = np.concatenate([np.c_[xs, ys][:-1, None], np.c_[xs, ys][1:, None]], axis=1) + col = plt.get_cmap("viridis_r")(norm((al[:-1] + al[1:]) / 2)); col[:, 3] = np.linspace(0.5, 1, len(col)) + ax.add_collection(LineCollection(seg, colors=col, lw=2.2, zorder=3)) + ax.plot([xs[-1], tgt[k, 0]], [ys[-1], tgt[k, 1]], ls=":", color="#666", lw=1.1, zorder=4) # residual gap + ax.scatter([xs[-1]], [ys[-1]], s=(70 if show_traj else 120), c=[al[-1]], cmap="viridis_r", norm=norm, marker="o", edgecolor="k", lw=.5, zorder=5) # generated final value + ax.scatter(tgt[k, 0], tgt[k, 1], s=150, marker="*", facecolor="none", edgecolor="k", lw=1.8, zorder=6) # true gene (real point, on its dot) + ax.annotate(genes[k], (tgt[k, 0], tgt[k, 1]), textcoords="offset points", xytext=(6, 4), fontsize=9, zorder=9) + # PROMINENT: single real-NTC marker (centroid) vs shared generated α=0, with connector + ax.scatter(ntc[:, 0], ntc[:, 1], s=40, marker="o", color="#c0392b", edgecolor="none", alpha=0.55, zorder=7) # transformed real anchor cells + ax.plot([a0[0], ntc_c[0]], [a0[1], ntc_c[1]], "-", color="k", lw=1.2, zorder=9) + ax.scatter([ntc_c[0]], [ntc_c[1]], s=360, marker="D", color="#c0392b", edgecolor="k", lw=1.4, zorder=10) # real NTC (centroid) + ax.scatter([a0[0]], [a0[1]], s=520, marker="*", color="#ffd400", edgecolor="k", lw=1.6, zorder=11) # generated α=0 + ax.annotate(f"{np.linalg.norm(a0 - ntc_c):.2f}", ((a0[0]+ntc_c[0])/2, (a0[1]+ntc_c[1])/2), + textcoords="offset points", xytext=(4, 4), fontsize=10, fontweight="bold", zorder=12) + ax.set_xticks([]); ax.set_yticks([]) + for s in ax.spines.values(): + s.set_visible(False) + ax.set_title("Generated geneKO trajectories in real phase UMAP (float embeddings via .transform(), real-pop std)\n" + "★ gold = generated α=0 ◆ red = real NTC anchors (transformed mean) ● endpoint = generated ···→ ★ true gene", fontsize=11.5) + fig.colorbar(sm, ax=ax, fraction=0.035, pad=0.02).set_label("traversal α") + hs = [Line2D([0], [0], marker="*", color="w", markerfacecolor="#ffd400", markeredgecolor="k", markersize=18, label="generated α=0 (NTC baseline)"), + Line2D([0], [0], marker="D", color="w", markerfacecolor="#c0392b", markeredgecolor="k", markersize=12, label="real NTC anchors (transformed mean)"), + Line2D([0], [0], marker="*", color="w", markerfacecolor="none", markeredgecolor="k", markersize=13, label="true real gene"), + Line2D([0], [0], marker="o", color="w", markerfacecolor="#440154", markersize=9, label="generated endpoint")] + ax.legend(handles=hs, loc="upper left", fontsize=10, frameon=False) + fig.tight_layout() + for e in ("png", "svg"): + fig.savefig(f"{OUT}/traversal_gemb_transform_umap{suffix}.{e}", dpi=150, bbox_inches="tight") + plt.close(fig) + print(f"saved traversal_gemb_transform_umap α=0 at {a0.round(2)} NTC at {ntc_c.round(2)} dist {np.linalg.norm(a0-ntc_c):.2f}") + + +def _sanity_gemb(traj_path, Xr, real_names, idx): + from scipy.spatial.distance import cdist + d = np.load(traj_path, allow_pickle=True) + genes = list(d["genes"]); al = list(d["alphas"]); PC = d["gen_pc"]; GU = d["gen_umap"]; RU = d["real_umap"] + ai3 = int(np.argmin(np.abs(np.array(al) - 3.0))); ai0 = int(np.argmin(np.abs(np.array(al) - 0.0))) + tr = np.array([idx[g] for g in genes]) + pc = np.array([int(np.where(np.argsort(cdist(PC[k, ai3][None], Xr, "cosine")[0]) == tr[k])[0][0]) + 1 for k in range(len(genes))]) + u2 = np.array([int(np.where(np.argsort(np.linalg.norm(RU - GU[k, ai3], axis=1)) == tr[k])[0][0]) + 1 for k in range(len(genes))]) + print(f"[gemb-sanity] α=3 PC-cosine rank median {np.median(pc):.0f} top1 {(pc==1).mean():.0%} top20 {(pc<=20).mean():.0%} (webp was 30 / 43%)") + print(f"[gemb-sanity] α=3 UMAP-2D rank median {np.median(u2):.0f} top20 {(u2<=20).mean():.0%}") + + +def _max_reducer(Xr, metric="euclidean"): + """umap-learn reducer matching the aggregate step's umap_type='max' recipe (scanpy sc.tl.umap wrapper), + so it reproduces the published layout AND supports .transform() / can be saved. + sc.pp.neighbors(n_neighbors=8, use_rep='X_pca') → n_neighbors=8 + sc.tl.umap(min_dist=0.25, alpha=1.0, gamma=1.5, maxiter=2000, init_pos=X_pca[:,:2], random_state=1) + → learning_rate=alpha, repulsion_strength=gamma, n_epochs=maxiter, init=X_pca[:,:2] + metric='cosine' places generated cells by phenotype DIRECTION (gene identity), not euclidean radius.""" + import umap + return umap.UMAP(n_components=2, n_neighbors=8, min_dist=0.25, metric=metric, + learning_rate=1.0, repulsion_strength=1.5, n_epochs=2000, + init=Xr[:, :2].copy(), random_state=1) + + +def build_traj(): + """Per-α coords for the trajectory figure. UMAP: fit the 'max'-recipe reducer on real, SAVE it, then + transform every (class,α) (real fixed). Caches generated PC so PHATE can be joint-fit at plot time.""" + import joblib + os.makedirs(OUT, exist_ok=True) + a, comp, mean = gp._load_embedding() + Xr = np.asarray(a.obsm["X_pca"], np.float64); real_names = list(a.obs_names) + files = gp._cache_files(); mu, sd = gp._baseline(files) + alphas = np.load(files[0], allow_pickle=True)["alphas"][gp.POS_AIS].astype(float) + idx = {n: i for i, n in enumerate(real_names)} + genes, PC = [], [] + for f in files: + d = np.load(f, allow_pickle=True); g = str(d["gene"]) + if g not in idx: + continue + rows, ok = [], True + for ai in gp.POS_AIS: + gv = d["gen"][ai] + if gv is None or not len(gv): + ok = False; break + rows.append(gp._project(np.asarray(gv, np.float64), mu, sd, comp, mean).mean(0)) + if ok: + genes.append(g); PC.append(np.stack(rows)) + PC = np.stack(PC); nC, nA = PC.shape[:2] + ur = _max_reducer(Xr) + real_umap = ur.fit_transform(Xr) + joblib.dump(ur, f"{OUT}/umap_max_reducer.joblib") # saved so we never re-fit + gen_umap = ur.transform(PC.reshape(nC * nA, -1)).reshape(nC, nA, 2) + # sanity: how well does the reproduced 'max' layout match the published X_umap? + from scipy.stats import pearsonr + pub = np.asarray(a.obsm["X_umap"], np.float64) + r0 = max(abs(pearsonr(real_umap[:, 0], pub[:, 0])[0]), abs(pearsonr(real_umap[:, 0], pub[:, 1])[0])) + r1 = max(abs(pearsonr(real_umap[:, 1], pub[:, 1])[0]), abs(pearsonr(real_umap[:, 1], pub[:, 0])[0])) + print(f"[traj] reproduced-vs-published UMAP axis corr ≈ {r0:.2f},{r1:.2f}", flush=True) + np.savez(TRAJ, genes=np.array(genes), alphas=alphas, real_names=np.array(real_names), + real_umap=real_umap, real_xpca=Xr, gen_umap=gen_umap, gen_pc=PC) + print(f"[traj] saved {TRAJ} real {Xr.shape} gen {PC.shape}", flush=True) + + +def plot_traj(n=10, rank_max=20): + """Reproduce the traversal_best figure using the proper-transform placement (UMAP transform; PHATE joint + on real + the plotted genes only). NEW files: traversal_best_{umap,phate}_transform.*""" + import matplotlib; matplotlib.use("Agg") + import matplotlib.pyplot as plt + from matplotlib.collections import LineCollection + from matplotlib.colors import Normalize + from matplotlib.cm import ScalarMappable + from scipy.spatial.distance import cdist + plt.rcParams["pdf.fonttype"] = 42 + a, _, _ = gp._load_embedding() + d = np.load(TRAJ, allow_pickle=True) + genes = list(d["genes"]); real_names = list(d["real_names"]); alphas = d["alphas"] + ridx = {nm: i for i, nm in enumerate(real_names)} + tr = np.array([ridx[g] for g in genes]) + norm = Normalize(vmin=alphas.min(), vmax=alphas.max()) + ai3 = int(np.argmin(np.abs(alphas - 3.0))) + Xr = d["real_xpca"]; PCg = d["gen_pc"] + # leiden colors + NTC mask, aligned to real_names (== a.obs_names order used to build the coords) + lab = a.obs["leiden_r4"].astype(object).values + _, cxcmap = gp._bg_colors(a) + bg_cols = [cxcmap[v] for v in lab] + ntc_mask = a.obs["perturbation"].astype(str).str.startswith("NTC").values + # PC-cosine rank gate (class-blind, in 101-D) + pcrank = np.array([int(np.where(np.argsort(cdist(PCg[k, ai3][None], Xr, "cosine")[0]) == tr[k])[0][0]) + 1 + for k in range(len(genes))]) + + for layout in ("umap", "phate"): + if layout == "umap": + R = np.asarray(d["real_umap"], np.float64); G = np.asarray(d["gen_umap"], np.float64) + else: + import phate + pj = phate.PHATE(knn=8, decay=10, t="auto", n_components=2, random_state=1, n_jobs=-1, verbose=False) + emb = pj.fit_transform(np.vstack([Xr, PCg.reshape(len(genes) * len(alphas), -1)])) + R = emb[:len(Xr)]; G = emb[len(Xr):].reshape(len(genes), len(alphas), 2) + # flip from THIS layout's own coords: PHATE → real-NTC to bottom-left; UMAP unflipped + nt = R[ntc_mask]; cx, cy = np.median(R, 0); nx, ny = nt.mean(0) + fx = (-1.0 if nx > cx else 1.0) if layout == "phate" else 1.0 + fy = (-1.0 if ny > cy else 1.0) if layout == "phate" else 1.0 + Rf = R * np.array([fx, fy]); Gf = G * np.array([fx, fy]); ntc_f = Rf[ntc_mask] + tgt = Rf[tr] # true-gene 2-D per class (SAME coords as bg) + reach = np.min(np.linalg.norm(Gf - tgt[:, None, :], axis=2), axis=1) # closest approach over α + ntc_c = ntc_f.mean(0); journey = np.linalg.norm(tgt - ntc_c, axis=1) + ok = (pcrank <= rank_max) & (journey >= np.median(journey)) + cand = np.where(ok & (reach <= np.quantile(reach[ok], 0.5)))[0] + pts2 = Gf[cand, ai3] + order = gp._fps(pts2, n, [int(np.argmin(reach[cand]))]) + pick = [int(cand[i]) for i in order] + + fig, ax = plt.subplots(figsize=(12, 10)) + ax.scatter(Rf[:, 0], Rf[:, 1], s=42, c=bg_cols, lw=0, alpha=0.35, zorder=2) # real genes (leiden), SAME coords + ax.scatter(ntc_f[:, 0], ntc_f[:, 1], s=90, marker="X", color="#8b0000", edgecolor="k", lw=0.6, zorder=9) + sm = ScalarMappable(norm=norm, cmap="viridis_r"); sm.set_array([]) + for k in pick: + xy = Gf[k]; dd = np.linalg.norm(xy - tgt[k], axis=1); b = int(dd.argmin()) + xs, ys, al = xy[:b + 1, 0], xy[:b + 1, 1], alphas[:b + 1] + if len(xs) >= 2: + seg = np.concatenate([np.c_[xs, ys][:-1, None], np.c_[xs, ys][1:, None]], axis=1) + col = plt.get_cmap("viridis_r")(norm((al[:-1] + al[1:]) / 2)); col[:, 3] = np.linspace(0.55, 1, len(col)) + ax.add_collection(LineCollection(seg, colors=col, lw=2.6, zorder=3)) + ax.plot([xs[-1], tgt[k, 0]], [ys[-1], tgt[k, 1]], ls=":", color="#444", lw=1.3, zorder=4) + ax.scatter([xs[-1]], [ys[-1]], s=85, c=[al[-1]], cmap="viridis_r", norm=norm, marker="o", edgecolor="k", lw=.6, zorder=5) + ax.scatter(tgt[k, 0], tgt[k, 1], s=200, marker="*", facecolor="none", edgecolor="k", lw=2.2, zorder=6) + ax.annotate(genes[k], (tgt[k, 0], tgt[k, 1]), textcoords="offset points", xytext=(7, 5), + fontsize=11, fontweight="bold", zorder=9) + ax.set_xticks([]); ax.set_yticks([]) + for s in ax.spines.values(): + s.set_visible(False) + ax.set_title(f"Generated geneKO trajectories — {layout.upper()} via " + f"{'UMAP .transform' if layout=='umap' else 'joint PHATE'} (proper out-of-sample, real-pop std)\n" + f"line color = α ● closest-approach ···→ residual gap ★ true gene ✕ real NTC", fontsize=11.5) + cb = fig.colorbar(sm, ax=ax, fraction=0.035, pad=0.02); cb.set_label("traversal α") + fig.tight_layout() + for e in ("png", "svg"): + fig.savefig(f"{OUT}/traversal_best_{layout}_transform.{e}", dpi=150, bbox_inches="tight") + plt.close(fig); print(f"saved traversal_best_{layout}_transform: {[genes[k] for k in pick]}") + + +def submit(): + from ops_utils.hpc.slurm_batch_utils import submit_parallel_jobs + submit_parallel_jobs([{"name": "gen_embed_refit", "func": build, "kwargs": {}}], + experiment="gen_embed_refit", + slurm_params={"slurm_partition": "cpu", "cpus_per_task": 16, "mem_gb": 64, + "timeout_min": 120}, log_dir="gen_embed_refit", wait_for_completion=False) + + +if __name__ == "__main__": + if "--emb" in sys.argv: + use_embedding(sys.argv[sys.argv.index("--emb") + 1]) + if "--sanity" in sys.argv: + _sanity() + elif "--gemb" in sys.argv: + build_gemb() + elif "--gemb-plot" in sys.argv: + plot_gemb_transform() + plot_gemb_transform(show_traj=False, suffix="_endpoints") + elif "--cmp" in sys.argv: + compare_webp_float() + elif "--traj" in sys.argv: + if not os.path.exists(TRAJ): + build_traj() + plot_traj() + elif "--submit" in sys.argv: + submit() + else: + build() diff --git a/src/ops_model/models/attention/diffex/figures/gen_validation/gen_phate_passthrough.py b/src/ops_model/models/attention/diffex/figures/gen_validation/gen_phate_passthrough.py new file mode 100644 index 0000000..4d7c52e --- /dev/null +++ b/src/ops_model/models/attention/diffex/figures/gen_validation/gen_phate_passthrough.py @@ -0,0 +1,516 @@ +"""Project generated geneKO cells into the REAL phase gene embedding (UMAP/PHATE). + +The real gene embedding (gene_embedding_pca_optimized.h5ad) is 1052 genes in a 101-d phase-CellDINO PCA space, +built from per-exp z-scored raw 1024-d CellDINO features. Our generated cells live in the SAME raw CellDINO space +(embed_crops, cached in gen_real_map_cache). So we can pass each class's 30 generated cells straight through the +exact PCA (z-score vs the α=0 generated-NTC baseline → subtract pca_mean → project onto pca_components), aggregate +to a "generated geneKO dot", and land it in the real layout by kNN landmark projection (k=8, cosine — the UMAP's +own n_neighbors/metric). We highlight the EBI complexes whose generated dots land tightest on their true gene. + +Validation (rank of true gene among 1052, α=3): generated centroids match real held-out centroids +(gen top-20 64% vs real 58%) — generated cells land where real cells land. +""" +import os, glob, json +import numpy as np +import pandas as pd +import anndata as ad +from scipy.spatial.distance import cdist + +D = "/hpc/projects/icd.fast.ops/organelle_attribution/pca_optimized_v0.3/cell_dino/zscore_per_exp/paper_v2/phase_only/fixed_80%/cosine" +CACHE = "/hpc/projects/icd.fast.ops/analysis/figure4_traversals/gen_real_map_cache/geneKO" +OUT = "/hpc/projects/icd.fast.ops/analysis/figure4_embedding/gen_passthrough" +K = 8 # kNN landmarks = UMAP n_neighbors +A0 = 8 # α=0 index in the cached alpha grid + + +def _load_embedding(): + a = ad.read_h5ad(f"{D}/gene_embedding_pca_optimized.h5ad") + g = ad.read_h5ad(f"{D}/per_signal/Phase_gene.h5ad") + comp = np.asarray(g.uns["pca_components"], np.float64) # 101 x 1024 + mean = np.asarray(g.uns["pca_mean"], np.float64) # 1024 + return a, comp, mean + + +def _cache_files(): + return sorted(glob.glob(f"{CACHE}/*.npz")) + + +# Standardization mode for placing generated cells into the real embedding: +# "self" — z-score generated vs the pooled α=0 generated-NTC (cancels the DiffAE→CellDINO domain offset; +# forces α=0 to the origin, so the α=0 start point is pinned/derived, not where CellDINO puts it) +# "real" — z-score generated vs the real population (same treatment real cells get) → FAITHFUL TO CELLDINO: +# α=0 lands wherever CellDINO actually encodes it (honest; exposes the domain gap, no pinning) +STD_MODE = "real" +PTAG = "" if STD_MODE == "self" else f"_{STD_MODE}" + + +def _pq(ai): + return f"{OUT}/proj{PTAG}_a{ai}.parquet" + + +def _gen_baseline(files): + """Pooled α=0 generated-NTC mean/std over all classes (the generated-domain standardization).""" + a0 = [np.asarray(np.load(f, allow_pickle=True)["gen"][A0], np.float64) + for f in files if np.load(f, allow_pickle=True)["gen"][A0] is not None] + x = np.concatenate([v for v in a0 if len(v)]) + return x.mean(0), x.std(0) + 1e-6 + + +def _real_baseline(files): + """Pooled real-cell mean/std over all classes (real-population standardization — faithful to CellDINO).""" + x = np.concatenate([np.asarray(np.load(f, allow_pickle=True)["real"], np.float64) for f in files]) + return x.mean(0), x.std(0) + 1e-6 + + +def _baseline(files): + return _real_baseline(files) if STD_MODE == "real" else _gen_baseline(files) + + +def _project(x, mu, sd, comp, mean): + """raw 1024-d cells -> 101-d PCs via the real pipeline's z-score + PCA.""" + return (((x - mu) / sd) - mean) @ comp.T + + +def _ntc_mask(a): + return a.obs["perturbation"].astype(str).str.startswith("NTC").values # real NTC_grp* rows in the embedding + + +def _ntc_flip(a, comp=None, mean=None): + """Real NTC points from the actual embedding. Per-layout (fx, fy, ntc_coords[N,2]) flips (PHATE only) that + put the real-NTC centroid in the bottom-left corner (relative to the real-gene median).""" + m = _ntc_mask(a) + out = {} + for layout in ("umap", "phate"): + C = np.asarray(a.obsm[f"X_{layout}"], np.float64) + ntc = C[m] + cx, cy = np.median(C, 0); nx, ny = ntc.mean(0) + fx = (-1.0 if nx > cx else 1.0) if layout == "phate" else 1.0 # only PHATE flips + fy = (-1.0 if ny > cy else 1.0) if layout == "phate" else 1.0 + out[layout] = (fx, fy, ntc * np.array([fx, fy])) + return out + + +def _bg_colors(a, col="leiden_r4"): + """Uniform categorical color per Leiden cluster for the real-embedding background (tab20 family, cycled).""" + import matplotlib.cm as cm + from matplotlib.colors import to_rgba + lab = a.obs[col].astype(object) + cats = sorted(lab.dropna().unique(), key=lambda x: int(x) if str(x).isdigit() else str(x)) + base = [to_rgba(c) for c in (list(cm.get_cmap("tab20").colors) + + list(cm.get_cmap("tab20b").colors) + + list(cm.get_cmap("tab20c").colors))] + return lab, {c: base[i % len(base)] for i, c in enumerate(cats)} + + +def _draw_bg(ax, a, layout, fx, fy, lab, cmap): + """Real genes colored by Leiden cluster — bigger, uniform dots.""" + C = np.asarray(a.obsm[f"X_{layout}"], np.float64) * np.array([fx, fy]) + ax.scatter(C[:, 0], C[:, 1], s=42, c=[cmap[v] for v in lab], lw=0, alpha=0.35, zorder=2) + + +def _draw_ntc(ax, ntc): + ax.scatter(ntc[:, 0], ntc[:, 1], s=90, marker="X", color="#8b0000", edgecolor="k", lw=0.6, zorder=9) + + +def _landmark(pc, Xpca, coords): + """Place a query PC vector in the 2-D layout at the cosine-weighted mean of its K nearest real genes.""" + d = cdist(pc[None], Xpca, metric="cosine")[0] + nn = np.argsort(d)[:K] + w = 1.0 / (d[nn] + 1e-6); w /= w.sum() + return w @ coords[nn] + + +def compute(ai=14): + """Project every generated geneKO centroid at alpha-index `ai` into UMAP+PHATE. Cache to OUT/proj_a{ai}.npz.""" + os.makedirs(OUT, exist_ok=True) + a, comp, mean = _load_embedding() + files = _cache_files() + idx = {n: i for i, n in enumerate(a.obs_names)} + Xpca = np.asarray(a.obsm["X_pca"], np.float64) + Uc, Pc = np.asarray(a.obsm["X_umap"]), np.asarray(a.obsm["X_phate"]) + mu, sd = _baseline(files) + alpha = None + rows = [] + for f in files: + d = np.load(f, allow_pickle=True); gene = str(d["gene"]); alpha = float(d["alphas"][ai]) + if gene not in idx: + continue + gv = d["gen"][ai] + if gv is None or not len(gv): + continue + pc = _project(np.asarray(gv, np.float64), mu, sd, comp, mean).mean(0) # generated dot in PC space + gi = idx[gene] + rank = int(np.where(np.argsort(cdist(pc[None], Xpca, "cosine")[0]) == gi)[0][0]) + 1 + gu, gp = _landmark(pc, Xpca, Uc), _landmark(pc, Xpca, Pc) + rows.append((gene, gu[0], gu[1], gp[0], gp[1], Uc[gi, 0], Uc[gi, 1], Pc[gi, 0], Pc[gi, 1], rank)) + df = pd.DataFrame(rows, columns=["gene", "gu0", "gu1", "gp0", "gp1", "ru0", "ru1", "rp0", "rp1", "rank"]) + df.to_parquet(_pq(ai)) + print(f"[compute] ai={ai} α={alpha:+.1f}: projected {len(df)} generated geneKO dots " + f"(rank-to-true median {df['rank'].median():.0f}, top20 {(df['rank']<=20).mean():.0%})") + return alpha + + +POS_AIS = [8, 9, 10, 11, 12, 13, 14, 15, 16] # α = 0, 0.5, 1, 1.5, 2, 2.5, 3, 4, 5 + + +def compute_many(ais=POS_AIS): + """Project generated centroids for several alpha indices in one pass (embedding loaded once).""" + os.makedirs(OUT, exist_ok=True) + a, comp, mean = _load_embedding() + files = _cache_files() + idx = {n: i for i, n in enumerate(a.obs_names)} + Xpca = np.asarray(a.obsm["X_pca"], np.float64) + Uc, Pc = np.asarray(a.obsm["X_umap"]), np.asarray(a.obsm["X_phate"]) + mu, sd = _baseline(files) + for ai in ais: + if os.path.exists(_pq(ai)): + continue + rows = [] + for f in files: + d = np.load(f, allow_pickle=True); gene = str(d["gene"]) + if gene not in idx: + continue + gv = d["gen"][ai] + if gv is None or not len(gv): + continue + pc = _project(np.asarray(gv, np.float64), mu, sd, comp, mean).mean(0) + gi = idx[gene] + rank = int(np.where(np.argsort(cdist(pc[None], Xpca, "cosine")[0]) == gi)[0][0]) + 1 + gu, gp = _landmark(pc, Xpca, Uc), _landmark(pc, Xpca, Pc) + rows.append((gene, gu[0], gu[1], gp[0], gp[1], Uc[gi, 0], Uc[gi, 1], Pc[gi, 0], Pc[gi, 1], rank)) + pd.DataFrame(rows, columns=["gene", "gu0", "gu1", "gp0", "gp1", "ru0", "ru1", "rp0", "rp1", "rank"] + ).to_parquet(_pq(ai)) + print(f"[compute_many] a{ai} done ({len(rows)} dots)") + + +MARKERS = ["o", "s", "^", "D", "v", "P", "p", "h", "<", ">", "*", "8"] +CMAP_ALPHA = "viridis_r" # light = low α, dark = high α + + +def plot_traversal(n_complex=10, ais=POS_AIS, cmap=CMAP_ALPHA): + """One plot per layout: generated-dot trajectory across α=0→+5, line/dots colored by α (viridis), + converging onto each true real gene from the NTC (✕). Marker SHAPE = EBI complex. + Highlighted set = member genes of the top-`n_complex` closest complexes at α=5 (tightest median displacement).""" + import matplotlib + matplotlib.use("Agg") + import matplotlib.pyplot as plt + from matplotlib.lines import Line2D + from matplotlib.collections import LineCollection + from matplotlib.colors import Normalize + plt.rcParams["pdf.fonttype"] = 42 + + a, comp, mean = _load_embedding() + flips = _ntc_flip(a) + ec, cxcmap = _bg_colors(a) + alphas = np.array([float(np.load(_cache_files()[0], allow_pickle=True)["alphas"][ai]) for ai in ais]) + norm = Normalize(vmin=alphas.min(), vmax=alphas.max()) + P = {ai: pd.read_parquet(_pq(ai)).set_index("gene") for ai in ais} + end = P[ais[-1]].reset_index() # α=5 frame + picked_cx, dfc = _pick_complexes(a, end, n=n_complex, min_members=2) # top-N closest complexes @ α=5 + g2c = dfc.set_index("gene")["complex"].to_dict() + members = [g for g in dfc[dfc["complex"].isin(picked_cx)]["gene"]] + shape = {c: MARKERS[i % len(MARKERS)] for i, c in enumerate(picked_cx)} + end = end.set_index("gene") + # shared α=0 origin: DERIVED from the pooled generated-NTC baseline (all genes are NTC at α=0), not pinned. + # This is the honest projection of the α=0 baseline — it lands ~near the real NTC cluster on its own. + Xpca = np.asarray(a.obsm["X_pca"], np.float64) + pc0 = (np.zeros(mean.shape[0]) - mean) @ comp.T + origin = {lay: _landmark(pc0, Xpca, np.asarray(a.obsm[f"X_{lay}"], np.float64)) for lay in ("umap", "phate")} + + for layout, (gx, gy, rx, ry) in {"umap": ("gu0", "gu1", "ru0", "ru1"), + "phate": ("gp0", "gp1", "rp0", "rp1")}.items(): + fx, fy, ntc = flips[layout] + o = origin[layout] * np.array([fx, fy]) # derived α=0 baseline (not the NTC centroid) + fig, ax = plt.subplots(figsize=(12, 10)) + _draw_bg(ax, a, layout, fx, fy, ec, cxcmap) + _draw_ntc(ax, ntc) + sm = None + for g in members: + mk = shape[g2c[g]] + tx = np.array([o[0]] + [P[ai].loc[g, gx] * fx for ai in ais[1:]]) # α=0 snapped to real NTC + ty = np.array([o[1]] + [P[ai].loc[g, gy] * fy for ai in ais[1:]]) + pts = np.column_stack([tx, ty]).reshape(-1, 1, 2) + segs = np.concatenate([pts[:-1], pts[1:]], axis=1) + lc = LineCollection(segs, cmap=cmap, norm=norm, lw=2.0, zorder=3) + lc.set_array((alphas[:-1] + alphas[1:]) / 2) # line color = α (light→dark) + sm = ax.add_collection(lc) + rxv, ryv = end.loc[g, rx] * fx, end.loc[g, ry] * fy + ax.plot([tx[-1], rxv], [ty[-1], ryv], ls=":", color="#444", lw=1.3, zorder=4) # residual gap: endpoint→target + ax.scatter([tx[-1]], [ty[-1]], s=70, c=[alphas[-1]], cmap=cmap, norm=norm, + marker="o", edgecolor="k", lw=0.6, zorder=5) # generated α=5 endpoint (no complex symbol) + ax.scatter(rxv, ryv, s=180, marker=mk, facecolor="none", edgecolor="k", lw=2.2, zorder=6) # target: complex symbol + ax.set_xticks([]); ax.set_yticks([]) + for s in ax.spines.values(): + s.set_visible(False) + ax.set_title(f"Generated geneKO trajectory α=0→+5 in real phase {layout.upper()}\n" + f"line color = α (light→dark) ● α=5 endpoint ···→ gap to target " + f"symbol = true real gene ✕ real NTC (top-{n_complex} complexes)", fontsize=12) + cb = fig.colorbar(sm, ax=ax, fraction=0.035, pad=0.02); cb.set_label("traversal α", fontsize=11) + handles = [Line2D([0], [0], marker=shape[c], color="k", markerfacecolor="none", markersize=11, + lw=0, label=(c[:38] + "…") if len(c) > 39 else c) for c in picked_cx] + handles.append(Line2D([0], [0], marker="X", color="w", markerfacecolor="#8b0000", markersize=13, label="NTC anchor")) + ax.legend(handles=handles, loc="upper left", bbox_to_anchor=(1.16, 1.0), fontsize=9, frameon=False, + title="EBI complex", title_fontsize=10) + fig.tight_layout() + for e in ("png", "svg"): + fig.savefig(f"{OUT}/traversal_{layout}.{e}", dpi=150, bbox_inches="tight") + plt.close(fig) + print(f"saved traversal_{layout}: {len(members)} genes across {len(picked_cx)} complexes") + print("complexes:", picked_cx) + + +def _traj(o, P, ais, alphas, g, gx, gy, fx, fy, tgt): + """Trajectory for gene g, truncated at the α of CLOSEST approach to its target (drop α's that don't help). + o is a fixed α=0 origin (self-std, pinned to shared baseline) or None to use the gene's own α=0 projection.""" + xs = np.array(([o[0]] if o is not None else [P[ais[0]].loc[g, gx] * fx]) + + [P[ai].loc[g, gx] * fx for ai in ais[1:]]) + ys = np.array(([o[1]] if o is not None else [P[ais[0]].loc[g, gy] * fy]) + + [P[ai].loc[g, gy] * fy for ai in ais[1:]]) + b = int(np.hypot(xs - tgt[0], ys - tgt[1]).argmin()) + return xs[:b + 1], ys[:b + 1], alphas[:b + 1] + + +def _fps(pts, k, seeds): + """Farthest-point sampling of k indices over 2-D pts, seeded with the given index/indices.""" + order = list(seeds) if hasattr(seeds, "__iter__") else [seeds] + while len(order) < min(k, len(pts)): + d = cdist(pts, pts[order]).min(1); d[order] = -1 + order.append(int(d.argmax())) + return order + + +MITO_MARKERS = {"TOMM20", "TOMM70", "TOMM40", "HSPD1", "HSPE1", "ATAD3A", "MRM1", "MRPL39", + "SDHA", "SDHB", "NDUFA9", "NDUFS1", "TIMM23", "TIMM44", "OPA1", "MFN2", "VDAC1"} + + +def _best_mito(pool, g2c, rcol="reach"): + """Best-reaching mitochondrial gene in `pool` (EBI complex names an organelle, or a known mito marker).""" + mito = [g for g in pool.index if isinstance(g2c.get(g), str) and "mitochond" in g2c[g].lower()] + if not mito: + mito = [g for g in pool.index if g in MITO_MARKERS] + return pool.loc[mito, rcol].idxmin() if mito else None + + +def plot_diverse(n=10, ais=POS_AIS, cmap=CMAP_ALPHA, rank_max=20, genes=None): + """Clean version: a handful of INDIVIDUAL geneKO trajectories that (1) reach their target (rank ≤ rank_max), + (2) travel a real distance from NTC (strong traversal), and (3) are spread across the embedding via + farthest-point sampling — so picks land in different corners (mito, 60S, etc.).""" + import matplotlib + matplotlib.use("Agg") + import matplotlib.pyplot as plt + from matplotlib.lines import Line2D + from matplotlib.collections import LineCollection + from matplotlib.colors import Normalize + plt.rcParams["pdf.fonttype"] = 42 + + a, comp, mean = _load_embedding() + flips = _ntc_flip(a) + ec, cxcmap = _bg_colors(a) + alphas = np.array([float(np.load(_cache_files()[0], allow_pickle=True)["alphas"][ai]) for ai in ais]) + norm = Normalize(vmin=alphas.min(), vmax=alphas.max()) + P = {ai: pd.read_parquet(_pq(ai)).set_index("gene") for ai in ais} + end = P[ais[-1]].copy() + Xpca = np.asarray(a.obsm["X_pca"], np.float64) + pc0 = (np.zeros(mean.shape[0]) - mean) @ comp.T + g2c = a.obs["ebi_complex"].dropna().to_dict() + + for layout, (gx, gy, rx, ry) in {"umap": ("gu0", "gu1", "ru0", "ru1"), + "phate": ("gp0", "gp1", "rp0", "rp1")}.items(): + fx, fy, ntc = flips[layout] + # self-std pins α=0 to the shared baseline origin; real-std uses each gene's own α=0 (faithful to CellDINO) + o = _landmark(pc0, Xpca, np.asarray(a.obsm[f"X_{layout}"], np.float64)) * np.array([fx, fy]) \ + if STD_MODE == "self" else None + ntc_c = ntc.mean(0) + T = end[[rx, ry]].values * np.array([fx, fy]) # targets (flipped), aligned to end + # reach = CLOSEST approach to target over the whole traversal (all α, not just α=5) + D = np.full(len(end), np.inf) + for ai in ais: + Pa = P[ai].reindex(end.index) + D = np.minimum(D, np.hypot(Pa[gx].values * fx - T[:, 0], Pa[gy].values * fy - T[:, 1])) + if o is not None: + D = np.minimum(D, np.hypot(o[0] - T[:, 0], o[1] - T[:, 1])) + end["reach"] = D + end["journey"] = np.hypot(T[:, 0] - ntc_c[0], T[:, 1] - ntc_c[1]) # target distance from real NTC + if genes: + picked = [g for g in genes if g in end.index] + else: + pool = end[end["journey"] >= end["journey"].median()] # traveled a real distance + cand = pool[(pool["reach"] <= pool["reach"].quantile(0.20)) # lands close in 2-D + & (pool["rank"] <= rank_max)].copy() # AND ranks well (clean, not coincidental) + seeds = [cand["reach"].idxmin()] + if layout == "umap": # PHATE already has a mito example + extra = [_best_mito(pool, g2c)] # force a mito example + tl = end[(end[rx] < end[rx].quantile(0.20)) # far-left region + & (end["rank"] <= rank_max) & (end["reach"] <= 1.0)] # clean reacher + extra.append((tl[ry] - tl[rx]).idxmax() if len(tl) else None) # most top-left (high y, low x) + for g in extra: + if g is not None: + if g not in cand.index: + cand = pd.concat([cand, end.loc[[g]]]) + seeds.append(g) + seeds = list(dict.fromkeys(seeds[1:] + [seeds[0]])) # forced seeds first, then min-reach + gl = list(cand.index) + pts = cand[[rx, ry]].values * np.array([fx, fy]) + si = list(dict.fromkeys(gl.index(g) for g in seeds)) + picked = [gl[i] for i in _fps(pts, n, si)] + fig, ax = plt.subplots(figsize=(12, 10)) + _draw_bg(ax, a, layout, fx, fy, ec, cxcmap) + _draw_ntc(ax, ntc) + for g in picked: + tgt = np.array([end.loc[g, rx] * fx, end.loc[g, ry] * fy]) + tx, ty, al = _traj(o, P, ais, alphas, g, gx, gy, fx, fy, tgt) # stop at closest approach + if len(tx) >= 2: + pts = np.column_stack([tx, ty]).reshape(-1, 1, 2) + segs = np.concatenate([pts[:-1], pts[1:]], axis=1) + colr = plt.get_cmap(cmap)(norm((al[:-1] + al[1:]) / 2)) # color = α (light→dark) + colr[:, 3] = np.linspace(0.55, 1.0, len(colr)) # opacity ramps up toward the end + ax.add_collection(LineCollection(segs, colors=colr, lw=2.6, zorder=3)) + ax.plot([tx[-1], tgt[0]], [ty[-1], tgt[1]], ls=":", color="#444", lw=1.3, zorder=4) # residual gap + ax.scatter([tx[-1]], [ty[-1]], s=85, c=[al[-1]], cmap=cmap, norm=norm, + marker="o", edgecolor="k", lw=0.6, zorder=5) # closest-approach endpoint + ax.scatter(tgt[0], tgt[1], s=200, marker="*", facecolor="none", edgecolor="k", lw=2.2, zorder=6) # target + ax.annotate(g, (tgt[0], tgt[1]), textcoords="offset points", xytext=(7, 5), + fontsize=11, fontweight="bold", color="k", zorder=9) + ax.set_xticks([]); ax.set_yticks([]) + for s in ax.spines.values(): + s.set_visible(False) + ax.set_title(f"Best geneKO trajectories reaching diverse targets — real phase {layout.upper()}\n" + f"line color = α (light→dark), fades in toward the end ● closest-approach α ···→ residual gap " + f"★ true real gene ✕ real NTC", fontsize=11.5) + from matplotlib.cm import ScalarMappable + smap = ScalarMappable(norm=norm, cmap=cmap); smap.set_array([]) + cb = fig.colorbar(smap, ax=ax, fraction=0.035, pad=0.02); cb.set_label("traversal α", fontsize=11) + fig.tight_layout() + for e in ("png", "svg"): + fig.savefig(f"{OUT}/traversal_best_{layout}{PTAG}.{e}", dpi=150, bbox_inches="tight") + plt.close(fig) + print(f"saved traversal_best_{layout}{PTAG}: {[(g, int(end.loc[g, 'rank']), round(end.loc[g, 'reach'], 1)) for g in picked]}") + + +def _pick_complexes(a, df, n=12, min_members=3): + """EBI complexes with >=min_members generated dots, ranked by tightest median gen->true UMAP displacement.""" + ec = a.obs["ebi_complex"] + g2c = ec.dropna().to_dict() + df = df.copy(); df["complex"] = df["gene"].map(g2c) + df = df.dropna(subset=["complex"]) + df["disp"] = np.hypot(df["gu0"] - df["ru0"], df["gu1"] - df["ru1"]) + agg = df.groupby("complex").agg(n=("gene", "size"), disp=("disp", "median")) + agg = agg[agg["n"] >= min_members].sort_values("disp") + return list(agg.index[:n]), df + + +def plot(ai=14, n_complexes=12, names=None): + import matplotlib + matplotlib.use("Agg") + import matplotlib.pyplot as plt + from matplotlib import cm + from matplotlib.lines import Line2D + plt.rcParams["pdf.fonttype"] = 42 + + a, comp, mean = _load_embedding() + flips = _ntc_flip(a) + df = pd.read_parquet(_pq(ai)) + alpha = float(np.load(_cache_files()[0], allow_pickle=True)["alphas"][ai]) + picked, dfc = _pick_complexes(a, df, n=n_complexes) + if names: + picked = [c for c in dfc["complex"].unique() if any(nm.lower() in c.lower() for nm in names)] + cols = cm.get_cmap("gist_rainbow")(np.linspace(0, 1, len(picked), endpoint=False)) + cmap = dict(zip(picked, cols)) + + for layout, (gx, gy, rx, ry) in {"umap": ("gu0", "gu1", "ru0", "ru1"), + "phate": ("gp0", "gp1", "rp0", "rp1")}.items(): + fx, fy, ntc = flips[layout] + Xr = np.asarray(a.obsm[f"X_{layout}"]) * np.array([fx, fy]) + fig, ax = plt.subplots(figsize=(11, 10)) + ax.scatter(Xr[:, 0], Xr[:, 1], s=10, c="#dddddd", lw=0, zorder=1) # all real genes (grey) + _draw_ntc(ax, ntc) + for cx in picked: + sub = dfc[dfc["complex"] == cx]; col = cmap[cx] + for _, r in sub.iterrows(): + ax.plot([r[rx] * fx, r[gx] * fx], [r[ry] * fy, r[gy] * fy], "-", color=col, lw=0.8, alpha=0.5, zorder=2) + ax.scatter(sub[rx] * fx, sub[ry] * fy, s=95, marker="o", facecolor="none", + edgecolor=col, lw=2.0, zorder=4) # real member gene + ax.scatter(sub[gx] * fx, sub[gy] * fy, s=150, marker="*", color=col, + edgecolor="k", lw=0.6, zorder=5) # generated dot + ax.set_xticks([]); ax.set_yticks([]) + for s in ax.spines.values(): + s.set_visible(False) + ax.set_title(f"Generated geneKO cells projected into real phase {layout.upper()} (α = {alpha:+.1f})\n" + f"★ generated dot ◯ true real gene — same complex", fontsize=13) + handles = [Line2D([0], [0], marker="o", color="w", markerfacecolor=cmap[c], markersize=10, + label=(c[:42] + "…") if len(c) > 43 else c) for c in picked] + ax.legend(handles=handles, loc="center left", bbox_to_anchor=(1.0, 0.5), fontsize=8.5, frameon=False) + fig.tight_layout() + for e in ("png", "svg"): + fig.savefig(f"{OUT}/passthrough_{layout}_a{ai}.{e}", dpi=150, bbox_inches="tight") + plt.close(fig) + print(f"saved passthrough_{layout}_a{ai} ({len(picked)} complexes highlighted)") + json.dump(picked, open(f"{OUT}/picked_a{ai}.json", "w")) + print("picked complexes:", picked) + + +def plot_top_genes(ai, n=5): + """Highlight the n individual geneKOs whose generated dot lands closest to its true real gene.""" + import matplotlib + matplotlib.use("Agg") + import matplotlib.pyplot as plt + from matplotlib import cm + from matplotlib.lines import Line2D + plt.rcParams["pdf.fonttype"] = 42 + + a, comp, mean = _load_embedding() + flips = _ntc_flip(a) + ec, cxcmap = _bg_colors(a) + df = pd.read_parquet(_pq(ai)).copy() + alpha = float(np.load(_cache_files()[0], allow_pickle=True)["alphas"][ai]) + df["disp"] = np.hypot(df["gu0"] - df["ru0"], df["gu1"] - df["ru1"]) + top = df.nsmallest(n, "disp").reset_index(drop=True) + cols = cm.get_cmap("gist_rainbow")(np.linspace(0, 1, n, endpoint=False)) + + for layout, (gx, gy, rx, ry) in {"umap": ("gu0", "gu1", "ru0", "ru1"), + "phate": ("gp0", "gp1", "rp0", "rp1")}.items(): + fx, fy, ntc = flips[layout] + fig, ax = plt.subplots(figsize=(11, 10)) + _draw_bg(ax, a, layout, fx, fy, ec, cxcmap) + _draw_ntc(ax, ntc) + for i, r in top.iterrows(): + col = cols[i] + ax.plot([r[rx] * fx, r[gx] * fx], [r[ry] * fy, r[gy] * fy], "-", color=col, lw=1.0, alpha=0.6, zorder=2) + ax.scatter(r[rx] * fx, r[ry] * fy, s=110, marker="o", facecolor="none", edgecolor=col, lw=2.2, zorder=4) + ax.scatter(r[gx] * fx, r[gy] * fy, s=170, marker="*", color=col, edgecolor="k", lw=0.6, zorder=5) + ax.annotate(r["gene"], (r[gx] * fx, r[gy] * fy), textcoords="offset points", xytext=(6, 6), + fontsize=10, fontweight="bold", color=col, zorder=6) + ax.set_xticks([]); ax.set_yticks([]) + for s in ax.spines.values(): + s.set_visible(False) + ax.set_title(f"Top-{n} closest generated→true geneKOs in real phase {layout.upper()} (α = {alpha:+.1f})\n" + f"★ generated dot ◯ true real gene", fontsize=13) + handles = [Line2D([0], [0], marker="*", color="w", markerfacecolor=cols[i], markersize=13, + label=f"{r['gene']} (rank {int(r['rank'])})") for i, r in top.iterrows()] + handles.append(Line2D([0], [0], marker="X", color="w", markerfacecolor="#8b0000", markersize=13, label="NTC anchor")) + ax.legend(handles=handles, loc="center left", bbox_to_anchor=(1.0, 0.5), fontsize=10, frameon=False) + fig.tight_layout() + for e in ("png", "svg"): + fig.savefig(f"{OUT}/topgenes_{layout}_a{ai}.{e}", dpi=150, bbox_inches="tight") + plt.close(fig) + print(f"saved topgenes_{layout}_a{ai}: {list(top['gene'])}") + + +if __name__ == "__main__": + import sys + ai = int(sys.argv[sys.argv.index("--ai") + 1]) if "--ai" in sys.argv else 14 + if "--diverse" in sys.argv: + compute_many() + plot_diverse(n=10) + elif "--traversal" in sys.argv: + compute_many() + plot_traversal(n_complex=10) + elif "--topgenes" in sys.argv: + if not os.path.exists(_pq(ai)): + compute(ai=ai) + plot_top_genes(ai=ai, n=5) + elif "--plot" in sys.argv: + plot(ai=ai) + else: + compute(ai=ai) + plot(ai=ai) diff --git a/src/ops_model/models/attention/diffex/figures/gen_validation/gen_real_centroid.py b/src/ops_model/models/attention/diffex/figures/gen_validation/gen_real_centroid.py new file mode 100644 index 0000000..16c5bcc --- /dev/null +++ b/src/ops_model/models/attention/diffex/figures/gen_validation/gen_real_centroid.py @@ -0,0 +1,360 @@ +"""Faithful real class centroids in the GENERATED (embed_crops CellDINO) space. + +Re-embed the top-N accuracy real cells per class with the SAME encoder the generated cells use (embed_crops via +_gather_class) and average → a low-variance centroid (fixes v2's noisy 30-cell centroid). Then score the cached +generated cells (nearest faithful centroid) → top-1/top-5 vs α, with the v2 per-domain standardization (gen vs +pooled gen-α0, centroids vs the real population). The training .pt store can't substitute — it's z-scored, a +different space than embed_crops (empirically 5-11% nearest-centroid, no better than raw). +""" +import os, json, glob +import numpy as np + +V5 = "/hpc/projects/icd.fast.ops/models/diffex/viewer_assets_v5" +RANK = {"geneKO": f"{V5}/_rankings/pma_v5_phase_geneKO.parquet", + "complex": f"{V5}/_rankings/pma_v5_phase_complex.parquet"} +CACHE = os.environ.get("GRC_CACHE", "/hpc/projects/icd.fast.ops/analysis/figure4_traversals/gen_real_map_cache") +OUT = os.environ.get("GRC_OUT", "/hpc/projects/icd.fast.ops/analysis/figure4_traversals/gen_real_centroid") +N = 1000 # top-N real cells per class → faithful centroid +PER_SHARD = 12 # classes per GPU shard (fine → high parallelism) + + +def _classes(grain): + # read the ORIGINAL class name from each npz (filenames are slugified; complex names have spaces) + return sorted(str(np.load(f, allow_pickle=True)["gene"]) for f in glob.glob(f"{CACHE}/{grain}/*.npz")) + + +def embed_centroids(grain, classes): + from ops_model.models.attention.diffex.viewer.precompute import _gather_class + from ops_model.models.attention.diffex.directions.config import DirConfig + from ops_model.models.attention.diffex.classifier.config import slugify + cfg = DirConfig(grain=grain, target=classes[0], device="cuda"); cfg.num_workers = 12 + cents, S, SS, n = {}, np.zeros(1024), np.zeros(1024), 0 + embs, lbl = [], [] + for c in classes: + try: + _, e = _gather_class(cfg, c, N, parquet=RANK[grain]) + e = np.asarray(e, np.float32) + cents[c] = e.mean(0).astype(np.float32) + embs.append(e); lbl += [c] * len(e) # keep individual embeddings for the proper-mAP gallery + S += e.sum(0); SS += (e.astype(np.float64) ** 2).sum(0); n += len(e) + except Exception as ex: + print(f"skip {c}: {ex}") + os.makedirs(f"{OUT}/shards", exist_ok=True) + names = list(cents) + np.savez(f"{OUT}/shards/{grain}_{slugify(classes[0])}.npz", + names=np.array(names), cents=np.stack([cents[k] for k in names]) if names else np.zeros((0, 1024)), + embs=np.concatenate(embs).astype(np.float32) if embs else np.zeros((0, 1024), np.float32), + labels=np.array(lbl), S=S, SS=SS, n=n) + return {"grain": grain, "n": len(cents), "cells": n} + + +def merge(grain): + names, C, S, SS, n = [], [], np.zeros(1024), np.zeros(1024), 0 + for f in glob.glob(f"{OUT}/shards/{grain}_*.npz"): + d = np.load(f, allow_pickle=True) + if len(d["names"]): + names += list(d["names"]); C.append(d["cents"]) + S += d["S"]; SS += d["SS"]; n += int(d["n"]) + C = np.concatenate(C); mu = S / n; sd = np.sqrt(np.clip(SS / n - mu ** 2, 1e-12, None)) + 1e-6 + np.savez(f"{OUT}/{grain}_centroids.npz", names=np.array(names), cents=C, mu=mu.astype(np.float32), sd=sd.astype(np.float32)) + return {"grain": grain, "classes": len(names), "cells": n} + + +def score(grain, cache=None, out=None, cap=None): + from ops_model.models.attention.diffex.classifier.config import slugify + CACHE_ = cache or CACHE; OUT_ = out or OUT # explicit args (SLURM-safe) override module globals + cap = cap if cap is not None else (int(os.environ.get("GRC_GEN_CAP", "0")) or None) # subsample gen bag to first `cap` cells/class + d = np.load(f"{OUT_}/{grain}_centroids.npz", allow_pickle=True) + names = list(d["names"]); cidx = {slugify(str(c)): i for i, c in enumerate(names)} # slug-key: cache genes may be slugged + mu_r, sd_r = d["mu"], d["sd"] + cz = (d["cents"] - mu_r) / sd_r + cz = cz / (np.linalg.norm(cz, axis=1, keepdims=True) + 1e-9) + caches = sorted(glob.glob(f"{CACHE_}/{grain}/*.npz")) + g0 = [] + for f in caches: + dd = np.load(f, allow_pickle=True) + a0i = int(np.argmin(np.abs(np.asarray(dd["alphas"], float)))) # α=0 baseline (grid may be one-sided, not symmetric) + z0 = dd["gen"][a0i] + if z0 is not None and len(z0): + g0.append(np.asarray(z0, np.float32)[:cap]) + g0 = np.concatenate(g0) + mu_g, sd_g = g0.mean(0), g0.std(0) + 1e-6 + by = {} + for f in caches: + dd = np.load(f, allow_pickle=True); g = str(dd["gene"]); al = dd["alphas"] + if slugify(g) not in cidx: + continue + ti = cidx[slugify(g)] + for ai in range(len(al)): + gv = dd["gen"][ai] + if gv is None or not len(gv): + continue + gz = (np.asarray(gv, np.float32)[:cap] - mu_g) / sd_g + gz = gz / (np.linalg.norm(gz, axis=1, keepdims=True) + 1e-9) + order = np.argsort(-(gz @ cz.T), axis=1) + rank_true = np.where(order == ti)[1] + 1 # 1-based rank of the true class centroid per cell + a = float(al[ai]); by.setdefault(a, {"top1": {}, "top5": {}, "map": {}}) + by[a]["top1"][g] = float(np.mean(order[:, 0] == ti)) + by[a]["top5"][g] = float(np.mean([ti in r[:5] for r in order])) + by[a]["map"][g] = float(np.mean(1.0 / rank_true)) # retrieval AP (1 positive = true centroid) → mean = mAP + json.dump({"alphas": sorted(by), "by_alpha": by, "n_classes": len(cidx)}, + open(f"{OUT_}/{grain}_scored.json", "w")) + return {"grain": grain, "n": len(cidx)} + + +DIST = "/hpc/projects/icd.fast.ops/analysis/figure4_traversals/gen_real_distinct" # real distinctiveness@20 per class + + +def ceiling(grain): + """Real cells (cached embed_crops) → nearest faithful centroid: per-class real mAP/top1/top5 (the ceiling).""" + from ops_model.models.attention.diffex.classifier.config import slugify + d = np.load(f"{OUT}/{grain}_centroids.npz", allow_pickle=True) + names = list(d["names"]); cidx = {slugify(str(c)): i for i, c in enumerate(names)} + cz = (d["cents"] - d["mu"]) / d["sd"]; cz = cz / (np.linalg.norm(cz, axis=1, keepdims=True) + 1e-9) + out = {} + for f in glob.glob(f"{CACHE}/{grain}/*.npz"): + dd = np.load(f, allow_pickle=True); c = str(dd["gene"]) + if slugify(c) not in cidx: + continue + ti = cidx[slugify(c)]; rz = (np.asarray(dd["real"], np.float32) - d["mu"]) / d["sd"] + rz = rz / (np.linalg.norm(rz, axis=1, keepdims=True) + 1e-9) + order = np.argsort(-(rz @ cz.T), axis=1); rank = np.where(order == ti)[1] + 1 + out[c] = {"top1": float(np.mean(order[:, 0] == ti)), "top5": float(np.mean([ti in o[:5] for o in order])), + "map": float(np.mean(1.0 / rank))} + json.dump(out, open(f"{OUT}/{grain}_ceiling.json", "w")) + return {"grain": grain, "n": len(out)} + + +def _acc_keep(grain, acc_thr): + """SetTransformer-distinguishable subset: real cells top1_acc > acc_thr @ bag20 (real_acc20.json), slug-keyed.""" + R = json.load(open("/hpc/projects/icd.fast.ops/models/diffex/viewer_assets_v5/real_acc20.json")) + pre = f"phase/{grain}/" + return {k[len(pre):] for k, v in R.items() if k.startswith(pre) and v > acc_thr} + + +def plot(min_dist=None, acc_thr=None, fname="centroid_topk", overlay=None, overlay_label="200-cell bag"): + """3 scores (mAP, top-1, top-5) vs α from the faithful centroids. min_dist: restrict to classes whose real + distinctiveness/EBI mAP@20 > min_dist. acc_thr: restrict to real top1_acc>acc_thr @bag20 (SetTransformer + subset). overlay: {grain: scored_dir} → dashed second line (e.g. the 200-cell bag) on that grain's axis.""" + from ops_model.models.attention.diffex.classifier.config import slugify + import matplotlib + matplotlib.use("Agg") + import matplotlib.pyplot as plt + plt.rcParams["pdf.fonttype"] = 42 + C1, C5, CM = "#c0392b", "#27ae60", "#2471a3" + fig, axes = plt.subplots(1, 2, figsize=(13, 5), sharey=True) + for ax, (grain, lbl) in zip(axes, [("geneKO", "Gene-level"), ("complex", "Protein complex")]): + d = json.load(open(f"{OUT}/{grain}_scored.json")); al = d["alphas"]; ba = d["by_alpha"] + keep = None + if acc_thr is not None: + ks = _acc_keep(grain, acc_thr); keep = {"__slug__"}; keep = ks # slug-keyed + elif min_dist is not None: + rd = json.load(open(f"{DIST}/{grain}_real.json")) + keep = {c for c, v in rd.items() if v > min_dist} + ink = lambda c: keep is None or c in keep or slugify(c) in keep + def vals_of(ba_, a, key): + return [v for c, v in ba_[str(a)][key].items() if ink(c)] + vals = lambda a, key: vals_of(ba, a, key) + n = len(vals(al[0], "top1")) + sem = lambda a, key: np.std(vals(a, key), ddof=1) / np.sqrt(max(len(vals(a, key)), 1)) + ceil = json.load(open(f"{OUT}/{grain}_ceiling.json")) if os.path.exists(f"{OUT}/{grain}_ceiling.json") else None + ov = None; n_ov = 0 + if overlay and grain in overlay: + ov = json.load(open(f"{overlay[grain]}/{grain}_scored.json")); n_ov = len(vals_of(ov["by_alpha"], ov["alphas"][0], "top1")) + gtag = "" if ov is None else " (45-cell)" + for key, col, name in [("map", CM, "MRR"), ("top1", C1, "top-1"), ("top5", C5, "top-5")]: + m = np.array([np.mean(vals(a, key)) for a in al]); se = np.array([sem(a, key) for a in al]) + ax.plot(al, m, "-", color=col, lw=2.6, label=f"{name} — generated{gtag}") + ax.fill_between(al, m - se, m + se, color=col, alpha=.18, lw=0) # mean ± SEM across classes + if ov is not None: # dashed = overlay bag (same centroids, more cells) + alo = ov["alphas"]; mo = np.array([np.mean(vals_of(ov["by_alpha"], a, key)) for a in alo]) + ax.plot(alo, mo, "--", color=col, lw=2.2, label=f"{name} — {overlay_label}") + if ceil: # dotted = real-cell ceiling (same centroids) + cvs = [v[key] for c, v in ceil.items() if ink(c)] + cv = np.mean(cvs); cse = np.std(cvs, ddof=1) / np.sqrt(max(len(cvs), 1)) + ax.axhline(cv, color=col, ls=":", lw=2.0, label=f"{name} — real ceiling") + ax.axhspan(cv - cse, cv + cse, color=col, alpha=.07, lw=0) + ax.axvline(0, color="#ccc", lw=1); ax.axvline(1, color="#bbb", lw=1, ls="--") + nlab = f"n={n}" if ov is None else f"n={n} (45c) / {n_ov} (200c)" + ax.set_title(f"{lbl} ({nlab})"); ax.set_xlabel("traversal α"); ax.grid(alpha=.25) + axes[0].set_ylabel(f"generated → nearest real centroid (top-{N} cells)"); axes[0].set_ylim(-0.02, 1.02); axes[0].legend(fontsize=9, ncol=2) + gate = "" if acc_thr is None else f" · real top1_acc@20 > {acc_thr}" + if min_dist is not None: + gate = f" · real distinctiveness@20 > {min_dist}" + fig.suptitle(f"Generated cells → faithful real centroids (top-{N} real cells/class){gate}", fontweight="bold") + fig.tight_layout() + for e in ("png", "svg"): + fig.savefig(f"{OUT}/{fname}.{e}", dpi=150, bbox_inches="tight") + plt.close(fig); print(f"saved {fname}") + + +def merge_embs(grain): + feats, lbl = [], [] + for f in glob.glob(f"{OUT}/shards/{grain}_*.npz"): + d = np.load(f, allow_pickle=True) + if len(d["labels"]): + feats.append(d["embs"]); lbl += [str(x) for x in d["labels"]] + feats = np.concatenate(feats).astype(np.float32) + np.savez(f"{OUT}/{grain}_embs.npz", feats=feats, labels=np.array(lbl)) + return {"grain": grain, "gallery_cells": len(feats), "classes": len(set(lbl))} + + +def proper_map_gpu(grain): + """Proper retrieval mAP with the full 1000-cell/class real gallery, on GPU (no copairs, no permutations). + gen queries → real gallery, AP over same-class positives; ceiling = held-out real queries (self masked).""" + import torch + dev = "cuda" if torch.cuda.is_available() else "cpu" + E = np.load(f"{OUT}/{grain}_embs.npz", allow_pickle=True) + cen = np.load(f"{OUT}/{grain}_centroids.npz", allow_pickle=True) + mu, sd = cen["mu"], cen["sd"] + labels = E["labels"]; cls = sorted(set(labels)); c2i = {c: i for i, c in enumerate(cls)} + G = (E["feats"] - mu) / sd + Gt = torch.tensor(G, device=dev, dtype=torch.float32); Gt = Gt / Gt.norm(dim=1, keepdim=True).clamp_min(1e-9) + glab = torch.tensor([c2i[c] for c in labels], device=dev) + + def ap(Q, ti, self_idx=None): # Q:(b,D) normalized, ti:(b,) class idx; self_idx:(b,) gallery row to mask + out = [] + for s in range(0, len(Q), 256): + q = Q[s:s + 256]; sim = q @ Gt.T + if self_idx is not None: + sim[torch.arange(len(q)), self_idx[s:s + 256]] = -1e9 + order = torch.argsort(sim, dim=1, descending=True) + pos = (glab[order] == ti[s:s + 256, None]).float() + csum = torch.cumsum(pos, 1); ranks = torch.arange(1, pos.shape[1] + 1, device=dev).float() + a = ((csum / ranks) * pos).sum(1) / pos.sum(1).clamp_min(1) + out.append(a.cpu().numpy()) + return np.concatenate(out) + + caches = sorted(glob.glob(f"{CACHE}/{grain}/*.npz")) + g0 = [] + for f in caches: + dd = np.load(f, allow_pickle=True) + a0i = int(np.argmin(np.abs(np.asarray(dd["alphas"], float)))) # α=0 baseline (grid may be one-sided, not symmetric) + z0 = dd["gen"][a0i] + if z0 is not None and len(z0): + g0.append(np.asarray(z0, np.float32)) + g0 = np.concatenate(g0); mu_g, sd_g = g0.mean(0), g0.std(0) + 1e-6 + gen = {} + for f in caches: + dd = np.load(f, allow_pickle=True); g = str(dd["gene"]); al = dd["alphas"] + if g not in c2i: + continue + ti0 = c2i[g] + for ai in range(len(al)): + gv = dd["gen"][ai] + if gv is None or not len(gv): + continue + gz = (np.asarray(gv, np.float32) - mu_g) / sd_g + Q = torch.tensor(gz, device=dev, dtype=torch.float32); Q = Q / Q.norm(dim=1, keepdim=True).clamp_min(1e-9) + aps = ap(Q, torch.full((len(Q),), ti0, device=dev)) + a = float(al[ai]); gen.setdefault(a, {})[g] = float(aps.mean()) + # ceiling: sample 30 real cells/class as queries vs full gallery, mask self + ceil = {} + rng = np.random.default_rng(0) + for c in cls: + idx = np.where(labels == c)[0] + qi = idx[:30] # first 30 gallery cells of the class as held-in queries (self masked) + Q = Gt[torch.tensor(qi, device=dev)] + aps = ap(Q, torch.full((len(qi),), c2i[c], device=dev), self_idx=torch.tensor(qi, device=dev)) + ceil[c] = float(aps.mean()) + json.dump({"alphas": sorted(gen), "gen": gen, "ceiling": ceil, "n_classes": len(c2i), + "gallery_cells": len(labels)}, open(f"{OUT}/{grain}_propermap1k.json", "w")) + return {"grain": grain, "n": len(c2i), "gallery": len(labels)} + + +def plot_proper1k(min_dist=None, fname="mAP_proper1k"): + import matplotlib + matplotlib.use("Agg") + import matplotlib.pyplot as plt + plt.rcParams["pdf.fonttype"] = 42 + CM = "#2471a3" + fig, axes = plt.subplots(1, 2, figsize=(13, 5)) + for ax, (grain, lbl) in zip(axes, [("geneKO", "Gene-level"), ("complex", "Protein complex")]): + d = json.load(open(f"{OUT}/{grain}_propermap1k.json")); al = d["alphas"]; gen = d["gen"]; ceil = d["ceiling"] + keep = None + if min_dist is not None: + rd = json.load(open(f"{DIST}/{grain}_real.json")); keep = {c for c, v in rd.items() if v > min_dist} + def vals(a): + return [v for c, v in gen[str(a)].items() if keep is None or c in keep] + m = np.array([np.mean(vals(a)) for a in al]); se = np.array([np.std(vals(a), ddof=1) / np.sqrt(len(vals(a))) for a in al]) + ax.plot(al, m, "-", color=CM, lw=2.6, label="generated (mAP, 1000-cell gallery)") + ax.fill_between(al, m - se, m + se, color=CM, alpha=.18, lw=0) + cvs = [v for c, v in ceil.items() if keep is None or c in keep] + cv = np.mean(cvs); cse = np.std(cvs, ddof=1) / np.sqrt(len(cvs)) + ax.axhline(cv, color=CM, ls=":", lw=2.0, label=f"real ceiling ({cv:.3f})"); ax.axhspan(cv - cse, cv + cse, color=CM, alpha=.07, lw=0) + ax.axvline(0, color="#ccc", lw=1); ax.axvline(1, color="#bbb", lw=1, ls="--") + n = len(vals(al[0])) + ax.set_title(f"{lbl} (n={n})"); ax.set_xlabel("traversal α"); ax.grid(alpha=.25) + axes[0].set_ylabel("real↔generated retrieval mAP (copairs-style)") + gate = "" if min_dist is None else f" · real distinctiveness@20 > {min_dist}" + fig.suptitle(f"Proper retrieval mAP: generated → full real gallery (1000 cells/class){gate}", fontweight="bold") + fig.tight_layout() + for e in ("png", "svg"): + fig.savefig(f"{OUT}/{fname}.{e}", dpi=150, bbox_inches="tight") + plt.close(fig); print(f"saved {fname}") + + +V2 = "/hpc/projects/icd.fast.ops/analysis/figure4_traversals/gen_real_map_v2" # proper copairs retrieval mAP (gen→real cells) + + +def plot_map_proper(min_dist=None, fname="mAP_proper"): + """Proper copairs retrieval mAP (generated cells → real CELLS gallery, multiple positives) from the v2 data, + same design as centroid_topk: solid generated (mean±SEM) + dotted real split-half ceiling. Own y-scale per grain.""" + import matplotlib + matplotlib.use("Agg") + import matplotlib.pyplot as plt + import pandas as pd + plt.rcParams["pdf.fonttype"] = 42 + CM = "#2471a3" + fig, axes = plt.subplots(1, 2, figsize=(13, 5)) + for ax, (grain, lbl) in zip(axes, [("geneKO", "Gene-level"), ("complex", "Protein complex")]): + files = sorted(glob.glob(f"{V2}/{grain}_a*.json"), key=lambda p: int(p.split("_a")[-1][:-5])) + al, cols = [], {} + for f in files: + d = json.load(open(f)); al.append(d["alpha"]); cols[d["alpha"]] = d["gen_map"] + df = pd.DataFrame(cols) # index=class, cols=alpha + keep = None + if min_dist is not None: + rd = json.load(open(f"{DIST}/{grain}_real.json")); keep = {c for c, v in rd.items() if v > min_dist} + df = df[df.index.map(lambda c: keep is not None and c in keep)] + m = df[al].mean(0).values; se = df[al].std(0, ddof=1).values / np.sqrt(len(df)) + ax.plot(al, m, "-", color=CM, lw=2.6, label="generated (copairs mAP)") + ax.fill_between(al, m - se, m + se, color=CM, alpha=.18, lw=0) + ceil = json.load(open(f"{V2}/{grain}_ceiling.json"))["map"] + cvs = [v for c, v in ceil.items() if keep is None or c in keep] + cv = np.mean(cvs); cse = np.std(cvs, ddof=1) / np.sqrt(len(cvs)) + ax.axhline(cv, color=CM, ls=":", lw=2.0, label=f"real split-half ceiling ({cv:.3f})") + ax.axhspan(cv - cse, cv + cse, color=CM, alpha=.07, lw=0) + ax.axvline(0, color="#ccc", lw=1); ax.axvline(1, color="#bbb", lw=1, ls="--") + ax.set_title(f"{lbl} (n={len(df)})"); ax.set_xlabel("traversal α"); ax.grid(alpha=.25) + axes[0].set_ylabel("real↔generated retrieval mAP (copairs)") + gate = "" if min_dist is None else f" · real distinctiveness@20 > {min_dist}" + fig.suptitle(f"Proper retrieval mAP: generated cells → real-cell gallery (30/class){gate}", fontweight="bold") + fig.tight_layout() + for e in ("png", "svg"): + fig.savefig(f"{OUT}/{fname}.{e}", dpi=150, bbox_inches="tight") + plt.close(fig); print(f"saved {fname}") + + +def main(): + from ops_utils.hpc.slurm_batch_utils import submit_parallel_jobs + jobs = [] + for grain in ["geneKO", "complex"]: + cls = _classes(grain) + for i in range(0, len(cls), PER_SHARD): + jobs.append({"name": f"grc_{grain}_{i}", "func": embed_centroids, "kwargs": {"grain": grain, "classes": cls[i:i + PER_SHARD]}}) + print(f"[gen-real-centroid] {len(jobs)} embed shards (N={N}, {PER_SHARD}/shard)") + submit_parallel_jobs(jobs, experiment="gen_real_centroid", + slurm_params={"slurm_partition": "gpu", "slurm_gres": "gpu:1", "cpus_per_task": 12, + "mem_gb": 64, "timeout_min": 300}, log_dir="gen_real_centroid", wait_for_completion=False) + + +if __name__ == "__main__": + import sys + if "--merge-score-plot" in sys.argv: + for g in ["geneKO", "complex"]: + print(merge(g)); print(score(g)); print(ceiling(g)) + plot() + plot(min_dist=0.1, fname="centroid_topk_dist01") + else: + main() diff --git a/src/ops_model/models/attention/diffex/figures/gen_validation/gen_real_distinct.py b/src/ops_model/models/attention/diffex/figures/gen_validation/gen_real_distinct.py new file mode 100644 index 0000000..1ca5545 --- /dev/null +++ b/src/ops_model/models/attention/diffex/figures/gen_validation/gen_real_distinct.py @@ -0,0 +1,182 @@ +"""Real-vs-Generated DISTINCTIVENESS mAP (within-domain), across all α. + +Per grain (geneKO = gene-level distinctiveness, complex = EBI): take the top-K accuracy cells per class and measure +the standard copairs distinctiveness mAP (each class's cells retrieve each other vs all other classes) on the REAL +cells (once) and on the GENERATED cells at EVERY α. Reuses gen_real_map_cache embeddings; no re-embedding. + +Within-domain, so the DiffAE domain offset (common to all generated cells) cancels — no standardization needed. +Outputs: {grain}_real.json + {grain}_gen_a{ai}.json → plot: gen distinctiveness vs α + real reference, and a +Real-vs-Generated violin at the α where generated distinctiveness peaks. +""" +import os, json, glob +import numpy as np +import pandas as pd + +CACHE = os.environ.get("GRD_CACHE", "/hpc/projects/icd.fast.ops/analysis/figure4_traversals/gen_real_map_cache") +OUT = os.environ.get("GRD_OUT", "/hpc/projects/icd.fast.ops/analysis/figure4_traversals/gen_real_distinct") +K = int(os.environ.get("GRD_K", 20)) # top-K cells/class (200 for the valid200 bag) +NA = int(os.environ.get("GRD_NA", 17)) # α grid points (7 for valid200) + + +def _distinct(feats, labels): + """copairs distinctiveness: per class, do its cells retrieve each other above all other-class cells → {class: mAP}.""" + from copairs import map as cm + meta = pd.DataFrame({"gene": labels}) + ap = cm.average_precision(meta, np.ascontiguousarray(feats, np.float32), pos_sameby=["gene"], pos_diffby=[], + neg_sameby=[], neg_diffby=["gene"], distance="cosine") + m = cm.mean_average_precision(ap, sameby=["gene"], null_size=200, threshold=0.05, seed=0) + return dict(zip(m["gene"], m["mean_average_precision"])) + + +def _load(grain): + return sorted(glob.glob(f"{CACHE}/{grain}/*.npz")) + + +def compute_real(grain): + rf, rl = [], [] + for c in _load(grain): + d = np.load(c, allow_pickle=True); g = str(d["gene"]) + r = np.asarray(d["real"], np.float32)[:K]; rf.append(r); rl += [g] * len(r) + os.makedirs(OUT, exist_ok=True) + real = _distinct(np.concatenate(rf), rl) + json.dump(real, open(f"{OUT}/{grain}_real.json", "w")) + return {"grain": grain, "n": len(real), "median": float(np.median(list(real.values())))} + + +def compute_gen(grain, ai): + gf, gl, alpha = [], [], None + for c in _load(grain): + d = np.load(c, allow_pickle=True); g = str(d["gene"]); alpha = float(d["alphas"][ai]) + gv = d["gen"][ai] + if gv is not None and len(gv): + gv = np.asarray(gv, np.float32)[:K]; gf.append(gv); gl += [g] * len(gv) + os.makedirs(OUT, exist_ok=True) + gen = _distinct(np.concatenate(gf), gl) + json.dump({"alpha": alpha, "gen": gen}, open(f"{OUT}/{grain}_gen_a{ai}.json", "w")) + return {"grain": grain, "ai": ai, "alpha": alpha, "n": len(gen), "median": float(np.median(list(gen.values())))} + + +_ALLGRAINS = [("geneKO", "Gene-level"), ("complex", "Protein\ncomplex")] +_KEYS = os.environ.get("GRD_GRAINS", "geneKO,complex").split(",") # restrict grains (valid200 = geneKO only) +GRAINS = [g for g in _ALLGRAINS if g[0] in _KEYS] + + +def _series(grain, base=None): + """→ (alphas, DataFrame[class × alpha] gen mAP, real dict).""" + base = base or OUT + real = json.load(open(f"{base}/{grain}_real.json")) + al, cols = [], {} + for f in sorted(glob.glob(f"{base}/{grain}_gen_a*.json"), key=lambda p: int(p.split("_a")[-1][:-5])): + d = json.load(open(f)); al.append(d["alpha"]); cols[d["alpha"]] = d["gen"] + return sorted(al), pd.DataFrame(cols), real + + +def plot(overlay=None, overlay_label="200-cell bag", overlay_K=50, fname="distinct_vs_alpha"): + """overlay: {grain: distinct_dir} → dashed second median line (e.g. the 200-cell bag) on that grain's axis.""" + import matplotlib + matplotlib.use("Agg") + import matplotlib.pyplot as plt + import matplotlib.patches as mp + plt.rcParams["pdf.fonttype"] = 42 + REAL_C, GEN_C, OV_C = "#8fa9c9", "#7fbf9a", "#2e8b57" + + # (1) curve: generated distinctiveness (median, IQR) vs α + real median reference + fig, axes = plt.subplots(1, 2, figsize=(13, 5), sharey=True) + peak_ai = {} + for ax, (grain, lbl) in zip(axes, GRAINS): + al, df, real = _series(grain) + med = df[al].median(0); q1 = df[al].quantile(.25); q3 = df[al].quantile(.75) + gtag = "" if not (overlay and grain in overlay) else f" (top-{K})" + ax.plot(al, med, "-", color=GEN_C, lw=2.6, label=f"Generated (median, IQR){gtag}") + ax.fill_between(al, q1, q3, color=GEN_C, alpha=.22, lw=0) + rm = np.median(list(real.values())) + ax.axhline(rm, color=REAL_C, ls="--", lw=2.2, label=f"Real top-{K} ({rm:.2f})") + if overlay and grain in overlay: # dashed = overlay bag (its own α grid + real) + alo, dfo, realo = _series(grain, base=overlay[grain]) + medo = dfo[alo].median(0) + ax.plot(alo, medo, "--", color=OV_C, lw=2.4, label=f"Generated — {overlay_label} (top-{overlay_K})") + ax.axhline(np.median(list(realo.values())), color=OV_C, ls=":", lw=1.6, + label=f"Real top-{overlay_K} ({np.median(list(realo.values())):.2f})") + ax.axvline(0, color="#ccc", lw=1); ax.axvline(1, color="#27ae60", lw=1, ls=":") + ax.set_title(lbl.replace("\n", " ")); ax.set_xlabel("traversal α"); ax.grid(alpha=.25) + peak_ai[grain] = int(np.argmax(med.values)) + axes[0].set_ylabel("distinctiveness / EBI mAP"); axes[0].set_ylim(-0.02, 1.02); axes[0].legend(fontsize=9) + fig.suptitle(f"Generated distinctiveness vs α (within-domain copairs, top-{K} cells/class)", fontweight="bold") + fig.tight_layout() + for e in ("png", "svg"): + fig.savefig(f"{OUT}/{fname}.{e}", dpi=150, bbox_inches="tight") + plt.close(fig); print(f"saved {fname}") + + # (2) Real-vs-Generated violin at the peak-α of each grain + fig, ax = plt.subplots(figsize=(7, 5.5)) + xt, xl = [], [] + for gi, (grain, lbl) in enumerate(GRAINS): + al, df, real = _series(grain) + ga = al[peak_ai[grain]] + pairs = [("Real", np.array(list(real.values())), REAL_C), ("Gen", df[ga].dropna().values, GEN_C)] + base = gi * 3 + for j, (_, vals, col) in enumerate(pairs): + pos = base + j + vp = ax.violinplot([vals], positions=[pos], widths=0.85, showextrema=False) + for b in vp["bodies"]: + b.set_facecolor(col); b.set_edgecolor("none"); b.set_alpha(0.9) + ax.hlines(np.median(vals), pos - 0.42, pos + 0.42, color="k", lw=3, zorder=5) + xt.append(base + 0.5); xl.append(f"{lbl}\n(gen α={ga:.0f})") + ax.set_xticks(xt); ax.set_xticklabels(xl, fontsize=13) + ax.set_ylabel("mAP score (distinctiveness / EBI)", fontsize=15); ax.set_ylim(-0.02, 1.02); ax.tick_params(labelsize=12) + for s in ("top", "right"): + ax.spines[s].set_visible(False) + ax.legend(handles=[mp.Patch(color=REAL_C, label=f"Real (top-{K} acc cells)"), + mp.Patch(color=GEN_C, label="Generated (peak α)"), + plt.Line2D([0], [0], color="k", lw=3, label="Median")], + loc="center left", bbox_to_anchor=(1.0, 0.5), fontsize=12, frameon=False) + fig.tight_layout() + for e in ("png", "svg"): + fig.savefig(f"{OUT}/distinct_violin.{e}", dpi=150, bbox_inches="tight") + plt.close(fig); print("saved distinct_violin") + + # (3) per-class scatter: real vs generated distinctiveness at peak α (paired by class) + from scipy.stats import spearmanr, pearsonr + fig, axes = plt.subplots(1, 2, figsize=(13, 6)) + for ax, (grain, lbl) in zip(axes, GRAINS): + al, df, real = _series(grain) + ga = al[peak_ai[grain]] + gen = df[ga].to_dict() + cls = [c for c in real if c in gen and not np.isnan(gen[c])] + x = np.array([real[c] for c in cls]); y = np.array([gen[c] for c in cls]) + ax.scatter(x, y, s=12, alpha=.5, color="#555") + lim = max(x.max(), y.max()) * 1.05 + ax.plot([0, lim], [0, lim], "--", color="#c0392b", lw=1.5, label="y = x") + rho = spearmanr(x, y)[0]; r = pearsonr(x, y)[0] + ax.set_title(f"{lbl.replace(chr(10),' ')} (gen α={ga:.0f})\nSpearman ρ={rho:.2f}, Pearson r={r:.2f}, n={len(cls)}") + ax.set_xlabel("Real distinctiveness mAP"); ax.set_ylabel("Generated distinctiveness mAP") + ax.set_xlim(-0.02, lim); ax.set_ylim(-0.02, lim); ax.grid(alpha=.25); ax.legend(fontsize=10) + fig.suptitle(f"Per-class distinctiveness: Real vs Generated (top-{K} cells/class)", fontweight="bold") + fig.tight_layout() + for e in ("png", "svg"): + fig.savefig(f"{OUT}/distinct_scatter.{e}", dpi=150, bbox_inches="tight") + plt.close(fig); print("saved distinct_scatter") + + +def main(): + from ops_utils.hpc.slurm_batch_utils import submit_parallel_jobs + keys = [g[0] for g in GRAINS] + mem = int(os.environ.get("GRD_MEM", 240)) # bump for high-K geneKO (1000-class copairs OOMs >240) + tmin = int(os.environ.get("GRD_TIME", 180)) + gen_only = os.environ.get("GRD_GEN_ONLY") == "1" # skip real (already computed; uses only 30 cached cells) + ais = [int(x) for x in os.environ["GRD_AIS"].split(",")] if os.environ.get("GRD_AIS") else list(range(NA)) # resubmit specific α + jobs = [] if gen_only else [{"name": f"grd_real_{g}", "func": compute_real, "kwargs": {"grain": g}} for g in keys] + for g in keys: + jobs += [{"name": f"grd_{g}_{ai}", "func": compute_gen, "kwargs": {"grain": g, "ai": ai}} for ai in ais] + print(f"[gen-real-distinct] {len(jobs)} jobs · mem={mem}GB · grains={keys}") + submit_parallel_jobs(jobs, experiment="gen_real_distinct", + slurm_params={"slurm_partition": "cpu", "cpus_per_task": 16, "mem_gb": mem, "timeout_min": tmin}, + log_dir="gen_real_distinct", wait_for_completion=False) + + +if __name__ == "__main__": + import sys + if "--plot" in sys.argv: + plot() + else: + main() diff --git a/src/ops_model/models/attention/diffex/figures/gen_validation/ntc_inverse_gap.py b/src/ops_model/models/attention/diffex/figures/gen_validation/ntc_inverse_gap.py new file mode 100644 index 0000000..6b0ecaf --- /dev/null +++ b/src/ops_model/models/attention/diffex/figures/gen_validation/ntc_inverse_gap.py @@ -0,0 +1,175 @@ +"""Real NTC vs inverse-α=0 NTC in CellDINO space — paired, same-cell. + +Uses the NEW DDIM-inverted v5 traversals (viewer_assets_v5_inv), NOT the old random-xT cache. +For the phase channel, geneKO anchors are NTC cells: anchor cellN (real.webp) and any gene's traversal cellN +frame_08 (α=0) are the SAME cell — so we get matched pairs (real NTC vs its α=0 reconstruction). We measure: + (1) per-cell gap: cosine(real_i, gen_i); + (2) offset consistency: are the N offset vectors (gen_i - real_i) parallel + equal-magnitude (rigid) or scattered; + (3) where each lands in the phase gene embedding (real-pop standardization) — does inverse-α=0 now sit on NTC? +α=0 is gene-independent (no direction), so one finished phase gene suffices. +""" +import os, glob, json +import numpy as np +from PIL import Image + +INV = "/hpc/projects/icd.fast.ops/models/diffex/viewer_assets_v5_inv" +A0 = 8 # α=0 frame index (17 frames, -5..+5) +OUT = "/hpc/projects/icd.fast.ops/analysis/figure4_embedding/ntc_inverse_gap" + + +def _celln(d): + return sorted(glob.glob(f"{d}/cell*"), key=lambda p: int(p.rsplit("cell", 1)[1])) + + +def _emb_webp(paths, cfg, embed_crops): + imgs = [np.asarray(Image.open(p).convert("L"), np.float32) / 255.0 * 2 - 1 for p in paths] + return embed_crops(np.stack(imgs)[:, None].astype(np.float32), cfg, cache_path=None) + + +def _first_done(channel=None): + """Find (channel, gene) of a finished geneKO traversal with α=0 frames + NTC anchors. Prefer phase.""" + chans = [channel] if channel else (["phase"] + sorted(os.listdir(INV))) + for ch in chans: + gk = f"{INV}/{ch}/geneKO" + if not os.path.isdir(gk) or not os.path.isdir(f"{INV}/{ch}/_anchors/NTC"): + continue + for g in sorted(os.listdir(gk)): + if glob.glob(f"{gk}/{g}/cell0/frame_{A0:02d}.webp"): + return ch, g + return None, None + + +def run(channel=None, gene=None): + import torch # noqa + from ops_model.models.attention.diffex.classifier.celldino_features import embed_crops + from ops_model.models.attention.diffex.directions.config import DirConfig + os.makedirs(OUT, exist_ok=True) + ch, g = _first_done(channel) + gene = gene or g + if ch is None: + print("[ntc-gap] no finished traversal yet — rerun when one lands"); return + print(f"[ntc-gap] channel={ch} gene={gene} for α=0 (NTC-anchor reconstruction, gene-independent)", flush=True) + + anc = _celln(f"{INV}/{ch}/_anchors/NTC") + trav = f"{INV}/{ch}/geneKO/{gene}" + n = min(len(anc), len(_celln(trav))) + real_paths = [f"{INV}/{ch}/_anchors/NTC/cell{i}/real.webp" for i in range(n)] + gen_paths = [f"{trav}/cell{i}/frame_{A0:02d}.webp" for i in range(n)] + keep = [i for i in range(n) if os.path.exists(real_paths[i]) and os.path.exists(gen_paths[i])] + real_paths = [real_paths[i] for i in keep]; gen_paths = [gen_paths[i] for i in keep] + print(f"[ntc-gap] {len(keep)} matched anchor/α=0 pairs", flush=True) + + cfg = DirConfig(grain="geneKO", target=gene, device="cuda") + R = np.asarray(_emb_webp(real_paths, cfg, embed_crops), np.float64) # real NTC (N,1024) + G = np.asarray(_emb_webp(gen_paths, cfg, embed_crops), np.float64) # inverse α=0 NTC (N,1024) + np.savez(f"{OUT}/emb_{ch}.npz", R=R, G=G, gene=gene, channel=ch) + analyze(ch) + + +def analyze(ch="phase"): + from numpy.linalg import norm + d = np.load(f"{OUT}/emb_{ch}.npz", allow_pickle=True) + R, G = d["R"], d["G"] + cos = lambda u, v: float(u @ v / (norm(u) * norm(v))) + # (1) per-cell paired gap + pc = np.array([cos(R[i], G[i]) for i in range(len(R))]) + # (2) offset consistency + off = G - R # per-cell offset vectors + mo = off.mean(0); mo_hat = mo / norm(mo) + cos_to_mean = np.array([cos(off[i], mo) for i in range(len(off))]) + mag = norm(off, axis=1) + # (3) centroid gap + print("\n=== Real NTC vs inverse-α=0 NTC (CellDINO, paired same-cell) ===") + print(f" pairs : {len(R)}") + print(f" (1) per-cell cosine(real,gen) : mean {pc.mean():.3f} min {pc.min():.3f} max {pc.max():.3f}") + print(f" (3) centroid cosine : {cos(R.mean(0), G.mean(0)):.3f}") + print(f" centroid ||offset|| : {norm(G.mean(0) - R.mean(0)):.2f}") + print(f" (2) offset ||·|| : mean {mag.mean():.2f} std {mag.std():.2f} (rigid ⇒ low std)") + print(f" offset dir consistency : cos(offset_i, mean offset) mean {cos_to_mean.mean():.3f} std {cos_to_mean.std():.3f}") + print(f" (cos→1 & low ||·|| std ⇒ CONSISTENT rigid offset, correctable by one translation)") + if ch == "phase": + _project(R, G) + else: + print(f"\n[note] channel={ch} (phase not finished yet); UMAP placement skipped (phase-only embedding).") + + +def _project(R, G): + """Place real-NTC and inverse-α=0 into the phase gene embedding (real-pop std) — distance to real NTC cluster.""" + import sys + sys.path.insert(0, os.path.dirname(os.path.abspath(__file__))) + import gen_phate_passthrough as gp + from numpy.linalg import norm + a, comp, mean = gp._load_embedding() + Xpca = np.asarray(a.obsm["X_pca"], np.float64) + U = np.asarray(a.obsm["X_umap"], np.float64) + m = gp._ntc_mask(a); ntc2d = U[m].mean(0) + mu, sd = gp._real_baseline(gp._cache_files()) # real-population standardization + def place(X): + pc = (((X - mu) / sd) - mean) @ comp.T + return gp._landmark(pc.mean(0), Xpca, U) + pr, pg = place(R), place(G) + print("\n=== Phase UMAP placement (real-pop std) ===") + print(f" real NTC cluster : {ntc2d.round(2)}") + print(f" real-anchor centroid lands : {pr.round(2)} dist to NTC {norm(pr - ntc2d):.2f} (sanity: should be small)") + print(f" inverse-α=0 centroid lands : {pg.round(2)} dist to NTC {norm(pg - ntc2d):.2f} (old random-xT was 8.43)") + + +CTRL = f"{INV}/phase/_anchors/NTC/ctrl.npz".replace("viewer_assets_v5_inv", "viewer_assets_v5") + + +def webp_ab(): + """8-bit-webp vs proper-float A/B: embed the SAME 45 real NTC crops (float, from ctrl.npz) two ways — + (A) float straight into CellDINO, (B) round-tripped through the traversal's 8-bit _save_webp path — + and compare CellDINO cosine. Answers: does saving as a proper (float/zarr) image bring generated closer to real?""" + import torch, tempfile # noqa + from ops_model.models.attention.diffex.classifier.celldino_features import embed_crops + from ops_model.models.attention.diffex.directions.config import DirConfig + from ops_model.models.attention.diffex.viewer.precompute import _save_webp + os.makedirs(OUT, exist_ok=True) + d = np.load(CTRL, allow_pickle=True) + imgs = d["anchor_imgs"].astype(np.float32) # (45,1,160,160) float [-1,1] + ctrl = d["ctrl_embs"].astype(np.float64) # pipeline float-path embeddings + cfg = DirConfig(grain="geneKO", target="NTC", device="cuda") + A = np.asarray(embed_crops(imgs, cfg, cache_path=None), np.float64) # (A) float path + tmp = tempfile.mkdtemp(); B = [] + for i in range(len(imgs)): + p = f"{tmp}/c{i}.webp"; _save_webp(p, imgs[i, 0], 256) # exact traversal 8-bit path + B.append(np.asarray(Image.open(p).convert("L"), np.float32) / 255.0 * 2 - 1) + B = np.asarray(embed_crops(np.stack(B)[:, None].astype(np.float32), cfg, cache_path=None), np.float64) + np.savez(f"{OUT}/webp_ab.npz", A=A, B=B, ctrl=ctrl) + _report_ab() + + +def _report_ab(): + from numpy.linalg import norm + d = np.load(f"{OUT}/webp_ab.npz"); A, B, ctrl = d["A"], d["B"], d["ctrl"] + cs = lambda U, V: np.array([float(U[i] @ V[i] / (norm(U[i]) * norm(V[i]))) for i in range(len(U))]) + fw = cs(A, B) + print("\n=== 8-bit webp vs proper float (same 45 real NTC cells, CellDINO) ===") + print(f" float-path vs 8bit-webp-path : cosine mean {fw.mean():.4f} min {fw.min():.4f} max {fw.max():.4f}") + print(f" (sanity) float-path vs pipeline ctrl_embs : mean {cs(A, ctrl).mean():.4f}") + print(" → cosine≈1 ⇒ webp is NOT the gap (proper/zarr image won't help; residual is generative/OOD)") + print(" cosine notably <1 ⇒ 8-bit webp shifts the embedding; saving as float/zarr would help") + + +def submit(func=run, name="ntc_inverse_gap"): + from ops_utils.hpc.slurm_batch_utils import submit_parallel_jobs + submit_parallel_jobs([{"name": name, "func": func, "kwargs": {}}], + experiment="ntc_inverse_gap", + slurm_params={"slurm_partition": "gpu", "slurm_gres": "gpu:1", "cpus_per_task": 8, + "mem_gb": 48, "timeout_min": 60}, log_dir="ntc_inverse_gap", + wait_for_completion=False) + + +if __name__ == "__main__": + import sys + if "--analyze" in sys.argv: + analyze() + elif "--report-ab" in sys.argv: + _report_ab() + elif "--webp" in sys.argv: + submit(func=webp_ab, name="ntc_webp_ab") + elif "--local" in sys.argv: + run() + else: + submit() diff --git a/src/ops_model/models/attention/diffex/figures/gen_validation/patch_cache_real.py b/src/ops_model/models/attention/diffex/figures/gen_validation/patch_cache_real.py new file mode 100644 index 0000000..ad1a0ed --- /dev/null +++ b/src/ops_model/models/attention/diffex/figures/gen_validation/patch_cache_real.py @@ -0,0 +1,56 @@ +"""Populate the `real` field for the 42 recovered geneKO in the new-v5 cache (they were built gen-only). +Gather each gene's top-30 accuracy real cells (dashed ranking name for KRTAP) → embed_crops → write into +gen_real_map_cache_v5new/geneKO/.npz so the pooled-centroid REAL ceiling covers the full 1000.""" +import os, glob +import numpy as np + +CV = "/hpc/projects/icd.fast.ops/analysis/figure4_traversals" +CACHE = f"{CV}/gen_real_map_cache_v5new/geneKO" +RANKP = "/hpc/projects/icd.fast.ops/models/diffex/viewer_assets_v5/_rankings/pma_v5_phase_geneKO.parquet" +KREAL = 30 + + +def _drop42(): + import glob as g + single = {os.path.basename(x) for x in g.glob("/hpc/projects/icd.fast.ops/models/diffex/viewer_assets_v5/phase/geneKO/*") + if os.path.isdir(x) and "__to__" not in x and not os.path.basename(x).startswith("_")} + old = {os.path.basename(f)[:-4] for f in g.glob(f"{CV}/gen_real_map_cache/geneKO/*.npz")} + return sorted(single - old) + + +def run(): + import pandas as pd + from ops_model.models.attention.diffex.viewer.precompute import _gather_class + from ops_model.models.attention.diffex.directions.config import DirConfig + from ops_model.models.attention.diffex.classifier.config import slugify + genes = _drop42() + orig = {slugify(str(x)): str(x) for x in pd.read_parquet(RANKP, columns=["gene"])["gene"].unique()} # slug→ranking name + cfg = DirConfig(grain="geneKO", target=genes[0], device="cuda"); cfg.num_workers = 12 + done = 0 + for g in genes: + cp = f"{CACHE}/{g}.npz" + if not os.path.exists(cp): + continue + d = np.load(cp, allow_pickle=True) + if len(np.asarray(d["real"])): # already has real + continue + name = orig.get(slugify(g), g) # dashed name for KRTAP, plain otherwise + _, embs = _gather_class(cfg, name, KREAL, parquet=RANKP) + if not len(embs): + print(f"no real for {g} ({name})"); continue + np.savez(cp, real=np.asarray(embs, np.float32)[:KREAL], gen=d["gen"], alphas=d["alphas"], gene=str(d["gene"])) + done += 1 + return {"patched": done, "n_genes": len(genes)} + + +def main(): + from ops_utils.hpc.slurm_batch_utils import submit_parallel_jobs + submit_parallel_jobs(jobs_to_submit=[{"name": "patchreal", "func": run, "kwargs": {}}], + experiment="patchreal", + slurm_params={"slurm_partition": "preempted", "slurm_gres": "gpu:1", "cpus_per_task": 12, + "mem_gb": 64, "timeout_min": 120, "slurm_constraint": "[a40|a6000|l40s]"}, + log_dir="patchreal", wait_for_completion=False) + + +if __name__ == "__main__": + main() diff --git a/src/ops_model/models/attention/diffex/figures/gen_validation/publish_multibag_page.py b/src/ops_model/models/attention/diffex/figures/gen_validation/publish_multibag_page.py new file mode 100644 index 0000000..5120f21 --- /dev/null +++ b/src/ops_model/models/attention/diffex/figures/gen_validation/publish_multibag_page.py @@ -0,0 +1,177 @@ +"""Create the sister Confluence page mirroring the v4/v5 DiffAE validation page, with the new multibag bag-sweep +figures. REST flow (create -> upload attachments -> PUT storage body with ). Idempotent-ish: pass an +existing PAGE_ID env to update instead of create.""" +import os, json, base64, urllib.request + +USER = "gav.sturm@czbiohub.org"; TOKEN = os.environ["CONFLUENCE_API_TOKEN"] +BASE = "https://czbiohub.atlassian.net/wiki" +SPACE_ID = "3319857206"; PARENT = "5538218009" +TITLE = "DiffAE multibag-traversal validation — bag-size sweep (Figure 4)" +PLOTS = "/hpc/projects/icd.fast.ops/analysis/figure4_traversals/bag_sweep_plots_v5new" +IMGS = ["settransformer_bagsweep_all.png", "settransformer_bagsweep_realdist.png", + "centroid_bagsweep.png", "centroid_bagsweep_global.png", "centroid_pooled_bagsweep.png", "centroid_pooled_bagsweep_global.png", + "control_halves_zscore.png", "anchor_halves.png", "st_anchor_halves.png", "retrieval_map_proper1k.png", + "distinct_sweep_median.png", "distinct_sweep_mean.png", "distinct_violin.png"] + + +def _auth(): + return "Basic " + base64.b64encode(f"{USER}:{TOKEN}".encode()).decode() + + +def _req(url, data=None, method="GET", ctype="application/json", raw=False): + h = {"Authorization": _auth()} + if ctype: + h["Content-Type"] = ctype + r = urllib.request.Request(url, data=data, method=method, headers=h) + with urllib.request.urlopen(r) as resp: + return resp.read() if raw else json.load(resp) + + +def create(): + body = {"type": "page", "title": TITLE, "space": {"key": "dashboard"}, + "ancestors": [{"id": PARENT}], "body": {"storage": {"value": "

scaffold

", "representation": "storage"}}} + d = _req(f"{BASE}/rest/api/content", json.dumps(body).encode(), "POST") + return d["id"] + + +def upload(pid, path): + """Upload new, or update-by-id if the filename already exists (idempotent re-runs). Returns the attachment id + (or None if missing) so the body can build a working REST download link.""" + import subprocess + if not os.path.exists(path): + print("skip (missing):", os.path.basename(path)); return None + fn = os.path.basename(path) + ex = _req(f"{BASE}/rest/api/content/{pid}/child/attachment?filename={fn}") + url = f"{BASE}/rest/api/content/{pid}/child/attachment" + if ex.get("results"): + url += f"/{ex['results'][0]['id']}/data" + subprocess.run(["curl", "-s", "-u", f"{USER}:{TOKEN}", "-X", "POST", "-H", "X-Atlassian-Token: nocheck", + "-F", f"file=@{path}", url], check=True, capture_output=True) + q = _req(f"{BASE}/rest/api/content/{pid}/child/attachment?filename={fn}") + return q["results"][0]["id"] if q.get("results") else None + + +PID = os.environ.get("PAGE_ID", "") +SVG_ATT = {} # svg filename -> attachment id (populated during upload) + + +def img(fn, w=1000): + svg = fn.rsplit(".", 1)[0] + ".svg" + att = SVG_ATT.get(svg) + link = (f'

' + f'⬇ download SVG (vector)

' if att else "") + return (f'' + f'{link}') + + +def img_cell(fn, w=620): + """Inline (non-wide) image for a table cell — two of these sit side by side in a 2-col table.""" + svg = fn.rsplit(".", 1)[0] + ".svg" + att = SVG_ATT.get(svg) + link = (f'

' + f'⬇ SVG

' if att else "") + return f'{link}' + + +def build_body(): + return f""" + +

Headline. The new multibag-ranked DiffAE phase traversals (viewer_assets_v5, 400 cells × 17 α, w=1.5, 100 DDIM steps) recover the strong v4-era validation that the earlier single-anchor "valid200" run had lost. Across all three independent checks the recovery signature is back: SetTransformer real-distinguishable recovery returns to ~46% top-1 (geneKO) / ~86% (complex) and P(target) near the real-cell ceiling; pooled centroid recovery (now per-bag α=0 standardized — the domain-honest metric established by the control below) reaches bag-200 top-1 ~48% (geneKO) / ~69% (complex) against a per-bag real-cell ceiling of 66% / 89%; and the classifier peak returns to α≈0.5–1.5.

+

This page is the bag-size-sweep sister of the v4/v5 DiffAE generative-cell validation page: every quantitative measure is re-scored at bag sizes {{20, 50, 100, 200, 400}} on the full 1,000 geneKO + 95 complex panel.

+
+ + +

MAJOR FINDING — the anchor-half "gap" is a centroid-metric standardization artifact, not phenotype loss. The 400-cell traversals concatenate two 200-cell anchor sets with identical directions and traversals: cells 0–199 = old hand-picked (curated favourite) NTC anchors; cells 200–399 = strict multibag top-200 NTC anchors. The apparent recovery gap appears only in the position-sensitive centroid metric, and only under its global standardization — it vanishes once standardization is matched to the classifier's:

+ + + + + + + + +
hand-picked (first 200)strict multibag NTC (second 200)gap
centroid top-1, global α=0 std — geneKO76%13%63 pts
centroid top-1, global α=0 std — complex86%40%46 pts
centroid top-1, per-bag α=0 std — geneKO48%42%7 pts
centroid top-1, per-bag α=0 std — complex69%64%5 pts
SetTransformer top-1 — geneKO10%9%~0
SetTransformer top-1 — complex52%54%~0
+

Mechanism. The centroid metric standardizes generated cells against a global / panel α=0 mean (shared {{grain}}_mu.npz), so a half-specific CellDINO offset survives and inflates the gap — it even pushes the second half's optimal α out to α=3 (extreme α needed to drag the offset half onto its target). score_embs_v5 (SetTransformer) instead self-standardizes each bag on its own α=0 generated frames, cancelling that offset — which is exactly why it sees no gap. Applying the same per-bag α=0 standardization to the centroid metric collapses the gap from 63→7 pts (geneKO) and 46→5 pts (complex), and restores the normal α≈0.5–1 peak.

+

Conclusion. The two anchor sets produce the same perturbation phenotype (the supervised classifier — the gold standard — cannot tell them apart). The 76/13 is a standardization choice in the centroid metric, so the bag=400 "dilution" is not real phenotype loss. The other bag-sweep plots still use bags 20–200 for a consistent global-standardized reference.

+
+{img("control_halves_zscore.png", 1000)} +

The control. Peak-α centroid top-1 under global (panel α=0) vs per-bag (each half's own α=0) standardization. The large first/second gap under global collapses under per-bag — definitively a metric standardization artifact.

+{img("anchor_halves.png", 1150)} +

Centroid recovery under the current global standardization — first-200 (hand-picked, blue) vs second-200 (strict multibag top-NTC, red). Same directions/traversals; only the anchor cells differ. The gap here is the artifact the control above dissolves.

+{img("st_anchor_halves.png", 1150)} +

SetTransformer on the same two halves (P(target) · median rank · mean rank · top-5/top-1). Near-identical — the self-standardizing classifier sees the same phenotype from both anchor sets.

+ +

Source on Bruno

+ + + + + + +
Cache buildops_model/.../diffex/figures/gen_validation/valid200_cache_build.py (V200_ASSETS=viewer_assets_v5, NCELL=400)
Scoringbag_sweep_score.py (SetTransformer) · centroid_bagsweep.py (centroid recovery) · gen_real_distinct.py (distinctiveness)
Plotsbag_sweep_plots.py
Cache/hpc/projects/icd.fast.ops/analysis/figure4_traversals/gen_real_map_cache_v5new/
Outputs/hpc/projects/icd.fast.ops/analysis/figure4_traversals/{{bag_sweep_v5new, centroid_bagsweep_v5new, gen_real_distinct_v5new_K*, bag_sweep_plots_v5new}}/
+ +

How to read these plots

+

Each measure is scored on bags of generated cells and plotted vs the traversal step α (x-axis, integer ticks −5…5; α=0 = real control, α≈1 = the class mean, |α|>1 = exaggerated). One line per bag size (viridis: dark=20 → yellow=400). geneKO / complex are the two rows. "Higher is better" on every y-axis except target rank (1 = the classifier's top pick, lower is better).

+

Bags are the first-B cells (nested), restricted to 20–200 = the hand-picked anchor half for a consistent reference (bag=400 mixes in the second-200 strict-multibag anchors; per the finding above their apparent "dilution" is a centroid-standardization artifact, not weaker phenotypes). Within that half, the pooled SetTransformer set-classifier and pooled centroid recovery both improve with bag size (more cells = more evidence), saturating by ~100–200.

+ +

1 · SetTransformer set-classifier (rank / P(target) / top-k)

+

An independently-trained v5 SetTransformer (Alex Lin; real cells only, never exposed to DiffAE traversals) scores each generated bag to its target class. Four columns: P(target), median target rank, mean target rank (outlier-sensitive), and % top-5 / top-1 recovered. Dashed real-cell references where applicable.

+{img("settransformer_bagsweep_all.png", 1200)} +

All classes (1,000 geneKO / 95 complex).

+{img("settransformer_bagsweep_realdist.png", 1200)} +

Real-distinguishable subset only (genes whose REAL cells score top-1 accuracy > 0.5 @ bag-20). Here generated traversals track close to the real-cell ceiling — geneKO ~46% top-1, complex ~86% top-1 at peak α≈1–1.5; complex P(target) rises to ~0.85 at bag 400.

+ +

2 · Centroid recovery (generated → nearest real centroid)

+

Classifier-independent: does a generated cell land nearest its true class's faithful real centroid (Cell-DINO), rather than a neighbour? mAP (1/rank of the true centroid), % top-1, % top-5. Now scored with per-bag α=0 standardization (each gene's own α=0 frames — matching score_embs_v5, the domain-honest metric established by the control above; the old global-mu numbers are preserved on disk).

+ + + +
Per-bag α=0 standardized (domain-honest — correct)Global-mu standardized (previous — inflated)
{img_cell("centroid_bagsweep.png")}{img_cell("centroid_bagsweep_global.png")}
+

Per-cell centroid recovery, per-bag (left) vs global-mu (right). Left (correct): geneKO mAP ~0.24 (top-1 ~15%) at α≈+4, complex ~0.53 (top-1 ~39%) at α≈+3; negative α recovers nothing. Right (old): geneKO mAP ~0.335 (the "v4 baseline"), complex ~0.65 — higher, but standardization-inflated per the control above, not truer recovery.

+ +

2b · Pooled bag-level centroid recovery (with per-bag real ceiling)

+

The bag-level view you actually want for a bag-size sweep: pool B cells → the generated class centroid → does that land nearest the true real centroid? More cells = a better centroid estimate, so recovery rises with bag; the dotted line is the matched per-bag real-cell ceiling (bootstrap B real cells → mean). Per-bag α=0 standardized.

+ + + +
Per-bag α=0 standardized (domain-honest — correct)Global-mu standardized (previous — inflated)
{img_cell("centroid_pooled_bagsweep.png")}{img_cell("centroid_pooled_bagsweep_global.png")}
+

Pooled bag-level recovery, per-bag (left) vs global-mu (right). Solid = generated, dotted = per-bag real-cell ceiling; all 1,000 geneKO / 95 complex. Left (correct): bag-200 top-1 ~48% (geneKO) / ~69% (complex), peak α≈+0.5, below the ~66% / ~89% real ceilings. Right (old): geneKO ~76% (slightly exceeds its ceiling), complex ~86% — the more impressive-looking result, but standardization-inflated per the control above.

+ +

3 · Distinctiveness (how separable generated cells are from other classes)

+

Within-domain copairs mAP: do a class's generated cells cluster together and apart from other classes? Dotted = real-cell reference. Both median (robust) and mean (outlier-sensitive) shown. K (cells/class) coverage differs by grain: complex (95 classes) extends to K = {{20, 50, 100, 200}}; geneKO (1,000 classes) caps at K = 50 — at K ≥ 100 the copairs pass OOMs (> 240 GB; 100–200k cells), and the cached real reference is only 30 cells/class, so higher-K geneKO is neither computable nor referenceable without a himem/GPU rewrite.

+{img("distinct_sweep_median.png", 1000)} +

Median distinctiveness mAP. (geneKO: K=20/50 only; complex: K=20/50/100/200.)

+{img("distinct_sweep_mean.png", 1000)} +

Mean distinctiveness mAP.

+{img("distinct_violin.png", 950)} +

Per-class distinctiveness distribution (top-50 cells/class) at each grain's peak α — real (blue) vs generated (green), median bar. Generated exceeds real for gene-KOs and matches/exceeds for complexes, reflecting the low-diversity effect (centroid-directed generated cells cluster tighter than real cells).

+ +

Not yet ported from the v4/v5 page

+

For completeness — two measures on the sister v4/v5 page are not reproduced here yet:

+
    +
  • Cross-domain retrieval mAP (mAP_proper1k) — for each generated class-X cell, the average precision of retrieving real class-X cells ahead of all others (gen→real 1000-cell gallery). A genuine third mAP-family measure. Complex is done; geneKO timed out (the full 1,000-cell/class GPU gallery pass hit the 2 h wall). It also uses the same global-mu path as §2 and so needs the per-bag redo; rebuild pending (capped gallery + per-bag standardization).
  • +
  • Reach-fraction vs bag (v5_reachfrac) — % of real-distinguishable classes whose generated accuracy reaches the real-cell level, per bag. Needs a fresh bag_scaling run on the new traversals (the old bagtest data has been cleared); real-cell per-bag accuracy is no longer on disk to reuse.
  • +
+

The embedding-projection sections (UMAP/PHATE placement, complex montages, §4–6 of the v4/v5 page) are separate analyses, not part of this bag-size sweep.

+""" + + +def main(): + pid = os.environ.get("PAGE_ID") or create() + print("page:", pid) + for fn in IMGS: + upload(pid, f"{PLOTS}/{fn}"); print("uploaded", fn) + svg = f"{fn.rsplit('.', 1)[0]}.svg" + aid = upload(pid, f"{PLOTS}/{svg}") # vector for download link + if aid: + SVG_ATT[svg] = aid + cur = _req(f"{BASE}/rest/api/content/{pid}?expand=version")["version"]["number"] + payload = {"version": {"number": cur + 1, "message": "multibag bag-sweep figures"}, + "title": TITLE, "type": "page", "body": {"storage": {"value": build_body(), "representation": "storage"}}} + v = _req(f"{BASE}/rest/api/content/{pid}", json.dumps(payload).encode(), "PUT")["version"]["number"] + print(f"PUT ok, version {v}") + print(f"URL: {BASE}/spaces/dashboard/pages/{pid}") + + +if __name__ == "__main__": + main() diff --git a/src/ops_model/models/attention/diffex/figures/gen_validation/rank_summary.py b/src/ops_model/models/attention/diffex/figures/gen_validation/rank_summary.py new file mode 100644 index 0000000..0767da1 --- /dev/null +++ b/src/ops_model/models/attention/diffex/figures/gen_validation/rank_summary.py @@ -0,0 +1,147 @@ +"""Population summary of the v5 SetTransformer target-rank across the DiffAE traversal. + +For every geneKO/complex traversal, rank_target[α] = the 1-indexed position of the true class in the +classifier's ranking of a bag-20 set of the α-frames. Aggregate mean & median rank across classes at +each α, for both anchor pools (accpool = 25 hand-picked accuracy NTCs = deployed default; attn = top-attention). +Rank 1 = perfect recovery; lower is better. + +Two figures: + rank_summary.png — all classes + rank_summary_realfilt.png — only classes whose REAL cells are distinguishable (real top1_acc > 0.5 @ bag20), + i.e. the feasible ceiling; measures recovery where a phenotype actually exists. +""" +import json, glob, os +import numpy as np +import matplotlib.pyplot as plt + +plt.rcParams["pdf.fonttype"] = 42 +plt.rcParams.update({ # figure-ready: ~2x default text + "font.size": 19, "axes.titlesize": 23, "axes.labelsize": 21, + "xtick.labelsize": 18, "ytick.labelsize": 18, "legend.fontsize": 14, + "figure.titlesize": 25, "axes.linewidth": 1.4, + "xtick.major.size": 7, "ytick.major.size": 7, "xtick.major.width": 1.4, "ytick.major.width": 1.4, +}) +XTICKS = list(range(-5, 6)) # integer α ticks +BASE = "/hpc/projects/icd.fast.ops/models/diffex" +OUTDIR = "/hpc/projects/icd.fast.ops/analysis/figure4_traversals/rank_summary" +os.makedirs(OUTDIR, exist_ok=True) +ASSETS = os.environ.get("RANK_ASSETS", "viewer_assets_v5_accpool") # default accpool; new multibag traversals = viewer_assets_v5 +OVERLAY = os.environ.get("RANK_OVERLAY", "") or None # optional dashed overlay tree (empty = none) +SUF = os.environ.get("RANK_SUF", "") # filename suffix (e.g. _v5new) +GRAINS = {"geneKO": ("geneKO", 1000), "complex": ("complex", 99)} +RED = "#c0392b" +REAL = json.load(open(f"{BASE}/viewer_assets_v5/real_acc20.json")) # real top1_acc@bag20 by asset_dir + + +def collect(assets, sub, keep=None): + """→ (alphas, matrix[n_class × n_alpha] of rank_target). keep: optional set of names to include.""" + rows, alphas = [], None + for d in sorted(glob.glob(f"{BASE}/{assets}/phase/{sub}/*")): + name = os.path.basename(d) + if not os.path.isdir(d) or "__to__" in name: + continue + if keep is not None and name not in keep: + continue + sf = f"{d}/scores_v5.json" + if not os.path.exists(sf): + continue + sc = json.load(open(sf)) + rk = sc.get("rank_target") + if not rk: + continue + alphas = sc["alphas"] + rows.append([np.nan if v is None else float(v) for v in rk]) + return np.array(alphas), np.array(rows) if rows else np.empty((0, 0)) + + +def real_keep(sub): + """dir-basenames whose real cells clear top1_acc>0.5 @bag20. Keyed by asset_dir 'phase/{sub}/{name}'.""" + pre = f"phase/{sub}/" + return {k[len(pre):] for k, v in REAL.items() if k.startswith(pre) and v > 0.5} + + +def make_fig(title, fname, errstyle): + """errstyle: 'mean_sem' (mean line + SEM band) or 'median_iqr' (median line + IQR band). Accpool only. + Overlays all classes (gray) vs real-distinguishable (real top1>0.5 @bag20, red) on the same axes.""" + clip = lambda a: np.maximum(a, 0.9) # keep bands positive on the log axis + fig, axes = plt.subplots(1, 2, figsize=(17, 6.2), constrained_layout=True) + for ax, (gname, (sub, nclass)) in zip(axes, GRAINS.items()): + ns = {} + for keep, ftag, color in [(None, "all classes", "#7f8c8d"), (real_keep(sub), "real top1 > 0.5", RED)]: + al, M = collect(ASSETS, sub, keep) + if M.size == 0: + continue + n = int(np.isfinite(M[:, len(al) // 2]).sum()); ns[ftag] = n + if errstyle == "mean_sem": + c = np.nanmean(M, axis=0); e = np.nanstd(M, axis=0) / np.sqrt(np.isfinite(M).sum(axis=0)) + lo, hi = clip(c - e), c + e + lab = f"{ftag} — mean ± SEM" + else: # median_iqr + c = np.nanmedian(M, axis=0) + lo, hi = clip(np.nanpercentile(M, 25, axis=0)), np.nanpercentile(M, 75, axis=0) + lab = f"{ftag} — median (IQR)" + ax.plot(al, c, "-", color=color, lw=3.4, label=lab) + ax.fill_between(al, lo, hi, color=color, alpha=0.22, lw=0) + alo, Mo = (collect(OVERLAY, sub, keep) if OVERLAY else (None, np.empty((0,0)))) # dashed = 200-cell bag (geneKO only) + if Mo.size: + co = np.nanmean(Mo, 0) if errstyle == "mean_sem" else np.nanmedian(Mo, 0) + ax.plot(alo, co, "--", color=color, lw=2.8, label=f"{ftag} — 200-cell bag") + ax.axvline(0, color="#bbb", lw=1, zorder=0) + ax.axhline(1, color="#27ae60", lw=1, ls=":", zorder=0) + ax.set_yscale("log") + ax.set_xlabel("traversal α") + ax.set_ylabel("target rank") + ax.set_title(f"{gname} · {ns.get('all classes', nclass)} classes ({ns.get('real top1 > 0.5', '?')} real-distinguishable)") + ax.set_xticks(XTICKS); ax.set_xlim(-5, 5) + ax.grid(alpha=0.25, which="both") + h, l = axes[0].get_legend_handles_labels() + fig.legend(h, l, loc="center left", bbox_to_anchor=(1.0, 0.5), fontsize=15, framealpha=0.9) + fig.suptitle(title, fontweight="bold") + out = f"{OUTDIR}/{fname}.png" + fig.savefig(out, dpi=150, bbox_inches="tight") + fig.savefig(out.replace(".png", ".svg"), bbox_inches="tight") + plt.close(fig) + print("saved", out) + + +def make_top5_fig(fname): + """% of classes whose target lands in the top-5 ranking at each α — all classes vs real-distinguishable + (real top1-acc > 0.5 @ bag20). Accpool only. top-5 hit = rank_target <= 5.""" + fig, axes = plt.subplots(1, 2, figsize=(17, 6.2), constrained_layout=True) + for ax, (gname, (sub, nclass)) in zip(axes, GRAINS.items()): + ns = {} + for keep, ftag, color in [(None, "all classes", "#7f8c8d"), (real_keep(sub), "real top1 > 0.5", RED)]: + al, M = collect(ASSETS, sub, keep) + if M.size == 0: + continue + valid = np.isfinite(M) + f5 = 100 * (valid & (M <= 5)).sum(axis=0) / valid.sum(axis=0) # target in top-5 + f1 = 100 * (valid & (M == 1)).sum(axis=0) / valid.sum(axis=0) # target is #1 (top-1) + ns[ftag] = int(valid[:, len(al) // 2].sum()) + ax.plot(al, f5, "-", color=color, lw=3.4, label=f"{ftag} — top-5") + ax.plot(al, f1, ":", color=color, lw=3.4, label=f"{ftag} — top-1") + b5, b1 = int(np.argmax(f5)), int(np.argmax(f1)) + print(f" {gname:8s} | {ftag:16s} | peak top-5 = {f5[b5]:.0f}% @α={al[b5]:+.1f} | peak top-1 = {f1[b1]:.0f}% @α={al[b1]:+.1f}") + alo, Mo = (collect(OVERLAY, sub, keep) if OVERLAY else (None, np.empty((0,0)))) # dashed = 200-cell bag top-5 (geneKO only) + if Mo.size: + vo = np.isfinite(Mo); f5o = 100 * (vo & (Mo <= 5)).sum(0) / vo.sum(0) + ax.plot(alo, f5o, "--", color=color, lw=2.8, label=f"{ftag} — 200-cell bag top-5") + ax.axvline(0, color="#bbb", lw=1, zorder=0) + ax.set_ylim(0, 100) + ax.set_xlabel("traversal α") + ax.set_ylabel("% recovered") + ax.set_title(f"{gname} · {ns.get('all classes', nclass)} classes ({ns.get('real top1 > 0.5', '?')} real-distinguishable)") + ax.set_xticks(XTICKS); ax.set_xlim(-5, 5) + ax.grid(alpha=0.25) + h, l = axes[0].get_legend_handles_labels() + fig.legend(h, l, loc="center left", bbox_to_anchor=(1.0, 0.5), fontsize=15, framealpha=0.9) + fig.suptitle("Top-5 (solid) / top-1 (dotted) recovery vs α (bag=20)", fontweight="bold") + out = f"{OUTDIR}/{fname}.png" + fig.savefig(out, dpi=150, bbox_inches="tight"); fig.savefig(out.replace(".png", ".svg"), bbox_inches="tight") + plt.close(fig); print("saved", out) + + +TITLE = "Target rank along the traversal (bag=20, accuracy anchors)" +make_fig(TITLE, "rank_summary_mean_sem"+SUF, "mean_sem") +make_fig(TITLE, "rank_summary_median_iqr"+SUF, "median_iqr") +make_top5_fig("top5_recovery_summary"+SUF) diff --git a/src/ops_model/models/attention/diffex/figures/gen_validation/st_halves_score.py b/src/ops_model/models/attention/diffex/figures/gen_validation/st_halves_score.py new file mode 100644 index 0000000..98dfa28 --- /dev/null +++ b/src/ops_model/models/attention/diffex/figures/gen_validation/st_halves_score.py @@ -0,0 +1,54 @@ +"""SetTransformer scored on the two anchor halves: cells [0:200] (hand-picked) vs [200:400] (strict multibag NTC). +Per gene → {first, second} scores_v5-dicts (P(target), rank, top1, top5 per α). Same score_embs_v5 as the bag +sweep, just fed each 200-cell half. GPU, sharded → per-gene json.""" +import os, glob, json +import numpy as np + +CV = "/hpc/projects/icd.fast.ops/analysis/figure4_traversals" +GRAIN = os.environ.get("BSS_GRAIN", "geneKO") +CACHE = f"{CV}/gen_real_map_cache_v5new/{GRAIN}" +OUT = f"{CV}/st_halves_v5new/{GRAIN}" + + +def score_shard(genes): + import torch + from ops_model.models.attention.diffex.viewer.score_generated import score_embs_v5 + from ops_model.models.attention.diffex.viewer.set_classifier import load_set_classifier, V5_CKPT_ROOT, V5_RUNS + from ops_model.models.attention.diffex.classifier.config import slugify + os.makedirs(OUT, exist_ok=True) + dev = "cuda" if torch.cuda.is_available() else "cpu" + run = V5_RUNS[("phase", "geneKO" if GRAIN == "geneKO" else "complex_ebionly")] + m, g2i, c2i = load_set_classifier(run=run, device=dev, root=V5_CKPT_ROOT) + ci = c2i.get("Phase2D", 0); slug2orig = {slugify(k): k for k in g2i} + done = 0 + for g in genes: + f = f"{CACHE}/{g}.npz"; outp = f"{OUT}/{g}.json" + if not os.path.exists(f) or os.path.exists(outp): + continue + d = np.load(f, allow_pickle=True); gene = str(d["gene"]); al = [float(a) for a in d["alphas"]] + tgt = gene if gene in g2i else slug2orig.get(slugify(gene), gene) + embs = [None if d["gen"][ai] is None else np.asarray(d["gen"][ai], np.float32) for ai in range(len(al))] + if any(e is not None and len(e) < 400 for e in embs): + continue + first = [None if e is None else e[:200] for e in embs] + second = [None if e is None else e[200:400] for e in embs] + rec = {"first": score_embs_v5(first, al, tgt, m, g2i, ci, run, device=dev, bag=200), + "second": score_embs_v5(second, al, tgt, m, g2i, ci, run, device=dev, bag=200)} + json.dump({"gene": gene, "alphas": al, **rec}, open(outp, "w")); done += 1 + return {"grain": GRAIN, "done": done} + + +def main(): + from ops_utils.hpc.slurm_batch_utils import submit_parallel_jobs + genes = sorted(os.path.basename(f)[:-4] for f in glob.glob(f"{CACHE}/*.npz")) + ch = 60; shards = [genes[i:i + ch] for i in range(0, len(genes), ch)] + jobs = [{"name": f"sth_{GRAIN}_{i}", "func": score_shard, "kwargs": {"genes": s}} for i, s in enumerate(shards)] + print(f"[st-halves] {GRAIN}: {len(genes)} genes → {len(jobs)} shards") + submit_parallel_jobs(jobs, experiment=f"sth_{GRAIN}", + slurm_params={"slurm_partition": "preempted", "slurm_gres": "gpu:1", "cpus_per_task": 8, + "mem_gb": 48, "timeout_min": 120, "slurm_constraint": "[a40|a6000|l40s]"}, + log_dir=f"sth_{GRAIN}", wait_for_completion=False) + + +if __name__ == "__main__": + main() diff --git a/src/ops_model/models/attention/diffex/figures/gen_validation/std_anchor_test.py b/src/ops_model/models/attention/diffex/figures/gen_validation/std_anchor_test.py new file mode 100644 index 0000000..366b80a --- /dev/null +++ b/src/ops_model/models/attention/diffex/figures/gen_validation/std_anchor_test.py @@ -0,0 +1,60 @@ +"""Is the baseline>>valid200 gap a STANDARDIZATION-anchor artifact? Re-score both caches on the same genes under +several gen-standardization schemes. If the gap closes under 'real_pop'/'none' vs 'gen_a0', the per-domain +gen-α=0 anchor (which changed meaning under DDIM inversion) was mismeasuring — not a generation regression. +""" +import json, glob, os +import numpy as np + +CV = "/hpc/projects/icd.fast.ops/analysis/figure4_traversals" +cen = np.load(f"{CV}/gen_real_centroid/geneKO_centroids.npz", allow_pickle=True) +names = list(cen["names"]); cidx = {c: i for i, c in enumerate(names)} +mu_r, sd_r = cen["mu"], cen["sd"] + +# same 15 genes as the ablation +GENES = sorted(os.path.basename(f)[:-4] for f in glob.glob(f"{CV}/gen_real_map_cache_stepabl_s100/geneKO/*.npz")) + + +def load(cache, cap): + """→ {gene: {alpha: (n,1024)}} + pooled α=0 for the 15 genes.""" + per, a0 = {}, [] + for g in GENES: + f = f"{cache}/geneKO/{g}.npz" + if not os.path.exists(f) or g not in cidx: continue + d = np.load(f, allow_pickle=True); al = list(np.asarray(d["alphas"], float)) + i0 = int(np.argmin(np.abs(np.array(al)))) + per[g] = {a: (None if d["gen"][i] is None else np.asarray(d["gen"][i], np.float32)[:cap]) for i, a in enumerate(al)} + if per[g][al[i0]] is not None: a0.append(per[g][al[i0]]) + return per, np.concatenate(a0) + + +def recover(per, mu_g, sd_g, scheme): + czr = (cen["cents"] - mu_r) / sd_r; czr = czr / (np.linalg.norm(czr, axis=1, keepdims=True) + 1e-9) + czn = cen["cents"] / (np.linalg.norm(cen["cents"], axis=1, keepdims=True) + 1e-9) + best = None + alphas = sorted({a for g in per for a in per[g]}) + for a in alphas: + t1s = [] + for g in per: + gv = per[g].get(a) + if gv is None or not len(gv): continue + if scheme == "gen_a0": gz = (gv - mu_g) / sd_g; cz = czr + elif scheme == "real_pop": gz = (gv - mu_r) / sd_r; cz = czr + elif scheme == "none": gz = gv; cz = czn + gz = gz / (np.linalg.norm(gz, axis=1, keepdims=True) + 1e-9) + order = np.argsort(-(gz @ cz.T), axis=1) + t1s.append(np.mean(order[:, 0] == cidx[g])) + m = np.mean(t1s) + if best is None or m > best[1]: best = (a, m) + return best + + +if __name__ == "__main__": + print(f"n genes = {len(GENES)} (cap=30, both caches have >=30 cells)\n") + caches = [("baseline 30c w=2 (old frames)", f"{CV}/gen_real_map_cache"), + ("valid200 w=1.5 INVERTED", f"{CV}/gen_real_map_cache_valid200")] + for scheme in ("gen_a0", "real_pop", "none"): + print(f"--- standardization = {scheme} ---") + for lbl, c in caches: + per, a0 = load(c, 30); mu_g, sd_g = a0.mean(0), a0.std(0) + 1e-6 + a, m = recover(per, mu_g, sd_g, scheme) + print(f" {lbl:34s}: peak α={a:+.1f} top1={m:.1%}") diff --git a/src/ops_model/models/attention/diffex/figures/gen_validation/stepabl_compare.py b/src/ops_model/models/attention/diffex/figures/gen_validation/stepabl_compare.py new file mode 100644 index 0000000..c706ff5 --- /dev/null +++ b/src/ops_model/models/attention/diffex/figures/gen_validation/stepabl_compare.py @@ -0,0 +1,71 @@ +"""Step-ablation comparison on the fixed 15-gene set: centroid recovery + SetTransformer, at matched 45-cell bag. +50-step arm = existing valid200 w=1.5 cache (capped to 45); 100/200 = stepabl caches. Baseline = orig 30-cell w=2. +""" +import json, glob, os +import numpy as np + +CV = "/hpc/projects/icd.fast.ops/analysis/figure4_traversals" +B = "/hpc/projects/icd.fast.ops/models/diffex" +cen = np.load(f"{CV}/gen_real_centroid/geneKO_centroids.npz", allow_pickle=True) +names = list(cen["names"]); cidx = {c: i for i, c in enumerate(names)} +cz = (cen["cents"] - cen["mu"]) / cen["sd"]; cz = cz / (np.linalg.norm(cz, axis=1, keepdims=True) + 1e-9) + +GENES = set(json.load(open(f"{CV}/gen_real_map_cache_stepabl_s100/geneKO/AACS.npz", allow_pickle=True).files) if False else + [os.path.basename(f)[:-4] for f in glob.glob(f"{CV}/gen_real_map_cache_stepabl_s100/geneKO/*.npz")]) + + +def centroid_recovery(cache, cap=45): + """per-α mean top1/top5/mAP over the 15 genes, α=0 baseline = argmin|α|, gen capped to `cap`.""" + caches = [f"{cache}/geneKO/{g}.npz" for g in GENES if os.path.exists(f"{cache}/geneKO/{g}.npz")] + g0 = [] + for f in caches: + d = np.load(f, allow_pickle=True); a0 = int(np.argmin(np.abs(np.asarray(d["alphas"], float)))); z = d["gen"][a0] + if z is not None and len(z): g0.append(np.asarray(z, np.float32)[:cap]) + mu, sd = np.concatenate(g0).mean(0), np.concatenate(g0).std(0) + 1e-6 + by = {} + for f in caches: + d = np.load(f, allow_pickle=True); g = str(d["gene"]); al = list(d["alphas"]) + if g not in cidx: continue + ti = cidx[g] + for ai, a in enumerate(al): + gv = d["gen"][ai] + if gv is None or not len(gv): continue + gz = (np.asarray(gv, np.float32)[:cap] - mu) / sd; gz = gz / (np.linalg.norm(gz, axis=1, keepdims=True) + 1e-9) + order = np.argsort(-(gz @ cz.T), axis=1); rk = np.where(order == ti)[1] + 1 + by.setdefault(a, {"t1": [], "t5": [], "mp": []}) + by[a]["t1"].append(np.mean(order[:, 0] == ti)); by[a]["t5"].append(np.mean([ti in r[:5] for r in order])) + by[a]["mp"].append(np.mean(1.0 / rk)) + al = sorted(by); mp = [np.mean(by[a]["mp"]) for a in al]; t1 = [np.mean(by[a]["t1"]) for a in al]; t5 = [np.mean(by[a]["t5"]) for a in al] + k = int(np.argmax(mp)); return al[k], mp[k], t1[k], t5[k] + + +def settransformer(tree): + """mean peak-α P(target), median rank, top5% over the 15 genes from scores_v5.json in `tree`.""" + P, RK, T5, al = [], [], [], None + for g in GENES: + f = f"{B}/{tree}/phase/geneKO/{g}/scores_v5.json" + if not os.path.exists(f): continue + s = json.load(open(f)); al = s["alphas"]; P.append(s["p_target"]); RK.append(s["rank_target"]); T5.append(s["top5_target"]) + if not P: return None + P = np.array(P, float); RK = np.array(RK, float); T5 = np.array(T5, float) + k = int(np.argmax(P.mean(0))); return al[k], P[:, k].mean(), np.median(RK[:, k]), T5[:, k].mean(), len(P) + + +if __name__ == "__main__": + print(f"n genes = {len(GENES)}") + print("\n=== CENTROID RECOVERY (peak α) — 15 genes, 45-cell bag ===") + for lbl, c in [("baseline 30c w=2 (orig)", f"{CV}/gen_real_map_cache"), + ("50-step w=1.5 (valid200)", f"{CV}/gen_real_map_cache_valid200"), + ("100-step w=1.5 INVERTED ", f"{CV}/gen_real_map_cache_stepabl_s100"), + ("100-step w=1.5 RANDOM-xT", f"{CV}/gen_real_map_cache_stepabl_s100_randxt"), + ("200-step w=1.5 INVERTED ", f"{CV}/gen_real_map_cache_stepabl_s200")]: + try: + a, mp, t1, t5 = centroid_recovery(c); print(f" {lbl}: α={a:+.1f} mAP={mp:.3f} top1={t1:.1%} top5={t5:.1%}") + except Exception as e: + print(f" {lbl}: ERR {e}") + print("\n=== SetTransformer (peak-α mean over 15 genes) ===") + for lbl, t in [("50-step w=1.5 (valid200, bag=200!)", "viewer_assets_valid200"), + ("100-step w=1.5 (stepabl, bag=45)", "viewer_assets_stepabl_s100"), + ("200-step w=1.5 (stepabl, bag=45)", "viewer_assets_stepabl_s200")]: + r = settransformer(t) + if r: a, p, rk, t5, n = r; print(f" {lbl}: α={a:+.1f} P(target)={p:.3f} medRank={rk:.0f} top5={t5:.0%} (n={n})") diff --git a/src/ops_model/models/attention/diffex/figures/gen_validation/valid200_alphastep.py b/src/ops_model/models/attention/diffex/figures/gen_validation/valid200_alphastep.py new file mode 100644 index 0000000..99b1872 --- /dev/null +++ b/src/ops_model/models/attention/diffex/figures/gen_validation/valid200_alphastep.py @@ -0,0 +1,38 @@ +"""Is the valid200 rightward-α shift a real generation property (compressed α-steps) or an analysis bug? +Independently of scores_v5 and of any standardization, measure the raw CellDINO displacement of each α-frame +from the α=0 frame, per gene, averaged — for v5 vs valid200 at the α values present in BOTH grids. If valid200's +per-α displacement is compressed (smaller step per α), the phenotype emerges later in α and EVERY downstream +measure inherits the shift → generation, not analysis. +""" +import numpy as np, glob + +CV = "/hpc/projects/icd.fast.ops/analysis/figure4_traversals" +SHARED = [0.0, 0.5, 1.0, 2.0, 3.0, 4.0, 5.0] # α present in both grids + + +def profile(cache): + disp = {a: [] for a in SHARED} + for f in sorted(glob.glob(f"{cache}/geneKO/*.npz")): + d = np.load(f, allow_pickle=True); al = list(np.asarray(d["alphas"], float)) + i0 = int(np.argmin(np.abs(np.array(al)))) + z0 = d["gen"][i0] + if z0 is None or not len(z0): + continue + c0 = np.asarray(z0, np.float32).mean(0) # gene's α=0 centroid in raw CellDINO + for a in SHARED: + if a not in al: + continue + za = d["gen"][al.index(a)] + if za is None or not len(za): + continue + disp[a].append(np.linalg.norm(np.asarray(za, np.float32).mean(0) - c0)) + return {a: (np.mean(v) if v else np.nan) for a, v in disp.items()} + + +if __name__ == "__main__": + v5 = profile(f"{CV}/gen_real_map_cache") + v2 = profile(f"{CV}/gen_real_map_cache_valid200") + print(f"{'alpha':>6} {'v5 disp':>10} {'v200 disp':>10} {'ratio v200/v5':>14}") + for a in SHARED: + r = v2[a] / v5[a] if v5[a] else np.nan + print(f"{a:6.1f} {v5[a]:10.2f} {v2[a]:10.2f} {r:14.2f}") diff --git a/src/ops_model/models/attention/diffex/figures/gen_validation/valid200_cache_build.py b/src/ops_model/models/attention/diffex/figures/gen_validation/valid200_cache_build.py new file mode 100644 index 0000000..0e7b29d --- /dev/null +++ b/src/ops_model/models/attention/diffex/figures/gen_validation/valid200_cache_build.py @@ -0,0 +1,85 @@ +"""Build a gen_real_map-format cache for the 200-cell validation bag (viewer_assets_valid200). + +Per gene: reuse the real top-K cells from the existing gen_real_map_cache, and embed the 200 generated +webp frames at each of the 7 valid200 α through Cell-DINO (loaded once per shard). Output npz per gene +{real, gen (list[7] of (200,1024)), alphas (7), gene} — the exact format gen_real_map / gen_real_centroid / +gen_real_distinct consume, so the mAP suite can run on the 200-cell bag by pointing CACHE here. +""" +import os, glob +import numpy as np + +GRAIN = os.environ.get("V200_GRAIN", "geneKO") # geneKO | complex +ASSETS = os.environ.get("V200_ASSETS", "viewer_assets_valid200") # new multibag: viewer_assets_v5 +OUTCACHE = os.environ.get("V200_OUTCACHE", "gen_real_map_cache_valid200") # new multibag: gen_real_map_cache_v5new +V = f"/hpc/projects/icd.fast.ops/models/diffex/{ASSETS}/phase/{GRAIN}" +OLD = f"/hpc/projects/icd.fast.ops/analysis/figure4_traversals/gen_real_map_cache/{GRAIN}" +OUT = f"/hpc/projects/icd.fast.ops/analysis/figure4_traversals/{OUTCACHE}/{GRAIN}" +NCELL = int(os.environ.get("V200_NCELL", "200")) + + +def _alphas(): + """Read the actual α grid from a gene's meta.json (new v5 = 17-pt [-5..5]); fallback to the 7-pt valid200 grid.""" + import json + for d in sorted(glob.glob(f"{V}/*")): + mp = f"{d}/meta.json" + if os.path.exists(mp): + return np.asarray(json.load(open(mp))["alphas"], np.float32) + return np.array([0.0, 0.5, 1.0, 2.0, 3.0, 4.0, 5.0], np.float32) + + +ALPHAS = _alphas() +NA = len(ALPHAS) + + +def build_shard(genes): + import torch + from PIL import Image + from ops_model.models.cell_dino import CellDinoModel + os.makedirs(OUT, exist_ok=True) + model = CellDinoModel(z_score=True) + + def emb(imgs): # (N,1,H,W) float32 → (N,1024) + out = [] + with torch.inference_mode(): + for i in range(0, len(imgs), 256): + out.append(model.extract_features({"data": torch.as_tensor(imgs[i:i + 256])}).float().cpu().numpy()) + return np.concatenate(out).astype(np.float32) + + for g in genes: + if os.path.exists(f"{OUT}/{g}.npz"): + continue # resume: already built + oc = f"{OLD}/{g}.npz" # reuse old real ref if present; else empty (real not needed for + real = np.asarray(np.load(oc, allow_pickle=True)["real"], np.float32) if os.path.exists(oc) else np.zeros((0, 1024), np.float32) # SetTransformer/distinct; centroid-recovery uses centroids.npz) + gen = [] + for ai in range(NA): + imgs = [] + for c in range(NCELL): + f = f"{V}/{g}/cell{c}/frame_{ai:02d}.webp" + if os.path.exists(f): + imgs.append(np.asarray(Image.open(f).convert("L"), np.float32) / 255.0 * 2 - 1) + gen.append(emb(np.stack(imgs)[:, None].astype(np.float32)) if imgs else None) + np.savez(f"{OUT}/{g}.npz", real=real, gen=np.array(gen, dtype=object), alphas=ALPHAS, gene=g) + return len(genes) + + +def main(): + from ops_utils.hpc.slurm_batch_utils import submit_parallel_jobs + gf = os.environ.get("V200_GENES_FILE") + if gf: # explicit gene list (e.g. the 42 dropped) + genes = [g.strip() for g in open(gf) if g.strip()] + else: + genes = sorted(os.path.basename(d) for d in glob.glob(f"{V}/*") + if os.path.isdir(d) and "__to__" not in os.path.basename(d) and not os.path.basename(d).startswith("_")) + ch = int(os.environ.get("V200_CHUNK", "40")) + shards = [genes[i:i + ch] for i in range(0, len(genes), ch)] + jobs = [{"name": f"v200cache_{i}", "func": build_shard, "kwargs": {"genes": s}} for i, s in enumerate(shards)] + print(f"[valid200-cache] {len(genes)} genes → {len(jobs)} GPU shards") + submit_parallel_jobs(jobs, experiment="valid200_cache", + slurm_params={"slurm_partition": "preempted", "slurm_gres": "gpu:1", "cpus_per_task": 10, + "mem_gb": 64, "timeout_min": 180, + "slurm_constraint": "[a40|a6000|l40s]"}, # preempted queue + weak GPUs — light embed job + log_dir="valid200_cache", wait_for_completion=False) + + +if __name__ == "__main__": + main() diff --git a/src/ops_model/models/attention/diffex/figures/gen_validation/valid200_capcheck.py b/src/ops_model/models/attention/diffex/figures/gen_validation/valid200_capcheck.py new file mode 100644 index 0000000..0ccb636 --- /dev/null +++ b/src/ops_model/models/attention/diffex/figures/gen_validation/valid200_capcheck.py @@ -0,0 +1,41 @@ +"""Diagnostic: is the valid200 effect-size gap to v5 driven by bag COMPOSITION (extra generic cells drag the +mean) or by the generation itself? Recompute α=5 nearest-centroid top-1 for valid200 using first-{45,100,200} +cells vs the v5 (~45-cell) bag, with the CORRECT α=0 baseline (argmin|α|). Faithful centroids reused from v5. +""" +import numpy as np, glob + +CV = "/hpc/projects/icd.fast.ops/analysis/figure4_traversals" +cen = np.load(f"{CV}/gen_real_centroid/geneKO_centroids.npz", allow_pickle=True) +names = list(cen["names"]); cidx = {c: i for i, c in enumerate(names)} +cz = (cen["cents"] - cen["mu"]) / cen["sd"]; cz = cz / (np.linalg.norm(cz, axis=1, keepdims=True) + 1e-9) + + +def collect(cache): + A0, A5, GN = [], [], [] + for f in sorted(glob.glob(f"{cache}/geneKO/*.npz")): + d = np.load(f, allow_pickle=True); g = str(d["gene"]); al = list(d["alphas"]) + if g not in cidx: + continue + i0 = int(np.argmin(np.abs(np.array(al)))); i5 = int(np.argmin(np.abs(np.array(al) - 5.0))) + z0, z5 = d["gen"][i0], d["gen"][i5] + if z0 is None or z5 is None or not len(z0) or not len(z5): + continue + A0.append(np.asarray(z0, np.float32)); A5.append(np.asarray(z5, np.float32)); GN.append(g) + return A0, A5, GN + + +def top1(A0, A5, GN, cap=None): + mu = np.concatenate([a[:cap] for a in A0]).mean(0); sd = np.concatenate([a[:cap] for a in A0]).std(0) + 1e-6 + t = [] + for z5, g in zip(A5, GN): + gz = (z5[:cap] - mu) / sd; gz = gz / (np.linalg.norm(gz, axis=1, keepdims=True) + 1e-9) + order = np.argsort(-(gz @ cz.T), axis=1); t.append(np.mean(order[:, 0] == cidx[g])) + return np.mean(t), len(t) + + +if __name__ == "__main__": + A0, A5, GN = collect(f"{CV}/gen_real_map_cache_valid200") + for cap in (45, 100, 200): + m, n = top1(A0, A5, GN, cap); print(f"valid200 first-{cap:3d} @a=5 top1={m:.1%} (n={n})") + A0v, A5v, GNv = collect(f"{CV}/gen_real_map_cache") + m, n = top1(A0v, A5v, GNv); print(f"v5 (~45-cell) @a=5 top1={m:.1%} (n={n})") diff --git a/src/ops_model/models/attention/diffex/figures/gen_validation/valid200_map_compare.py b/src/ops_model/models/attention/diffex/figures/gen_validation/valid200_map_compare.py new file mode 100644 index 0000000..aea1814 --- /dev/null +++ b/src/ops_model/models/attention/diffex/figures/gen_validation/valid200_map_compare.py @@ -0,0 +1,105 @@ +"""Compare the 200-cell validation bag (valid200) vs the 45-cell bag on the Cell-DINO mAP measures, both grains. + +Centroid recovery (gen cells -> nearest faithful 1000-real centroid; same centroids for both bags, only the +generated bag differs) and within-domain distinctiveness. Rows = grain (geneKO, complex); cols = centroid +mAP / top-1 / top-5 / distinctiveness. geneKO has both bags; complex only exists at the 45-cell bag (no +200-cell complex generation). thr!=None restricts to the SetTransformer-distinguishable subset: classes whose +REAL cells clear top1_acc > thr @ bag20 (real_acc20.json) — a FIXED set (same classes for both bags). Reads +gen_real_centroid.score + gen_real_distinct JSONs per cache. +""" +import json, glob, os +import numpy as np +from ops_model.models.attention.diffex.classifier.config import slugify + +CV = "/hpc/projects/icd.fast.ops/analysis/figure4_traversals" +OUT = f"{CV}/valid200_metrics" +BIG, SM = "#c0392b", "#888" +REAL_ACC = json.load(open("/hpc/projects/icd.fast.ops/models/diffex/viewer_assets_v5/real_acc20.json")) + +# (label, color, ls, centroid_dir, distinct_dir) per grain; distinct na inferred from glob +BAGS = { + "geneKO": [("200-cell bag", BIG, "-", f"{CV}/gen_real_centroid_valid200", f"{CV}/gen_real_distinct_valid200"), + ("45-cell bag", SM, "--", f"{CV}/gen_real_centroid", f"{CV}/gen_real_distinct")], + "complex": [("45-cell bag", SM, "--", f"{CV}/gen_real_centroid", f"{CV}/gen_real_distinct")], +} +GRAIN_LBL = {"geneKO": "Gene-level (distinctiveness)", "complex": "Protein complex (EBI)"} + + +def _keep(grain, thr): + """SetTransformer-distinguishable subset: real cells top1_acc > thr @ bag20. Slug-keyed (matches complex + names with spaces). None = all classes.""" + if thr is None: + return None + pre = f"phase/{grain}/" + return {k[len(pre):] for k, v in REAL_ACC.items() if k.startswith(pre) and v > thr} + + +def _in(c, keep): + return keep is None or slugify(c) in keep + + +def _cent(centdir, grain, keep): + d = json.load(open(f"{centdir}/{grain}_scored.json")); al = d["alphas"]; by = d["by_alpha"] + def f(k): + return np.array([np.mean([v for c, v in by[str(a)][k].items() if _in(c, keep)]) for a in al]) + n = sum(_in(c, keep) for c in by[str(al[0])]["map"]) + return np.array(al), f("map"), f("top1"), f("top5"), n + + +def _dist(distdir, grain, keep): + na = len(glob.glob(f"{distdir}/{grain}_gen_a*.json")) + real = json.load(open(f"{distdir}/{grain}_real.json")) + rv = [v for c, v in real.items() if _in(c, keep)] + rm = float(np.median(rv)) if rv else np.nan + al, med = [], [] + for ai in range(na): + g = json.load(open(f"{distdir}/{grain}_gen_a{ai}.json")) + vv = [v for c, v in g["gen"].items() if _in(c, keep)] + al.append(g["alpha"]); med.append(np.median(vv) if vv else np.nan) + return np.array(al), np.array(med), rm + + +def plot(thr=None, suffix=""): + import matplotlib; matplotlib.use("Agg") + import matplotlib.pyplot as plt + plt.rcParams["pdf.fonttype"] = 42 + os.makedirs(OUT, exist_ok=True) + grains = list(BAGS) + fig, axes = plt.subplots(len(grains), 4, figsize=(21, 5 * len(grains)), squeeze=False) + for gi, grain in enumerate(grains): + ax = axes[gi] + for lbl, c, ls, cd, dd in BAGS[grain]: + keep = _keep(grain, thr) + al, mp, t1, t5, n = _cent(cd, grain, keep) + tag = f"{lbl} (n={n})" + ax[0].plot(al, mp, ls, color=c, lw=2.4, label=tag) + ax[1].plot(al, t1 * 100, ls, color=c, lw=2.4, label=tag) + ax[2].plot(al, t5 * 100, ls, color=c, lw=2.4, label=tag) + ald, med, rm = _dist(dd, grain, keep) + ax[3].plot(ald, med, ls, color=c, lw=2.4, label=lbl) + ax[3].axhline(rm, color=c, ls=":", lw=1.5, label=f"{lbl} real ({rm:.3f})") + ax[0].set_ylabel(f"{GRAIN_LBL[grain]}\n\nmAP (1/rank true centroid)") + ax[0].set_title("centroid-recovery mAP vs α"); ax[1].set_title("centroid top-1 vs α") + ax[2].set_title("centroid top-5 vs α"); ax[3].set_title("distinctiveness median mAP vs α") + ax[1].set_ylabel("% cells nearest true centroid"); ax[2].set_ylabel("% cells within top-5") + ax[3].set_ylabel("median distinctiveness / EBI mAP") + for a in ax: + a.set_xlabel("traversal α"); a.grid(alpha=.25); a.axvline(0, color="#ccc", lw=1); a.legend(fontsize=9) + ttl = "Cell-DINO mAP: 200-cell vs 45-cell bag (faithful real centroids)" + if thr is not None: + ttl += f" — SetTransformer-distinguishable subset (real top1_acc > {thr} @ bag20)" + fig.suptitle(ttl, fontweight="bold", fontsize=15) + fig.tight_layout() + for e in ("png", "svg"): + fig.savefig(f"{OUT}/valid200_map_compare{suffix}.{e}", dpi=150, bbox_inches="tight") + plt.close(fig); print(f"saved valid200_map_compare{suffix}") + for grain in grains: + for lbl, c, ls, cd, dd in BAGS[grain]: + keep = _keep(grain, thr) + al, mp, t1, t5, n = _cent(cd, grain, keep); k = int(np.argmax(mp)) + print(f" {grain:8s} {lbl:12s} n={n:4d}: peak α={al[k]:+.1f} mAP={mp[k]:.3f} top1={t1[k]:.1%} top5={t5[k]:.1%}") + + +if __name__ == "__main__": + plot() # all classes + plot(thr=0.5, suffix="_acc50") # SetTransformer-distinguishable subset (real top1_acc>0.5 @bag20) diff --git a/src/ops_model/models/attention/diffex/figures/gen_validation/valid200_metrics.py b/src/ops_model/models/attention/diffex/figures/gen_validation/valid200_metrics.py new file mode 100644 index 0000000..75e7257 --- /dev/null +++ b/src/ops_model/models/attention/diffex/figures/gen_validation/valid200_metrics.py @@ -0,0 +1,72 @@ +"""SetTransformer rank / set-accuracy on the 200-cell validation bag (viewer_assets_valid200), vs the small bag. + +valid200 = 1000 phase geneKO traversals with 200 generated cells/class (bigger bag → more evidence for the +set classifier). scores_v5.json is already computed inline (P(target), rank_target, top1/top5_target per α). +This aggregates them across all classes and overlays the small-bag reference (viewer_assets_v5, ~45-cell bag). + +mAP is separate (needs Cell-DINO embeddings of the 200-cell bags → gen_real_map); this file is the classifier +rank / set-accuracy only (no GPU — pure JSON aggregation). +""" +import json, glob, os +import numpy as np + +BASE = "/hpc/projects/icd.fast.ops/models/diffex" +BIG = f"{BASE}/viewer_assets_valid200/phase/geneKO" # 200-cell bag +SMALL = f"{BASE}/viewer_assets_v5/phase/geneKO" # ~45-cell bag (reference) +OUT = "/hpc/projects/icd.fast.ops/analysis/figure4_traversals/valid200_metrics" + + +def _collect(d): + """→ (alphas, P, RK, T1, T5) each (n_class, n_alpha) from scores_v5.json under dir d.""" + al, P, RK, T1, T5, genes = None, [], [], [], [], [] + for f in sorted(glob.glob(f"{d}/*/scores_v5.json")): + s = json.load(open(f)); al = s["alphas"] + P.append(s["p_target"]); RK.append(s["rank_target"]); T1.append(s["top1_target"]); T5.append(s["top5_target"]) + genes.append(os.path.basename(os.path.dirname(f))) + return np.array(al), np.array(P), np.array(RK, float), np.array(T1, float), np.array(T5, float), genes + + +def summary(): + os.makedirs(OUT, exist_ok=True) + al, P, RK, T1, T5, genes = _collect(BIG) + bag = json.load(open(f"{BIG}/{genes[0]}/scores_v5.json")).get("bag") + pk = int(np.argmax(P.mean(0))) # peak α by mean P(target) + print(f"[valid200] {len(genes)} geneKO, bag={bag}, α={list(np.round(al,1))}") + print(f" peak α={al[pk]:+.1f}: mean P(target)={P[:,pk].mean():.2f} median rank={np.median(RK[:,pk]):.0f} " + f"top1={T1[:,pk].mean():.0%} top5={T5[:,pk].mean():.0%}") + # real-distinguishable-ish subset: classes that ever reach top-1 (min rank == 1) + dist = RK.min(1) == 1 + print(f" among classes reaching top-1 at some α (n={dist.sum()}): peak top5={T5[dist, pk].mean():.0%}") + return al, P, RK, T1, T5, bag + + +def plot(): + import matplotlib; matplotlib.use("Agg") + import matplotlib.pyplot as plt + plt.rcParams["pdf.fonttype"] = 42 + os.makedirs(OUT, exist_ok=True) + alB, PB, RKB, T1B, T5B, gB = _collect(BIG); bagB = json.load(open(f"{BIG}/{gB[0]}/scores_v5.json")).get("bag") + alS, PS, RKS, T1S, T5S, gS = _collect(SMALL); bagS = json.load(open(f"{SMALL}/{gS[0]}/scores_v5.json")).get("bag") + fig, ax = plt.subplots(1, 3, figsize=(17, 5)) + for al, P, RK, T1, T5, bag, c, ls in [(alB, PB, RKB, T1B, T5B, bagB, "#c0392b", "-"), + (alS, PS, RKS, T1S, T5S, bagS, "#888", "--")]: + lbl = f"bag={bag}" + m = P.mean(0); q1, q3 = np.percentile(P, [25, 75], 0) + ax[0].plot(al, m, ls, color=c, lw=2.4, label=lbl); ax[0].fill_between(al, q1, q3, color=c, alpha=.12, lw=0) + ax[1].plot(al, np.median(RK, 0), ls, color=c, lw=2.4, label=lbl) + ax[2].plot(al, (T5).mean(0) * 100, ls, color=c, lw=2.4, label=f"{lbl} top-5") + ax[2].plot(al, (T1).mean(0) * 100, ls, color=c, lw=1.4, alpha=.7, label=f"{lbl} top-1") + ax[0].set_title("P(target class) vs α (mean, IQR)"); ax[0].set_ylabel("P(target)"); ax[0].set_ylim(-.02, 1.02) + ax[1].set_title("median target rank vs α"); ax[1].set_ylabel("rank (1 = top pick)"); ax[1].set_yscale("log"); ax[1].axhline(1, color="#ccc", lw=1) + ax[2].set_title("% recovered vs α"); ax[2].set_ylabel("% of geneKOs"); ax[2].set_ylim(-2, 102) + for a in ax: + a.set_xlabel("traversal α"); a.grid(alpha=.25); a.legend(fontsize=8) + fig.suptitle(f"v5 SetTransformer on the 200-cell bag (n={len(gB)} geneKO) vs the {bagS}-cell bag", fontweight="bold") + fig.tight_layout() + for e in ("png", "svg"): + fig.savefig(f"{OUT}/valid200_setacc.{e}", dpi=150, bbox_inches="tight") + plt.close(fig); print(f"saved {OUT}/valid200_setacc") + + +if __name__ == "__main__": + summary(); plot() diff --git a/src/ops_model/models/attention/diffex/figures/nc_ratio.py b/src/ops_model/models/attention/diffex/figures/nc_ratio.py new file mode 100644 index 0000000..649b585 --- /dev/null +++ b/src/ops_model/models/attention/diffex/figures/nc_ratio.py @@ -0,0 +1,219 @@ +"""Nuclear/cytoplasmic proteasome-intensity ratio (PSMB7) — real cells + DiffEx traversal, replacing the +tubular seg (which over-fragments as signal shifts nucleus<->cytoplasm). + +Masks are proteasome-INDEPENDENT and framed IDENTICALLY to the generated frames (same materialize_crops +pipeline the traversal build used): nucleus = the aligned `nuclei_prediction` DNN channel; cytoplasm = +cell foreground minus nucleus. Metric = mean GFP(nucleus) / mean GFP(cytoplasm) — scale-invariant, so the +8-bit display frames are valid. Cell index == rank order of the generation ranking, so each generated cell +reuses its own anchor's nucleus. + +Outputs a full panel (real | α0 | α1 | α3 image row + nucleus/cytoplasm overlay row) + the N/C violin. +Run (SLURM): python nc_ratio.py --submit +""" +import json +import os +import sys + +import matplotlib + +matplotlib.use("Agg") +import matplotlib.pyplot as plt +import numpy as np +import pandas as pd +from PIL import Image +from scipy import ndimage as ndi +from skimage.filters import threshold_otsu, gaussian +from skimage.morphology import binary_closing, disk + +from _setacc_common import _materialize + +plt.rcParams["pdf.fonttype"] = 42 +plt.rcParams["svg.fonttype"] = "none" + +VA = "/hpc/projects/icd.fast.ops/models/diffex/viewer_assets_v5" +OUT = "/hpc/projects/icd.fast.ops/analysis/figure4_nc_ratio" +MARKER_DIR = "proteasome_PSMB7" +MC = "proteasome_PSMB7" # marker_channel (DirConfig); channel to READ is passed separately +TARGET = "PSMB6" +GRAIN = "geneKO" +RANK = f"{VA}/_rankings/fluor/geneKO/proteasome_PSMB7.parquet" # the ranking the traversal was generated from +N_REAL = 300 +COLORS = {"real": "#999999", "KO": "#2e8b57", "α=0": "#c6dbef", "α=1": "#6baed6", "α=3": "#08519c"} +EXAMPLE_CELL = 0 + + +def _rank_df(gene, n=None): + d = pd.read_parquet(RANK) + d = d[d["gene"].astype(str) == gene] + if "rank_type" in d.columns: + d = d[d["rank_type"] == "top"] + d = d.sort_values("rank").reset_index(drop=True) + return d.head(n) if n else d + + +def _central(binmask): + lab, n = ndi.label(binmask) + if n == 0: + return binmask + cy, cx = np.array(binmask.shape) // 2 + cen = lab[cy, cx] or (1 + int(np.argmax(np.bincount(lab.ravel())[1:]))) + return lab == cen + + +def _nucleus(pred): + """Nucleus mask from the aligned nuclei_prediction crop: blurred Otsu, fill, central component.""" + g = gaussian(pred, 2); v = g[g > 0] + if v.size == 0: + return np.zeros(pred.shape, bool) + m = ndi.binary_fill_holes(binary_closing(g > threshold_otsu(v), disk(3))) + return _central(m) + + +def _foreground(gray): + """Cell foreground (drops background so 'all other pixels' don't include empty corners).""" + g = gaussian(gray, 2); v = g[g > 0] + if v.size == 0: + return np.ones(gray.shape, bool) + fg = ndi.binary_fill_holes(binary_closing(g > np.percentile(g, 55), disk(3))) + return _central(fg) if fg.any() else fg + + +def _ratio(gray, nucleus, cyto): + if nucleus.sum() < 20 or cyto.sum() < 20: + return np.nan + return float(gray[nucleus].mean() / (gray[cyto].mean() + 1e-6)) + + +def _crops(df, channel): + """materialize_crops the given channel for these cells — identical framing to the generated frames.""" + raw, recs = _materialize(df.assign(gene=TARGET), MC, channel, TARGET) + return raw[:, 0], recs # (N, H, W) + + +def real_nc(gene, n): + df = _rank_df(gene, n * 2) + gfp, recs = _crops(df, "GFP") + nuc, _ = _crops(df, "nuclei_prediction") + out = [] + for i in range(min(len(gfp), len(nuc))): + nm = _nucleus(nuc[i]); fg = _foreground(gfp[i]); cy = fg & ~nm + v = _ratio(gfp[i], nm, cy) + if np.isfinite(v): + out.append(v) + if len(out) >= n: + break + return np.array(out) + + +def gen_nc(alphas_show=(0, 1, 3)): + md = f"{VA}/{MARKER_DIR}/{GRAIN}/{TARGET}" + al = np.array(json.load(open(f"{md}/meta.json"))["alphas"]) + idxs = {a: int(np.argmin(np.abs(al - a))) for a in alphas_show} + ncell = len([d for d in os.listdir(md) if d.startswith("cell")]) + anch = _rank_df(TARGET, ncell) + nuc, _ = _crops(anch, "nuclei_prediction") # anchor nucleus per cell, aligned to frames + per = {a: [] for a in alphas_show} + for c in range(min(ncell, len(nuc))): + nm = _nucleus(nuc[c]) + for a, i in idxs.items(): + fp = f"{md}/cell{c}/frame_{i:02d}.webp" + if not os.path.exists(fp): + continue + gray = np.asarray(Image.open(fp).convert("L"), np.float32) + nmr = nm if nm.shape == gray.shape else _rs(nm, gray.shape) + fg = _foreground(gray); cy = fg & ~nmr + v = _ratio(gray, nmr, cy) + if np.isfinite(v): + per[a].append(v) + return {a: np.array(v) for a, v in per.items()}, idxs, md, nuc + + +def _rs(mask, shape): + from skimage.transform import resize + return resize(mask.astype(float), shape) > 0.5 + + +def _ov(gray, nm, cy): + rgb = np.stack([gray] * 3, -1) / max(gray.max(), 1e-6) + rgb[nm] = 0.5 * rgb[nm] + 0.5 * np.array([0.15, 0.55, 1.0]) + rgb[cy] = 0.75 * rgb[cy] + 0.25 * np.array([1.0, 0.5, 0.1]) + return np.clip(rgb, 0, 1) + + +def panel(gen, idxs, md, nuc, real_gfp_ex, real_nuc_ex): + """Full panel: real | α0/α1/α3 image row + nucleus(blue)/cytoplasm(orange) overlay row, violin at right.""" + from skimage.transform import resize + c = EXAMPLE_CELL + + def _fit(pred, shape): + return pred if pred.shape == shape else resize(pred, shape, preserve_range=True) + cols = [("real", real_gfp_ex, real_nuc_ex)] + for a, i in idxs.items(): + g = np.asarray(Image.open(f"{md}/cell{c}/frame_{i:02d}.webp").convert("L"), np.float32) + cols.append((f"α={a}", g, _fit(nuc[c], g.shape))) + nc = len(cols) + fig = plt.figure(figsize=(nc * 2.0 + 4.2, 4.4), facecolor="white") + gs = fig.add_gridspec(2, nc + 2, width_ratios=[1] * nc + [0.25, 2.4], hspace=0.06, wspace=0.06, + left=0.02, right=0.985, top=0.9, bottom=0.06) + for j, (t, g, pred) in enumerate(cols): + nm = _nucleus(pred); fg = _foreground(g); cy = fg & ~nm + ax = fig.add_subplot(gs[0, j]); ax.imshow(g, cmap="gray"); ax.set_title(t, fontsize=13); ax.axis("off") + ax2 = fig.add_subplot(gs[1, j]); ax2.imshow(_ov(g, nm, cy)); ax2.axis("off") + axv = fig.add_subplot(gs[:, nc + 1]) + data = [gen.get(0, []), gen.get(1, []), gen.get(3, [])] + _violin_into(axv, PANEL_REAL, data) + os.makedirs(OUT, exist_ok=True) + for ext in ("png", "svg"): + fig.savefig(f"{OUT}/PSMB7_nc_ratio_panel.{ext}", dpi=220, bbox_inches="tight", facecolor="white") + plt.close(fig) + print(f"saved {OUT}/PSMB7_nc_ratio_panel", flush=True) + + +PANEL_REAL = {} + + +def _violin_into(ax, real, gen_data): + data = [real.get("NTC", []), real.get("KO", []), *gen_data] + labels = ["real", "KO", "α=0", "α=1", "α=3"] + keep = [i for i, d in enumerate(data) if len(d)] + parts = ax.violinplot([np.asarray(data[i]) for i in keep], positions=keep, showextrema=False, widths=0.82) + for pc, i in zip(parts["bodies"], keep): + pc.set_facecolor(COLORS[labels[i]]); pc.set_alpha(0.6); pc.set_edgecolor(COLORS[labels[i]]); pc.set_linewidth(1.5) + for i in keep: + ax.hlines(np.mean(data[i]), i - 0.34, i + 0.34, color="#222", lw=3, zorder=5) + ax.axhline(1.0, color="#999", lw=2, ls="--") + pooled = np.concatenate([np.asarray(data[i]) for i in keep]) + ylo, yhi = np.percentile(pooled, (1, 98)); pad = 0.05 * (yhi - ylo + 1e-9) + ax.set_ylim(ylo - pad, yhi + pad) + ax.set_xticks(range(len(labels))); ax.set_xticklabels(labels, fontsize=20) + ax.set_ylabel("Nuclear / cytoplasmic\nproteasome intensity", fontsize=18) + ax.tick_params(axis="y", labelsize=18, width=2, length=7); ax.tick_params(axis="x", length=0) + for s in ("top", "right"): + ax.spines[s].set_visible(False) + + +def build(): + ntc = real_nc("NTC", N_REAL); ko = real_nc(TARGET, N_REAL) + PANEL_REAL["NTC"] = ntc; PANEL_REAL["KO"] = ko + gen, idxs, md, nuc = gen_nc() + # example real crop for the panel's "real" column (top KO cell) + ex = _rank_df(TARGET, 1) + rg, _ = _crops(ex, "GFP"); rn, _ = _crops(ex, "nuclei_prediction") + m = lambda d: round(float(np.mean(d)), 3) if len(d) else None + print(f" real NTC {m(ntc)} (n{len(ntc)}) / KO {m(ko)} (n{len(ko)}) | gen α0 {m(gen.get(0,[]))} " + f"α1 {m(gen.get(1,[]))} α3 {m(gen.get(3,[]))}", flush=True) + panel(gen, idxs, md, nuc, rg[0], rn[0]) + + +def submit(): + import pathlib + from ops_utils.hpc.slurm_batch_utils import submit_parallel_jobs + figdir = str(pathlib.Path(__file__).resolve().parent) + os.environ["PYTHONPATH"] = figdir + os.pathsep + os.environ.get("PYTHONPATH", "") + submit_parallel_jobs([{"name": "psmb7_nc", "func": build, "kwargs": {}}], experiment="diffex_nc", + slurm_params={"slurm_partition": "cpu", "cpus_per_task": 8, "mem_gb": 64, "timeout_min": 120}, + log_dir="diffex_nc", wait_for_completion=False) + + +if __name__ == "__main__": + submit() if (len(sys.argv) > 1 and sys.argv[1] == "--submit") else build() diff --git a/src/ops_model/models/attention/diffex/figures/ntc_anchor_compare.py b/src/ops_model/models/attention/diffex/figures/ntc_anchor_compare.py new file mode 100644 index 0000000..b3a4727 --- /dev/null +++ b/src/ops_model/models/attention/diffex/figures/ntc_anchor_compare.py @@ -0,0 +1,63 @@ +"""Quick comparison grid: top-50 NTC anchor cells from the OLD v5-accuracy ranking vs the NEW phase +multirank (shap_screen). Crops the phase channel from phenotyping_v3.zarr (same source as the viewer).""" +import numpy as np +import pandas as pd +import matplotlib +matplotlib.use("Agg"); matplotlib.rcParams["pdf.fonttype"] = 42 +import matplotlib.pyplot as plt +import zarr +from ..viewer.build_pc_crops_masked import BASE, CROP_SIZE, PHASE_CHANNEL, _crop, _render_gray, _zarr_patch + +R = "/hpc/projects/icd.fast.ops/models/diffex/viewer_assets_v5/_rankings" +N = 50 + + +def top_ntc(parquet, n=N): + df = pd.read_parquet(parquet, columns=["gene", "experiment", "well", "x_pheno", "y_pheno", "rank"]) + return df[df["gene"].astype(str) == "NTC"].sort_values("rank").head(n).reset_index(drop=True) + + +def crops(rows): + _zarr_patch(); half = CROP_SIZE // 2; cache = {}; out = [] + for r in rows.itertuples(): + key = (r.experiment, r.well) + if key not in cache: + pos = f"{BASE}/{r.experiment}/3-assembly/phenotyping_v3.zarr/{r.well[0]}/{r.well[1:]}/0" + try: + cache[key] = zarr.open(f"{pos}/0", mode="r") + except Exception: + cache[key] = None + img = cache[key] + if img is None: + out.append(None); continue + try: + out.append(_render_gray(_crop(img, PHASE_CHANNEL, int(round(r.x_pheno)), int(round(r.y_pheno)), half))) + except Exception: + out.append(None) + return out + + +def panel(ax_grid, imgs, title): + for i in range(N): + a = ax_grid[i] + if i < len(imgs) and imgs[i] is not None: + a.imshow(imgs[i]); + a.set_xticks([]); a.set_yticks([]) + a.set_title(str(i + 1), fontsize=5, pad=1) + + +old = top_ntc(f"{R}/pma_v5_phase_geneKO.parquet") +new = top_ntc(f"{R}/pma_shap_phase_geneKO.parquet") +oi, ni = crops(old), crops(new) + +fig = plt.figure(figsize=(20, 22)) +outer = fig.add_gridspec(2, 1, hspace=0.12) +for row, (imgs, title) in enumerate([(oi, "OLD v5-accuracy ranking — top-50 NTC anchors"), + (ni, "NEW phase multirank (shap_screen) — top-50 NTC anchors")]): + inner = outer[row].subgridspec(5, 10, hspace=0.25, wspace=0.05) + axs = [fig.add_subplot(inner[j]) for j in range(N)] + panel(axs, imgs, title) + fig.text(0.5, 0.905 - row * 0.485, title, ha="center", fontsize=15, fontweight="bold") +out = "/hpc/projects/icd.fast.ops/analysis/ntc_anchor_old_vs_multirank.png" +fig.savefig(out, dpi=110, bbox_inches="tight"); plt.close(fig) +print("wrote", out) diff --git a/src/ops_model/models/attention/diffex/figures/phase_montages.py b/src/ops_model/models/attention/diffex/figures/phase_montages.py new file mode 100644 index 0000000..9587ab8 --- /dev/null +++ b/src/ops_model/models/attention/diffex/figures/phase_montages.py @@ -0,0 +1,12 @@ +"""Top-N phase montages (KO per group + a single phase NTC) for picking cells for panel E.""" +from _setacc_common import _materialize, slugify +from _setacc_phase import COLS_PHASE, PHASE_CH, phase_df, phase_ntc +from debug_setacc_top100 import render_montage + +for c in COLS_PHASE: + raw, recs = _materialize(phase_df(c["block"], c["key"]).head(100), None, PHASE_CH, c["key"]) + render_montage(raw, recs, f"KO — {c['top_label'].replace(chr(10),' ')} (phase) rank-ordered set-accuracy", + f"debug_phase_KO_{slugify(c['key'])[:30]}") + +raw, recs = _materialize(phase_ntc().head(100), None, PHASE_CH, "NTC") +render_montage(raw, recs, "NTC — phase rank-ordered set-accuracy", "debug_phase_NTC") diff --git a/src/ops_model/models/attention/diffex/figures/phase_multibag_montages.py b/src/ops_model/models/attention/diffex/figures/phase_multibag_montages.py new file mode 100644 index 0000000..e7fa696 --- /dev/null +++ b/src/ops_model/models/attention/diffex/figures/phase_multibag_montages.py @@ -0,0 +1,68 @@ +"""Top-100 PHASE cell montages for hand-picking representative cells, ranked by the multibag SHAP ranking +(pma_shap_phase_geneKO). One montage per gene, rank-ordered with rank + cell-key + conf, inverse blue seg +mask (reuses _materialize + render_montage). For the fig-4 new-phenotype review. + +Run: OPS_DIFFEX_ASSETS=viewer_assets_v5 python phase_multibag_montages.py [GENE ...] +""" +import sys + +import pandas as pd + +from _setacc_common import _materialize +import debug_setacc_top100 as D + +OUT_DIR = "/hpc/projects/icd.fast.ops/analysis/figure4_shap_montages" # shared SHAP-montage review dir (phase + fluor) +RANK = "/hpc/projects/icd.fast.ops/models/diffex/viewer_assets_v5/_rankings/pma_shap_phase_geneKO.parquet" +PHASE_CH = "Phase2D" +N = 100 + +GENES = [ # gene -> literature-backed KO phenotype to look for + ("KIF23", "multi-nucleation"), + ("CAPZB", "stretched morphology"), + ("SNRPD1", "dark vacuoles"), + ("SAMM50", "globular mitochondria"), + ("RAB7A", "enlarged & increased vesicles / lysosomes"), + ("NTC", "control (phase geneKO + complex NTC pool)"), +] + + +def main(genes=None): + D.OUT = OUT_DIR # set here (runs in SLURM worker; module-level is skipped by cloudpickle) + want = set(genes) if genes else None + for g, ph in GENES: + if want and g not in want: + continue + d = pd.read_parquet(RANK, filters=[("gene", "==", g)]) + if "rank_type" in d.columns: + d = d[d["rank_type"] == "top"] + d = d.sort_values("rank").head(N) + raw, recs = _materialize(d, None, PHASE_CH, g) + D.render_montage(raw, recs, f"KO — {g} ({ph}) · phase · multibag SHAP rank", f"phase_multibag_{g}") + + +def _job(gene): + import os + os.environ.setdefault("OPS_DIFFEX_ASSETS", "viewer_assets_v5") + main([gene]) + + +def submit(genes=None): + """One SLURM cpu job per gene (crop materialization is memory-heavy — dies on the login node).""" + import os + import pathlib + from ops_utils.hpc.slurm_batch_utils import submit_parallel_jobs + figdir = str(pathlib.Path(__file__).resolve().parent) + os.environ["PYTHONPATH"] = figdir + os.pathsep + os.environ.get("PYTHONPATH", "") + os.environ.setdefault("OPS_DIFFEX_ASSETS", "viewer_assets_v5") + gs = genes or [g for g, _ in GENES] + jobs = [{"name": f"phmont_{g}", "func": _job, "kwargs": {"gene": g}} for g in gs] + submit_parallel_jobs(jobs, experiment="diffex_phmont", + slurm_params={"slurm_partition": "cpu", "cpus_per_task": 8, "mem_gb": 64, "timeout_min": 90}, + log_dir="diffex_phmont", wait_for_completion=False) + + +if __name__ == "__main__": + if len(sys.argv) > 1 and sys.argv[1] == "--submit": + submit(sys.argv[2:] or None) + else: + main(sys.argv[1:] or None) diff --git a/src/ops_model/models/attention/diffex/figures/phase_sample_montages.py b/src/ops_model/models/attention/diffex/figures/phase_sample_montages.py new file mode 100644 index 0000000..34ca3d3 --- /dev/null +++ b/src/ops_model/models/attention/diffex/figures/phase_sample_montages.py @@ -0,0 +1,23 @@ +"""Phase KO montages for a sample of very-strong set-accuracy geneKOs/complexes to pick panel-E +gene/complex columns from. NTC = shared debug_phase_NTC.png.""" +from _setacc_common import _materialize, slugify +from _setacc_phase import PHASE_CH, phase_df +from debug_setacc_top100 import render_montage + +SAMPLE = [ + ("genes", "MICOS13"), ("genes", "KIF23"), ("genes", "CAPZB"), ("genes", "SAMM50"), + ("genes", "SON"), ("genes", "RAB7A"), ("genes", "SRSF3"), ("genes", "MTOR"), + ("complexes", "DNA-directed RNA polymerase I complex"), + ("complexes", "ESCRT-III complex"), + ("complexes", "TRAPP II complex, TRAPPC2 variant"), + ("complexes", "Nuclear pore complex"), + ("complexes", "COP9 signalosome variant 1"), +] + +for block, key in SAMPLE: + try: + raw, recs = _materialize(phase_df(block, key).head(100), None, PHASE_CH, key) + render_montage(raw, recs, f"KO — {key} (phase) rank-ordered set-accuracy", + f"debug_phase_KO_{slugify(key)[:34]}") + except Exception as e: + print(f"skip {key}: {e}") diff --git a/src/ops_model/models/attention/diffex/figures/rab_candidate_montages.py b/src/ops_model/models/attention/diffex/figures/rab_candidate_montages.py new file mode 100644 index 0000000..fc9bdb4 --- /dev/null +++ b/src/ops_model/models/attention/diffex/figures/rab_candidate_montages.py @@ -0,0 +1,22 @@ +"""Top-N montages (KO complex + per-marker NTC) for each cis-Golgi/Rab-slot candidate, so specific +cells can be picked. Complexes only have top-30 cells.""" +from cis_golgi_alternatives import CANDS +from debug_setacc_top100 import montage +from ops_model.models.attention.diffex.classifier.config import slugify + +ntc_done = set() +for c in CANDS: + try: + montage(c["mc"], c["ch"], "complexes", c["key"], + f"KO — {c['top_label']} · {c['marker_label'].replace(chr(10),' ')} rank-ordered set-accuracy", + f"rab_cand_KO_{c['slug']}_{slugify(c['key'])[:24]}") + except Exception as e: + print(f"skip KO {c['slug']}: {e}") + if c["slug"] not in ntc_done: + ntc_done.add(c["slug"]) + try: + montage(c["mc"], c["ch"], "complexes", "NTC", + f"NTC — {c['marker_label'].replace(chr(10),' ')} marker rank-ordered set-accuracy", + f"rab_cand_NTC_{c['slug']}") + except Exception as e: + print(f"skip NTC {c['slug']}: {e}") diff --git a/src/ops_model/models/attention/diffex/figures/rebuild_traversals_n100.py b/src/ops_model/models/attention/diffex/figures/rebuild_traversals_n100.py new file mode 100644 index 0000000..ab2677c --- /dev/null +++ b/src/ops_model/models/attention/diffex/figures/rebuild_traversals_n100.py @@ -0,0 +1,181 @@ +"""Rebuild the 8 fig-4 morpho-traversal examples at n=100 cells (tight SEM). Per marker: back up + clear +the NTC anchor cache, re-gather 100 anchors + regenerate the picked target traversals (force), reusing the +EXACT v5 build config (phase = accuracy_parquet; fluor = fluor_rank_parquet). The NTC anchor cache is +per-modality (shared geneKO+complex), so the geneKO pass builds the 100-cell cache from the geneKO parquet +(thousands of NTC) and the complex pass reuses it (complex parquets only have ~30 NTC). _gather_class is +rank-deterministic → cells 0-39 preserved, other genes' existing traversals stay valid. + +Then remeasure metrics: OPS_DIFFEX_ASSETS=viewer_assets_v5 submit phase-morpho --n-cells 100 +Usage: python rebuild_traversals_n100.py [modality ...] (default = all 6) +""" +import os +import shutil +import sys +from pathlib import Path + +from ops_model.models.attention.diffex.viewer import catalog as C +from ops_model.models.attention.diffex.classifier.config import slugify + +ASSETS = "viewer_assets_v5" +RANK = f"{C.OUT}/{ASSETS}/_rankings/fluor" +V5G = f"{C.OUT}/{ASSETS}/_rankings/pma_v5_phase_geneKO.parquet" +V5C = f"{C.OUT}/{ASSETS}/_rankings/pma_v5_phase_complex.parquet" +PHASE_CK = f"{C.DD}/phase_v1/diffae_best.pt" +NCELLS = 100 + +FLUOR = { # modality -> (diffae_dir, marker_channel, channel) + "nucleus_NucleoLIVE_Live_Cell_dye": ("fluor_NucleoLive", "nucleus_NucleoLIVE Live Cell dye", "mCherry"), + "nucleolus_GC_NPM3": ("fluor_NPM3", "nucleolus-GC_NPM3", "GFP"), + "mitochondria_ChromaLIVE_561_excitation": ("fluor_ChromaLIVE_mito", "mitochondria_ChromaLIVE 561 excitation", "mCherry"), + "ER_Golgi_COP_II_SEC23A": ("fluor_ER_Golgi_COP_II_SEC23A", "ER/Golgi COP-II_SEC23A", "GFP"), + "actin_filament_FastAct_SPY555_Live_Cell_Dye": ("fluor_FastAct", "actin filament_FastAct_SPY555 Live Cell Dye", "mCherry"), + "lysosome_LysoTracker_live_cell_dye": ("fluor_LysoTracker", "lysosome_LysoTracker live-cell dye", "GFP"), + "lipid_droplet_BODIPY_live_cell_dye": ("fluor_lipid_droplet_BODIPY_live_cell_dye", "lipid droplet_BODIPY live cell dye", "GFP"), + "clathrin_vesicles_CLTA": ("fluor_clathrin_vesicles_CLTA", "clathrin vesicles_CLTA", "GFP"), + "stress_granule_G3BP1": ("fluor_stress_granule_G3BP1", "stress granule_G3BP1", "GFP"), + "chromatin_H2BC21": ("fluor_chromatin_H2BC21", "chromatin_H2BC21", "mCherry"), + "nucleolus_DFC_FBL": ("fluor_nucleolus_DFC_FBL", "nucleolus-DFC_FBL", "GFP"), + "autophagosome_MAP1LC3B": ("fluor_autophagosome_MAP1LC3B", "autophagosome_MAP1LC3B", "GFP"), + "lysosome_LAMP1": ("fluor_lysosome_LAMP1", "lysosome_LAMP1", "GFP"), + "proteasome_PSMB7": ("fluor_proteasome_PSMB7", "proteasome_PSMB7", "GFP"), + "F_actin_Phalloidin": ("fluor_F_actin_Phalloidin", "F-actin_Phalloidin", "CP1_f_actin_Phalloidin"), +} +TARGETS = { # modality -> {grain: [class,...]} (the 8 picked fig-4 examples) + "phase": {"geneKO": ["TOMM20", "MICOS13", "SAMM50"]}, + "nucleus_NucleoLIVE_Live_Cell_dye": {"geneKO": ["KIF23"]}, + "nucleolus_GC_NPM3": {"geneKO": ["POLR1B"], "complex": ["Chaperonin-containing T-complex"]}, + "mitochondria_ChromaLIVE_561_excitation": {"geneKO": ["TOMM20"], + "complex": ["TIM23 mitochondrial inner membrane pre-sequence translocase complex, TIM17A variant"]}, + "ER_Golgi_COP_II_SEC23A": {"geneKO": ["GBF1"]}, + "actin_filament_FastAct_SPY555_Live_Cell_Dye": {"geneKO": ["CAPZB"]}, + "lysosome_LysoTracker_live_cell_dye": {"geneKO": ["LAMTOR2"]}, + "lipid_droplet_BODIPY_live_cell_dye": {"geneKO": ["RAB7A"]}, + "clathrin_vesicles_CLTA": {"geneKO": ["AP2M1"]}, + "stress_granule_G3BP1": {"geneKO": ["EIF2S2"]}, + "chromatin_H2BC21": {"geneKO": ["AURKB"]}, + "nucleolus_DFC_FBL": {"geneKO": ["NOP56"]}, + "autophagosome_MAP1LC3B": {"geneKO": ["ATG9A"]}, + "lysosome_LAMP1": {"geneKO": ["ATP6V1B2"]}, + "proteasome_PSMB7": {"geneKO": ["PSMB6"]}, + "F_actin_Phalloidin": {"geneKO": ["CAPZB"]}, +} + + +def _clear_anchor(modality): + """Back up the old (40-45 cell) ctrl.npz so the next gather rebuilds the cache fresh at NCELLS. + Idempotent: if the cache is already >= NCELLS, keep it (so re-runs reuse the same 100 anchors).""" + import numpy as np + ad = Path(C.OUT) / ASSETS / modality / "_anchors" / "NTC" + ck = ad / "ctrl.npz" + if ck.exists(): + z = np.load(ck) + have = z["anchor_imgs"].shape[0] if "anchor_imgs" in z.files else 0 + if have >= NCELLS: + print(f"[keep] {modality}: anchor cache already {have} >= {NCELLS} — reuse"); return + bak = ad / "ctrl.npz.n40bak" + if bak.exists(): + bak.unlink() + ck.rename(bak) + print(f"[clear] {modality}: ctrl.npz ({have}) -> ctrl.npz.n40bak (rebuild @ {NCELLS})") + + +def rebuild_marker(modality): + os.environ["OPS_DIFFEX_ASSETS"] = ASSETS + from ops_model.models.attention.diffex.viewer import precompute as P + P._ASSETS = ASSETS + _clear_anchor(modality) + tg = TARGETS[modality] + if modality == "phase": + P.precompute_marker(grain="geneKO", targets=tg["geneKO"], ckpt=PHASE_CK, out_root=C.OUT, + control="NTC", accuracy_parquet=V5G, v5_score=True, n_cells=NCELLS, force=True) + return + d, mc, ch = FLUOR[modality] + if "geneKO" in tg: # geneKO pass FIRST → builds the 100-cell NTC cache + P.precompute_marker(grain="geneKO", targets=tg["geneKO"], ckpt=f"{C.DD}/{d}/diffae_best.pt", out_root=C.OUT, + marker_channel=mc, channel=ch, control="NTC", n_cells=NCELLS, + fluor_rank_parquet=f"{RANK}/geneKO/{slugify(mc)}.parquet", v5_score=True, + load_workers=12, force=True) + if "complex" in tg: # reuses the cache built above + P.precompute_marker(grain="complex", targets=tg["complex"], ckpt=f"{C.DD}/{d}/diffae_best.pt", out_root=C.OUT, + marker_channel=mc, channel=ch, control="NTC", n_cells=NCELLS, + fluor_rank_parquet=f"{RANK}/complex/{slugify(mc)}.parquet", v5_score=True, + load_workers=12, force=True) + + +def gen_phase(targets): + """Generate specific phase geneKO targets at n=100 (reuses the existing 100-cell phase anchor cache).""" + os.environ["OPS_DIFFEX_ASSETS"] = ASSETS + from ops_model.models.attention.diffex.viewer import precompute as P + P._ASSETS = ASSETS + _clear_anchor("phase") # idempotent: keeps the 100-cell cache + P.precompute_marker(grain="geneKO", targets=targets, ckpt=PHASE_CK, out_root=C.OUT, + control="NTC", accuracy_parquet=V5G, v5_score=True, n_cells=NCELLS, force=True) + + +def gen_phase_complex(targets): + """Generate phase COMPLEX targets at n=100 (full complex names; reuses the 100-cell phase anchor cache).""" + os.environ["OPS_DIFFEX_ASSETS"] = ASSETS + from ops_model.models.attention.diffex.viewer import precompute as P + P._ASSETS = ASSETS + _clear_anchor("phase") + P.precompute_marker(grain="complex", targets=targets, ckpt=PHASE_CK, out_root=C.OUT, + control="NTC", accuracy_parquet=V5C, v5_score=True, n_cells=NCELLS, force=True) + + +def submit_phase_targets(targets, tag, func=None): + """One SLURM job PER target (parallel across GPUs) — the phase-100 anchor cache already exists and + _clear_anchor is idempotent (>=100 → keep), so per-target jobs reuse it read-only (no re-gather race).""" + from ops_utils.hpc.slurm_batch_utils import submit_parallel_jobs + fn = func or gen_phase + jobs = [{"name": f"trav100_{tag}_{slugify(str(t))[:14]}", "func": fn, "kwargs": {"targets": [t]}} for t in targets] + submit_parallel_jobs(jobs, experiment="diffex_trav100", + slurm_params={"slurm_partition": "gpu", "slurm_gres": "gpu:1", "cpus_per_task": 12, + "mem_gb": 96, "timeout_min": 300}, + log_dir="diffex_trav100", wait_for_completion=False) + + +def gen_phase_chunk(targets, cell_range): + """Generate a CELL-RANGE slice of phase geneKO targets. Parallel chunks share the cached 100-anchor + the + cached direction (both ckpt-independent), each writing its own cell{c} dirs → parallelizes the per-cell + DDIM inversion across GPUs instead of one long serial job. v5 scoring skipped (whole-target only).""" + os.environ["OPS_DIFFEX_ASSETS"] = ASSETS + from ops_model.models.attention.diffex.viewer import precompute as P + P._ASSETS = ASSETS + _clear_anchor("phase") # idempotent: 100-anchor cache kept + P.precompute_marker(grain="geneKO", targets=targets, ckpt=PHASE_CK, out_root=C.OUT, + control="NTC", accuracy_parquet=V5G, v5_score=False, n_cells=NCELLS, force=True, + cell_range=cell_range) + + +def submit_phase_chunked(targets, tag, n_chunks=5): + """One SLURM job per (target, cell-chunk) — parallelize NCELLS across GPUs (not longer walltime).""" + from ops_utils.hpc.slurm_batch_utils import submit_parallel_jobs + step = (NCELLS + n_chunks - 1) // n_chunks + jobs = [] + for t in targets: + for i in range(n_chunks): + s, e = i * step, min((i + 1) * step, NCELLS) + if s >= e: + continue + jobs.append({"name": f"tv_{tag}_{slugify(str(t))[:8]}_{s}", "func": gen_phase_chunk, + "kwargs": {"targets": [t], "cell_range": (s, e)}}) + print(f"[chunked] {len(jobs)} jobs ({len(targets)} targets x {n_chunks} chunks of ~{step} cells)") + submit_parallel_jobs(jobs, experiment="diffex_trav100", + slurm_params={"slurm_partition": "gpu", "slurm_gres": "gpu:1", "cpus_per_task": 12, + "mem_gb": 96, "timeout_min": 120}, + log_dir="diffex_trav100", wait_for_completion=False) + + +def main(): + from ops_utils.hpc.slurm_batch_utils import submit_parallel_jobs + mods = sys.argv[1:] or list(TARGETS) + jobs = [{"name": f"trav100_{slugify(m)[:16]}", "func": rebuild_marker, "kwargs": {"modality": m}} for m in mods] + print(f"[trav-n100] {len(jobs)} marker job(s) @ n_cells={NCELLS}: {mods}") + submit_parallel_jobs(jobs, experiment="diffex_trav100", + slurm_params={"slurm_partition": "gpu", "slurm_gres": "gpu:1", "cpus_per_task": 12, + "mem_gb": 96, "timeout_min": 720}, + log_dir="diffex_trav100", wait_for_completion=False) + + +if __name__ == "__main__": + main() diff --git a/src/ops_model/models/attention/diffex/figures/traversal_montage_schematic.py b/src/ops_model/models/attention/diffex/figures/traversal_montage_schematic.py new file mode 100644 index 0000000..99fcef3 --- /dev/null +++ b/src/ops_model/models/attention/diffex/figures/traversal_montage_schematic.py @@ -0,0 +1,142 @@ +"""Figure 4 schematic: DiffAE counterfactual phenotype traversal (single row, 3 steps). + + 1. Train DiffAE — semantic encoder -> semantic latent z (Cell-DINO); forward noise -> x_T; + conditional U-Net denoiser reverses it (DDIM) back to the cell, conditioned on z. + 2. Pick the traversal destination — supervised NTC -> gene-KO direction from the top set-accuracy + cells (kept cells solid; excluded cells faded). + 3. Traverse — decode the control cell along that direction (alpha 0 -> +5). + +The z heat-strip shows the *values* of the semantic embedding vector (not categories). + + python traversal_montage_schematic.py +""" +import matplotlib.pyplot as plt +import numpy as np +from matplotlib.patches import FancyArrowPatch, Ellipse, FancyBboxPatch, Rectangle, Circle + +plt.rcParams["pdf.fonttype"] = 42 +plt.rcParams["svg.fonttype"] = "none" +plt.rcParams["font.family"] = "sans-serif" +plt.rcParams["font.sans-serif"] = ["Arial", "Helvetica", "DejaVu Sans"] + +INK = "#1a1a1a"; GREY = "#7a7a7a" +TEAL = "#0b6b73"; DORANGE = "#b35900"; PINK = "#d94f9a" +CELL_FILL = "#d7d7d7"; NUC = "#4a4a4a" + + +def _arrow(ax, x0, y0, x1, y1, lw=1.6, color=INK, mut=11): + ax.add_patch(FancyArrowPatch((x0, y0), (x1, y1), arrowstyle="-|>", mutation_scale=mut, + lw=lw, color=color, shrinkA=0, shrinkB=0, zorder=7)) + + +def _cell(ax, cx, cy, s, morph=0.0, edge=INK, lw=1.5): + """Nucleus size (small→large) and speckle roundness (streaky→round) illustrate the phenotype shift.""" + ax.add_patch(Rectangle((cx - s, cy - s), 2 * s, 2 * s, facecolor=CELL_FILL, edgecolor=edge, lw=lw, zorder=4)) + nr = s * (0.30 + 0.36 * morph) + ax.add_patch(Circle((cx, cy), nr, facecolor=NUC, edgecolor="none", zorder=5)) + rnd = 0.28 + 0.72 * morph; dw = s * 0.15 + for k in range(3): + ang = 2.2 * k + 0.7; rr = nr * 0.5 + ax.add_patch(Ellipse((cx + rr * np.cos(ang), cy + rr * np.sin(ang)), 2 * dw, 2 * dw * rnd, + angle=35, facecolor="white", edgecolor="none", zorder=6)) + + +def _noise_tile(ax, cx, cy, s, level=1.0, edge=INK, lw=1.5): + ax.add_patch(Rectangle((cx - s, cy - s), 2 * s, 2 * s, facecolor=CELL_FILL, edgecolor="none", zorder=3)) + if level < 0.99: + ax.add_patch(Circle((cx, cy), s * 0.5, facecolor=NUC, edgecolor="none", alpha=1 - level, zorder=4)) + ax.imshow(np.random.rand(14, 14), extent=[cx - s, cx + s, cy - s, cy + s], cmap="gray", vmin=0, vmax=1, + alpha=level, zorder=5, aspect="auto", interpolation="nearest") + ax.add_patch(Rectangle((cx - s, cy - s), 2 * s, 2 * s, fill=False, edgecolor=edge, lw=lw, zorder=6)) + + +def _vec(ax, x, y, w=1.0, h=0.15, n=9): + """Semantic embedding vector: cells shaded by value (grey shades = different numbers, not categories).""" + for i, v in enumerate(np.random.rand(n)): + ax.add_patch(Rectangle((x + i * w / n, y), w / n * 0.9, h, facecolor=str(0.22 + 0.6 * v), + edgecolor="white", lw=0.5, zorder=7)) + + +def _cluster(ax, cx, cy, color, n_sel=6, n_un=8): + """Kept (top set-accuracy) cells solid; excluded cells faded/hollow. Centroid = kept cells.""" + un = np.random.normal([cx, cy], 0.34, (n_un, 2)) + ax.scatter(un[:, 0], un[:, 1], s=20, facecolors="none", edgecolors=color, linewidths=1.1, alpha=0.4, zorder=3) + sel = np.random.normal([cx, cy], 0.19, (n_sel, 2)) + ax.scatter(sel[:, 0], sel[:, 1], s=24, c=color, edgecolors="white", linewidths=0.4, zorder=4) + ax.add_patch(Circle((cx, cy), 0.09, facecolor=color, edgecolor=INK, lw=1.3, zorder=6)) + + +def build(outstem="traversal_montage_schematic"): + fig, ax = plt.subplots(figsize=(15.2, 4.4)) + ax.set_xlim(0, 15.2); ax.set_ylim(0, 4.4); ax.axis("off") + ax.text(0.12, 4.28, "B", fontsize=28, fontweight="bold", va="top") + ax.text(0.68, 4.26, "Counterfactual phenotype traversal", fontsize=20, va="top") + np.random.seed(1) + hy = 3.55; yc = 2.25 + + # ===== 1. Train DiffAE ===== + ax.text(0.3, hy, "1 Train DiffAE", fontsize=14, fontweight="bold", va="center") + _cell(ax, 0.7, yc, 0.26, morph=1.0) + ax.text(0.7, yc - 0.4, "real cell x0", fontsize=8.5, color=GREY, ha="center", va="top") + _arrow(ax, 0.85, yc + 0.24, 1.22, yc + 0.5, lw=1.2, mut=8) + ax.text(1.28, yc + 0.53, "z =", fontsize=11, ha="left", va="center", fontweight="bold", color="#a3266f", zorder=8) + _vec(ax, 1.68, yc + 0.46, w=0.85) + ax.text(2.1, yc + 0.9, "semantic latent (Cell-DINO)", fontsize=8.5, ha="center", color="#a3266f") + _arrow(ax, 0.98, yc, 1.45, yc, lw=1.2, mut=8) + ax.text(1.2, yc + 0.15, "noise", fontsize=7.5, color=GREY, ha="center") + _noise_tile(ax, 1.72, yc, 0.26, level=1.0) + ax.text(1.72, yc - 0.4, "x_T", fontsize=8.5, color=GREY, ha="center", va="top") + _arrow(ax, 1.98, yc, 2.36, yc, lw=1.2, mut=8) + _noise_tile(ax, 2.66, yc, 0.26, level=0.5) + _arrow(ax, 2.92, yc, 3.3, yc, lw=1.2, mut=8) + _cell(ax, 3.6, yc, 0.26, morph=1.0) # = α=0 control reconstruction in step 3 + ax.text(3.6, yc - 0.4, "generated cell", fontsize=8.5, color=GREY, ha="center", va="top") + ax.add_patch(FancyArrowPatch((2.35, yc + 0.46), (2.64, yc + 0.32), arrowstyle="-|>", mutation_scale=8, + lw=1.0, color=GREY, ls="--", shrinkA=0, shrinkB=2, zorder=6)) + ax.text(0.3, 1.28, "conditional U-Net denoises x_T → cell\n(reverse diffusion, DDIM), conditioned on z;\nfixed noise x_T carries cell identity", + fontsize=8.5, color=GREY, ha="left", va="top") + + _arrow(ax, 4.15, yc, 4.75, yc, lw=1.8, mut=13) + + # ===== 2. Pick destination ===== + ax.text(5.0, hy, "2 Pick destination", fontsize=14, fontweight="bold", va="center") + _cluster(ax, 5.7, yc - 0.05, GREY) + _cluster(ax, 7.15, yc + 0.2, PINK) + _arrow(ax, 5.88, yc, 6.95, yc + 0.17, lw=1.9) + ax.text(5.7, yc - 0.52, "NTC", fontsize=9, ha="center", va="top", color=GREY) + ax.text(7.62, yc + 0.2, "gene-KO", fontsize=9, ha="left", va="center", color="#a3266f") + # legend: kept vs excluded + ax.scatter([5.35], [1.5], s=24, c=PINK, edgecolors="white", linewidths=0.4, zorder=5) + ax.text(5.5, 1.5, "top set-accuracy cells (kept → centroid)", fontsize=8, ha="left", va="center") + ax.scatter([5.35], [1.22], s=20, facecolors="none", edgecolors=PINK, linewidths=1.1, alpha=0.6, zorder=5) + ax.text(5.5, 1.22, "other cells (excluded)", fontsize=8, ha="left", va="center", color=GREY) + ax.text(5.35, 0.9, "→ supervised NTC → gene-KO direction", fontsize=8.5, ha="left", color=INK) + + _arrow(ax, 8.5, yc, 9.1, yc, lw=1.8, mut=13) + + # ===== 3. Traverse ===== + ax.text(9.35, hy, "3 Traverse", fontsize=14, fontweight="bold", va="center") + axs = np.linspace(9.9, 14.2, 6) + for xx, a in zip(axs, [0, 1, 2, 3, 4, 5]): + edge = PINK if a == 1 else DORANGE if a == 5 else INK + _cell(ax, xx, yc, 0.34, morph=1 - a / 5.0, edge=edge, lw=2.6 if a in (1, 5) else 1.5) + ax.text(xx, yc - 0.6, f"α = {a}" if a == 0 else f"+{a}", fontsize=10, ha="center") + ax.text(axs[1], yc + 0.52, "gene-KO\ncentroid", fontsize=8.5, ha="center", va="bottom", + color="#a3266f", fontweight="bold") + ax.text(axs[5], yc + 0.52, "exaggerated\nphenotype", fontsize=8.5, ha="center", va="bottom", + color=DORANGE, fontweight="bold") + ax.annotate("", (axs[-1] + 0.42, 1.42), (axs[0] - 0.42, 1.42), + arrowprops=dict(arrowstyle="-|>", color=INK, lw=1.7)) + ax.text(axs[0] - 0.42, 1.16, "NTC", fontsize=11, ha="left", color=GREY, fontweight="bold") + ax.text(axs[-1] + 0.42, 1.16, "+KO", fontsize=11, ha="right", color=DORANGE, fontweight="bold") + ax.text(12.05, 0.72, "α interpolates the control cell toward the knockdown state; |α| > 1 exaggerates", + fontsize=9.5, ha="center", color=GREY) + + for ext in ("png", "svg"): + fig.savefig(f"{outstem}.{ext}", dpi=200, bbox_inches="tight", facecolor="white") + plt.close(fig) + print(f"[schematic] -> {outstem}.png / .svg") + + +if __name__ == "__main__": + build() diff --git a/src/ops_model/models/attention/diffex/figures/virtual_staining_schematic.py b/src/ops_model/models/attention/diffex/figures/virtual_staining_schematic.py new file mode 100644 index 0000000..ea3c1fe --- /dev/null +++ b/src/ops_model/models/attention/diffex/figures/virtual_staining_schematic.py @@ -0,0 +1,202 @@ +"""Figure schematic: spatially-conditioned DiffAE for multi-channel virtual staining (phase → marker). + +Variant of traversal_montage_schematic.py (same visual language) for the cross-channel task: + + 1. Spatially-conditioned DiffAE — a label-free Phase2D crop conditions the generator TWO ways: + - semantic z: phase → frozen Cell-DINO (ViT, patch-tokenise → pooled 1024-d) → global FiLM + conditioning (+ learned marker id) — content, but NO spatial layout; + - dense pixels: the raw phase image concatenated as an extra U-Net input channel (in_channels 1→2) + — keeps the layout, so the output stays pixel-registered to the input (the lift 0.13 → 0.78). + The conditional U-Net denoises x_T → the chosen fluorescent marker. + 2. Render any marker — from the SAME phase cell, switching the marker id renders every trained marker, + each co-registered to the input; compared to the real marker (held-out Pearson). + +One model over 42 live markers; trained on paired (phase, marker) crops from phenotyping_v3.zarr (both +channels in one stitched frame → registered). 4i/CP markers excluded (fixed later → misregistered). + + python virtual_staining_schematic.py +""" +import matplotlib.pyplot as plt +import numpy as np +from matplotlib.patches import FancyArrowPatch, FancyBboxPatch, Rectangle, Circle, Ellipse + +plt.rcParams["pdf.fonttype"] = 42 +plt.rcParams["svg.fonttype"] = "none" +plt.rcParams["font.family"] = "sans-serif" +plt.rcParams["font.sans-serif"] = ["Arial", "Helvetica", "DejaVu Sans"] + +INK = "#1a1a1a"; GREY = "#7a7a7a"; PURPLE = "#a3266f"; TEAL = "#0b6b73" +PHASE_FILL = "#d7d7d7"; NUC = "#4a4a4a" +# marker channel colours (distinct organelles rendered from the same phase) +MARKERS = [("mito", "#e8791a"), ("ER", "#2ca089"), ("nucleus", "#3a6fd8"), ("lysosome", "#d94f4f")] + + +def _arrow(ax, x0, y0, x1, y1, lw=1.6, color=INK, mut=11, ls="-"): + ax.add_patch(FancyArrowPatch((x0, y0), (x1, y1), arrowstyle="-|>", mutation_scale=mut, + lw=lw, color=color, ls=ls, shrinkA=0, shrinkB=0, zorder=7)) + + +def _phase_cell(ax, cx, cy, s, edge=INK, lw=1.5): + """Label-free phase cell: grey square, dark nucleus, faint internal texture (no fluor signal).""" + ax.add_patch(Rectangle((cx - s, cy - s), 2 * s, 2 * s, facecolor=PHASE_FILL, edgecolor=edge, lw=lw, zorder=4)) + ax.add_patch(Circle((cx, cy), s * 0.42, facecolor=NUC, edgecolor="none", alpha=0.85, zorder=5)) + for k in range(4): + ang = 1.7 * k + 0.4; rr = s * 0.55 + ax.add_patch(Ellipse((cx + rr * np.cos(ang), cy + rr * np.sin(ang)), s * 0.34, s * 0.16, + angle=40 * k, facecolor="#c2c2c2", edgecolor="none", zorder=5)) + + +def _marker_tile(ax, cx, cy, s, kind, color, edge=INK, lw=1.5): + """Stylised fluorescent marker on black: a characteristic organelle pattern in the channel colour.""" + ax.add_patch(Rectangle((cx - s, cy - s), 2 * s, 2 * s, facecolor="black", edgecolor=edge, lw=lw, zorder=4)) + rng = np.random.default_rng(abs(hash(kind)) % 2**32) + if kind == "mito": # tubular streaks + for _ in range(9): + a = rng.uniform(0, np.pi); p = rng.uniform(-0.6, 0.6, 2) * s + ax.add_patch(Ellipse((cx + p[0], cy + p[1]), s * 0.5, s * 0.13, angle=np.degrees(a), + facecolor=color, edgecolor="none", alpha=0.9, zorder=5)) + elif kind == "ER": # reticular mesh + for _ in range(11): + p = rng.uniform(-0.7, 0.7, 2) * s + ax.add_patch(Circle((cx + p[0], cy + p[1]), s * 0.18, facecolor="none", + edgecolor=color, lw=1.0, alpha=0.85, zorder=5)) + elif kind == "nucleus": # filled nuclear blob + ax.add_patch(Ellipse((cx, cy), s * 1.1, s * 0.9, facecolor=color, edgecolor="none", alpha=0.85, zorder=5)) + else: # puncta (lysosome) + for _ in range(13): + p = rng.uniform(-0.75, 0.75, 2) * s + ax.add_patch(Circle((cx + p[0], cy + p[1]), s * 0.09, facecolor=color, edgecolor="none", zorder=5)) + + +def _noise_tile(ax, cx, cy, s, level=1.0, edge=INK, lw=1.5): + ax.add_patch(Rectangle((cx - s, cy - s), 2 * s, 2 * s, facecolor="black", edgecolor="none", zorder=3)) + ax.imshow(np.random.rand(14, 14), extent=[cx - s, cx + s, cy - s, cy + s], cmap="gray", vmin=0, vmax=1, + alpha=level, zorder=5, aspect="auto", interpolation="nearest") + ax.add_patch(Rectangle((cx - s, cy - s), 2 * s, 2 * s, fill=False, edgecolor=edge, lw=lw, zorder=6)) + + +def _vec(ax, x, y, w=1.0, h=0.15, n=9, seed=0): + """Semantic embedding vector: cells shaded by value.""" + rng = np.random.default_rng(seed) + for i, v in enumerate(rng.random(n)): + ax.add_patch(Rectangle((x + i * w / n, y), w / n * 0.9, h, facecolor=str(0.22 + 0.6 * v), + edgecolor="white", lw=0.5, zorder=7)) + + +def _marker_selector(ax, x, y, w, sel_i): + """Row of marker-id tokens; the selected one highlighted (which channel to render).""" + n = len(MARKERS); cw = w / n + for i, (name, col) in enumerate(MARKERS): + on = i == sel_i + ax.add_patch(FancyBboxPatch((x + i * cw, y), cw * 0.82, 0.22, boxstyle="round,pad=0.01,rounding_size=0.03", + facecolor=col if on else "white", edgecolor=col, lw=1.6 if on else 1.0, zorder=7)) + ax.text(x + i * cw + cw * 0.41, y + 0.11, name, fontsize=6.8, ha="center", va="center", + color="white" if on else col, fontweight="bold" if on else "normal", zorder=8) + + +def _patch_cell(ax, cx, cy, s, edge=INK): + """Phase cell overlaid with a faint patch grid — the ViT tokenisation used only for the semantic branch.""" + _phase_cell(ax, cx, cy, s, edge=edge, lw=1.3) + for k in range(1, 4): + g = -s + 2 * s * k / 4 + ax.plot([cx - s, cx + s], [cy + g, cy + g], color=PURPLE, lw=0.5, alpha=0.55, zorder=7) + ax.plot([cx + g, cx + g], [cy - s, cy + s], color=PURPLE, lw=0.5, alpha=0.55, zorder=7) + + +def build(outstem="/hpc/projects/icd.fast.ops/analysis/figure4_schematic/virtual_staining_schematic"): + fig, ax = plt.subplots(figsize=(14.6, 4.9)) + ax.set_xlim(0, 14.6); ax.set_ylim(0, 4.9); ax.axis("off") + ax.text(0.12, 4.72, "C", fontsize=28, fontweight="bold", va="top") + ax.text(0.68, 4.70, "Multi-channel virtual staining (label-free phase → fluorescent marker)", + fontsize=19, va="top") + np.random.seed(3) + hy = 3.95; yc = 2.15 + + # ===== 1. Spatially-conditioned DiffAE (encode + generate) ===== + ax.text(0.3, hy, "1 Spatially-conditioned DiffAE", fontsize=14, fontweight="bold", va="center") + + # phase cell (shared input to BOTH conditioning paths) + _phase_cell(ax, 0.8, yc, 0.30) + ax.text(0.8, yc - 0.46, "phase cell", fontsize=8.5, color=GREY, ha="center", va="top") + + # --- semantic branch (up): phase -> Cell-DINO ViT -> pooled z (global, no layout) --- + _arrow(ax, 1.0, yc + 0.28, 1.45, yc + 0.72, lw=1.1, mut=8, color=PURPLE) + _patch_cell(ax, 1.75, yc + 0.95, 0.19) + _arrow(ax, 1.98, yc + 0.95, 2.35, yc + 0.95, lw=1.0, mut=7, color=PURPLE) + ax.add_patch(FancyBboxPatch((2.4, yc + 0.72), 0.95, 0.46, boxstyle="round,pad=0.02,rounding_size=0.05", + facecolor="white", edgecolor=PURPLE, lw=1.3, zorder=6)) + ax.text(2.88, yc + 0.95, "Cell-DINO\n(ViT, frozen)", fontsize=7.6, ha="center", va="center", color=PURPLE, zorder=7) + _arrow(ax, 3.35, yc + 0.95, 3.65, yc + 0.95, lw=1.0, mut=7, color=PURPLE) + _vec(ax, 3.7, yc + 0.87, w=0.72, seed=2) + ax.text(4.06, yc + 1.32, "z: pooled 1024-d (semantic, no spatial layout)", fontsize=7.8, ha="center", color=PURPLE) + ax.text(1.68, yc + 0.62, "patch-tokenise → pool", fontsize=6.8, ha="center", color=PURPLE, style="italic") + + # --- spatial branch (down): phase raw pixels concatenated to the noisy input = SPATIAL CONDITIONING --- + _arrow(ax, 1.02, yc - 0.26, 1.42, yc - 0.5, lw=1.1, mut=8, color=TEAL) + ax.text(1.28, yc - 0.62, "raw pixels", fontsize=7.0, color=TEAL, ha="center", va="top", style="italic") + # boxed, labelled callout so it's unmistakable which element is the spatial conditioning: + # the U-Net INPUT = [ noisy x_T ⊕ phase image ] (2 channels) + ax.add_patch(FancyBboxPatch((1.78, yc - 1.02), 1.02, 1.16, boxstyle="round,pad=0.02,rounding_size=0.05", + facecolor="#e7f2f1", edgecolor=TEAL, lw=1.9, zorder=3)) + ax.text(2.29, yc + 0.30, "SPATIAL CONDITIONING", fontsize=7.8, color=TEAL, ha="center", va="center", + fontweight="bold", zorder=8) + _noise_tile(ax, 2.14, yc - 0.32, 0.185); ax.text(2.44, yc - 0.32, "x_T", fontsize=6.8, color=GREY, ha="left", va="center", zorder=8) + # circled-plus (concat) between the two channel tiles + ax.add_patch(Circle((2.14, yc - 0.58), 0.058, facecolor="white", edgecolor=TEAL, lw=1.1, zorder=8)) + ax.plot([2.104, 2.176], [yc - 0.58, yc - 0.58], color=TEAL, lw=1.0, zorder=9) + ax.plot([2.14, 2.14], [yc - 0.616, yc - 0.544], color=TEAL, lw=1.0, zorder=9) + _phase_cell(ax, 2.14, yc - 0.84, 0.185, edge=TEAL, lw=1.3); ax.text(2.44, yc - 0.84, "phase", fontsize=6.8, color=TEAL, ha="left", va="center", zorder=8) + ax.text(2.29, yc + 0.05, "2-channel U-Net input", fontsize=6.6, color=TEAL, ha="center", va="center", zorder=8) + _arrow(ax, 2.82, yc - 0.5, 3.2, yc - 0.42, lw=1.3, mut=9, color=TEAL) + + # U-Net + ax.add_patch(FancyBboxPatch((3.25, yc - 0.68), 1.2, 1.02, boxstyle="round,pad=0.02,rounding_size=0.06", + facecolor="#f2f2f2", edgecolor=INK, lw=1.7, zorder=5)) + ax.text(3.85, yc - 0.02, "conditional\nU-Net", fontsize=9, ha="center", va="center", fontweight="bold", zorder=6) + ax.text(3.85, yc - 0.46, "DDIM", fontsize=7.5, ha="center", va="center", color=GREY, zorder=6) + + # z (+ marker id) -> U-Net as FiLM class-embedding (from top) + _arrow(ax, 4.05, yc + 0.82, 3.95, yc + 0.36, lw=1.1, mut=8, color=PURPLE, ls="--") + ax.text(4.2, yc + 0.6, "z + marker id\n(FiLM)", fontsize=7.4, color=PURPLE, ha="left", va="center", fontweight="bold") + + # denoise -> predicted marker + _arrow(ax, 4.5, yc - 0.02, 4.95, yc - 0.02, lw=1.3, mut=10) + _marker_tile(ax, 5.3, yc, 0.30, "mito", MARKERS[0][1]) + ax.text(5.3, yc - 0.48, "predicted marker", fontsize=8.5, color=GREY, ha="center", va="top") + + # marker-id selector + ax.text(0.35, 0.55, "marker id:", fontsize=8.5, ha="left", va="center", fontweight="bold") + _marker_selector(ax, 1.2, 0.43, 2.4, sel_i=0) + # the two conditioning roles, spelled out + ax.text(6.0, yc + 0.95, "semantic z → global FiLM conditioning\n(what to render — content, no layout)", + fontsize=8.1, color=PURPLE, ha="left", va="center") + ax.text(6.0, yc - 0.55, "phase pixels → dense concat into U-Net input\n(keeps layout → output pixel-registered to input)", + fontsize=8.1, color=TEAL, ha="left", va="center", fontweight="bold") + + _arrow(ax, 8.7, yc, 9.35, yc, lw=1.8, mut=13) + + # ===== 2. Render any marker ===== + ax.text(9.55, hy, "2 Render any marker", fontsize=14, fontweight="bold", va="center") + ax.text(9.55, hy - 0.36, "same phase cell → switch marker id", fontsize=8.3, color=GREY, va="center") + xs = np.linspace(10.1, 14.2, 4) + rs = [0.78, 0.74, 0.81, 0.69] + for i, (xx, (name, col)) in enumerate(zip(xs, MARKERS)): + _marker_tile(ax, xx, yc + 0.15, 0.33, name, col, edge=col, lw=2.2) + ax.text(xx, yc + 0.62, name, fontsize=9, ha="center", color=col, fontweight="bold") + ax.text(xx, yc - 0.32, f"r = {rs[i]:.2f}", fontsize=8.2, ha="center", color=INK) + ax.text(12.15, yc - 0.62, "held-out Pearson(pred, real)", fontsize=8, ha="center", color=GREY) + ax.annotate("", (xs[-1] + 0.45, 1.15), (xs[0] - 0.45, 1.15), + arrowprops=dict(arrowstyle="-|>", color=INK, lw=1.5)) + ax.text(12.15, 0.85, "one model over 42 live markers (dyes + FP tags), all from a single phase image", + fontsize=8.6, ha="center", color=INK) + + import os + os.makedirs(os.path.dirname(outstem), exist_ok=True) + for ext in ("png", "svg"): + fig.savefig(f"{outstem}.{ext}", dpi=200, bbox_inches="tight", facecolor="white") + plt.close(fig) + print(f"[schematic] -> {outstem}.png / .svg") + + +if __name__ == "__main__": + build() diff --git a/src/ops_model/models/attention/diffex/kyle_pcs/build_static_explorer.py b/src/ops_model/models/attention/diffex/kyle_pcs/build_static_explorer.py new file mode 100644 index 0000000..9423ae9 --- /dev/null +++ b/src/ops_model/models/attention/diffex/kyle_pcs/build_static_explorer.py @@ -0,0 +1,436 @@ +#!/usr/bin/env python3 +"""Build a static, self-contained PC Strip Explorer HTML file. + +Embeds all crop images as base64 and all data as inline JSON, +producing a single HTML file that works from file:// or any static host. + +Required inputs (all in --artifacts-dir): + gene_names.json — list of gene names + gene_pc_scores.npy — (n_genes, n_pcs) mean PC scores per gene + gene_pc_analysis.json — must contain "explained_variance" list + representatives.json — per-PC representative cells with metadata + crops_png/ — cell crop PNGs named pc{NNN}_bin{NN}_row{N}.png + +Usage: + python3 build_static_explorer.py + python3 build_static_explorer.py --artifacts-dir /path/to/data --output explorer.html +""" + +import argparse +import base64 +import json +from pathlib import Path + +import numpy as np + + +def build(artifacts_dir, output_path): + artifacts_dir = Path(artifacts_dir) + crops_dir = artifacts_dir / "crops_png" + + with open(artifacts_dir / "gene_names.json") as f: + gene_names = json.load(f) + scores = np.load(artifacts_dir / "gene_pc_scores.npy") + with open(artifacts_dir / "gene_pc_analysis.json") as f: + analysis = json.load(f) + with open(artifacts_dir / "representatives.json") as f: + reps_data = json.load(f) + + ev = analysis.get("explained_variance", []) + + overview = { + "n_genes": len(gene_names), + "n_pcs": scores.shape[1], + "n_bins": reps_data.get("cells_per_row", 15), + "n_rows": reps_data.get("n_rows", 3), + "crop_size": reps_data.get("crop_size", 96), + "total_variance": round(sum(ev) * 100, 1), + "explained_variance": [round(v * 100, 2) for v in ev], + } + + reps_by_pc = {} + for r in reps_data["representatives"]: + reps_by_pc.setdefault(r["pc"], []).append(r) + + pc_data = {} + for pc_num in range(1, scores.shape[1] + 1): + idx = pc_num - 1 + col = scores[:, idx] + top_high = np.argsort(col)[::-1][:15].tolist() + top_low = np.argsort(col)[:15].tolist() + + entries = reps_by_pc.get(idx, []) + bins = {} + for r in entries: + bins.setdefault(r["bin"], []).append(r) + + strip = [] + for bi in sorted(bins.keys()): + rows = sorted(bins[bi], key=lambda x: x["row"]) + strip.append({ + "bin": bi, + "cells": [{ + "gene": r["gene"], + "score": round(r["score"], 2), + "experiment": r["experiment"], + "well": r["well"], + "x": round(r["x"], 1), + "y": round(r["y"], 1), + "has_crop": r["has_crop"], + "img": f"pc{idx:03d}_bin{bi:02d}_row{r['row']}.png", + } for r in rows] + }) + + pc_data[pc_num] = { + "pc": pc_num, + "explained_variance": round(ev[idx] * 100, 2) if idx < len(ev) else None, + "high_genes": [{"gene": gene_names[i], "score": round(float(col[i]), 3)} for i in top_high], + "low_genes": [{"gene": gene_names[i], "score": round(float(col[i]), 3)} for i in top_low], + "strip": strip, + } + + gene_data = {} + for gi, name in enumerate(gene_names): + profile = scores[gi].tolist() + top_pcs = np.argsort(np.abs(scores[gi]))[::-1][:15].tolist() + gene_data[name] = { + "gene": name, + "top_pcs": [{"pc": int(pc + 1), "score": round(float(profile[pc]), 3)} for pc in top_pcs], + } + + print("Encoding crop images...") + crop_b64 = {} + for png in sorted(crops_dir.glob("*.png")): + crop_b64[png.name] = base64.b64encode(png.read_bytes()).decode("ascii") + print(f" {len(crop_b64)} images encoded") + + print("Building HTML...") + html = build_html(overview, pc_data, gene_data, gene_names, crop_b64) + + output_path = Path(output_path) + output_path.write_text(html) + size_mb = output_path.stat().st_size / 1e6 + print(f"Written {output_path} ({size_mb:.1f} MB)") + + +def build_html(overview, pc_data, gene_data, gene_names, crop_b64): + return r""" + +PC Strip Explorer + + +
+
+

PC Strip Explorer

+
+ +
+ + +
+
+ +
+
+ +
+ + +
+
+ + + +""" + + +if __name__ == "__main__": + parser = argparse.ArgumentParser(description=__doc__, + formatter_class=argparse.RawDescriptionHelpFormatter) + default_artifacts = Path(__file__).resolve().parent.parent / "artifacts" + parser.add_argument("--artifacts-dir", default=str(default_artifacts), + help=f"Directory with input data (default: {default_artifacts})") + parser.add_argument("--output", default=None, + help="Output HTML path (default: /pc_explorer_static.html)") + args = parser.parse_args() + + output = args.output or str(Path(args.artifacts_dir) / "pc_explorer_static.html") + build(args.artifacts_dir, output) diff --git a/src/ops_model/models/attention/diffex/kyle_pcs/compute_pc_strips.py b/src/ops_model/models/attention/diffex/kyle_pcs/compute_pc_strips.py new file mode 100644 index 0000000..5e9df19 --- /dev/null +++ b/src/ops_model/models/attention/diffex/kyle_pcs/compute_pc_strips.py @@ -0,0 +1,442 @@ +#!/usr/bin/env python3 +"""Compute PCA strip artifacts from cDINO embeddings and zarr image stores. + +For each of N principal components, selects representative cells spanning +low-to-high along that axis and extracts Phase2D crop images. Outputs a +directory of artifacts consumed by build_static_explorer.py. + +Outputs (in --output-dir): + representatives.json — selected cells with PC, bin, gene, score, position + crops_png/ — 96×96 Phase2D crops named pc{NNN}_bin{NN}_row{N}.png + gene_names.json — ordered list of gene names + gene_pc_scores.npy — (n_genes, n_pcs) mean PC score per gene + +Data format: + --emb-dir: per-gene .pt files. + Each .pt: {embeddings: tensor[N,D], cell_metadata} + cell_metadata: {experiment, well, x_pheno, y_pheno, segmentation_id} + (each as list-of-lists, one inner list per FOV group) + + --zarr-base: experiment dirs, each with: + 3-assembly/phenotyping_v3.zarr/{row}/{col}/0/0 (zarr v3 array) + Shape [1, C, 1, Y, X]. Phase2D = channel 0 by default. + +Examples: + # Fit PCA from scratch and save model + python3 compute_pc_strips.py \\ + --emb-dir /path/to/embeddings \\ + --zarr-base /path/to/zarrs \\ + --output-dir /path/to/artifacts \\ + --n-components 97 \\ + --save-model /path/to/pca_model + + # Reuse a saved PCA model (skip 40-min fit) + python3 compute_pc_strips.py \\ + --emb-dir /path/to/embeddings \\ + --zarr-base /path/to/zarrs \\ + --output-dir /path/to/artifacts \\ + --load-model /path/to/pca_model +""" + +import argparse +import heapq +import json +import os +import time + +import numpy as np +import torch + +# ── Zarr v3 compatibility patch ── +try: + import zarr + from zarr.core.metadata.v3 import ArrayV3Metadata + _orig_from_dict = ArrayV3Metadata.from_dict.__func__ + @classmethod + def _patched_from_dict(cls, data): + if isinstance(data, dict): + data.pop("storage_transformers", None) + return _orig_from_dict(cls, data) + ArrayV3Metadata.from_dict = _patched_from_dict +except Exception: + pass + + +def list_gene_files(emb_dir): + return sorted([ + f for f in os.listdir(emb_dir) + if f.endswith(".pt") and f != "metadata.pt" + ]) + + +def load_gene_embeddings(path): + d = torch.load(path, weights_only=False) + return d["embeddings"].numpy() + + +def load_gene_metadata(path): + d = torch.load(path, weights_only=False) + cm = d["cell_metadata"] + n = d["embeddings"].shape[0] + flat = {k: [] for k in ("experiment", "well", "x_pheno", "y_pheno")} + for g in range(len(cm["experiment"])): + for k in flat: + flat[k].extend(cm[k][g]) + assert len(flat["experiment"]) == n + return flat + + +# ── PCA ── + +def load_or_fit_pca(emb_dir, n_components, batch_size, save_model=None, load_model=None): + from sklearn.decomposition import IncrementalPCA + + gene_files = list_gene_files(emb_dir) + + if load_model: + print(f" Loading PCA model from {load_model}") + components = np.load(os.path.join(load_model, "components.npy")) + mean = np.load(os.path.join(load_model, "mean.npy")) + ev = np.load(os.path.join(load_model, "explained_variance.npy")) + n_components = components.shape[0] + else: + pca = IncrementalPCA(n_components=n_components) + print(f" PCA fit: {len(gene_files)} genes, n_components={n_components}") + total = 0 + t0 = time.time() + for gi, gf in enumerate(gene_files): + emb = load_gene_embeddings(os.path.join(emb_dir, gf)) + for start in range(0, emb.shape[0], batch_size): + batch = emb[start:start + batch_size] + if batch.shape[0] < n_components: + continue + pca.partial_fit(batch) + total += emb.shape[0] + if (gi + 1) % 200 == 0: + print(f" [{gi+1}/{len(gene_files)}] {total:,} cells, {time.time()-t0:.0f}s") + print(f" Fit complete: {total:,} cells, {time.time()-t0:.0f}s") + print(f" Explained variance: {pca.explained_variance_ratio_.sum()*100:.1f}%") + components = pca.components_ + mean = pca.mean_ + ev = pca.explained_variance_ratio_ + + if save_model: + os.makedirs(save_model, exist_ok=True) + np.save(os.path.join(save_model, "components.npy"), components) + np.save(os.path.join(save_model, "mean.npy"), mean) + np.save(os.path.join(save_model, "explained_variance.npy"), ev) + print(f" Saved PCA model to {save_model}") + + def scorer(embeddings): + return (embeddings - mean) @ components.T + + return scorer, n_components, ev.tolist() + + +# ── Representative selection ── + +def select_representatives(emb_dir, scorer, n_strips, n_bins, n_rows, batch_size): + gene_files = list_gene_files(emb_dir) + + print(" Pass 1/3: counting cells...") + total_cells = 0 + for gf in gene_files: + total_cells += load_gene_embeddings(os.path.join(emb_dir, gf)).shape[0] + print(f" {total_cells:,} cells across {len(gene_files)} genes") + + print(" Pass 2/3: computing percentile boundaries (200K subsample)...") + rng = np.random.RandomState(42) + subsample_scores = [] + for gf in gene_files: + emb = load_gene_embeddings(os.path.join(emb_dir, gf)) + mask = rng.random(emb.shape[0]) < min(1.0, 200_000 / max(total_cells, 1)) + if mask.sum() > 0: + subsample_scores.append(scorer(emb[mask])) + subsample_scores = np.concatenate(subsample_scores, axis=0) + print(f" Subsample: {subsample_scores.shape[0]:,} cells") + + percentiles = np.linspace(0, 100, n_bins) + bin_targets = np.zeros((n_strips, n_bins)) + for s in range(n_strips): + bin_targets[s] = np.percentile(subsample_scores[:, s], percentiles) + del subsample_scores + + print(f" Pass 3/3: selecting representatives ({n_strips} x {n_bins} x {n_rows})...") + heaps = [[[] for _ in range(n_bins)] for _ in range(n_strips)] + gene_names = [] + gene_mean_scores = [] + t0 = time.time() + + for gi, gf in enumerate(gene_files): + gene_name = gf.replace(".pt", "") + emb = load_gene_embeddings(os.path.join(emb_dir, gf)) + gene_names.append(gene_name) + gene_score_sum = np.zeros(n_strips) + gene_score_count = 0 + + for start in range(0, emb.shape[0], batch_size): + batch = emb[start:start + batch_size] + scores = scorer(batch) + gene_score_sum += scores.sum(axis=0) + gene_score_count += scores.shape[0] + + for s in range(n_strips): + s_scores = scores[:, s] + dists_all = np.abs(s_scores[:, None] - bin_targets[s][None, :]) + nearest_bin = np.argmin(dists_all, axis=1) + for bi in range(n_bins): + mask = (nearest_bin == bi) + if not mask.any(): + continue + indices = np.where(mask)[0] + dists = dists_all[indices, bi] + k = min(n_rows, len(indices)) + if k >= len(indices): + top_k = range(len(indices)) + else: + top_k = np.argpartition(dists, k)[:k] + for ti in top_k: + idx = int(indices[ti]) + d = float(dists[ti]) + local_i = start + idx + entry = (-d, gene_name, local_i, float(scores[idx, s])) + heap = heaps[s][bi] + if len(heap) < n_rows: + heapq.heappush(heap, entry) + elif -d > heap[0][0]: + heapq.heapreplace(heap, entry) + + gene_mean_scores.append(gene_score_sum / max(gene_score_count, 1)) + + if (gi + 1) % 200 == 0: + print(f" [{gi+1}/{len(gene_files)}] {time.time()-t0:.0f}s") + + reps = [] + for s in range(n_strips): + for bi in range(n_bins): + entries = sorted(heaps[s][bi], key=lambda x: x[3]) + for row, (neg_d, gene, local_i, score) in enumerate(entries): + reps.append((s, bi, row, gene, local_i, score)) + + print(f" Selected {len(reps)} representatives in {time.time()-t0:.0f}s") + return reps, total_cells, gene_names, np.array(gene_mean_scores) + + +# ── Crop extraction ── + +def extract_crops(emb_dir, zarr_base, reps, crop_size, phase_channel): + print(f" Loading metadata for {len(reps)} representatives...") + + needed = {} + for s, bi, row, gene, local_i, score in reps: + needed.setdefault(gene, set()).add(local_i) + + cell_info = {} + for gene, indices in needed.items(): + meta = load_gene_metadata(os.path.join(emb_dir, gene + ".pt")) + for idx in indices: + cell_info[(gene, idx)] = { + "experiment": meta["experiment"][idx], + "well": meta["well"][idx], + "x": meta["x_pheno"][idx], + "y": meta["y_pheno"][idx], + } + + # Deduplicate by spatial position within each strip + seen_per_strip = {} + deduped_reps = [] + for r in reps: + s, bi, row, gene, local_i, score = r + info = cell_info[(gene, local_i)] + pos_key = (info["experiment"], info["well"], + int(round(info["x"])), int(round(info["y"]))) + seen = seen_per_strip.setdefault(s, set()) + if pos_key in seen: + continue + seen.add(pos_key) + deduped_reps.append(r) + if len(deduped_reps) < len(reps): + print(f" Removed {len(reps) - len(deduped_reps)} spatial duplicates") + reps = deduped_reps + + print(f" Extracting {len(reps)} crops...") + crop_half = crop_size // 2 + zarr_cache = {} + crops = {} + failures = 0 + t0 = time.time() + + for ri, (s, bi, row, gene, local_i, score) in enumerate(reps): + info = cell_info[(gene, local_i)] + x = int(round(info["x"])) + y = int(round(info["y"])) + exp = info["experiment"] + well = info["well"] + well_row, well_col = well[0], well[1:] + + key = (exp, well) + if key not in zarr_cache: + zpath = os.path.join( + zarr_base, exp, + "3-assembly", "phenotyping_v3.zarr", + well_row, well_col, "0", "0" + ) + try: + zarr_cache[key] = zarr.open(zpath, mode="r") + except Exception: + zarr_cache[key] = None + + arr = zarr_cache[key] + crop = None + if arr is not None: + try: + _, _, _, H, W = arr.shape + y0, y1 = max(0, y - crop_half), min(H, y + crop_half) + x0, x1 = max(0, x - crop_half), min(W, x + crop_half) + raw = np.array(arr[0, phase_channel, 0, y0:y1, x0:x1]) + if raw.shape != (crop_size, crop_size): + padded = np.zeros((crop_size, crop_size), dtype=raw.dtype) + py = (crop_size - raw.shape[0]) // 2 + px = (crop_size - raw.shape[1]) // 2 + padded[py:py+raw.shape[0], px:px+raw.shape[1]] = raw + raw = padded + crop = raw + except Exception: + failures += 1 + + crops[(s, bi, row)] = crop + if (ri + 1) % 1000 == 0: + print(f" [{ri+1}/{len(reps)}] {failures} failures, {time.time()-t0:.0f}s") + + ok = sum(1 for v in crops.values() if v is not None) + print(f" Done: {ok} crops, {failures} failures") + return reps, crops, cell_info + + +# ── Save artifacts ── + +def save_artifacts(output_dir, reps, crops, cell_info, gene_names, + gene_mean_scores, explained_variance, n_bins, n_rows, crop_size): + from PIL import Image + + os.makedirs(output_dir, exist_ok=True) + crops_dir = os.path.join(output_dir, "crops_png") + os.makedirs(crops_dir, exist_ok=True) + + print(f" Saving crop PNGs to {crops_dir}") + saved = 0 + for (s, bi, row), crop in sorted(crops.items()): + if crop is None: + continue + lo, hi = np.percentile(crop, (1, 99)) + if hi - lo < 1e-6: + hi = lo + 1 + normed = np.clip((crop - lo) / (hi - lo), 0, 1) + img = Image.fromarray((normed * 255).astype(np.uint8), mode="L") + img.save(os.path.join(crops_dir, f"pc{s:03d}_bin{bi:02d}_row{row}.png")) + saved += 1 + print(f" {saved} PNGs saved") + + rep_list = [] + for s, bi, row, gene, local_i, score in reps: + info = cell_info.get((gene, local_i), {}) + crop = crops.get((s, bi, row)) + rep_list.append({ + "pc": s, + "bin": bi, + "row": row, + "gene": gene, + "local_i": local_i, + "score": round(float(score), 4), + "experiment": info.get("experiment", ""), + "well": info.get("well", ""), + "x": round(float(info.get("x", 0)), 1), + "y": round(float(info.get("y", 0)), 1), + "has_crop": crop is not None, + }) + with open(os.path.join(output_dir, "representatives.json"), "w") as f: + json.dump({ + "representatives": rep_list, + "cells_per_row": n_bins, + "n_rows": n_rows, + "crop_size": crop_size, + }, f) + print(f" representatives.json: {len(rep_list)} entries") + + with open(os.path.join(output_dir, "gene_names.json"), "w") as f: + json.dump(gene_names, f) + np.save(os.path.join(output_dir, "gene_pc_scores.npy"), gene_mean_scores) + print(f" gene_names.json: {len(gene_names)} genes") + print(f" gene_pc_scores.npy: {gene_mean_scores.shape}") + + with open(os.path.join(output_dir, "gene_pc_analysis.json"), "w") as f: + json.dump({"explained_variance": explained_variance}, f) + print(f" gene_pc_analysis.json saved") + + +def main(): + parser = argparse.ArgumentParser( + description=__doc__, + formatter_class=argparse.RawDescriptionHelpFormatter, + ) + parser.add_argument("--emb-dir", required=True, + help="Per-gene .pt files with embeddings + cell_metadata") + parser.add_argument("--zarr-base", required=True, + help="Base directory of experiment zarr stores") + parser.add_argument("--output-dir", required=True, + help="Directory for output artifacts") + + pca = parser.add_argument_group("PCA model") + pca.add_argument("--n-components", type=int, default=97, + help="Number of PCA components (default: 97)") + pca.add_argument("--save-model", + help="Save fitted PCA model to this directory") + pca.add_argument("--load-model", + help="Load a previously saved PCA model (skip fitting)") + + vis = parser.add_argument_group("strip layout") + vis.add_argument("--cells-per-row", type=int, default=15, + help="Bins along each strip (default: 15)") + vis.add_argument("--n-rows", type=int, default=3, + help="Cells per bin (default: 3)") + vis.add_argument("--crop-size", type=int, default=96, + help="Crop size in pixels (default: 96)") + vis.add_argument("--phase-channel", type=int, default=0, + help="Zarr channel index for Phase2D (default: 0)") + vis.add_argument("--batch-size", type=int, default=50000, + help="Batch size for streaming (default: 50000)") + + args = parser.parse_args() + + gene_files = list_gene_files(args.emb_dir) + print(f"Found {len(gene_files)} gene files in {args.emb_dir}") + + print(f"\n=== PCA ({args.n_components} components) ===") + scorer, n_strips, ev = load_or_fit_pca( + args.emb_dir, args.n_components, args.batch_size, + save_model=args.save_model, load_model=args.load_model + ) + + print(f"\n=== Selecting representatives ===") + reps, total_cells, gene_names, gene_mean_scores = select_representatives( + args.emb_dir, scorer, n_strips, + args.cells_per_row, args.n_rows, args.batch_size + ) + + print(f"\n=== Extracting crops ===") + reps, crops, cell_info = extract_crops( + args.emb_dir, args.zarr_base, reps, args.crop_size, + phase_channel=args.phase_channel + ) + + print(f"\n=== Saving artifacts to {args.output_dir} ===") + save_artifacts(args.output_dir, reps, crops, cell_info, + gene_names, gene_mean_scores, ev, + args.cells_per_row, args.n_rows, args.crop_size) + + print(f"\nDone. {total_cells:,} cells, {len(gene_names)} genes, {n_strips} PCs.") + print(f"Next: python3 build_static_explorer.py --artifacts-dir {args.output_dir}") + + +if __name__ == "__main__": + main() diff --git a/src/ops_model/models/attention/diffex/viewer/__init__.py b/src/ops_model/models/attention/diffex/viewer/__init__.py new file mode 100644 index 0000000..475e90f --- /dev/null +++ b/src/ops_model/models/attention/diffex/viewer/__init__.py @@ -0,0 +1 @@ +"""DiffEx traversal viewer: precompute α-frame assets + manifest for a static MOPS-style tool.""" diff --git a/src/ops_model/models/attention/diffex/viewer/_altanchor_build.py b/src/ops_model/models/attention/diffex/viewer/_altanchor_build.py new file mode 100644 index 0000000..4a2402d --- /dev/null +++ b/src/ops_model/models/attention/diffex/viewer/_altanchor_build.py @@ -0,0 +1,123 @@ +"""Rebuild the A->B alt-anchor traversals in viewer_assets_v5 with the MULTIRANK ranking: + - the existing directed pairs (from each __to__ dir's meta), both directions + - 45 anchor cells per traversal, ranked top-45 by the MULTIRANK (pma_shap_phase_{geneKO,complex}.parquet) + for phase, and the per-marker fluor_shap ranking for fluorescence + - then score each with the v5 SetTransformer → P(target B) + rank_target +Sharded by anchor class so no two jobs write the same _anchors/ real-cell dir (force rebuilds it). +""" +import os, json, glob +from . import catalog as C +from ..classifier.config import slugify + +ASSETS = "viewer_assets_v5" +ROOT = C.OUT +V5G = f"{ROOT}/{ASSETS}/_rankings/pma_shap_phase_geneKO.parquet" # MULTIRANK (was pma_v5) +V5C = f"{ROOT}/{ASSETS}/_rankings/pma_shap_phase_complex.parquet" +FRANK = f"{ROOT}/{ASSETS}/_rankings/fluor_shap" # per-marker: {grain}/{slug}.parquet +PHASE_CK = f"{C.DD}/phase_v1/diffae_best.pt" +NEW = "" # curated-new-pairs json (optional; rebuild uses the existing __to__ pairs) +N_CELLS = 45 + + +PAIRS_JSON = os.path.join(os.path.dirname(__file__), "altanchor_pairs.json") # canonical committed pair list + + +def _all_directed(grain): + """Canonical undirected pairs (altanchor_pairs.json) → both directions. Returns (directed_pairs, classes).""" + undirected = {tuple(sorted((a, b))) for a, b in json.load(open(PAIRS_JSON))[grain] if a != b} + directed = [d for a, b in undirected for d in ((a, b), (b, a))] + return directed, sorted({c for p in directed for c in p}) + + +def gen_shard(grain, classes, pairs): + os.environ["OPS_DIFFEX_ASSETS"] = ASSETS + from . import precompute as P + P._ASSETS = ASSETS + parq = V5G if grain == "geneKO" else V5C + P.precompute_anchors_marker(grain=grain, classes=classes, ckpt=PHASE_CK, out_root=ROOT, + marker_channel=None, channel="Phase2D", n_cells=N_CELLS, + pairs=pairs, accuracy_parquet=parq, force=True, v5_score=True) + + +def gen_shard_fluor(d, marker_channel, channel, grain, classes, pairs): + """Fluor A→B alt-anchors for ONE marker: anchor cells from the marker's fluor_shap MULTIRANK ranking, + morphed in the marker channel with that marker's DiffAE checkpoint.""" + os.environ["OPS_DIFFEX_ASSETS"] = ASSETS + from . import precompute as P + P._ASSETS = ASSETS + parq = f"{FRANK}/{grain}/{slugify(marker_channel)}.parquet" + P.precompute_anchors_marker(grain=grain, classes=classes, ckpt=f"{C.DD}/{d}/diffae_best.pt", out_root=ROOT, + marker_channel=marker_channel, channel=channel, n_cells=N_CELLS, + pairs=pairs, accuracy_parquet=parq, force=True, v5_score=True) + + +def submit_fluor(): + """Same canonical A→B pairs for every fluor marker, filtered to pairs whose BOTH classes exist in that + marker's fluor_shap ranking (a marker only covers the geneKOs it's distinctive for). One shard per marker×grain.""" + import pandas as pd + from ops_utils.hpc.slurm_batch_utils import submit_parallel_jobs + jobs = [] + for d, mc, ch in C.complete_markers(): + slug = slugify(mc) + for grain in ["geneKO", "complex"]: + parq = f"{FRANK}/{grain}/{slug}.parquet" + if not os.path.exists(parq): + continue + cc = "gene" if grain == "geneKO" else "predicted_class" + avail = set(pd.read_parquet(parq, columns=[cc])[cc].astype(str).unique()) + directed, _ = _all_directed(grain) + pairs = [(a, b) for a, b in directed if a in avail and b in avail] + if not pairs: + continue + classes = sorted({c for p in pairs for c in p}) + jobs.append({"name": f"falt_{grain[0]}_{slug[:12]}", "func": gen_shard_fluor, + "kwargs": {"d": d, "marker_channel": mc, "channel": ch, "grain": grain, "classes": classes, "pairs": pairs}}) + print(f"[altanchor-fluor] {len(jobs)} marker×grain shards (canonical pairs, marker-filtered)") + submit_parallel_jobs(jobs, experiment="diffex_altanchor_fluor", + slurm_params={"slurm_partition": "preempted", "gpus_per_node": 1, "cpus_per_task": 12, + "mem_gb": 96, "timeout_min": 300, "slurm_constraint": "[a100_80|h100|h200|6000_blackwell]", + "slurm_additional_parameters": {"requeue": True}, + "slurm_setup": ["export OPS_DIFFEX_ASSETS=viewer_assets_v5"]}, + log_dir="diffex_altanchor_fluor", wait_for_completion=False) + + +def score_shard(grain): + os.environ["OPS_DIFFEX_ASSETS"] = ASSETS + from . import score_generated as SG + SG.V5_BASE = f"{ROOT}/{ASSETS}/phase" + SG.score_anchor_traversals(grain) + + +def main(n_shards=10): + os.environ["OPS_DIFFEX_ASSETS"] = ASSETS + from ops_utils.hpc.slurm_batch_utils import submit_parallel_jobs + jobs = [] + for grain in ["geneKO", "complex"]: + directed, _ = _all_directed(grain) + by_anchor = {} + for a, b in directed: + by_anchor.setdefault(a, []).append((a, b)) + anchors = sorted(by_anchor) + shards = [anchors[i::n_shards] for i in range(n_shards)] + for i, sh in enumerate(shards): + if not sh: + continue + pairs = [p for a in sh for p in by_anchor[a]] + classes = sorted({c for p in pairs for c in p}) + jobs.append({"name": f"altanchor_{grain}_{i}", "func": gen_shard, + "kwargs": {"grain": grain, "classes": classes, "pairs": pairs}}) + print(f"[altanchor] {grain}: {len(directed)} directed pairs, {len(anchors)} anchors") + print(f"[altanchor] submitting {len(jobs)} phase gen shards (MULTIRANK)") + submit_parallel_jobs(jobs, experiment="diffex_altanchor", + slurm_params={"slurm_partition": "gpu", "gpus_per_node": 1, "cpus_per_task": 12, + "mem_gb": 96, "timeout_min": 300, + "slurm_setup": ["export OPS_DIFFEX_ASSETS=viewer_assets_v5"]}, + log_dir="diffex_altanchor", wait_for_completion=False) + + +if __name__ == "__main__": + import sys + if len(sys.argv) > 1 and sys.argv[1] == "fluor": + submit_fluor() + else: + main() # phase A→B alt-anchors (multirank) diff --git a/src/ops_model/models/attention/diffex/viewer/_anchortest.py b/src/ops_model/models/attention/diffex/viewer/_anchortest.py new file mode 100644 index 0000000..8bbd073 --- /dev/null +++ b/src/ops_model/models/attention/diffex/viewer/_anchortest.py @@ -0,0 +1,19 @@ +"""DIAGNOSTIC (removable): regenerate a couple of geneKO traversals with v4 NTC anchor cells +(attention-ranked → typical, in the DiffAE training distribution) but v5 KD centroids, to test +whether the v5 graininess / extreme-alpha blowout comes from selecting NTC anchors by set-accuracy +(which picks the most extreme/atypical control cells) instead of attention. + +Launch with OPS_DIFFEX_ASSETS=viewer_assets_v5_anchortest and OPS_DIFFEX_V5 UNSET so the control +(NTC) gather reads the v4 parquet while the KD gather uses the v5 accuracy parquet override. +Output is isolated under viewer_assets_v5_anchortest/. Delete this module after the diagnosis.""" +from . import catalog as C +from .precompute import precompute_marker + +V5_GENEKO = "/hpc/projects/icd.fast.ops/models/diffex/viewer_assets_v5/_rankings/pma_v5_phase_geneKO.parquet" +PHASE_CK = f"{C.DD}/phase_v1/diffae_best.pt" + + +def run_anchor_test(targets=("HSPA5", "SRSF3")): + # control="NTC" → v4 parquet (OPS_DIFFEX_V5 unset); accuracy_parquet → v5 KD centroids. + return precompute_marker(grain="geneKO", targets=list(targets), ckpt=PHASE_CK, out_root=C.OUT, + control="NTC", accuracy_parquet=V5_GENEKO, score=False, force=True) diff --git a/src/ops_model/models/attention/diffex/viewer/_build_stepablation.py b/src/ops_model/models/attention/diffex/viewer/_build_stepablation.py new file mode 100644 index 0000000..f1bc3b4 --- /dev/null +++ b/src/ops_model/models/attention/diffex/viewer/_build_stepablation.py @@ -0,0 +1,74 @@ +"""DDIM-step ablation: does the inverted-anchor scheme need >50 steps to resolve? Regenerate a small fixed set +of high-recovery geneKO at ddim_steps in {50,100,200}, everything else identical to valid200 (inverted, w=1.5, +same anchors + v5 directions). Score centroid recovery + SetTransformer (inline v5) + a sharpness metric to see +whether 100/200 steps recovers the loss the inversion introduced at 50. Small set — 100/200 steps are 2x/4x slow. + + VAL_STEPS=100 python -m ...viewer._build_stepablation gen +""" +import os, sys, json +from pathlib import Path +import numpy as np + +from . import catalog as C +from .precompute import precompute_marker + +STEPS = int(os.environ.get("VAL_STEPS", "50")) +INVERT = int(os.environ.get("VAL_INVERT", "1")) # 0 = random-xT (non-inverted, the old scheme) +DIRS = os.environ.get("VAL_DIRS", "viewer_assets_v5") # which tree's _directions to use (old vs current) +_TREE = f"viewer_assets_stepabl_s{STEPS}" + ("" if INVERT else "_randxt") + ("" if DIRS == "viewer_assets_v5" else f"_dirs_{DIRS.replace('viewer_assets_','')}") +_SRC = "viewer_assets_valid200" # reuse its 200-NTC anchors +PHASE_CK = f"{C.DD}/phase_v1/diffae_best.pt" +V5G = f"{C.OUT}/viewer_assets_v5/_rankings/pma_v5_phase_geneKO.parquet" +ALPHAS = (0.0, 0.5, 1.0, 2.0, 3.0, 4.0, 5.0) +NCELLS = 45 # match the original v5 SetTransformer bag=45 (bag matters for set-accuracy) +FT = "/hpc/projects/icd.fast.ops/analysis/figure4_traversals" + + +def sel(n=15): + d = json.load(open(f"{FT}/gen_real_centroid/geneKO_scored.json")); al = d["alphas"]; by = d["by_alpha"] + best = {} + for a in al: + for g, v in by[str(a)]["top1"].items(): + best[g] = max(best.get(g, 0), v) + return [g for g, _ in sorted(best.items(), key=lambda x: -x[1])[:n]] + + +def setup(): + root = Path(C.OUT) / _TREE; (root / "phase").mkdir(parents=True, exist_ok=True) + for name, tgt in [("_directions", Path(C.OUT) / DIRS / "_directions"), + ("phase/_anchors", Path(C.OUT) / _SRC / "phase" / "_anchors")]: + ln = root / name + if not ln.exists(): + ln.symlink_to(tgt); print(f"[stepabl s{STEPS}] symlinked {name}") + + +def build(targets): + os.environ["OPS_DIFFEX_ASSETS"] = _TREE + from . import precompute as P + P._ASSETS = _TREE + return precompute_marker(grain="geneKO", targets=list(targets), ckpt=PHASE_CK, out_root=C.OUT, + control="NTC", n_cells=NCELLS, alphas=ALPHAS, invert_anchors=bool(INVERT), w=1.5, + force=False, v5_score=True, accuracy_parquet=V5G, load_workers=12, + ddim_steps=STEPS) + + +def main(): + if sys.argv[1:2] == ["gen"]: + from ops_utils.hpc.slurm_batch_utils import submit_parallel_jobs + setup() + genes = sel(15); print(f"[stepabl s{STEPS}] genes: {genes}") + tmo = {50: 90, 100: 120, 200: 200}[STEPS] # per-gene timeout (1 gene/shard → parallel) + jobs = [{"name": f"stepabl_s{STEPS}_{g}", "func": build, "kwargs": {"targets": [g]}} for g in genes] + submit_parallel_jobs(jobs_to_submit=jobs, + experiment=f"stepabl_s{STEPS}", + slurm_params={"slurm_partition": "preempted", "gpus_per_node": 1, "cpus_per_task": 12, + "mem_gb": 96, "timeout_min": tmo, "slurm_constraint": "[a40|a6000|l40s]", + "slurm_setup": [f"export OPS_DIFFEX_ASSETS={_TREE}", f"export VAL_STEPS={STEPS}", f"export VAL_INVERT={INVERT}", f"export VAL_DIRS={DIRS}", + "export PYTORCH_CUDA_ALLOC_CONF=expandable_segments:True"]}, + log_dir=f"stepabl_s{STEPS}", wait_for_completion=False) + else: + raise SystemExit("usage: _build_stepablation gen (set VAL_STEPS)") + + +if __name__ == "__main__": + main() diff --git a/src/ops_model/models/attention/diffex/viewer/_build_v5_inverted.py b/src/ops_model/models/attention/diffex/viewer/_build_v5_inverted.py new file mode 100644 index 0000000..adfb096 --- /dev/null +++ b/src/ops_model/models/attention/diffex/viewer/_build_v5_inverted.py @@ -0,0 +1,586 @@ +"""FINAL v5 traversals with DDIM-inverted anchors (faithful α=0, w=1.5) into viewer_assets_v5. + +MARKERS: anchor = top-40 accuracy NTC from the per-marker v5 ranking (_rankings/fluor/geneKO/.parquet += fluor_rank_parquet). The existing v5 ctrl.npz (embeddings only) is dropped so the SAME accuracy cells are +re-materialized LOSSLESSLY (verified identical to v5: anchor-match 1.000) — enabling faithful inversion +(α=0-vs-real 0.988) instead of inverting the lossy real.webp. Run under OPS_DIFFEX_ASSETS=viewer_assets_v5. + +COVERAGE: phase gets all 1000 geneKO + 98 complex; each fluor marker gets only the 35-423 geneKO (+~36-54 +complex) whose per-cell rankings exist. That is Alex Lin's top1_acc>0.5@100-cell distinctiveness filter (see +_fluor_v5_build.py) — NOT missing data (cells exist for all 1000×55). Lower-acc combos need Alex to gen more. + + python -m ops_model.models.attention.diffex.viewer._build_v5_inverted markers +""" +import json +import os +import sys +from pathlib import Path + +from . import catalog as C +from ..classifier.config import slugify +from .precompute import precompute_marker + +FRP_DIR = f"{C.OUT}/viewer_assets_v5/_rankings/fluor_shap/geneKO" # NEW shap_screen rankings (robust bin-size top-acc); old at _rankings/fluor/geneKO_OLD_qualifying backup +# OUTPUT tree: a FRESH dir so force=False gives skip-done resume (timeouts harmless) without clobbering the old +# non-inverted v5 traversals. _directions + each /_anchors are SYMLINKED back to v5 (setup_v5inv below), so +# the 11,219 ckpt/w-independent directions + prebuilt lossless anchors are reused with no re-gather. +_V5 = "viewer_assets_v5" # single production assets tree (was viewer_assets_v5_inv; merged + retired) +_V5EMB = "viewer_assets_v5_emb" # embeddings-only tree: CellDINO of float frames, no webp (parallel to _V5) + + +def _use_v5(): + """Force the target tree to _V5 at RUNTIME (inside the job) — the import-time snapshot of + precompute._ASSETS is unreliable across submitit workers. Returns the assets dirname to use for paths.""" + os.environ["OPS_DIFFEX_ASSETS"] = _V5 + from . import precompute as P + P._ASSETS = _V5 + return _V5 + + +def _genes_for(frp): + import pandas as pd + g = pd.read_parquet(frp, columns=["gene"])["gene"].astype(str) + return sorted(x for x in set(g) if not x.startswith("NTC")) + + +def build_marker(d, marker_channel, channel): + """Drop the embeddings-only v5 ctrl.npz → precompute_marker re-materializes the SAME top-40 accuracy + anchors LOSSLESS (fresh gather from fluor_rank_parquet) + inverts them; directions recompute from the + same accuracy parquet (v5-equivalent).""" + frp = f"{FRP_DIR}/{slugify(marker_channel)}.parquet" + if not os.path.exists(frp): + return f"skip {marker_channel}: no fluor_rank_parquet" + ctrl = Path(C.OUT) / _V5 / slugify(marker_channel) / "_anchors" / "NTC" / "ctrl.npz" + if ctrl.exists(): + os.remove(ctrl) # force fresh lossless anchor gather (same v5 cells) + return precompute_marker(grain="geneKO", targets=_genes_for(frp), marker_channel=marker_channel, + channel=channel, ckpt=f"{C.DD}/{d}/diffae_best.pt", out_root=C.OUT, + control="NTC", fluor_rank_parquet=frp, n_cells=40, invert_anchors=True, + w=1.5, force=True, v5_score=True) + + +CFRP_DIR = f"{C.OUT}/viewer_assets_v5/_rankings/fluor_shap/complex" # NEW shap_screen EBI-pooled complex rankings +SEL25 = "/hpc/projects/icd.fast.ops/analysis/figure4_traversals/ntc_accanchor_selected25.csv" +PHASE_CK = f"{C.DD}/phase_v1/diffae_best.pt" +PMA_SHAP_G = f"{C.OUT}/viewer_assets_v5/_rankings/pma_shap_phase_geneKO.parquet" # NEW shap_screen phase rankings +PMA_SHAP_C = f"{C.OUT}/viewer_assets_v5/_rankings/pma_shap_phase_complex.parquet" + + +def build_phase_shap(grain, targets, tree="viewer_assets_v5", save_gemb=False, ddim_steps=100, alphas=None, + n_cells=45, cell_range=None): + """PHASE traversals from the NEW shap_screen rankings: force=True recomputes each target's direction from + the new-shap cells (accuracy_parquet), keeping the cached NTC anchors. `tree` = output assets dir + (production viewer_assets_v5 for the real rebuild, or a scratch tree + save_gemb for the 15-gene test). + `alphas` overrides the α grid (test uses the 7 forward α to match the step-ablation arms; None → full 17). + `cell_range` (lo,hi) + force → REUSE cached directions and only sample that cell slice (the 200-cell top-up).""" + os.environ["OPS_DIFFEX_ASSETS"] = tree + from . import precompute as P + P._ASSETS = tree + acc = PMA_SHAP_G if grain == "geneKO" else PMA_SHAP_C + kw = {} if alphas is None else {"alphas": tuple(alphas)} + if cell_range is not None: + kw["cell_range"] = tuple(cell_range) + return precompute_marker(grain=grain, targets=list(targets), ckpt=PHASE_CK, out_root=C.OUT, + control="NTC", n_cells=n_cells, invert_anchors=True, w=1.5, force=True, + v5_score=True, accuracy_parquet=acc, ddim_steps=ddim_steps, + save_gemb=save_gemb, load_workers=12, **kw) + + +def build_phase_anchor_200(): + """Extend the phase NTC anchor to 200 cells for the 200-cell top-up: cells 0..99 keep the existing 100 + (attention+accuracy) anchors; cells 100..199 append the next top-accuracy NTC (aligned real imgs + embs) + so inversion stays faithful across all 200. Overwrites _anchors/NTC ctrl.npz + real.webp cell100..199.""" + import numpy as np, pandas as pd + from pathlib import Path + from concurrent.futures import ThreadPoolExecutor + from .precompute import _gather_class, _save_webp + from ..diffae.data import normalize + from ..directions.config import DirConfig + _use_v5() + rd = Path(C.OUT) / _V5 / "phase" / "_anchors" / "NTC" + z = dict(np.load(rd / "ctrl.npz")) + have = len(z["anchor_imgs"]) + if have >= 200: + return {"note": f"anchor already {have} cells"} + cfg = DirConfig(grain="geneKO", target="NTC", control="NTC", device="cuda") + imgs, emb = _gather_class(cfg, "NTC", 200) # 200 top-accuracy NTC (aligned imgs↔embs) + real200 = normalize(imgs[:200]); emb200 = emb[:200] + real200[:have] = z["anchor_imgs"][:have] # keep the existing 0..have anchors identical + emb200[:have] = z["ctrl_embs"][:have] + tp = ThreadPoolExecutor(8) + for c in range(have, 200): + (rd / f"cell{c}").mkdir(parents=True, exist_ok=True); tp.submit(_save_webp, rd / f"cell{c}" / "real.webp", real200[c, 0], 256) + tp.shutdown(wait=True) + np.savez(rd / "ctrl.npz", ctrl_embs=emb200, mu_ctrl=emb200.mean(0), anchor_imgs=real200) + print(f"[phase-anchor] extended {have} → 200 anchors -> {rd/'ctrl.npz'}") + return {"from": have, "to": 200} + + +def build_phase_topup(grain, targets, lo=45, hi=200, tree="viewer_assets_v5"): + """200-cell top-up: for targets missing the top cell (hi-1), sample cell_range=(lo,hi) reusing the cached + new-shap directions (cells 0..lo-1 kept). Idempotent (skips genes already topped up). The filter must check + the SAME tree build_phase_shap writes to (tree, default viewer_assets_v5) — NOT _V5 (viewer_assets_v5_inv), + or every round re-inverts all genes.""" + base = Path(C.OUT) / tree / "phase" / grain + todo = [t for t in targets if not (base / slugify(t) / f"cell{hi - 1}" / "frame_00.webp").exists()] + if not todo: + return {"grain": grain, "done": 0, "note": "all topped up"} + print(f"[phase-topup] {grain}: {len(todo)}/{len(targets)} targets → cells {lo}..{hi - 1}") + return build_phase_shap(grain, todo, tree=tree, n_cells=hi, cell_range=(lo, hi)) + + +def build_phase_anchor_multirank(hi=400): + """MULTI-BAG anchor: append cells 200..hi-1 = top-(hi-200) NTC from the MULTIRANK ranking (PMA_SHAP_G), + keeping the existing single-bag anchors 0..199 intact. This is the fix for the anchor gather using the old + pma_v5 ranking — here we gather NTC by the multirank so cells 200-399 are the clean multi_bag anchors.""" + import numpy as np + from concurrent.futures import ThreadPoolExecutor + from .precompute import _gather_class, _save_webp + from ..diffae.data import normalize + from ..directions.config import DirConfig + _use_v5() + rd = Path(C.OUT) / _V5 / "phase" / "_anchors" / "NTC" + z = dict(np.load(rd / "ctrl.npz")) + have = len(z["anchor_imgs"]) + if have >= hi: + return {"note": f"anchor already {have} cells"} + n_new = hi - 200 + cfg = DirConfig(grain="geneKO", target="NTC", control="NTC", device="cuda") + imgs, emb = _gather_class(cfg, "NTC", n_new, parquet=PMA_SHAP_G) # multirank NTC ranks 1..n_new + real_new = normalize(imgs[:n_new]) + real_all = np.concatenate([z["anchor_imgs"][:200], real_new], axis=0) # 0..199 single_bag kept; 200.. multirank + emb_all = np.concatenate([z["ctrl_embs"][:200], emb[:n_new]], axis=0) + tp = ThreadPoolExecutor(8) + for c in range(200, hi): + (rd / f"cell{c}").mkdir(parents=True, exist_ok=True) + tp.submit(_save_webp, rd / f"cell{c}" / "real.webp", real_all[c, 0], 256) + tp.shutdown(wait=True) + np.savez(rd / "ctrl.npz", ctrl_embs=emb_all, mu_ctrl=emb_all.mean(0), anchor_imgs=real_all) + print(f"[phase-anchor-mr] appended multirank {have}→{hi} anchors -> {rd/'ctrl.npz'}") + return {"from": have, "to": hi} + + +def _submit_multibag(gk_partition="preempted", cx_partition="gpu"): + """MULTI-BAG phase traversals (cells 200-399, multirank anchor + cached directions, 100-step inverted). + Anchor built once (idempotent — skipped if already 400). Frame shards SPLIT by grain: geneKO scavenges + gk_partition (default preempted, +requeue), complex runs on cx_partition (default gpu). single_bag 0-199 untouched.""" + import numpy as np + from ops_utils.hpc.slurm_batch_utils import submit_parallel_jobs + ap = Path(C.OUT) / _V5 / "phase" / "_anchors" / "NTC" / "ctrl.npz" + have = len(np.load(ap)["anchor_imgs"]) if ap.exists() else 0 + after = None + if have < 400: # build the multirank anchor first + sp_a = dict(_gpu_sp(120)); sp_a["slurm_partition"] = cx_partition; sp_a["slurm_constraint"] = "[a40|a6000|l40s]" + r = submit_parallel_jobs(jobs_to_submit=[{"name": "mbag_anchor", "func": build_phase_anchor_multirank, "kwargs": {"hi": 400}}], + experiment="diffex_v5inv", slurm_params=sp_a, log_dir="diffex_v5inv", wait_for_completion=False) + after = r["base_job_id"] + for grain, targets, tag, part in [("geneKO", C.all_genes(), "g", gk_partition), ("complex", C.ebi_complexes(), "c", cx_partition)]: + chunk = 2 if grain == "geneKO" else 6 # 200 cells/gene is heavy → 2 geneKO/shard fits the 12h limit (6 timed out) + jobs = [{"name": f"mbag_{tag}{i}", "func": build_phase_topup, "kwargs": {"grain": grain, "targets": targets[i:i + chunk], "lo": 200, "hi": 400}} + for i in range(0, len(targets), chunk)] + sp = dict(_gpu_sp(720)); sp["slurm_partition"] = part; sp["slurm_constraint"] = "[a40|a6000|l40s]" + addl = {} + if after: + addl["dependency"] = f"afterany:{after}" + if part == "preempted": + addl["requeue"] = True # survive preemption instead of FAIL + if addl: + sp["slurm_additional_parameters"] = addl + print(f"[multibag] {grain}: {len(jobs)} frame shards on {part} (cells 200-399)" + (f" [after {after}]" if after else "")) + submit_parallel_jobs(jobs_to_submit=jobs, experiment="diffex_v5inv", slurm_params=sp, log_dir="diffex_v5inv", wait_for_completion=False) + + +def build_phase_shap_resume(grain, targets, cutoff, tree="viewer_assets_v5", ddim_steps=100): + """Resume-aware phase-shap build for the preempted chain: regenerate only targets NOT already rebuilt this + run — meta.json missing, older than `cutoff` (this rebuild's launch time), or not at `ddim_steps`. Idempotent + across rounds, so a preempted round is picked up by the next without redoing finished targets.""" + base = Path(C.OUT) / tree / "phase" / grain + todo = [] + for t in targets: + m = base / slugify(t) / "meta.json" + if not m.exists() or m.stat().st_mtime < cutoff: + todo.append(t); continue + try: + if json.load(open(m)).get("ddim_steps") != ddim_steps: + todo.append(t) + except Exception: + todo.append(t) + if not todo: + return {"grain": grain, "done": 0, "note": "all fresh"} + print(f"[phase-shap resume] {grain}: {len(todo)}/{len(targets)} targets to (re)build → {ddim_steps} steps") + return build_phase_shap(grain, todo, tree=tree, ddim_steps=ddim_steps) + + +def resume_marker(d, marker_channel, channel, target_steps=100): + """Resume/upgrade a marker: KEEP the (step-independent) anchors and regenerate only the genes not already + built at `target_steps`. A gene is regenerated if its meta.json is missing, was built at a different + ddim_steps (e.g. old 50-step frames → the 100-step relaunch), or predates the anchor gather. Genes already + at target_steps are skipped, so this is idempotent across preemption/timeout rounds — and it correctly + distinguishes 50- vs 100-step frames via the stamped meta['ddim_steps'] (mtime alone cannot).""" + frp = f"{FRP_DIR}/{slugify(marker_channel)}.parquet" + if not os.path.exists(frp): + return f"skip {marker_channel}: no fluor_rank_parquet" + base = Path(C.OUT) / _V5 / slugify(marker_channel) + ctrl = base / "_anchors" / "NTC" / "ctrl.npz" + if not ctrl.exists(): + return build_marker(d, marker_channel, channel) # no anchors yet → full build + cutoff = ctrl.stat().st_mtime + stale = [] + for g in _genes_for(frp): + m = base / "geneKO" / slugify(g) / "meta.json" + if not m.exists() or m.stat().st_mtime < cutoff: + stale.append(g); continue + try: + steps = json.load(open(m)).get("ddim_steps") + except Exception: + steps = None + if steps != target_steps: # built at a different step count (50→100) or unstamped (old) + stale.append(g) + if not stale: + return {"marker": marker_channel, "stale": 0, "note": f"already {target_steps}-step fresh"} + print(f"[resume] {marker_channel}: regenerating {len(stale)} genes → {target_steps} steps (keeping anchors)") + return precompute_marker(grain="geneKO", targets=stale, marker_channel=marker_channel, + channel=channel, ckpt=f"{C.DD}/{d}/diffae_best.pt", out_root=C.OUT, + control="NTC", fluor_rank_parquet=frp, n_cells=40, invert_anchors=True, + w=1.5, force=True, v5_score=True, ddim_steps=target_steps) + + +def build_phase_anchor(): + """Pre-build phase _anchors/NTC ctrl.npz with the 45 MIXED anchors: 20 attention (phase geneKO parquet + top-20) + 25 accuracy (SEL25). Per-anchor z0 embeddings + LOSSLESS anchor_imgs so inversion is faithful + for BOTH pools (the accuracy cells are NOT in the attention-ranked embeddings, so this is required).""" + import numpy as np, pandas as pd + from pathlib import Path + from concurrent.futures import ThreadPoolExecutor + from .precompute import _gather_class, _save_webp + from ..diffae.data import normalize + from ..directions.config import DirConfig + _use_v5() + cfg = DirConfig(grain="geneKO", target="NTC", control="NTC", device="cuda") + a_imgs, a_emb = _gather_class(cfg, "NTC", 20) # 20 attention + sel = pd.read_csv(SEL25) + parq = pd.DataFrame({"gene": "NTC", "experiment": sel.experiment, "well": sel.well, + "segmentation": sel.segmentation, "x_pheno": sel.x_pheno, "y_pheno": sel.y_pheno, + "pma_attention": sel.pma_attention, "rank": range(1, len(sel) + 1), "rank_type": "top"}) + tmp = f"{C.OUT}/{_V5}/_ntc25.parquet"; Path(tmp).parent.mkdir(parents=True, exist_ok=True); parq.to_parquet(tmp) + b_imgs, b_emb = _gather_class(cfg, "NTC", 25, parquet=tmp) # 25 accuracy (SEL25) + embs = np.concatenate([a_emb[:20], b_emb[:25]], 0) + real = normalize(np.concatenate([a_imgs[:20], b_imgs[:25]], 0)) # 45 lossless, cell0..44 + rd = Path(C.OUT) / _V5 / "phase" / "_anchors" / "NTC"; rd.mkdir(parents=True, exist_ok=True) + tp = ThreadPoolExecutor(8) + for c in range(len(real)): + (rd / f"cell{c}").mkdir(parents=True, exist_ok=True); tp.submit(_save_webp, rd / f"cell{c}" / "real.webp", real[c, 0], 256) + tp.shutdown(wait=True) + np.savez(rd / "ctrl.npz", ctrl_embs=embs, mu_ctrl=embs.mean(0), anchor_imgs=real) + print(f"[phase-anchor] pre-built 45 mixed anchors (20 attn + 25 acc): {real.shape}") + + +def build_phase(grain, targets): + """Gen phase frames inverted into the fresh _V5 tree, REUSING cached v5 directions (via symlinked + _directions) + the pre-built 45 anchors (symlinked _anchors). force=False → skip already-done targets + (resume) and reuse cached directions; the fresh dir has no stale meta to drop.""" + _use_v5() + return precompute_marker(grain=grain, targets=list(targets), ckpt=PHASE_CK, out_root=C.OUT, + control="NTC", n_cells=45, invert_anchors=True, w=1.5, force=False, v5_score=True) + + +def _use_v5emb(): + os.environ["OPS_DIFFEX_ASSETS"] = _V5EMB + from . import precompute as P + P._ASSETS = _V5EMB + return _V5EMB + + +def setup_v5inv_emb(): + """Embeddings-only tree: symlink _directions + phase/_anchors from viewer_assets_v5 so directions and the + 45-cell lossless phase anchor are reused (no re-gather). Traversal dirs written fresh; we save gemb.npz only.""" + v5 = Path(C.OUT) / "viewer_assets_v5" + emb = Path(C.OUT) / _V5EMB + emb.mkdir(parents=True, exist_ok=True) + if not (emb / "_directions").exists(): + (emb / "_directions").symlink_to(v5 / "_directions") + (emb / "phase").mkdir(parents=True, exist_ok=True) + if not (emb / "phase" / "_anchors").exists(): + (emb / "phase" / "_anchors").symlink_to(v5 / "phase" / "_anchors") + print(f"[v5inv-emb] setup {emb}: _directions + phase/_anchors symlinked from v5") + + +def build_phase_embed(grain, targets): + """Re-decode the inverted phase traversals and save the in-memory float CellDINO embeddings (gemb.npz), + NO webp — the webp viewer assets come from the parallel _V5 run. Reuses cached directions + anchors.""" + _use_v5emb() + return precompute_marker(grain=grain, targets=list(targets), ckpt=PHASE_CK, out_root=C.OUT, + control="NTC", n_cells=45, invert_anchors=True, w=1.5, force=False, + score=False, v5_score=False, save_gemb=True, skip_webp=True) + + +_V5EMBC = "viewer_assets_v5_emb_cmp" # test tree: gemb.npz with BOTH float + webp-roundtrip embeddings + + +def build_phase_embed_cmp(grain, targets): + """Same as build_phase_embed but also stores the webp-round-tripped embedding (webp_compare=True) so we can + compare float-vs-webp mapping on the IDENTICAL inverted frames. Separate tree; force=True (small subset).""" + os.environ["OPS_DIFFEX_ASSETS"] = _V5EMBC + from . import precompute as P + P._ASSETS = _V5EMBC + v5 = Path(C.OUT) / "viewer_assets_v5"; emb = Path(C.OUT) / _V5EMBC + emb.mkdir(parents=True, exist_ok=True) + if not (emb / "_directions").exists(): + (emb / "_directions").symlink_to(v5 / "_directions") + (emb / "phase").mkdir(parents=True, exist_ok=True) + if not (emb / "phase" / "_anchors").exists(): + (emb / "phase" / "_anchors").symlink_to(v5 / "phase" / "_anchors") + return precompute_marker(grain=grain, targets=list(targets), ckpt=PHASE_CK, out_root=C.OUT, + control="NTC", n_cells=45, invert_anchors=True, w=1.5, force=True, + score=False, v5_score=False, save_gemb=True, skip_webp=True, webp_compare=True) + + +def phase_test(grain, targets): + build_phase_anchor() + return build_phase(grain, targets) + + +def build_marker_complex(d, marker_channel, channel): + """Complex traversals: REUSE the geneKO-built anchors (ctrl.npz already has lossless anchor_imgs) — do + NOT drop it — and take complex directions from the per-marker complex ranking. Run after geneKO.""" + frp = f"{CFRP_DIR}/{slugify(marker_channel)}.parquet" + if not os.path.exists(frp): + return f"skip {marker_channel}: no complex ranking" + return precompute_marker(grain="complex", targets=C.ebi_complexes(), marker_channel=marker_channel, + channel=channel, ckpt=f"{C.DD}/{d}/diffae_best.pt", out_root=C.OUT, + control="NTC", fluor_rank_parquet=frp, n_cells=40, invert_anchors=True, + w=1.5, force=True, v5_score=True) + + +def _submit(kind, func, after=None): + from ops_utils.hpc.slurm_batch_utils import submit_parallel_jobs + fdir = FRP_DIR if kind == "markers" else CFRP_DIR + jobs = [{"name": f"v5{kind[:3]}_{slugify(mc)[:14]}", "func": func, + "kwargs": {"d": d, "marker_channel": mc, "channel": ch}} + for d, mc, ch in C.complete_markers() + if os.path.exists(f"{fdir}/{slugify(mc)}.parquet")] + sp = {"slurm_partition": "gpu", "gpus_per_node": 1, "cpus_per_task": 12, "mem_gb": 96, + "timeout_min": 720, "slurm_constraint": "[a100_80|h100|h200|6000_blackwell]", + "slurm_setup": ["export OPS_DIFFEX_ASSETS=viewer_assets_v5", + "export PYTORCH_CUDA_ALLOC_CONF=expandable_segments:True"]} + if after: + sp["slurm_additional_parameters"] = {"dependency": f"afterany:{after}"} + print(f"[v5inv] submitting {len(jobs)} {kind} builds → viewer_assets_v5 (inverted, w=1.5)" + + (f" [after {after}]" if after else "")) + submit_parallel_jobs(jobs_to_submit=jobs, experiment="diffex_v5inv", slurm_params=sp, + log_dir="diffex_v5inv", wait_for_completion=False) + + +def _submit_phase(): + """Full phase: all ~1000 geneKO + 98 complexes, sharded. Anchor ctrl.npz (45 mixed) is already + pre-built; each shard reuses it + the cached v5 accuracy directions (force=False), inverting.""" + from ops_utils.hpc.slurm_batch_utils import submit_parallel_jobs + genes, cx = C.all_genes(), C.ebi_complexes() + ch = lambda lst, n: [lst[i:i + n] for i in range(0, len(lst), n)] + jobs = [{"name": f"v5ph_g{i}", "func": build_phase, "kwargs": {"grain": "geneKO", "targets": s}} + for i, s in enumerate(ch(genes, 40))] + jobs += [{"name": f"v5ph_c{i}", "func": build_phase, "kwargs": {"grain": "complex", "targets": s}} + for i, s in enumerate(ch(cx, 40))] + print(f"[v5inv] submitting phase: {len(genes)} geneKO + {len(cx)} complex → {len(jobs)} shards (inverted, w=1.5)") + submit_parallel_jobs( + jobs_to_submit=jobs, experiment="diffex_v5inv", + slurm_params={"slurm_partition": "gpu", "gpus_per_node": 1, "cpus_per_task": 12, "mem_gb": 96, + "timeout_min": 720, "slurm_constraint": "[a100_80|h100|h200|6000_blackwell]", + "slurm_setup": ["export OPS_DIFFEX_ASSETS=viewer_assets_v5", + "export PYTORCH_CUDA_ALLOC_CONF=expandable_segments:True"]}, + log_dir="diffex_v5inv", wait_for_completion=False) + + +def _submit_altanchors(): + """Rebuild the EXISTING phase A→B alt-anchors (80 geneKO + 80 complex) with inversion — reuse the + sharded generator (V5G/V5C accuracy cells; precompute_anchors_marker now inverts by default). Sharded + by anchor class so no two jobs write the same _anchors/ real-cell dir.""" + import glob, json + from ops_utils.hpc.slurm_batch_utils import submit_parallel_jobs + from ._altanchor_build import gen_shard + jobs = [] + for grain in ["geneKO", "complex"]: + by_anchor = {} + for d in glob.glob(f"{C.OUT}/viewer_assets_v5/phase/{grain}/*__to__*"): + m = json.load(open(f"{d}/meta.json")); by_anchor.setdefault(m["control"], []).append((m["control"], m["target"])) + anchors = sorted(by_anchor); nsh = max(1, min(10, len(anchors))) + for i in range(nsh): + sh = anchors[i::nsh] + if not sh: + continue + ps = [p for a in sh for p in by_anchor[a]] + jobs.append({"name": f"v5alt_{grain}_{i}", "func": gen_shard, + "kwargs": {"grain": grain, "classes": sorted({c for p in ps for c in p}), "pairs": ps}}) + print(f"[v5alt] {grain}: {sum(len(v) for v in by_anchor.values())} pairs across {len(anchors)} anchors") + print(f"[v5alt] submitting {len(jobs)} alt-anchor shards → viewer_assets_v5 (inverted, w=1.5)") + submit_parallel_jobs( + jobs_to_submit=jobs, experiment="diffex_v5inv", + slurm_params={"slurm_partition": "gpu", "gpus_per_node": 1, "cpus_per_task": 12, "mem_gb": 96, + "timeout_min": 300, "slurm_constraint": "[a100_80|h100|h200|6000_blackwell]", + "slurm_setup": ["export OPS_DIFFEX_ASSETS=viewer_assets_v5", + "export PYTORCH_CUDA_ALLOC_CONF=expandable_segments:True"]}, + log_dir="diffex_v5inv", wait_for_completion=False) + + +# ============ FAST ARCHITECTURE: pre-build lossless anchors ONCE, then gene-chunk gen-shards ============ +CHUNK = 25 # small genes/shard: better parallelism packing + tiny tail; resume covers any timeout + + +def prebuild_marker_anchor(d, marker_channel, channel): + """Explicitly write ctrl.npz with LOSSLESS anchor_imgs (top-40 accuracy NTC) + 200 embeddings (for + directions) + real.webp. No reliance on drop/fresh-gather-branch — this IS the anchor build.""" + import numpy as np, pandas as pd + from concurrent.futures import ThreadPoolExecutor + from .precompute import _gather_class, _save_webp + from ..diffae.data import normalize + from ..directions.config import DirConfig + frp = f"{FRP_DIR}/{slugify(marker_channel)}.parquet" + if not os.path.exists(frp): + return f"skip {marker_channel}" + _use_v5() + cfg = DirConfig(grain="geneKO", target="NTC", control="NTC", device="cuda") + cfg.marker_channel = marker_channel; cfg.channel = channel; cfg._fluor_rows = pd.read_parquet(frp) + imgs, embs = _gather_class(cfg, "NTC", 200) # top-200 accuracy NTC (embs → directions) + n = min(40, len(embs)); real = normalize(imgs[:n]) # 40 anchors, lossless + rd = Path(C.OUT) / _V5 / slugify(marker_channel) / "_anchors" / "NTC"; rd.mkdir(parents=True, exist_ok=True) + tp = ThreadPoolExecutor(8) + for c in range(n): + (rd / f"cell{c}").mkdir(parents=True, exist_ok=True); tp.submit(_save_webp, rd / f"cell{c}" / "real.webp", real[c, 0], 256) + tp.shutdown(wait=True) + np.savez(rd / "ctrl.npz", ctrl_embs=embs, mu_ctrl=embs.mean(0), anchor_imgs=real) + return f"prebuilt {marker_channel}: {n} anchors" + + +def genshard(d, marker_channel, channel, grain, targets): + """Gen frames for a chunk of targets, reusing the pre-built lossless anchors (invert) + accuracy dirs.""" + _use_v5() + frp = f"{(FRP_DIR if grain == 'geneKO' else CFRP_DIR)}/{slugify(marker_channel)}.parquet" + return precompute_marker(grain=grain, targets=list(targets), marker_channel=marker_channel, channel=channel, + ckpt=f"{C.DD}/{d}/diffae_best.pt", out_root=C.OUT, control="NTC", + fluor_rank_parquet=frp, n_cells=40, invert_anchors=True, w=1.5, force=False, v5_score=True) + + +def _gpu_sp(timeout, parallel=64): + return {"slurm_partition": "gpu", "gpus_per_node": 1, "cpus_per_task": 12, "mem_gb": 96, + "timeout_min": timeout, "slurm_constraint": "[a100_80|h100|h200|6000_blackwell]", + "slurm_array_parallelism": parallel, + "slurm_setup": ["export OPS_DIFFEX_ASSETS=viewer_assets_v5", + "export PYTORCH_CUDA_ALLOC_CONF=expandable_segments:True"]} + + +def _submit_prebuild(): + from ops_utils.hpc.slurm_batch_utils import submit_parallel_jobs + jobs = [{"name": f"v5pre_{slugify(mc)[:14]}", "func": prebuild_marker_anchor, + "kwargs": {"d": d, "marker_channel": mc, "channel": ch}} + for d, mc, ch in C.complete_markers() if os.path.exists(f"{FRP_DIR}/{slugify(mc)}.parquet")] + jobs.append({"name": "v5pre_phase", "func": build_phase_anchor, "kwargs": {}}) + print(f"[v5pre] {len(jobs)} anchor pre-builds → viewer_assets_v5 (lossless, fast)") + submit_parallel_jobs(jobs_to_submit=jobs, experiment="diffex_v5inv", slurm_params=_gpu_sp(60), + log_dir="diffex_v5inv", wait_for_completion=False) + + +def setup_v5inv(): + """Stand up the viewer_assets_v5 output tree: symlink _directions + each modality's _anchors back + to viewer_assets_v5 so cached directions (ckpt/w-independent) and prebuilt lossless anchors are reused with + no re-gather. Traversal dirs (geneKO/, complex/) are written fresh. Idempotent.""" + v5 = Path(C.OUT) / "viewer_assets_v5" + inv = Path(C.OUT) / _V5 + inv.mkdir(parents=True, exist_ok=True) + dl = inv / "_directions" + if not dl.exists(): + dl.symlink_to(v5 / "_directions") + mods = [slugify(mc) for d, mc, ch in C.complete_markers() if os.path.exists(f"{FRP_DIR}/{slugify(mc)}.parquet")] + ["phase"] + n = 0 + for m in mods: + src = v5 / m / "_anchors" + if not src.exists(): + print(f" WARN no anchors in v5 for {m}"); continue + (inv / m).mkdir(parents=True, exist_ok=True) + al = inv / m / "_anchors" + if not al.exists(): + al.symlink_to(src); n += 1 + print(f"[v5inv] setup {inv}: _directions symlinked, {n}/{len(mods)} modality _anchors symlinked") + + +def _submit_gen(after): + from ops_utils.hpc.slurm_batch_utils import submit_parallel_jobs + ch = lambda lst, n: [lst[i:i + n] for i in range(0, len(lst), n)] + cx = C.ebi_complexes(); jobs = [] + for d, mc, chan in C.complete_markers(): + if not os.path.exists(f"{FRP_DIR}/{slugify(mc)}.parquet"): + continue + for i, s in enumerate(ch(_genes_for(f"{FRP_DIR}/{slugify(mc)}.parquet"), CHUNK)): + jobs.append({"name": f"g_{slugify(mc)[:10]}_{i}", "func": genshard, + "kwargs": {"d": d, "marker_channel": mc, "channel": chan, "grain": "geneKO", "targets": s}}) + if os.path.exists(f"{CFRP_DIR}/{slugify(mc)}.parquet"): + jobs.append({"name": f"c_{slugify(mc)[:10]}", "func": genshard, + "kwargs": {"d": d, "marker_channel": mc, "channel": chan, "grain": "complex", "targets": cx}}) + for i, s in enumerate(ch(C.all_genes(), CHUNK)): + jobs.append({"name": f"ph_g{i}", "func": build_phase, "kwargs": {"grain": "geneKO", "targets": s}}) + jobs.append({"name": "ph_c", "func": build_phase, "kwargs": {"grain": "complex", "targets": cx}}) + sp = _gpu_sp(240) + if after: + sp["slurm_additional_parameters"] = {"dependency": f"afterany:{after}"} + print(f"[v5gen] {len(jobs)} gen-shards (chunk={CHUNK}) → viewer_assets_v5, parallelism 64" + (f" [after {after}]" if after else "")) + submit_parallel_jobs(jobs_to_submit=jobs, experiment="diffex_v5gen", slurm_params=sp, + log_dir="diffex_v5gen", wait_for_completion=False) + + +def _submit_fluor(partition="gpu"): + """FLUOR buildout only (no phase): ONE genshard per marker per grain (geneKO all genes + complex), no + chunking — matches the working 35083270/35083684 chains. genshard needs an 80GB+ GPU (weak a40/a6000/l40s + OOM at ~1min), so keep the strong constraint from _gpu_sp. MUST be launched from a shared-fs cwd (repo), + not /tmp, or compute nodes cannot write submitit logs and every task dies at ~1min.""" + from ops_utils.hpc.slurm_batch_utils import submit_parallel_jobs + cx = C.ebi_complexes(); jobs = [] + for d, mc, chan in C.complete_markers(): + frp = f"{FRP_DIR}/{slugify(mc)}.parquet" + if not os.path.exists(frp): + continue + jobs.append({"name": f"g_{slugify(mc)[:12]}", "func": genshard, + "kwargs": {"d": d, "marker_channel": mc, "channel": chan, "grain": "geneKO", "targets": _genes_for(frp)}}) + if os.path.exists(f"{CFRP_DIR}/{slugify(mc)}.parquet"): + jobs.append({"name": f"c_{slugify(mc)[:12]}", "func": genshard, + "kwargs": {"d": d, "marker_channel": mc, "channel": chan, "grain": "complex", "targets": cx}}) + sp = dict(_gpu_sp(720)); sp["slurm_partition"] = partition + print(f"[fluor] {len(jobs)} per-marker gen-shards ({partition}, strong GPU) → viewer_assets_v5") + submit_parallel_jobs(jobs_to_submit=jobs, experiment="diffex_v5gen", slurm_params=sp, + log_dir="diffex_v5gen", wait_for_completion=False) + + +def _submit_topup(partition="gpu"): + """PHASE 200-cell top-up: build_phase_topup per gene-chunk for genes still missing cell199 (idempotent). + Light job → weak GPUs are fine. Launch from a shared-fs cwd (repo), not /tmp.""" + from ops_utils.hpc.slurm_batch_utils import submit_parallel_jobs + genes = C.all_genes() + jobs = [{"name": f"topup_{i}", "func": build_phase_topup, "kwargs": {"grain": "geneKO", "targets": genes[i:i + 6]}} + for i in range(0, len(genes), 6)] + # 155-cell guided inversion is ~30min/gene on weak GPUs → 6 genes/shard, 12h timeout (5h timed out at 8/shard) + sp = dict(_gpu_sp(720)); sp["slurm_partition"] = partition; sp["slurm_constraint"] = "[a40|a6000|l40s]" + print(f"[topup] {len(jobs)} shards ({partition}, weak GPU ok) → viewer_assets_v5/phase/geneKO") + submit_parallel_jobs(jobs_to_submit=jobs, experiment="diffex_v5inv", slurm_params=sp, + log_dir="diffex_v5inv", wait_for_completion=False) + + +if __name__ == "__main__": + # NOTE: run from a shared-fs cwd (the repo), NOT /tmp — submitit writes logs relative to cwd and compute + # nodes cannot read /tmp, so /tmp-launched jobs all die at ~1min with empty logs. + cmd = sys.argv[1] if len(sys.argv) > 1 else "" + arg2 = sys.argv[2] if len(sys.argv) > 2 else None + if cmd == "prebuild": + _submit_prebuild() + elif cmd == "setup": + setup_v5inv() + elif cmd == "gen": + setup_v5inv() + _submit_gen(arg2) + elif cmd == "fluor": + _submit_fluor(arg2 or "gpu") + elif cmd == "topup": + _submit_topup(arg2 or "gpu") + elif cmd == "multibag": + _submit_multibag(arg2 or "preempted", (sys.argv[3] if len(sys.argv) > 3 else "gpu")) # gk_partition cx_partition + elif cmd == "altanchors": + _submit_altanchors() + else: + print("usage: _build_v5_inverted.py prebuild | setup | gen [after_jobid] | fluor [partition] | topup [partition] | altanchors") diff --git a/src/ops_model/models/attention/diffex/viewer/_build_v5_montages.py b/src/ops_model/models/attention/diffex/viewer/_build_v5_montages.py new file mode 100644 index 0000000..259e8c6 --- /dev/null +++ b/src/ops_model/models/attention/diffex/viewer/_build_v5_montages.py @@ -0,0 +1,119 @@ +"""Per-marker embedding montages for completed v5 markers. + +Same as the v4 montage build (alphas 1-5, cells 0-19, umap+phate) EXCEPT each marker's tiles are laid out +on ITS OWN gene embedding (marker_leaves.embedding_h5ad → paper_v2/markers//…) instead of the shared +phase embedding. Reads the merged inverted frames from viewer_assets_v5//geneKO and writes tiles to +viewer_assets_v5/_montage/. Only builds markers whose geneKO is 100% present in viewer_assets_v5. + + python -m ops_model.models.attention.diffex.viewer._build_v5_montages +""" +import glob +import os + +from . import catalog as C +from . import marker_leaves as ML +from .build_umap_montage import OUT +from ..classifier.config import slugify + +V5 = "viewer_assets_v5" +CELLS = list(range(20)) +ALPHAS = [1.0, 2.0, 3.0, 4.0, 5.0] +EMBS = ["umap", "phate"] + + +BACKUP = f"{OUT}/viewer_assets_v5_preinvert_backup" # only markers we MERGED (inverted) have a backup here + + +def _completed_in_v5(): + """(marker_channel, slug) for markers whose INVERTED geneKO was merged into viewer_assets_v5 — i.e. the + old build was moved to viewer_assets_v5_preinvert_backup// AND geneKO is 100% present in v5. This + excludes markers still showing the OLD non-inverted build (no backup) so we only montage inverted frames.""" + from ._build_v5_inverted import _genes_for, FRP_DIR + out = [] + for d, mc, ch in C.complete_markers(): + s = slugify(mc) + frp = f"{FRP_DIR}/{s}.parquet" + if not os.path.exists(frp): # all v5 fluor is inverted now (_inv retired) → gate only on a ranking + complete geneKO + continue + exp = len(_genes_for(frp)) + dn = len([x for x in glob.glob(f"{OUT}/{V5}/{s}/geneKO/*") + if os.path.isdir(x) and "__to__" not in os.path.basename(x) and os.path.exists(f"{x}/meta.json")]) + if exp > 0 and dn >= exp: + out.append((mc, s)) + return out + + +def mont_job(marker_channel, slug, cell, alpha, emb): + """One (marker, cell, α, emb) montage → viewer_assets_v5/_montage, laid out on the marker's own embedding.""" + os.environ["OPS_DIFFEX_ASSETS"] = V5 + from . import build_umap_montage as BM + BM._ASSETS = V5 # runtime override (import-time snapshot is unreliable) + h5 = ML.embedding_h5ad(marker_channel) or f"{ML.PHASE_LEAF}/gene_embedding_pca_optimized.h5ad" + oz = f"{OUT}/{V5}/_montage/{slug}_geneKO_{emb}_cell{cell}_a{alpha:g}.zarr" + return BM.build_montage_web(h5ad=h5, out_zarr=oz, cell=cell, alpha=alpha, embedding=emb, modality=slug) + + +PHASE_CELLS = list(range(45)) # single_bag anchors (cells 0-44) +MULTIBAG_CELLS = list(range(200, 250)) # multi_bag anchors: first 50 of the 200 multirank cells (disk 200-249 → display rank 1-50) + + +def phase_mont_job(cell, alpha, emb): + """Phase montage on the phase gene embedding. Reads the inverted frames from viewer_assets_v5/phase + (phase swapped into production 2026-07-23), writes tiles to viewer_assets_v5/_montage.""" + os.environ["OPS_DIFFEX_ASSETS"] = V5 + from . import build_umap_montage as BM + BM._ASSETS = V5 + h5 = ML.embedding_h5ad(None) # phase_only leaf gene embedding + oz = f"{OUT}/{V5}/_montage/phase_geneKO_{emb}_cell{cell}_a{alpha:g}.zarr" + return BM.build_montage_web(h5ad=h5, out_zarr=oz, cell=cell, alpha=alpha, embedding=emb, modality="phase") + + +def submit_phase(cells=PHASE_CELLS): + from ops_utils.hpc.slurm_batch_utils import submit_parallel_jobs + jobs = [{"name": f"mtg5_phase_{emb[:2]}_c{cell}_a{a:g}", "func": phase_mont_job, + "kwargs": {"cell": cell, "alpha": a, "emb": emb}} + for emb in EMBS for cell in cells for a in ALPHAS] + print(f"[v5mont] phase: {len(jobs)} montage jobs ({len(cells)} cells × {len(ALPHAS)} α × {len(EMBS)} emb) → viewer_assets_v5/_montage") + submit_parallel_jobs( + jobs_to_submit=jobs, experiment="diffex_v5mont", + slurm_params={"slurm_partition": "cpu", "cpus_per_task": 4, "mem_gb": 24, "timeout_min": 60, + "slurm_array_parallelism": 100, + "slurm_setup": ["export OPS_DIFFEX_ASSETS=viewer_assets_v5"]}, + log_dir="diffex_v5mont", wait_for_completion=False) + + +def main(): + import sys + if len(sys.argv) > 1 and sys.argv[1] == "phase": + submit_phase(); return + if len(sys.argv) > 1 and sys.argv[1] == "phase_multibag": + submit_phase(cells=MULTIBAG_CELLS); return + force = "force" in sys.argv[1:] # mtime skip is unreliable (frames overwritten in place don't bump the dir mtime) → force a full rebuild + from ops_utils.hpc.slurm_batch_utils import submit_parallel_jobs + comp = _completed_in_v5() + per_marker = sum(1 for mc, _ in comp if ML.embedding_h5ad(mc)) + print(f"[v5mont] {len(comp)} completed markers in v5 ({per_marker} with own embedding, " + f"{len(comp) - per_marker} phase-fallback){' [FORCE]' if force else ''}") + jobs = [] + for mc, s in comp: + gk = f"{OUT}/{V5}/{s}/geneKO" + cm = os.path.getmtime(gk) if os.path.isdir(gk) else 0 # content-aware skip: rebuild only if stale + for emb in EMBS: + for cell in CELLS: + for a in ALPHAS: + tj = f"{OUT}/{V5}/_montage/{s}_geneKO_{emb}_cell{cell}_a{a:g}_tiles/tiles.json" + if not force and os.path.exists(tj) and os.path.getmtime(tj) >= cm: + continue # montage already reflects current cache + jobs.append({"name": f"mtg5_{s[:10]}_{emb[:2]}_c{cell}_a{a:g}", "func": mont_job, + "kwargs": {"marker_channel": mc, "slug": s, "cell": cell, "alpha": a, "emb": emb}}) + print(f"[v5mont] {len(jobs)} montage jobs ({len(comp)}×{len(EMBS)}×{len(CELLS)}×{len(ALPHAS)}) → viewer_assets_v5/_montage") + submit_parallel_jobs( + jobs_to_submit=jobs, experiment="diffex_v5mont", + slurm_params={"slurm_partition": "cpu", "cpus_per_task": 4, "mem_gb": 24, "timeout_min": 60, + "slurm_array_parallelism": 100, + "slurm_setup": ["export OPS_DIFFEX_ASSETS=viewer_assets_v5"]}, + log_dir="diffex_v5mont", wait_for_completion=False) + + +if __name__ == "__main__": + main() diff --git a/src/ops_model/models/attention/diffex/viewer/_build_valid200.py b/src/ops_model/models/attention/diffex/viewer/_build_valid200.py new file mode 100644 index 0000000..6a2a232 --- /dev/null +++ b/src/ops_model/models/attention/diffex/viewer/_build_valid200.py @@ -0,0 +1,169 @@ +"""Scaled-up PHASE validation traversals: 200 NTC anchor cells × all ~1000 geneKO, forward alphas only +(0, 0.5, 1, 2, 3, 4, 5), DDIM-inverted anchors (faithful α=0, w=1.5) — same generator as the v5 production +build, just more cells and fewer alphas, written to a SEPARATE sibling tree so it stays out of the viewer. + +Output: {C.OUT}/viewer_assets_valid200/phase/geneKO//cell/frame_.webp (+ scores_v5.json per target) + (sibling of viewer_assets_v5 — can be symlinked into the viewer later if wanted.) + +Reuses the v5 per-class directions (d_vec + gap are the PRODUCTION values we are validating) via a symlinked +_directions tree, so no direction re-fit — only the 200-anchor inversion + 200×7 decodes per target. + + python -m ops_model.models.attention.diffex.viewer._build_valid200 anchor # 1 GPU: build the 200-cell NTC anchor + python -m ops_model.models.attention.diffex.viewer._build_valid200 submit # shard all 1000 geneKO (after anchor) + python -m ops_model.models.attention.diffex.viewer._build_valid200 all # anchor job -> shards (afterok dep) +""" +import os +import sys +from pathlib import Path + +from ..classifier.config import slugify +from . import catalog as C +from .precompute import precompute_marker + +W = float(os.environ.get("VALID200_W", "1.5")) # CFG guidance weight (baseline recovery used w=2.0) +_VALID = "viewer_assets_valid200" if W == 1.5 else f"viewer_assets_valid200_w{W:g}" +PHASE_CK = f"{C.DD}/phase_v1/diffae_best.pt" +V5G = f"{C.OUT}/viewer_assets_v5/_rankings/pma_v5_phase_geneKO.parquet" # target-cell ranking (only used if a direction is uncached) +V5C = f"{C.OUT}/viewer_assets_v5/_rankings/pma_v5_phase_complex.parquet" # complex target-cell ranking +NCELLS = 200 +VALID_ALPHAS = (0.0, 0.5, 1.0, 2.0, 3.0, 4.0, 5.0) # forward-only morph strengths +CHUNK = 12 # genes/shard (200×7 ≈ 1.8× v5 per-target → half the v5 chunk) + + +def _use_valid(): + """Point precompute at the sibling tree (env is unreliable across submitit workers → set _ASSETS too).""" + os.environ["OPS_DIFFEX_ASSETS"] = _VALID + from . import precompute as P + P._ASSETS = _VALID + return _VALID + + +def setup_dirs(): + """Symlink _directions from viewer_assets_v5 so the 1000 production phase-geneKO directions (d_vec+gap) + are reused verbatim — validation measures the production traversal, not a refit.""" + root = Path(C.OUT) / _VALID + root.mkdir(parents=True, exist_ok=True) + ln = root / "_directions" + v5dir = Path(C.OUT) / "viewer_assets_v5" / "_directions" + if not ln.exists(): + ln.symlink_to(v5dir) + print(f"[valid200] symlinked _directions -> {v5dir}") + else: + print(f"[valid200] _directions already present ({ln})") + if _VALID != "viewer_assets_valid200": # w-variant: reuse the SAME 200 anchors (only w differs) + (root / "phase").mkdir(parents=True, exist_ok=True) + aln = root / "phase" / "_anchors" + if not aln.exists(): + aln.symlink_to(Path(C.OUT) / "viewer_assets_valid200" / "phase" / "_anchors") + print(f"[valid200] symlinked phase/_anchors -> viewer_assets_valid200 (shared 200 anchors)") + + +def build_anchor(): + """Build the 200-cell phase NTC anchor: top-200 ACCURACY-ranked NTC (from V5G, the same accuracy table the + production build uses — NOT top-attention). Saves per-anchor CellDINO embeddings + LOSSLESS anchor_imgs + (needed for faithful DDIM inversion) + real.webp, once.""" + import numpy as np + from concurrent.futures import ThreadPoolExecutor + from .precompute import _gather_class, _save_webp + from ..diffae.data import normalize + from ..directions.config import DirConfig + _use_valid() + setup_dirs() + cfg = DirConfig(grain="geneKO", target="NTC", control="NTC", device="cuda") + imgs, embs = _gather_class(cfg, "NTC", NCELLS, parquet=V5G) # top-200 ACCURACY-ranked NTC + n = min(NCELLS, len(embs)) + real = normalize(imgs[:n]) # lossless anchors, cell0..n-1 + rd = Path(C.OUT) / _VALID / "phase" / "_anchors" / "NTC"; rd.mkdir(parents=True, exist_ok=True) + tp = ThreadPoolExecutor(8) + for c in range(n): + (rd / f"cell{c}").mkdir(parents=True, exist_ok=True) + tp.submit(_save_webp, rd / f"cell{c}" / "real.webp", real[c, 0], 256) + tp.shutdown(wait=True) + np.savez(rd / "ctrl.npz", ctrl_embs=embs[:n], mu_ctrl=embs[:n].mean(0), anchor_imgs=real) + print(f"[valid200] built {n}-cell NTC anchor: {real.shape} -> {rd/'ctrl.npz'}") + return {"n_anchor": n} + + +def build_valid(targets): + """Generate 200-cell × forward-α phase traversals for a chunk of geneKO into the sibling tree, reusing the + pre-built 200 anchors (inverted at startup) + symlinked v5 directions. force=False → resumable.""" + _use_valid() + return precompute_marker(grain="geneKO", targets=list(targets), ckpt=PHASE_CK, out_root=C.OUT, + control="NTC", n_cells=NCELLS, alphas=VALID_ALPHAS, invert_anchors=True, w=W, + force=False, v5_score=True, accuracy_parquet=V5G, load_workers=12) + + +def build_valid_complex(targets): + """Same recipe as build_valid but grain='complex': reuses the SAME 200-NTC anchor (control='NTC', + grain-independent _anchors/NTC) + the production phase/complex directions. → 200-cell complex traversals.""" + _use_valid() + return precompute_marker(grain="complex", targets=list(targets), ckpt=PHASE_CK, out_root=C.OUT, + control="NTC", n_cells=NCELLS, alphas=VALID_ALPHAS, invert_anchors=True, w=W, + force=False, v5_score=True, accuracy_parquet=V5C, load_workers=12) + + +def submit_complex(after=None): + """Shard the 98 EBI complexes (200 cells × 7 α). Anchor already built by the geneKO run — no dep needed.""" + from ops_utils.hpc.slurm_batch_utils import submit_parallel_jobs + cx = C.ebi_complexes() + ch = lambda l, n: [l[i:i + n] for i in range(0, len(l), n)] + jobs = [{"name": f"val200c_{i}", "func": build_valid_complex, "kwargs": {"targets": s}} + for i, s in enumerate(ch(cx, CHUNK))] + sp = _sp() + if after: + sp["slurm_additional_parameters"] = {"dependency": f"afterok:{after}"} + print(f"[valid200] {len(cx)} complexes → {len(jobs)} shards (chunk {CHUNK}, 200 cells × {len(VALID_ALPHAS)} α, inverted w=1.5)") + return submit_parallel_jobs(jobs_to_submit=jobs, experiment="diffex_valid200c", + slurm_params=sp, log_dir="diffex_valid200c", wait_for_completion=False) + + +def _sp(timeout=720): + return {"slurm_partition": "gpu", "gpus_per_node": 1, "cpus_per_task": 12, "mem_gb": 96, + "timeout_min": timeout, "slurm_constraint": "[a100_80|h100|h200|6000_blackwell]", + "slurm_setup": [f"export OPS_DIFFEX_ASSETS={_VALID}", f"export VALID200_W={W:g}", + "export PYTORCH_CUDA_ALLOC_CONF=expandable_segments:True"]} + + +def submit_shards(after=None): + from ops_utils.hpc.slurm_batch_utils import submit_parallel_jobs + genes = C.all_genes() + ch = lambda l, n: [l[i:i + n] for i in range(0, len(l), n)] + jobs = [{"name": f"val200_g{i}", "func": build_valid, "kwargs": {"targets": s}} + for i, s in enumerate(ch(genes, CHUNK))] + sp = _sp() + if after: + sp["slurm_additional_parameters"] = {"dependency": f"afterok:{after}"} + print(f"[valid200] {len(genes)} geneKO → {len(jobs)} shards (chunk {CHUNK}, 200 cells × {len(VALID_ALPHAS)} α, inverted w=1.5)" + + (f" [after {after}]" if after else "")) + return submit_parallel_jobs(jobs_to_submit=jobs, experiment="diffex_valid200", + slurm_params=sp, log_dir="diffex_valid200", wait_for_completion=False) + + +def submit_anchor(): + from ops_utils.hpc.slurm_batch_utils import submit_parallel_jobs + setup_dirs() + return submit_parallel_jobs(jobs_to_submit=[{"name": "val200_anchor", "func": build_anchor, "kwargs": {}}], + experiment="diffex_valid200", slurm_params=_sp(timeout=180), + log_dir="diffex_valid200", wait_for_completion=False) + + +def main(): + cmd = sys.argv[1] if len(sys.argv) > 1 else "all" + if cmd == "anchor": + build_anchor() + elif cmd == "setup": + setup_dirs() + elif cmd == "submit": + submit_shards() + elif cmd == "complex": + submit_complex() + elif cmd == "all": + r = submit_anchor() + aid = str(r.get("base_job_id") or r.get("job_id")) + submit_shards(after=aid) + else: + raise SystemExit(f"unknown cmd {cmd!r}") + + +if __name__ == "__main__": + main() diff --git a/src/ops_model/models/attention/diffex/viewer/_consolidate_cells.py b/src/ops_model/models/attention/diffex/viewer/_consolidate_cells.py new file mode 100644 index 0000000..1c1e65c --- /dev/null +++ b/src/ops_model/models/attention/diffex/viewer/_consolidate_cells.py @@ -0,0 +1,106 @@ +"""Consolidate the top-accuracy anchor cells into the main v5 traversal dirs so all 45 NTC-anchor cells +live in one self-contained place (S3-deployable) instead of two pools + a viewer toggle. + +Per traversal in viewer_assets_v5/phase/{sub}/{name} (that also exists in the accpool): + - copy accpool cell0..24 -> cell20..44 (frames + real.webp) + - meta.json: n_cells = 45, cell_source = ['attention']*20 + ['accuracy']*25 + - scores.json (per-cell linear): first-20 attention + 25 accuracy -> 45 entries + - scores_v5.json (set-score): replaced with the accuracy pool's (bag-20 of the accuracy cells), + tagged score_source='accuracy'; the prior attention set-score is preserved as scores_v5_attention.json +Idempotent: re-running skips already-copied cell dirs and rebuilds meta/scores from the source slices. +""" +import os, glob, json, shutil + +ROOT = "/hpc/projects/icd.fast.ops/models/diffex" +V5, ACC = f"{ROOT}/viewer_assets_v5", f"{ROOT}/viewer_assets_v5_accpool" +ATTN_N = 20 # attention-anchored cells already present as cell0..19 + + +def consolidate_one(sub, name): + v5d, accd = f"{V5}/phase/{sub}/{name}", f"{ACC}/phase/{sub}/{name}" + if not (os.path.isdir(v5d) and os.path.isdir(accd)): + return f"skip {sub}/{name}: missing dir" + acc_cells = sorted((c for c in os.listdir(accd) if c.startswith("cell")), key=lambda x: int(x[4:])) + acc_n = len(acc_cells) + for k in range(acc_n): # copy accpool cell{k} -> v5 cell{ATTN_N+k} + dst = f"{v5d}/cell{ATTN_N + k}" + if not os.path.exists(dst): + shutil.copytree(f"{accd}/cell{k}", dst) + total = ATTN_N + acc_n + # meta + meta = json.load(open(f"{v5d}/meta.json")) + meta["n_cells"] = total + meta["cell_source"] = ["attention"] * ATTN_N + ["accuracy"] * acc_n + json.dump(meta, open(f"{v5d}/meta.json", "w")) + # per-cell linear scores.json: attention[:20] + accuracy[:acc_n] + if os.path.exists(f"{v5d}/scores.json") and os.path.exists(f"{accd}/scores.json"): + v5s, accs = json.load(open(f"{v5d}/scores.json")), json.load(open(f"{accd}/scores.json")) + v5s["scores"] = v5s["scores"][:ATTN_N] + accs["scores"][:acc_n] + json.dump(v5s, open(f"{v5d}/scores.json", "w")) + # set-score scores_v5.json: use accuracy pool's; preserve attention's once + if os.path.exists(f"{accd}/scores_v5.json"): + if os.path.exists(f"{v5d}/scores_v5.json") and not os.path.exists(f"{v5d}/scores_v5_attention.json"): + shutil.copy(f"{v5d}/scores_v5.json", f"{v5d}/scores_v5_attention.json") + acc_sv = json.load(open(f"{accd}/scores_v5.json")); acc_sv["score_source"] = "accuracy" + json.dump(acc_sv, open(f"{v5d}/scores_v5.json", "w")) + return f"ok {sub}/{name}: {total} cells" + + +def shard(items): + return [consolidate_one(sub, name) for sub, name in items] + + +def _targets(): + out = [] + for sub in ["geneKO", "complex"]: + for d in sorted(glob.glob(f"{ACC}/phase/{sub}/*")): + if os.path.isdir(d) and "__to__" not in os.path.basename(d): + out.append((sub, os.path.basename(d))) + return out + + +def consolidate_anchors(): + """Consolidate the shared NTC real-cell anchors (real_dir=phase/_anchors/NTC): copy accpool cell0..24 + -> cell20..44 so the 'show real cells' row has all 45. Other _anchors/* (alt-anchor A cells) stay as-is.""" + v5d, accd = f"{V5}/phase/_anchors/NTC", f"{ACC}/phase/_anchors/NTC" + acc = sorted((c for c in os.listdir(accd) if c.startswith("cell")), key=lambda x: int(x[4:])) + for k in range(len(acc)): + dst = f"{v5d}/cell{ATTN_N + k}" + if not os.path.exists(dst): + shutil.copytree(f"{accd}/cell{k}", dst) + print(f"[anchors] NTC real anchors: {ATTN_N + len(acc)} cells") + + +def finalize_manifest(): + """Set n_cells=45 in manifest.json for every merged target (matches the consolidated meta).""" + mf = f"{V5}/manifest.json" + m = json.load(open(mf)) + merged = 0 + for mk in m["markers"]: + for t in mk.get("targets", []): + v5d = f"{V5}/{t['asset_dir']}" + mp = f"{v5d}/meta.json" + if os.path.exists(mp): + nc = json.load(open(mp)).get("n_cells") + if nc and nc != t.get("n_cells"): + t["n_cells"] = nc; merged += 1 + json.dump(m, open(mf, "w")) + print(f"[manifest] updated n_cells on {merged} targets") + + +def main(n_shards=32): + from ops_utils.hpc.slurm_batch_utils import submit_parallel_jobs + items = _targets() + shards = [s for s in (items[i::n_shards] for i in range(n_shards)) if s] + jobs = [{"name": f"consolidate_{i}", "func": shard, "kwargs": {"items": s}} for i, s in enumerate(shards)] + print(f"[consolidate] {len(items)} traversals across {len(jobs)} shards") + submit_parallel_jobs( + jobs, experiment="diffex_consolidate_cells", + slurm_params={"slurm_partition": "cpu", "cpus_per_task": 4, "mem_gb": 16, "timeout_min": 90}, + log_dir="diffex_consolidate_cells", wait_for_completion=True) + consolidate_anchors() + finalize_manifest() + + +if __name__ == "__main__": + main() diff --git a/src/ops_model/models/attention/diffex/viewer/_fluor_complex_build.py b/src/ops_model/models/attention/diffex/viewer/_fluor_complex_build.py new file mode 100644 index 0000000..52e23a4 --- /dev/null +++ b/src/ops_model/models/attention/diffex/viewer/_fluor_complex_build.py @@ -0,0 +1,136 @@ +"""Build v5 fluorescence COMPLEX-level (EBI) Top Cells for all 55 markers. + +Alex supplied only a per-GENE EBI cell ranking (gene_marker_ebi_complexqual.compact.parquet, cells labeled by +member gene) plus a gene->complex label map (fluor_ebi_bychannel_pergene.csv: gene_name -> label_name). There is +no complex-labeled cell ranking, so we build one: pool each complex's member-gene TOP cells (saturation bag), +re-rank by model score, keep top-N. Reuses crop_marker_shard (key="complexes") to merge a `complexes` block into +each marker's existing top-cells index.json; register_complex() adds grain="complex" targets to the manifest. +""" +import os, re, json +import pandas as pd +from . import catalog as C +from ..classifier.config import slugify +from ._fluor_topcells import crop_marker_shard, TOP_N, OUT + +F = "/hpc/projects/icd.fast.ops/models/alex_lin_attention/v5/fluorescence" +EBI = f"{F}/misc/gene_marker_ebi_complexqual.compact.parquet" # per-GENE EBI cells (all 55 channels) +G2C = f"{F}/fluor_ebi_bychannel_pergene.csv" # gene_name -> label_name (complex) +ASSETS = "viewer_assets_v5" +RANKDIR = f"{C.OUT}/{ASSETS}/_rankings/fluor/complex" +_COLS = ["channel_name", "gene", "bag_size", "rank", "score", "experiment", "well", "x_pheno", "y_pheno", "segmentation_id", "_pool"] + + +def build_rankings(): + """Per-channel complex ranking parquets (geneKO schema, `gene` col = complex name) in RANKDIR. → {channel: (parquet, [complexes])}.""" + os.makedirs(RANKDIR, exist_ok=True) + g2c = (pd.read_csv(G2C, usecols=["gene_name", "label_name"]).dropna() + .drop_duplicates("gene_name").set_index("gene_name")["label_name"].to_dict()) + df = pd.read_parquet(EBI, columns=_COLS) + df = df[df["_pool"] == "top"] # top cells (not the random pool) + mb = df.groupby(["channel_name", "gene"])["bag_size"].transform("max") # saturation bag per member gene + df = df[df["bag_size"] == mb].copy() + df["complex"] = df["gene"].map(g2c) + df = df.dropna(subset=["complex"]) + out = {} + for ch, gch in df.groupby("channel_name"): + parts = [] + for cx, gcx in gch.groupby("complex"): + g = (gcx.drop_duplicates(["experiment", "well", "x_pheno", "y_pheno"]) # one row per cell + .sort_values("score", ascending=False).head(TOP_N).copy()) # re-rank member cells by score + g["rank"] = range(1, len(g) + 1) + g["gene"] = cx # grouping key -> complex name + parts.append(g) + if not parts: + continue + o = (pd.concat(parts, ignore_index=True)[["channel_name", "gene", "rank", "score", "experiment", "well", + "x_pheno", "y_pheno", "segmentation_id"]] + .rename(columns={"segmentation_id": "segmentation", "score": "pma_attention"})) + o["rank_type"] = "top" + o["predicted_class"] = o["gene"] # complex class_col for grain="complex" (top-cells crop uses `gene`) + p = f"{RANKDIR}/{slugify(ch)}.parquet"; o.to_parquet(p) + out[ch] = (p, sorted(o["gene"].unique())) + print(f"[complex-rank] {len(out)} channels; complexes/channel: {[len(v[1]) for v in out.values()][:5]}...") + return out + + +def register_complex(): + """Register grain='complex' targets for every fluor marker (from the `complexes` block of its index.json). + Complexes with a generated traversal (meta.json) get its asset_dir/alphas; the rest are top-cells-only.""" + V5 = f"{C.OUT}/{ASSETS}" + man = json.load(open(f"{V5}/manifest.json")) + total = full = 0 + for mk in man["markers"]: + mc = mk.get("marker_channel") + if not mc or re.match(r"(?i)phase", mc): + continue + mod = slugify(mc) + tci = f"{V5}/top_cells/markers/{mod}/index.json" + cx = sorted(json.load(open(tci)).get("complexes", {})) if os.path.exists(tci) else [] + keep = [t for t in mk["targets"] if t["grain"] != "complex"] # keep geneKO/PC; rebuild complex + cxt = [] + for c in cx: + mp = f"{V5}/{mod}/complex/{slugify(c)}/meta.json" + if os.path.exists(mp): # traversal generated → full target + m = json.load(open(mp)) + cxt.append({"grain": "complex", "target": c, "slug": c, "control": None, "has_real": m.get("has_real", True), + "real_dir": m.get("real_dir"), "n_cells": m.get("n_cells"), "asset_dir": m["asset_dir"], + "alphas": m["alphas"], "dist_map": None, "desc": ""}); full += 1 + else: # top-cells only + cxt.append({"grain": "complex", "target": c, "slug": c, "control": None, "has_real": False, + "real_dir": None, "n_cells": 0, "asset_dir": None, "alphas": [], "dist_map": None, "desc": ""}) + mk["targets"] = keep + cxt + total += len(cxt) + json.dump(man, open(f"{V5}/manifest.json", "w")) + print(f"[register] {total} fluor complex targets ({full} with traversals) across markers") + + +N_CELLS = 40 # NTC accuracy anchor cells (matches the geneKO fleet) + + +def gen_complex_shard(mc, d, ch, targets, parq, force=False): + """One marker's complex traversals: NTC-anchored, complex-KD direction, fluor complex SetTransformer scoring.""" + os.environ["OPS_DIFFEX_ASSETS"] = ASSETS + from . import precompute as P + P._ASSETS = ASSETS + P.precompute_marker(grain="complex", targets=targets, ckpt=f"{C.DD}/{d}/diffae_best.pt", out_root=C.OUT, + marker_channel=mc, channel=ch, control="NTC", n_cells=N_CELLS, + fluor_rank_parquet=parq, v5_score=True, load_workers=12, force=force) + + +def launch_traversals(): + """Fan out one GPU shard per marker over its complex ranking parquet (built by build_rankings).""" + from ops_utils.hpc.slurm_batch_utils import submit_parallel_jobs + os.environ["OPS_DIFFEX_ASSETS"] = ASSETS + cm = {mc: (d, ch) for d, mc, ch in C.complete_markers()} + jobs = [] + for mc, (d, ch) in cm.items(): + p = f"{RANKDIR}/{slugify(mc)}.parquet" + if not os.path.exists(p): + continue + cxs = sorted(pd.read_parquet(p, columns=["gene"]).gene.unique()) + jobs.append({"name": f"fluorcxv5_{slugify(mc)[:14]}", "func": gen_complex_shard, + "kwargs": {"mc": mc, "d": d, "ch": ch, "targets": cxs, "parq": p}}) + print(f"[fluor-complex-v5] {len(jobs)} marker traversal shards") + submit_parallel_jobs(jobs, experiment="diffex_fluor_cx_v5", + slurm_params={"slurm_partition": "gpu", "slurm_gres": "gpu:1", "cpus_per_task": 12, + "mem_gb": 96, "timeout_min": 600}, + log_dir="diffex_fluor_cx_v5", wait_for_completion=False) + + +def main(): + from ops_utils.hpc.slurm_batch_utils import submit_parallel_jobs + build_rankings() + cm = {mc: (d, ch) for d, mc, ch in C.complete_markers()} + jobs = [] + for mc, (d, ch) in cm.items(): + if os.path.exists(f"{RANKDIR}/{slugify(mc)}.parquet"): + jobs.append({"name": f"ftcx_{slugify(mc)[:20]}", "func": crop_marker_shard, + "kwargs": {"mc": mc, "ch": ch, "rankdir": RANKDIR, "block": "complexes"}}) + print(f"[fluor-complex-topcells] {len(jobs)} marker crop jobs") + submit_parallel_jobs(jobs, experiment="diffex_fluor_cx_topcells", + slurm_params={"slurm_partition": "cpu", "cpus_per_task": 8, "mem_gb": 32, "timeout_min": 150}, + log_dir="diffex_fluor_cx_topcells", wait_for_completion=False) + + +if __name__ == "__main__": + main() diff --git a/src/ops_model/models/attention/diffex/viewer/_fluor_topcells.py b/src/ops_model/models/attention/diffex/viewer/_fluor_topcells.py new file mode 100644 index 0000000..464beea --- /dev/null +++ b/src/ops_model/models/attention/diffex/viewer/_fluor_topcells.py @@ -0,0 +1,114 @@ +"""Per-marker fluorescence Top Cells (top-30 by v5 accuracy) — crops the actual MARKER channel +(reuses materialize_crops, name-based out_channels=[cfg.channel], exactly like the fluor traversals), +PLUS a tiny transparent-inside / blue-outside seg overlay per cell so the viewer can toggle the mask +client-side (no doubled crop cache). + +Output: viewer_assets_v5/top_cells/markers//{crops/*.webp, overlays/*.png, index.json} +""" +import os, json +import numpy as np +import pandas as pd +import zarr +from PIL import Image +from . import catalog as C +from ..classifier.config import slugify, GRAINS +from ..classifier.data import make_labels_df, materialize_crops +from ..directions.config import DirConfig +from ..diffae.data import normalize +from .precompute import _save_webp +from .build_pc_crops_masked import BASE, CROP_SIZE, MASK_DILATION, OVERLAY_RGB, OVERLAY_ALPHA, _crop, _zarr_patch + +ASSETS = "viewer_assets_v5" +RANKDIR = f"{C.OUT}/{ASSETS}/_rankings/fluor_shap/geneKO" # NEW shap_screen per-channel rankings (same cells as traversals) +OUT = f"{C.OUT}/{ASSETS}/top_cells/markers" +TOP_N = 40 + + +def _overlay_rgba(seg, half): + """Transparent inside the (dilated) center cell, blue+alpha outside → RGBA uint8 (the toggleable mask layer).""" + from scipy.ndimage import binary_dilation + center = seg[half, half] + if center == 0: + c = seg[half - 12:half + 12, half - 12:half + 12]; nz = c[c > 0] + center = np.bincount(nz).argmax() if nz.size else 0 + rgba = np.zeros((*seg.shape, 4), np.uint8) + if center != 0: + inv = ~binary_dilation(seg == center, iterations=MASK_DILATION) + rgba[inv, 0], rgba[inv, 1], rgba[inv, 2] = [int(v * 255) for v in OVERLAY_RGB] + rgba[inv, 3] = int(OVERLAY_ALPHA * 255) + return rgba + + +def crop_marker_shard(mc, ch, top_n=TOP_N, rankdir=RANKDIR, block="genes"): + """Crop top-N accuracy cells per class for one marker channel → per-marker crops/ + overlays/ + index.json. + block="genes" (geneKO) or "complexes" (EBI); complex crops MERGE into the marker's existing index.json.""" + _zarr_patch() + df = pd.read_parquet(f"{rankdir}/{slugify(mc)}.parquet") + df = df[df["rank"] <= top_n] # include NTC (its top-N controls) — pinned by default in the tab + recs = df.rename(columns={"gene": "cls", "pma_attention": "score"}).copy() + recs["label"] = 0 + cfg = DirConfig(grain="geneKO", target=recs["cls"].iloc[0], device="cpu") + cfg.marker_channel = mc; cfg.channel = ch; cfg.num_workers = 8 + raw, _, exps = materialize_crops(make_labels_df(recs, cfg), cfg, cache_path=None) # marker channel (raw intensity); drops failed-store experiments + recs = recs[recs["experiment"].isin(set(exps))].reset_index(drop=True) # realign to surviving cells (materialize drops whole failed experiments, in order) + pc = normalize(raw) # per-cell z-score (current default) + lo, hi = np.percentile(raw, (1, 99)) # marker-global intensity window (over ALL this marker's cells) + if hi - lo < 1e-6: + hi = lo + 1 + out = f"{OUT}/{slugify(mc)}"; cdir, ndir, odir = f"{out}/crops", f"{out}/crops_norm", f"{out}/overlays" + for dd in (cdir, ndir, odir): + os.makedirs(dd, exist_ok=True) + half = CROP_SIZE // 2 + segcache, genes = {}, {} + for i, r in recs.iterrows(): + if i >= len(raw): + break + key = f"{r['experiment']}_{r['well']}_{int(round(r['x_pheno']))}_{int(round(r['y_pheno']))}".replace("/", "-") + _save_webp(f"{cdir}/{key}.webp", pc[i, 0], 256) # per-cell normalized + gnorm = np.clip((raw[i, 0] - lo) / (hi - lo), 0, 1) # marker-global normalized (intensity comparable) + Image.fromarray((gnorm * 255).astype(np.uint8)).resize((256, 256), Image.BILINEAR).save(f"{ndir}/{key}.webp") + # seg overlay (same store/coords → aligned with the marker crop) + ek = (r["experiment"], r["well"]) + if ek not in segcache: + pos = f"{BASE}/{r['experiment']}/3-assembly/phenotyping_v3.zarr/{str(r['well'])[0]}/{str(r['well'])[1:]}/0" + try: + segcache[ek] = zarr.open(f"{pos}/labels/cell_seg/0", mode="r") + except Exception: + segcache[ek] = None + if segcache[ek] is not None: + try: + seg = _crop(segcache[ek], None, int(round(r["x_pheno"])), int(round(r["y_pheno"])), half) + ov = Image.fromarray(_overlay_rgba(seg, half)).resize((256, 256), Image.NEAREST) + ov.save(f"{odir}/{key}.png") + except Exception: + pass + genes.setdefault(str(r["cls"]), []).append( + {"img": f"{key}.webp", "ov": f"{key}.png", "exp": r["experiment"], "well": str(r["well"]), + "x": int(round(r["x_pheno"])), "y": int(round(r["y_pheno"])), + "rank": int(r["rank"]), "conf": round(float(r.get("score", 0) or 0), 5)}) + for g in genes: + genes[g] = sorted(genes[g], key=lambda c: c["rank"])[:top_n] + ipath = f"{out}/index.json" + prev = json.load(open(ipath)) if os.path.exists(ipath) else {} + idx = {k: prev[k] for k in ("genes", "complexes", "top_n") if k in prev} # merge; drop any stray keys + idx[block] = {g: {"attention": [], "accuracy": genes[g]} for g in sorted(genes)} + idx["top_n"] = top_n + json.dump(idx, open(ipath, "w")) + return {"marker": mc, block: len(genes), "cells": sum(len(v) for v in genes.values())} + + +def main(): + from ops_utils.hpc.slurm_batch_utils import submit_parallel_jobs + cm = {mc: (d, ch) for d, mc, ch in C.complete_markers()} + jobs = [] + for mc, (d, ch) in cm.items(): + if os.path.exists(f"{RANKDIR}/{slugify(mc)}.parquet"): + jobs.append({"name": f"ftc_{slugify(mc)[:20]}", "func": crop_marker_shard, "kwargs": {"mc": mc, "ch": ch}}) + print(f"[fluor-topcells] {len(jobs)} marker crop jobs (top {TOP_N} + overlays)") + submit_parallel_jobs(jobs, experiment="diffex_fluor_topcells", + slurm_params={"slurm_partition": "cpu", "cpus_per_task": 8, "mem_gb": 32, "timeout_min": 150}, + log_dir="diffex_fluor_topcells", wait_for_completion=False) + + +if __name__ == "__main__": + main() diff --git a/src/ops_model/models/attention/diffex/viewer/_fluor_v5_build.py b/src/ops_model/models/attention/diffex/viewer/_fluor_v5_build.py new file mode 100644 index 0000000..257ec84 --- /dev/null +++ b/src/ops_model/models/attention/diffex/viewer/_fluor_v5_build.py @@ -0,0 +1,160 @@ +"""Build v5 fluorescence geneKO NTC traversals for all 55 markers: + - anchor = top-40 accuracy NTC cells for that channel (from Alex's v5 gene_marker_1K_qualifying) + - direction = v5-accuracy KD centroid (same ranking), per-marker DiffAE checkpoint + - inline v5 FLUOR SetTransformer scoring (P(target)+rank) via the fixed modality-aware v5ctx +Complexes are a separate follow-up (per-gene table needs gene->complex pooling). geneKO only here. +""" +import os, re, json +import pandas as pd +from . import catalog as C +from ..classifier.config import slugify + +ASSETS = "viewer_assets_v5" +F = "/hpc/projects/icd.fast.ops/models/alex_lin_attention/v5/fluorescence" +# Alex Lin's celldino rankings (src /bio/projects/katamari/alex.lin/paper_celldino_rankings_v2/fluorescence): +# - accuracies exist for ALL 1000 gene × 55 marker combos (fluor_bychannel_*_pergene.csv, all 1001 genes). +# - per-CELL rankings (coords below) only generated where top1_acc > 0.5 @ 100 cells -> this is the ONLY +# reason a marker has 35-423 geneKO (not 1000): a DISTINCTIVENESS filter, NOT missing cells (the 65M-cell +# screen has thousands/class·marker). Alex can generate rankings for lower-acc combos on request. +GENE1K = f"{F}/misc/gene_marker_1K_qualifying.compact.parquet" +CP_CSV = f"{F}/misc/gene_marker_1K_CP.csv" # 7 Cell-Painting markers (TOMM20, Tubulin, ...) — authoritative ranking +RANKDIR = f"{C.OUT}/{ASSETS}/_rankings/fluor/geneKO" +N_CELLS = 40 +_COLS = ["channel_name", "gene", "rank", "score", "experiment", "well", "x_pheno", "y_pheno", "segmentation_id"] + + +def build_rankings(): + """Per-channel geneKO ranking parquets in the _fluor_rows schema (incl. NTC). → {channel: (parquet, [genes])}. + Cell-Painting markers (CP_CSV) override the qualifying-set ranking for their channels; 4i markers stay in qualifying.""" + os.makedirs(RANKDIR, exist_ok=True) + q = pd.read_parquet(GENE1K, columns=_COLS) + cp = pd.read_csv(CP_CSV, usecols=_COLS) + cp_ch = set(cp["channel_name"].unique()) + df = pd.concat([q[~q["channel_name"].isin(cp_ch)], cp], ignore_index=True) # CP CSV wins for its 7 channels + df = df.rename(columns={"segmentation_id": "segmentation", "score": "pma_attention"}) + df["rank_type"] = "top" + out = {} + for ch, g in df.groupby("channel_name"): + p = f"{RANKDIR}/{slugify(ch)}.parquet" + g.to_parquet(p) + genes = sorted(x for x in g["gene"].unique() if not str(x).startswith("NTC")) + out[ch] = (p, genes) + return out + + +def gen_marker_shard(mc, d, ch, targets, parq, force=False): + os.environ["OPS_DIFFEX_ASSETS"] = ASSETS + from . import precompute as P + P._ASSETS = ASSETS + P.precompute_marker(grain="geneKO", targets=targets, ckpt=f"{C.DD}/{d}/diffae_best.pt", out_root=C.OUT, + marker_channel=mc, channel=ch, control="NTC", n_cells=N_CELLS, + fluor_rank_parquet=parq, v5_score=True, load_workers=12, force=force) + + +CP_MARKERS = ["Endoplasmic Reticulum_Concanavalin A", "F-actin_Phalloidin", "Microtubules_Tubulin", + "Mitochondria_TOMM20", "Nucleoli_NPM1", "Nucleus_Hoechst", "Plasma Membrane_Wheat Germ Agglutinin"] + + +def regen_cp_markers(): + """Force-regenerate the 7 Cell-Painting markers' geneKO traversals with the CP-CSV ranking (their original + jobs read the qualifying parquets before the CP fix). Submits force=True jobs.""" + from ops_utils.hpc.slurm_batch_utils import submit_parallel_jobs + os.environ["OPS_DIFFEX_ASSETS"] = ASSETS + rk = build_rankings() + cm = {mc: (d, ch) for d, mc, ch in C.complete_markers()} + jobs = [] + for mc in CP_MARKERS: + if mc in rk and mc in cm: + d, ch = cm[mc]; parq, genes = rk[mc] + jobs.append({"name": f"fluorv5cp_{slugify(mc)[:16]}", "func": gen_marker_shard, + "kwargs": {"mc": mc, "d": d, "ch": ch, "targets": genes, "parq": parq, "force": True}}) + print(f"[fluor-v5-cp] force-regen {len(jobs)} CP markers") + submit_parallel_jobs(jobs, experiment="diffex_fluor_v5", + slurm_params={"slurm_partition": "gpu", "slurm_gres": "gpu:1", + "cpus_per_task": 12, "mem_gb": 96, "timeout_min": 600}, + log_dir="diffex_fluor_v5", wait_for_completion=False) + + +def main(): + import re # noqa + from ops_utils.hpc.slurm_batch_utils import submit_parallel_jobs + + +def register_manifest(): + """Register every fluor marker's FULL qualifying geneKO set (from its top-cells index) as targets, so the + perturbation dropdown / Top Cells is complete regardless of traversal-fleet progress. Genes with a generated + traversal get its asset_dir/alphas; the rest are top-cells-only (n_cells=0, empty Traversal). Idempotent.""" + V5 = f"{C.OUT}/{ASSETS}" + man = json.load(open(f"{V5}/manifest.json")) + acc = pd.read_csv(f"{F}/fluor_bychannel_paperv2gene_cps_pergene.csv", usecols=["channel_name", "gene_name", "top1_acc"]) + accmap = {(r.channel_name, r.gene_name): float(r.top1_acc) for r in acc.itertuples()} + desc = json.load(open(f"{V5}/gene_desc.json")) if os.path.exists(f"{V5}/gene_desc.json") else {} + total = 0 + for mk in man["markers"]: + mc = mk.get("marker_channel") + if not mc or re.match(r"(?i)phase", mc): + continue + mod = slugify(mc) + tci = f"{V5}/top_cells/markers/{mod}/index.json" + genes = sorted(g for g in json.load(open(tci))["genes"] if g != "NTC") if os.path.exists(tci) else [] # viewer adds NTC itself + keep = [t for t in mk["targets"] if t["grain"] != "geneKO"] # keep PC; rebuild geneKO + gk = [] + for g in genes: + mp = f"{V5}/{mod}/geneKO/{g}/meta.json" + if os.path.exists(mp): # traversal generated → full target + m = json.load(open(mp)) + gk.append({"grain": "geneKO", "target": g, "slug": g, "control": None, "has_real": m.get("has_real", True), + "real_dir": m.get("real_dir"), "n_cells": m.get("n_cells"), "asset_dir": m["asset_dir"], + "alphas": m["alphas"], "dist_map": accmap.get((mc, g)), "desc": desc.get(g, "")}) + else: # top-cells only (traversal not built yet) + gk.append({"grain": "geneKO", "target": g, "slug": g, "control": None, "has_real": False, + "real_dir": None, "n_cells": 0, "asset_dir": None, "alphas": [], + "dist_map": accmap.get((mc, g)), "desc": desc.get(g, "")}) + mk["targets"] = keep + gk + total += len(gk) + json.dump(man, open(f"{V5}/manifest.json", "w")) + print(f"[register] {total} fluor geneKO targets (full qualifying set) across {len(man['markers']) - 1} markers") + + +def build_real_acc20(): + """Extend real_acc20.json with fluor real-cell top1_acc@bag20, keyed by traversal asset_dir so the viewer's + real-ceiling reference lights up for the 55 markers (like phase). geneKO = per-(channel,gene); complex = mean of + member-gene acc grouped by Alex's EBI label_name. Idempotent (overwrites the fluor keys).""" + V5 = f"{C.OUT}/{ASSETS}" + ra = json.load(open(f"{V5}/real_acc20.json")) + gk = pd.read_csv(f"{F}/fluor_bychannel_paperv2gene_cps_pergene.csv", usecols=["channel_name", "n_cells", "gene_name", "top1_acc"]) + gk = gk[gk["n_cells"] == 20] + for r in gk.itertuples(): + ra[f"{slugify(r.channel_name)}/geneKO/{r.gene_name}"] = float(r.top1_acc) + cx = pd.read_csv(f"{F}/fluor_ebi_bychannel_pergene.csv", usecols=["channel_name", "n_cells", "label_name", "top1_acc"]) + cx = cx[cx["n_cells"] == 20] + nc = 0 + for (ch, lbl), g in cx.groupby(["channel_name", "label_name"]): + ra[f"{slugify(ch)}/complex/{lbl}"] = float(g["top1_acc"].mean()); nc += 1 + json.dump(ra, open(f"{V5}/real_acc20.json", "w")) + print(f"[real_acc20] +{len(gk)} fluor geneKO, +{nc} fluor complex keys (total {len(ra)})") + + +def main(): + import re # noqa + from ops_utils.hpc.slurm_batch_utils import submit_parallel_jobs + os.environ["OPS_DIFFEX_ASSETS"] = ASSETS + rk = build_rankings() + cm = {mc: (d, ch) for d, mc, ch in C.complete_markers()} # channel -> (diffae dir, raw channel) + jobs, skipped = [], [] + for ch_name, (parq, genes) in rk.items(): + if ch_name not in cm: + skipped.append(ch_name); continue + d, rawch = cm[ch_name] + jobs.append({"name": f"fluorv5_{slugify(ch_name)[:18]}", "func": gen_marker_shard, + "kwargs": {"mc": ch_name, "d": d, "ch": rawch, "targets": genes, "parq": parq}}) + print(f"[fluor-v5] {len(jobs)} marker jobs ({sum(len(g) for _, (p, g) in rk.items())} gene×channel pairs)") + print(f"[fluor-v5] skipped (no complete DiffAE / name mismatch): {skipped}") + submit_parallel_jobs(jobs, experiment="diffex_fluor_v5", + slurm_params={"slurm_partition": "gpu", "slurm_gres": "gpu:1", + "cpus_per_task": 12, "mem_gb": 96, "timeout_min": 600}, + log_dir="diffex_fluor_v5", wait_for_completion=False) + + +if __name__ == "__main__": + main() diff --git a/src/ops_model/models/attention/diffex/viewer/_migrate_v4_to_v5.py b/src/ops_model/models/attention/diffex/viewer/_migrate_v4_to_v5.py new file mode 100644 index 0000000..62fe01e --- /dev/null +++ b/src/ops_model/models/attention/diffex/viewer/_migrate_v4_to_v5.py @@ -0,0 +1,49 @@ +"""Migrate the v4 viewer components that v5 models will NOT regenerate into viewer_assets_v5, so they work +under the v5 cache: the PC viewer (pcs/) and the minibinder + PC-axis traversals (phase/minibinder, phase/pc). + +Left in v4 (still awaiting v5 model outputs): fluorescence marker traversals/montage and attention heads. +""" +import os, json, shutil, subprocess + +ROOT = "/hpc/projects/icd.fast.ops/models/diffex" +V4, V5 = f"{ROOT}/viewer_assets", f"{ROOT}/viewer_assets_v5" +DIRS = ["pcs", "phase/minibinder", "phase/pc"] # big trees → one SLURM rsync each +FILES = ["_minibinder_meta.json"] + + +def rsync_dir(rel): + src, dst = f"{V4}/{rel}/", f"{V5}/{rel}/" + os.makedirs(dst, exist_ok=True) + subprocess.run(["rsync", "-a", src, dst], check=True) + return f"{rel}: {subprocess.run(['du','-sh',dst],capture_output=True,text=True).stdout.split()[0]}" + + +def merge_manifest(): + """Append the v4 minibinder + pc targets (phase) into the v5 phase marker's target list (idempotent).""" + v4 = json.load(open(f"{V4}/manifest.json")) + v5 = json.load(open(f"{V5}/manifest.json")) + # collect v4 minibinder+pc targets (they live under the phase marker) + add = [t for mk in v4["markers"] for t in mk["targets"] if t["grain"] in ("minibinder", "pc")] + m5 = v5["markers"][0] # v5 is phase-only, single marker + have = {(t["grain"], t["asset_dir"]) for t in m5["targets"]} + new = [t for t in add if (t["grain"], t["asset_dir"]) not in have] + m5["targets"].extend(new) + json.dump(v5, open(f"{V5}/manifest.json", "w")) + print(f"[manifest] added {len(new)} targets (minibinder+pc); v5 marker now {len(m5['targets'])} targets") + + +def main(): + from ops_utils.hpc.slurm_batch_utils import submit_parallel_jobs + jobs = [{"name": f"migrate_{r.replace('/','_')}", "func": rsync_dir, "kwargs": {"rel": r}} for r in DIRS] + print(f"[migrate] copying {DIRS} v4 -> v5 via {len(jobs)} rsync jobs") + submit_parallel_jobs( + jobs, experiment="diffex_migrate_v5", + slurm_params={"slurm_partition": "cpu", "cpus_per_task": 4, "mem_gb": 16, "timeout_min": 120}, + log_dir="diffex_migrate_v5", wait_for_completion=True) + for f in FILES: + shutil.copy(f"{V4}/{f}", f"{V5}/{f}"); print(f"[file] {f}") + merge_manifest() + + +if __name__ == "__main__": + main() diff --git a/src/ops_model/models/attention/diffex/viewer/_phase_vs.py b/src/ops_model/models/attention/diffex/viewer/_phase_vs.py new file mode 100644 index 0000000..446979c --- /dev/null +++ b/src/ops_model/models/attention/diffex/viewer/_phase_vs.py @@ -0,0 +1,427 @@ +"""Combine phase counterfactual traversals with the multi-marker virtual-staining (VS) system. + +Pipeline: a real phase cell's traversal frame at α (already built, viewer_assets_v5/phase/geneKO//cell/ +frame_.webp) → CellDINO embed → multi-marker VS model → the SAME synthesized cell rendered in every one of +the 42 live markers. So one phase cell yields, per geneKO and α, the phase phenotype + all 42 marker phenotypes. + +Stages: + proto() — sanity montage: a few geneKOs × 42 markers at α=5 (tests VS on GENERATED phase). + stain_shard()— stain a chunk of geneKOs (all 42 markers) at (cell, α) → per-(marker,gene) webp on disk. + submit_stain()— shard the ~1000 geneKOs across GPUs. +Then render_montage_scales.render_composed sources these stained tiles → one composed montage per channel. +""" +import glob +import json +import os + +import numpy as np +import torch +from PIL import Image + +from ..classifier.config import slugify + +VS_OUT = "/hpc/projects/icd.fast.ops/analysis/virtual_staining/multi_marker" +V5 = "/hpc/projects/icd.fast.ops/models/diffex/viewer_assets_v5" +STAINED = "/hpc/projects/icd.fast.ops/analysis/figure4_embedding/phase_vs_combine/stained" # /_c_a.webp +AI_OF = {a: i for i, a in enumerate([-5, -4, -3, -2.5, -2, -1.5, -1, -0.5, 0, 0.5, 1, 1.5, 2, 2.5, 3, 4, 5])} + + +def load_vs(dev): + from ..diffae.config import DiffAEConfig + from ..diffae.model import DiffAE + markers = json.load(open(f"{VS_OUT}/markers.json")) + cfg = DiffAEConfig(spatial_cond=True, n_markers=len(markers), device="cuda", epochs=1) + ema = DiffAE(cfg).to(dev).eval() + st = torch.load(f"{VS_OUT}/train_state.pt", map_location=dev) + ema.load_state_dict(st["ema"]) + return ema, markers, cfg, st.get("epoch") + + +def _load_phase(gene, cell, ai, H): + p = f"{V5}/phase/geneKO/{gene}/cell{cell}/frame_{ai:02d}.webp" + if not os.path.exists(p): + return None + im = Image.open(p).convert("L").resize((H, H)) + return (np.asarray(im, np.float32) / 255.0 * 2 - 1)[None, None] # (1,1,H,H) in [-1,1] + + +@torch.no_grad() +def stain(ema, markers, cfg, dev, phase_np, seed=0): + """phase_np (1,1,H,H) in [-1,1] → {marker_idx: pred (H,H)} for all markers (fixed xT seed).""" + from ..diffae.virtstain_multi import _sample_marker + from ..classifier.celldino_features import embed_crops + H = cfg.crop_size + emb = torch.as_tensor(embed_crops(phase_np, cfg), dtype=torch.float32, device=dev) + ci = torch.as_tensor(phase_np, dtype=torch.float32, device=dev) + out = {} + for mid in range(len(markers)): + g = torch.Generator(device=dev).manual_seed(seed) + xT = torch.randn(1, 1, H, H, generator=g, device=dev) + mk = torch.as_tensor([mid], dtype=torch.long, device=dev) + out[mid] = _sample_marker(ema, xT, emb, ci, mk, cfg, dev).cpu().numpy()[0, 0] + return out + + +def _save(path, arr): + os.makedirs(os.path.dirname(path), exist_ok=True) + Image.fromarray((np.clip((arr + 1) / 2, 0, 1) * 255).astype("uint8")).resize((256, 256)).save(path, quality=90, method=6) + + +def stain_shard(genes, cell=1, alphas=(1.0, 2.0, 3.0, 4.0, 5.0), cells=None): + """Stain a chunk of geneKOs into all 42 markers at (each cell, each α); write /_c_a.webp. + Model loaded once per shard; all requested cells×alphas stained per gene (skip-guard resumes). + `cells` (list) overrides `cell` to cover multiple anchor cells in one shard (e.g. multibag 200-209).""" + dev = torch.device("cuda") + ema, markers, cfg, ep = load_vs(dev) + last = slugify(markers[-1]); done = skip = 0 # gene×cell×α complete once the LAST marker webp exists + for c in (cells if cells is not None else [cell]): + for g in genes: + for a in alphas: + an = f"a{a:g}" + if os.path.exists(f"{STAINED}/{last}/{slugify(g)}_c{c}_{an}.webp"): + skip += 1; continue # resume: already stained + ph = _load_phase(g, c, AI_OF[a], cfg.crop_size) + if ph is None: + continue + preds = stain(ema, markers, cfg, dev, ph) + for mid, name in enumerate(markers): + _save(f"{STAINED}/{slugify(name)}/{slugify(g)}_c{c}_{an}.webp", preds[mid]) + done += 1 + print(f"[phasevs] stained {done} gene×cell×α ({skip} already done) × {len(markers)} markers (VS ep{ep}) -> {STAINED}") + return {"done": done, "skipped": skip, "markers": len(markers)} + + +def submit_stain(cell=1, alphas=(1.0, 2.0, 3.0, 4.0, 5.0), chunk=8, parallel=64): + """Shard the ~1000 geneKOs across GPUs; each shard stains its genes into all 42 markers at (cell, each α). + Resumable (skip-guard) → re-run to fill gaps. Small chunks + high parallelism for short wall time.""" + from . import catalog as C + from ops_utils.hpc.slurm_batch_utils import submit_parallel_jobs + genes = C.all_genes() + ch = lambda l, n: [l[i:i + n] for i in range(0, len(l), n)] + jobs = [{"name": f"pvs_c{cell}_{i}", "func": stain_shard, + "kwargs": {"genes": s, "cell": cell, "alphas": list(alphas)}} for i, s in enumerate(ch(genes, chunk))] + print(f"[phasevs] {len(genes)} geneKOs → {len(jobs)} stain shards (chunk {chunk}, parallel {parallel}, cell {cell}, α={list(alphas)})") + submit_parallel_jobs( + jobs_to_submit=jobs, experiment="diffex_phasevs", + slurm_params={"slurm_partition": "gpu", "gpus_per_node": 1, "cpus_per_task": 8, "mem_gb": 64, + "timeout_min": 90, "slurm_constraint": "[a100_80|h100|h200|6000_blackwell]", + "slurm_array_parallelism": parallel}, + log_dir="diffex_phasevs", wait_for_completion=False) + + +MULTIBAG_CELLS = list(range(200, 210)) # VS on the top-10 multi_bag anchor cells (disk 200-209 → display rank 1-10) + + +def submit_stain_multibag(cells=MULTIBAG_CELLS, alphas=(1.0, 2.0, 3.0, 4.0, 5.0), chunk=8, parallel=64): + """Stain the multibag anchor cells (200-209) → all 42 markers at each α. One shard per (cell, gene-chunk) + to keep shards short + parallelism high. Resumable (skip-guard). Also stains the α0 NTC anchor per cell.""" + from . import catalog as C + from ops_utils.hpc.slurm_batch_utils import submit_parallel_jobs + genes = C.all_genes() + ch = lambda l, n: [l[i:i + n] for i in range(0, len(l), n)] + jobs = [{"name": f"pvs_c{c}_{i}", "func": stain_shard, + "kwargs": {"genes": s, "cell": c, "alphas": list(alphas)}} + for c in cells for i, s in enumerate(ch(genes, chunk))] + jobs += [{"name": f"pvsntc_c{c}", "func": stain_ntc, "kwargs": {"cell": c}} for c in cells] + print(f"[phasevs] multibag: {len(cells)} cells × {len(genes)} genes → {len(jobs)} stain shards " + f"(chunk {chunk}, parallel {parallel}, α={list(alphas)}) + {len(cells)} NTC anchors") + submit_parallel_jobs( + jobs_to_submit=jobs, experiment="diffex_phasevs_mb", + slurm_params={"slurm_partition": "gpu", "gpus_per_node": 1, "cpus_per_task": 8, "mem_gb": 64, + "timeout_min": 90, "slurm_constraint": "[a100_80|h100|h200|6000_blackwell]", + "slurm_array_parallelism": parallel}, + log_dir="diffex_phasevs_mb", wait_for_completion=False) + + +COMPOSED = "/hpc/projects/icd.fast.ops/analysis/figure4_embedding/phase_vs_combine/composed" +STAINED_ALL = "/hpc/projects/icd.fast.ops/analysis/figure4_embedding/phase_vs_combine/stained_all" # RGB 42-marker merge + + +def marker_colors(n): + """n distinct hues (unique color per marker), full saturation/value → (n,3) RGB in [0,1].""" + import matplotlib.pyplot as plt + return plt.get_cmap("hsv")(np.linspace(0, 1, n, endpoint=False))[:, :3] + + +def compose_all_shard(genes, cell=1, alpha=5.0): + """Merge the 42 stained marker tiles per gene into ONE RGB image: unique hue per marker, fluorescence-style + false-color on BLACK. Each stain sits at a high flat baseline (~0.49), so a naive additive sum of 42 + channels clips to white everywhere — instead background-subtract + stretch each marker, then max/lighten + blend (each pixel = brightest marker's color). → stained_all/_c_a.webp (+ __NTC).""" + markers = json.load(open(f"{VS_OUT}/markers.json")) + slugs = [slugify(m) for m in markers] + cols = marker_colors(len(markers)) # (42,3) + a5 = f"a{alpha:g}"; done = 0 + for g in list(genes) + ["__NTC"]: + imgs, ok = [], True + for s in slugs: + p = f"{STAINED}/{s}/__NTC_c{cell}.webp" if g == "__NTC" else f"{STAINED}/{s}/{slugify(g)}_c{cell}_{a5}.webp" + if not os.path.exists(p): + ok = False; break + imgs.append(np.asarray(Image.open(p).convert("L"), np.float32) / 255.0) + if not ok: + continue + stack = np.stack(imgs, 0) # (M,H,W) in [0,1] + lo = np.percentile(stack, 70, axis=(1, 2), keepdims=True) # per-marker background + hi = np.percentile(stack, 99.5, axis=(1, 2), keepdims=True) + xs = np.clip((stack - lo) / np.clip(hi - lo, 1e-6, None), 0, 1) # isolate bright structures → black bg + rgb = (xs[..., None] * cols[:, None, None, :]).max(axis=0) # (H,W,3) max/lighten blend + out = f"{STAINED_ALL}/{'__NTC_c%d' % cell if g == '__NTC' else slugify(g) + '_c%d_%s' % (cell, a5)}.webp" + os.makedirs(os.path.dirname(out), exist_ok=True) + Image.fromarray((rgb * 255).astype("uint8")).resize((256, 256)).save(out, quality=90, method=6) + done += 1 + print(f"[phasevs] merged {done} all-marker RGB tiles → {STAINED_ALL}") + return {"done": done} + + +def submit_compose_all(cell=1, alphas=(1.0, 2.0, 3.0, 4.0, 5.0), chunk=60, cells=None): + from . import catalog as C + from ops_utils.hpc.slurm_batch_utils import submit_parallel_jobs + genes = C.all_genes() + ch = lambda l, n: [l[i:i + n] for i in range(0, len(l), n)] + cs = cells if cells is not None else [cell] + jobs = [{"name": f"pvsall_c{c}_a{a:g}_{i}", "func": compose_all_shard, "kwargs": {"genes": s, "cell": c, "alpha": a}} + for c in cs for a in alphas for i, s in enumerate(ch(genes, chunk))] + print(f"[phasevs] {len(genes)} genes × {len(cs)} cells × {len(alphas)} α → {len(jobs)} all-marker merge shards") + submit_parallel_jobs(jobs_to_submit=jobs, experiment="diffex_phasevs", + slurm_params={"slurm_partition": "cpu", "cpus_per_task": 4, "mem_gb": 32, "timeout_min": 40, + "slurm_array_parallelism": 32}, + log_dir="diffex_phasevs", wait_for_completion=False) + + +def render_channel(channel, cell=1, alpha=5.0, level=4): + """Render ONE composed montage for `channel` ("phase" or a marker slug) on the shared phase embedding, + matching the no-marks reference (phate, L4). Phase tiles from viewer_assets_v5; marker tiles from STAINED.""" + from .render_montage_scales import render_composed + render_composed([alpha], cell=cell, level=level, out_dir=f"{COMPOSED}/{channel}", marks=False, bg="white", + channel=channel, stained_dir=STAINED, assets="viewer_assets_v5") # white canvas; fluor tiles magma-composited + return {"channel": channel} + + +def submit_montages(cell=1, alpha=5.0, level=4): + """Fan out the phase + 42 marker composed montages (all on the phase embedding).""" + from ops_utils.hpc.slurm_batch_utils import submit_parallel_jobs + markers = json.load(open(f"{VS_OUT}/markers.json")) + chans = ["phase"] + [slugify(m) for m in markers] + jobs = [{"name": f"pvsmtg_{c[:12]}", "func": render_channel, + "kwargs": {"channel": c, "cell": cell, "alpha": alpha, "level": level}} for c in chans] + print(f"[phasevs] {len(jobs)} composed montages (phase + {len(markers)} markers) → {COMPOSED}/") + submit_parallel_jobs( + jobs_to_submit=jobs, experiment="diffex_phasevs", + slurm_params={"slurm_partition": "cpu", "cpus_per_task": 8, "mem_gb": 64, "timeout_min": 60, + "slurm_array_parallelism": 43}, + log_dir="diffex_phasevs", wait_for_completion=False) + + +def render_all(cell=1, alpha=5.0, level=4): + """Render the ALL-MARKERS composed figure (42-marker RGB merge) on the shared phase embedding, white bg.""" + from .render_montage_scales import render_composed + render_composed([alpha], cell=cell, level=level, out_dir=f"{COMPOSED}/__allmarkers__", marks=False, bg="black", + channel="__allmarkers__", stained_dir=STAINED_ALL, assets="viewer_assets_v5") + return {"channel": "__allmarkers__"} + + +def submit_all_figure(cell=1, alphas=(1.0, 2.0, 3.0, 4.0, 5.0), level=4): + """Fan out the all-markers composed figure per α.""" + from ops_utils.hpc.slurm_batch_utils import submit_parallel_jobs + jobs = [{"name": f"pvsallfig_a{a:g}", "func": render_all, "kwargs": {"cell": cell, "alpha": a, "level": level}} + for a in alphas] + print(f"[phasevs] {len(jobs)} all-marker composed figures → {COMPOSED}/__allmarkers__") + submit_parallel_jobs( + jobs_to_submit=jobs, experiment="diffex_phasevs", + slurm_params={"slurm_partition": "cpu", "cpus_per_task": 8, "mem_gb": 64, "timeout_min": 60, + "slurm_array_parallelism": 5}, + log_dir="diffex_phasevs", wait_for_completion=False) + + +def montage_vs_tiles(marker_slug, cell=1, alpha=5.0, embedding="phate", tile=256, ppu=5600): + """Interactive VS montage tiles for the viewer: place each geneKO's STAINED tile at its PHASE-embedding + coordinate → OME-zarr → PNG tile pyramid at viewer_assets_v5/_montage_vs/__cell_a_tiles/. + Same layout as the phase montage (NTC nodes use the stained α0 anchor).""" + import shutil + import anndata as ad + from latent_lens import MontageConfig, build_montage + from .build_umap_montage import _embed_coords, montage_to_tiles, ZARR_SCRATCH, OUT as MOUT + from .render_montage_scales import UMAP_H5AD + a5 = f"a{alpha:g}" + ann = ad.read_h5ad(UMAP_H5AD) + coords_all = _embed_coords(ann, embedding) + gc = {str(g): coords_all[i] for i, g in enumerate(ann.obs["perturbation"])} + genes, coords, srcs = [], [], [] + for g, xy in gc.items(): # real genes with a stained tile (same set as phase) + if str(g).startswith("NTC"): + continue + if os.path.exists(f"{STAINED}/{marker_slug}/{slugify(g)}_c{cell}_{a5}.webp"): + genes.append(g); coords.append(xy); srcs.append(slugify(g)) + for g, xy in gc.items(): # NTC nodes → stained α0 anchor + if str(g).startswith("NTC"): + genes.append(g); coords.append(xy); srcs.append("__NTC") + coords = np.asarray(coords, np.float32) + + def crops(i): + p = (f"{STAINED}/{marker_slug}/__NTC_c{cell}.webp" if srcs[i] == "__NTC" + else f"{STAINED}/{marker_slug}/{srcs[i]}_c{cell}_{a5}.webp") + return np.asarray(Image.open(p).convert("L")) + + os.makedirs(ZARR_SCRATCH, exist_ok=True) + oz = f"{ZARR_SCRATCH}/vs_{marker_slug}_{embedding}_c{cell}_{a5}.zarr" + build_montage(umap_coords=coords, crops=crops, categories=np.array(["m"] * len(genes)), + category_colors={"m": (1.0, 1.0, 1.0)}, output_path=oz, labels=np.array(genes), + config=MontageConfig(crop_size=tile, px_per_umap=ppu, border_width=max(4, tile // 40))) + tiles = f"{MOUT}/viewer_assets_v5/_montage_vs/{marker_slug}_{embedding}_cell{cell}_{a5}_tiles" + montage_to_tiles(oz, UMAP_H5AD, out_dir=tiles, placed=set(genes), embedding=embedding) + shutil.rmtree(oz, ignore_errors=True) + print(f"[phasevs] VS tiles {marker_slug} {embedding} → {tiles}") + return {"marker": marker_slug, "genes": len(genes)} + + +def submit_vs_tiles(cell=1, alphas=(1.0, 2.0, 3.0, 4.0, 5.0), cells=None): + """Build interactive VS montage tiles for all 42 markers × {umap, phate} × each cell × α → viewer_assets_v5/_montage_vs/.""" + from ops_utils.hpc.slurm_batch_utils import submit_parallel_jobs + markers = json.load(open(f"{VS_OUT}/markers.json")) + cs = cells if cells is not None else [cell] + jobs = [{"name": f"vstile_{slugify(m)[:10]}_{e[:2]}_c{c}_a{a:g}", "func": montage_vs_tiles, + "kwargs": {"marker_slug": slugify(m), "cell": c, "alpha": a, "embedding": e}} + for m in markers for c in cs for e in ("umap", "phate") for a in alphas] + print(f"[phasevs] {len(jobs)} VS montage-tile jobs (42 markers × {len(cs)} cells × 2 emb × {len(alphas)} α)") + submit_parallel_jobs( + jobs_to_submit=jobs, experiment="diffex_phasevs", + slurm_params={"slurm_partition": "cpu", "cpus_per_task": 6, "mem_gb": 48, "timeout_min": 60, + "slurm_array_parallelism": 48}, + log_dir="diffex_phasevs", wait_for_completion=False) + + +def montage_all_tiles(cell=1, alpha=5.0, embedding="phate", tile=256, ppu=5600): + """Interactive ALL-MARKERS montage tiles: place each geneKO's pre-composited 42-marker RGB tile at its + phase-embedding coord → OSD pyramid at _montage_vs/__allmarkers____cell_a_tiles/. + build_montage only tints grayscale by one color, so its two crop helpers are monkeypatched to pass RGB through.""" + import shutil + import anndata as ad + from latent_lens import MontageConfig, build_montage + from latent_lens import montage as _M + from .build_umap_montage import _embed_coords, montage_to_tiles, ZARR_SCRATCH, OUT as MOUT + from .render_montage_scales import UMAP_H5AD + an = f"a{alpha:g}" + ann = ad.read_h5ad(UMAP_H5AD) + coords_all = _embed_coords(ann, embedding) + gc = {str(g): coords_all[i] for i, g in enumerate(ann.obs["perturbation"])} + genes, coords, srcs = [], [], [] + for g, xy in gc.items(): # real genes with a composited all-marker tile + if str(g).startswith("NTC"): + continue + if os.path.exists(f"{STAINED_ALL}/{slugify(g)}_c{cell}_{an}.webp"): + genes.append(g); coords.append(xy); srcs.append(slugify(g)) + for g, xy in gc.items(): # NTC nodes → composited α0 anchor + if str(g).startswith("NTC"): + genes.append(g); coords.append(xy); srcs.append("__NTC") + coords = np.asarray(coords, np.float32) + + def crops(i): + p = (f"{STAINED_ALL}/__NTC_c{cell}.webp" if srcs[i] == "__NTC" + else f"{STAINED_ALL}/{srcs[i]}_c{cell}_{an}.webp") + return np.asarray(Image.open(p).convert("RGB"), np.float32) / 255.0 # (H,W,3) in [0,1] + + def _norm_rgb(crop, cs): # RGB-aware pad/normalize (grayscale falls back to orig) + if crop.ndim == 2: + return _norm_orig(crop, cs) + if crop.shape[0] != cs or crop.shape[1] != cs: + pad = np.zeros((cs, cs, 3), np.float32); h = min(crop.shape[0], cs); w = min(crop.shape[1], cs) + pad[:h, :w] = crop[:h, :w]; crop = pad + return crop.astype(np.float32) + + def _tint_rgb(crop, color): # already colored → (3,H,W), skip single-color tint + return np.transpose(crop, (2, 0, 1)) if crop.ndim == 3 else _tint_orig(crop, color) + + _norm_orig, _tint_orig = _M._normalize_crop, _M.tint_crop + _M._normalize_crop, _M.tint_crop = _norm_rgb, _tint_rgb + try: + os.makedirs(ZARR_SCRATCH, exist_ok=True) + oz = f"{ZARR_SCRATCH}/vsall_{embedding}_c{cell}_{an}.zarr" + build_montage(umap_coords=coords, crops=crops, categories=np.array(["m"] * len(genes)), + category_colors={"m": (1.0, 1.0, 1.0)}, output_path=oz, labels=np.array(genes), + config=MontageConfig(crop_size=tile, px_per_umap=ppu, border_width=max(4, tile // 40))) + finally: + _M._normalize_crop, _M.tint_crop = _norm_orig, _tint_orig + tiles = f"{MOUT}/viewer_assets_v5/_montage_vs/__allmarkers___{embedding}_cell{cell}_{an}_tiles" + montage_to_tiles(oz, UMAP_H5AD, out_dir=tiles, placed=set(genes), embedding=embedding) + shutil.rmtree(oz, ignore_errors=True) + print(f"[phasevs] ALL-marker tiles {embedding} a{alpha:g} → {tiles}") + return {"genes": len(genes)} + + +def submit_all_tiles(cell=1, alphas=(1.0, 2.0, 3.0, 4.0, 5.0), cells=None): + """Build interactive ALL-MARKERS montage tiles for {umap, phate} × each cell × α → viewer_assets_v5/_montage_vs/.""" + from ops_utils.hpc.slurm_batch_utils import submit_parallel_jobs + cs = cells if cells is not None else [cell] + jobs = [{"name": f"vsalltile_{e[:2]}_c{c}_a{a:g}", "func": montage_all_tiles, + "kwargs": {"cell": c, "alpha": a, "embedding": e}} + for c in cs for e in ("umap", "phate") for a in alphas] + print(f"[phasevs] {len(jobs)} all-marker montage-tile jobs ({len(cs)} cells × 2 emb × {len(alphas)} α)") + submit_parallel_jobs( + jobs_to_submit=jobs, experiment="diffex_phasevs", + slurm_params={"slurm_partition": "cpu", "cpus_per_task": 6, "mem_gb": 48, "timeout_min": 60, + "slurm_array_parallelism": 10}, + log_dir="diffex_phasevs", wait_for_completion=False) + + +def stain_ntc(cell=1): + """Stain the α=0 anchor (control) cell → all 42 markers, for the composed montage's NTC nodes. + Saved as /__NTC_c.webp (α=0 content, shared by every NTC grid node).""" + dev = torch.device("cuda") + ema, markers, cfg, ep = load_vs(dev) + a0 = AI_OF[0.0] + g0 = next(os.path.basename(os.path.dirname(os.path.dirname(p))) + for p in sorted(glob.glob(f"{V5}/phase/geneKO/*/cell{cell}/frame_{a0:02d}.webp"))) + preds = stain(ema, markers, cfg, dev, _load_phase(g0, cell, a0, cfg.crop_size)) + for mid, name in enumerate(markers): + _save(f"{STAINED}/{slugify(name)}/__NTC_c{cell}.webp", preds[mid]) + print(f"[phasevs] stained NTC anchor (α0 from {g0}) → {len(markers)} markers") + return {"markers": len(markers)} + + +def proto(): + """Sanity montage: a few geneKOs × 42 markers at α=5. Writes phase_vs_combine/proto.png.""" + dev = torch.device("cuda") + ema, markers, cfg, ep = load_vs(dev); H = cfg.crop_size + cand = ["TP53", "KRAS", "TUBB", "ACTB", "POLR1B", "KIF23", "MYC", "CTNNB1"] + genes = [g for g in cand if os.path.exists(f"{V5}/phase/geneKO/{g}/cell1/frame_16.webp")] + if len(genes) < 4: + genes = [os.path.basename(os.path.dirname(os.path.dirname(p))) + for p in sorted(glob.glob(f"{V5}/phase/geneKO/*/cell1/frame_16.webp"))[:8]] + rows = [(g, _load_phase(g, 1, 16, H)[0, 0], stain(ema, markers, cfg, dev, _load_phase(g, 1, 16, H))) for g in genes] + import matplotlib + matplotlib.use("Agg"); matplotlib.rcParams["pdf.fonttype"] = 42 + import matplotlib.pyplot as plt + ncol = 1 + len(markers) + fig, ax = plt.subplots(len(rows), ncol, figsize=(ncol * 0.85, len(rows) * 0.95), squeeze=False) + for r_ in range(len(rows)): + for c_ in range(ncol): + ax[r_][c_].set_xticks([]); ax[r_][c_].set_yticks([]) + for r, (g, ph, preds) in enumerate(rows): + ax[r][0].imshow(ph, cmap="gray", vmin=-1, vmax=1, aspect="auto"); ax[r][0].set_ylabel(g, fontsize=7) + for m in range(len(markers)): + ax[r][1 + m].imshow(preds[m], cmap="magma", vmin=-1, vmax=1, aspect="auto") + ax[0][0].set_title("phase α5", fontsize=6) + for m, name in enumerate(markers): + ax[0][1 + m].set_title(name.split("_")[0][:10], fontsize=5, rotation=90, va="bottom") + fig.suptitle(f"Phase traversal (α=5) virtually stained → all {len(markers)} markers · VS ep{ep} · {len(rows)} geneKOs", fontsize=10) + fig.subplots_adjust(wspace=0.03, hspace=0.05, top=0.9) + out = os.path.dirname(STAINED); os.makedirs(out, exist_ok=True) + fig.savefig(f"{out}/proto.png", dpi=150, bbox_inches="tight"); plt.close(fig) + print(f"[phasevs] proto -> {out}/proto.png ({genes})") + return {"genes": genes} + + +if __name__ == "__main__": + import sys + cmd = sys.argv[1] if len(sys.argv) > 1 else "proto" + if cmd == "stain_multibag": + submit_stain_multibag() + elif cmd == "compose_multibag": # 42-marker RGB overlay for cells 200-209 (α5) → STAINED_ALL + submit_compose_all(cells=MULTIBAG_CELLS, alphas=(5.0,)) + elif cmd == "vstiles_multibag": # interactive per-marker montage tiles (reads STAINED, done) + submit_vs_tiles(cells=MULTIBAG_CELLS, alphas=(5.0,)) + elif cmd == "alltiles_multibag": # interactive all-markers overlay tiles (reads STAINED_ALL → run AFTER compose_multibag) + submit_all_tiles(cells=MULTIBAG_CELLS, alphas=(5.0,)) + else: + proto() diff --git a/src/ops_model/models/attention/diffex/viewer/_rebuild_v5.py b/src/ops_model/models/attention/diffex/viewer/_rebuild_v5.py new file mode 100644 index 0000000..4353051 --- /dev/null +++ b/src/ops_model/models/attention/diffex/viewer/_rebuild_v5.py @@ -0,0 +1,59 @@ +"""Rebuild v5 phase traversals with the CORRECTED anchor: v4 typical NTC anchor cells (pre-copied into +the shared _anchors/NTC cache, so the control gather is skipped) + v5 KD centroids (accuracy_parquet), +scoring the v5 SetTransformer inline (reuses gemb → no separate re-decode/re-embed pass). force=True +recomputes directions (v5_KD − v4_NTC) and frames. Removable after the rebuild.""" +from . import catalog as C +from .precompute import precompute_marker + +V5G = "/hpc/projects/icd.fast.ops/models/diffex/viewer_assets_v5/_rankings/pma_v5_phase_geneKO.parquet" +V5C = "/hpc/projects/icd.fast.ops/models/diffex/viewer_assets_v5/_rankings/pma_v5_phase_complex.parquet" +PHASE_CK = f"{C.DD}/phase_v1/diffae_best.pt" + + +def rebuild_shard(grain, targets): + parq = V5G if grain == "geneKO" else V5C + return precompute_marker(grain=grain, targets=list(targets), ckpt=PHASE_CK, out_root=C.OUT, + control="NTC", accuracy_parquet=parq, v5_score=True, force=True) + + +def accanchor_shard(grain, targets): + """Reproduce the ORIGINAL v5 (accuracy-selected NTC anchor + v5 KD): run with OPS_DIFFEX_V5=1 so GRAINS + parquet = v5 for BOTH control and KD (no accuracy_parquet override), isolated in viewer_assets_v5_accanchor.""" + return precompute_marker(grain=grain, targets=list(targets), ckpt=PHASE_CK, out_root=C.OUT, + control="NTC", v5_score=True, force=True) + + +SEL25 = "/hpc/projects/icd.fast.ops/analysis/figure4_traversals/ntc_accanchor_selected25.csv" + + +def build_accpool_anchor(): + """Accuracy-pool NTC anchor = the 25 hand-picked quality cells (z0 + control). Materialize+embed them and + write ctrl.npz + cell0-24 real.webp into OPS_DIFFEX_ASSETS/phase/_anchors/NTC/ (viewer_assets_v5_accpool).""" + import numpy as np, pandas as pd + from pathlib import Path + from concurrent.futures import ThreadPoolExecutor + from .precompute import _gather_class, _ASSETS, _save_webp + from ..diffae.data import normalize + from ..directions.config import DirConfig + sel = pd.read_csv(SEL25) + parq = pd.DataFrame({"gene": "NTC", "experiment": sel.experiment, "well": sel.well, "segmentation": sel.segmentation, + "x_pheno": sel.x_pheno, "y_pheno": sel.y_pheno, "pma_attention": sel.pma_attention, + "rank": range(1, len(sel) + 1), "rank_type": "top"}) + ptmp = f"{C.OUT}/{_ASSETS}/_ntc25.parquet"; Path(ptmp).parent.mkdir(parents=True, exist_ok=True); parq.to_parquet(ptmp) + cfg = DirConfig(grain="geneKO", target="NTC", device="cuda") + imgs, embs = _gather_class(cfg, "NTC", 25, parquet=ptmp) + realdir = Path(C.OUT) / _ASSETS / "phase" / "_anchors" / "NTC"; realdir.mkdir(parents=True, exist_ok=True) + real = normalize(imgs); tp = ThreadPoolExecutor(8) + for c in range(len(real)): + (realdir / f"cell{c}").mkdir(parents=True, exist_ok=True) + tp.submit(_save_webp, realdir / f"cell{c}" / "real.webp", real[c, 0], 256) + tp.shutdown(wait=True) + np.savez(realdir / "ctrl.npz", ctrl_embs=embs, mu_ctrl=embs.mean(0)) + print(f"accpool anchor built: {embs.shape} -> {realdir}") + + +def accpool_shard(grain, targets): + """Generate accuracy-pool traversals: pre-built 25 hand-picked anchor + v5 direction (OPS_DIFFEX_V5=1), + the SAME ±5 α grid, n_cells=25, set-acc over a fixed 20-cell bag (v5_bag=20) for cross-approach comparison.""" + return precompute_marker(grain=grain, targets=list(targets), ckpt=PHASE_CK, out_root=C.OUT, + control="NTC", n_cells=25, v5_score=True, v5_bag=20, force=True) diff --git a/src/ops_model/models/attention/diffex/viewer/_rescore_rank.py b/src/ops_model/models/attention/diffex/viewer/_rescore_rank.py new file mode 100644 index 0000000..0d4603b --- /dev/null +++ b/src/ops_model/models/attention/diffex/viewer/_rescore_rank.py @@ -0,0 +1,66 @@ +"""Re-score all v5 traversals to add `rank_target` (1-indexed rank of the target class) to scores_v5.json. + +score_embs_v5 now emits rank_target alongside p_target/top1/top5; this pass re-runs the SetTransformer scorer +over both anchor pools so the viewer's "target rank" overlay has data. One pass fully subsumes the old top5 backfill. +""" +import os, glob, json + +BASE = "/hpc/projects/icd.fast.ops/models/diffex" +POOLS = ["viewer_assets_v5", "viewer_assets_v5_accpool"] + + +def _retarget(assets): + """Point the score module at `assets` and return the module (V5_BASE is read at call time).""" + import ops_model.models.attention.diffex.viewer.score_generated as SG + SG.V5_BASE = f"{BASE}/{assets}/phase" + return SG + + +def rescore_shard(assets, grain, targets, bag=20): + SG = _retarget(assets) + SG.score_targets(grain, targets, bag=bag) + + +def anchor_shard(assets, grain): + SG = _retarget(assets) + SG.score_anchor_traversals(grain) + + +def _targets(assets, grain): + """NTC-anchor target names (from each traversal's meta.json) for a pool/grain, excluding A→B dirs.""" + sub = "geneKO" if grain == "geneKO" else "complex" + out = [] + for d in sorted(glob.glob(f"{BASE}/{assets}/phase/{sub}/*")): + if not os.path.isdir(d) or "__to__" in os.path.basename(d): + continue + mp = f"{d}/meta.json" + if os.path.exists(mp): + out.append(json.load(open(mp))["target"]) + return out + + +def main(n_shards=24): + from ops_utils.hpc.slurm_batch_utils import submit_parallel_jobs + jobs = [] + for assets in POOLS: + for grain in ["geneKO", "complex"]: + names = _targets(assets, grain) + shards = [s for s in (names[i::n_shards] for i in range(n_shards)) if s] + for i, s in enumerate(shards): + jobs.append({"name": f"rank_{assets[-4:]}_{grain}_{i}", "func": rescore_shard, + "kwargs": {"assets": assets, "grain": grain, "targets": s, "bag": 20}}) + # A→B alt-anchor dirs live only in the attention pool + if glob.glob(f"{BASE}/{assets}/phase/geneKO/*__to__*"): + for grain in ["geneKO", "complex"]: + jobs.append({"name": f"rank_{assets[-4:]}_alt_{grain}", "func": anchor_shard, + "kwargs": {"assets": assets, "grain": grain}}) + print(f"[rank-rescore] submitting {len(jobs)} jobs across {POOLS}") + submit_parallel_jobs( + jobs, experiment="diffex_rank_rescore", + slurm_params={"slurm_partition": "gpu", "slurm_gres": "gpu:1", + "cpus_per_task": 8, "mem_gb": 64, "timeout_min": 180}, + log_dir="diffex_rank_rescore", wait_for_completion=False) + + +if __name__ == "__main__": + main() diff --git a/src/ops_model/models/attention/diffex/viewer/_score_v4.py b/src/ops_model/models/attention/diffex/viewer/_score_v4.py new file mode 100644 index 0000000..4b12ed8 --- /dev/null +++ b/src/ops_model/models/attention/diffex/viewer/_score_v4.py @@ -0,0 +1,37 @@ +"""DIAGNOSTIC (removable): score the EXISTING v4 traversals (viewer_assets) with the v5 SetTransformer, +to compare v4 vs v5 phenotype accuracy — are the v4 (in-distribution) traversals better phenotypes? +Reads v4 frames directly (no writes into the live v4 assets); collects peak/α0/α5 P(target) per traversal.""" +import json +import os + +V4_BASE = "/hpc/projects/icd.fast.ops/models/diffex/viewer_assets/phase" + + +def score_v4_shard(grain, targets, out_json): + from .set_classifier import load_set_classifier, V5_CKPT_ROOT, V5_RUNS + from .score_generated import _emb_frames, score_embs_v5 + from ..directions.config import DirConfig + from ..classifier.celldino_features import embed_crops + from ..classifier.config import slugify + run = V5_RUNS[("phase", "geneKO" if grain == "geneKO" else "complex_ebionly")] + model, g2i, ci_map = load_set_classifier(run=run, device="cuda", root=V5_CKPT_ROOT) + ci = ci_map.get("Phase2D", 0) + sub = "geneKO" if grain == "geneKO" else "complex" + res = {} + for tgt in targets: + trav = f"{V4_BASE}/{sub}/{tgt if grain == 'geneKO' else slugify(tgt)}" + mp = f"{trav}/meta.json" + if not os.path.exists(mp): + continue + alphas = json.load(open(mp))["alphas"] + cfg = DirConfig(grain=grain, target=(tgt if grain == "geneKO" else "NTC"), device="cuda") + embs = [_emb_frames(cfg, trav, ai, embed_crops) for ai in range(len(alphas))] + d = score_embs_v5(embs, alphas, tgt, model, g2i, ci, run, "cuda") + if d: + p = d["p_target"]; z0 = len(alphas) // 2 + fin = [v for v in p if v is not None] + res[tgt] = {"alphas": d["alphas"], "p_target": p, # full curve for the accuracy-vs-α plot + "peak": max(fin) if fin else None, "a0": p[z0], "a5": p[-1]} + os.makedirs(os.path.dirname(out_json), exist_ok=True) + json.dump(res, open(out_json, "w")) + return {"grain": grain, "n": len(res), "out": out_json} diff --git a/src/ops_model/models/attention/diffex/viewer/_v4acc_test.py b/src/ops_model/models/attention/diffex/viewer/_v4acc_test.py new file mode 100644 index 0000000..e6ab223 --- /dev/null +++ b/src/ops_model/models/attention/diffex/viewer/_v4acc_test.py @@ -0,0 +1,17 @@ +"""DIAGNOSTIC (removable): generate the 40S→60S A→B traversal anchored on v4-ACCURACY-ranked complex +cells, to disentangle 'attention vs accuracy ranking' from 'new v5 cells'. Compared against the +existing v4-attention (viewer_assets) and v5-accuracy (viewer_assets_v5) 40S→60S. Monkeypatches the +complex parquet to a v4-accuracy ribosomal parquet; output isolated under viewer_assets_v4acc_test.""" +from . import catalog as C + +V4ACC = "/hpc/projects/icd.fast.ops/models/diffex/viewer_assets_v5/_rankings/pma_v4acc_phase_complex_ribo.parquet" +C40 = "40S cytosolic small ribosomal subunit" +C60 = "60S cytosolic large ribosomal subunit" + + +def run(): + from ..classifier.config import GRAINS + GRAINS["complex"]["parquet"] = V4ACC # both anchor (A) and target (B) from v4-accuracy cells + from .precompute import precompute_anchors_marker + return precompute_anchors_marker(grain="complex", classes=[C40, C60], + ckpt=f"{C.DD}/phase_v1/diffae_best.pt", out_root=C.OUT) diff --git a/src/ops_model/models/attention/diffex/viewer/_verify_pt_space.py b/src/ops_model/models/attention/diffex/viewer/_verify_pt_space.py new file mode 100644 index 0000000..3188184 --- /dev/null +++ b/src/ops_model/models/attention/diffex/viewer/_verify_pt_space.py @@ -0,0 +1,47 @@ +"""GPU check: does Alex's per-gene .pt CellDINO embedding == embed_crops (the DiffAE conditioning +space)? Embed top cells via the current gather, match to .pt by segmentation_id, compare. +Run as a one-off SLURM job; prints cosine + gap agreement. Delete after.""" +import torch, numpy as np +from .precompute import _gather_class +from ..directions.config import DirConfig +from ..directions.data import _top_cells +from ..classifier.config import GRAINS + +PT = "/hpc/projects/icd.fast.ops/models/alex_lin_attention/v4/train_ops_zstdcontrol_cdino_v2" + + +def run(): + cfg = DirConfig(grain="geneKO", target="KIF11", control="NTC", device="cuda") + cfg.num_workers = 12 + pq = GRAINS["geneKO"]["parquet"] + for g in ["NTC", "KIF11"]: + rows = _top_cells(pq, "gene", g, 40).reset_index(drop=True) + imgs, embs = _gather_class(cfg, g, 40) # fresh embed_crops + rows = rows.iloc[:len(embs)] + # COMPOSITE key (segmentation_id is only unique within an image) + key_fresh = list(zip(rows["experiment"].astype(str), rows["well"].astype(str), + rows["segmentation"].astype(np.int64))) + o = torch.load(f"{PT}/{g}.pt", map_location="cpu") + E = np.asarray(o["embeddings"], np.float32) + md = o["cell_metadata"] + exp_pt = [x for bag in md["experiment"] for x in bag] + well_pt = [x for bag in md["well"] for x in bag] + seg_pt = [x for bag in md["segmentation_id"] for x in bag] + pt_by_key = {(str(e), str(w), int(s)): E[i] for i, (e, w, s) in enumerate(zip(exp_pt, well_pt, seg_pt))} + cos, nr = [], [] + for e, k in zip(embs, key_fresh): + if k in pt_by_key: + p = pt_by_key[k] + cos.append(float(e @ p / (np.linalg.norm(e) * np.linalg.norm(p) + 1e-9))) + nr.append(np.linalg.norm(e) / (np.linalg.norm(p) + 1e-9)) + cos = np.array(cos) + print(f"[{g}] fresh embed_crops |row|={np.linalg.norm(embs,axis=1).mean():.2f} " + f".pt |row|={np.linalg.norm(E,axis=1).mean():.2f} " + f"matched {len(cos)} cos(fresh,.pt) mean={cos.mean():.4f} min={cos.min():.4f} " + f"norm-ratio={np.mean(nr):.3f}") + print("VERDICT: cos≈1.0 + norm-ratio≈1.0 → .pt IS embed_crops space (safe drop-in). " + "cos≈1 but ratio≠1 → same direction, rescale needed. cos<0.9 → different space, do NOT use.") + + +if __name__ == "__main__": + run() diff --git a/src/ops_model/models/attention/diffex/viewer/_verify_score_bridge.py b/src/ops_model/models/attention/diffex/viewer/_verify_score_bridge.py new file mode 100644 index 0000000..6ba459f --- /dev/null +++ b/src/ops_model/models/attention/diffex/viewer/_verify_score_bridge.py @@ -0,0 +1,26 @@ +"""Does embed_crops + z-standardize-on-control land in the SetTransformer's input space? +If yes, we can score generated cells: embed_crops(gen) → (x-μ_NTC)/σ_NTC → classifier bag. +Run as a GPU job; prints P(KIF11) for raw vs zstd-control embed_crops bags. Delete after.""" +import numpy as np + +from ..directions.config import DirConfig +from .precompute import _gather_class +from .set_classifier import load_set_classifier, score_bags + + +def run(): + cfg = DirConfig(grain="geneKO", target="KIF11", control="NTC", device="cuda") + cfg.num_workers = 12 + _, ntc = _gather_class(cfg, "NTC", 400) # embed_crops (per-image z-score) features + _, kif = _gather_class(cfg, "KIF11", 400) + mu, sd = ntc.mean(0), ntc.std(0) + 1e-6 + m, g2i, c2i = load_set_classifier("miwkg1cy", device="cuda") + ci, tgt = c2i["Phase2D"], g2i["KIF11"] + rng = np.random.default_rng(0) + print(f"embed_crops |row|={np.linalg.norm(kif, axis=1).mean():.1f}") + for name, feats in [("raw embed_crops", kif), ("zstd-on-control", (kif - mu) / sd)]: + bags = np.stack([feats[rng.choice(len(feats), 100)] for _ in range(5)]) + p = score_bags(m, bags, ci, device="cuda") + print(f" {name}: P(KIF11)={p[:, tgt].mean():.3f} top1=={int((p.argmax(1) == tgt).sum())}/5 " + f"argmax={[list(g2i)[i] for i in p.argmax(1)]}") + print("VERDICT: zstd-on-control P(KIF11) high → bridge works (embed_crops+zstd = classifier space).") diff --git a/src/ops_model/models/attention/diffex/viewer/altanchor_pairs.json b/src/ops_model/models/attention/diffex/viewer/altanchor_pairs.json new file mode 100644 index 0000000..a5e2ecb --- /dev/null +++ b/src/ops_model/models/attention/diffex/viewer/altanchor_pairs.json @@ -0,0 +1,326 @@ +{ + "geneKO": [ + [ + "ARPC2", + "ARPC4" + ], + [ + "ARPC4", + "CAPZB" + ], + [ + "ATP6V0A1", + "ATP6V1B2" + ], + [ + "AURKB", + "INCENP" + ], + [ + "CCT3", + "CCT7" + ], + [ + "COPA", + "COPB1" + ], + [ + "COX7B", + "COX7C" + ], + [ + "DDX18", + "DDX46" + ], + [ + "DYNC1H1", + "KIF11" + ], + [ + "DYNLL1", + "DYNLL2" + ], + [ + "EIF3A", + "EIF4A1" + ], + [ + "EIF4A1", + "EIF5B" + ], + [ + "HNRNPH1", + "HNRNPH2" + ], + [ + "HNRNPH1", + "HNRNPK" + ], + [ + "HSPA14", + "HSPA5" + ], + [ + "HSPA5", + "HSPA9" + ], + [ + "KIF11", + "KIF20A" + ], + [ + "KIF11", + "KIF23" + ], + [ + "KIF14", + "KIF23" + ], + [ + "KPNB1", + "XPO1" + ], + [ + "MCM3", + "MCM5" + ], + [ + "MCM3", + "MCM7" + ], + [ + "MRPL20", + "MRPL39" + ], + [ + "NUP54", + "NUP98" + ], + [ + "POLR1B", + "POLR3E" + ], + [ + "POLR2B", + "POLR2C" + ], + [ + "PSMC2", + "PSMD14" + ], + [ + "RAB10", + "RAB11A" + ], + [ + "RAB11A", + "RAB7A" + ], + [ + "RAB4A", + "RAB5C" + ], + [ + "RAB5C", + "RAB7A" + ], + [ + "RAB6A", + "RAB7A" + ], + [ + "RPL23", + "RPS14" + ], + [ + "SEC23B", + "SEC24A" + ], + [ + "SEC23B", + "SEC61B" + ], + [ + "SNRPC", + "SNRPD3" + ], + [ + "SNRPD1", + "SNRPD3" + ], + [ + "SRSF3", + "SRSF7" + ], + [ + "SRSF7", + "SRSF8" + ], + [ + "TIMM23", + "TOMM20" + ] + ], + "complex": [ + [ + "19S proteasome regulatory complex", + "COP9 signalosome variant 1" + ], + [ + "19S proteasome regulatory complex", + "Chaperonin-containing T-complex" + ], + [ + "39S mitochondrial large ribosomal subunit", + "40S cytosolic small ribosomal subunit" + ], + [ + "39S mitochondrial large ribosomal subunit", + "60S cytosolic large ribosomal subunit" + ], + [ + "40S cytosolic small ribosomal subunit", + "60S cytosolic large ribosomal subunit" + ], + [ + "40S cytosolic small ribosomal subunit", + "Small ribosomal subunit processome" + ], + [ + "60S cytosolic large ribosomal subunit", + "UFM1 ribosome E3 ligase complex" + ], + [ + "AP-2 Adaptor complex, alpha1 variant", + "Ubiquitous AP-1 Adaptor complex, sigma1a variant" + ], + [ + "BLOC-1 complex", + "Retromer complex, VPS26A variant" + ], + [ + "CCR4-NOT mRNA deadenylase complex, CNOT6L-CNOT7 variant", + "Nucleolar exosome complex, EXOSC10 variant" + ], + [ + "COG tethering complex", + "COPI vesicle coat complex, COPG1-COPZ1 variant" + ], + [ + "COPI vesicle coat complex, COPG1-COPZ1 variant", + "ESCRT-III complex" + ], + [ + "Chaperonin-containing T-complex", + "HIR histone chaperone complex, UBN1 variant" + ], + [ + "Chaperonin-containing T-complex", + "TTT complex" + ], + [ + "Chromosomal passenger complex, AURKB variant", + "Kinetochore CCAN complex" + ], + [ + "DNA polymerase alpha:primase complex", + "DNA polymerase epsilon complex" + ], + [ + "DNA-directed RNA polymerase I complex", + "DNA-directed RNA polymerase II complex" + ], + [ + "DNA-directed RNA polymerase II complex", + "DNA-directed RNA polymerase III complex, POLR3G variant" + ], + [ + "Dynactin complex", + "Dynein-1 complex, variant 2" + ], + [ + "Eukaryotic translation initiation factor 2 complex", + "Eukaryotic translation initiation factor 2B complex" + ], + [ + "Eukaryotic translation initiation factor 2 complex", + "Eukaryotic translation initiation factor 3 complex" + ], + [ + "Eukaryotic translation initiation factor 3 complex", + "Eukaryotic translation initiation factor 4F, EIF4A1 and EIF4G1 variant" + ], + [ + "GINS complex", + "MCM complex" + ], + [ + "GINS complex", + "Replication fork protection complex" + ], + [ + "INO80 chromatin remodeling complex", + "SWI/SNF ATP-dependent chromatin remodeling complex, ACTL6A-ARID1A-SMARCA2 variant" + ], + [ + "Intron Lariat Spliceosome, type 1 complex", + "SF3B complex" + ], + [ + "Mitochondrial isocitrate dehydrogenase complex (NAD+)", + "Mitochondrial proton-transporting ATP synthase complex" + ], + [ + "Mitochondrial proton-transporting ATP synthase complex", + "Mitochondrial respiratory chain complex IV" + ], + [ + "Mitochondrial proton-transporting ATP synthase complex", + "TIM23 mitochondrial inner membrane pre-sequence translocase complex, TIM17A variant" + ], + [ + "Mitochondrial proton-transporting ATP synthase complex", + "Vacuolar proton translocating ATPase complex, ATP6V0A1 variant" + ], + [ + "NSL histone acetyltransferase complex", + "NuA4 histone acetyltransferase complex" + ], + [ + "Nuclear pore complex", + "TIM23 mitochondrial inner membrane pre-sequence translocase complex, TIM17A variant" + ], + [ + "Nuclear pore complex", + "TREX transcription-export complex, DX39B variant" + ], + [ + "Ragulator complex", + "mTORC1 complex" + ], + [ + "SEC61 protein-conducting channel complex, SEC1A1 variant", + "Signal recognition particle" + ], + [ + "SF3B complex", + "Sm complex" + ], + [ + "Signal recognition particle", + "Signal recognition particle receptor complex" + ], + [ + "Small ribosomal subunit processome", + "UTP-B complex" + ], + [ + "TSC1-TSC2 complex", + "mTORC1 complex" + ], + [ + "U1 small nuclear ribonucleoprotein complex", + "U2 small nuclear ribonucleoprotein complex" + ] + ] +} \ No newline at end of file diff --git a/src/ops_model/models/attention/diffex/viewer/anchor_cells.py b/src/ops_model/models/attention/diffex/viewer/anchor_cells.py new file mode 100644 index 0000000..b82a746 --- /dev/null +++ b/src/ops_model/models/attention/diffex/viewer/anchor_cells.py @@ -0,0 +1,75 @@ +"""Enumerate every viewer ANCHOR cell (the NTC/anchor cells each traversal morphs) with the metadata +needed to recreate the crop — a handoff list for Ritvik to compute SetTransformer attention-head +pixel-patches on the real cells (to compare what the classifier attends to vs the generative morph). + +Anchor cells are the top-N-by-attention cells per (marker, anchor), so this is complete for ALL +markers/anchor types regardless of whether that traversal's frames are cached yet. Phase geneKO + +complex share ONE NTC phase-anchor set; each fluor marker has its own NTC set (its channel's ranking). +""" +from __future__ import annotations + +import pandas as pd + +from ..classifier.config import PMA_PHASE_EBI +from ..classifier.data import _BASE_COLS +from . import catalog as C + +OUT_CSV = f"{C.OUT}/viewer_assets/anchor_cells_for_attention.csv" +# geneKO (`gene`) + pathway membership (`ebi_complex`) columns included per your ask. +# NOTE: guideRNA is NOT in the pma source (no sgRNA column) — would need a per-cell join to +# guide-call data by (experiment, well, segmentation); omitted rather than faked. +# segmentation_id is the pma-source `segmentation` value → (experiment, well, segmentation_id) is the +# unique cell lookup key in the pma CSVs/parquets. +COLS = ["marker_channel", "anchor", "gene", "ebi_complex", "cell_index", "experiment", "well", + "segmentation_id", "x_pheno", "y_pheno", "rank", "pma_attention"] + + +def build(n_cells=20, n_anchors=8, out_csv=OUT_CSV): + """Every (marker, anchor) top-`n_cells` anchor set. Markers = phase + ALL fluor channels in the + data (not just trained ones). Anchors = NTC + the marker's top-`n_anchors` complexes by + EBI complex mAP (cross-phenotype A→B comparison). Future-proof: covers everything the cache holds.""" + import yaml + cdist = C.complex_dist() # complex(name) × reporter EBI mAP + dist = C.dist_matrix() + y = yaml.safe_load(open(C.EBI_YAML)) or {} + gene2cx = {g: v["name"] for v in y.values() if isinstance(v, dict) for g in (v.get("genes") or [])} + parts = [] + + def top_cx(reporter): + if not reporter or reporter not in cdist.columns: + return [] + return list(cdist[reporter].dropna().sort_values(ascending=False).head(n_anchors).index) + + def add(df, mc, anchor): + if df is None or not len(df): + return + d = df.sort_values("rank").head(n_cells).copy() + d["marker_channel"] = mc; d["anchor"] = anchor + d["ebi_complex"] = d["gene"].astype(str).map(gene2cx).fillna("") # the geneKO's own EBI complex membership + d["cell_index"] = range(len(d)); parts.append(d) + + # PHASE: NTC + top EBI-mAP complexes (from the EBI phase parquet; reporter col = "Phase") + ph = pd.read_parquet(PMA_PHASE_EBI, columns=["gene", "predicted_class", "rank_type", *_BASE_COLS]) + ph = ph[ph["rank_type"] == "top"] + add(ph[ph["gene"].astype(str) == "NTC"], "phase", "NTC") + for cx in top_cx("Phase"): + add(ph[ph["predicted_class"].astype(str) == cx], "phase", cx) + + # FLUOR: every channel × (NTC + its top EBI-mAP complexes for that reporter) + fr = pd.read_csv(C.EBI_FLUOR_CSV, usecols=["gene", "channel", "predicted_class", "rank_type", *_BASE_COLS]) + fr = fr[fr["rank_type"] == "top"] + for mc in sorted(fr["channel"].dropna().astype(str).unique()): + sub = fr[fr["channel"] == mc] + add(sub[sub["gene"].astype(str) == "NTC"], mc, "NTC") + for cx in top_cx(C.rep_of(dist, mc)): + add(sub[sub["predicted_class"].astype(str) == cx], mc, cx) + + df = pd.concat(parts, ignore_index=True).rename(columns={"segmentation": "segmentation_id"})[COLS] + df.to_csv(out_csv, index=False) + print(f"[anchor_cells] {len(df)} cells | {df['marker_channel'].nunique()} markers | " + f"{df.groupby('marker_channel')['anchor'].nunique().mean():.1f} anchors/marker -> {out_csv}") + return out_csv + + +if __name__ == "__main__": + build() diff --git a/src/ops_model/models/attention/diffex/viewer/build_attention_heads.py b/src/ops_model/models/attention/diffex/viewer/build_attention_heads.py new file mode 100644 index 0000000..9ee8308 --- /dev/null +++ b/src/ops_model/models/attention/diffex/viewer/build_attention_heads.py @@ -0,0 +1,203 @@ +"""Render Kevin's CellDINO attention-head pixel-attribution npz into static WebP tiles the viewer +overlays (inferno) on the real phenotype cells. Reproducible + SLURM-parallel (one job per shard). + +Kevin's dump (`{SRC}`) has four pixel-attribution trees, all with the same npz schema: + maps (n_cells, n_heads, 128, 128) f16 · crops (n_cells,128,128) f32 (z-scored) + heads (n_heads,2) int32 = (layer,head) · patch_masks (n_cells,196) bool + - phase geneKO : pixel_attribution_cache/.npz + - phase complex : complex_pixel_attribution/phase/pixel_attribution_cache/.npz + - fluor geneKO : fluorescence_pixel_attribution//pixel_attribution_cache/.npz + - fluor complex : complex_pixel_attribution/fluorescence//pixel_attribution_cache/.npz +Ranking metrics (auroc/spec) only exist for phase geneKO (`head_rankings_per_gene.json`); the npz +`heads` array carries the (layer,head) pairs for the rest. + +Overlay pipeline matches Ritvik: Gaussian-smooth each map (σ=2), mask to the CELL (patch_masks → +14×14 → 128 nearest, outside→0). We ship grayscale crop + mask + per-head maps; the webapp applies the +inferno LUT + clim + alpha live. Output (uniform, addressable by the viewer's marker×grain×target): + {AH}////cell/{crop,mask,head0..5}.webp + heads.json + {AH}/index.json {global_max, assets:{:{:[keys]}}} +where modality = "phase" | slugify(marker_channel), grain = geneKO|complex, key = gene | complex-slug. + + python -m ops_model.models.attention.diffex.viewer.build_attention_heads render # SLURM (all trees) + python -m ops_model.models.attention.diffex.viewer.build_attention_heads render --local # serial, no SLURM + python -m ops_model.models.attention.diffex.viewer.build_attention_heads render --dry-run + python -m ops_model.models.attention.diffex.viewer.build_attention_heads index # (re)aggregate index.json +""" +from __future__ import annotations + +import argparse +import glob +import json +import os +from concurrent.futures import ThreadPoolExecutor +from pathlib import Path + +import numpy as np +from PIL import Image +from scipy.ndimage import gaussian_filter + +from ..classifier.config import slugify +from . import catalog as C + +AH_ROOT = f"{C.OUT}/viewer_assets/attention_heads" +SRC = f"{AH_ROOT}/celldino_attention_head_analysis" # Kevin's persistent source dump +RANKINGS = f"{SRC}/head_rankings_per_gene.json" + + +def _save_gray(path, u8, upsize): + im = Image.fromarray(u8) + if upsize: + im = im.resize((upsize, upsize), Image.BILINEAR) + im.save(path, quality=90, method=6) + + +def _crop_u8(crop): + """z-scored crop → uint8 via robust (1–99 pct) min-max, matching typical grayscale display.""" + lo, hi = np.percentile(crop, (1, 99)) + if hi <= lo: + hi = lo + 1e-6 + return (np.clip((crop - lo) / (hi - lo), 0, 1) * 255).astype("uint8") + + +def sources(): + """(modality, grain, cache_dir) for every pixel-attribution tree present in Kevin's dump.""" + out = [("phase", "geneKO", f"{SRC}/pixel_attribution_cache")] + pcx = f"{SRC}/complex_pixel_attribution" + if os.path.isdir(f"{pcx}/phase/pixel_attribution_cache"): + out.append(("phase", "complex", f"{pcx}/phase/pixel_attribution_cache")) + for md in sorted(glob.glob(f"{SRC}/fluorescence_pixel_attribution/*/pixel_attribution_cache")): + out.append((slugify(Path(md).parent.name), "geneKO", md)) + for md in sorted(glob.glob(f"{pcx}/fluorescence/*/pixel_attribution_cache")): + out.append((slugify(Path(md).parent.name), "complex", md)) + return out + + +def _render_one(path, out_dir, rankings, upsize, n_workers): + d = np.load(path, allow_pickle=True) + maps, crops, heads, pmasks = d["maps"].astype(np.float32), d["crops"], d["heads"], d["patch_masks"] + n_cells, n_heads = maps.shape[:2] + H = maps.shape[-1] + grid = int(round(pmasks.shape[1] ** 0.5)) + idx = np.minimum((np.arange(H) * grid // H), grid - 1) # nearest upscale 14→128 + cell_masks = np.empty((n_cells, H, H), bool) + pmaps = np.empty_like(maps) + for c in range(n_cells): + cm = pmasks[c].reshape(grid, grid)[np.ix_(idx, idx)] + cell_masks[c] = cm + for h in range(n_heads): + pmaps[c, h] = np.where(cm, gaussian_filter(maps[c, h], sigma=2.0), 0.0) + gene_max = float(pmaps.max()) + scale = 255.0 / gene_max if gene_max > 0 else 0.0 + pool = ThreadPoolExecutor(max_workers=n_workers) + for c in range(n_cells): + cdir = out_dir / f"cell{c}" + cdir.mkdir(parents=True, exist_ok=True) + pool.submit(_save_gray, cdir / "crop.webp", _crop_u8(crops[c]), upsize) + pool.submit(_save_gray, cdir / "mask.webp", (cell_masks[c] * 255).astype("uint8"), upsize) + for h in range(n_heads): + pool.submit(_save_gray, cdir / f"head{h}.webp", np.clip(pmaps[c, h] * scale, 0, 255).astype("uint8"), upsize) + pool.shutdown(wait=True) + key = out_dir.name + rk = {(r["layer"], r["head"]): r for r in rankings.get(key, [])} # metrics by (layer,head), phase-geneKO only + head_meta = [] + for h in range(n_heads): + layer, head = int(heads[h][0]), int(heads[h][1]) + m = rk.get((layer, head), {}) + head_meta.append({"layer": layer, "head": head, "feature": m.get("feature"), + "spec_p10": m.get("spec_p10"), "spec_min": m.get("spec_min"), + "auroc_vs_ntc": m.get("auroc_vs_ntc")}) + (out_dir / "heads.json").write_text(json.dumps( + {"gene": key, "n_cells": n_cells, "gene_max": gene_max, "heads": head_meta})) + return gene_max + + +def render_shard(modality, grain, npz_paths, out_root=AH_ROOT, upsize=256, n_workers=8, use_rankings=False, force=False): + """SLURM job unit: render a list of npz for one (modality, grain) into ////. + Incremental by default (skips keys whose heads.json already exists); `force` re-renders. Skips + unreadable npz loudly (corrupt-at-source). Returns a small summary.""" + rankings = json.load(open(RANKINGS)) if (use_rankings and os.path.exists(RANKINGS)) else {} + base = Path(out_root) / modality / grain + done, bad, skip = [], [], 0 + for p in npz_paths: + key = Path(p).stem + out_dir = base / key + if not force and (out_dir / "heads.json").exists(): + skip += 1 + continue + try: + _render_one(p, out_dir, rankings, upsize, n_workers) + done.append(key) + except Exception as e: + bad.append(key) + print(f"[attn] SKIP {modality}/{grain}/{key}: {type(e).__name__}") + print(f"[attn] {modality}/{grain}: {len(done)} rendered, {skip} already-present, {len(bad)} unreadable") + return {"modality": modality, "grain": grain, "rendered": len(done), "skipped": skip, "bad": bad} + + +def build_index(out_root=AH_ROOT): + """Aggregate every rendered heads.json into one index.json (availability + fixed-norm scale). + Source of truth = the rendered dirs, so re-running any shard just re-scans the filesystem.""" + assets, gmax = {}, 0.0 + for hj in glob.glob(f"{out_root}/*/*/*/heads.json"): + p = Path(hj) + modality, grain, key = p.parents[2].name, p.parents[1].name, p.parent.name + assets.setdefault(modality, {}).setdefault(grain, []).append(key) + try: + gmax = max(gmax, float(json.loads(p.read_text()).get("gene_max", 0.0))) + except Exception: + pass + for m in assets: + for g in assets[m]: + assets[m][g] = sorted(assets[m][g]) + (Path(out_root) / "index.json").write_text(json.dumps({"global_max": gmax, "assets": assets})) + n = sum(len(v) for m in assets.values() for v in m.values()) + print(f"[attn] index: {len(assets)} modalities, {n} keys, global_max={gmax:.4f} -> {out_root}/index.json") + return f"{out_root}/index.json" + + +def submit(dry_run=False, parallel=40, chunk=150, local=False, upsize=256, force=False): + """Fan out render_shard over all four trees (chunked), then aggregate index.json. Incremental by + default — only new/unrendered npz are processed, so re-running picks up Kevin's newly-dumped markers.""" + jobs = [] + for modality, grain, cache in sources(): + paths = sorted(glob.glob(f"{cache}/*.npz")) + use_r = modality == "phase" and grain == "geneKO" + for i in range(0, len(paths), chunk): + jobs.append({"name": f"ah_{modality[:12]}_{grain[:2]}_{i // chunk}", + "func": render_shard, + "kwargs": dict(modality=modality, grain=grain, npz_paths=paths[i:i + chunk], + use_rankings=use_r, upsize=upsize, force=force), + "metadata": {"modality": modality, "grain": grain}}) + total = sum(len(glob.glob(f'{c}/*.npz')) for _, _, c in sources()) + print(f"[attn] {len(jobs)} shard jobs across {len(sources())} trees ({total} npz, chunk={chunk})") + if local: + for j in jobs: + j["func"](**j["kwargs"]) + return build_index() + from ops_utils.hpc.slurm_batch_utils import submit_parallel_jobs + return submit_parallel_jobs( + jobs_to_submit=jobs, experiment="diffex_attn_heads", + slurm_params={"slurm_partition": "cpu", "cpus_per_task": 8, "mem_gb": 32, "timeout_min": 120, + "slurm_array_parallelism": parallel}, + log_dir="diffex_attn_heads", dry_run=dry_run, + post_completion_callback=lambda *_: build_index()) + + +if __name__ == "__main__": + ap = argparse.ArgumentParser() + sub = ap.add_subparsers(dest="cmd") + r = sub.add_parser("render", help="fan out render_shard on SLURM (default), then build index") + r.add_argument("--local", action="store_true", help="run serially in this process, no SLURM") + r.add_argument("--dry-run", action="store_true") + r.add_argument("--parallel", type=int, default=40) + r.add_argument("--chunk", type=int, default=150) + r.add_argument("--upsize", type=int, default=256) + r.add_argument("--force", action="store_true", help="re-render even if heads.json already exists") + sub.add_parser("index", help="(re)aggregate index.json from rendered dirs") + args = ap.parse_args() + if args.cmd == "index": + build_index() + else: + submit(dry_run=getattr(args, "dry_run", False), parallel=getattr(args, "parallel", 40), + chunk=getattr(args, "chunk", 150), local=getattr(args, "local", False), + upsize=getattr(args, "upsize", 256), force=getattr(args, "force", False)) diff --git a/src/ops_model/models/attention/diffex/viewer/build_complex_ebi_map.py b/src/ops_model/models/attention/diffex/viewer/build_complex_ebi_map.py new file mode 100644 index 0000000..dd7c838 --- /dev/null +++ b/src/ops_model/models/attention/diffex/viewer/build_complex_ebi_map.py @@ -0,0 +1,58 @@ +"""Compute the per-marker EBI complex mAP matrix (complex × reporter) by running the paper's +`phenotypic_consistency_ebi` (copairs, EBI Complex Portal annotations) on each marker's gene +embedding. This is the CORRECT EBI metric per reporter — the `complex_reporter_chad_consistency.csv` +is a different (consistency) metric. Cached to CSV; `catalog.complex_dist()` reads it. +""" +from __future__ import annotations + +import glob +import os + +import anndata as ad +import pandas as pd + +from ops_utils.analysis.map_scores import phenotypic_consistency_ebi + +PS = ("/hpc/projects/icd.fast.ops/organelle_attribution/pca_optimized_v0.3/cell_dino/" + "zscore_per_exp/paper_v2/with_cp/with_4i/all_livecell/fixed_80%/cosine/per_signal") +PHASE_GENE = ("/hpc/projects/icd.fast.ops/organelle_attribution/pca_optimized_v0.3/cell_dino/" + "zscore_per_exp/paper_v2/phase_only/fixed_80%/cosine/gene_embedding_pca_optimized.h5ad") +OUT = "/hpc/projects/icd.fast.ops/models/diffex/complex_reporter_ebi_map.csv" + + +def _ebi_map(h5ad): + a = ad.read_h5ad(h5ad) + if a.X is None or a.X.shape[1] != a.obsm["X_pca"].shape[1]: + a.X = a.obsm["X_pca"] # copairs runs on adata.X + df, _ = phenotypic_consistency_ebi(a, plot_results=False, cache_similarity=True, null_size=1000) + # map complex_num -> name so columns/index match the rest of the app + import yaml + y = yaml.safe_load(open("/hpc/projects/icd.fast.ops/configs/gene_clusters/EBI_complexes_v1_updated_gene_names.yaml")) or {} + n2n = {int(k): v["name"] for k, v in y.items() if isinstance(v, dict) and v.get("name")} + df = df.copy(); df["name"] = df["complex_num"].map(n2n) + return df.dropna(subset=["name"]).set_index("name")["mean_average_precision"] + + +def build(out=OUT): + cols = {} + cols["Phase"] = _ebi_map(PHASE_GENE) + print("[ebi] Phase done") + for h in sorted(glob.glob(f"{PS}/*_gene.h5ad")): + sweep = h.replace("_gene.h5ad", "_sweep.csv") + try: + rep = pd.read_csv(sweep, usecols=["signal"])["signal"].iloc[0] # canonical reporter name + except Exception: + rep = os.path.basename(h).replace("_gene.h5ad", "") + try: + cols[rep] = _ebi_map(h) + print(f"[ebi] {rep}: {int((cols[rep] >= 0.01).sum())}/{len(cols[rep])} >=0.01") + except Exception as e: + print(f"[ebi] {rep} FAILED {e}") + mat = pd.DataFrame(cols) + mat.to_csv(out) + print(f"[ebi] complex x reporter EBI mAP -> {out} shape {mat.shape}") + return out + + +if __name__ == "__main__": + build() diff --git a/src/ops_model/models/attention/diffex/viewer/build_fluor_shap_rankings.py b/src/ops_model/models/attention/diffex/viewer/build_fluor_shap_rankings.py new file mode 100644 index 0000000..a8a8203 --- /dev/null +++ b/src/ops_model/models/attention/diffex/viewer/build_fluor_shap_rankings.py @@ -0,0 +1,114 @@ +"""Split Alex's new shap_screen fluor ranking (one combined CSV, 55 channels, per-cell, robust bin-size +ranking) into the per-marker parquets the v5 traversal build consumes (FRP_DIR/.parquet), in the +SAME schema as the old qualifying rankings so build_marker is a drop-in. + + new CSV cols: gene, channel_name, rank, shap, ..., experiment, well, x_pheno, y_pheno, segmentation_id + old schema: channel_name, gene, rank, pma_attention, experiment, well, x_pheno, y_pheno, segmentation, rank_type + + python -m ops_model.models.attention.diffex.viewer.build_fluor_shap_rankings # local (needs ~64GB) + python -m ops_model.models.attention.diffex.viewer.build_fluor_shap_rankings --submit # SLURM cpu, mem 96 +""" +from __future__ import annotations + +import argparse +import os + +import pandas as pd + +from ..classifier.config import slugify + +CSV = ("/hpc/projects/icd.fast.ops/models/alex_lin_attention/v5/multi_rank/" + "shap_screen/shap_screen_fluor_all.csv") +OUT_DIR = "/hpc/projects/icd.fast.ops/models/diffex/viewer_assets_v5/_rankings/fluor_shap/geneKO" + +# --- fluor COMPLEX (EBI) ranking: per-GENE cells → pool member genes into each complex --- +EBI_CSV = ("/hpc/projects/icd.fast.ops/models/alex_lin_attention/v5/multi_rank/" + "shap_screen_ebi_fluor_all.csv") +# the shap CSV uses OLD gene symbols (RARS/DARS/MARS/… not RARS1/…), so the gene→complex map MUST come from the +# old-gene-names yaml — the updated one silently drops the aminoacyl-tRNA-synthetase complex (11 genes). +EBI_YAML = "/hpc/projects/icd.fast.ops/configs/gene_clusters/EBI_complexes_v1_old_gene_names.yaml" +CX_OUT = "/hpc/projects/icd.fast.ops/models/diffex/viewer_assets_v5/_rankings/fluor_shap/complex" +TOP_N_CX = 500 # top cells per (channel, complex) for a robust centroid (display cap is separate) + + +def _gene_to_complex(): + import yaml + y = yaml.safe_load(open(EBI_YAML)) or {} + m = {} + for _, v in y.items(): + for g in v.get("genes", []): + m[str(g)] = v["name"] + return m + + +def build_complex(): + """Pool the per-gene EBI shap cells into per-complex centroid rankings (gene col = complex name), same + schema build_marker_complex consumes. FAIL LOUD if any non-NTC gene has no complex.""" + os.makedirs(CX_OUT, exist_ok=True) + g2c = _gene_to_complex() + use = ["gene", "channel_name", "rank", "shap", "experiment", "well", "x_pheno", "y_pheno", "segmentation_id"] + print(f"[shap-cx] reading {EBI_CSV} ...", flush=True) + df = pd.read_csv(EBI_CSV, usecols=use) + df = df[~df["gene"].astype(str).str.startswith("NTC")] # NTC = control anchor, not a complex target + df["complex"] = df["gene"].astype(str).map(g2c) + unmapped = sorted(df.loc[df["complex"].isna(), "gene"].astype(str).unique()) + if unmapped: # no good reason a screened gene lacks a complex + raise ValueError(f"[shap-cx] {len(unmapped)} gene(s) unmapped in {os.path.basename(EBI_YAML)}: {unmapped}") + n = 0 + for ch, gch in df.groupby("channel_name"): + parts = [] + for cx, gcx in gch.groupby("complex"): + g = (gcx.drop_duplicates(["experiment", "well", "x_pheno", "y_pheno"]) # one row per cell + .sort_values("shap", ascending=False).head(TOP_N_CX).copy()) # pool members, re-rank by score + g["rank"] = range(1, len(g) + 1) + g["gene"] = cx # grouping key → complex name + parts.append(g) + o = (pd.concat(parts, ignore_index=True)[["channel_name", "gene", "rank", "shap", "experiment", "well", + "x_pheno", "y_pheno", "segmentation_id"]] + .rename(columns={"segmentation_id": "segmentation", "shap": "pma_attention"})) + o["rank_type"] = "top"; o["predicted_class"] = o["gene"] + p = f"{CX_OUT}/{slugify(ch)}.parquet"; o.to_parquet(p) + n += 1 + print(f" [{n:2d}] {slugify(ch):40s} {o.gene.nunique()} complexes, {len(o):>6,} cells -> {os.path.basename(p)}", + flush=True) + print(f"[shap-cx] wrote {n} per-marker complex parquets -> {CX_OUT}") + return {"markers": n} + + +def build(): + os.makedirs(OUT_DIR, exist_ok=True) + use = ["gene", "channel_name", "rank", "shap", "experiment", "well", "x_pheno", "y_pheno", "segmentation_id"] + print(f"[shap-rank] reading {CSV} ...", flush=True) + df = pd.read_csv(CSV, usecols=use) + df = df.rename(columns={"shap": "pma_attention", "segmentation_id": "segmentation"}) + df["rank_type"] = "top" + cols = ["channel_name", "gene", "rank", "pma_attention", "experiment", "well", + "x_pheno", "y_pheno", "segmentation", "rank_type"] + df = df[cols] + print(f"[shap-rank] {len(df):,} rows, {df.channel_name.nunique()} channels", flush=True) + n = 0 + for ch, sub in df.groupby("channel_name"): + p = f"{OUT_DIR}/{slugify(ch)}.parquet" + sub.reset_index(drop=True).to_parquet(p) + n += 1 + print(f" [{n:2d}] {slugify(ch):40s} {len(sub):>8,} rows ({sub.gene.nunique()} classes) -> {os.path.basename(p)}", + flush=True) + print(f"[shap-rank] wrote {n} per-marker parquets -> {OUT_DIR}") + return {"markers": n, "rows": int(len(df))} + + +if __name__ == "__main__": + ap = argparse.ArgumentParser() + ap.add_argument("--submit", action="store_true") + ap.add_argument("--complex", action="store_true", help="build the EBI complex rankings instead of geneKO") + a = ap.parse_args() + fn = build_complex if a.complex else build + if a.submit: + from ops_utils.hpc.slurm_batch_utils import submit_parallel_jobs + submit_parallel_jobs( + jobs_to_submit=[{"name": "fluor_shap_cx" if a.complex else "fluor_shap_split", "func": fn, "kwargs": {}}], + experiment="diffex_shaprank", + slurm_params={"slurm_partition": "cpu", "cpus_per_task": 8, "mem_gb": 96, "timeout_min": 60}, + log_dir="diffex_shaprank", wait_for_completion=False) + else: + fn() diff --git a/src/ops_model/models/attention/diffex/viewer/build_montage_features.py b/src/ops_model/models/attention/diffex/viewer/build_montage_features.py new file mode 100644 index 0000000..94c2d15 --- /dev/null +++ b/src/ops_model/models/attention/diffex/viewer/build_montage_features.py @@ -0,0 +1,61 @@ +"""Export per-gene OP/CP morphometric feature values for the Embedding tab's 'Color by' feature coloring. + +Reads gene_feature_means.h5ad (genes × 3113 OP/CP phase features), collapses variants to base names +(same _base_name dedup as the PC features), robustly normalizes each feature to 0–1 across genes (2–98 +pct) for the colormap, and writes a compact lookup the montage colors points by. + + viewer_assets/montage_features.json + {"features": [base names], "range": {feat: [lo, hi]}, "values": {gene: [0..1 per feature | null]}} + + python -m ops_model.models.attention.diffex.viewer.build_montage_features +""" +from __future__ import annotations + +import json +import os +from collections import OrderedDict + +import numpy as np + +from . import catalog as C +from .build_pc_features import _base_name + +FM = "/hpc/projects/icd.fast.ops/analysis/pc_feature_correlation/phase_only/gene_feature_means.h5ad" +_ASSETS = os.environ.get("OPS_DIFFEX_ASSETS", "viewer_assets") # isolated v5 build → viewer_assets_v5 +OUT = f"{C.OUT}/{_ASSETS}/montage_features.json" + + +def build(fm_path=FM, out=OUT): + import anndata as ad + fm = ad.read_h5ad(fm_path) + genes = [str(g) for g in fm.obs_names] + feats = [str(v) for v in fm.var_names] + X = np.asarray(fm.X, dtype=np.float64) # genes × features + + groups = OrderedDict() # collapse variants → base name (mean of members) + for j, f in enumerate(feats): + groups.setdefault(_base_name(f), []).append(j) + base = sorted(groups) + D = np.column_stack([np.nanmean(X[:, groups[b]], axis=1) for b in base]) # genes × base + print(f"[montage-feat] {len(feats)} features → {len(base)} deduped base names over {len(genes)} genes") + + lo = np.nanpercentile(D, 2, axis=0) + hi = np.nanpercentile(D, 98, axis=0) + span = np.where(hi > lo, hi - lo, 1.0) + N = np.clip((D - lo) / span, 0, 1) # robust 0–1 per feature for the colormap + + fin = lambda v, nd: (round(float(v), nd) if np.isfinite(v) else None) # non-finite → null (valid JSON; NaN is not) + values = {} + for i, g in enumerate(genes): + values[g] = [fin(N[i, j], 3) for j in range(len(base))] + rng = {b: [fin(lo[j], 4), fin(hi[j], 4)] for j, b in enumerate(base)} + + os.makedirs(os.path.dirname(out), exist_ok=True) + with open(out, "w") as f: + json.dump({"features": base, "range": rng, "values": values}, f, allow_nan=False) # fail loud if any NaN slips through + print(f"[montage-feat] -> {out} ({os.path.getsize(out) / 1e6:.1f} MB)") + return out + + +if __name__ == "__main__": + build() diff --git a/src/ops_model/models/attention/diffex/viewer/build_pc_crops_masked.py b/src/ops_model/models/attention/diffex/viewer/build_pc_crops_masked.py new file mode 100644 index 0000000..4eac1dc --- /dev/null +++ b/src/ops_model/models/attention/diffex/viewer/build_pc_crops_masked.py @@ -0,0 +1,200 @@ +"""Re-render the PC-strip crops with the attention_atlas blue cell-mask overlay. + +Kyle's compute_pc_strips.py baked plain 96px Phase2D crops. This re-crops the SAME +representative cells (positions already in pcs/index.json) at a larger size and paints +everything OUTSIDE the target cell translucent blue — matching attention_atlas.py's +"negative overlay" (RGB 0.30/0.40/0.85, alpha 0.55, mask dilated 15px) so the cell of +interest reads in natural phase-gray against a blue surround. + +Source store (per reference_phenotyping_v3_store): + {BASE}/{exp}/3-assembly/phenotyping_v3.zarr/{row}/{col}/0/0 image [1,C,1,Y,X], Phase2D=ch0 + {BASE}/{exp}/3-assembly/phenotyping_v3.zarr/{row}/{col}/0/labels/cell_seg/0 int32 labels + + python -m ops_model.models.attention.diffex.viewer.build_pc_crops_masked --sample 24 # preview + python -m ops_model.models.attention.diffex.viewer.build_pc_crops_masked # full (overwrites crops/) +""" +from __future__ import annotations + +import argparse +import json +import os + +import numpy as np + +from . import catalog as C + +BASE = "/hpc/projects/intracellular_dashboard/fast_ops" +PCS_OUT = f"{C.OUT}/viewer_assets/pcs" +CROP_SIZE = 160 # native px re-crop (was 96); crisper at the 150px display + shows surround +PHASE_CHANNEL = 0 +MASK_DILATION = 15 # matches attention_atlas.MASK_DILATION +OVERLAY_RGB = (0.30, 0.40, 0.85) # attention_atlas negative-overlay blue (#4D66D9) +OVERLAY_ALPHA = 0.55 + + +def _zarr_patch(): + try: + import zarr + from zarr.core.metadata.v3 import ArrayV3Metadata + _o = ArrayV3Metadata.from_dict.__func__ + + @classmethod + def _p(cls, data): + if isinstance(data, dict): + data.pop("storage_transformers", None) + return _o(cls, data) + ArrayV3Metadata.from_dict = _p + except Exception: + pass + + +def _cells(idx): + """Flatten index.json pcData → ALL representative cells with a position + filename (incl. the ones Kyle + left has_crop=false; the zarr is reachable now so most re-crop fine), dedup by output filename.""" + out, seen = [], set() + for d in idx["pcData"].values(): + for b in d["strip"]: + for c in b["cells"]: + if c and c.get("img") and c.get("experiment") and c["img"] not in seen: + seen.add(c["img"]); out.append(c) + return out + + +def _is_blank(phase): + """A crop is blank if it lands in empty/edge space: mostly exact-zeros (stitch gaps / edge pad) or flat.""" + return float((phase == 0).mean()) > 0.4 or float(phase.std()) < 0.02 + + +def _crop(arr, ch, x, y, half): + """Centered crop (channel ch) with zero-pad at the edges → (2*half, 2*half).""" + _, _, _, H, W = arr.shape + y0, y1 = max(0, y - half), min(H, y + half) + x0, x1 = max(0, x - half), min(W, x + half) + raw = np.array(arr[0, ch, 0, y0:y1, x0:x1]) if ch is not None else np.array(arr[0, 0, 0, y0:y1, x0:x1]) + n = 2 * half + if raw.shape != (n, n): + pad = np.zeros((n, n), dtype=raw.dtype) + py, px = (n - raw.shape[0]) // 2, (n - raw.shape[1]) // 2 + pad[py:py + raw.shape[0], px:px + raw.shape[1]] = raw + raw = pad + return raw + + +def _render(phase, seg, half): + """Gray phase (1–99 pct) with blue overlay outside the dilated center cell → uint8 RGB.""" + from scipy.ndimage import binary_dilation + lo, hi = np.percentile(phase, (1, 99)) + if hi - lo < 1e-6: + hi = lo + 1 + g = np.clip((phase - lo) / (hi - lo), 0, 1) + rgb = np.stack([g, g, g], axis=-1) + + center = seg[half, half] + if center == 0: # center on background → use most common label in the central 24px box + c = seg[half - 12:half + 12, half - 12:half + 12] + nz = c[c > 0] + center = np.bincount(nz).argmax() if nz.size else 0 + if center != 0: + inv = ~binary_dilation(seg == center, iterations=MASK_DILATION) + for k in range(3): + rgb[..., k][inv] = rgb[..., k][inv] * (1 - OVERLAY_ALPHA) + OVERLAY_RGB[k] * OVERLAY_ALPHA + return (rgb * 255).astype(np.uint8) + + +def _render_gray(phase): + """Raw grayscale phase (1–99 pct) → uint8 RGB, no overlay (the toggle-off base image).""" + lo, hi = np.percentile(phase, (1, 99)) + if hi - lo < 1e-6: + hi = lo + 1 + g = np.clip((phase - lo) / (hi - lo), 0, 1) + return (np.stack([g, g, g], axis=-1) * 255).astype(np.uint8) + + +def _overlay_rgba(seg, half): + """Transparent inside the (dilated) center cell, blue+alpha outside → RGBA uint8 (toggleable mask layer).""" + from scipy.ndimage import binary_dilation + center = seg[half, half] + if center == 0: + c = seg[half - 12:half + 12, half - 12:half + 12]; nz = c[c > 0] + center = np.bincount(nz).argmax() if nz.size else 0 + rgba = np.zeros((*seg.shape, 4), np.uint8) + if center != 0: + inv = ~binary_dilation(seg == center, iterations=MASK_DILATION) + rgba[inv, 0], rgba[inv, 1], rgba[inv, 2] = [int(v * 255) for v in OVERLAY_RGB] + rgba[inv, 3] = int(OVERLAY_ALPHA * 255) + return rgba + + +def build(sample=0, out=PCS_OUT): + import zarr + from PIL import Image + _zarr_patch() + idx = json.load(open(f"{out}/index.json")) + cells = _cells(idx) + if sample: + cells = cells[:sample] + crops_dir = f"{out}/_crops_sample" + else: + crops_dir = f"{out}/crops" + os.makedirs(crops_dir, exist_ok=True) + half = CROP_SIZE // 2 + cache, ok, blank, fail, valid = {}, 0, 0, 0, set() + for i, c in enumerate(cells): + exp, well = c["experiment"], c["well"] + wr, wc = well[0], well[1:] + key = (exp, well) + if key not in cache: + pos = f"{BASE}/{exp}/3-assembly/phenotyping_v3.zarr/{wr}/{wc}/0" + try: + cache[key] = (zarr.open(f"{pos}/0", mode="r"), zarr.open(f"{pos}/labels/cell_seg/0", mode="r")) + except Exception as e: + cache[key] = None + print(f"[crops] open failed {exp}/{well}: {e}") + if cache[key] is None: + fail += 1; continue + img, seg = cache[key] + x, y = int(round(c["x"])), int(round(c["y"])) + try: + phase = _crop(img, PHASE_CHANNEL, x, y, half) + if _is_blank(phase): # empty/edge region → leave slot as placeholder (no valid crop to show) + blank += 1; continue + segc = _crop(seg, None, x, y, half) + Image.fromarray(_render(phase, segc, half)).save(f"{crops_dir}/{c['img']}") + ok += 1; valid.add(c["img"]) + except Exception as e: + fail += 1 + if fail <= 5: + print(f"[crops] crop failed {c['img']} ({exp}/{well} {x},{y}): {e}") + if (i + 1) % 500 == 0: + print(f"[crops] {i + 1}/{len(cells)} ok={ok} blank={blank} fail={fail}") + print(f"[crops] done: {ok} written, {blank} blank(placeholder), {fail} failed -> {crops_dir}") + if not sample: + _sync_has_crop(idx, valid, out) + return crops_dir + + +def _sync_has_crop(idx, valid, out): + """Update index.json has_crop to reflect what actually rendered (recovers Kyle's false-blanks that now + crop; demotes any that came back blank), so the app shows real crops and only truly-empty slots stay blank.""" + recovered = demoted = 0 + for d in idx["pcData"].values(): + for b in d["strip"]: + for c in b["cells"]: + if not (c and c.get("img")): + continue + now = c["img"] in valid + if now and not c.get("has_crop"): + recovered += 1 + elif not now and c.get("has_crop"): + demoted += 1 + c["has_crop"] = now + with open(f"{out}/index.json", "w") as f: + json.dump(idx, f) + print(f"[crops] index.json has_crop synced: +{recovered} recovered, -{demoted} demoted blank") + + +if __name__ == "__main__": + ap = argparse.ArgumentParser() + ap.add_argument("--sample", type=int, default=0, help="only render first N crops to _crops_sample/ (preview)") + ap.add_argument("--out", default=PCS_OUT) + build(ap.parse_args().sample, ap.parse_args().out) diff --git a/src/ops_model/models/attention/diffex/viewer/build_pc_features.py b/src/ops_model/models/attention/diffex/viewer/build_pc_features.py new file mode 100644 index 0000000..d530975 --- /dev/null +++ b/src/ops_model/models/attention/diffex/viewer/build_pc_features.py @@ -0,0 +1,202 @@ +"""Build the PC ↔ morphometric-feature panel for the OPSin viewer 'PCs' tab. + +The 'Features' toggle (next to tf-idf) swaps the ontology-enrichment bars for +morphometric-feature enrichment, sourced from +`organelle_profiler/scripts/pc_feature_correlation/pc_feature_correlation.py` +(PC × OP/CP phase-feature Pearson r + TF-IDF distinctiveness). Per PC we emit: + - top +corr / -corr features (raw) and top TF-IDF-distinctive features + - feature-class composition and organelle-group composition (both raw & tf-idf) + +Sign alignment: the correlation matrix comes from a gene-level PCA +(gene_embedding_pca_optimized.h5ad), while the viewer strips are binned by Kyle's +own cell-level PCA (viewer geneData 'score' == his gene_pc_scores). Same axes up +to a per-PC sign flip, so we flip each PC's r by sign(pearson(kyle_score, P_corr)) +to make "+corr features" line up with the strip's high bins. Compositions are +unsigned (|r| / tf-idf) so they need no flip. + + python -m ops_model.models.attention.diffex.viewer.build_pc_features +""" +from __future__ import annotations + +import argparse +import importlib.util +import json +import os + +import numpy as np +import pandas as pd + +from . import catalog as C + +PCS_OUT = f"{C.OUT}/viewer_assets/pcs" +CORR_DIR = ("/hpc/projects/icd.fast.ops/analysis/pc_feature_correlation/phase_only") +EMB_H5AD = ("/hpc/projects/icd.fast.ops/organelle_attribution/pca_optimized_v0.3/" + "cell_dino/zscore_per_exp/paper_v1/phase_only/fixed_80%/cosine/" + "gene_embedding_pca_optimized.h5ad") +_SRC = ("/hpc/mydata/gav.sturm/ops_mono/organelle_profiler/scripts/" + "pc_feature_correlation/pc_feature_correlation.py") + +TOP_FEATS = 8 # +corr / -corr / distinctive features shown per PC +COMP_TOP_N = 50 # features whose group shares make up the composition bars + + +def _load_helpers(): + """Import _feat_class/_organelle_group/_composition + taxonomy from the analysis script.""" + spec = importlib.util.spec_from_file_location("pc_feature_correlation", _SRC) + m = importlib.util.module_from_spec(spec) + spec.loader.exec_module(m) + return m + + +def _pc_flips(idx, n_pcs): + """Per PC (1..n_pcs): +1 if the correlation-embedding axis matches the strip's + low→high direction, else -1. Determined by pearson between Kyle's gene scores + (viewer geneData, sparse top_pcs) and the correlation embedding's gene projections.""" + import anndata as ad + emb = ad.read_h5ad(EMB_H5AD) + genes = [str(g) for g in emb.obs_names] + gi = {g: i for i, g in enumerate(genes)} + var = {str(v): i for i, v in enumerate(emb.var_names)} + X = np.asarray(emb.X) + # sparse Kyle scores from geneData.top_pcs → {pc: [(kyle_score, proj)]} + pairs = {p: [] for p in range(1, n_pcs + 1)} + for g, gd in idx["geneData"].items(): + if g not in gi: + continue + for tp in gd["top_pcs"]: + col = f"Phase_PC{tp['pc'] - 1}" + if col in var: + pairs[tp["pc"]].append((tp["score"], X[gi[g], var[col]])) + flips, conf, weak = {}, {}, [] + for p in range(1, n_pcs + 1): + a = np.array(pairs[p]) + if len(a) >= 8: + r = np.corrcoef(a[:, 0], a[:, 1])[0, 1] + flips[p] = -1.0 if r < 0 else 1.0 + conf[p] = abs(r) >= 0.5 + else: + flips[p] = 1.0 + conf[p] = False + if not conf[p]: + weak.append(p) + print(f"[pcfeat] sign flips resolved for {n_pcs} PCs; {len(weak)} low-confidence " + f"(|pearson|<0.5 or sparse): {weak[:12]}{'...' if len(weak) > 12 else ''}") + return flips, conf + + +def _source_group(name): # feature-name prefix → profiling tool + return "CellProfiler" if name.startswith("cp") else "OrganelleProfiler" if name.startswith("op") else "other" + + +# statistic/aggregation suffixes stripped (with trailing numeric param tokens) to get a feature's base name +STAT_TOKENS = {"mean", "median", "std", "max", "min", "sum", "var", "variance", "q1", "q3", "q25", "q75", + "p25", "p75", "iqr", "mad", "sem", "cv", "range", "count", "total", "integrated", "avg"} + + +def _base_name(name): + """Collapse variants of a measurement to one base: drop trailing stat suffixes and numeric params, so + e.g. AngularSecondMoment_3_00_256 / _3_02_256 / eccentricity_mean / _std all group under one name.""" + toks = name.split("_") + while len(toks) > 2 and (toks[-1].isdigit() or toks[-1].lower() in STAT_TOKENS): + toks.pop() + return "_".join(toks) + + +def _dedup_matrix(M): # collapse columns sharing a base name → mean across the variants + groups = {} + for c in M.columns: + groups.setdefault(_base_name(c), []).append(c) + return pd.DataFrame({b: M[cols].mean(axis=1) for b, cols in groups.items()}, index=M.index) + + +def _composition_norm(M, group_of, groups, top_n): + """Like the analysis _composition (share of a PC's top-N features per group) but corrected for how many + features each group has: divide each group's top-N count by that group's total feature count, then + renormalize to sum 1. Removes the base-rate advantage of groups we simply measure more of.""" + gsize = {g: sum(1 for f in M.columns if group_of.get(f) == g) for g in groups} + out = pd.DataFrame(0.0, index=M.index, columns=groups) + for pc in M.index: + top = M.loc[pc].sort_values(ascending=False).head(top_n) + cnt = {} + for f in top.index: + g = group_of.get(f) + cnt[g] = cnt.get(g, 0) + 1 + vals = {g: (cnt.get(g, 0) / gsize[g] if gsize[g] else 0.0) for g in groups} + s = sum(vals.values()) or 1.0 + for g in groups: + out.loc[pc, g] = vals[g] / s + return out + + +def _comp_dict(comp_row, groups): + """PC composition row → {group: fraction} dropping zero groups, biggest first.""" + d = {g: round(float(comp_row[g]), 3) for g in groups if comp_row[g] > 0} + return dict(sorted(d.items(), key=lambda kv: -kv[1])) + + +def _panels(Rm, TFm, H, flips, n_pcs): + """Per-PC feature panel for one matrix pair: top +r/-r + distinctive features and the class/organelle/source + compositions. Works identically on the full feature set or the deduped (base-name-collapsed) one.""" + feats = list(Rm.columns) + cls_of = {f: H._feat_class(f) for f in feats} + org_of = {f: H._organelle_group(f) for f in feats} + src_of = {f: _source_group(f) for f in feats} + cls_g = [g for g in H.FEATURE_CLASSES if any(v == g for v in cls_of.values())] + org_g = [g for g in H.ORGANELLE_GROUPS if any(v == g for v in org_of.values())] + src_g = [g for g in ("CellProfiler", "OrganelleProfiler", "other") if any(v == g for v in src_of.values())] + Rabs = Rm.abs() + comp = H._composition # raw share of top-N features per group + cR = {"cls": comp(Rabs, cls_of, cls_g, COMP_TOP_N), "org": comp(Rabs, org_of, org_g, COMP_TOP_N), "src": comp(Rabs, src_of, src_g, COMP_TOP_N)} + cT = {"cls": comp(TFm, cls_of, cls_g, COMP_TOP_N), "org": comp(TFm, org_of, org_g, COMP_TOP_N), "src": comp(TFm, src_of, src_g, COMP_TOP_N)} + nR = {"cls": _composition_norm(Rabs, cls_of, cls_g, COMP_TOP_N), "org": _composition_norm(Rabs, org_of, org_g, COMP_TOP_N), "src": _composition_norm(Rabs, src_of, src_g, COMP_TOP_N)} + nT = {"cls": _composition_norm(TFm, cls_of, cls_g, COMP_TOP_N), "org": _composition_norm(TFm, org_of, org_g, COMP_TOP_N), "src": _composition_norm(TFm, src_of, src_g, COMP_TOP_N)} + out = {} + for p in range(1, n_pcs + 1): + rn = f"Phase_PC{p - 1}" + s = (Rm.loc[rn] * flips[p]).sort_values(ascending=False) # sign-aligned to strip low→high + pos = [{"f": f, "r": round(float(v), 3)} for f, v in s.head(TOP_FEATS).items() if v > 0] + neg = [{"f": f, "r": round(float(v), 3)} for f, v in s.tail(TOP_FEATS).items() if v < 0][::-1] + td = TFm.loc[rn].sort_values(ascending=False).head(TOP_FEATS) + dist = [{"f": f, "tfidf": round(float(v), 3), "r": round(float(Rm.loc[rn, f] * flips[p]), 3)} for f, v in td.items()] + out[p] = {"raw": {"pos": pos, "neg": neg, + "cls": _comp_dict(cR["cls"].loc[rn], cls_g), "org": _comp_dict(cR["org"].loc[rn], org_g), "src": _comp_dict(cR["src"].loc[rn], src_g), + "clsN": _comp_dict(nR["cls"].loc[rn], cls_g), "orgN": _comp_dict(nR["org"].loc[rn], org_g), "srcN": _comp_dict(nR["src"].loc[rn], src_g)}, + "tfidf": {"dist": dist, + "cls": _comp_dict(cT["cls"].loc[rn], cls_g), "org": _comp_dict(cT["org"].loc[rn], org_g), "src": _comp_dict(cT["src"].loc[rn], src_g), + "clsN": _comp_dict(nT["cls"].loc[rn], cls_g), "orgN": _comp_dict(nT["org"].loc[rn], org_g), "srcN": _comp_dict(nT["src"].loc[rn], src_g)}} + return out, (cls_g, org_g, src_g) + + +def write_features_json(R, TF, flips, conf, n_pcs, out, H=None): + """Assemble + write features.json from a PC×feature correlation matrix R and tf-idf matrix TF (rows named + Phase_PC0..). Shared by the phase build and the per-marker build (build_pcs_marker).""" + H = H or _load_helpers() + full, (cls_g, org_g, src_g) = _panels(R, TF, H, flips, n_pcs) + dedup, _ = _panels(_dedup_matrix(R), _dedup_matrix(TF), H, flips, n_pcs) # base-name-collapsed (default view) + out_data = {"meta": {"classes": cls_g, "orgGroups": org_g, "srcGroups": src_g, + "compTopN": COMP_TOP_N, "topFeats": TOP_FEATS}} + for p in range(1, n_pcs + 1): + out_data[str(p)] = {"dirConf": bool(conf[p]), + "raw": {"full": full[p]["raw"], "dedup": dedup[p]["raw"]}, + "tfidf": {"full": full[p]["tfidf"], "dedup": dedup[p]["tfidf"]}} + path = f"{out}/features.json" + with open(path, "w") as f: + json.dump(out_data, f) + print(f"[pcfeat] {n_pcs} PCs feature panel -> {path} ({os.path.getsize(path) / 1024:.0f} KB)") + return path + + +def build(out=PCS_OUT): + H = _load_helpers() + R = pd.read_parquet(f"{CORR_DIR}/raw/pc_feature_corr_matrix.parquet") # PC × feature signed r + TF = pd.read_parquet(f"{CORR_DIR}/tfidf/tfidf_matrix.parquet") # PC × feature tf-idf (>=0) + idx = json.load(open(f"{out}/index.json")) + n_pcs = idx["overview"]["n_pcs"] + flips, conf = _pc_flips(idx, n_pcs) + return write_features_json(R, TF, flips, conf, n_pcs, out, H) + + +if __name__ == "__main__": + ap = argparse.ArgumentParser() + ap.add_argument("--out", default=PCS_OUT) + build(ap.parse_args().out) diff --git a/src/ops_model/models/attention/diffex/viewer/build_pc_walks.py b/src/ops_model/models/attention/diffex/viewer/build_pc_walks.py new file mode 100644 index 0000000..ed7a50c --- /dev/null +++ b/src/ops_model/models/attention/diffex/viewer/build_pc_walks.py @@ -0,0 +1,161 @@ +"""PC latent-walk morphs: for a marker (or phase), decode a DiffAE traversal along each principal +component of that marker's CellDINO feature space, showing what each unsupervised PC axis looks like +as a cell morph — the generative analogue of the PC-strip tab (which bins real cells). + +Per marker: refit PCA on its per-cell CellDINO features (same fit as build_pcs_marker), take one base +control cell's CellDINO embedding z0 (from the cached anchor embeddings) with a fixed diffusion seed, and +for each PC p decode z0 + (α·σ_p)·(v_p ⊙ sd) across α ∈ [−n_std, +n_std] with the marker's DiffAE. σ_p is +the PC score std; v_p is the (z-scored) eigenvector, mapped back to raw CellDINO space by the mean per-exp +sd. Output: one composite figure per marker (rows = PCs, cols = α). + + python -m ops_model.models.attention.diffex.viewer.build_pc_walks --markers "Mitochondria_TOMM20" + python -m ops_model.models.attention.diffex.viewer.build_pc_walks --all # SLURM, every marker +""" +from __future__ import annotations + +import argparse +import os + +import numpy as np +import torch + +OUT_DIR = "/hpc/projects/icd.fast.ops/analysis/figure4_pc_walks" + + +@torch.no_grad() +def pc_walks_marker(marker_channel, channel, ckpt, out_root, n_pcs=20, n_cells=10, + w=2.0, device="cuda", batch=48, upsize=256, force=False): + """Write PC-walk traversals into the viewer cache as a `pc` grain: one target per PC + (viewer_assets//pc/PC##/cell/frame_.webp + meta.json), so the viewer treats each PC + like any other perturbation (α scrub, cell selection, pinning). α uses the same VIEWER_ALPHAS grid + (−5…+5) as every other traversal, here in units of the PC score's σ.""" + import json + from pathlib import Path + from concurrent.futures import ThreadPoolExecutor + from sklearn.decomposition import PCA + from .build_pcs_marker import _marker_meta, _fp, _load, FIT_N + from .precompute import DirConfig, load_diffae, _sample_guided, _save_webp, VIEWER_ALPHAS + from ..classifier.config import slugify + + dev = torch.device(device if torch.cuda.is_available() else "cpu") + modality = slugify(marker_channel) if marker_channel else "phase" + + # 1) refit PCA on the CellDINO features (per-experiment z-score, pooled) — as build_pcs_marker + if marker_channel: + leaf, reporter, chan, exps, fdir = _marker_meta(marker_channel) + fps = [_fp(e, reporter, fdir) for e in exps] + channel = channel or chan + else: # phase: pooled features_processed_Phase (cell_dino v2) + import glob as _g + fps = sorted(_g.glob("/hpc/projects/icd.fast.ops/*/3-assembly/cell_dino_features_v2/" + "anndata_objects/features_processed_Phase.h5ad"))[:12] + channel = channel or "Phase2D" + rng = np.random.RandomState(0) + zparts, sds = [], [] + for fp in fps: + if not fp or not os.path.exists(fp): + continue + X, _ = _load(fp, [], FIT_N, rng) + mu, sd = X.mean(0), X.std(0) + 1e-8 + zparts.append((X - mu) / sd); sds.append(sd) + pca = PCA(n_components=n_pcs, svd_solver="randomized", random_state=0).fit(np.vstack(zparts)) + sd_mean = np.mean(sds, 0) + evr = (pca.explained_variance_ratio_ * 100) + std_p = np.sqrt(pca.explained_variance_) # PC score std (walk magnitude unit) + print(f"[pcw] {marker_channel}: PCA {n_pcs} PCs; PC1-5 var {evr[:5].round(2)}") + + # 2) base control cells (cached anchor CellDINO embs) + fixed diffusion seeds + cfg = DirConfig(grain="geneKO", target="NTC", control="NTC", device=device) + if ckpt: + cfg.diffae_ckpt = ckpt + if marker_channel: + cfg.marker_channel = marker_channel + if channel: + cfg.channel = channel + H = cfg.crop_size + ctrl = np.load(f"{out_root}/viewer_assets/{modality}/_anchors/NTC/ctrl.npz")["ctrl_embs"] + ncell = min(n_cells, len(ctrl)) + z0 = torch.as_tensor(ctrl[:ncell], dtype=torch.float32, device=dev) + xT = torch.stack([torch.randn(1, H, H, generator=torch.Generator(device=dev).manual_seed(1234 + c), device=dev) + for c in range(ncell)]) + diffae = load_diffae(cfg, dev) + null = diffae.null_emb.detach()[None].to(dev) + alphas = list(VIEWER_ALPHAS) # same −5…+5 grid as every other traversal (σ units) + n_frames = len(alphas) + + # 3) per PC: decode each base cell across the α (σ) sweep, write frames + meta (viewer `pc` grain) + for p in range(n_pcs): + slug = f"PC{p + 1:02d}" + adir = Path(out_root) / "viewer_assets" / modality / "pc" / slug + if (adir / "meta.json").exists() and not force: + continue + vp = torch.as_tensor(pca.components_[p] * sd_mean, dtype=torch.float32, device=dev)[None] + conds, keys = [], [] + for c in range(ncell): + for i, a in enumerate(alphas): + conds.append(z0[c:c + 1] + (a * std_p[p]) * vp); keys.append((c, i)) + gen = np.empty((ncell, n_frames, H, H), np.float32) + for i0 in range(0, len(conds), batch): + cb = torch.cat(conds[i0:i0 + batch], 0) + xb = torch.cat([xT[c:c + 1] for c, _ in keys[i0:i0 + batch]], 0) + outb = _sample_guided(diffae, xb, cb, null.expand(cb.shape[0], -1), w, cfg).cpu().numpy()[:, 0] + for k, (c, i) in enumerate(keys[i0:i0 + batch]): + gen[c, i] = outb[k] + fp2 = ThreadPoolExecutor(max_workers=8) + for c in range(ncell): + (adir / f"cell{c}").mkdir(parents=True, exist_ok=True) + for i in range(n_frames): + fp2.submit(_save_webp, adir / f"cell{c}" / f"frame_{i:02d}.webp", gen[c, i], upsize) + fp2.shutdown(wait=True) + (adir / "meta.json").write_text(json.dumps({ + "grain": "pc", "target": slug, "modality": modality, "control": None, + "marker_channel": marker_channel, "channel": channel, "slug": slug, "w": w, + "alphas": alphas, "n_cells": ncell, "has_scores": False, "has_real": True, + "real_dir": f"{modality}/_anchors/NTC", "asset_dir": f"{modality}/pc/{slug}", + "explained_variance": round(float(evr[p]), 2)})) + print(f"[pcw] {modality}/{slug}: {ncell}×{n_frames} ({evr[p]:.1f}% var)") + print(f"[pc-walks] {modality}: {n_pcs} PC targets written") + return f"{out_root}/viewer_assets/{modality}/pc" + + +def _marker_jobs(n_pcs, n_cells, force): + from . import catalog as C + from ..classifier.config import slugify + jobs = [] + for d, mc, ch in C.complete_markers(): + jobs.append({"name": f"pcw_{slugify(mc)[:18]}", "func": pc_walks_marker, + "kwargs": dict(marker_channel=mc, channel=ch, ckpt=f"{C.DD}/{d}/diffae_best.pt", + out_root=C.OUT, n_pcs=n_pcs, n_cells=n_cells, force=force)}) + return jobs + + +if __name__ == "__main__": + ap = argparse.ArgumentParser() + ap.add_argument("--markers", nargs="+", help="marker_channel(s) to run locally (needs a GPU)") + ap.add_argument("--all", action="store_true", help="submit every complete marker via SLURM (GPU)") + ap.add_argument("--n-pcs", type=int, default=20) + ap.add_argument("--n-cells", type=int, default=10) + ap.add_argument("--force", action="store_true") + a = ap.parse_args() + if a.all: + from ops_utils.hpc.slurm_batch_utils import submit_parallel_jobs + from . import catalog as C + jobs = _marker_jobs(a.n_pcs, a.n_cells, a.force) + jobs.append({"name": "pcw_phase", "func": pc_walks_marker, # phase embedding too + "kwargs": dict(marker_channel=None, channel="Phase2D", + ckpt=f"{C.DD}/phase_v1/diffae_best.pt", out_root=C.OUT, + n_pcs=a.n_pcs, n_cells=a.n_cells, force=a.force)}) + print(f"[pc-walks] submitting {len(jobs)} jobs (markers + phase)") + submit_parallel_jobs(jobs_to_submit=jobs, experiment="diffex_pc_walks", + slurm_params={"slurm_partition": "gpu", "gpus_per_node": 1, "cpus_per_task": 12, + "mem_gb": 64, "timeout_min": 1000}, # ~200 PCs @ ~3.4min/PC + log_dir="diffex_pc_walks", wait_for_completion=False) + elif a.markers: + from . import catalog as C + by_mc = {mc: (d, ch) for d, mc, ch in C.complete_markers()} + for mc in a.markers: + d, ch = by_mc[mc] + pc_walks_marker(mc, ch, f"{C.DD}/{d}/diffae_best.pt", C.OUT, + n_pcs=a.n_pcs, n_cells=a.n_cells, force=a.force) + else: + ap.error("pass --markers or --all") diff --git a/src/ops_model/models/attention/diffex/viewer/build_pcs.py b/src/ops_model/models/attention/diffex/viewer/build_pcs.py new file mode 100644 index 0000000..63cc84a --- /dev/null +++ b/src/ops_model/models/attention/diffex/viewer/build_pcs.py @@ -0,0 +1,165 @@ +"""Build PC Strip Explorer assets for the OPSin viewer 'PCs' tab. + +Kyle's `kyle_pcs/build_static_explorer.py` bakes everything (inline JSON + base64 crops) into one +41 MB HTML. This splits that into the static-asset layout the OPSin app uses: + viewer_assets/pcs/index.json {overview, pcData, geneData, geneNames} + viewer_assets/pcs/crops/.png the per-cell strip crops (pc{NNN}_bin{NN}_row{N}.png) + +Source = the self-contained `pc_explorer_static.html` (the assembled data + crops live inline there; +the raw artifacts dir is gone). If Kyle regenerates artifacts, re-run his build_static_explorer.py +then point --html at the fresh output. + + python -m ops_model.models.attention.diffex.viewer.build_pcs + python -m ops_model.models.attention.diffex.viewer.build_pcs --html /path/to/pc_explorer_static.html +""" +from __future__ import annotations + +import argparse +import base64 +import json +import os + +from . import catalog as C + +DEFAULT_HTML = os.path.join(os.path.dirname(__file__), "..", "kyle_pcs", "pc_explorer_static.html") +PCS_OUT = f"{C.OUT}/viewer_assets/pcs" + + +def build_from_html(html_path=DEFAULT_HTML, out=PCS_OUT): + """Extract the inline `const overview/pcData/geneData/geneNames/cropB64 = …;` blocks (one per line) + → index.json + decoded crop PNGs. No re-computation; the HTML is the assembled source of truth.""" + consts, crop_b64 = {}, {} + for line in open(html_path): + for v in ("overview", "pcData", "geneData", "geneNames"): + pre = f"const {v} = " + if line.startswith(pre): + consts[v] = json.loads(line[len(pre):].rstrip()[:-1]) # strip trailing ';' + if line.startswith("const cropB64 = "): + crop_b64 = json.loads(line[len("const cropB64 = "):].rstrip()[:-1]) + missing = [v for v in ("overview", "pcData", "geneData", "geneNames") if v not in consts] + if missing: + raise SystemExit(f"could not extract {missing} from {html_path}") + + os.makedirs(f"{out}/crops", exist_ok=True) + for name, b64 in crop_b64.items(): + with open(f"{out}/crops/{name}", "wb") as f: + f.write(base64.b64decode(b64)) + with open(f"{out}/index.json", "w") as f: + json.dump({k: consts[k] for k in ("overview", "pcData", "geneData", "geneNames")}, f) + ov = consts["overview"] + print(f"[pcs] {ov['n_pcs']} PCs · {ov['n_genes']} genes · {len(crop_b64)} crops -> {out}/index.json") + return f"{out}/index.json" + + +def _pc_gene_sets(gene_data, top_n=50, tfidf=False): + """Per PC, split genes by the SIGN of their loading (from geneData top_pcs) → high/low gene lists, + each capped at top_n. Default ranks by |loading|. With tfidf=True, rank by |loading|·idf where + idf = log(nPCs/(1+df)) and df = # PC-directions the gene is in the raw top-n — so genes shared + across many PCs sink and PC-unique genes rise (→ the enrichment reflects PC-unique biology).""" + import math + high, low = {}, {} + for gene, gd in gene_data.items(): + for tp in gd["top_pcs"]: + (high if tp["score"] > 0 else low).setdefault(tp["pc"], []).append((gene, abs(tp["score"]))) + n_pcs = max([*high, *low], default=1) + if not tfidf: + top = lambda d: {p: [g for g, _ in sorted(v, key=lambda x: -x[1])[:top_n]] for p, v in d.items()} + return top(high), top(low) + df = {} # document frequency over the raw top-n sets (both directions) + for d in (high, low): + for p, v in d.items(): + for g, _ in sorted(v, key=lambda x: -x[1])[:top_n]: + df[g] = df.get(g, 0) + 1 + idf = lambda g: math.log(n_pcs / (1 + df.get(g, 0))) + rr = lambda d: {p: [g for g, _ in sorted([(g, s * idf(g)) for g, s in v], key=lambda x: -x[1])[:top_n]] for p, v in d.items()} + return rr(high), rr(low) + + +# the 4 the viewer shows (labels below); GO_BP, GO_compartments, KEGG, Reactome +ENRICH_LIBS = ("GO_Biological_Process_2025", "GO_Cellular_Component_2025", "Reactome_2022", "KEGG_2026") +LIB_LABEL = {"GO_Biological_Process_2025": "GO BP", "GO_Cellular_Component_2025": "GO compartment", + "Reactome_2022": "Reactome", "KEGG_2026": "KEGG"} + + +def _n_overlap(overlap): + return len([x for x in str(overlap).strip("[]").split(",") if x.strip()]) + + +def annotate_term_sizes(out=PCS_OUT): + """Add K (total genes in each ontology term) to every enrichment record so the app can show k/K %. + Enrichr's API returns only the overlap gene list, not K — it lives in the library GMT. Fetch each + library's GMT once, key term sizes by the same display name we stored (GO id stripped), annotate.""" + import urllib.request + sizes = {} + for lib in ENRICH_LIBS: + url = f"https://maayanlab.cloud/Enrichr/geneSetLibrary?mode=text&libraryName={lib}" + try: + txt = urllib.request.urlopen(url, timeout=120).read().decode() + except Exception as e: + print(f"[pcs] GMT fetch failed for {lib}: {e}"); continue + m = {} + for line in txt.splitlines(): + p = line.split("\t") + if len(p) < 3: + continue + m[p[0].split(" (GO:")[0]] = len([g for g in p[2:] if g.strip()]) # strip GO id to match stored term + sizes[LIB_LABEL[lib]] = m + print(f"[pcs] {lib}: {len(m)} term sizes") + enr = json.load(open(f"{out}/enrichment.json")) + for dirs in enr.values(): + for libs in dirs.values(): + for lib, terms in libs.items(): + sm = sizes.get(lib, {}) + for t in terms: + t["K"] = sm.get(t["term"]) + with open(f"{out}/enrichment.json", "w") as f: + json.dump(enr, f) + print(f"[pcs] annotated term sizes -> {out}/enrichment.json") + + +def build_enrichment(out=PCS_OUT, top_n=50, top_terms=8): + """Enrichr (speedrichr) per PC × direction × library on the PC-loading gene sets → enrichment.json: + {pc: {high: {lib_label: [{term, adjp, n_overlap}]}, low: {...}}}, terms ranked by adjusted p-value. + Reuses embedding_overlays._run_cluster_enrichment (GO BP + GO compartments + Reactome + KEGG).""" + from ops_model.post_process.combination.embedding_overlays import _run_cluster_enrichment + idx = json.load(open(f"{out}/index.json")) + bg = idx["geneNames"] + rawH, rawL = _pc_gene_sets(idx["geneData"], top_n, tfidf=False) + tfH, tfL = _pc_gene_sets(idx["geneData"], top_n, tfidf=True) + c2g = {} # raw high/low (H/L) + tf-idf high/low (h/l) per PC — cache both so the toggle needs no API + for p, g in rawH.items(): c2g[f"H{p}"] = g + for p, g in rawL.items(): c2g[f"L{p}"] = g + for p, g in tfH.items(): c2g[f"h{p}"] = g + for p, g in tfL.items(): c2g[f"l{p}"] = g + print(f"[pcs] enrichment on {len(c2g)} PC×direction sets (raw + tf-idf) × {len(ENRICH_LIBS)} libraries...") + res = _run_cluster_enrichment(c2g, background_genes=bg, libraries=ENRICH_LIBS, top_n_terms=top_terms) + + def compact(key): + bl = (res.get(key) or {}).get("by_library", {}) + return {LIB_LABEL[lib]: [{"term": t["term"].split(" (GO:")[0], "adjp": t["adj_pvalue"], "n": _n_overlap(t["overlap"])} + for t in bl.get(lib, [])] for lib in ENRICH_LIBS if lib in bl} + enr = {str(p): {"high": compact(f"H{p}"), "low": compact(f"L{p}"), + "high_tfidf": compact(f"h{p}"), "low_tfidf": compact(f"l{p}")} + for p in range(1, idx["overview"]["n_pcs"] + 1)} + with open(f"{out}/enrichment.json", "w") as f: + json.dump(enr, f) + print(f"[pcs] enrichment for {len(enr)} PCs -> {out}/enrichment.json") + annotate_term_sizes(out) # add K (term sizes) for k/K % + return f"{out}/enrichment.json" + + +if __name__ == "__main__": + ap = argparse.ArgumentParser() + ap.add_argument("--html", default=DEFAULT_HTML) + ap.add_argument("--out", default=PCS_OUT) + ap.add_argument("--enrich", action="store_true", help="also run GO/KEGG/Reactome enrichment per PC") + ap.add_argument("--enrich-only", action="store_true", help="skip html extract; just (re)run enrichment") + ap.add_argument("--sizes-only", action="store_true", help="just annotate term sizes (K) onto existing enrichment.json") + args = ap.parse_args() + if args.sizes_only: + annotate_term_sizes(args.out) + else: + if not args.enrich_only: + build_from_html(args.html, args.out) + if args.enrich or args.enrich_only: + build_enrichment(args.out) diff --git a/src/ops_model/models/attention/diffex/viewer/build_pcs_marker.py b/src/ops_model/models/attention/diffex/viewer/build_pcs_marker.py new file mode 100644 index 0000000..a99a427 --- /dev/null +++ b/src/ops_model/models/attention/diffex/viewer/build_pcs_marker.py @@ -0,0 +1,356 @@ +"""Per-marker PC-strip build (Kyle-style) from paper_v2 cell_dino_features_v2. + +For a viewer marker we refit PCA on that marker's per-cell CellDINO features (which carry the cell's +position, exactly like Kyle's .pt inputs), bin cells low->high along each PC, and crop that marker's own +fluor channel from phenotyping_v3.zarr with the blue negative cell-mask overlay. Gene chips/loadings come +from the same refit (gene-mean PC scores), so strips + chips are self-consistent. + +Output is ADDITIVE and isolated (does not touch the phase pcs/ cache that is syncing to S3): + viewer_assets/pcs/markers//index.json (same schema as the phase pcs/index.json) + viewer_assets/pcs/markers//crops/pc###_bin##_row#.png + + python -m ops_model.models.attention.diffex.viewer.build_pcs_marker --marker "autophagosome_MAP1LC3B" +""" +from __future__ import annotations + +import argparse +import glob +import json +import os + +import numpy as np + +from . import catalog as C +from . import marker_leaves as ML +from .build_pc_crops_masked import CROP_SIZE, _crop, _is_blank, _render, _zarr_patch + +FOPS = "/hpc/projects/intracellular_dashboard/fast_ops" +PCS_OUT = f"{C.OUT}/viewer_assets/pcs/markers" +N_PCS, N_BINS, N_ROWS = 40, 15, 3 # PCs shown; strip bins; cells per bin +FIT_N, SEL_N = 120_000, 300_000 # cells subsampled per experiment for PCA fit / representative selection + + +def _exp_dir(exp): + d = sorted(glob.glob(f"{FOPS}/{exp}_*")) + return d[0] if d else None + + +def _features_dir(leaf): # CP / 4i channel CellDINO features live in adjacent _cp / _4i folders + return "cell_dino_features_v2_cp" if leaf.endswith("_cp") else "cell_dino_features_v2_4i" if leaf.endswith("_4i") else "cell_dino_features_v2" + + +def _marker_meta(marker_channel): + """Resolve a viewer marker → (reporter, channel_name, [exp prefixes]) using its leaf + a features_processed file.""" + import h5py + leaf = ML.resolve_leaf(marker_channel) + if not leaf: + raise SystemExit(f"no paper_v2 leaf for {marker_channel}") + manifest = f"{ML.leaf_dir(leaf)}/downsampled_manifest.csv" + fdir = _features_dir(leaf) + import csv + exps = [] + with open(manifest) as f: + for row in csv.DictReader(f): + exps = [e.strip() for e in row["experiments"].split(",")] + # find the marker's features_processed_.h5ad (non-Phase) in the first available experiment + reporter = channel = None + for e in exps: + d = _exp_dir(e) + if not d: + continue + cand = [f for f in glob.glob(f"{d}/3-assembly/{fdir}/anndata_objects/features_processed_*.h5ad") + if "_Phase" not in f] + for fp in cand: + with h5py.File(fp, "r") as h: + ch = h["uns"]["channel"][()] + ch = ch.decode() if isinstance(ch, bytes) else ch + reporter = os.path.basename(fp)[len("features_processed_"):-len(".h5ad")] + channel = ch + break + if reporter: + break + if not reporter: + raise SystemExit(f"no features_processed for {marker_channel} ({leaf})") + return leaf, reporter, channel, exps, fdir + + +def _fp(exp, reporter, fdir="cell_dino_features_v2"): + d = _exp_dir(exp) + return f"{d}/3-assembly/{fdir}/anndata_objects/features_processed_{reporter}.h5ad" if d else None + + +def _channel_index(exp, channel_name): + """Index of `channel_name` in this experiment's phenotyping_v3.zarr channel list.""" + import zarr + g = zarr.open_group(f"{_exp_dir(exp)}/3-assembly/phenotyping_v3.zarr/A/1/0", mode="r") + labels = [c["label"] for c in dict(g.attrs["ome"])["omero"]["channels"]] + return labels.index(channel_name) + + +def _load(fp, cols, n=None, rng=None): + """Read features X (subsampled) + obs columns from a features_processed h5ad.""" + import h5py + with h5py.File(fp, "r") as h: + N = h["X"].shape[0] + idx = np.arange(N) if (n is None or n >= N) else np.sort(rng.choice(N, n, replace=False)) + X = h["X"][idx].astype(np.float64) + out = {} + for c in cols: + g = h["obs"][c] + if isinstance(g, h5py.Group): # categorical + cats = [x.decode() if isinstance(x, bytes) else x for x in g["categories"][:]] + out[c] = np.array([cats[i] for i in g["codes"][idx]]) + else: + out[c] = g[idx] + return X, out + + +def build_marker(marker_channel, out=PCS_OUT, with_enrich=True): + from PIL import Image + import zarr + _zarr_patch() + leaf, reporter, channel, exps, fdir = _marker_meta(marker_channel) + slug = ML._norm(marker_channel) + print(f"[pcm] {marker_channel} → leaf {leaf}, reporter {reporter}, channel {channel}, {len(exps)} exps") + rng = np.random.RandomState(0) + + # 1) fit PCA on z-scored features (per-experiment z-score), pooled subsample + from sklearn.decomposition import PCA + zparts, stats = [], {} + for e in exps: + fp = _fp(e, reporter, fdir) + if not fp or not os.path.exists(fp): + continue + X, _ = _load(fp, [], FIT_N, rng) + mu, sd = X.mean(0), X.std(0) + 1e-8 + stats[e] = (mu, sd) + zparts.append((X - mu) / sd) + pca = PCA(n_components=N_PCS, svd_solver="randomized", random_state=0).fit(np.vstack(zparts)) + ev = (pca.explained_variance_ratio_ * 100).round(3).tolist() + print(f"[pcm] PCA fit; PC0-4 var {ev[:5]}") + + # 2) selection sample (with positions) across experiments → PC scores + scores, well, xs, ys, pert, exp_of = [], [], [], [], [], [] + for e in exps: + fp = _fp(e, reporter, fdir) + if e not in stats or not fp: + continue + X, obs = _load(fp, ["well", "x_position", "y_position", "perturbation"], SEL_N, rng) + mu, sd = stats[e] + sc = ((X - mu) / sd - pca.mean_) @ pca.components_.T + scores.append(sc); xs.append(obs["x_position"]); ys.append(obs["y_position"]) + pert.append(obs["perturbation"]); exp_of += [e] * len(sc) + well.append(np.array([w.split("_")[0] for w in obs["well"]])) + S = np.vstack(scores); X0 = np.concatenate(xs); Y0 = np.concatenate(ys) + W = np.concatenate(well); P = np.concatenate(pert); E = np.array(exp_of) + print(f"[pcm] selection sample {S.shape}") + + # 3) gene loadings = gene-mean PC scores (self-consistent with the strips) + genes = sorted(set(P)) + gmean = {g: S[P == g].mean(0) for g in genes} + geneData = {} + for g in genes: + prof = gmean[g] + order = np.argsort(-np.abs(prof))[:15] + geneData[g] = {"top_pcs": [{"pc": int(p) + 1, "score": round(float(prof[p]), 3)} for p in order], + "profile": [round(float(v), 3) for v in prof]} + + # 4) representatives per PC×bin + crop the marker channel + crops_dir = f"{out}/{slug}/crops"; os.makedirs(crops_dir, exist_ok=True) + half = CROP_SIZE // 2 + zc, ci = {}, {} # zarr handle cache, channel-index cache + pcData = {}; ok = 0 + for p in range(N_PCS): + sc = S[:, p] + tgt = np.percentile(sc, np.linspace(2, 98, N_BINS)) + strip = [] + high = sorted(genes, key=lambda g: -gmean[g][p])[:15] + low = sorted(genes, key=lambda g: gmean[g][p])[:15] + for bi, t in enumerate(tgt): + near = np.argsort(np.abs(sc - t))[:N_ROWS * 4] # candidates; take first N_ROWS that crop OK + cells = [] + for j in near: + if len(cells) >= N_ROWS: + break + e = E[j]; w = W[j]; x, y = int(round(X0[j])), int(round(Y0[j])) + r, col = w.split("/")[0], w.split("/")[1] + key = (e, r, col) + if key not in zc: + b = f"{_exp_dir(e)}/3-assembly/phenotyping_v3.zarr/{r}/{col}/0" + try: + zc[key] = (zarr.open(f"{b}/0", mode="r"), zarr.open(f"{b}/labels/cell_seg/0", mode="r")) + ci[e] = ci.get(e) or _channel_index(e, channel) + except Exception: + zc[key] = None + if zc[key] is None: + continue + img, seg = zc[key] + try: + ph = _crop(img, ci[e], x, y, half) + if _is_blank(ph): + continue + fn = f"pc{p:03d}_bin{bi:02d}_row{len(cells)}.png" + Image.fromarray(_render(ph, _crop(seg, None, x, y, half), half)).save(f"{crops_dir}/{fn}") + cells.append({"gene": str(P[j]), "score": round(float(sc[j]), 2), "experiment": e, + "well": w, "x": round(float(X0[j]), 1), "y": round(float(Y0[j]), 1), + "has_crop": True, "img": fn}) + ok += 1 + except Exception: + continue + while len(cells) < N_ROWS: + cells.append(None) + strip.append({"cells": cells}) + pcData[str(p + 1)] = {"pc": p + 1, "explained_variance": ev[p], "strip": strip, + "high_genes": [{"gene": g, "score": round(float(gmean[g][p]), 3)} for g in high], + "low_genes": [{"gene": g, "score": round(float(gmean[g][p]), 3)} for g in low]} + if (p + 1) % 10 == 0: + print(f"[pcm] {p + 1}/{N_PCS} PCs, {ok} crops") + + index = {"overview": {"n_genes": len(genes), "n_pcs": N_PCS, "n_bins": N_BINS, "n_rows": N_ROWS, + "crop_size": CROP_SIZE, "marker": marker_channel, "channel": channel, + "total_variance": round(sum(ev), 1), "explained_variance": ev}, + "pcData": pcData, "geneData": geneData, "geneNames": genes} + with open(f"{out}/{slug}/index.json", "w") as f: + json.dump(index, f) + print(f"[pcm] {marker_channel}: {ok} crops, {len(genes)} genes -> {out}/{slug}/index.json") + if with_enrich: + from .build_pcs import build_enrichment + build_enrichment(out=f"{out}/{slug}") # ontology enrichment (speedrichr) + term sizes + build_marker_features(slug, out) # morphometric features + return slug + + +GENE_FEATURE_MEANS = "/hpc/projects/icd.fast.ops/analysis/pc_feature_correlation/phase_only/gene_feature_means.h5ad" + + +def build_marker_features(slug, out=PCS_OUT): + """Per-marker morphometric features: correlate the marker's gene PC loadings vs the OP/CP gene feature means. + Strips + loadings share one PCA here, so no sign flip is needed (flips = +1, dirConf = True).""" + import anndata as ad + import pandas as pd + from .build_pc_features import _load_helpers, write_features_json + idx = json.load(open(f"{out}/{slug}/index.json")) + n_pcs = idx["overview"]["n_pcs"] + gd = idx["geneData"] + genes = [g for g in idx["geneNames"] if g in gd] + P = pd.DataFrame([gd[g]["profile"][:n_pcs] for g in genes], index=genes, + columns=[f"Phase_PC{i}" for i in range(n_pcs)]) + fm = ad.read_h5ad(GENE_FEATURE_MEANS) + fm_genes = [str(g) for g in fm.obs_names] + common = [g for g in genes if g in set(fm_genes)] + P = P.loc[common] + F = pd.DataFrame(np.asarray(fm[[g for g in common], :].X), index=common, columns=[str(v) for v in fm.var_names]) + H = _load_helpers() + R = H.correlate(P, F) # n_pcs × feature signed Pearson r + TF, _ = H.tfidf_distinctive(R) # tf-idf distinctiveness + R.index = [f"Phase_PC{i}" for i in range(n_pcs)]; TF.index = R.index + flips = {p: 1.0 for p in range(1, n_pcs + 1)}; conf = {p: True for p in range(1, n_pcs + 1)} + return write_features_json(R, TF, flips, conf, n_pcs, f"{out}/{slug}", H) + + +def build_marker_layouts(mont_dir=None): + """Per-marker Live-mode layouts: reposition dots by each marker's own embedding (X_umap/X_phate), reusing + the phase layout's rich gene annotations for color-by. Cheap (no tiles) → Live renders these per marker.""" + import anndata as ad + from . import build_umap_montage as BM + mont_dir = mont_dir or f"{C.OUT}/viewer_assets_v5/_montage" + phase = json.load(open(f"{mont_dir}/layout_umap.json")) + ann_of = {g["g"]: {k: v for k, v in g.items() if k not in ("g", "nx", "ny")} for g in phase["genes"]} + cfields = phase["color_fields"] + import re + jss = lambda s: re.sub(r"[^A-Za-z0-9]", "_", str(s)).strip("_") # matches the app's jsSlug (montage modality) + n = 0 + for mk in _all_markers(): + h5 = ML.embedding_h5ad(mk); slug = jss(mk) + a = ad.read_h5ad(h5) + for emb in ("umap", "phate"): + c = BM._embed_coords(a, emb); lo = c.min(0); rng = c.max(0) - lo; rng[rng == 0] = 1 + genes = [] + for i, g in enumerate(a.obs["perturbation"]): + g = str(g) + rec = {"g": g, "nx": float((c[i, 0] - lo[0]) / rng[0]), "ny": float((c[i, 1] - lo[1]) / rng[1])} + rec.update(ann_of.get(g, {})) + genes.append(rec) + with open(f"{mont_dir}/layout_{slug}_{emb}.json", "w") as f: + json.dump({"embedding": emb, "color_fields": cfields, "genes": genes}, f) + n += 1 + print(f"[pcm] wrote per-marker Live layouts for {n} markers → {mont_dir}/layout__.json") + + +def build_marker_job(marker_channel, out=PCS_OUT): + """SLURM job: PC strips + morphometric features for one marker (no external calls).""" + slug = build_marker(marker_channel, out, with_enrich=False) + build_marker_features(slug, out) + return {"marker": marker_channel, "slug": slug} + + +def enrich_shard(markers, out=PCS_OUT): + """SLURM job: ontology enrichment (speedrichr) for a handful of markers, serially (few shards = throttled).""" + from .build_pcs import build_enrichment + done = [] + for mk in markers: + slug = ML._norm(mk) + try: + build_enrichment(out=f"{out}/{slug}"); done.append(slug) + except Exception as e: + print(f"[pcm] enrich failed {slug}: {e}") + return {"done": done} + + +def _all_markers(): + m = json.load(open(f"{C.OUT}/viewer_assets_v5/manifest.json")) + mks = [x["marker_channel"] for x in m["markers"] if x.get("marker_channel") and ML.resolve_leaf(x["marker_channel"])] + return sorted(set(mks)) + + +def build_all_layouts(mont_dir=None): + """Live-mode layout assets into viewer_assets_v5/_montage: the shared phase layout (gene annotations + + phase positions) + per-marker layouts repositioned by each marker's own embedding. Cheap, no SLURM.""" + from . import build_umap_montage as BM + mont_dir = mont_dir or f"{C.OUT}/viewer_assets_v5/_montage" + BM.build_layout(ML.embedding_h5ad(None), mont_dir) # shared phase layout_{umap,phate}.json (color-fields + phase positions) + build_marker_layouts(mont_dir) + + +def submit_strips(out=PCS_OUT): + """Fan out PC strips + features for all 55 markers (parallel, no external calls).""" + from ops_utils.hpc.slurm_batch_utils import submit_parallel_jobs + mks = _all_markers() + jobs = [{"name": f"pcm_{ML._norm(mk)[:18]}", "func": build_marker_job, "kwargs": {"marker_channel": mk, "out": out}} for mk in mks] + print(f"[pcm] submitting {len(jobs)} marker strip+feature jobs") + submit_parallel_jobs(jobs, experiment="pcs_markers", + slurm_params={"slurm_partition": "cpu", "cpus_per_task": 8, "mem_gb": 64, "timeout_min": 150}, + log_dir="pcs_markers", wait_for_completion=True) + + +def submit_enrich(out=PCS_OUT, n_shards=4): + """Throttled ontology enrichment for all markers (few shards so speedrichr isn't hammered).""" + from ops_utils.hpc.slurm_batch_utils import submit_parallel_jobs + mks = _all_markers() + shards = [mks[i::n_shards] for i in range(n_shards)] + jobs = [{"name": f"pcm_enrich_{i}", "func": enrich_shard, "kwargs": {"markers": s, "out": out}} for i, s in enumerate(shards) if s] + print(f"[pcm] submitting {len(jobs)} throttled enrichment shards for {len(mks)} markers") + submit_parallel_jobs(jobs, experiment="pcs_markers_enrich", + slurm_params={"slurm_partition": "cpu", "cpus_per_task": 4, "mem_gb": 16, "timeout_min": 300}, + log_dir="pcs_markers_enrich", wait_for_completion=True) + + +if __name__ == "__main__": + ap = argparse.ArgumentParser() + ap.add_argument("--marker", help="viewer marker_channel → build one marker (strips+enrich+features)") + ap.add_argument("--features-slug", help="build per-marker features.json for an already-built slug") + ap.add_argument("--submit-strips", action="store_true", help="SLURM fan-out strips+features for all 55") + ap.add_argument("--submit-enrich", action="store_true", help="SLURM throttled enrichment for all 55") + ap.add_argument("--layouts", action="store_true", help="build Live-mode layouts (shared phase + per-marker) into viewer_assets_v5/_montage") + ap.add_argument("--out", default=PCS_OUT) + a = ap.parse_args() + if a.layouts: + build_all_layouts() + elif a.submit_strips: + submit_strips(a.out) + elif a.submit_enrich: + submit_enrich(a.out) + elif a.features_slug: + build_marker_features(a.features_slug, a.out) + else: + build_marker(a.marker, a.out) diff --git a/src/ops_model/models/attention/diffex/viewer/build_phase_shap_rankings.py b/src/ops_model/models/attention/diffex/viewer/build_phase_shap_rankings.py new file mode 100644 index 0000000..f649d94 --- /dev/null +++ b/src/ops_model/models/attention/diffex/viewer/build_phase_shap_rankings.py @@ -0,0 +1,93 @@ +"""Split Alex's new shap_screen PHASE rankings into the pma_v5 parquet format the phase traversal build +consumes (single Phase2D channel → one parquet each). NON-DESTRUCTIVE: writes to pma_shap_phase_* (the +production pma_v5_phase_* are left in place) so we can validate before repointing GRAINS. + + geneKO → pma_shap_phase_geneKO.parquet (schema: gene, experiment, well, x_pheno, y_pheno, segmentation, + pma_attention, rank, rank_type) + complex → pma_shap_phase_complex.parquet (adds predicted_class=complex, gene=member gene; EBI-pooled) + + python -m ops_model.models.attention.diffex.viewer.build_phase_shap_rankings --geneko --submit + python -m ops_model.models.attention.diffex.viewer.build_phase_shap_rankings --complex --submit +""" +from __future__ import annotations + +import argparse +import os + +import pandas as pd + +M = "/hpc/projects/icd.fast.ops/models/alex_lin_attention/v5/multi_rank" +GENEKO_CSV = f"{M}/shap_screen_phase_all.csv" +EBI_CSV = f"{M}/shap_screen_ebi_phase_all.csv" +# the shap CSVs use OLD gene symbols → old-names yaml (the updated one silently drops the tRNA-synthetase complex) +EBI_YAML = "/hpc/projects/icd.fast.ops/configs/gene_clusters/EBI_complexes_v1_old_gene_names.yaml" +RANK = "/hpc/projects/icd.fast.ops/models/diffex/viewer_assets_v5/_rankings" +GENEKO_OUT = f"{RANK}/pma_shap_phase_geneKO.parquet" +COMPLEX_OUT = f"{RANK}/pma_shap_phase_complex.parquet" +TOP_N_CX = 500 + + +def _gene_to_complex(): + import yaml + y = yaml.safe_load(open(EBI_YAML)) or {} + return {str(g): v["name"] for _, v in y.items() for g in v.get("genes", [])} + + +def build_geneko(): + use = ["gene", "rank", "shap", "experiment", "well", "x_pheno", "y_pheno", "segmentation_id"] + print(f"[phase-shap] reading {GENEKO_CSV} ...", flush=True) + df = (pd.read_csv(GENEKO_CSV, usecols=use) + .rename(columns={"shap": "pma_attention", "segmentation_id": "segmentation"})) + df["rank_type"] = "top" + df = df[["gene", "experiment", "well", "x_pheno", "y_pheno", "segmentation", "pma_attention", "rank", "rank_type"]] + df.reset_index(drop=True).to_parquet(GENEKO_OUT) + print(f"[phase-shap] geneKO: {len(df):,} rows, {df.gene.nunique()} classes -> {GENEKO_OUT}") + return {"rows": int(len(df)), "classes": int(df.gene.nunique())} + + +def build_complex(): + g2c = _gene_to_complex() + use = ["gene", "rank", "shap", "experiment", "well", "x_pheno", "y_pheno", "segmentation_id"] + print(f"[phase-shap] reading {EBI_CSV} ...", flush=True) + df = pd.read_csv(EBI_CSV, usecols=use) + df = df[~df["gene"].astype(str).str.startswith("NTC")] + df["complex"] = df["gene"].astype(str).map(g2c) + unmapped = sorted(df.loc[df["complex"].isna(), "gene"].astype(str).unique()) + if unmapped: # no good reason a screened gene lacks a complex + raise ValueError(f"[phase-shap] {len(unmapped)} gene(s) unmapped in {os.path.basename(EBI_YAML)}: {unmapped}") + parts = [] + for cx, gcx in df.groupby("complex"): + g = (gcx.drop_duplicates(["experiment", "well", "x_pheno", "y_pheno"]) + .sort_values("shap", ascending=False).head(TOP_N_CX).copy()) + g["rank"] = range(1, len(g) + 1) + g["predicted_class"] = cx + parts.append(g) + o = (pd.concat(parts, ignore_index=True) + .rename(columns={"shap": "pma_attention", "segmentation_id": "segmentation"}) + [["predicted_class", "gene", "experiment", "well", "segmentation", "x_pheno", "y_pheno", + "pma_attention", "rank"]]) + o["rank_type"] = "top" + o.to_parquet(COMPLEX_OUT) + print(f"[phase-shap] complex: {len(o):,} rows, {o.predicted_class.nunique()} complexes -> {COMPLEX_OUT}") + return {"rows": int(len(o)), "complexes": int(o.predicted_class.nunique())} + + +if __name__ == "__main__": + ap = argparse.ArgumentParser() + ap.add_argument("--geneko", action="store_true") + ap.add_argument("--complex", action="store_true") + ap.add_argument("--submit", action="store_true") + a = ap.parse_args() + fns = ([build_geneko] if a.geneko else []) + ([build_complex] if a.complex else []) + if not fns: + ap.error("pass --geneko and/or --complex") + if a.submit: + from ops_utils.hpc.slurm_batch_utils import submit_parallel_jobs + submit_parallel_jobs( + jobs_to_submit=[{"name": f"phase_shap_{fn.__name__}", "func": fn, "kwargs": {}} for fn in fns], + experiment="diffex_shaprank", + slurm_params={"slurm_partition": "cpu", "cpus_per_task": 8, "mem_gb": 200, "timeout_min": 90}, + log_dir="diffex_shaprank", wait_for_completion=False) + else: + for fn in fns: + fn() diff --git a/src/ops_model/models/attention/diffex/viewer/build_phate_figure.py b/src/ops_model/models/attention/diffex/viewer/build_phate_figure.py new file mode 100644 index 0000000..b3ac3ee --- /dev/null +++ b/src/ops_model/models/attention/diffex/viewer/build_phate_figure.py @@ -0,0 +1,252 @@ +"""Figure 4 embedding — 3-panel reproduction of the paper's phase gene-PHATE (paper_v2): + E major biological-process arms + F mitochondrial sub-groups (membrane translocation/folding, electron transport, 39S mito ribosome) + G transcription / RNA-processing sub-groups (spliceosome U-snRNP & Prp19-LSm, Pol I/II/III) +Each panel: the same PHATE scatter (grey), that panel's groups colored + leader-labelled with the +single-cell generated morph (NTC cell1 → group, alpha=+5). NTC original shown top-left of panel E. + + python -m ops_model.models.attention.diffex.viewer.build_phate_figure +""" +from __future__ import annotations + +import argparse +import json +import math +import os + +import numpy as np +import matplotlib +matplotlib.use("Agg") +import matplotlib.pyplot as plt +from matplotlib.gridspec import GridSpec +from matplotlib.offsetbox import AnnotationBbox, OffsetImage, TextArea, VPacker + +plt.rcParams["pdf.fonttype"] = 42 + +VA = "/hpc/projects/icd.fast.ops/models/diffex/viewer_assets" +LAYOUT = f"{VA}/_montage/layout_phate.json" +MORPH = f"{VA}/phase/geneKO/{{gene}}/cell1/frame_16.webp" # alpha=+5 (index 16 of 17) +# NTC original = the SAME cell1 at alpha=0 (frame_08, the traversal midpoint = unmorphed base). +# Every morph above is this exact cell pushed toward its gene, so this is the honest reference. +NTC_IMG = f"{VA}/phase/geneKO/FANCC/cell1/frame_08.webp" +OUT_DIR = "/hpc/projects/icd.fast.ops/analysis/figure4_embedding" + +# each panel: {label: (color, [match substrings])}. First matching label wins within a panel. +PANEL_E = { + "spliceosome": ("#e15759", ["spliceosome", "snRNP", "snRNA", "RNA splicing", "Prp19", "LSm"]), + "RNA transcription": ("#7b4173", ["rna polymerase", "mediator", "transcription factor TFII"]), + "DNA replication": ("#bcbd22", ["DNA replication", "MCM", "replicative", "origin recognition", "replisome"]), + "mitochondria & ox. phos.": ("#8c6d31", ["mitochondrial", "electron transport", "respiratory chain", "39S mito"]), + "ER-Golgi transport": ("#2ca02c", ["ER-Golgi", "COPI", "COPII", "COP-II", "golgi", "endoplasmic reticulum", "SEC23", "SEC61"]), + "dynein motors": ("#17becf", ["dynein", "dynactin", "microtubule motor"]), + "translation initiation": ("#4c78a8", ["translational initiation", "translation factor", "eIF", "eukaryotic initiation"]), + "ribosome biogenesis": ("#e377c2", ["ribosomal subunit processome", "ribosome biogenesis", "rRNA processing", "nucleolar"]), + "60S ribosome": ("#ff7f0e", ["60S cytosolic large ribosomal", "large ribosomal subunit"]), + "40S ribosome": ("#6baed6", ["40S cytosolic small ribosomal", "small ribosomal subunit"]), + "proteasome": ("#9467bd", ["proteasome", "PA700", "ubiquitin-dependent protein catabolic"]), + "mTORC1": ("#1b9e77", ["mtorc1", "mtorc2", "mtor complex"]), +} +PANEL_F = { + "mito. membrane translocation & folding": ("#8c6d31", ["tim23", "tom complex", "mitochondrial import", "presequence translocase", "translocase of the", "chaperonin", "hsp60"]), + "electron transport chain": ("#4c78a8", ["electron transport", "respiratory chain", "atp synthase", "cytochrome c oxidase", "nadh dehydrogenase"]), + "39S mito. ribosome": ("#8bc34a", ["39s mitochondrial", "mitochondrial large ribosomal", "55S ribosome, mitochondrial"]), +} +PANEL_G = { + "spliceosome U1-5 snRNPs": ("#4c78a8", ["u1 snrnp", "u2 snrnp", "u4", "u5 snrnp", "u1-5", "u11/u12", "u2-type spliceosomal", "snrnp"]), + "spliceosome Prp19 / LSm": ("#8bc34a", ["prp19", "lsm", "nineteen complex", "intron lariat"]), + "Pol-I RNA polymerase": ("#c9a227", ["rna polymerase i complex", "polymerase i "]), + "Pol-II RNA polymerase": ("#8c6d31", ["rna polymerase ii", "polymerase ii complex"]), + "Pol-III RNA polymerase": ("#e15759", ["rna polymerase iii", "polymerase iii complex"]), +} + + +def _img(path, frac=0.80): + from PIL import Image + if not os.path.exists(path): + return None + a = np.asarray(Image.open(path).convert("L"), dtype=np.float32) / 255.0 + h, w = a.shape; ch, cw = int(h * (1 - frac) / 2), int(w * (1 - frac) / 2) + return a[ch:h - ch, cw:w - cw] + + +def _framed(a, color, bw_frac=0.045): + """Grey image -> RGB with a solid group-colored border around the FOV.""" + from matplotlib.colors import to_rgb + if a is None: + return None + rgb = np.repeat(a[:, :, None], 3, axis=2) + c = np.array(to_rgb(color)); b = max(2, int(a.shape[0] * bw_frac)) + rgb[:b, :] = c; rgb[-b:, :] = c; rgb[:, :b] = c; rgb[:, -b:] = c + return rgb + + +# hand-picked representative gene per arm (must be an EBI member of that arm with a morph) +REP_OVERRIDE = { + "mitochondria & ox. phos.": "TIMM23", + "RNA transcription": "POLR1B", + "40S ribosome": "RPS16", + "mTORC1": "MTOR", +} + +# explicit arm angle (radians, 0=+x CCW) for arms whose cluster is too central to have a natural direction +ANG_OVERRIDE = {"mTORC1": math.radians(205)} + +# per-label directional bias (span units, +x right / +y up), applied before declutter +NUDGE = { + "ER-Golgi transport": (0.34, 0.16), + "proteasome": (0.10, 0.32), + "RNA transcription": (0.22, 0.05), + "40S ribosome": (0.06, 0.00), + "DNA replication": (0.00, 0.12), +} + + +EBI_YAML = "/hpc/projects/icd.fast.ops/configs/gene_clusters/EBI_complexes_v1_updated_gene_names.yaml" +_G2C = None + + +def _ebi_complex(sym): + """Authoritative gene -> EBI complex name (310 genes); None if the gene is in no EBI complex.""" + global _G2C + if _G2C is None: + import yaml + _G2C = {} + for v in yaml.safe_load(open(EBI_YAML)).values(): + for g in (v.get("genes") or []): + _G2C.setdefault(g, v["name"]) + return _G2C.get(sym) + + +def _assign(g, panel): + ebi = _ebi_complex(g.get("g")) # EBI complex membership is the only source of truth + if not ebi: + return None # gene in no EBI complex -> grey/unassigned + e = ebi.lower() + for label, (_, subs) in panel.items(): + if any(s.lower() in e for s in subs): + return label + return None + + +def _arm_tip(pts, cen): + """Robust outer tip of an arm = mean of its 5 points farthest from the cloud center.""" + rad = np.linalg.norm(pts - cen, axis=1) + return pts[np.argsort(-rad)[:min(5, len(pts))]].mean(0) + + +def _declutter(box, cen, minsep, P, pad, iters=800, fixed=()): + """Force-separate insets; keep each box just outside the scatter's edge *in its own direction* + (hug the cloud, minimal whitespace). `fixed` boxes never move but still push the others away.""" + keys = list(box); fixed = set(fixed) + for _ in range(iters): + moved = False + for i in range(len(keys)): + for j in range(i + 1, len(keys)): + a, b = keys[i], keys[j] + v = box[a] - box[b]; d = np.hypot(*v) + if d < minsep: + u = v / (d or 1); sh = (minsep - d) / 2 + fa, fb = a in fixed, b in fixed + if fa and fb: + continue + if fb: + box[a] = box[a] + u * 2 * sh + elif fa: + box[b] = box[b] - u * 2 * sh + else: + box[a] = box[a] + u * sh; box[b] = box[b] - u * sh + moved = True + for l in keys: # keep box just past the cloud edge along its dir + if l in fixed: + continue + dirv = box[l] - cen; rb = np.hypot(*dirv); u = dirv / (rb or 1) + rmin = ((P - cen) @ u).max() + pad + if rb < rmin: + box[l] = cen + u * rmin + if not moved: + break + return box + + +def build(out_dir=None, thumb=130, minsep=0.42, ntc_scale=1.7, alpha=5): + out_dir = out_dir or OUT_DIR + os.makedirs(out_dir, exist_ok=True) + frame = max(0, min(16, int(round(8 + alpha / 5.0 * 8)))) # 17 frames span alpha=-5..+5; 8 = base + morph = MORPH.replace("frame_16", f"frame_{frame:02d}") + genes = json.load(open(LAYOUT))["genes"] + xy = {g["g"]: (g["nx"], g["ny"]) for g in genes} + ntc = np.array([[g["nx"], g["ny"]] for g in genes if str(g["g"]).startswith("NTC")]) + real = [g for g in genes if not str(g["g"]).startswith("NTC")] + P = np.array([[g["nx"], g["ny"]] for g in real]) + panel = PANEL_E + grp = {g["g"]: _assign(g, panel) for g in real} + labels = [l for l in panel if any(v == l for v in grp.values())] + + fig, ax = plt.subplots(figsize=(18, 16)) + ax.scatter(P[:, 0], P[:, 1], s=15, c="#dcdcdc", edgecolors="none", zorder=1) + for l in labels: + pts = np.array([xy[g] for g in grp if grp[g] == l]) + ax.scatter(pts[:, 0], pts[:, 1], s=43, c=panel[l][0], edgecolors="white", linewidths=0.4, zorder=3) + if len(ntc): + c = ntc.mean(0); r = 1.6 * np.percentile(np.linalg.norm(ntc - c, axis=1), 90) + 0.01 + ax.scatter(ntc[:, 0], ntc[:, 1], s=18, c="#000", edgecolors="none", zorder=5) + ax.annotate("NTCs", c + [0, -r * 1.4], ha="center", va="top", fontsize=11, color="#000", weight="bold") + + cen = P.mean(0); span = max(np.ptp(P[:, 0]), np.ptp(P[:, 1])) + pad = 0.12 * span # gap from cloud edge to inset + tip, box, rep, img, ang = {}, {}, {}, {}, {} + for l in labels: + members = [g for g in grp if grp[g] == l] + pts = np.array([xy[g] for g in members]); gc = pts.mean(0) + tip[l] = _arm_tip(pts, cen) + withm = [g for g in members if os.path.exists(morph.format(gene=g))] + ov = REP_OVERRIDE.get(l) + rep[l] = ov if ov in members and os.path.exists(morph.format(gene=ov)) \ + else min(withm or members, key=lambda g: np.hypot(*(np.array(xy[g]) - gc))) + img[l] = _img(morph.format(gene=rep[l])) + v = tip[l] - cen + ang[l] = ANG_OVERRIDE.get(l, math.atan2(v[1], v[0]) if np.hypot(*v) > 1e-6 else 0.0) + u = np.array([math.cos(ang[l]), math.sin(ang[l])]) + box[l] = cen + u * (((P - cen) @ u).max() + pad) # just beyond the cloud edge along this arm + # pin 40S & 60S together: 40S at its slot, 60S stacked just below (close, non-overlapping) + pair = [l for l in ("40S ribosome", "60S ribosome") if l in labels] + if len(pair) == 2: + box["60S ribosome"] = box["40S ribosome"] + np.array([0.0, -0.40 * span]) + for l, (dx, dy) in NUDGE.items(): # directional bias before declutter resolves overlaps + if l in box and l not in pair: + box[l] = box[l] + np.array([dx, dy]) * span + box = _declutter(box, cen, minsep * span, P, pad, fixed=pair) + for l in labels: + col = panel[l][0] + kids = [TextArea(l, textprops=dict(color=col, size=12, weight="bold", ha="center"))] + fr = _framed(img[l], col) + if fr is not None: + kids += [OffsetImage(fr, zoom=thumb / fr.shape[0]), + TextArea(rep[l], textprops=dict(color=col, size=13, weight="bold", ha="center"))] + ax.add_artist(AnnotationBbox(VPacker(children=kids, align="center", pad=0, sep=2), tip[l], xybox=box[l], + xycoords="data", boxcoords="data", frameon=False, annotation_clip=False, + arrowprops=dict(arrowstyle="-", color=col, lw=1.3, shrinkA=0, shrinkB=3))) + nimg = _framed(_img(NTC_IMG), "#111") # original NTC — bigger, top-left corner reference + if nimg is not None: + kids = [TextArea("NTC (original)", textprops=dict(color="#111", size=14, weight="bold", ha="center")), + OffsetImage(nimg, zoom=ntc_scale * thumb / nimg.shape[0])] + ax.add_artist(AnnotationBbox(VPacker(children=kids, align="center", pad=0, sep=3), (0.005, 0.995), + xycoords="axes fraction", box_alignment=(0, 1), frameon=False)) + ax.set_xticks([]); ax.set_yticks([]) + for s in ax.spines.values(): + s.set_visible(False) + ax.margins(0.42) + stem = f"phate_arms_morph_a{alpha:g}" + for ext in ("png", "svg"): + fig.savefig(f"{out_dir}/{stem}.{ext}", dpi=180, bbox_inches="tight") + plt.close(fig) + print(f"[phate-fig] single-panel figure (alpha={alpha}, frame_{frame:02d}) -> {out_dir}/{stem}.png / .svg") + + +if __name__ == "__main__": + ap = argparse.ArgumentParser() + ap.add_argument("--out", default=None) + ap.add_argument("--thumb", type=int, default=130) + ap.add_argument("--alpha", type=float, default=5) + a = ap.parse_args() + build(a.out, a.thumb, alpha=a.alpha) diff --git a/src/ops_model/models/attention/diffex/viewer/build_setacc_bins.py b/src/ops_model/models/attention/diffex/viewer/build_setacc_bins.py new file mode 100644 index 0000000..1d78e6d --- /dev/null +++ b/src/ops_model/models/attention/diffex/viewer/build_setacc_bins.py @@ -0,0 +1,62 @@ +"""Per-marker SetTransformer real-cell set-accuracy across ALL bag sizes (bins) for the viewer's +Top-cells 'show classification accuracy' overlay (adaptive bin dropdown) + the 'by SET ACC' sort. + +Same sources as build_setacc_bymarker.py but keeps every n_cells (bag size) instead of only 100. +Bins differ by modality (phase has up to 5000; fluor up to 500), so the viewer reads the per-marker +`bins` list and adapts the dropdown. + +Output (compact, bins + aligned arrays): + { ""|"phase": { "bins": [10,20,...], "acc": { "": [acc@10, acc@20, ...] } } } +Written to viewer_assets_v5/_montage/setacc_bins_bymarker.json (new file; leaves the @100-only +setacc_bymarker.json in place so older app builds keep working).""" +import json, os, re +import pandas as pd + +V5 = "/hpc/projects/icd.fast.ops/models/alex_lin_attention/v5" +OUT = "/hpc/projects/icd.fast.ops/models/diffex/viewer_assets_v5/_montage/setacc_bins_bymarker.json" +slug = lambda s: re.sub(r"[^A-Za-z0-9]", "_", str(s)) + + +def build(): + raw = {} # marker -> perturbation -> {bin: acc} + + def add(marker, pert, nb, acc): + raw.setdefault(marker, {}).setdefault(pert, {})[int(nb)] = round(float(acc), 4) + + # fluor per-marker geneKO + g = pd.read_csv(f"{V5}/fluorescence/fluor_bychannel_paperv2gene_cps_pergene.csv") + for _, r in g.iterrows(): + add(slug(r.channel_name), r.gene_name, r.n_cells, r.top1_acc) + + # fluor per-marker complex: mean member-gene acc per (channel, label, bin) + ce = pd.read_csv(f"{V5}/fluorescence/fluor_ebi_bychannel_pergene.csv") + for r in ce.groupby(["channel_name", "label_name", "n_cells"]).top1_acc.mean().reset_index().itertuples(index=False): + add(slug(r[0]), r[1], r[2], r[3]) + + # phase geneKO + p = pd.read_csv(f"{V5}/phase/eval_phase_e200_pergene_val.csv") + for _, r in p.iterrows(): + add("phase", r.gene_name, r.n_cells, r.top1_acc) + + # phase complex: mean member-gene acc per (label, bin) + pe = pd.read_csv(f"{V5}/phase/eval_phase_ebionly_e200_pergene_val.csv") + for r in pe.groupby(["label_name", "n_cells"]).top1_acc.mean().reset_index().itertuples(index=False): + add("phase", r[0], r[1], r[2]) + + out = {} + for marker, perts in raw.items(): + bins = sorted({b for d in perts.values() for b in d}) + out[marker] = {"bins": bins, "acc": {p: [d.get(b) for b in bins] for p, d in perts.items()}} + + os.makedirs(os.path.dirname(OUT), exist_ok=True) + json.dump(out, open(OUT, "w")) + nmk = len(out) - 1 + print(f"[setacc_bins] {nmk} markers + phase -> {OUT}") + print(f" phase bins={out['phase']['bins']} ({len(out['phase']['acc'])} perturbations)") + ex = next(k for k in out if k != "phase") + print(f" fluor '{ex}' bins={out[ex]['bins']} ({len(out[ex]['acc'])} perturbations)") + return out + + +if __name__ == "__main__": + build() diff --git a/src/ops_model/models/attention/diffex/viewer/build_setacc_bymarker.py b/src/ops_model/models/attention/diffex/viewer/build_setacc_bymarker.py new file mode 100644 index 0000000..04b57c6 --- /dev/null +++ b/src/ops_model/models/attention/diffex/viewer/build_setacc_bymarker.py @@ -0,0 +1,48 @@ +"""Per-marker SetTransformer real-cell set-accuracy (top1_acc @ bag=100) for the viewer's Perturbation +'by SET ACC' sort. Keys = slugified marker channel (matches attnModality()/VS slug) + "phase". +Each value = {perturbation_name: top1_acc} covering geneKO (gene_name) and complex (label_name, mean over members). +Written to viewer_assets_v5/_montage/setacc_bymarker.json.""" +import json, os, re +import pandas as pd + +V5 = "/hpc/projects/icd.fast.ops/models/alex_lin_attention/v5" +OUT = "/hpc/projects/icd.fast.ops/models/diffex/viewer_assets_v5/_montage/setacc_bymarker.json" +BAG = 100 +slug = lambda s: re.sub(r"[^A-Za-z0-9]", "_", str(s)) + + +def build(): + out = {} + + # --- fluor per-marker geneKO: {channel_slug: {gene: acc}} --- + g = pd.read_csv(f"{V5}/fluorescence/fluor_bychannel_paperv2gene_cps_pergene.csv") + g = g[g.n_cells == BAG] + for ch, sub in g.groupby("channel_name"): + out.setdefault(slug(ch), {}).update(dict(zip(sub.gene_name, sub.top1_acc.round(4)))) + + # --- fluor per-marker complex: mean member-gene acc per (channel, label_name) --- + ce = pd.read_csv(f"{V5}/fluorescence/fluor_ebi_bychannel_pergene.csv") + ce = ce[ce.n_cells == BAG] + for (ch, lab), sub in ce.groupby(["channel_name", "label_name"]): + out.setdefault(slug(ch), {})[lab] = round(float(sub.top1_acc.mean()), 4) + + # --- phase geneKO + complex --- + p = pd.read_csv(f"{V5}/phase/eval_phase_e200_pergene_val.csv") + p = p[p.n_cells == BAG] + out["phase"] = dict(zip(p.gene_name, p.top1_acc.round(4))) + pe = pd.read_csv(f"{V5}/phase/eval_phase_ebionly_e200_pergene_val.csv") + pe = pe[pe.n_cells == BAG] + for lab, sub in pe.groupby("label_name"): + out["phase"][lab] = round(float(sub.top1_acc.mean()), 4) + + os.makedirs(os.path.dirname(OUT), exist_ok=True) + json.dump(out, open(OUT, "w")) + nmk = len(out) - 1 + print(f"[setacc_bymarker] {nmk} markers + phase; e.g. phase geneKO={len(out['phase'])} entries -> {OUT}") + ex = next(k for k in out if k != "phase") + print(f" sample marker '{ex}': {list(out[ex].items())[:3]}") + return out + + +if __name__ == "__main__": + build() diff --git a/src/ops_model/models/attention/diffex/viewer/build_top_cells.py b/src/ops_model/models/attention/diffex/viewer/build_top_cells.py new file mode 100644 index 0000000..bed2423 --- /dev/null +++ b/src/ops_model/models/attention/diffex/viewer/build_top_cells.py @@ -0,0 +1,199 @@ +"""Build the PHASE 'Top Cells' tab assets: per-class top-N phenotype cells from the shap_screen rankings +(pma_shap_phase_{geneKO,complex}.parquet) — the EXACT cells the traversal directions are built from, so the +tab stays consistent with the traversals. Accuracy-only (no attention head), top_n per class by rank. + +Each ranked cell has experiment/well/x_pheno/y_pheno/segmentation → we crop it from phenotyping_v3.zarr with +the SAME 150px + blue negative cell-mask overlay as the PC tab (reuses build_pc_crops_masked) and emit: + viewer_assets_v5/top_cells/index.json {"top_n", "genes"|"complexes": {CLASS: {"accuracy": [rec...]}}} + viewer_assets_v5/top_cells/crops/.png + + python -m ops_model.models.attention.diffex.viewer.build_top_cells geneKO # SLURM crop shards + finalize + python -m ops_model.models.attention.diffex.viewer.build_top_cells complex --finalize # rebuild index only +""" +from __future__ import annotations + +import argparse +import json +import os + +from . import catalog as C +from .build_pc_crops_masked import BASE, CROP_SIZE, PHASE_CHANNEL, _crop, _is_blank, _render, _render_gray, _overlay_rgba, _zarr_patch + +TOP_N = 40 + + +def _pos_key(exp, well, x, y): + return f"{exp}__{well}__{int(round(float(x)))}__{int(round(float(y)))}" + + +def _crop_cells(attn, acc, crops_dir): + """Crop every unique cell (dedup by position across both rankings) from the zarr with the blue mask.""" + import zarr + from PIL import Image + _zarr_patch() + uniq = {} + for d in (attn, acc): + for recs in d.values(): + for c in recs: + uniq[_pos_key(c["exp"], c["well"], c["x"], c["y"])] = c + print(f"[topcells] {len(uniq)} unique cells to crop") + os.makedirs(crops_dir, exist_ok=True) + ov_dir = os.path.join(os.path.dirname(crops_dir), "overlays"); os.makedirs(ov_dir, exist_ok=True) + half = CROP_SIZE // 2 + cache, ok, blank, fail, valid = {}, 0, 0, 0, set() + for i, (pk, c) in enumerate(uniq.items()): + exp, well = c["exp"], c["well"] + key = (exp, well) + if key not in cache: + pos = f"{BASE}/{exp}/3-assembly/phenotyping_v3.zarr/{well[0]}/{well[1:]}/0" + try: + cache[key] = (zarr.open(f"{pos}/0", mode="r"), zarr.open(f"{pos}/labels/cell_seg/0", mode="r")) + except Exception as e: + cache[key] = None; print(f"[topcells] open failed {exp}/{well}: {e}") + if cache[key] is None: + fail += 1; continue + img, seg = cache[key] + x, y = int(round(c["x"])), int(round(c["y"])) + try: + phase = _crop(img, PHASE_CHANNEL, x, y, half) + if _is_blank(phase): + blank += 1; continue + Image.fromarray(_render_gray(phase)).save(f"{crops_dir}/{pk}.png") # raw grayscale (toggle-off) + Image.fromarray(_overlay_rgba(_crop(seg, None, x, y, half), half)).save(f"{ov_dir}/{pk}.png") # blue-outside overlay + ok += 1; valid.add(pk) + except Exception as e: + fail += 1 + if fail <= 5: + print(f"[topcells] crop failed {pk}: {e}") + if (i + 1) % 2000 == 0: + print(f"[topcells] {i + 1}/{len(uniq)} ok={ok} blank={blank} fail={fail}") + print(f"[topcells] crops: {ok} written, {blank} blank, {fail} failed") + return valid + + +def _recs(d, extra, valid): # attach crop filename; drop cells without a valid crop + out_d = {} + for g, cells in d.items(): + lst = [{"img": f"{_pos_key(c['exp'], c['well'], c['x'], c['y'])}.png", "ov": f"{_pos_key(c['exp'], c['well'], c['x'], c['y'])}.png", + "gene": g, "exp": c["exp"], + "well": c["well"], "x": round(c["x"], 1), "y": round(c["y"], 1), "rank": c["rank"], extra: c["score"]} + for c in cells if _pos_key(c["exp"], c["well"], c["x"], c["y"]) in valid] + if lst: + out_d[g] = lst + return out_d + + +def _merge_index(out, top_n, key, entries): + """Merge entries under index[key] ('genes' or 'complexes'), preserving the other key.""" + path = f"{out}/index.json" + idx = json.load(open(path)) if os.path.exists(path) else {"marker": "phase"} + idx["marker"] = "phase"; idx["top_n"] = top_n + idx.setdefault(key, {}).update(entries) + with open(path, "w") as f: + json.dump(idx, f) + print(f"[topcells] {key}: +{len(entries)} → {len(idx[key])} total; {os.path.getsize(path) / 1024:.0f} KB") + return path + + +def _cap(cells, top_n): + """Base-name aggregation can merge several variant labels → dedup by position, keep the top_n by score, re-rank.""" + seen, uniq = set(), [] + for c in sorted(cells, key=lambda d: -d["score"]): + pk = _pos_key(c["exp"], c["well"], c["x"], c["y"]) + if pk in seen: + continue + seen.add(pk); uniq.append(c) + if len(uniq) >= top_n: + break + return [{**c, "rank": i} for i, c in enumerate(uniq, 1)] + + +# --------------------------------------------------------------------------- +# Top-Accuracy cells (shap_screen rankings). +# +# NOTE ON BAG SIZE (read before touching this): there is NO single bag size for the v5 top +# cells. Alex's v5 set-accuracy ranking assigns ONE bag size PER CLASS — the bag at which that +# class saturates. Strong classes (HSPA5/CAPZB and every complex) rank at bag=10; a weak single +# gene KO like AACS only produces any signal at bag=500 (and even there rank-1 score ≈ 0.07). +# So each perturbation's cells come from its own bag. We take them straight from the slim viewer +# parquets (top-N by rank, already per-class-bag), which are the EXACT cells the traversal +# centroids are built from — keeping the Top-Cells tab and the traversals consistent. v5 has an +# accuracy ranking only (no attention head), so "attention" is left empty. +# --------------------------------------------------------------------------- +V5_RANK = "/hpc/projects/icd.fast.ops/models/diffex/viewer_assets_v5/_rankings" +V5_PARQUET = {"geneKO": f"{V5_RANK}/pma_shap_phase_geneKO.parquet", # NEW shap_screen phase rankings (same cells as traversals) + "complex": f"{V5_RANK}/pma_shap_phase_complex.parquet"} +V5_CLASS_COL = {"geneKO": "gene", "complex": "predicted_class"} +OUT_V5 = f"{C.OUT}/viewer_assets_v5/top_cells" +V5_RECORDS = OUT_V5 + "/_v5_records_{grain}.json" + + +def _v5_records(grain, top_n, names=None): + """{class: [top_n cell records]} straight from the slim v5 parquet (already per-class-bag ranked). + Complexes collapse to the base complex name (v4 convention) then dedup-by-position + cap to top_n.""" + import pandas as pd + ccol = V5_CLASS_COL[grain] + df = pd.read_parquet(V5_PARQUET[grain], + columns=[ccol, "experiment", "well", "x_pheno", "y_pheno", "segmentation", "pma_attention", "rank"]) + df = df[df["rank"] <= top_n * 3 if grain == "complex" else df["rank"] <= top_n] # complex: extra headroom for dedup + if grain == "complex": + df[ccol] = df[ccol].str.split(",").str[0].str.strip() + if names: + df = df[df[ccol].isin(set(names))] + recs = {} + for cls, g in df.groupby(ccol): + cells = [{"exp": r.experiment, "well": r.well, "x": float(r.x_pheno), "y": float(r.y_pheno), + "seg": str(r.segmentation), "rank": int(r.rank), "score": round(float(r.pma_attention), 5)} + for r in g.itertuples()] + recs[cls] = _cap(cells, top_n) # dedup by position, keep top_n by score, re-rank 1..N + return recs + + +def prepare_v5_records(grain, out=OUT_V5, top_n=TOP_N, names=None): + recs = _v5_records(grain, top_n, names) + os.makedirs(out, exist_ok=True) + with open(V5_RECORDS.format(grain=grain), "w") as f: + json.dump(recs, f) + print(f"[topcells-v5] {grain}: prepared {len(recs)} classes (top {top_n}) -> {V5_RECORDS.format(grain=grain)}") + return list(recs) + + +def crop_v5_shard(grain, classes, out=OUT_V5): + """SLURM job: crop this shard's classes' cells into the shared crops/ dir (additive, unique per position).""" + recs = json.load(open(V5_RECORDS.format(grain=grain))) + acc = {c: recs[c] for c in classes if c in recs} + _crop_cells({}, acc, f"{out}/crops") + return {"grain": grain, "classes": len(acc), "cells": sum(len(v) for v in acc.values())} + + +def finalize_v5_index(grain, out=OUT_V5, top_n=TOP_N): + """Build index[genes|complexes] from prepared records + whichever crops exist; accuracy-only, attention empty.""" + recs = json.load(open(V5_RECORDS.format(grain=grain))) + valid = {f[:-4] for f in os.listdir(f"{out}/crops") if f.endswith(".png")} + acc = _recs(recs, "conf", valid) + entries = {g: {"attention": [], "accuracy": acc.get(g, [])} for g in sorted(acc)} + key = "genes" if grain == "geneKO" else "complexes" + return _merge_index(out, top_n, key, entries) + + +def submit_v5(grain, out=OUT_V5, top_n=TOP_N, n_shards=24): + """Prepare records, fan out crop shards on SLURM, then finalize the index for this grain.""" + from ops_utils.hpc.slurm_batch_utils import submit_parallel_jobs + names = prepare_v5_records(grain, out, top_n) + shards = [s for s in ([names[i::n_shards] for i in range(n_shards)]) if s] + jobs = [{"name": f"topcells_v5_{grain}_{i}", "func": crop_v5_shard, "kwargs": {"grain": grain, "classes": s, "out": out}} + for i, s in enumerate(shards)] + print(f"[topcells-v5] submitting {len(jobs)} shards for {len(names)} {grain} classes (top {top_n})") + submit_parallel_jobs(jobs, experiment=f"topcells_v5_{grain}", + slurm_params={"slurm_partition": "cpu", "cpus_per_task": 8, "mem_gb": 32, "timeout_min": 120}, + log_dir=f"topcells_v5_{grain}", wait_for_completion=True) + finalize_v5_index(grain, out, top_n) + + +if __name__ == "__main__": + ap = argparse.ArgumentParser() + ap.add_argument("grain", choices=["geneKO", "complex"], help="top-accuracy cells for this grain (SLURM shards)") + ap.add_argument("--finalize", action="store_true", help="rebuild the index for this grain from existing crops") + ap.add_argument("--top-n", type=int, default=TOP_N) + a = ap.parse_args() + (finalize_v5_index if a.finalize else submit_v5)(a.grain, OUT_V5, a.top_n) diff --git a/src/ops_model/models/attention/diffex/viewer/build_umap_montage.py b/src/ops_model/models/attention/diffex/viewer/build_umap_montage.py new file mode 100644 index 0000000..8c886c9 --- /dev/null +++ b/src/ops_model/models/attention/diffex/viewer/build_umap_montage.py @@ -0,0 +1,241 @@ +"""latent-lens multiscale UMAP montage: place each gene's ALREADY-GENERATED traversal frame at its +gene-UMAP coordinate, so panning the embedding shows one cell morphed toward each neighborhood. + +Harvests the traversal CACHE (`viewer_assets////cell/frame_.webp`) — the +frames were decoded with the correct top-attention directions, so no re-embedding / no gene_bulked +centroids (those are all-cell means → ~13× too weak → collapsed morphs). Layout = the phase gene UMAP +`X_umap`. One α (frame index) per montage. +""" +from __future__ import annotations + +import os +from pathlib import Path + +import numpy as np +import anndata as ad +from PIL import Image + +from latent_lens import MontageConfig, build_montage + +from ..classifier.config import slugify +from .precompute import VIEWER_ALPHAS + +OUT = "/hpc/projects/icd.fast.ops/models/diffex" +import os +_ASSETS = os.environ.get("OPS_DIFFEX_ASSETS", "viewer_assets") # isolated v5 build → viewer_assets_v5 +ZARR_SCRATCH = f"{OUT}/_montage_zarr" # transient montage zarrs live OUTSIDE viewer_assets so they never sync to the app + + +def _embed_coords(ann, embedding, span=12.0): + """obsm X_ → coords auto-oriented so NTC sits bottom-left, rescaled to a common `span` + (so UMAP and PHATE montages are comparable size and share px_per_umap regardless of native scale).""" + c = np.asarray(ann.obsm[f"X_{embedding}"]).astype(float).copy() + pert = ann.obs["perturbation"].astype(str).values + ntc = np.array([p.startswith("NTC") for p in pert]) + lo, hi = c.min(0), c.max(0); mid = (lo + hi) / 2 + if ntc.any(): + nc = c[ntc].mean(0) + if nc[0] > mid[0]: c[:, 0] = lo[0] + hi[0] - c[:, 0] # NTC → left (small x) + if nc[1] < mid[1]: c[:, 1] = lo[1] + hi[1] - c[:, 1] # NTC → large y (bottom, since y maps top→bottom) + lo = c.min(0); s = span / max((c.max(0) - lo).max(), 1e-9) + return (c - lo) * s + + +def build_layout(h5ad, out_dir, embeddings=("umap", "phate")): + """Emit the SHARED gene layout the live viewer places cache frames onto: gene → (nx, ny) in [0,1] + + categorical color fields, one small JSON per embedding. Replaces the per-montage tile precompute — + the layout is identical across every marker/cell/α (it's the phase gene embedding), so it's built once.""" + import json + ann = ad.read_h5ad(h5ad); obs = ann.obs + color_fields = [] # same auto-detect as montage_to_tiles + for c in obs.columns: + s = obs[c] + if s.dtype.kind in "fiu": + continue + v = s.astype(str); n = v.nunique() + if 2 <= n <= 300 and v.str.startswith("[").mean() <= 0.3 and v.str.len().mean() <= 60: + color_fields.append(c) + _cat = lambda x: "" if str(x) in ("nan", "NaN", "None", "") else str(x) + os.makedirs(out_dir, exist_ok=True) + outs = [] + for emb in embeddings: + c = _embed_coords(ann, emb); lo = c.min(0); rng = c.max(0) - lo; rng[rng == 0] = 1 + genes = [] + for i, g in enumerate(obs["perturbation"]): + rec = {"g": str(g), "nx": float((c[i, 0] - lo[0]) / rng[0]), "ny": float((c[i, 1] - lo[1]) / rng[1])} + for cf in color_fields: + rec[cf] = _cat(obs[cf].iloc[i]) + genes.append(rec) + p = f"{out_dir}/layout_{emb}.json" + json.dump({"embedding": emb, "color_fields": color_fields, "genes": genes}, open(p, "w")) + print(f"[layout] {p}: {len(genes)} genes, {len(color_fields)} color fields") + outs.append(p) + return outs + + +def montage_from_cache(h5ad, out_zarr, cell=0, alpha=2.0, modality="phase", grain="geneKO", + tile=256, px_per_umap=5600, embedding="umap", border_field=None, border_width=None): + """Build the montage from the traversal cache: each gene tile = its cell-`cell`, α=`alpha` frame, + placed at the gene's position in `embedding` (obsm X_, e.g. umap or phate). + Crops kept GRAYSCALE (white category tint). crop_size=256 sharp; px_per_umap≈22×crop fills canvas. + `border_field` (an obs column, e.g. 'leiden_r4') draws a per-cell colored border keyed on that group.""" + al = list(VIEWER_ALPHAS) + ai = int(np.argmin([abs(a - alpha) for a in al])) # frame index for the requested α + a0 = int(np.argmin([abs(a) for a in al])) # α=0 frame index (shared anchor recon) + ann = ad.read_h5ad(h5ad) + coords_all = _embed_coords(ann, embedding) + gc = {str(g): coords_all[i] for i, g in enumerate(ann.obs["perturbation"])} + va = f"{OUT}/{_ASSETS}/{modality}/{grain}" + + genes, coords, srcs, ntc = [], [], [], [] + for g, xy in gc.items(): + if str(g).startswith("NTC"): # NTC is split into ~50 NTC_grp* embedding nodes + ntc.append((g, xy)); continue + if not os.path.exists(f"{va}/{slugify(g)}/cell{cell}/frame_{ai:02d}.webp"): # cheap stat, not Image.open + continue + genes.append(g); coords.append(xy); srcs.append((slugify(g), cell, ai)) + ntc_ref = srcs[0][0] if srcs else None + if ntc_ref: # NTC nodes = the SAME cell `cell` at α=0 (base NTC recon, + for g, xy in ntc: # no morph): NTC→NTC_group direction is ~0, so α=0 is exact + genes.append(g); coords.append(xy); srcs.append((ntc_ref, cell, a0)) + coords = np.asarray(coords, dtype=np.float32) + print(f"[montage] {len(srcs) - len(ntc)} genes + {len(ntc)} NTC nodes (cell {cell}, α={al[ai]:g}, {embedding})") + + def crops(i): + slug, c, fi = srcs[i] + return np.asarray(Image.open(f"{va}/{slug}/cell{c}/frame_{fi:02d}.webp").convert("L")) + + border_colors = border_groups = None + if border_field: # per-cell colored border keyed on an obs group + import matplotlib.pyplot as _plt + gv = {str(g): str(v) for g, v in zip(ann.obs["perturbation"], ann.obs[border_field])} + border_groups = np.array([gv.get(g, "") for g in genes]) + uniq = sorted({v for v in border_groups if v not in ("", "nan", "None")}, key=lambda s: (len(s), s)) + cmap = _plt.get_cmap("hsv") + border_colors = {v: tuple(cmap(i / max(1, len(uniq) - 1))[:3]) for i, v in enumerate(uniq)} + + cfg = MontageConfig(crop_size=tile, px_per_umap=px_per_umap, + border_width=border_width or max(4, tile // 40)) + build_montage(umap_coords=coords, crops=crops, categories=np.array(["geneKO"] * len(genes)), + category_colors={"geneKO": (1.0, 1.0, 1.0)}, output_path=out_zarr, # white = no tint → grayscale + labels=np.array(genes), config=cfg, border_colors=border_colors, border_groups=border_groups) + print(f"[montage] wrote {out_zarr}: {len(genes)} genes") + return out_zarr, genes + + +def build_montage_grid(h5ad, montage_dir, modality, embedding, cells, alphas, force=False): + """One SLURM job: build every (cell, alpha) montage for one (marker, embedding). Content-aware skip: + a montage is rebuilt only if the marker's geneKO cache changed after it was last built (or force), + so re-runs after the cache grows only touch what's stale — nothing redundant.""" + gk = f"{OUT}/{_ASSETS}/{modality}/geneKO" + cache_mtime = os.path.getmtime(gk) if os.path.isdir(gk) else 0 # bumps when a new gene traversal is added + outs = skipped = 0 + for cell in cells: + for a in alphas: + oz = f"{montage_dir}/{modality}_geneKO_{embedding}_cell{cell}_a{a:g}.zarr" + tj = f"{oz[:-5]}_tiles/tiles.json" + if not force and os.path.exists(tj) and os.path.getmtime(tj) >= cache_mtime: + skipped += 1; continue # montage already reflects the current cache + build_montage_web(h5ad, oz, cell=cell, alpha=a, embedding=embedding, modality=modality) + outs += 1 + print(f"[montage] {modality}/{embedding}: built {outs}, {skipped} up-to-date") + return {"modality": modality, "embedding": embedding, "built": outs, "skipped": skipped} + + +def build_montage_web(h5ad, out_zarr, cell=0, alpha=2.0, embedding="umap", modality="phase"): + """One step for the viewer: harvest the cache → montage zarr → PNG tiles + labels (in `embedding`). + `modality` selects which traversal frames to place (phase | slugified marker); the LAYOUT always + comes from the shared `h5ad` (phase gene embedding) so every marker shares the same gene positions. + `out_zarr` names the output; the served `_tiles/` go next to it (in viewer_assets), but the transient + `.zarr` intermediate is written to ZARR_SCRATCH (outside viewer_assets) and deleted after transcoding — + so no zarr ever lands in the served dir even if this job is killed mid-run.""" + import shutil + tiles_dir = str(out_zarr)[:-5] + "_tiles" # served output (viewer_assets/_montage/_tiles) + os.makedirs(ZARR_SCRATCH, exist_ok=True) + scratch_zarr = f"{ZARR_SCRATCH}/{os.path.basename(out_zarr)}" # transient zarr, outside viewer_assets + _, placed = montage_from_cache(h5ad, scratch_zarr, cell=cell, alpha=alpha, embedding=embedding, modality=modality) + tiles = montage_to_tiles(scratch_zarr, h5ad, out_dir=tiles_dir, placed=set(placed), embedding=embedding) + shutil.rmtree(scratch_zarr, ignore_errors=True) # ~225MB each × thousands — never persisted + return tiles + + +def montage_to_tiles(zarr_path, h5ad, out_dir=None, tile=512, placed=None, + embedding="umap", color_fields=None): + """Transcode the montage OME-Zarr RGB pyramid → PNG tiles + tiles.json + labels.json for the + web viewer (OpenSeadragon). Avoids blosc-in-browser; served same-origin from viewer_assets.""" + import json + import zarr + zg = zarr.open(str(zarr_path), mode="r") + at = dict(zg.attrs) + ext = at["umap_extent"] + levels = sorted((d["path"] for d in at["multiscales"][0]["datasets"]), key=int) + import shutil + out = Path(out_dir or (str(zarr_path)[:-5] + "_tiles")) + if out.exists(): + shutil.rmtree(out) # clear stale tiles from prior builds (else orphan crops persist) + out.mkdir(parents=True, exist_ok=True) + ann = ad.read_h5ad(h5ad); obs = ann.obs + # Only transcode tiles that actually contain a crop (sparse montage) — compute occupied (col,row) + # per level from the placed gene pixel positions instead of scanning the whole (mostly-empty) canvas. + coords = _embed_coords(ann, embedding) + pu = ext["px_per_umap"]; crop0 = int(at.get("crop_size", 256)) + placed = placed or set() + px0 = [((coords[i, 0] - ext["xmin"]) * pu, (coords[i, 1] - ext["ymin"]) * pu) + for i, g in enumerate(obs["perturbation"]) if str(g) in placed] + W0 = H0 = 0 + for k in levels: + a = zg[k] # (3, Y, X) float32 in [0,1] + _, Y, X = a.shape + if k == "0": + W0, H0 = X, Y + ld = out / f"L{k}"; ld.mkdir(exist_ok=True) + ds = 2 ** int(k); half = crop0 / ds / 2 + tile # +1 tile margin so we never clip a crop + occ = set() + for x, y in px0: + cx, cy = x / ds, y / ds + for col in range(max(0, int((cx - half) // tile)), int((cx + half) // tile) + 1): + for row in range(max(0, int((cy - half) // tile)), int((cy + half) // tile) + 1): + if col * tile < X and row * tile < Y: + occ.add((col, row)) + for col, row in occ: + blk = np.asarray(a[:, row * tile:(row + 1) * tile, col * tile:(col + 1) * tile]) + if not blk.any(): + continue # margin over-includes; drop the truly-empty ones + img = (np.clip(blk.transpose(1, 2, 0), 0, 1) * 255).astype(np.uint8) + Image.fromarray(img).save(ld / f"{col}_{row}.png") + # color-by fields: ALL usable categorical obs columns (auto-detected) — non-numeric, 2–300 distinct, + # not list-like, not free-text. Covers complexes, all leiden/ontology resolutions, GO/Reactome/KEGG, etc. + if color_fields is None: + color_fields = [] + for c in obs.columns: + s = obs[c] + if s.dtype.kind in "fiu": + continue + v = s.astype(str); n = v.nunique() + if 2 <= n <= 300 and v.str.startswith("[").mean() <= 0.3 and v.str.len().mean() <= 60: + color_fields.append(c) + fields = [c for c in color_fields if c in obs.columns] + (out / "tiles.json").write_text(json.dumps( + {"width": W0, "height": H0, "tileSize": tile, "levels": [int(k) for k in levels], + "embedding": f"phase gene {embedding.upper()} (CellDINO)", "color_fields": fields})) + + # gene labels: umap coord → level-0 pixel (y flipped to match the montage build), normalized by width. + # Each carries has-crop (else the viewer draws a dot) + categorical fields for color-by overlays. + coords = _embed_coords(ann, embedding) # same auto-orient + rescale as the build + placed = placed or set() + pu = ext["px_per_umap"] + def _cat(v): + s = str(v) + return "" if s in ("nan", "NaN", "None", "") else s + labels = [] + for i, g in enumerate(obs["perturbation"]): + g = str(g) # NTC included (its α=0 anchor tile is placed) + rec = {"g": g, "nx": float((coords[i, 0] - ext["xmin"]) * pu / W0), + "ny": float((coords[i, 1] - ext["ymin"]) * pu / W0), "crop": g in placed} + for c in fields: + rec[c] = _cat(obs[c].iloc[i]) if c in obs else "" + labels.append(rec) + (out / "labels.json").write_text(json.dumps(labels)) + print(f"[tiles] {out}: {len(levels)} levels, {len(labels)} labels " + f"({sum(l['crop'] for l in labels)} with crops), canvas {W0}x{H0}") + return str(out) diff --git a/src/ops_model/models/attention/diffex/viewer/catalog.py b/src/ops_model/models/attention/diffex/viewer/catalog.py new file mode 100644 index 0000000..f503476 --- /dev/null +++ b/src/ops_model/models/attention/diffex/viewer/catalog.py @@ -0,0 +1,196 @@ +"""Shared selection logic for the DiffEx viewer cache: distinctiveness matrices, per-marker +top-gene ranking, the complete-marker catalog (marker_channel + raw channel + generator ckpt), +and the dist/description maps the manifest attaches. Imported by `submit.py` so every cache +build uses the same targets — no duplicated one-off scripts. +""" +from __future__ import annotations + +import json + +import numpy as np +import pandas as pd + +from ..classifier.config import slugify + +OUT = "/hpc/projects/icd.fast.ops/models/diffex" +DD = f"{OUT}/diffae" +_DIST_BASE = ("/hpc/projects/icd.fast.ops/organelle_attribution/pca_optimized_v0.3/cell_dino/" + "zscore_per_exp/paper_v2/with_cp/with_4i") +_DIST_RELS = ["all_livecell"] # v2 with_cp/with_4i: single 56-reporter matrix (live + CP + 4i) +LAUNCH_JSON = f"{OUT}/directions/_ranking/fluor_marker_launch.json" +GENE_PANEL = "/hpc/projects/icd.fast.ops/configs/annotated_gene_panel_July2025.csv" +EBI_YAML = "/hpc/projects/icd.fast.ops/configs/gene_clusters/EBI_complexes_v1_updated_gene_names.yaml" +EBI_FLUOR_CSV = "/hpc/projects/icd.fast.ops/models/alex_lin_attention/v4/pma_fluorescent_cells_ebi_all.csv" +GENE_EMB_H5AD = ("/hpc/projects/icd.fast.ops/organelle_attribution/pca_optimized_v0.3/cell_dino/" + "zscore_per_exp/paper_v2/phase_only/fixed_80%/cosine/gene_embedding_pca_optimized.h5ad") + +# fixed-cell reporter (distinctiveness matrix col) for CP/4i marker_channels +FIXED_REP = { + "Endoplasmic Reticulum_Concanavalin A": "ER, ConA (cp)", "F-actin_Phalloidin": "f-actin, Phalloidin (cp)", + "Microtubules_Tubulin": "microtubules, Tubulin (cp)", "Mitochondria_TOMM20": "mitochondria, TOMM20 (cp)", + "Nucleus_Hoechst": "nuclei, Hoechst (cp)", "Nucleoli_NPM1": "nucleoli, NPM1 (cp)", + "Plasma Membrane_Wheat Germ Agglutinin": "plasma membrane, WGA (cp)", "NFkB_NFkB (mouse-488)": "NFkB (4i)", + "p53_p53 (mouse-488)": "p53 (4i)", "pRb_pRb (rabbit-647)": "pRb (4i)", "pS6_pS6 (rabbit-647)": "pS6 (4i)", + "p21_p21 (rabbit-647)": "p21 (4i)", "b-catenin_b-catenin (mouse-488)": "b-catenin (4i)", "c-Myc_c-Myc (mouse-488)": "c-Myc (4i)", +} +# launch-complete fluorescent markers (ep≥98) by generator-dir slug +COMPLETE_LAUNCH = { + "fluor_autophagosome_ATG101", "fluor_lysosome_LAMP1", "fluor_nucleolus_DFC_FBL", "fluor_stress_granule_G3BP1", + "fluor_microtubules_MAP4", "fluor_mitochondria_TOMM70A", "fluor_ER_NCLN", "fluor_Endoplasmic_Reticulum_Concanavalin_A", + "fluor_Fe2__FeRhoNox_live_cell_dye", "fluor_Microtubules_Tubulin", "fluor_Mitochondria_TOMM20", "fluor_5xUPRE", + "fluor_ChromaLIVE_488_excitation", "fluor_ER_Golgi_COPE", "fluor_ER_Golgi_COP_II_SEC23A", "fluor_ER_SEC61B", + "fluor_ER_golgi_bridge_VAPA", "fluor_F_actin_Phalloidin", "fluor_cell_proliferation_marker_MKI67", + "fluor_Nucleus_Hoechst", "fluor_Nucleoli_NPM1", "fluor_Plasma_Membrane_Wheat_Germ_Agglutinin", "fluor_NFkB_NFkB__mouse_488", +} +# extra fluorescent markers not in launch.json (recovered from training pickles / built after the ranking CSV): +# (generator dir, marker_channel, raw channel). complete_markers() gates each on a trained checkpoint. cisGolgi/ +# VIM/LMNB1 are first-class here (they have DiffAE checkpoints + v5 fluor rankings like every other marker). +EARLY_MARKERS = [ + ("fluor_NucleoLive", "nucleus_NucleoLIVE Live Cell dye", "mCherry"), + ("fluor_NPM3", "nucleolus-GC_NPM3", "GFP"), + ("fluor_FastAct", "actin filament_FastAct_SPY555 Live Cell Dye", "mCherry"), + ("fluor_LysoTracker", "lysosome_LysoTracker live-cell dye", "GFP"), + ("fluor_ChromaLIVE_mito", "mitochondria_ChromaLIVE 561 excitation", "mCherry"), + ("fluor_cisGolgi_mStayGold", "cis-Golgi_mStayGold-CENPRaltORF", "GFP"), + ("fluor_VIM", "intermediate filaments_VIM", "GFP"), + ("fluor_LMNB1", "laminin_LMNB1", "GFP"), +] + + +def ebi_complexes(): + """Canonical EBI complex names — from the config yaml (do NOT re-derive from cells).""" + import yaml + y = yaml.safe_load(open(EBI_YAML)) or {} + return [v["name"] for v in y.values() if isinstance(v, dict) and v.get("name")] + + +def all_genes(): + """All ~1000 geneKO classes (dist-matrix index, excl NTC).""" + return [g for g in dist_matrix().index if not str(g).startswith("NTC")] + + +def dist_matrix(): + return pd.concat([pd.read_csv(f"{_DIST_BASE}/{v}/fixed_80%/cosine/plots/marker_overlay/" + "gene_reporter_distinctiveness_raw.csv", index_col=0) for v in _DIST_RELS], + axis=1) + + +COMPLEX_EBI_MAP_CSV = f"{OUT}/complex_reporter_ebi_map.csv" + + +def complex_dist(): + """Complex(name) × reporter EBI mAP — the copairs/EBI-Complex-Portal metric computed per-marker + by `build_complex_ebi_map` (NOT the chad_consistency file). Columns = reporters (match dist_matrix).""" + return pd.read_csv(COMPLEX_EBI_MAP_CSV, index_col=0) + + +def _assignment(dist): + d = dist.drop(index=[g for g in dist.index if str(g).startswith("NTC")], errors="ignore") + vals = d.values; order = np.argsort(-vals, axis=1, kind="stable") + bmap = np.take_along_axis(vals, order[:, :1], axis=1).ravel() + smap = np.take_along_axis(vals, order[:, 1:2], axis=1).ravel() + return pd.DataFrame({"gene": d.index, "best": d.columns[order[:, 0]], "bestmap": bmap, "margin": bmap - smap}), d + + +def top_genes(dist, rep, n=8): + """Top-n marker-specific genes for a distinctiveness reporter (best-marker assignment, margin-broken).""" + asg, d = _assignment(dist) + c = asg[asg.best == rep].sort_values(["bestmap", "margin"], ascending=False) + g = list(c["gene"][:n]) + if len(g) < n and rep in d: + g += [x for x in d[rep].dropna().sort_values(ascending=False).index if x not in g][:n - len(g)] + return g + + +def rep_of(dist, marker_channel): + """Distinctiveness-matrix column (reporter) for a marker_channel (live: normalized; fixed: hand map).""" + return {c.replace(", ", "_"): c for c in dist.columns}.get(marker_channel) or FIXED_REP.get(marker_channel) + + +def complete_markers(min_ep=0): + """[(generator_dir, marker_channel, raw_channel)] for fluorescent markers with a trained generator + (diffae_best.pt + train_state on disk). Default min_ep=0 = no epoch gate — epoch count is NOT a quality + signal (cond_ratio peaks ~ep40-55 then declines; diffae_best.pt banks the peak). Pass min_ep only to + re-impose a floor. `submit seed` auto-includes any marker with a checkpoint.""" + import os + import torch + launch = json.load(open(LAUNCH_JSON)) + cand = {"fluor_" + slugify(e["marker"]): (e["marker"], e["channel"]) for e in launch} + for d, mc, ch in EARLY_MARKERS: + cand[d] = (mc, ch) + out = [] + for d, (mc, ch) in cand.items(): + sp = f"{DD}/{d}/diffae_train_state.pt" + if not os.path.exists(f"{DD}/{d}/diffae_best.pt") or not os.path.exists(sp): + continue + try: + ep = torch.load(sp, map_location="cpu", mmap=True).get("epoch", 0) + except Exception: + ep = 0 + if ep >= min_ep: + out.append((d, mc, ch)) + return out + + +def desc_map(): + """{class_name: description} — per-gene function + known systems (GO/Reactome/KEGG/CORUM) pulled from + the gene-embedding h5ad (populated for all ~1000 genes, unlike the sparse panel) + complex members (EBI). + Sections are joined with ' || ' so the viewer can render each as a labeled block.""" + out = {} + try: + import anndata as ad + ob = ad.read_h5ad(GENE_EMB_H5AD, backed="r").obs + flds = [("LongName", None), ("go_bp_term", "GO biological process"), ("go_cc_term", "GO cellular component"), + ("reactome_term", "Reactome"), ("kegg_term", "KEGG"), ("corum_complex_term", "CORUM complex")] + def _v(i, c): + s = str(ob[c].iloc[i]) if c in ob.columns else "" + return "" if s in ("nan", "None", "") else s.strip() + for i, g in enumerate(ob["perturbation"].astype(str)): + if g.startswith("NTC"): + continue + parts = [(v if lbl is None else f"{lbl}: {v}") for c, lbl in flds if (v := _v(i, c))] + if parts: + out[g] = " || ".join(parts) + except Exception as e: + print("gene desc failed:", e) + try: + import yaml + for _, v in (yaml.safe_load(open(EBI_YAML)) or {}).items(): + if isinstance(v, dict) and v.get("name"): + mem = list(v.get("genes", []) or []) + out[v["name"]] = (f"Members ({len(mem)}): " + ", ".join(mem)) if mem else v["name"] + except Exception as e: + print("complex desc failed:", e) + return out + + +def dist_map_for_assets(viewer_assets): + """{(modality, grain, slug): distinctiveness mAP} for every rendered target (manifest sorting).""" + import glob + dist = dist_matrix() + dd = dist.drop(index=[g for g in dist.index if str(g).startswith("NTC")], errors="ignore") + try: + cx = complex_dist() # complex × reporter EBI mAP (complexes aren't in the gene matrix) + except Exception: + cx = None + try: + mb = json.load(open(f"{viewer_assets}/_minibinder_meta.json")) # minibinders → per-binder cell_score + except Exception: + mb = {} + out = {} + for mj in glob.glob(f"{viewer_assets}/*/*/*/meta.json"): + m = json.load(open(mj)) + rep = (rep_of(dist, m["marker_channel"]) if m.get("marker_channel") else "Phase") + if m["grain"] == "minibinder": # minibinders → cell_score (no mAP) + if m["slug"] in mb: + out[(m["modality"], m["grain"], m["slug"])] = float(mb[m["slug"]]["cell_score"]) + elif m["grain"] == "complex": # complexes → EBI complex mAP + if cx is not None and m["target"] in cx.index and rep in cx.columns: + v = cx.at[m["target"], rep] + if pd.notna(v): + out[(m["modality"], m["grain"], m["slug"])] = float(v) + elif rep in dd.columns and m["target"] in dd.index: # geneKO → distinctiveness mAP + v = dd.at[m["target"], rep] + if pd.notna(v): + out[(m["modality"], m["grain"], m["slug"])] = float(v) + return out diff --git a/src/ops_model/models/attention/diffex/viewer/deploy/README.md b/src/ops_model/models/attention/diffex/viewer/deploy/README.md new file mode 100644 index 0000000..342ae0f --- /dev/null +++ b/src/ops_model/models/attention/diffex/viewer/deploy/README.md @@ -0,0 +1,24 @@ +# DiffEx viewer — Argus/S3 deploy + +Review copies of the infra PR to `chanzuckerberg/sfbiohub-infra` (write access granted). +Argus is stateless → viewer assets live in S3, the deployment downloads on boot. + +- [diffex-viewer-dev.tf](diffex-viewer-dev.tf) — the one new file the PR adds: `diffex-viewer-dev` + S3 bucket + read-only IRSA role for the nonprod Argus cluster (`argus-diffex-viewer-rdev`/`diffex-viewer`) + + a readwrite uploader role for our `aws s3 sync`. Mirrors `proteohub-argus-s3-reader-dev.tf`. +- [PR_BODY.md](PR_BODY.md) — PR description (requests a 1 TB ceiling). + +These are copies for in-workspace review. The actual PR working copy (branch `diffex-viewer-dev`) is a +clone of the infra repo at `/hpc/mydata/gav.sturm/sfbiohub-infra` — push from there: + +```bash +cd /hpc/mydata/gav.sturm/sfbiohub-infra +git add terraform/accounts/biohub-nonprod/diffex-viewer-dev.tf +git commit -m "Add diffex-viewer-dev S3 bucket + Argus read-only role" +git push -u origin diffex-viewer-dev +gh pr create --repo chanzuckerberg/sfbiohub-infra --base main \ + --title "diffex-viewer-dev: S3 bucket + Argus read-only" --body-file .diffex_pr_body.md +``` + +Confirm the Argus namespace/SA (`argus-diffex-viewer-rdev`/`diffex-viewer`) matches the registered +deployment before pushing — that comes from the app-readiness step (`czbiohub-sf/biohub-argus-example-app`). diff --git a/src/ops_model/models/attention/diffex/viewer/marker_leaves.py b/src/ops_model/models/attention/diffex/viewer/marker_leaves.py new file mode 100644 index 0000000..0b95b18 --- /dev/null +++ b/src/ops_model/models/attention/diffex/viewer/marker_leaves.py @@ -0,0 +1,83 @@ +"""Resolve a viewer marker (manifest `marker_channel`) → its paper_v2 per-marker embedding leaf. + +55 of 58 viewer markers have an isolated single-channel embedding under paper_v2/markers/: + live fluorescent markers//all_livecell/fixed_80%/cosine/ + Cell Painting markers/_cp/only_cp/all_livecell/fixed_80%/cosine/ + 4i markers/_4i/only_4i/all_livecell/fixed_80%/cosine/ +Phase is separate: paper_v2/phase_only/fixed_80%/cosine/. +(NFkB, Rb, gH2AX have no leaf — they keep the phase embedding.) + +Each leaf holds gene_embedding_pca_optimized.h5ad (genes×101 PCs, obsm X_umap/X_phate/X_pca) + +per_signal/_cells.h5ad (per-cell) + leiden_cache.pkl + metrics/. +""" +from __future__ import annotations + +import os +import re + +PAPER_V2 = "/hpc/projects/icd.fast.ops/organelle_attribution/pca_optimized_v0.3/cell_dino/zscore_per_exp/paper_v2" +MARKERS_ROOT = f"{PAPER_V2}/markers" +PHASE_LEAF = f"{PAPER_V2}/phase_only/fixed_80%/cosine" + +# viewer marker_channel names whose CP leaf is spelled differently than the auto-normalizer would guess +_ALIAS = { + "Endoplasmic Reticulum_Concanavalin A": "ER_ConA_cp", + "Nucleus_Hoechst": "nuclei_Hoechst_cp", + "Plasma Membrane_Wheat Germ Agglutinin": "plasma_membrane_WGA_cp", +} + + +def _norm(s): + return re.sub(r"[^a-z0-9]", "", s.lower()) + + +def _leaf_base(leaf): + return re.sub(r"_(cp|4i)$", "", leaf) + + +def build_map(markers_root=MARKERS_ROOT): + """{viewer_marker_channel: leaf_dir_name} for every viewer marker that has a paper_v2 leaf.""" + leaves = sorted(os.listdir(markers_root)) + by_base = {_norm(_leaf_base(l)): l for l in leaves} + return leaves, by_base + + +def resolve_leaf(marker_channel, markers_root=MARKERS_ROOT): + """Viewer marker_channel → leaf dir name (or None if no per-marker embedding exists).""" + if marker_channel in _ALIAS: + return _ALIAS[marker_channel] + _, by_base = build_map(markers_root) + for key in (_norm(marker_channel), _norm(marker_channel.split("_")[0])): # full, then compartment token (4i = _) + if key in by_base: + return by_base[key] + return None + + +def leaf_dir(leaf, markers_root=MARKERS_ROOT): + """Leaf name → its cosine embedding directory (handles the only_cp / only_4i / all_livecell nesting).""" + sub = "only_cp/all_livecell" if leaf.endswith("_cp") else "only_4i/all_livecell" if leaf.endswith("_4i") else "all_livecell" + return f"{markers_root}/{leaf}/{sub}/fixed_80%/cosine" + + +def embedding_h5ad(marker_channel): + """Viewer marker_channel → its gene_embedding_pca_optimized.h5ad path (phase falls back to the phase leaf).""" + if not marker_channel or marker_channel.lower() in ("phase", "phase2d"): + return f"{PHASE_LEAF}/gene_embedding_pca_optimized.h5ad" + leaf = resolve_leaf(marker_channel) + return f"{leaf_dir(leaf)}/gene_embedding_pca_optimized.h5ad" if leaf else None + + +if __name__ == "__main__": + import json + m = json.load(open("/hpc/projects/icd.fast.ops/models/diffex/viewer_assets/manifest.json")) + n = miss = 0 + for x in m["markers"]: + mc = x.get("marker_channel") + if not mc: + continue + p = embedding_h5ad(mc) + ok = bool(p and os.path.exists(p)) + n += ok + if not ok: + miss += 1; print(f" no leaf: {mc}") + print(f"{n} markers resolve to an existing embedding h5ad; {miss} without") diff --git a/src/ops_model/models/attention/diffex/viewer/mimic_alex_embed.py b/src/ops_model/models/attention/diffex/viewer/mimic_alex_embed.py new file mode 100644 index 0000000..c5421f9 --- /dev/null +++ b/src/ops_model/models/attention/diffex/viewer/mimic_alex_embed.py @@ -0,0 +1,461 @@ +"""Mimic Alex's CellDINO extraction so our cells land in the SetTransformer classifier's +input space (see katamari extract_embeddings_phase_celldino_fast.yaml + CellDinoWrapper). + +Pipeline per cell: 128x128 Phase2D crop @ (x_pheno,y_pheno) → seg-mask (cell_seg==id) + → percentile-norm (x-p1)/(p99-p1) → CellDINO (resize224 + per-image z-score, our embed_crops) + → z-standardize per (channel,experiment) on NTC-control stats. + +The CellDINO forward is identical to our existing `embed_crops`; the added mask + percentile are +cheap pixel ops, and crops are chunk-local windowed reads — so cost/cell ≈ the CellDINO forward +we already pay. `validate()` reproduces a few genes in one experiment, scores bags with the real +classifier, and prints accuracy + timing/cell. +""" +from __future__ import annotations + +import time +import types + +import numpy as np +import pandas as pd +import torch +import zarr + +from .set_classifier import load_set_classifier, score_bags + +PT_PHASE = "/hpc/projects/icd.fast.ops/models/alex_lin_attention/v4/val_ops_zstdcontrol_cdino_v2" +ZARR = "/hpc/projects/icd.fast.ops/{exp}/3-assembly/phenotyping_v3.zarr" +SIZE = 128 +PHASE_CH = 0 # Phase2D channel index +PCT_LEVEL = "4" # pyramid level for the per-well p1/p99 estimate + + +def flatten_pt(gene, pt_root=PT_PHASE): + """Alex's per-gene .pt cell_metadata (bags of lists) → (flat per-cell df, his per-cell embeddings + aligned to df rows). His embeddings are the ground-truth target space (already control-z-std).""" + d = torch.load(f"{pt_root}/{gene}.pt", map_location="cpu") + cm = d["cell_metadata"] + cols = ["experiment", "well", "x_pheno", "y_pheno", "segmentation_id"] + flat = {c: [v for bag in cm[c] for v in bag] for c in cols} + df = pd.DataFrame(flat) + df["gene"] = gene + alex = np.asarray(d["embeddings"]) + if len(alex) != len(df): + print(f" [warn] {gene}: {len(alex)} embs vs {len(df)} flattened cells — alignment off") + return df, alex + + +def _well_pos(root, well): + """well 'A3' -> position group 'A/3/0'.""" + return root[f"{well[0]}/{well[1:]}/0"] + + +def _pct(pos): + """per-well (p1, p99) of Phase2D from a low-res pyramid level (one cheap full read).""" + lo = np.asarray(pos[PCT_LEVEL][0, PHASE_CH, 0]) + return np.percentile(lo, [1, 99]) + + +def load_raw(exp, cells, size=SIZE): + """cells for ONE experiment → (raw crops (N,1,s,s), masks (N,s,s) bool, per-well (p1,p99), keep). + Returns the components uncomposed so `compose()` can build mask/percentile variants.""" + root = zarr.open(ZARR.format(exp=exp), mode="r") + h = size // 2 + crops, masks, pcts, keep = [], [], [], [] + for well, g in cells.groupby("well"): + pos = _well_pos(root, well) + img = pos["0"]; seg = pos["labels/cell_seg/0"] + Y, X = img.shape[-2:] + p1, p99 = _pct(pos) + for idx, r in g.iterrows(): + y, x = int(round(r.y_pheno)), int(round(r.x_pheno)) + if y - h < 0 or x - h < 0 or y + h > Y or x + h > X: + continue + crop = np.asarray(img[0, PHASE_CH, 0, y - h:y + h, x - h:x + h]).astype(np.float32) + m = np.asarray(seg[0, 0, 0, y - h:y + h, x - h:x + h]) == int(r.segmentation_id) + crops.append(crop); masks.append(m); pcts.append((p1, p99)); keep.append(idx) + return (np.stack(crops), np.stack(masks), np.array(pcts), keep) if crops else (None, None, None, []) + + +def compose(crops, masks, pcts, mask=True, pct="well"): + """Build (N,1,s,s) from raw crops. mask: apply seg mask. pct: 'well'|'crop'|'none' intensity norm.""" + x = crops.copy() + if mask: + x = x * masks + if pct == "well": + p1, p99 = pcts[:, 0][:, None, None], pcts[:, 1][:, None, None] + x = (x - p1) / (p99 - p1 + 1e-6) + elif pct == "crop": + p1 = np.percentile(crops.reshape(len(x), -1), 1, axis=1)[:, None, None] + p99 = np.percentile(crops.reshape(len(x), -1), 99, axis=1)[:, None, None] + x = (x - p1) / (p99 - p1 + 1e-6) + return x[:, None].astype(np.float32) + + +def load_masked_crops(exp, cells, size=SIZE): + crops, masks, pcts, keep = load_raw(exp, cells, size) + if crops is None: + return np.zeros((0, 1, size, size), np.float32), [] + return compose(crops, masks, pcts), keep + + +def _celldino(crops, batch=256): + """(N,1,H,W) → (N,1024) via the SAME CellDinoModel embed_crops uses. Returns (embs, sec/cell).""" + from ops_model.models.cell_dino import CellDinoModel + model = CellDinoModel(z_score=True) + embs = [] + t0 = time.time() + with torch.inference_mode(): + for i in range(0, len(crops), batch): + out = model.extract_features({"data": torch.as_tensor(crops[i:i + batch])}) + embs.append(out.float().cpu().numpy()) + e = np.concatenate(embs).astype(np.float32) + return e, (time.time() - t0) / max(len(crops), 1) + + +def zstd_control(embs, control_mask): + """z-standardize per feature using CONTROL-only mean/std (matches z_standardize_control_only).""" + ctrl = embs[control_mask] + mu, sd = ctrl.mean(0), ctrl.std(0) + 1e-6 + return (embs - mu) / sd + + +def validate(exp="ops0031_20250424", genes=("HSPA5", "KIF11", "POLR1B", "TIMM23"), + per_gene=150, run="miwkg1cy"): + """Reproduce cells for a few genes (+NTC for control stats) in ONE experiment, embed via the + mimic, z-std on NTC, score bags → per-gene argmax accuracy + timing/cell.""" + m, g2i, c2i = load_set_classifier(run) + i2g = {v: k for k, v in g2i.items()} + genes = list(genes) + (["NTC"] if "NTC" in g2i else []) + + frames, alex_rows = [], [] + io_t0 = time.time() + for gene in genes: + df, alex = flatten_pt(gene) + mask_exp = (df.experiment == exp).to_numpy() + df = df[mask_exp].head(per_gene); alex = alex[mask_exp][:per_gene] + if not len(df): + print(f" {gene}: no cells in {exp} — skip"); continue + crops, keep = load_masked_crops(exp, df) + if not len(crops): + print(f" {gene}: all crops OOB — skip"); continue + pos = [df.index.get_loc(i) for i in keep] # positional idx into df/alex + sub = df.loc[keep].copy(); sub["_crops"] = list(crops) + frames.append(sub); alex_rows.append(alex[pos]) + print(f" {gene}: {len(sub)} cells loaded") + all_df = pd.concat(frames, ignore_index=True) + alex_emb = np.concatenate(alex_rows) # his (control-z-std) embeddings, aligned + io_per = (time.time() - io_t0) / len(all_df) + + crops = np.stack(list(all_df["_crops"])) + raw, cd_per = _celldino(crops) + embs = zstd_control(raw, (all_df.gene == "NTC").to_numpy()) + + # fidelity: my reproduction vs Alex's SAME-cell embedding (cosine) + def _cos(a, b): + a = a / (np.linalg.norm(a, axis=1, keepdims=True) + 1e-9) + b = b / (np.linalg.norm(b, axis=1, keepdims=True) + 1e-9) + return (a * b).sum(1) + cos_all = _cos(embs, alex_emb) + print(f"\n[fidelity] mean cosine(my zstd, Alex .pt) = {cos_all.mean():.3f} " + f"(p50 {np.median(cos_all):.3f}, >0.9: {(cos_all>0.9).mean()*100:.0f}%)") + + print(f"\n[timing] I/O+mask+pct: {io_per*1000:.1f} ms/cell | CellDINO: {cd_per*1000:.1f} ms/cell " + f"| total ~{(io_per+cd_per)*1000:.1f} ms/cell (N={len(all_df)})") + print(f"[scale] ~{(io_per+cd_per):.3f} s/cell → 119k cells ≈ {(io_per+cd_per)*119000/3600:.1f} GPU-hr\n") + + rng = np.random.default_rng(0) + ci = c2i.get("Phase", 0) + print(f"{'gene':8s} | {'mine argmax':12s} P(t) | {'Alex argmax':12s} P(t) | cos") + for gene in genes: + gi = np.where((all_df.gene == gene).to_numpy())[0] + if len(gi) < 20: + continue + sel = [rng.choice(gi, min(100, len(gi))) for _ in range(6)] + pm = score_bags(m, np.stack([embs[s] for s in sel]), channel_idx=ci) + pa = score_bags(m, np.stack([alex_emb[s] for s in sel]), channel_idx=ci) + tm, ta = i2g[int(pm.mean(0).argmax())], i2g[int(pa.mean(0).argmax())] + print(f" {gene:6s} | {tm:12s} {pm[:, g2i[gene]].mean():.3f} {'HIT' if tm==gene else ' '} " + f"| {ta:12s} {pa[:, g2i[gene]].mean():.3f} {'HIT' if ta==gene else ' '} " + f"| {cos_all[gi].mean():.3f}") + + +def _crop_on_centroid(masked, masks, size=128): + """Crop size×size CENTERED ON each mask's centroid (matches training crops centered on the cell + at x_pheno,y_pheno) — not the frame center. Returns (cropped masked imgs, cropped masks).""" + from scipy import ndimage as ndi + H, W = masked.shape[-2:]; h = size // 2 + oi, om = [], [] + for im, mk in zip(masked, masks): + cy, cx = ndi.center_of_mass(mk) if mk.any() else (H / 2, W / 2) + cy = int(np.clip(round(cy), h, H - h)); cx = int(np.clip(round(cx), h, W - h)) + oi.append(im[cy - h:cy + h, cx - h:cx + h]); om.append(mk[cy - h:cy + h, cx - h:cx + h]) + return np.stack(oi), np.stack(om) + + +def cellpose_masks(crops01, diameter=None, flow_threshold=0.4, batch=64): + """Segment GENERATED phase crops (N,H,W) in [0,1] with Cellpose-SAM → central-cell boolean masks. + The morphed cell is centered, so keep the label covering the crop centre (fallback: nearest label). + Lets us mask fake images the same way training masks real cells (the load-bearing step).""" + import torch + from cellpose import models + from scipy import ndimage as ndi + m = models.CellposeModel(gpu=True, device=torch.device("cuda")) + H, W = crops01.shape[-2:]; cy, cx = H // 2, W // 2 + out = [] + for i in range(0, len(crops01), batch): + labs = m.eval(list(crops01[i:i + batch]), diameter=diameter, flow_threshold=flow_threshold)[0] + for lab in labs: + c = lab[cy, cx] + if c == 0 and lab.max() > 0: # centre is background → nearest cell + idx = ndi.distance_transform_edt(lab == 0, return_distances=False, return_indices=True) + c = lab[tuple(v[cy, cx] for v in idx)] + if c > 0: + out.append(lab == c) + else: # cellpose found nothing → central-disk fallback + yy, xx = np.ogrid[:H, :W] + out.append((yy - cy) ** 2 + (xx - cx) ** 2 <= (min(H, W) * 0.35) ** 2) + return np.stack(out) + + +def test_cellpose(genes=("HSPA5", "POLR1B", "TIMM23"), n=20, alpha_frame=8): + """Load real GENERATED α-frames from the viewer cache, Cellpose-SAM segment them, keep the + central cell → report mask coverage + timing/cell (feasibility of masking fake images).""" + from PIL import Image + cache = "/hpc/projects/icd.fast.ops/models/diffex/viewer_assets/phase/geneKO" + crops = [] + for g in genes: + for c in range(n): + f = f"{cache}/{g}/cell{c}/frame_{alpha_frame:02d}.webp" + try: + crops.append(np.asarray(Image.open(f).convert("L"), np.float32) / 255.0) + except Exception: + pass + crops = np.stack(crops) + print(f"loaded {len(crops)} generated crops {crops.shape[1:]}") + t0 = time.time() + masks = cellpose_masks(crops) + per = (time.time() - t0) / len(crops) + cov = masks.reshape(len(masks), -1).mean(1) + nonempty = (cov > 0).mean() + print(f"[cellpose] {per*1000:.1f} ms/cell | central-cell found: {nonempty*100:.0f}% | " + f"mean coverage {cov.mean()*100:.1f}% of the {crops.shape[1]}px crop") + print(f"[scale] {per:.3f} s/cell → 119k ≈ {per*119000/3600:.1f} GPU-hr for masking") + + +def _zstd_per_exp(embs, exps, is_ntc): + """z-standardize per experiment on that experiment's NTC control (fallback: all cells).""" + out = embs.copy() + for e in np.unique(exps): + me = exps == e + ctrl = embs[me & is_ntc] + if len(ctrl) < 10: + ctrl = embs[me] + out[me] = (embs[me] - ctrl.mean(0)) / (ctrl.std(0) + 1e-6) + return out + + +def validate_bags(genes=("HSPA5", "POLR1B", "KIF11", "TIMM23"), per_exp=30, max_exp=25, + bag=100, draws=20, run="miwkg1cy"): + """FIDELITY-AT-SCALE test: reproduce cells across MANY experiments (per-exp z-std), then compare + mimic vs Alex's own embeddings at realistic bag sizes — does 0.91 fidelity give correct argmax?""" + m, g2i, c2i = load_set_classifier(run) + i2g = {v: k for k, v in g2i.items()} + glist = list(genes) + (["NTC"] if "NTC" in g2i else []) + crops_l, mk_l, pc_l, alex_l, gcol, ecol = [], [], [], [], [], [] + for gene in glist: + df, alex = flatten_pt(gene) + # cap per experiment, cap #experiments — keeps I/O bounded + df = df.reset_index(drop=True) + parts = [] + for e, g in df.groupby("experiment"): + parts.append(g.head(per_exp)) + if len(parts) >= max_exp: + break + df = pd.concat(parts); alex = alex[df.index.to_numpy()] + for e, g in df.groupby("experiment"): + c, mk, pc, keep = load_raw(e, g) + if c is None: + continue + pos = [g.index.get_loc(i) for i in keep] + crops_l.append(c); mk_l.append(mk); pc_l.append(pc) + alex_l.append(alex[[df.index.get_loc(i) for i in keep]]) + gcol += [gene] * len(keep); ecol += [e] * len(keep) + print(f" {gene}: {sum(x==gene for x in gcol)} cells across experiments") + crops = np.concatenate(crops_l); masks = np.concatenate(mk_l); pcts = np.concatenate(pc_l) + alex_emb = np.concatenate(alex_l); gcol = np.array(gcol); ecol = np.array(ecol) + + raw, per = _celldino(compose(crops, masks, pcts, mask=True, pct="none")) + print(f"[timing] {per*1000:.1f} ms/cell CellDINO, N={len(raw)}") + is_ntc = gcol == "NTC" + mine = _zstd_per_exp(raw, ecol, is_ntc) + alex_z = alex_emb # already control-z-std by Alex + cos = ((mine/ (np.linalg.norm(mine,axis=1,keepdims=True)+1e-9)) * + (alex_z/(np.linalg.norm(alex_z,axis=1,keepdims=True)+1e-9))).sum(1) + print(f"[fidelity] mean cos(mine, Alex) = {cos.mean():.3f}\n") + + ci = c2i.get("Phase", 0); rng = np.random.default_rng(0) + print(f"{'gene':8s} | {'MINE hit-rate meanP':22s} | {'ALEX hit-rate meanP':22s}") + for gene in genes: + gi = np.where(gcol == gene)[0] + if len(gi) < bag: + print(f" {gene}: only {len(gi)} cells (<{bag}) — skip"); continue + def _eval(E): + bags = np.stack([E[rng.choice(gi, bag, replace=False)] for _ in range(draws)]) + p = score_bags(m, bags, channel_idx=ci) + hits = np.mean([i2g[int(r.argmax())] == gene for r in p]) + return hits, p[:, g2i[gene]].mean() + hm, pm = _eval(mine); ha, pa = _eval(alex_z) + print(f" {gene:6s} | {hm*100:5.0f}% P={pm:.3f} | {ha*100:5.0f}% P={pa:.3f}") + + +CTRL_REF = "/hpc/projects/icd.fast.ops/models/diffex/control_ref_phase.npz" +CACHE = "/hpc/projects/icd.fast.ops/models/diffex/viewer_assets" + + +def build_control_ref(per_exp=40, max_exp=25, out=CTRL_REF): + """Global NTC control reference (per-feature mean/std of mimic-embedded REAL NTC cells) → used to + z-standardize GENERATED cells into the classifier space (they have no experiment of their own).""" + import os + if os.path.exists(out): + d = np.load(out); return d["mu"], d["sd"] + df, _ = flatten_pt("NTC") + parts = [g.head(per_exp) for _, g in df.groupby("experiment")][:max_exp] + df = pd.concat(parts) + crops = [] + for e, g in df.groupby("experiment"): + c, mk, pc, keep = load_raw(e, g) + if c is not None: + crops.append(compose(c, mk, pc, mask=True, pct="none")) + x = np.concatenate(crops) + raw, _ = _celldino(x) + mu, sd = raw.mean(0), raw.std(0) + 1e-6 + np.savez(out, mu=mu, sd=sd) + print(f"[control_ref] {len(raw)} NTC cells -> {out}") + return mu, sd + + +def _load_gen(gene, grain="geneKO", modality="phase", n_cells=20): + """Load a traversal's cache α-frames → (n_alpha, n_cells, 256, 256) FULL frames + the α list. + Cellpose segments on the FULL FOV (robust); the crop to the 128 training FOV happens AFTER masking.""" + import json + from PIL import Image + base = f"{CACHE}/{modality}/{grain}/{gene}" + meta = json.load(open(f"{base}/cell0/meta.json")) if __import__("os").path.exists(f"{base}/cell0/meta.json") \ + else json.load(open(f"{base}/meta.json")) + alphas = meta["alphas"] + frames = [[] for _ in alphas] + for c in range(n_cells): + for ai in range(len(alphas)): + frames[ai].append(np.asarray(Image.open(f"{base}/cell{c}/frame_{ai:02d}.webp").convert("L"), np.float32) / 255.0) + return np.array(frames), alphas + + +def score_generated(genes=("HSPA5", "POLR1B", "KIF11", "TIMM23"), grain="geneKO", + modality="phase", n_cells=20, run="miwkg1cy"): + """End-to-end per-α classifier score for GENERATED traversals: cellpose-mask each fake frame → + mimic-CellDINO → z-std(control ref) → bag P(target) per α. Prints the α-curve (should rise toward + the target as |α| increases if the counterfactual convinces the classifier).""" + m, g2i, c2i = load_set_classifier(run) + i2g = {v: k for k, v in g2i.items()} + mu, sd = build_control_ref() + ci = c2i.get("Phase", 0) + for gene in genes: + if gene not in g2i: + print(f"{gene}: not in classifier"); continue + try: + frames, alphas = _load_gen(gene, grain, modality, n_cells) + except Exception as e: + print(f"{gene}: load failed ({e})"); continue + na, nc = frames.shape[:2] + flat = frames.reshape(na * nc, *frames.shape[2:]) # full 256 frames + masks = cellpose_masks(flat) # segment on the FULL FOV (robust) + masked, maskc = _crop_on_centroid(flat * masks, masks, 128) # 128 crop CENTERED ON THE CELL + raw, _ = _celldino(masked[:, None].astype(np.float32)) + cov = maskc.reshape(len(maskc), -1).mean(1) + print(f"[diag] {gene}: mask coverage {cov.mean()*100:.1f}% (empty {int((cov<0.01).sum())}/{len(cov)}) " + f"| raw-emb norm {np.linalg.norm(raw,axis=1).mean():.1f}") + emb = ((raw - mu) / sd).reshape(na, nc, -1) + ptarget = [float(score_bags(m, e[None], channel_idx=ci)[0, g2i[gene]]) for e in emb] + top = [i2g[int(score_bags(m, e[None], channel_idx=ci)[0].argmax())] for e in emb] + print(f"\n{gene} (grain={grain}) per-α P(target):") + for a, p, t in zip(alphas, ptarget, top): + bar = "#" * int(p * 40) + print(f" α={a:+.1f} P={p:.3f} {bar:40s} argmax={t}{' /0`); each position is a horizontal strip of the n_cells crops +(padded) in the marker channel. Seg runs per position → labels/_seg, read back per crop. +""" +from __future__ import annotations + +import json +import os + +import numpy as np + +CACHE = f"/hpc/projects/icd.fast.ops/models/diffex/{os.environ.get('OPS_DIFFEX_ASSETS', 'viewer_assets')}" +SYNTH_BASE = "/hpc/projects/icd.fast.ops/models/diffex/morpho_synth" +PAD = 24 +CROP = 256 +GEN_CROP = 160 # DiffEx cfg.crop_size: the generated crops are 160 px (native) upsized to 256 — real ref crops must match this window + + +def _json_safe(o): # non-finite floats (NaN/±Inf) → None so the browser's strict JSON.parse accepts the file + import math + if isinstance(o, float): + return o if math.isfinite(o) else None + if isinstance(o, dict): + return {k: _json_safe(v) for k, v in o.items()} + if isinstance(o, (list, tuple)): + return [_json_safe(v) for v in o] + return o + + +def _clip_border(lab, m=5): + """Zero the m-px border band: segmentations touching the crop edge are shrunk to end ~m px inside it, and + we measure on the clipped objects. Simpler than per-object edge/reach filtering — no whole-object drops.""" + if m <= 0: + return lab + out = lab.copy() + out[:m, :] = 0; out[-m:, :] = 0; out[:, :m] = 0; out[:, -m:] = 0 + return out + + +# Masked-Object (MO) nucleoli seg — copied from coding_exps/nucleoli_roundness (_segment_threshold), +# reusing the real apply_intensity_threshold. NPM3 nucleoli are round blobs → frangi vesselness under- +# detects them; MO (intensity threshold + per-object local adjust) is the right detector. No nucleus +# tile_mask here (the generated crop is already a single cell). +MO_PARAMS = {"threshold_method": "masked_object", "threshold_factor": 1.0, + "mo_global_method": "triangle", "mo_local_adjust": 1.3, + "mo_object_min_area_px": 15, "min_object_size": 15} + + +def _nucleus_mask(img): + """Central-cell nucleus mask for NPM3 crops (no separate nucleus channel): blurred Otsu, flood-guarded + (if it covers >45% of the crop, fall back to the 88th percentile), keep only the central component — + kills off-nucleus background specks that plain MO would otherwise label on noisy high-α frames.""" + from skimage.filters import threshold_otsu, gaussian + from skimage.morphology import binary_closing, disk + from scipy import ndimage as ndi + g = gaussian(img, 2) + v = g[g > 0] + if v.size == 0: + return np.zeros(img.shape, bool) + nuc = ndi.binary_fill_holes(binary_closing(g > threshold_otsu(v), disk(3))) + if nuc.mean() > 0.45: + nuc = ndi.binary_fill_holes(binary_closing(g > np.percentile(g, 88), disk(3))) + lab, n = ndi.label(nuc) + if n == 0: + return nuc + cy, cx = np.array(img.shape) // 2 + cen = lab[cy, cx] or (1 + int(np.argmax(np.bincount(lab.ravel())[1:]))) + return lab == cen + + +def _fit_ellipse_mask(nuc, scale=1.0): + """Ovular mask fit to the nucleus's 2nd moments (right centroid/orientation/aspect, NOT a perfect circle), + scaled by `scale` (1.0 ≈ nucleus extent; <1 shrinks it). Replaces ragged erosion so perinuclear protrusions + of the raw mask are cut off — objects are kept only inside the clean shrunk ellipse.""" + ys, xs = np.nonzero(nuc) + if len(ys) < 10: + return nuc + cy, cx = ys.mean(), xs.mean() + cov = np.cov(np.vstack([ys - cy, xs - cx])) + inv = np.linalg.pinv(cov) + Y, X = np.mgrid[0:nuc.shape[0], 0:nuc.shape[1]] + dy = Y - cy; dx = X - cx + md2 = inv[0, 0] * dy * dy + 2 * inv[0, 1] * dy * dx + inv[1, 1] * dx * dx + return md2 <= (2.0 * scale) ** 2 # md2<=4 (scale=1) reproduces the uniform ellipse + + +def _vs_h2b_nucleus_npz(marker_dir, target, grain, out_npz, n_cells, force=False): + """VS-predict H2B (nuclear) from each generated phase frame with the all-channel DiffAE, Cellpose-SAM the + virtual nucleus, keep the largest (central) object → cache (n_cells, n_alpha, CROP, CROP) uint8 nucleus masks. + Cached (masks depend only on the traversal frames). Replaces the Otsu/ellipse nucleus mask with a real one.""" + import json + import numpy as np + import torch + from PIL import Image + if os.path.exists(out_npz) and not force: + print(f"[vs-nuc] cache {out_npz}"); return out_npz + from ..diffae.config import DiffAEConfig + from ..diffae.model import DiffAE + from ..classifier.celldino_features import embed_crops + from diffusers import DDIMScheduler + from cellpose import models as cpm + from skimage.transform import resize + VOUT = "/hpc/projects/icd.fast.ops/analysis/virtual_staining/multi_marker" + dev = torch.device("cuda") + markers = json.load(open(f"{VOUT}/markers.json")); h2b = markers.index("chromatin_H2BC21") + cfg = DiffAEConfig(spatial_cond=True, n_markers=len(markers), device="cuda"); Hg = cfg.crop_size + model = DiffAE(cfg).to(dev).eval(); model.load_state_dict(torch.load(f"{VOUT}/diffae_best.pt", map_location=dev)) + cp = cpm.CellposeModel(gpu=True) + src = f"{CACHE}/{marker_dir}/{grain}/{target}" + mp = f"{src}/cell0/meta.json" if os.path.exists(f"{src}/cell0/meta.json") else f"{src}/meta.json" + na = len(json.load(open(mp))["alphas"]) + out = np.zeros((n_cells, na, CROP, CROP), np.uint8) + + @torch.no_grad() + def _vs_batch(P, E): # batched DDIM sample of H2B for a whole α strip + fwd = DDIMScheduler(num_train_timesteps=cfg.train_timesteps); fwd.set_timesteps(cfg.ddim_steps) + emb = torch.as_tensor(E, device=dev); ci = torch.as_tensor(P, device=dev) + mk = torch.full((P.shape[0],), h2b, dtype=torch.long, device=dev) + c = model.cond(emb, mk); x = torch.randn(P.shape[0], 1, Hg, Hg, device=dev) + for t in fwd.timesteps: + x = fwd.step(model.denoise(x, t, c, ci), t, x).prev_sample + return x.cpu().numpy()[:, 0] + + for ai in range(na): + idx, Ps = [], [] + for c in range(n_cells): + f = f"{src}/cell{c}/frame_{ai:02d}.webp" + if not os.path.exists(f): + continue + im = np.asarray(Image.open(f).convert("L").resize((Hg, Hg)), np.float32) / 255.0 + idx.append(c); Ps.append(im * 2 - 1) + if not idx: + continue + P = np.stack(Ps)[:, None].astype(np.float32) + pred = _vs_batch(P, embed_crops(P, cfg)) + for j, c in enumerate(idx): + m = cp.eval(np.clip((pred[j] + 1) / 2, 0, 1), diameter=None, flow_threshold=0.4, cellprob_threshold=0.0)[0] + if m.max() > 0: + ids, cnt = np.unique(m, return_counts=True); ids, cnt = ids[1:], cnt[1:] + m = (m == ids[cnt.argmax()]) # largest = the central nucleus + out[c, ai] = resize(np.asarray(m, float), (CROP, CROP), order=0).astype(np.uint8) + print(f"[vs-nuc] {target} α{ai}: {len(idx)} cells", flush=True) + np.savez_compressed(out_npz, masks=out) + print(f"[vs-nuc] saved {out_npz} {out.shape}"); return out_npz + + +def _seg_masked_object(img, tp=MO_PARAMS, nucleus=False, nucleus_scale=1.0, nucleus_override=None, override_erode=0): + """MO intensity seg on one 2D crop → int32 labels (fill_holes → CC label → min-size). + nucleus=True → constrain to the nucleus. nucleus_override (a precomputed binary, e.g. VS→Cellpose) is used + directly if given (shrunk by override_erode px); else fall back to the Otsu mask fit to an ellipse (nucleus_scale).""" + from scipy import ndimage as ndi + from organelle_profiler.organelle_seg.thresholding import apply_intensity_threshold + binary = apply_intensity_threshold( + img, method=tp["threshold_method"], threshold_factor=tp.get("threshold_factor", 1.0), + mo_global_method=tp.get("mo_global_method", "triangle"), + mo_local_adjust=tp.get("mo_local_adjust", 0.98), + mo_object_min_area_px=tp.get("mo_object_min_area_px", 100)) + if nucleus: + if nucleus_override is not None: + nuc = nucleus_override.astype(bool) + if override_erode > 0 and nuc.any(): + from skimage.morphology import binary_erosion, disk + nuc = binary_erosion(nuc, disk(override_erode)) + else: + nm = _nucleus_mask(img) + nuc = _fit_ellipse_mask(nm, nucleus_scale) if nm.any() else nm + if not nuc.any(): + return np.zeros(img.shape, np.int32) + binary = binary & nuc + binary = ndi.binary_fill_holes(binary) + fp = ndi.generate_binary_structure(binary.ndim, 1) + objs, _ = ndi.label(binary, structure=fp) + ms = tp.get("min_object_size", 0) + if ms > 0 and objs.max() > 0: + ids, cnt = np.unique(objs, return_counts=True) + small = ids[(ids > 0) & (cnt < ms)] + if small.size: + objs[np.isin(objs, small)] = 0 + objs, _ = ndi.label(objs > 0, structure=fp) + return objs.astype(np.int32) + + +def _seg_strip_mo(img, n_cells, nucleus=False, nucleus_scale=1.0, nuc_masks=None, vs_erode=0): + """Run MO per generated crop within the α strip (percentile-normalized per crop, like the standalone + _prep_npm3), assembling a strip-wide int32 label array with globally-unique ids. nuc_masks (n_cells,CROP,CROP) + = precomputed per-crop nucleus overrides (VS→Cellpose) for this α, else the internal Otsu/ellipse mask.""" + Y, W = img.shape + strip = np.zeros((Y, W), np.int32) + for c in range(n_cells): + x0 = c * (CROP + PAD) + crop = img[:, x0:x0 + CROP] + if crop.max() <= 0: + continue + lo, hi = np.percentile(crop, [1, 99.5]) + cn = np.clip((crop - lo) / max(hi - lo, 1e-6), 0, 1).astype(np.float32) + ov = nuc_masks[c] if nuc_masks is not None else None + objs = _seg_masked_object(cn, nucleus=nucleus, nucleus_scale=nucleus_scale, nucleus_override=ov, override_erode=vs_erode) + m = objs > 0 + if m.any(): + objs[m] += int(strip.max()) + strip[:, x0:x0 + CROP][m] = objs[m] + return strip + + +def _run_seg_masked_object(zpath, n_alpha, label_name, nucleus=False, nucleus_scale=1.0, vs_nucleus_npz=None, vs_erode=0): + """MO seg branch of run_seg: read each α strip (marker=last channel), segment per crop, and write the + labels into the mini-zarr via the production label writer so readback/full_features read them unchanged. + vs_nucleus_npz → precomputed VS→Cellpose nucleus masks (n_cells,n_alpha,CROP,CROP) used as per-crop overrides.""" + import zarr + from ops_utils.io.zarr_labels import _init_organelle_label_array, _update_labels_metadata + root = zarr.open(zpath, mode="r") + W = int(np.asarray(root["A/0/0/0"]).shape[-1]) + n_cells = round((W + PAD) / (CROP + PAD)) + vs_masks = np.load(vs_nucleus_npz)["masks"] if vs_nucleus_npz else None + out = [] + for ai in range(n_alpha): + img = np.asarray(root[f"A/{ai}/0/0"][0, -1, 0]).astype(np.float32) # marker channel = last + nm = vs_masks[:, ai] if vs_masks is not None else None + lab = _seg_strip_mo(img, n_cells, nucleus=nucleus, nucleus_scale=nucleus_scale, nuc_masks=nm, vs_erode=vs_erode) + Y, Wx = lab.shape + _init_organelle_label_array(zpath, f"A/{ai}/0", label_name, shape=(1, 1, 1, Y, Wx)) + store = zarr.open(zpath, mode="r+") + store[f"A/{ai}/0/labels/{label_name}/0"][0, 0, 0] = lab + _update_labels_metadata(zpath, f"A/{ai}/0", label_name) + out.append((ai, True, int(lab.max()), label_name, None)) + print(f" α_idx {ai}: MO n_obj={int(lab.max())} label={label_name}") + return out + + +def build_mini_zarr(marker_dir, target, grain, real_exp, channel_names, n_cells=6, base_dir=SYNTH_BASE): + """Stage a traversal's generated α-frames as a phenotyping_v3-style zarr under //. + Generated crop → the LAST channel (the marker channel, e.g. GFP); others zero. Returns (zpath, n_alpha).""" + from iohub import open_ome_zarr + from PIL import Image + src = f"{CACHE}/{marker_dir}/{grain}/{target}" + mp = f"{src}/cell0/meta.json" if os.path.exists(f"{src}/cell0/meta.json") else f"{src}/meta.json" + alphas = json.load(open(mp))["alphas"] + na = len(alphas) + W = n_cells * CROP + (n_cells - 1) * PAD + zpath = f"{base_dir}/{real_exp}/3-assembly/phenotyping_v3.zarr" + os.makedirs(os.path.dirname(zpath), exist_ok=True) + import shutil + if os.path.exists(zpath): + shutil.rmtree(zpath) + C = len(channel_names) + # version="0.5" → NGFF 0.5 = zarr v3 (matches real phenotyping_v3; the org-seg label writer needs v3 sharding) + with open_ome_zarr(zpath, layout="hcs", mode="w", channel_names=channel_names, version="0.5") as ds: + for ai in range(na): + strip = np.zeros((C, CROP, W), np.float32) + for c in range(n_cells): + f = f"{src}/cell{c}/frame_{ai:02d}.webp" + if not os.path.exists(f): + continue + im = np.asarray(Image.open(f).convert("L"), np.float32) / 255.0 + x0 = c * (CROP + PAD) + strip[-1, :, x0:x0 + CROP] = im # marker channel = last + pos = ds.create_position("A", str(ai), "0") + pos.create_image("0", strip[None, :, None]) # (T=1,C,Z=1,Y,X) + print(f"[mini-zarr] {zpath} {na} α × {n_cells} cells, channels={channel_names}") + return zpath, na + + +def _nucleus_binary(img): + """Nucleus binary (blurred Otsu + close + fill), NO central-only restriction — the generated crop is a + single masked cell, so ALL components are that cell's nuclei (captures KIF23 multinucleation).""" + from skimage.filters import threshold_otsu, gaussian + from skimage.morphology import binary_closing, disk + from scipy import ndimage as ndi + g = gaussian(img, 2); v = g[g > 0] + if v.size == 0: + return np.zeros(img.shape, bool) + nuc = ndi.binary_fill_holes(binary_closing(g > threshold_otsu(v), disk(3))) + if nuc.mean() > 0.45: + nuc = ndi.binary_fill_holes(binary_closing(g > np.percentile(g, 88), disk(3))) + return nuc + + +def _run_seg_nucleus(zpath, n_alpha, label_name, min_area=50): + """Nucleus-mask seg branch of run_seg: per α strip, per crop → nucleus binary → CC label (min-size), + write labels into the mini-zarr (same writer as the MO branch).""" + import zarr + from scipy import ndimage as ndi + from ops_utils.io.zarr_labels import _init_organelle_label_array, _update_labels_metadata + root = zarr.open(zpath, mode="r") + W = int(np.asarray(root["A/0/0/0"]).shape[-1]) + n_cells = round((W + PAD) / (CROP + PAD)) + out = [] + for ai in range(n_alpha): + img = np.asarray(root[f"A/{ai}/0/0"][0, -1, 0]).astype(np.float32) + Y, Wx = img.shape + strip = np.zeros((Y, Wx), np.int32) + for c in range(n_cells): + x0 = c * (CROP + PAD); crop = img[:, x0:x0 + CROP] + if crop.max() <= 0: + continue + lo, hi = np.percentile(crop, [1, 99.5]) + cn = np.clip((crop - lo) / max(hi - lo, 1e-6), 0, 1).astype(np.float32) + objs, _ = ndi.label(_nucleus_binary(cn)) + if objs.max(): # drop specks below min_area + ids, cnt = np.unique(objs, return_counts=True) + small = ids[(ids > 0) & (cnt < min_area)] + if small.size: + objs[np.isin(objs, small)] = 0 + objs, _ = ndi.label(objs > 0) + m = objs > 0 + if m.any(): + objs[m] += int(strip.max()); strip[:, x0:x0 + CROP][m] = objs[m] + _init_organelle_label_array(zpath, f"A/{ai}/0", label_name, shape=(1, 1, 1, Y, Wx)) + store = zarr.open(zpath, mode="r+") + store[f"A/{ai}/0/labels/{label_name}/0"][0, 0, 0] = strip + _update_labels_metadata(zpath, f"A/{ai}/0", label_name) + out.append((ai, True, int(strip.max()), label_name, None)) + print(f" α_idx {ai}: nucleus n_obj={int(strip.max())} label={label_name}") + return out + + +def run_seg(real_exp, marker_channel, n_alpha, structure_type=None, base_dir=SYNTH_BASE, frangi_params=None, method=None, label_name=None, mo_nucleus=False, mo_nucleus_scale=1.0, vs_nucleus_npz=None, vs_erode=0): + """Run the REAL production org-seg on each α position of the mini-zarr (config resolved from the + real experiment's channel map). `frangi_params` overrides the resolved frangi config (e.g. to switch + to the ADAPTIVE dynamic threshold on the generated images). Writes labels; returns per-α results. + method="masked_object" → the MO intensity-threshold path (not wired through segment_single_position_channel); + writes labels under `label_name` directly. mo_nucleus=True → nucleus-constrained MO (NPM3 nucleoli).""" + os.environ["OPS_OUTPUT_BASE_DIR"] = base_dir + if method == "masked_object": + return _run_seg_masked_object(f"{base_dir}/{real_exp}/3-assembly/phenotyping_v3.zarr", + n_alpha, label_name or "organelle_seg", nucleus=mo_nucleus, nucleus_scale=mo_nucleus_scale, vs_nucleus_npz=vs_nucleus_npz, vs_erode=vs_erode) + if method == "nucleus": + return _run_seg_nucleus(f"{base_dir}/{real_exp}/3-assembly/phenotyping_v3.zarr", + n_alpha, label_name or "nucleus_seg") + from organelle_profiler.organelle_seg.organelle_segmentation import segment_single_position_channel + out = [] + for ai in range(n_alpha): + r = segment_single_position_channel(experiment=real_exp, position=f"A/{ai}/0", + channel_key=marker_channel, structure_type=structure_type, + use_clahe=True, frangi_params=frangi_params, method=method) + out.append((ai, r.get("success"), r.get("num_objects"), r.get("output_label"), r.get("error"))) + print(f" α_idx {ai}: success={r.get('success')} n_obj={r.get('num_objects')} " + f"label={r.get('output_label')} {r.get('error') or ''}") + return out + + +OPCP = "/hpc/projects/icd.fast.ops/analysis/op_cp_features/op_cp_features_{store}.h5ad" +REF_SUFFIX = {"count": "count", "total_area": "area_sum", "mean_area": "area_mean", + "mean_int": "intensity_mean_mean", "mean_ecc": "eccentricity_mean"} + + +def _real_ref(store_marker, org_prefix, target, ref_map=None): + """NTC(empty gene) + target mean±SEM for each aggregate feature, from the precomputed op_cp_features + store (per-cell, 60M-cell scale). → {agg: {ntc:[mean,sem], ko:[mean,sem]}}. + ref_map (optional) = {agg: full store feature name}, overrides the default op__.""" + import anndata as ad + st = ad.read_h5ad(OPCP.format(store=store_marker), backed="r") + gn = st.obs["gene_name"].astype(str).values + ntc, ko = gn == "", gn == target + items = ref_map.items() if ref_map else {ak: f"op_{org_prefix}_{suf}" for ak, suf in REF_SUFFIX.items()}.items() + ref = {} + for ak, fn in items: + if fn not in st.var_names: + continue + col = np.asarray(st[:, fn].X).ravel().astype(float) + def ms(m): + v = col[m]; v = v[np.isfinite(v)] + return [float(v.mean()), float(v.std() / max(len(v) ** 0.5, 1.0))] if len(v) else [None, None] + ref[ak] = {"ntc": ms(ntc), "ko": ms(ko)} + return ref + + +def readback_cache(marker_dir, target, grain, real_exp, label_name, n_cells, base_dir=SYNTH_BASE, + store_marker=None, org_prefix=None, network=False, ref_map=None): + """Read the REAL org-seg labels the pipeline wrote into the mini-zarr, split per (cell, α) crop, + measure per-organelle features (regionprops on the real labels) + per-α aggregates, and cache for + the viewer overlay/plot. → viewer_assets/_morphometrics////""" + import json + import zarr + from PIL import Image + from skimage.measure import label as relabel, regionprops + zpath = f"{base_dir}/{real_exp}/3-assembly/phenotyping_v3.zarr" + src = f"{CACHE}/{marker_dir}/{grain}/{target}" + mp = f"{src}/cell0/meta.json" if os.path.exists(f"{src}/cell0/meta.json") else f"{src}/meta.json" + alphas = json.load(open(mp))["alphas"] + root = zarr.open(zpath, mode="r") + out = f"{CACHE}/_morphometrics/{marker_dir}/{grain}/{target}"; os.makedirs(out, exist_ok=True) + # network markers (mito/tubular): count = num fragments (num_objects), skel = per-component skeleton + # length (branch-length proxy). Else vesicular morphology. + AGG = (["count", "total_area", "total_skel", "mean_skel", "mean_int"] if network + else ["count", "total_area", "mean_area", "mean_int", "mean_ecc"]) + agg = {k: [] for k in AGG} + if network: + from skimage.morphology import skeletonize + for ai in range(len(alphas)): + lab = np.asarray(root[f"A/{ai}/0/labels/{label_name}/0"][0, 0, 0]) # (Y, W) real labels + img = np.asarray(root[f"A/{ai}/0/0"][0, -1, 0]) # marker channel strip + per_cell = [] + for c in range(n_cells): + x0 = c * (CROP + PAD) + lc = relabel(_clip_border(lab[:, x0:x0 + CROP]) > 0) # clip border band, then relabel within crop + ic = img[:, x0:x0 + CROP] + rp = [r for r in regionprops(lc, intensity_image=ic) if r.area >= 3] # border already clipped; just drop tiny + remap = {r.label: i + 1 for i, r in enumerate(rp)} # sequential 1..K for the 16-bit mask + mask = np.zeros(lc.shape, np.uint16) # 16-bit: no 255-object cap (was clipping bottom of dense segs) + feats = {} + for r in rp: + nl = remap[r.label]; mask[lc == r.label] = nl + per = float(r.perimeter) if r.perimeter else 0.0 + d = {"area": float(r.area), "area_filled": float(r.area_filled), "mean_int": float(r.intensity_mean), + "ecc": float(r.eccentricity), "extent": float(r.extent), "solidity": float(r.solidity), + "axis_major_length": float(r.axis_major_length), "axis_minor_length": float(r.axis_minor_length), + "circularity": float(4 * np.pi * r.area / (per * per)) if per > 0 else 0.0} # per-object props for feature-matched overlay + if network: + d["skel"] = float(skeletonize(lc == r.label).sum()) # per-component branch-length proxy + feats[str(nl)] = d # keyed by the mask's pixel value + cdir = f"{out}/cell{c}"; os.makedirs(cdir, exist_ok=True) + Image.fromarray(mask).save(f"{cdir}/a{ai:02d}_labels.png") # 16-bit component-index mask (real shape) + json.dump(feats, open(f"{cdir}/a{ai:02d}_feats.json", "w")) + A = [f["area"] for f in feats.values()]; I = [f["mean_int"] for f in feats.values()] + pc = {"count": len(feats), "total_area": float(np.sum(A) if A else 0), + "mean_int": float(np.mean(I) if I else 0)} + if network: + S = [f["skel"] for f in feats.values()] + pc["total_skel"] = float(np.sum(S) if S else 0); pc["mean_skel"] = float(np.mean(S) if S else 0) + else: + E = [f["ecc"] for f in feats.values()] + pc["mean_area"] = float(np.mean(A) if A else 0); pc["mean_ecc"] = float(np.mean(E) if E else 0) + per_cell.append(pc) + for k in AGG: + agg[k].append(float(np.mean([pc[k] for pc in per_cell]))) + real_ref = _real_ref(store_marker, org_prefix, target, ref_map) if (store_marker and (org_prefix or ref_map)) else {} + json.dump({"marker_dir": marker_dir, "target": target, "grain": grain, "alphas": alphas, + "n_cells": n_cells, "label_name": label_name, "agg": agg, "features": AGG, + "real_ref": real_ref}, open(f"{out}/morpho.json", "w")) + if real_ref: + print(f" real_ref (NTC vs {target}): " + + ", ".join(f"{k}={v['ntc'][0]:.1f}/{v['ko'][0]:.1f}" for k, v in real_ref.items() if v['ntc'][0])) + print(f"[readback] {marker_dir}/{target}: REAL-seg aggregates cached -> {out}") + print(f" count α-series: {[round(v,1) for v in agg['count']]}") + print(f" total_area α-series: {[round(v,0) for v in agg['total_area']]}") + return out + + +def reference_cells(marker_dir, target, real_exp, marker_channel, structure_type=None, + groups=("NTC",), n_cells=6, base_dir=SYNTH_BASE): + """Segment REAL reference cells (the traversal's cached `_anchors//cell*/real.webp`) through + the SAME org-seg pipeline → per-organelle feats, so the viewer can show real cells with the identical + overlay next to the generated ones. groups e.g. ("NTC","AP2M1"). → _morphometrics/...//_ref.json""" + import json + import shutil + import zarr + from iohub import open_ome_zarr + from PIL import Image + from skimage.measure import label as relabel, regionprops + grp_cells = {} # {group: [crop arrays]} + for g in groups: + crops = [] + for c in range(n_cells): + f = f"{CACHE}/{marker_dir}/_anchors/{g}/cell{c}/real.webp" + if os.path.exists(f): + crops.append(np.asarray(Image.open(f).convert("L"), np.float32) / 255.0) + if crops: + grp_cells[g] = crops + if not grp_cells: + print("[ref] no real cells cached for", groups); return + # stage: one position per group, strip of its cells + zpath = f"{base_dir}/{real_exp}/3-assembly/phenotyping_v3.zarr" + if os.path.exists(zpath): + shutil.rmtree(zpath) + os.makedirs(os.path.dirname(zpath), exist_ok=True) + glist = list(grp_cells) + with open_ome_zarr(zpath, layout="hcs", mode="w", channel_names=["Phase2D", marker_channel], version="0.5") as ds: + for gi, g in enumerate(glist): + crops = grp_cells[g]; W = len(crops) * CROP + (len(crops) - 1) * PAD + strip = np.zeros((2, CROP, W), np.float32) + for c, cr in enumerate(crops): + strip[-1, :, c * (CROP + PAD):c * (CROP + PAD) + CROP] = cr + ds.create_position("A", str(gi), "0").create_image("0", strip[None, :, None]) + os.environ["OPS_OUTPUT_BASE_DIR"] = base_dir + from organelle_profiler.organelle_seg.organelle_segmentation import segment_single_position_channel + root = None + ref = {} + for gi, g in enumerate(glist): + r = segment_single_position_channel(experiment=real_exp, position=f"A/{gi}/0", + channel_key=marker_channel, structure_type=structure_type, use_clahe=True) + if not r.get("success"): + print(f"[ref] {g} seg failed: {r.get('error')}"); continue + if root is None: + root = zarr.open(zpath, mode="r") + lab = np.asarray(root[f"A/{gi}/0/labels/{r['output_label']}/0"][0, 0, 0]) + img = np.asarray(root[f"A/{gi}/0/0"][0, -1, 0]) + cells = [] + for c in range(len(grp_cells[g])): + x0 = c * (CROP + PAD); lc = relabel(lab[:, x0:x0 + CROP] > 0); ic = img[:, x0:x0 + CROP] + EM = 6 + orgs = [{"cx": float(rr.centroid[1]), "cy": float(rr.centroid[0]), "r": float((rr.area / 3.14159) ** 0.5), + "area": float(rr.area), "mean_int": float(rr.intensity_mean), "ecc": float(rr.eccentricity)} + for rr in regionprops(lc, intensity_image=ic) + if rr.area >= 3 and rr.bbox[0] >= EM and rr.bbox[1] >= EM and rr.bbox[2] <= CROP - EM and rr.bbox[3] <= CROP - EM] + cells.append({"cell": c, "organelles": orgs}) + ref[g] = cells + out = f"{CACHE}/_morphometrics/{marker_dir}/geneKO/{target}" + os.makedirs(out, exist_ok=True) + json.dump({"groups": glist, "n_cells": n_cells, "cells": ref}, open(f"{out}/_ref.json", "w")) + print(f"[ref] {marker_dir}/{target}: real cells {[(g, len(ref.get(g,[]))) for g in glist]} -> {out}/_ref.json") + + +def reference_from_direction(marker_dir, target, real_exp, marker_channel, structure_type=None, + network=False, n_cells=4, base_dir=SYNTH_BASE, adaptive=True, grain="geneKO"): + """Real reference cells for targets with NO cached anchors (e.g. MICOS13): pull the direction-build's + real crops (`directions//geneKO//cache/crops__*.npz` — labels 0=control, + 1=KD), seg them through the SAME pipeline, save crop webps + per-organelle feats. → /_ref.""" + import glob + import json + import shutil + import zarr + from iohub import open_ome_zarr + from PIL import Image + from skimage.measure import label as relabel, regionprops + from skimage.transform import resize as skresize + npz = glob.glob(f"/hpc/projects/icd.fast.ops/models/diffex/directions/{marker_dir}/{grain}/{target}/cache/crops_{target}_*.npz") + if not npz: + print(f"[ref-dir] no crops npz for {marker_dir}/{target}"); return + d = np.load(npz[0], allow_pickle=True) + imgs = d["images"]; labs = d["labels"] + imgs = imgs[:, 0] if imgs.ndim == 4 else imgs # (N,H,W) + def norm256(a): + a = skresize(a.astype(np.float32), (CROP, CROP), preserve_range=True, anti_aliasing=True) + lo, hi = np.percentile(a, [1, 99]); return np.clip((a - lo) / (hi - lo + 1e-6), 0, 1) + grp_cells = {"NTC": [norm256(x) for x in imgs[labs == 0][:n_cells]], + target: [norm256(x) for x in imgs[labs == 1][:n_cells]]} + grp_cells = {g: v for g, v in grp_cells.items() if v} + zpath = f"{base_dir}/{real_exp}/3-assembly/phenotyping_v3.zarr" + if os.path.exists(zpath): + shutil.rmtree(zpath) + os.makedirs(os.path.dirname(zpath), exist_ok=True) + glist = list(grp_cells) + chans0 = ["Phase2D"] if marker_channel == "Phase2D" else ["Phase2D", marker_channel] # phase: single channel (crop IS phase; avoid empty-slot collision) + with open_ome_zarr(zpath, layout="hcs", mode="w", channel_names=chans0, version="0.5") as ds: + for gi, g in enumerate(glist): + cr = grp_cells[g]; W = len(cr) * CROP + (len(cr) - 1) * PAD; strip = np.zeros((len(chans0), CROP, W), np.float32) + for c, x in enumerate(cr): + strip[-1, :, c * (CROP + PAD):c * (CROP + PAD) + CROP] = x + ds.create_position("A", str(gi), "0").create_image("0", strip[None, :, None]) + os.environ["OPS_OUTPUT_BASE_DIR"] = base_dir + from organelle_profiler.organelle_seg.organelle_segmentation import segment_single_position_channel + fp = None + if network: + from skimage.morphology import skeletonize # used below for per-object skel features + if adaptive: # match the generated: adaptive dynamic threshold (adaptive=False → perfected config) + _, dp = _resolve_seg(real_exp, marker_channel); fp = dict(dp) + fp["threshold"] = None; fp["threshold_mult"] = 0.1; fp["pixel_size_um"] = 0.065 # match generated seg + out = f"{CACHE}/_morphometrics/{marker_dir}/{grain}/{target}"; os.makedirs(out, exist_ok=True) + root, ref = None, {} + for gi, g in enumerate(glist): + r = segment_single_position_channel(experiment=real_exp, position=f"A/{gi}/0", + channel_key=marker_channel, structure_type=structure_type, + use_clahe=True, frangi_params=fp) + if not r.get("success"): + print(f"[ref-dir] {g} seg failed: {r.get('error')}"); continue + root = root or zarr.open(zpath, mode="r") + lab = np.asarray(root[f"A/{gi}/0/labels/{r['output_label']}/0"][0, 0, 0]); img = np.asarray(root[f"A/{gi}/0/0"][0, -1, 0]) + gd = f"{out}/_ref/{g}"; os.makedirs(gd, exist_ok=True); cells = [] + for c in range(len(grp_cells[g])): + x0 = c * (CROP + PAD); lc = relabel(lab[:, x0:x0 + CROP] > 0); ic = img[:, x0:x0 + CROP]; EM = 6 + Image.fromarray((grp_cells[g][c] * 255).astype(np.uint8)).save(f"{gd}/cell{c}.webp") + keep = [rr for rr in regionprops(lc, intensity_image=ic) + if rr.area >= 3 and rr.bbox[0] >= EM and rr.bbox[1] >= EM and rr.bbox[2] <= CROP - EM and rr.bbox[3] <= CROP - EM][:255] + remap = {rr.label: i + 1 for i, rr in enumerate(keep)} + mask8 = np.zeros(lc.shape, np.uint8); feats = {} + for rr in keep: + nl = remap[rr.label]; mask8[lc == rr.label] = nl + o = {"area": float(rr.area), "mean_int": float(rr.intensity_mean), "ecc": float(rr.eccentricity)} + if network: + o["skel"] = float(skeletonize(lc == rr.label).sum()) + feats[str(nl)] = o + Image.fromarray(mask8).save(f"{gd}/cell{c}_labels.png") + cells.append({"cell": c, "feats": feats}) + ref[g] = cells + json.dump({"groups": glist, "n_cells": n_cells, "cells": ref, "img_dir": "_ref"}, open(f"{out}/_ref.json", "w")) + print(f"[ref-dir] {marker_dir}/{target}: real cells {[(g, len(ref.get(g, []))) for g in glist]} + webps -> {out}/_ref") + + +def reference_from_store(marker_dir, target, grain, image_channel, org_label, n_cells=4, network=True, cache=CACHE, store_exps=None): + """Cache real-cell reference thumbnails + PRODUCTION org-label overlays for the morpho demo. Each cell is + cropped ONCE (build-time) from its experiment's phenotyping_v3.zarr — marker image channel + the on-disk + production seg label — and written as static webp/png (the _v3 store is NOT on S3; the web reads the cache). + Cells = top-1k ACCURACY KO (member genes for a complex) + top-attention NTC, matching the plot's real_ref.""" + import json + import re + import pandas as pd + import zarr + from PIL import Image + from skimage.measure import label as relabel, regionprops + from ..classifier.config import GRAINS + from skimage.morphology import skeletonize + GK = GRAINS["geneKO"]["parquet"] # attention parquet (has NTC + coords) + ACC = "/hpc/projects/icd.fast.ops/models/diffex/accuracy_ranking/phase_geneKO_topacc_ALL_top1000.parquet" + COLS = ["gene", "experiment", "well", "segmentation", "x_pheno", "y_pheno", "rank", "rank_type"] + + def _cells(pq, genes, n): # top-n ranked 'top' cells for gene(s) + d = pd.read_parquet(pq, columns=COLS) + d = d[(d["gene"].astype(str).isin([str(x) for x in genes])) & (d["rank_type"] == "top")] + if store_exps: # only experiments where the marker channel IS this reporter (fluor markers are experiment-specific) + d = d[d["experiment"].astype(str).isin(set(store_exps))] + return d.sort_values("rank").head(n) + if grain == "complex": + import yaml + from ..classifier.config import slugify + y = yaml.safe_load(open("/hpc/projects/icd.fast.ops/configs/gene_clusters/EBI_complexes_v1_updated_gene_names.yaml")) + members = next((e["genes"] for e in y.values() if slugify(e["name"]) == target), []) + ko = _cells(ACC, members, n_cells) + else: + ko = _cells(ACC, [target], n_cells) + ntc = _cells(GK, ["NTC"], n_cells) + groups = {"NTC": ntc, target: ko} + half = CROP // 2 + import shutil + out = f"{cache}/_morphometrics/{marker_dir}/{grain}/{target}"; os.makedirs(out, exist_ok=True) + shutil.rmtree(f"{out}/_ref", ignore_errors=True) # clear stale ref cells (e.g. a prior recompute build) + ref, glist = {}, [] + for g, df in groups.items(): + gd = f"{out}/_ref/{g}"; os.makedirs(gd, exist_ok=True); cells = [] + for c, (_, row) in enumerate(df.iterrows()): + exp = str(row["experiment"]); w = str(row["well"]).strip() + m = re.match(r"^([A-Za-z]+)(\d+)$", w); pos = w if w.count("/") == 2 else (f"{m.group(1)}/{m.group(2)}/0" if m else w) + zp = f"/hpc/projects/icd.fast.ops/{exp}/3-assembly/phenotyping_v3.zarr" + if not os.path.exists(zp): + continue + try: + z = zarr.open(zp, mode="r"); P = z[pos] + chans = [ch.get("label") for ch in dict(P.attrs).get("ome", {}).get("omero", {}).get("channels", [])] + ci = chans.index(image_channel) if image_channel in chans else 0 + from skimage.transform import resize as _rs + y, x = int(round(float(row["y_pheno"]))), int(round(float(row["x_pheno"]))) + h = GEN_CROP // 2 # SAME fixed window as the generated crops (cfg.crop_size=160), then upsize to CROP — matches traversal framing exactly + imc = np.asarray(P["0"][0, ci, 0, max(0, y - h):y + h, max(0, x - h):x + h]).astype(np.float32) + lbc = np.asarray(P["labels"][org_label]["0"][0, 0, 0, max(0, y - h):y + h, max(0, x - h):x + h]).astype(np.int32) + im = _rs(imc, (CROP, CROP), preserve_range=True, anti_aliasing=True).astype(np.float32) + lb = _rs(lbc.astype(np.float32), (CROP, CROP), order=0, preserve_range=True, anti_aliasing=False).astype(np.int32) + except Exception as e: + print(f"[ref-store] {g} cell{c} {exp} {pos} crop failed: {repr(e)[:100]}"); continue + if im.shape != (CROP, CROP): # pad edge crops to full tile + im = np.pad(im, [(0, CROP - im.shape[0]), (0, CROP - im.shape[1])]); lb = np.pad(lb, [(0, CROP - lb.shape[0]), (0, CROP - lb.shape[1])]) + lo, hi = np.percentile(im, [1, 99]); imn = np.clip((im - lo) / (hi - lo + 1e-6), 0, 1) + Image.fromarray((imn * 255).astype(np.uint8)).save(f"{gd}/cell{c}.webp") + lc = relabel(_clip_border(lb) > 0) # clip border band, then measure clipped objects + keep = [r for r in regionprops(lc, intensity_image=im) if r.area >= 3][:255] + remap = {r.label: i + 1 for i, r in enumerate(keep)}; mask8 = np.zeros(lc.shape, np.uint8); feats = {} + for r in keep: + mask8[lc == r.label] = remap[r.label] + o = {"area": float(r.area), "mean_int": float(r.intensity_mean), "ecc": float(r.eccentricity)} + if network: + o["skel"] = float(skeletonize(lc == r.label).sum()) + feats[str(remap[r.label])] = o + Image.fromarray(mask8).save(f"{gd}/cell{c}_labels.png") + cells.append({"cell": c, "feats": feats}) + if cells: + ref[g] = cells; glist.append(g) + json.dump({"groups": glist, "n_cells": n_cells, "cells": ref, "img_dir": "_ref"}, open(f"{out}/_ref.json", "w")) + print(f"[ref-store] {marker_dir}/{grain}/{target}: real cells {[(g, len(ref.get(g, []))) for g in glist]} (production {org_label}) -> {out}/_ref") + + +def sweep_grid(marker_dir, target, real_exp, marker_channel, pix=(0.065, 0.13, 0.185, 0.37), + thr=(0.1, 0.5, 1.0), ai=8, out="/hpc/projects/icd.fast.ops/models/diffex/morpho_grid_sweep.png"): + """Sweep the two frangi knobs — pixel_size_um (rows) × threshold_mult (cols) — on one generated crop, + render seg boundaries + object counts, so we can pick the cleanest (signal, no noise). No postprocess.""" + import matplotlib + matplotlib.use("Agg") + import matplotlib.pyplot as plt + import zarr + from skimage.segmentation import find_boundaries + plt.rcParams["pdf.fonttype"] = 42 + zp, _ = build_mini_zarr(marker_dir, target, "geneKO", real_exp, ["Phase2D", marker_channel], 1) + os.environ["OPS_OUTPUT_BASE_DIR"] = SYNTH_BASE + from organelle_profiler.organelle_seg.organelle_segmentation import segment_single_position_channel + _, dp = _resolve_seg(real_exp, marker_channel) + root = zarr.open(zp, mode="r"); img = np.asarray(root[f"A/{ai}/0/0"][0, -1, 0])[:, :CROP] + fig, ax = plt.subplots(len(pix), len(thr) + 1, figsize=((len(thr) + 1) * 2.3, len(pix) * 2.3)) + for ri, px in enumerate(pix): + ax[ri, 0].imshow(img, cmap="gray"); ax[ri, 0].set_ylabel(f"px={px}", fontsize=9) + ax[ri, 0].set_xticks([]); ax[ri, 0].set_yticks([]) + if ri == 0: ax[ri, 0].set_title("raw", fontsize=9) + for cj, t in enumerate(thr): + fp = dict(dp); fp["threshold"] = None; fp["threshold_mult"] = t; fp["pixel_size_um"] = px + r = segment_single_position_channel(experiment=real_exp, position=f"A/{ai}/0", channel_key=marker_channel, + structure_type=None, use_clahe=True, frangi_params=fp) + lab = np.asarray(root[f"A/{ai}/0/labels/{r['output_label']}/0"][0, 0, 0])[:, :CROP] + n = len(np.unique(lab)) - 1; b = find_boundaries(lab) + a = ax[ri, cj + 1]; a.imshow(img, cmap="gray"); a.imshow(np.ma.masked_where(~b, b), cmap="autumn"); a.axis("off") + if ri == 0: a.set_title(f"thr={t}", fontsize=9) + a.text(4, 20, str(n), color="cyan", fontsize=8) + fig.suptitle(f"{target} frangi: pixel_size (rows) × threshold_mult (cols) — obj count cyan", fontsize=11) + fig.tight_layout(); fig.savefig(out, dpi=125, bbox_inches="tight"); print(f"[sweep] {out}") + + +def _resolve_seg(real_exp, marker_channel): + """Config-matched (method, detection_params) for this channel — from org_seg_params, no hardcoding.""" + os.environ.setdefault("OPS_OUTPUT_BASE_DIR", SYNTH_BASE) + from organelle_profiler.organelle_seg.channel_processor import resolve_single_channel_info + from organelle_profiler.organelle_seg.metadata import _determine_processing_params + ci = resolve_single_channel_info(marker_channel, ["Phase2D", marker_channel], experiment=real_exp) + dp, _, m = _determine_processing_params(organelle_key=ci.get("organelle_key", marker_channel), + source_channel=marker_channel, structure_type=None, ch_info=ci, frangi_params=None, + clahe_params=None, post_clahe_smoothing_sigma=None) + return m, (dp or {}) + + +def _auto_ref_map(network, channel): + """Store (op_cp_features) feature names for the aggregates — network markers use op_network__*.""" + ch = channel.lower() + if network: + return {"count": f"op_network_{ch}_num_skeleton_components", "total_skel": f"op_network_{ch}_total_branch_length", + "mean_skel": f"op_network_{ch}_branch_length_mean", "total_area": f"op_{ch}_area_sum", + "mean_int": f"op_{ch}_intensity_mean_mean"} + return {"count": f"op_{ch}_count", "total_area": f"op_{ch}_area_sum", "mean_area": f"op_{ch}_area_mean", + "mean_int": f"op_{ch}_intensity_mean_mean", "mean_ecc": f"op_{ch}_eccentricity_mean"} + + +def run_target(marker_dir, target, real_exp, marker_channel, store_marker, grain="geneKO", n_cells=6, + refs=True, adaptive_mult=0.1, fake_pixel_um=0.065, adaptive=True, frangi_override=None, structure_type=None, seg_method=None, base_dir=SYNTH_BASE, org_label=None, mo_nucleus=False, mo_nucleus_scale=1.0, vs_nucleus_npz=None, vs_erode=0): + """Fully config-driven: seg method + params AUTO from org_seg_params; network-vs-vesicular feature set + + store ref-map AUTO from the resolved method. For frangi on the GENERATED (fake) images, switch to the + ADAPTIVE dynamic threshold (compute_frangi_threshold) — the config's fixed threshold mis-fits the fake + intensity/noise. mini-zarr → REAL seg → readback + real_ref (+ ref cells).""" + method, dp = _resolve_seg(real_exp, marker_channel) + network = (structure_type == "tubular") if structure_type else (method == "frangi") # only tubular has a skeleton (vesicular = blob features) + ref_map = _auto_ref_map(network, marker_channel) + fp = None + if network and adaptive: # adapt frangi to the fake resolution (adaptive=False → perfected config params) + fp = dict(dp); fp["threshold"] = None; fp["threshold_mult"] = adaptive_mult + fp["pixel_size_um"] = fake_pixel_um # smaller px → larger sigmas → coarse network, not noise + if frangi_override: # per-target tweaks on the resolved config (e.g. lower pixel_size for NPM3 nucleoli) + fp = dict(fp if fp else dp); fp.update(frangi_override) + print(f"[run_target] {marker_dir}/{target}: method={method} network={network} adaptive={adaptive} override={frangi_override}") + chans0 = ["Phase2D"] if marker_channel == "Phase2D" else ["Phase2D", marker_channel] # phase: single channel (avoid empty-slot collision) + zpath, na = build_mini_zarr(marker_dir, target, grain, real_exp, chans0, n_cells, base_dir=base_dir) + res = run_seg(real_exp, marker_channel, na, structure_type=structure_type, frangi_params=fp, method=seg_method, base_dir=base_dir, label_name=org_label, mo_nucleus=mo_nucleus, mo_nucleus_scale=mo_nucleus_scale, vs_nucleus_npz=vs_nucleus_npz, vs_erode=vs_erode) # config auto-resolves; fp=adaptive/override + label_name = next((r[3] for r in res if r[1] and r[3]), None) + if not label_name: + print("[run_target] seg produced no label — abort"); return + readback_cache(marker_dir, target, grain, real_exp, label_name, n_cells, base_dir=base_dir, + store_marker=store_marker, network=network, ref_map=ref_map) + if refs: + reference_from_direction(marker_dir, target, real_exp, marker_channel, structure_type=None, + network=network, n_cells=min(4, n_cells), adaptive=adaptive, grain=grain) + + +def validate(marker_dir="lysosome_LAMP1", target="ABCE1", grain="geneKO", + real_exp="ops0047_20250612", marker_channel="GFP", structure_type="vesicular", n_cells=6): + """Build the mini-zarr + run the REAL org-seg on a demo traversal → per-α num_objects (should + amplify with α if the phenotype does). Proves the reuse works end-to-end.""" + zpath, na = build_mini_zarr(marker_dir, target, grain, real_exp, ["Phase2D", marker_channel], n_cells) + res = run_seg(real_exp, marker_channel, na, structure_type) + nobj = [r[2] for r in res if r[1]] + print(f"\n[validate] {marker_dir}/{target}: real org-seg num_objects per α = {nobj}") + + +def full_features(marker_dir, target, real_exp, marker_channel, grain="geneKO", n_cells=6, + fake_pixel_um=0.05, adaptive_mult=0.1, adaptive=True, out_root=None, frangi_override=None, structure_type=None, seg_method=None, base_dir=SYNTH_BASE, org_label=None, mo_nucleus=False, mo_nucleus_scale=1.0, vs_nucleus_npz=None, vs_erode=0): + """REAL org-profiler feature extraction on the generated (cell, α) crops — NO skimage shortcut. + Stage traversal (mini-zarr) → production org-seg → per crop run `process_single_cell` + (extract_organelle_features + calculate_network_features), aggregate objects→cell with the pipeline's + AGGREGATION_FUNCTIONS. Writes a full per-(cell,α) feature table (op_cp-comparable) → cache parquet. + Each generated crop is treated as one cell (whole-crop cell mask, per design).""" + import pandas as pd + import zarr + from organelle_profiler.feature_extraction.fe_workers import process_single_cell + from organelle_profiler.feature_extraction.fe_constants import AGGREGATION_FUNCTIONS + method, dp = _resolve_seg(real_exp, marker_channel) + network = (structure_type == "tubular") if structure_type else (method == "frangi") # only tubular has a skeleton (vesicular = blob features) + fp = None + if network and adaptive: # adaptive frangi override for fake images + fp = dict(dp); fp["threshold"] = None; fp["threshold_mult"] = adaptive_mult; fp["pixel_size_um"] = fake_pixel_um + # adaptive=False → fp stays None → run_seg uses the perfected config params (e.g. Phase2D tubular) as-is + if frangi_override: # per-target tweaks on the resolved config (e.g. real NPM3/NucleoLIVE settings) + fp = dict(fp if fp else dp); fp.update(frangi_override) + chans0 = ["Phase2D"] if marker_channel == "Phase2D" else ["Phase2D", marker_channel] # phase: single channel (generated frame IS phase; avoids the empty Phase2D-slot collision) + zpath, na = build_mini_zarr(marker_dir, target, grain, real_exp, chans0, n_cells, base_dir=base_dir) + res = run_seg(real_exp, marker_channel, na, structure_type=structure_type, frangi_params=fp, method=seg_method, base_dir=base_dir, label_name=org_label, mo_nucleus=mo_nucleus, mo_nucleus_scale=mo_nucleus_scale, vs_nucleus_npz=vs_nucleus_npz, vs_erode=vs_erode) + label_name = next((r[3] for r in res if r[1] and r[3]), None) + if not label_name: + print("[full] seg produced no label — abort"); return + px = (fp or dp).get("pixel_size_um", fake_pixel_um) # match feature spacing to the seg's pixel size + organelle, chans, sp = label_name, chans0, (px, px) + netorg = [organelle] if network else [] + orgmap = {organelle: marker_channel} + root = zarr.open(zpath, mode="r") + rows = []; n_empty = 0 + for ai in range(na): + lab = np.asarray(root[f"A/{ai}/0/labels/{label_name}/0"][0, 0, 0]).astype(np.int32) # (Y, W) + img = np.asarray(root[f"A/{ai}/0/0"][0, :, 0]) # (C, Y, W) + for c in range(n_cells): + x0 = c * (CROP + PAD) + org_crop = _clip_border(lab[:, x0:x0 + CROP]) # clip border band → measure clipped objects + if org_crop.max() == 0: + n_empty += 1; continue # empty seg → skip (logged below) + inten = img[:, :, x0:x0 + CROP].astype(np.float32) # (C, Y, CROP) + cf, of, _nf = process_single_cell( + cell_info={"global_cell_id": f"{target}_a{ai:02d}_c{c}", "well": f"A/{ai}/0"}, + cell_specific_mask=np.ones((CROP, CROP), np.uint8), + organelle_mask_arrays={organelle: org_crop}, intensity_image=inten, + frangi_image_arrays={}, organelles_to_process=[organelle], network_organelles=netorg, + spacing=sp, channel_names=chans, organelle_map=orgmap, full_features=True) + row = {k: v for k, v in cf.items()} + row["alpha_idx"], row["cell"] = ai, c + odf = of.get(organelle) if of else None + if odf is not None and len(odf): # aggregate per-object features → cell (pipeline AGG_FUNCS) + num = odf.select_dtypes(include="number") + for feat in num.columns: + for fn in AGGREGATION_FUNCTIONS: + row[f"obj_{feat}_{fn}"] = int(num[feat].count()) if fn == "count" else float(getattr(num[feat], fn)()) + rows.append(row) + df = pd.DataFrame(rows) + out = out_root or f"{CACHE}/_morphometrics/{marker_dir}/{grain}/{target}" + os.makedirs(out, exist_ok=True) + fp_out = f"{out}/full_features.parquet" + df.to_parquet(fp_out) + print(f"[full] {marker_dir}/{grain}/{target}: {len(df)} (cell,α) rows × {df.shape[1]} cols " + f"({n_empty} empty-seg cell×α dropped) -> {fp_out}") + print(f" sample feature cols: {[c for c in df.columns if 'network' in c or c.startswith('obj_area')][:8]}") + return fp_out + + +def real_percell(marker_dir, target, grain, store_marker, features, store_channel=None, n_top=1000): + """Per-cell real NTC & KO values for the given generated `features`, from the op_cp_features store, + restricted to the **top-`n_top` set-accuracy cells** (KO + NTC) — the strongest-phenotype real cells + (capped at whatever the ranking/store provide). phase → accuracy/attention parquets; fluor → the + per-marker set-accuracy rank parquet (geneKO/complex). → {feature: {ntc:arr, ko:arr}}. For the violin.""" + import re + import anndata as ad + import numpy as np + import pandas as pd + from ..classifier.config import GRAINS, slugify + feats = list(features) + CH = "mcherry" + for f in feats: + mm = re.match(r"network_(\w+?)_seg_", f) + if mm: + CH = mm.group(1); break + + def store_col(f): + mm = re.match(r"network_(\w+?)_seg_(.+)", f) + if store_marker == "phase": + org = store_channel or "phase2d_tubular" # store_channel routes the seg organelle (e.g. 'nuclei' for nucleus seg) + if mm: + return f"op_network_{org}_{mm.group(2)}" + if f.startswith("obj_"): + return f"op_{org}_count" if f.endswith("_count") else f"op_{org}_{f[len('obj_'):]}" # object count → store's bare count + return None + if mm: + return f"op_network_{store_channel or mm.group(1)}_{mm.group(2)}" + if f.startswith("obj_"): + org = store_channel or CH + return f"op_{org}_count" if f.endswith("_count") else f"op_{org}_{f[len('obj_'):]}" + return None + + a = ad.read_h5ad(OPCP.format(store=store_marker), backed="r") + var = set(map(str, a.var_names)) + mapped = {f: store_col(f) for f in feats} + scols = sorted({v for v in mapped.values() if v in var}) + if not scols: + return {} + gn = a.obs["gene_name"].astype(str).values + obs = a.obs + ekey = obs["experiment"].astype(str).values; wkey = obs["well"].astype(str).values; skey = obs["segmentation"].astype(str).values + + def _nw(w): + w = str(w).strip() + if w.count("/") == 2: + return w + m = re.match(r"^([A-Za-z]+)(\d+)$", w); return f"{m.group(1)}/{m.group(2)}/0" if m else w + + def _topn(idx, pq, val, n): + """Walk class `val`'s cells in set-accuracy rank order and collect store-matched cells (by exp/well/seg) + until n are found — headroom past the top-n covers ranked cells absent from the store.""" + try: + d = pd.read_parquet(pq) + except Exception as e: + print(f" [topn] {val}: {e}"); return np.array([], int) + gcol = "gene" if "gene" in d.columns else "predicted_class" + if "rank_type" in d.columns: + d = d[d["rank_type"] == "top"] + d = d[d[gcol].astype(str) == str(val)].sort_values("rank").drop_duplicates(["experiment", "well", "segmentation"]) + k2i = {} + for i in idx: + k2i.setdefault((ekey[i], _nw(wkey[i]), skey[i]), i) + matched = [] + for e, w, s in zip(d["experiment"].astype(str), d["well"].map(_nw), d["segmentation"].astype(str)): + i = k2i.get((e, w, s)) + if i is not None: + matched.append(i) + if len(matched) >= n: + break + return np.array(matched, int) + + i_ntc = np.where(gn == "")[0] + ko_val = target + if grain == "complex": + import yaml + y = yaml.safe_load(open("/hpc/projects/icd.fast.ops/configs/gene_clusters/EBI_complexes_v1_updated_gene_names.yaml")) + ent = next((e for e in y.values() if slugify(e["name"]) == target), None) + members = ent["genes"] if ent else [] + ko_val = ent["name"] if ent else target # complex rank parquet keys by full complex name + i_ko = np.where(np.isin(gn, members))[0] + else: + i_ko = np.where(gn == target)[0] + + if store_marker == "phase": # phase: accuracy (KO) + top-attention (NTC) + ko_pq = "/hpc/projects/icd.fast.ops/models/diffex/accuracy_ranking/phase_geneKO_topacc_ALL_top1000.parquet" + ntc_pq = GRAINS["geneKO"]["parquet"]; ko_val = target + else: # fluor: the marker's set-accuracy rank parquet + ko_pq = ntc_pq = f"{CACHE}/_rankings/fluor/{'complex' if grain == 'complex' else 'geneKO'}/{marker_dir}.parquet" + ko2, ntc2 = _topn(i_ko, ko_pq, ko_val, n_top), _topn(i_ntc, ntc_pq, "NTC", n_top) + i_ko = ko2 if len(ko2) else i_ko[:n_top] + i_ntc = ntc2 if len(ntc2) else i_ntc[:n_top] + if not len(ko2): + print(f" [real] {target}: KO not accuracy-ranked in {ko_pq} — first {len(i_ko)} store cells") + if not len(ntc2): + print(f" [real] NTC not accuracy-ranked — first {len(i_ntc)} store cells") + + rows_idx = np.concatenate([i_ntc, i_ko]) + sub = a[rows_idx, scols].to_memory().to_df() + lab = np.array([""] * len(i_ntc) + [target] * len(i_ko)) + out = {} + for f, sc in mapped.items(): + if sc in scols: + out[f] = {"ntc": sub.loc[lab == "", sc].values.astype(float), + "ko": sub.loc[lab == target, sc].values.astype(float)} + return out + + +def full_features_cache(marker_dir, target, grain, store_marker, store_channel=None): + """Turn full_features.parquet into a demo-navigable JSON: per-α mean trajectory for EVERY feature + + the real NTC→KO mean±SEM reference (from op_cp_features) for each feature that maps to a store column + + grouped feature lists for the dropdown. → /full_features.json (read by morpho_demo.html).""" + import json + import anndata as ad + import numpy as np + import pandas as pd + base = f"{CACHE}/_morphometrics/{marker_dir}/{grain}/{target}" + df = pd.read_parquet(f"{base}/full_features.parquet") + src = f"{CACHE}/{marker_dir}/{grain}/{target}" + mp = f"{src}/cell0/meta.json" if os.path.exists(f"{src}/cell0/meta.json") else f"{src}/meta.json" + alphas = json.load(open(mp))["alphas"]; na = len(alphas) + num = df.select_dtypes(include="number") + feats = [c for c in num.columns if c not in ("alpha_idx", "cell")] + per = df.groupby("alpha_idx") + agg = {f: [float(per[f].mean().get(ai, np.nan)) for ai in range(na)] for f in feats} + gsem = per[feats].std(ddof=1) / per[feats].count() ** 0.5 # per-α SEM over the generated cells (for demo/figure error bars) + agg_sem = {f: [float(gsem[f].get(ai, 0.0)) for ai in range(na)] for f in feats} + + import re + CH = "mcherry" # infer the channel from the network__seg_ columns + for f in feats: + mm = re.match(r"network_(\w+?)_seg_", f) + if mm: + CH = mm.group(1); break + + def store_col(f): # generated feature name → op_cp store column (channel-aware) + mm = re.match(r"network_(\w+?)_seg_(.+)", f) + if store_marker == "phase": # phase op_cp store: seg organelle via store_channel (tubular default, 'nuclei' for nucleus seg) + org = store_channel or "phase2d_tubular" + if mm: + return f"op_network_{org}_{mm.group(2)}" + if f.startswith("obj_"): + return f"op_{org}_count" if f.endswith("_count") else f"op_{org}_{f[len('obj_'):]}" # object count → store's bare count + return None + if mm: + return f"op_network_{store_channel or mm.group(1)}_{mm.group(2)}" # store_channel: seg organelle name ≠ store's physical channel (e.g. NPM3 nucleoli → 'gfp') + if f.startswith("obj_"): + org = store_channel or CH + return f"op_{org}_count" if f.endswith("_count") else f"op_{org}_{f[len('obj_'):]}" + return None + + from ..classifier.config import slugify + def msem(v): + v = np.asarray(v, float); v = v[~np.isnan(v)] + return [float(v.mean()), float(v.std() / max(1, len(v)) ** 0.5)] if len(v) else None + # TOP-1k set-accuracy NTC/KO (same cells the violin uses) — via real_percell + rp = real_percell(marker_dir, target, grain, store_marker, feats, store_channel=store_channel, n_top=1000) + # ALL-cells NTC/KO (sampled) + population centroid — the second reference alongside top-1k + a = ad.read_h5ad(OPCP.format(store=store_marker), backed="r") + var = set(map(str, a.var_names)) + scols = sorted({store_col(f) for f in feats if store_col(f) in var}) + gn = a.obs["gene_name"].astype(str).values + rng = np.random.default_rng(0) + if grain == "complex": + import yaml + y = yaml.safe_load(open("/hpc/projects/icd.fast.ops/configs/gene_clusters/EBI_complexes_v1_updated_gene_names.yaml")) + members = next((e["genes"] for e in y.values() if slugify(e["name"]) == target), []) + i_ko_all = np.where(np.isin(gn, members))[0] + else: + i_ko_all = np.where(gn == target)[0] + def _samp(idx, n): return rng.choice(idx, n, replace=False) if len(idx) > n else idx + i_ntc_all = _samp(np.where(gn == "")[0], 20000); i_ko_all = _samp(i_ko_all, 20000); i_cen = _samp(np.arange(a.n_obs), 50000) + allm = {"ntc": {}, "ko": {}, "cen": {}} + if scols: + s = a[np.concatenate([i_ntc_all, i_ko_all, i_cen]), scols].to_memory().to_df() + lab = np.array(["ntc"] * len(i_ntc_all) + ["ko"] * len(i_ko_all) + ["cen"] * len(i_cen)) + for g_ in allm: + allm[g_] = {f: msem(s.loc[lab == g_, store_col(f)].values) for f in feats if store_col(f) in scols} + real_ref = {} + for f in feats: + r = rp.get(f) + if r is not None and len(r["ntc"]) and len(r["ko"]): + n, k = msem(r["ntc"]), msem(r["ko"]) + if n and k: + real_ref[f] = {"ntc": n, "ko": k, "ntc_all": allm["ntc"].get(f), + "ko_all": allm["ko"].get(f), "cen": allm["cen"].get(f)} + + def grp(f): + if f.startswith("network_"): + return "Network" + if any(k in f.lower() for k in ["haralick", "glcm", "zernike", "moment", "texture"]): + return "Texture" + if "intensity" in f: + return "Intensity" + if f.startswith("obj_"): + return "Object morphology" + return "Cell" + groups = {} + for f in feats: + groups.setdefault(grp(f), []).append(f) + ncell = int(df["cell"].max()) + 1 if "cell" in df.columns and len(df) else 0 + json.dump(_json_safe({"alphas": alphas, "agg": agg, "agg_sem": agg_sem, "real_ref": real_ref, "groups": groups, "target": target, + "marker_dir": marker_dir, "n_cells": ncell, "n_features": len(feats), "n_with_ref": len(real_ref)}), + open(f"{base}/full_features.json", "w")) + print(f"[full-cache] {marker_dir}/{grain}/{target}: {len(feats)} features ({len(real_ref)} with real NTC→KO ref) " + f"-> {base}/full_features.json") + return f"{base}/full_features.json" + + +# morpho demo targets (SLURM-friendly, picklable). adaptive=False → the perfected org_seg config params. +MORPHO_TARGETS = { + "MICOS13": dict(marker_dir="phase", target="MICOS13", real_exp="ops0047_20250612", marker_channel="Phase2D", + store_marker="phase", grain="geneKO", image_channel="Phase2D", org_label="phase2d_tubular_seg"), + "TOMM20": dict(marker_dir="phase", target="TOMM20", real_exp="ops0047_20250612", marker_channel="Phase2D", + store_marker="phase", grain="geneKO", image_channel="Phase2D", org_label="phase2d_tubular_seg"), + "SAMM50": dict(marker_dir="phase", target="SAMM50", real_exp="ops0047_20250612", marker_channel="Phase2D", + store_marker="phase", grain="geneKO", image_channel="Phase2D", org_label="phase2d_tubular_seg"), + "KIF23_PHASE": dict(marker_dir="phase", target="KIF23", real_exp="ops0047_20250612", marker_channel="Phase2D", + store_marker="phase", grain="geneKO", image_channel="Phase2D", seg_method="nucleus", + org_label="phase2d_nucleus_seg", store_channel="nuclei"), # nuclear-mask seg → real ref from op_nuclei_* store features + "MTOR_PHASE": dict(marker_dir="phase", target="MTOR", real_exp="ops0047_20250612", marker_channel="Phase2D", + store_marker="phase", grain="geneKO", image_channel="Phase2D", org_label="phase2d_tubular_seg"), + "INTRON_SPLICEOSOME_PHASE": dict(marker_dir="phase", target="Intron_Lariat_Spliceosome__type_1_complex", + real_exp="ops0047_20250612", marker_channel="Phase2D", store_marker="phase", grain="complex", + image_channel="Phase2D", structure_type="tubular", org_label="phase2d_tubular_seg"), + "SNRNP200_PHASE": dict(marker_dir="phase", target="SNRNP200", real_exp="ops0047_20250612", marker_channel="Phase2D", + store_marker="phase", grain="geneKO", image_channel="Phase2D", structure_type="vesicular_dark", seg_method="blob", + org_label="phase2d_vesicular_dark_seg", store_channel="phase2d_vesicular_dark", + frangi_override={"min_radius_um": 0.25, "max_radius_um": 0.5}), # drop too-small, max just above original (0.4) + "LAMTOR2_LYSO": dict(marker_dir="lysosome_LysoTracker_live_cell_dye", target="LAMTOR2", real_exp="ops0047_20250612", + marker_channel="lysosome_LysoTracker live-cell dye", store_marker="lysosome_lysotracker_live-cell_dye", + grain="geneKO", image_channel="GFP", structure_type="vesicular", seg_method="blob", org_label="gfp_seg", + store_channel="gfp"), # lysosomes (LysoTracker) — bright vesicular blob + "RAB7A_BODIPY": dict(marker_dir="lipid_droplet_BODIPY_live_cell_dye", target="RAB7A", real_exp="ops0047_20250612", + marker_channel="lipid droplet_BODIPY live cell dye", store_marker="lipid_droplet_bodipy_live_cell_dye", + grain="geneKO", image_channel="GFP", structure_type="vesicular", seg_method="blob", org_label="gfp_seg", + store_channel="gfp"), # lipid droplets (BODIPY) — bright vesicular blob + # --- batch: 10 distinct-phenotype phase geneKO + 3 phase complex + 1 fluor (phenotype-matched seg) --- + "KIF11_PHASE": dict(marker_dir="phase", target="KIF11", real_exp="ops0047_20250612", marker_channel="Phase2D", + store_marker="phase", grain="geneKO", image_channel="Phase2D", seg_method="nucleus", + org_label="phase2d_nucleus_seg", store_channel="nuclei"), # mitotic arrest / rounding + "ATP6V1B2_PHASE": dict(marker_dir="phase", target="ATP6V1B2", real_exp="ops0047_20250612", marker_channel="Phase2D", + store_marker="phase", grain="geneKO", image_channel="Phase2D", structure_type="vesicular_dark", seg_method="blob", + org_label="phase2d_vesicular_dark_seg", store_channel="phase2d_vesicular_dark"), # V-ATPase → vacuoles + "HGS_PHASE": dict(marker_dir="phase", target="HGS", real_exp="ops0047_20250612", marker_channel="Phase2D", + store_marker="phase", grain="geneKO", image_channel="Phase2D", structure_type="vesicular", seg_method="blob", + org_label="phase2d_vesicular_seg", store_channel="phase2d_vesicular"), # ESCRT → enlarged endosomes + "RRM1_PHASE": dict(marker_dir="phase", target="RRM1", real_exp="ops0047_20250612", marker_channel="Phase2D", + store_marker="phase", grain="geneKO", image_channel="Phase2D", seg_method="nucleus", + org_label="phase2d_nucleus_seg", store_channel="nuclei"), # replication stress → nuclei + "RRN3_PHASE": dict(marker_dir="phase", target="RRN3", real_exp="ops0047_20250612", marker_channel="Phase2D", + store_marker="phase", grain="geneKO", image_channel="Phase2D", org_label="phase2d_tubular_seg"), # Pol I / nucleolar + "SEC61A1_PHASE": dict(marker_dir="phase", target="SEC61A1", real_exp="ops0047_20250612", marker_channel="Phase2D", + store_marker="phase", grain="geneKO", image_channel="Phase2D", seg_method="nucleus", + org_label="phase2d_nucleus_seg", store_channel="nuclei"), # object area from nucleus seg (like SON) + "SON_PHASE": dict(marker_dir="phase", target="SON", real_exp="ops0047_20250612", marker_channel="Phase2D", + store_marker="phase", grain="geneKO", image_channel="Phase2D", seg_method="nucleus", + org_label="phase2d_nucleus_seg", store_channel="nuclei"), # nuclear speckles + "AP2M1_PHASE": dict(marker_dir="phase", target="AP2M1", real_exp="ops0047_20250612", marker_channel="Phase2D", + store_marker="phase", grain="geneKO", image_channel="Phase2D", structure_type="tubular", + org_label="phase2d_tubular_seg"), # clathrin/endocytosis → frangi tubular (fused elongated tubules) + "GOLGA2_PHASE": dict(marker_dir="phase", target="GOLGA2", real_exp="ops0047_20250612", marker_channel="Phase2D", + store_marker="phase", grain="geneKO", image_channel="Phase2D", org_label="phase2d_tubular_seg"), # Golgi fragmentation + "SMC2_PHASE": dict(marker_dir="phase", target="SMC2", real_exp="ops0047_20250612", marker_channel="Phase2D", + store_marker="phase", grain="geneKO", image_channel="Phase2D", seg_method="nucleus", + org_label="phase2d_nucleus_seg", store_channel="nuclei"), # condensin → chromosome + "NPC_PHASE": dict(marker_dir="phase", target="Nuclear_pore_complex", real_exp="ops0047_20250612", marker_channel="Phase2D", + store_marker="phase", grain="complex", image_channel="Phase2D", seg_method="nucleus", + org_label="phase2d_nucleus_seg", store_channel="nuclei"), # nuclear pore complex + "PROTEASOME_PHASE": dict(marker_dir="phase", target="19S_proteasome_regulatory_complex", real_exp="ops0047_20250612", + marker_channel="Phase2D", store_marker="phase", grain="complex", image_channel="Phase2D", + structure_type="vesicular_dark", seg_method="blob", org_label="phase2d_vesicular_dark_seg", + store_channel="phase2d_vesicular_dark"), # proteasome → aggregates + "HAUS_PHASE": dict(marker_dir="phase", target="HAUS_complex", real_exp="ops0047_20250612", marker_channel="Phase2D", + store_marker="phase", grain="complex", image_channel="Phase2D", org_label="phase2d_tubular_seg"), # HAUS → spindle/MT + "AP2M1_CLTA": dict(marker_dir="clathrin_vesicles_CLTA", target="AP2M1", real_exp="ops0047_20250612", + marker_channel="clathrin vesicles_CLTA", store_marker="clathrin_vesicles_clta", grain="geneKO", + image_channel="GFP", structure_type="vesicular", seg_method="blob", org_label="gfp_seg", store_channel="gfp"), # clathrin puncta + "EIF2S2_SG": dict(marker_dir="stress_granule_G3BP1", target="EIF2S2", real_exp="ops0047_20250612", + marker_channel="stress granule_G3BP1", store_marker="stress_granule_g3bp1", grain="geneKO", + image_channel="GFP", structure_type="vesicular", seg_method="blob", org_label="gfp_seg", store_channel="gfp"), # stress granules + "AURKB_CHROMATIN": dict(marker_dir="chromatin_H2BC21", target="AURKB", real_exp="ops0047_20250612", + marker_channel="chromatin_H2BC21", store_marker="chromatin_h2bc21", grain="geneKO", + image_channel="mCherry", structure_type="tubular", seg_method="frangi", org_label="mcherry_seg", store_channel="mcherry", + frangi_override={"pixel_size_um": 0.185}), # chromatin (H2BC21) — frangi tubular per org_seg config + "NOP56_FBL": dict(marker_dir="nucleolus_DFC_FBL", target="NOP56", real_exp="ops0047_20250612", + marker_channel="nucleolus-DFC_FBL", store_marker="nucleolus-dfc_fbl", grain="geneKO", + image_channel="GFP", structure_type="vesicular", seg_method="frangi", org_label="gfp_seg", store_channel="gfp", + frangi_override={"pixel_size_um": 0.85, "beta": 0.1, "min_radius_um": 0.5, "max_radius_um": 5.0, "min_object_size": 8}), # nucleolus-DFC (FBL) — frangi vesicular per org_seg config + "ATG9A_AUTOPHAGO": dict(marker_dir="autophagosome_MAP1LC3B", target="ATG9A", real_exp="ops0047_20250612", + marker_channel="autophagosome_MAP1LC3B", store_marker="autophagosome_map1lc3b", grain="geneKO", + image_channel="GFP", structure_type="vesicular", seg_method="blob", org_label="gfp_seg", store_channel="gfp", + frangi_override={"threshold": 0.10, "min_radius_um": 0.10, "max_radius_um": 1.0}), # autophagosomes (LC3 puncta) — blob per MAP1LC3B config (ops0090) + "ATP6V1B2_LAMP1": dict(marker_dir="lysosome_LAMP1", target="ATP6V1B2", real_exp="ops0047_20250612", + marker_channel="lysosome_LAMP1", store_marker="lysosome_lamp1", grain="geneKO", + image_channel="GFP", structure_type="vesicular", seg_method="blob", org_label="gfp_seg", store_channel="gfp"), # lysosomes (LAMP1) + "RAB7A_PHASE": dict(marker_dir="phase", target="RAB7A", real_exp="ops0047_20250612", marker_channel="Phase2D", + store_marker="phase", grain="geneKO", image_channel="Phase2D", structure_type="vesicular", seg_method="blob", + org_label="phase2d_vesicular_seg", store_channel="phase2d_vesicular", + frangi_override={"threshold": 0.03, "min_radius_um": 0.1}), # lower blob thr/min-radius so it fires on the α=0 frame + "HSPA5_PHASE": dict(marker_dir="phase", target="HSPA5", real_exp="ops0047_20250612", marker_channel="Phase2D", + store_marker="phase", grain="geneKO", image_channel="Phase2D", structure_type="vesicular_dark", seg_method="blob", + org_label="phase2d_vesicular_dark_seg", store_channel="phase2d_vesicular_dark"), # dark vacuoles (ER stress) — blob, not frangi + # CCT/NPM3: nucleoli are round blobs → frangi vesselness under-/over-detects them; use the MASKED-OBJECT + # intensity path (from coding_exps/nucleoli_roundness, tuned NPM3 MO_PARAMS). pixel_size only sets feature spacing. + "CCT": dict(marker_dir="nucleolus_GC_NPM3", target="Chaperonin_containing_T_complex", real_exp="ops0092_20251027", + marker_channel="nucleolus-GC_NPM3", store_marker="nucleolus-gc_npm3", grain="complex", + image_channel="GFP", org_label="gfp_seg", store_channel="gfp", structure_type="vesicular", + seg_method="masked_object", mo_nucleus=True, frangi_override={"pixel_size_um": 0.825}), + # POLR1B geneKO, same NPM3 nucleoli marker + nucleus-constrained MO seg as CCT + "POLR1B": dict(marker_dir="nucleolus_GC_NPM3", target="POLR1B", real_exp="ops0092_20251027", + marker_channel="nucleolus-GC_NPM3", store_marker="nucleolus-gc_npm3", grain="geneKO", + image_channel="GFP", org_label="gfp_seg", store_channel="gfp", structure_type="vesicular", + seg_method="masked_object", mo_nucleus=True, frangi_override={"pixel_size_um": 0.825}), + # GBF1 geneKO on ER/Golgi COP-II (SEC23A) puncta — MO intensity seg (punctate blobs, like nucleoli) + "GBF1": dict(marker_dir="ER_Golgi_COP_II_SEC23A", target="GBF1", real_exp="ops0081_20250924", + marker_channel="ER_Golgi_COP-II_SEC23A", store_marker="er_golgi_cop-ii_sec23a", grain="geneKO", + image_channel="GFP", org_label="gfp_seg", store_channel="gfp", structure_type="vesicular", + seg_method="masked_object"), + # NucleoLIVE (org_seg_params NucleoLIVE variant): tubular on mCherry; lower pixel than real 0.185 to suppress generated background + "KIF23_NUCLEOLIVE": dict(marker_dir="nucleus_NucleoLIVE_Live_Cell_dye", target="KIF23", real_exp="ops0120_20260204", + marker_channel="nucleus_NucleoLIVE Live Cell dye", store_marker="nuclei_nucleolive_live_cell_dye", grain="geneKO", + image_channel="mCherry", org_label="mcherry_seg", store_channel="mcherry", structure_type="tubular", + frangi_override={"pixel_size_um": 0.09, "min_object_size": 2}), + "TIM23_PHASE": dict(marker_dir="phase", target="TIM23_mitochondrial_inner_membrane_pre_sequence_translocase_complex__TIM17A_variant", + real_exp="ops0047_20250612", marker_channel="Phase2D", store_marker="phase", grain="complex", + image_channel="Phase2D", org_label="phase2d_tubular_seg", structure_type="tubular"), + "CHROMALIVE_TIM23": dict(marker_dir="mitochondria_ChromaLIVE_561_excitation", + target="TIM23_mitochondrial_inner_membrane_pre_sequence_translocase_complex__TIM17A_variant", real_exp="ops0122_20260211", + marker_channel="mitochondria_ChromaLIVE 561 excitation", store_marker="mitochondria_chromalive_561_excitation", + grain="complex", image_channel="mCherry", org_label="mcherry_seg", store_channel="mcherry", structure_type="tubular"), + "CHROMALIVE_MICOS13": dict(marker_dir="mitochondria_ChromaLIVE_561_excitation", target="MICOS13", real_exp="ops0122_20260211", + marker_channel="mitochondria_ChromaLIVE 561 excitation", store_marker="mitochondria_chromalive_561_excitation", + grain="geneKO", image_channel="mCherry", org_label="mcherry_seg", store_channel="mcherry", structure_type="tubular"), + "CHROMALIVE_TOMM20": dict(marker_dir="mitochondria_ChromaLIVE_561_excitation", target="TOMM20", real_exp="ops0122_20260211", + marker_channel="mitochondria_ChromaLIVE 561 excitation", store_marker="mitochondria_chromalive_561_excitation", + grain="geneKO", image_channel="mCherry", org_label="mcherry_seg", store_channel="mcherry", structure_type="tubular"), + # FastAct actin filaments (mCherry, tubular → real config = frangi, like ChromaLIVE mito) + "ARP23_FASTACT": dict(marker_dir="actin_filament_FastAct_SPY555_Live_Cell_Dye", + target="Actin_related_protein_2_3_complex__ARPC1A_ACTR3B_ARPC5_variant", real_exp="ops0076_20250917", + marker_channel="actin filament_FastAct SPY555 Live Cell Dye", store_marker="actin_filament_fastact_spy555_live_cell_dye", + grain="complex", image_channel="mCherry", org_label="mcherry_seg", store_channel="mcherry", structure_type="tubular", + frangi_override={"pixel_size_um": 0.06}), # lower px → coarser frangi → traces filaments, drops generated-image noise + "CAPZB_PHALLOIDIN": dict(marker_dir="F_actin_Phalloidin", target="CAPZB", real_exp="ops0094_20251217", + marker_channel="CP1_f_actin_Phalloidin", store_marker="f-actin_phalloidin", grain="geneKO", + image_channel="CP1_f_actin_Phalloidin", org_label="cp1_f_actin_phalloidin_seg", + store_channel="cp1_f_actin_phalloidin", structure_type="tubular", + frangi_override={"pixel_size_um": 0.06}), # F-actin filaments (tubular), like FastAct + "CAPZB_FASTACT": dict(marker_dir="actin_filament_FastAct_SPY555_Live_Cell_Dye", target="CAPZB", real_exp="ops0076_20250917", + marker_channel="actin filament_FastAct SPY555 Live Cell Dye", store_marker="actin_filament_fastact_spy555_live_cell_dye", + grain="geneKO", image_channel="mCherry", org_label="mcherry_seg", store_channel="mcherry", structure_type="tubular", + frangi_override={"pixel_size_um": 0.06}), + # PSMB6 geneKO on proteasome PSMB7 marker — frangi (per PSMB7 org_seg config); phenotype = nucleus→cytoplasm reorg + "PSMB6_PROTEASOME": dict(marker_dir="proteasome_PSMB7", target="PSMB6", real_exp="ops0047_20250612", + marker_channel="proteasome_PSMB7", store_marker="proteasome_psmb7", grain="geneKO", + image_channel="GFP", org_label="gfp_seg", store_channel="gfp", structure_type="tubular", seg_method="frangi", + frangi_override={"pixel_size_um": 0.175, "min_object_size": 30}), + # --- phase NUCLEOLI (nucleus-masked MO intensity → dense intranuclear bodies); real ref = op_nucleoli_phase2d_* --- + "POLR1B_NUCLEOLI_PHASE": dict(marker_dir="phase", target="POLR1B", real_exp="ops0047_20250612", marker_channel="Phase2D", + store_marker="phase", grain="geneKO", image_channel="Phase2D", org_label="phase2d_nucleoli_seg", + store_channel="nucleoli_phase2d", structure_type="vesicular", seg_method="masked_object", mo_nucleus=True, mo_vs_nucleus=True, mo_vs_erode=4), + "ZNRD1_NUCLEOLI_PHASE": dict(marker_dir="phase", target="ZNRD1", real_exp="ops0047_20250612", marker_channel="Phase2D", + store_marker="phase", grain="geneKO", image_channel="Phase2D", org_label="phase2d_nucleoli_seg", + store_channel="nucleoli_phase2d", structure_type="vesicular", seg_method="masked_object", mo_nucleus=True, mo_nucleus_scale=0.6), # POLR1H = ZNRD1 + "PROTEASOME_NUCLEOLI_PHASE": dict(marker_dir="phase", target="19S_proteasome_regulatory_complex", real_exp="ops0047_20250612", + marker_channel="Phase2D", store_marker="phase", grain="complex", image_channel="Phase2D", org_label="phase2d_nucleoli_seg", + store_channel="nucleoli_phase2d", structure_type="vesicular", seg_method="masked_object", mo_nucleus=True, mo_nucleus_scale=0.6), +} + + +def build_morpho(marker_dir, target, real_exp, marker_channel, store_marker, grain="geneKO", n_cells=12, + image_channel=None, org_label=None, adaptive=False, store_channel=None, frangi_override=None, structure_type=None, seg_method=None, mo_nucleus=False, mo_nucleus_scale=1.0, mo_vs_nucleus=False, mo_vs_erode=0): + """One morpho-demo target end to end (the wrapper that bundles EVERYTHING into one target dir, so nothing is + scattered): generated seg-mask overlays (a*_labels.png) + per-object feats (a*_feats.json) + full org-profiler + feature table (full_features.parquet) + demo JSON w/ per-α trajectory + top-accuracy store real-ref stats + (full_features.json) + CACHED production-label real-cell images & seg overlays (_ref/). adaptive=False = perfected config. + mo_vs_nucleus=True → nucleus mask from VS-H2B → Cellpose (cached npz), not the Otsu/ellipse mask.""" + from ..classifier.config import slugify + base_dir = f"{SYNTH_BASE}/job_{slugify(marker_dir)}_{grain}_{slugify(target)}" # per-(marker,grain,target) staging → same gene on 2 markers won't race the shared zarr + vs_npz = None + if mo_vs_nucleus: # BEFORE the OPS_OUTPUT_BASE_DIR override — CellDino/VS ckpt paths key off it + vs_npz = f"{CACHE}/_morphometrics/{marker_dir}/{grain}/{target}/vs_nucleus.npz" + os.makedirs(os.path.dirname(vs_npz), exist_ok=True) + _vs_h2b_nucleus_npz(marker_dir, target, grain, vs_npz, n_cells) + os.environ["OPS_OUTPUT_BASE_DIR"] = base_dir + run_target(marker_dir, target, real_exp, marker_channel, store_marker=store_marker, grain=grain, n_cells=n_cells, + refs=False, adaptive=adaptive, frangi_override=frangi_override, structure_type=structure_type, seg_method=seg_method, base_dir=base_dir, org_label=org_label, mo_nucleus=mo_nucleus, mo_nucleus_scale=mo_nucleus_scale, vs_nucleus_npz=vs_npz, vs_erode=mo_vs_erode) # generated seg-mask overlays + per-object feats + full_features(marker_dir, target, real_exp, marker_channel, grain=grain, n_cells=n_cells, adaptive=adaptive, + frangi_override=frangi_override, structure_type=structure_type, seg_method=seg_method, base_dir=base_dir, org_label=org_label, mo_nucleus=mo_nucleus, mo_nucleus_scale=mo_nucleus_scale, vs_nucleus_npz=vs_npz, vs_erode=mo_vs_erode) + full_features_cache(marker_dir, target, grain, store_marker=store_marker, store_channel=store_channel) + if image_channel and org_label: # cached real-cell images w/ PRODUCTION org labels (built-time crop from _v3) + exps = None + if store_marker != "phase": # fluor markers: restrict real cells to experiments where this channel IS the reporter + import anndata as ad + exps = set(ad.read_h5ad(OPCP.format(store=store_marker), backed="r").obs["experiment"].astype(str).unique()) + reference_from_store(marker_dir, target, grain, image_channel, org_label, + n_cells=min(4, n_cells), network=("tubular" in org_label), store_exps=exps) + return f"{CACHE}/_morphometrics/{marker_dir}/{grain}/{target}" + + +def build_morpho_target(key, n_cells=12): + return build_morpho(n_cells=n_cells, **MORPHO_TARGETS[key]) + + +def build_phase_morpho(target, real_exp="ops0047_20250612", n_cells=12, grain="geneKO"): # back-compat wrapper + return build_morpho("phase", target, real_exp, "Phase2D", "phase", grain=grain, n_cells=n_cells, + image_channel="Phase2D", org_label="phase2d_tubular_seg") + + +if __name__ == "__main__": + validate() diff --git a/src/ops_model/models/attention/diffex/viewer/morphometrics.py b/src/ops_model/models/attention/diffex/viewer/morphometrics.py new file mode 100644 index 0000000..44cbf4c --- /dev/null +++ b/src/ops_model/models/attention/diffex/viewer/morphometrics.py @@ -0,0 +1,170 @@ +"""Morphometrics on GENERATED cells: run each marker's organelle segmentation (org_seg_params) on +the generated α-frames (whole crop, no cell mask), measure classical org-profiler features, and +compare their α-trajectory to the REAL NTC→KO shift (op_cp_features) to show the generated images +AMPLIFY the real morphometric change. + +`preview()` renders the design: org-seg masks + per-organelle feature colormaps across α, so we can +eyeball that segmentation + feature measurement work before scaling the full compute/cache. +""" +from __future__ import annotations + +import numpy as np +import yaml + +ORG_SEG_YAML = "/hpc/projects/intracellular_dashboard/fast_ops/configs/org_seg_params.yaml" +CACHE = "/hpc/projects/icd.fast.ops/models/diffex/viewer_assets" +PIXEL_UM = 0.325 # phenotype native pixel size + + +_DEFAULT_FRANGI = {"method": "frangi", "frangi": {"pixel_size_um": 0.1, "min_object_size": 10, "postprocess": True}} + + +def seg_config(marker_key): + """org_seg_params block for a marker key (e.g. 'LAMP1') → (method, params dict). Markers with no + config (e.g. ChromaLIVE) fall back to a generic frangi (tubular/network-ish).""" + y = yaml.safe_load(open(ORG_SEG_YAML)) + for blk in y.get(marker_key, []): + if isinstance(blk, dict) and "segmentation_config" in blk: + sc = blk["segmentation_config"] + return sc.get("method", "frangi"), sc + return "frangi", _DEFAULT_FRANGI + + +def seg_organelles(crop, marker_key, pixel_um=PIXEL_UM): + """Generated crop (H,W float) → labeled organelle mask via the marker's org_seg_params method. + blob → LoG (discrete vesicles); frangi/tubular → vesselness ridge → connected components (fragments, + so `count` = network fragmentation for mito/MICOS).""" + method, sc = seg_config(marker_key) + if method == "blob": + from organelle_profiler.organelle_seg.blob_detection import _segment_blob_log + return _segment_blob_log(crop.astype(np.float32), pixel_um, sc["blob"]) + from skimage.filters import frangi, threshold_otsu + from skimage.measure import label + from skimage.morphology import remove_small_objects + v = frangi(crop.astype(np.float32), black_ridges=False) + if not (v > 0).any(): + return np.zeros(crop.shape, np.int32) + m = v > threshold_otsu(v[v > 0]) + m = remove_small_objects(m, min_size=8) + return label(m).astype(np.int32) + + +def _org_feats(labels, intensity): + """Per-organelle regionprops → dict of {label: {area, mean_int, eccentricity}} + aggregates.""" + from skimage.measure import regionprops + rp = regionprops(labels, intensity_image=intensity) + feats = {r.label: {"area": r.area, "mean_int": r.intensity_mean, + "ecc": r.eccentricity if r.area >= 5 else 0.0} for r in rp} + agg = {"count": len(rp), "total_area": sum(f["area"] for f in feats.values()), + "mean_int": float(np.mean([f["mean_int"] for f in feats.values()])) if feats else 0.0} + return feats, agg + + +def _paint(labels, feats, key, cmap): + """Color each organelle by feats[label][key] → RGB image (background black).""" + import matplotlib.cm as cm + vals = np.array([f[key] for f in feats.values()]) if feats else np.array([0.0]) + lo, hi = float(vals.min()), float(vals.max() + 1e-9) + rgb = np.zeros((*labels.shape, 3), np.float32) + mp = cm.get_cmap(cmap) + for lab, f in feats.items(): + rgb[labels == lab] = mp((f[key] - lo) / (hi - lo))[:3] + return rgb + + +AGG_KEYS = ["count", "total_area", "mean_area", "mean_int", "mean_ecc"] + + +def _aggregate(feats): + if not feats: + return {k: 0.0 for k in AGG_KEYS} + A = [f["area"] for f in feats.values()]; I = [f["mean_int"] for f in feats.values()]; E = [f["ecc"] for f in feats.values()] + return {"count": len(feats), "total_area": float(np.sum(A)), "mean_area": float(np.mean(A)), + "mean_int": float(np.mean(I)), "mean_ecc": float(np.mean(E))} + + +def compute_target(marker_key, marker_dir, target, grain="geneKO", n_cells=6): + """Cache morphometrics for one traversal: per (cell, α) the org-seg label mask (PNG) + per-organelle + features (JSON) + per-α bag aggregates (for the plot). Generated side only (real NTC/KO ref added + once process_single_cell matches the op_cp_features schema). → viewer_assets/_morphometrics////""" + import json + import os + from PIL import Image + base = f"{CACHE}/{marker_dir}/{grain}/{target}" + mp = f"{base}/cell0/meta.json" if os.path.exists(f"{base}/cell0/meta.json") else f"{base}/meta.json" + alphas = json.load(open(mp))["alphas"] + out = f"{CACHE}/_morphometrics/{marker_dir}/{grain}/{target}" + os.makedirs(out, exist_ok=True) + agg_per_cell = [] # [cell][alpha] = agg dict + for c in range(n_cells): + cdir = f"{out}/cell{c}"; os.makedirs(cdir, exist_ok=True) + cell_feats, cell_agg = {}, [] + for ai in range(len(alphas)): + f = f"{base}/cell{c}/frame_{ai:02d}.webp" + if not os.path.exists(f): + cell_agg.append({k: 0.0 for k in AGG_KEYS}); continue + im = np.asarray(Image.open(f).convert("L"), np.float32) / 255.0 + labels = seg_organelles(im, marker_key) + feats, _ = _org_feats(labels, im) + Image.fromarray(labels.astype(np.uint16)).save(f"{cdir}/a{ai:02d}_labels.png") + cell_feats[ai] = {str(k): {kk: float(vv) for kk, vv in v.items()} for k, v in feats.items()} + cell_agg.append(_aggregate(feats)) + json.dump(cell_feats, open(f"{cdir}/feats.json", "w")) + agg_per_cell.append(cell_agg) + # bag mean per α across cells + agg = {k: [float(np.mean([agg_per_cell[c][ai][k] for c in range(n_cells)])) for ai in range(len(alphas))] for k in AGG_KEYS} + json.dump({"marker_key": marker_key, "marker_dir": marker_dir, "target": target, "grain": grain, + "alphas": alphas, "n_cells": n_cells, "agg": agg, "features": AGG_KEYS}, + open(f"{out}/morpho.json", "w")) + print(f"[morpho] {marker_dir}/{grain}/{target}: {n_cells} cells × {len(alphas)} α cached -> {out}") + print(f" count α-series: {[round(v,1) for v in agg['count']]}") + print(f" total_area α-series: {[round(v,0) for v in agg['total_area']]}") + return out + + +def preview(marker_key="LAMP1", marker_dir="lysosome_LAMP1", target="ABCE1", cell=0, + alpha_idxs=(0, 6, 8, 12, 16), out="/hpc/projects/icd.fast.ops/models/diffex/morpho_preview.png"): + """Design preview: rows = α; cols = [raw crop | org-seg outlines | organelles colored by area | + by mean-intensity | by eccentricity]. Title per row = aggregate (count, total area).""" + import matplotlib + matplotlib.use("Agg") + import matplotlib.pyplot as plt + from PIL import Image + from skimage.segmentation import find_boundaries + plt.rcParams["pdf.fonttype"] = 42 + + import json + base = f"{CACHE}/{marker_dir}/geneKO/{target}" + meta = json.load(open(f"{base}/cell{cell}/meta.json")) if __import__("os").path.exists(f"{base}/cell{cell}/meta.json") \ + else json.load(open(f"{base}/meta.json")) + alphas = meta["alphas"] + cols = ["raw crop", "org-seg", "→ area", "→ mean intensity", "→ eccentricity"] + fig, ax = plt.subplots(len(alpha_idxs), len(cols), figsize=(len(cols) * 2.4, len(alpha_idxs) * 2.4)) + for r, ai in enumerate(alpha_idxs): + im = np.asarray(Image.open(f"{base}/cell{cell}/frame_{ai:02d}.webp").convert("L"), np.float32) / 255.0 + labels = seg_organelles(im, marker_key) + feats, agg = _org_feats(labels, im) + panels = [(im, "gray", None), (find_boundaries(labels), "gray", None), + (_paint(labels, feats, "area", "viridis"), None, None), + (_paint(labels, feats, "mean_int", "inferno"), None, None), + (_paint(labels, feats, "ecc", "plasma"), None, None)] + for c, (img, cmap, _) in enumerate(panels): + a = ax[r, c] + if c == 1: + a.imshow(im, cmap="gray"); a.imshow(np.ma.masked_where(~img, img), cmap="autumn", alpha=0.9) + else: + a.imshow(img, cmap=cmap) + a.set_xticks([]); a.set_yticks([]) + if r == 0: + a.set_title(cols[c], fontsize=9) + if c == 0: + a.set_ylabel(f"α={alphas[ai]:+.0f}\n{agg['count']} org\narea {int(agg['total_area'])}", fontsize=8) + fig.suptitle(f"Morphometrics preview — {marker_key} / {target} / cell{cell}: org-seg + per-organelle feature colormaps across α", + fontsize=11) + fig.tight_layout() + fig.savefig(out, dpi=130, bbox_inches="tight") + print(f"[morpho-preview] {out}") + + +if __name__ == "__main__": + preview() diff --git a/src/ops_model/models/attention/diffex/viewer/nway_clf.py b/src/ops_model/models/attention/diffex/viewer/nway_clf.py new file mode 100644 index 0000000..1963bf9 --- /dev/null +++ b/src/ops_model/models/attention/diffex/viewer/nway_clf.py @@ -0,0 +1,120 @@ +"""N-way per-(marker, grain) single-cell classifiers for the DiffEx viewer score. + +The viewer badge should read "does the classifier call this generated cell the target +CLASS, out of all classes" (1-of-N distinctiveness) — NOT "target vs NTC". So for each +(marker, grain) we train one MLPHead over ALL classes on CellDINO features of top-attention +cells. A generated cell is then scored: image → CellDINO (embed_crops) → MLP → softmax → +P(target). + +Trained on embed_crops features so it lives in the SAME space the viewer re-encodes +generated cells into (no domain mismatch). Outputs under +/_clf///: mlp.pt, classes.json, metrics.json. +""" +from __future__ import annotations + +import json +from pathlib import Path + +import numpy as np +import pandas as pd +import torch +import torch.nn as nn +from torch.utils.data import DataLoader, TensorDataset + +from ..classifier.celldino_features import embed_crops +from ..classifier.config import Config, GRAINS, slugify +from ..classifier.data import _BASE_COLS, make_labels_df, materialize_crops +from ..classifier.models import MLPHead + + +def _all_class_table(cfg, marker_channel, fluor_csv, n_per_class): + """Top-attention cells of EVERY class (incl NTC) with integer class labels.""" + cc = cfg.class_col + if marker_channel: + cols = list(dict.fromkeys([cc, *_BASE_COLS, "channel", "rank_type"])) + rows = pd.read_csv(fluor_csv, usecols=cols) + rows = rows[(rows["channel"] == marker_channel) & (rows["rank_type"] == "top")] + else: + rows = pd.read_parquet(cfg.pma_parquet, filters=[("rank_type", "==", "top")], + columns=[cc, *_BASE_COLS]) + rows = rows.sort_values("rank").groupby(cc, group_keys=False).head(n_per_class) + classes = sorted(map(str, rows[cc].unique())) + idx = {c: i for i, c in enumerate(classes)} + rows = rows.rename(columns={cc: "cls"}).copy() + rows["cls"] = rows["cls"].astype(str) + rows["label"] = rows["cls"].map(idx) + return rows, classes + + +def _grouped_split(experiment, seed=0, val_frac=0.15): + """Hold out whole experiments for val (confound guard); fall back to random if too few.""" + exps = np.array(sorted(set(experiment))) + rng = np.random.RandomState(seed) + if len(exps) >= 4: + rng.shuffle(exps) + n_val = max(1, int(round(len(exps) * val_frac))) + val_exps = set(exps[:n_val]) + va = np.array([e in val_exps for e in experiment]) + else: + va = rng.rand(len(experiment)) < val_frac + return ~va, va + + +def _topk(model, X, y, dev, bs=512, k=5): + model.eval() + t1 = t5 = 0 + with torch.no_grad(): + for i in range(0, len(X), bs): + logits = model(torch.as_tensor(X[i:i + bs]).to(dev)) + top = logits.topk(min(k, logits.shape[1]), -1).indices.cpu().numpy() + yb = y[i:i + bs] + t1 += int((top[:, 0] == yb).sum()) + t5 += int([yy in tr for yy, tr in zip(yb, top)].count(True)) + return t1 / len(X), t5 / len(X) + + +def train_nway(grain, out_root, marker_channel=None, channel="Phase2D", fluor_csv=None, + n_per_class=100, epochs=30, device="cuda", load_workers=12, hidden=256): + dev = torch.device(device if torch.cuda.is_available() else "cpu") + cfg = Config(class_col=GRAINS[grain]["class_col"], pma_parquet=GRAINS[grain]["parquet"], + channel=channel, num_workers=load_workers) + modality = slugify(marker_channel) if marker_channel else "phase" + out = Path(out_root) / "_clf" / modality / grain + out.mkdir(parents=True, exist_ok=True) + + df, classes = _all_class_table(cfg, marker_channel, fluor_csv, n_per_class) + print(f"[nway] {modality}/{grain}: {len(classes)} classes, {len(df)} cells") + ldf = make_labels_df(df, cfg) + images, labels, experiment = materialize_crops(ldf, cfg, cache_path=str(out / "crops.npz")) + X = embed_crops(images, cfg, cache_path=str(out / "celldino.npz")).astype(np.float32) + y = labels.astype(np.int64) + + tr, va = _grouped_split(experiment, seed=cfg.seed) + net = MLPHead(in_dim=X.shape[1], hidden=hidden, n_classes=len(classes)).to(dev) + opt = torch.optim.AdamW(net.parameters(), lr=cfg.lr, weight_decay=cfg.weight_decay) + crit = nn.CrossEntropyLoss() + loader = DataLoader(TensorDataset(torch.as_tensor(X[tr]), torch.as_tensor(y[tr])), + batch_size=cfg.batch_size, shuffle=True) + best = {"top1": -1, "state": None} + for ep in range(epochs): + net.train() + for xb, yb in loader: + opt.zero_grad(); loss = crit(net(xb.to(dev)), yb.to(dev)); loss.backward(); opt.step() + t1, t5 = _topk(net, X[va], y[va], dev) + if t1 > best["top1"]: + best = {"top1": t1, "top5": t5, "epoch": ep, + "state": {k: v.detach().cpu() for k, v in net.state_dict().items()}} + print(f" ep{ep:02d}: val top1={t1:.3f} top5={t5:.3f} (best {best['top1']:.3f})") + net.load_state_dict(best["state"]) + torch.save({"state": best["state"], "in_dim": int(X.shape[1]), "hidden": hidden, + "n_classes": len(classes)}, out / "mlp.pt") + (out / "classes.json").write_text(json.dumps(classes)) + meta = {"grain": grain, "modality": modality, "marker_channel": marker_channel, + "channel": channel, "n_classes": len(classes), "n_cells": int(len(X)), + "n_per_class": n_per_class, "val_top1": best["top1"], "val_top5": best["top5"], + "best_epoch": best["epoch"]} + (out / "metrics.json").write_text(json.dumps(meta, indent=2)) + # crops cache is large + one-time; drop it (keep celldino features, small + reusable) + (out / "crops.npz").unlink(missing_ok=True) + print(json.dumps(meta, indent=2)) + return meta diff --git a/src/ops_model/models/attention/diffex/viewer/phenotype_cells.py b/src/ops_model/models/attention/diffex/viewer/phenotype_cells.py new file mode 100644 index 0000000..3f5e66c --- /dev/null +++ b/src/ops_model/models/attention/diffex/viewer/phenotype_cells.py @@ -0,0 +1,106 @@ +"""Enumerate the PHENOTYPE cells the viewer traverses toward: for every (marker × perturbation), +the top-`n_cells` attention-ranked real cells. Perturbations = geneKOs + EBI complexes; markers = +phase + all fluor channels. A handoff list for Ritvik to compute SetTransformer attention pixel-patches +on the real phenotype cells (to compare what the classifier attends to vs the generative morph). + +Grid = n_cells × (1000 geneKO + ~100 complex) × markers. `segmentation_id` = the pma-source +`segmentation` value → (experiment, well, segmentation_id) is the unique cell lookup key. +Sources: phase = v4 parquets (rank pushdown); fluor = v4 CSVs (geneKO=pma_fluorescent_cells_all, +complex=pma_fluorescent_cells_ebi_all), read chunked + filtered to the top ranks. +""" +from __future__ import annotations + +import glob +import os + +import pandas as pd +import pyarrow.parquet as pq + +from ..classifier.config import PMA_PHASE_EBI, PMA_PHASE_GENEKO +from ..classifier.data import _BASE_COLS +from . import catalog as C + +_V4 = os.path.dirname(C.EBI_FLUOR_CSV) +FLUOR_GENEKO_CSV = f"{_V4}/pma_fluorescent_cells_all.csv" +FLUOR_COMPLEX_CSV = C.EBI_FLUOR_CSV +OUT_CSV = f"{C.OUT}/viewer_assets/phenotype_cells_for_attention.csv" +# rank_source: "model" = ranked by the marker's SetTransformer attention (rank/pma_attention are real); +# "fallback" = marker/perturbation not in the attention model (e.g. cisGolgi) → cells from another source, +# rank/pma_attention/map_score are NOT model-derived (marked so consumers don't trust them as attention). +COLS = ["marker_channel", "grain", "perturbation", "geneKO", "ebi_complex", "map_score", "rank_source", + "cell_index", "experiment", "well", "segmentation_id", "x_pheno", "y_pheno", "rank", "pma_attention"] + + +def _gene2cx(): + import yaml + y = yaml.safe_load(open(C.EBI_YAML)) or {} + return {g: v["name"] for v in y.values() if isinstance(v, dict) for g in (v.get("genes") or [])} + + +def _fmt(df, mc, grain, pert_col, gene2cx, n_cells, score=None, rank_source="model"): + df = df.sort_values("rank").copy() + df["marker_channel"] = mc; df["grain"] = grain; df["rank_source"] = rank_source + df["perturbation"] = df[pert_col].astype(str) + if "gene" not in df: + df["gene"] = df["perturbation"] + df["ebi_complex"] = df["gene"].astype(str).map(gene2cx).fillna("") + df["map_score"] = df["perturbation"].map(score) if score is not None else float("nan") # dist(geneKO)/EBI mAP(complex) + df["cell_index"] = df.groupby(["marker_channel", "grain", "perturbation"]).cumcount() + df = df[df["cell_index"] < n_cells] + return df.rename(columns={"segmentation": "segmentation_id", "gene": "geneKO"})[COLS] + + +def _csv_top(path, group_cols, n_cells, chunk=3_000_000): + """Chunked read of a big fluor CSV → each MARKER's top-`n_cells` cells per group. `rank` is GLOBAL + per-geneKO (not per marker), so we re-rank WITHIN each (channel, group): take the n_cells + lowest-rank (highest-attention) cells present in that channel. Two-pass head() keeps memory bounded.""" + use = list(dict.fromkeys(["gene", "channel", "predicted_class", "rank_type", *_BASE_COLS])) + keep = [] + for ch in pd.read_csv(path, usecols=use, chunksize=chunk): + ch = ch[ch["rank_type"] == "top"] + keep.append(ch.sort_values("rank").groupby(group_cols, sort=False).head(n_cells)) + return pd.concat(keep, ignore_index=True).sort_values("rank").groupby(group_cols, sort=False).head(n_cells) + + +def build(n_cells=20, out_csv=OUT_CSV, fluor=True, map_thr=0.01): + """Phase = ALL perturbations. Fluor = only perturbations the marker distinguishes: geneKO by gene + distinctiveness >= map_thr, complex by EBI complex mAP >= map_thr (per the marker's reporter).""" + g2c = _gene2cx() + dist = C.dist_matrix(); cdist = C.complex_dist() + parts = [] + filt = [("rank_type", "==", "top"), ("rank", "<=", n_cells)] + + gk = pq.read_table(PMA_PHASE_GENEKO, columns=["gene", *_BASE_COLS], filters=filt).to_pandas() + parts.append(_fmt(gk, "phase", "geneKO", "gene", g2c, n_cells, score=dist.get("Phase"))) + cx = pq.read_table(PMA_PHASE_EBI, columns=["gene", "predicted_class", *_BASE_COLS], filters=filt).to_pandas() + parts.append(_fmt(cx, "phase", "complex", "predicted_class", g2c, n_cells, score=cdist.get("Phase"))) + print(f"[phenotype] phase (ALL): {len(parts[0])} geneKO + {len(parts[1])} complex cells") + + if fluor: + fg = _csv_top(FLUOR_GENEKO_CSV, ["channel", "gene"], n_cells) + for mc, sub in fg.groupby("channel"): # geneKO: keep genes with distinctiveness >= thr + rep = C.rep_of(dist, str(mc)) + if rep not in dist.columns: # no distinctiveness mAP (e.g. 4i antibodies) → exclude + continue + sc = dist[rep] + sub = sub[sub["gene"].astype(str).isin(set(sc.index[sc >= map_thr]))] + parts.append(_fmt(sub, str(mc), "geneKO", "gene", g2c, n_cells, score=sc)) + fc = _csv_top(FLUOR_COMPLEX_CSV, ["channel", "predicted_class"], n_cells) + for mc, sub in fc.groupby("channel"): # complex: keep complexes with EBI mAP >= thr + rep = C.rep_of(dist, str(mc)) + if not rep or rep not in cdist.columns: # no EBI mAP → exclude + continue + sc = cdist[rep] + sub = sub[sub["predicted_class"].astype(str).isin(set(sc.index[sc >= map_thr]))] + parts.append(_fmt(sub, str(mc), "complex", "predicted_class", g2c, n_cells, score=sc)) + print(f"[phenotype] fluor filtered @ mAP>={map_thr}") + + df = pd.concat(parts, ignore_index=True) + df.to_csv(out_csv, index=False) + print(f"[phenotype] {len(df)} cells | {df['marker_channel'].nunique()} markers | " + f"{df.groupby(['marker_channel', 'grain'])['perturbation'].nunique().sum()} (marker,grain,pert) groups -> {out_csv}") + return out_csv + + +if __name__ == "__main__": + build() diff --git a/src/ops_model/models/attention/diffex/viewer/precompute.py b/src/ops_model/models/attention/diffex/viewer/precompute.py new file mode 100644 index 0000000..4e89a8f --- /dev/null +++ b/src/ops_model/models/attention/diffex/viewer/precompute.py @@ -0,0 +1,500 @@ +"""Precompute per-(marker, target, cell) traversal frames + manifest for the DiffEx viewer. + +A shareable MOPS-style static viewer can't run GPU diffusion live, so we precompute the +α-frame sequence of every traversal and let the frontend scrub it. w is FIXED (default 2.0, +the validated default); α is the scrub axis; marker / geneKO|complex / cell are routing. + +Each traversal emits raw decoded frames (no label overlay — the viewer draws its own α axis +and −KO/NTC/+KO cues) at: + /viewer_assets////cell/frame_.webp +plus a per-target meta.json. `build_manifest` aggregates all meta.json into one manifest.json +(marker -> grain -> targets -> cells + α list) that the static S3 viewer reads. + +The α seed noise xT is fixed per cell, so identity is anchored and only the phenotype shifts +across frames — smooth to scrub. α=0 (center) is the true NTC; +α = toward the KO phenotype, +−α = pushed to the opposite extreme. +""" +from __future__ import annotations + +import json +import os +from pathlib import Path + +# Assets subdir under out_root. Default = the live "viewer_assets"; the isolated v5 build sets +# OPS_DIFFEX_ASSETS=viewer_assets_v5 so nothing is written into the live tree until the final swap. +_ASSETS = os.environ.get("OPS_DIFFEX_ASSETS", "viewer_assets") + +import numpy as np +import pandas as pd +import torch +from PIL import Image + +from concurrent.futures import ThreadPoolExecutor + +from ..classifier.celldino_features import embed_crops +from ..classifier.config import GRAINS, slugify +from ..classifier.data import _BASE_COLS, make_labels_df, materialize_crops +from ..diffae.data import normalize +from ..directions.config import DirConfig +from ..directions.data import _top_cells +from ..directions.make_gifs import _pair_slug, _setup, _sample_guided +from ..directions.rank import supervised_direction +from ..directions.traverse import _ddim_guided, load_diffae +from ..directions.rank import supervised_direction + +# scrub axis; α in units of the control→KD gap (α=±1 ≈ full traversal). Dense in ±3 where the +# phenotype resolves, reaching ±5 for the subtle markers where extreme α still adds signal. +VIEWER_ALPHAS = (-5.0, -4.0, -3.0, -2.5, -2.0, -1.5, -1.0, -0.5, 0.0, + 0.5, 1.0, 1.5, 2.0, 2.5, 3.0, 4.0, 5.0) + + +def _save_webp(path, arr, upsize): + im = Image.fromarray((np.clip((arr + 1) / 2, 0, 1) * 255).astype("uint8")) + if upsize: + im = im.resize((upsize, upsize), Image.BILINEAR) + im.save(path, quality=90, method=6) + + +def _gather_class(cfg, value, n, parquet=None): + """Materialize + CellDINO-embed the top-n cells of ONE class → (images, embs). `parquet` overrides + the grain's attention parquet (phase only) — e.g. an accuracy-ranked table for the top-accuracy variant.""" + cc = GRAINS[cfg.grain]["class_col"] + if getattr(cfg, "marker_channel", None): + pre = getattr(cfg, "_fluor_rows", None) # per-marker preloaded rows (read the 12GB CSV ONCE) + if pre is None: + cols = list(dict.fromkeys([cc, *_BASE_COLS, "channel", "rank_type"])) + pre = pd.read_csv(cfg.fluor_csv, usecols=cols) + pre = pre[(pre["channel"] == cfg.marker_channel) & (pre["rank_type"] == "top")] + rows = pre[pre[cc].astype(str) == str(value)].sort_values("rank").head(n) + rows = rows.rename(columns={cc: "cls"}).copy() + else: + rows = _top_cells(parquet or GRAINS[cfg.grain]["parquet"], cc, value, n) # returns 'cls' + base cols + if rows.empty: + H = cfg.crop_size + return np.empty((0, 1, H, H), np.float32), np.empty((0, 1024), np.float32) + rows = rows.copy(); rows["label"] = 0 + ldf = make_labels_df(rows, cfg) + imgs, _, _ = materialize_crops(ldf, cfg, cache_path=None) + return imgs, embed_crops(imgs, cfg, cache_path=None) + + +@torch.no_grad() +def precompute_marker(grain, targets, ckpt, out_root, marker_channel=None, channel="Phase2D", + fluor_csv=None, control="NTC", n_cells=20, w=1.5, alphas=VIEWER_ALPHAS, + device="cuda", upsize=256, score=True, batch=48, n_workers=8, + load_workers=12, n_per_class=1000, + accuracy_parquet=None, variant=None, accuracy_fluor_csv=None, force=False, + v5_score=False, v5_bag=None, fluor_rank_parquet=None, invert_anchors=True, + save_gemb=False, skip_webp=False, webp_compare=False, cell_range=None, ddim_steps=100): + """Per-marker driver: gather the shared control/anchor cells ONCE and reuse across every + `target` (all a marker's geneKOs/complexes share the same NTC/anchor base cells + seeds). + Saves the ~n_cells real cells once under /_anchors//. Amortizes the + ckpt load + control gather; each target only gathers its own KD cells for the direction.""" + dev = torch.device(device if torch.cuda.is_available() else "cpu") + cfg = DirConfig(grain=grain, target=targets[0], control=control, device=device) + cfg.num_workers = load_workers + if ddim_steps: + cfg.ddim_steps = ddim_steps # traversal default 100 (DiffAEConfig=50): inverted-anchor + # round-trip needs ~100 to resolve (saturates there); 50 under-resolves + if ckpt: cfg.diffae_ckpt = ckpt + if marker_channel: cfg.marker_channel = marker_channel + if channel: cfg.channel = channel + if fluor_csv: cfg.fluor_csv = fluor_csv + if marker_channel: # preload the marker's cell table ONCE (all targets share it) + cc = GRAINS[grain]["class_col"] + if fluor_rank_parquet: # v5 accuracy ranking, pre-built per-channel (cc + base cols + rank_type) + cfg._fluor_rows = pd.read_parquet(fluor_rank_parquet) + print(f"[fluor-v5] {len(cfg._fluor_rows)} accuracy rows for '{marker_channel}' ({cfg._fluor_rows[cc].nunique()} classes)") + elif accuracy_fluor_csv: # fluor ACCURACY ranking (ebi_class_channel): label_name→class, per-complex rank + _a = pd.read_csv(accuracy_fluor_csv, usecols=["label_name", "channel", "experiment", "well", "x_pheno", "y_pheno", "segmentation_id", "score", "rank"]) + _a = _a[_a["channel"] == marker_channel].rename(columns={"label_name": cc, "segmentation_id": "segmentation", "score": "pma_attention"}) + _a["rank_type"] = "top"; cfg._fluor_rows = _a + print(f"[fluor-acc] {len(cfg._fluor_rows)} accuracy rows for '{marker_channel}' ({cfg._fluor_rows[cc].nunique()} complexes)") + else: # read the 12GB fluor CSV ONCE for this marker + _cols = list(dict.fromkeys([cc, *_BASE_COLS, "channel", "rank_type"])) + _all = pd.read_csv(cfg.fluor_csv, usecols=_cols) + cfg._fluor_rows = _all[(_all["channel"] == marker_channel) & (_all["rank_type"] == "top")] + print(f"[fluor] preloaded {len(cfg._fluor_rows)} '{marker_channel}' top rows (1 CSV read for all targets)") + modality = (slugify(marker_channel) if marker_channel else "phase") + (f"_{variant}" if variant else "") + al = sorted(alphas); A = len(al); H = cfg.crop_size + anchor = "NTC" if (not control or str(control).upper() == "NTC") else slugify(control) + realdir = Path(out_root) / _ASSETS / modality / "_anchors" / anchor + + # --- shared control/anchor cells: gather ONCE, cache the 1000-cell embeddings so every + # rebuild (new α/cells/anchor/ckpt) skips the gather (CellDINO embeds are ckpt-independent) --- + acache = realdir / "ctrl.npz" + if acache.exists(): + # PRE-SELECTED anchors (e.g. v5's accuracy cells): load embeddings; NEVER re-gather — the selected + # cells (and their count/order) must not change. Anchor images for inversion come from real.webp below. + z = np.load(acache); ctrl_embs = z["ctrl_embs"]; mu_ctrl = z["mu_ctrl"] + anchor_imgs = z["anchor_imgs"] if "anchor_imgs" in z.files else None + print(f"[cache] control embs <- {acache} {ctrl_embs.shape} (anchors fixed)") + else: + # no anchors selected yet (new marker) → gather + select fresh, save real.webp + embeddings + ctrl_imgs, ctrl_embs = _gather_class(cfg, control, n_per_class) + mu_ctrl = ctrl_embs.mean(0) + real = normalize(ctrl_imgs[:min(n_cells, len(ctrl_embs))]) + anchor_imgs = real + rp = ThreadPoolExecutor(max_workers=n_workers) + for c in range(len(real)): + (realdir / f"cell{c}").mkdir(parents=True, exist_ok=True) + rp.submit(_save_webp, realdir / f"cell{c}" / "real.webp", real[c, 0], upsize) + rp.shutdown(wait=True) + np.savez(acache, ctrl_embs=ctrl_embs, mu_ctrl=mu_ctrl, anchor_imgs=real) + print(f"[cache] control embs -> {acache} {ctrl_embs.shape}") + ncell = min(n_cells, len(ctrl_embs)) + cells = list(range(ncell)) if cell_range is None else [c for c in range(cell_range[0], min(cell_range[1], ncell))] + z0 = torch.as_tensor(ctrl_embs[:ncell], dtype=torch.float32, device=dev) + + diffae = load_diffae(cfg, dev) + null_base = diffae.null_emb.detach()[None].to(dev) + # anchor identity: DDIM-INVERT each SELECTED anchor (LOSSLESS anchor_imgs) to its own xT so α=0 is the + # real cell. anchor_imgs come from the fresh gather (this run) or a pre-built ctrl.npz. We NEVER invert + # the lossy real.webp — if invert is on but no lossless anchor_imgs are present, fail loud. + if invert_anchors: + if anchor_imgs is None: + raise RuntimeError(f"invert_anchors=True but no lossless anchor_imgs in {acache}; pre-build the " + f"anchor cache first (lossy real.webp inversion is disabled).") + x0a = torch.as_tensor(anchor_imgs[:ncell], dtype=torch.float32, device=dev) + xT, _nc = {}, 0 # cache each cell's inverted xT (ckpt/w/steps-fixed) → invert ONCE, reuse across shards + for c in cells: + xtp = realdir / f"cell{c}" / f"xT_w{w:g}_s{ddim_steps}.npz" + if xtp.exists(): + xT[c] = torch.as_tensor(np.load(xtp)["xT"], dtype=torch.float32, device=dev); _nc += 1 + else: + xt = _ddim_guided(diffae, x0a[c:c + 1], z0[c:c + 1], null_base, w, cfg, inverse=True) + xtp.parent.mkdir(parents=True, exist_ok=True) + tmp = xtp.parent / f".xT_{os.getpid()}_{c}.npz" # atomic write → concurrent shards can't corrupt the cache + np.savez(tmp, xT=xt.detach().cpu().numpy()); os.replace(tmp, xtp) + xT[c] = xt + print(f"[invert] anchored xT for {len(cells)} cell(s) ({_nc} from cache, w={w})") + else: + xT = {c: torch.randn(1, 1, H, H, generator=torch.Generator(device=dev).manual_seed(1234 + c), device=dev) for c in cells} + + v5ctx = None # inline v5 SetTransformer scoring (reuses gemb; no re-decode) + if v5_score: + from .set_classifier import load_set_classifier, V5_CKPT_ROOT, V5_RUNS + from .score_generated import score_embs_v5 + _vmod = "fluor" if marker_channel else "phase" # fluor markers score with the fluor SetTransformer + their channel idx + _vrun = V5_RUNS[(_vmod, "geneKO" if grain == "geneKO" else "complex_ebionly")] + _vm, _vg2i, _vc2i = load_set_classifier(run=_vrun, device=dev, root=V5_CKPT_ROOT) + v5ctx = (_vm, _vg2i, _vc2i.get(marker_channel or "Phase2D", 0), _vrun, score_embs_v5) + + done = 0 + for tgt in targets: + slug = slugify(tgt) if anchor == "NTC" else f"{anchor}__to__{slugify(tgt)}" + adir = Path(out_root) / _ASSETS / modality / grain / slug + if not force and (adir / "meta.json").exists(): # resume: target already rendered (force = overwrite in place) + print(f"[done] {modality}/{grain}/{slug}"); done += 1; continue + # direction cache: (d_vec, gap) is CellDINO-derived + ckpt-independent → gather once, reuse + dcache = Path(out_root) / _ASSETS / "_directions" / modality / grain / f"{slug}.npz" + if dcache.exists() and (cell_range is not None or not force): # direction is ckpt-independent → chunks always reuse + z = np.load(dcache); d_vec = z["d_vec"]; gap = float(z["gap"]); lr_w = z["lr_w"]; lr_b = float(z["lr_b"]) + print(f"[cache] direction <- {dcache}") + else: + kd_imgs, kd_embs = _gather_class(cfg, tgt, n_per_class, parquet=accuracy_parquet) + if not len(kd_embs): + print(f"[skip] {tgt}: no cells"); continue + embs = np.concatenate([kd_embs, ctrl_embs], 0) + labels = np.concatenate([np.ones(len(kd_embs)), np.zeros(len(ctrl_embs))]).astype(int) + d_vec, lr_w, lr_b, _ = supervised_direction(embs, labels, cfg) + gap = float(np.linalg.norm(kd_embs.mean(0) - mu_ctrl)) + dcache.parent.mkdir(parents=True, exist_ok=True) + np.savez(dcache, d_vec=d_vec, gap=gap, lr_w=lr_w, lr_b=lr_b) + fixed_dir = torch.as_tensor(d_vec, dtype=torch.float32, device=dev)[None] + + conds, xts, keys = [], [], [] + for c in cells: + for ai, a in enumerate(al): + conds.append(z0[c:c + 1] + (a * gap) * fixed_dir); xts.append(xT[c]); keys.append((c, ai)) + gen = np.empty((ncell, A, H, H), np.float32) + for i0 in range(0, len(conds), batch): + cb = torch.cat(conds[i0:i0 + batch], 0); xb = torch.cat(xts[i0:i0 + batch], 0) + outb = _sample_guided(diffae, xb, cb, null_base.expand(cb.shape[0], -1), w, cfg).cpu().numpy()[:, 0] + for j, (c, ai) in enumerate(keys[i0:i0 + batch]): + gen[c, ai] = outb[j] + scores, v5d, gemb = None, None, None + if (score or v5ctx is not None or save_gemb) and cell_range is None: # gemb/scores are whole-target → skip for cell chunks + gemb = embed_crops(gen.reshape(-1, 1, H, H).astype(np.float32), cfg, cache_path=None) + if score and cell_range is None: + scores = (1.0 / (1.0 + np.exp(-(gemb @ lr_w + lr_b)))).reshape(ncell, A) + if v5ctx is not None and cell_range is None: + try: + adir.mkdir(parents=True, exist_ok=True) # scores_v5.json is written before the frame loop creates adir + _vm, _vg2i, _vci, _vrun, _score_embs = v5ctx + g = gemb.reshape(ncell, A, -1) + v5d = _score_embs([g[:, ai, :] for ai in range(A)], al, tgt, _vm, _vg2i, _vci, _vrun, dev, v5_bag) + if v5d is not None: + (adir / "scores_v5.json").write_text(json.dumps(v5d)) + except Exception as e: + print(f"[v5score ERR] {tgt}: {repr(e)[:120]}") + + if save_gemb and cell_range is None: # persist the in-memory float CellDINO embeddings (no webp round-trip) + adir.mkdir(parents=True, exist_ok=True) + out = {"gemb": gemb.reshape(ncell, A, -1).astype(np.float32), + "alphas": np.asarray(al, np.float32), "target": tgt, "ncell": ncell} + if webp_compare: # ALSO embed the SAME frames after an 8-bit webp round-trip + import io # only the α we compare (0 and 3) → ~8× faster than all 17 + cmp_ais = [ai for ai in (8, 14) if ai < A] + def _rt(arr): + u8 = (np.clip((arr + 1) / 2, 0, 1) * 255).astype(np.uint8) + b = io.BytesIO(); Image.fromarray(u8).resize((upsize, upsize), Image.BILINEAR).save(b, "webp", quality=90) + b.seek(0); return np.asarray(Image.open(b).convert("L"), np.float32) / 255.0 * 2 - 1 + gw = np.stack([_rt(gen[c, ai]) for c in range(ncell) for ai in cmp_ais])[:, None].astype(np.float32) + gwe = embed_crops(gw, cfg, cache_path=None).reshape(ncell, len(cmp_ais), -1) + full = np.zeros((ncell, A, gwe.shape[-1]), np.float32) + for j, ai in enumerate(cmp_ais): + full[:, ai, :] = gwe[:, j, :] + out["gemb_webp"] = full.astype(np.float32); out["gemb_webp_ais"] = np.asarray(cmp_ais) + np.savez(adir / "gemb.npz", **out) + if not skip_webp: + fp = ThreadPoolExecutor(max_workers=n_workers) + for c in cells: + cdir = adir / f"cell{c}"; cdir.mkdir(parents=True, exist_ok=True) + for ai in range(A): + fp.submit(_save_webp, cdir / f"frame_{ai:02d}.webp", gen[c, ai], upsize) + fp.shutdown(wait=True) + if scores is not None: + (adir / "scores.json").write_text(json.dumps({"alphas": al, "scores": np.round(scores, 3).tolist()})) + meta = {"grain": grain, "target": tgt, "modality": modality, + "control": None if anchor == "NTC" else control, + "marker_channel": marker_channel, "channel": channel, "slug": slug, "w": w, + "ddim_steps": cfg.ddim_steps, # stamp step count so resume can tell 50- vs 100-step frames apart + "alphas": al, "gap": gap, "n_cells": ncell, "has_scores": scores is not None, + "has_scores_v5": v5d is not None, + "has_real": True, "real_dir": f"{modality}/_anchors/{anchor}", + "asset_dir": f"{modality}/{grain}/{slug}"} + (adir / "meta.json").write_text(json.dumps(meta)); done += 1 + print(f"[viewer] {modality}/{grain}/{slug}: {ncell}×{A}" + (" +scores" if score else "")) + print(f"[marker] {modality}/{grain} ({anchor}-anchored): {done}/{len(targets)} targets, real cells shared") + return {"modality": modality, "grain": grain, "anchor": anchor, "n_targets": done} + + +@torch.no_grad() +def precompute_anchors_marker(grain, classes, ckpt, out_root, marker_channel=None, channel="Phase2D", + fluor_csv=None, n_cells=20, w=1.5, alphas=VIEWER_ALPHAS, device="cuda", + upsize=256, batch=48, n_workers=8, load_workers=12, n_per_class=1000, + pairs=None, accuracy_parquet=None, force=False, + invert_anchors=True, v5_score=False, v5_bag=None): + """A→B anchors among `classes` for ONE marker: gather each class's cells ONCE (single CSV read), + then generate every ordered pair a→b (anchor a's cells morphed toward b's centroid). The gather is + amortized across all K·(K−1) pairs; the mean-diff direction is b_centroid − a_centroid.""" + dev = torch.device(device if torch.cuda.is_available() else "cpu") + cfg = DirConfig(grain=grain, target=classes[0], device=device) + cfg.num_workers = load_workers + if ckpt: cfg.diffae_ckpt = ckpt + if marker_channel: cfg.marker_channel = marker_channel + if channel: cfg.channel = channel + if fluor_csv: cfg.fluor_csv = fluor_csv + modality = slugify(marker_channel) if marker_channel else "phase" + if marker_channel: # preload the marker's cell table ONCE + cc = GRAINS[grain]["class_col"] + _cols = list(dict.fromkeys([cc, *_BASE_COLS, "channel", "rank_type"])) + _all = pd.read_csv(cfg.fluor_csv, usecols=_cols) + cfg._fluor_rows = _all[(_all["channel"] == marker_channel) & (_all["rank_type"] == "top")] + cache = {} # gather each class ONCE (in-memory filter, no re-read) + for cls in classes: + imgs, embs = _gather_class(cfg, cls, n_per_class, parquet=accuracy_parquet) # accuracy_parquet → top-n by v5 accuracy + if len(embs): + cache[cls] = (imgs, embs) + al = sorted(alphas); A = len(al); H = cfg.crop_size + diffae = load_diffae(cfg, dev); null_base = diffae.null_emb.detach()[None].to(dev) + v5ctx = None # inline v5 SetTransformer scoring — same as precompute_marker + if v5_score: + from .set_classifier import load_set_classifier, V5_CKPT_ROOT, V5_RUNS + from .score_generated import score_embs_v5 + _vmod = "fluor" if marker_channel else "phase" + _vrun = V5_RUNS[(_vmod, "geneKO" if grain == "geneKO" else "complex_ebionly")] + _vm, _vg2i, _vc2i = load_set_classifier(run=_vrun, device=dev, root=V5_CKPT_ROOT) + v5ctx = (_vm, _vg2i, _vc2i.get(marker_channel or "Phase2D", 0), _vrun, score_embs_v5) + done = 0 + ordered = list(pairs) if pairs is not None else [(a, b) for a in classes for b in classes if a != b] + setup = {} + def _anchor_setup(a): # per-anchor z0/xT/mu + real cells, built once (force → rebuild reals) + if a in setup: + return setup[a] + a_imgs, a_embs = cache[a]; ncell = min(n_cells, len(a_embs)) + z0 = torch.as_tensor(a_embs[:ncell], dtype=torch.float32, device=dev) + if invert_anchors: # DDIM-invert the LOSSLESS source cells → faithful α=0 + x0a = torch.as_tensor(normalize(a_imgs[:ncell]), dtype=torch.float32, device=dev) + xT = torch.cat([_ddim_guided(diffae, x0a[c:c + 1], z0[c:c + 1], null_base, w, cfg, inverse=True) + for c in range(ncell)], 0) + else: + xT = torch.stack([torch.randn(1, H, H, generator=torch.Generator(device=dev).manual_seed(1234 + c), device=dev) + for c in range(ncell)]) + mu_a = a_embs.mean(0); anchor = slugify(a) + realdir = Path(out_root) / _ASSETS / modality / "_anchors" / anchor + if force or not (realdir / "cell0" / "real.webp").exists(): + real = normalize(a_imgs[:ncell]); rp = ThreadPoolExecutor(max_workers=n_workers) + for c in range(ncell): + (realdir / f"cell{c}").mkdir(parents=True, exist_ok=True) + rp.submit(_save_webp, realdir / f"cell{c}" / "real.webp", real[c, 0], upsize) + rp.shutdown(wait=True) + setup[a] = (z0, xT, mu_a, ncell, anchor) + return setup[a] + for a, b in ordered: + if a not in cache or b not in cache or a == b: + continue + z0, xT, mu_a, ncell, anchor = _anchor_setup(a) + slug = f"{anchor}__to__{slugify(b)}" + adir = Path(out_root) / _ASSETS / modality / grain / slug + if (adir / "meta.json").exists() and not force: + done += 1; continue + if True: + d = cache[b][1].mean(0) - mu_a; gap = float(np.linalg.norm(d)) + fixed_dir = torch.as_tensor(d / (gap + 1e-9), dtype=torch.float32, device=dev)[None] + conds, xts, keys = [], [], [] + for c in range(ncell): + for ai, av in enumerate(al): + conds.append(z0[c:c + 1] + (av * gap) * fixed_dir); xts.append(xT[c:c + 1]); keys.append((c, ai)) + gen = np.empty((ncell, A, H, H), np.float32) + for i0 in range(0, len(conds), batch): + cb = torch.cat(conds[i0:i0 + batch], 0); xb = torch.cat(xts[i0:i0 + batch], 0) + outb = _sample_guided(diffae, xb, cb, null_base.expand(cb.shape[0], -1), w, cfg).cpu().numpy()[:, 0] + for j, (c, ai) in enumerate(keys[i0:i0 + batch]): + gen[c, ai] = outb[j] + fp = ThreadPoolExecutor(max_workers=n_workers) + for c in range(ncell): + cdir = adir / f"cell{c}"; cdir.mkdir(parents=True, exist_ok=True) + for ai in range(A): + fp.submit(_save_webp, cdir / f"frame_{ai:02d}.webp", gen[c, ai], upsize) + fp.shutdown(wait=True) + v5d = None + if v5ctx is not None: # inline P(target b) + rank, reusing the generated frames (no re-decode) + try: + gemb = embed_crops(gen.reshape(-1, 1, H, H).astype(np.float32), cfg, cache_path=None).reshape(ncell, A, -1) + _vm, _vg2i, _vci, _vrun, _score_embs = v5ctx + v5d = _score_embs([gemb[:, ai, :] for ai in range(A)], al, b, _vm, _vg2i, _vci, _vrun, dev, v5_bag) + if v5d is not None: + (adir / "scores_v5.json").write_text(json.dumps(v5d)) + except Exception as e: + print(f"[v5score ERR] {b}: {repr(e)[:120]}") + (adir / "meta.json").write_text(json.dumps( + {"grain": grain, "target": b, "modality": modality, "control": a, "marker_channel": marker_channel, + "channel": channel, "slug": slug, "w": w, "alphas": al, "gap": gap, "n_cells": ncell, + "has_scores": False, "has_scores_v5": v5d is not None, "has_real": True, "real_dir": f"{modality}/_anchors/{anchor}", + "asset_dir": f"{modality}/{grain}/{slug}"})) + done += 1 + print(f"[anchor] {modality}/{grain}/{slug}: {ncell}×{A}") + print(f"[anchors] {modality}/{grain}: {done} pairs across {len(cache)} classes") + return {"modality": modality, "grain": grain, "n_pairs": done} + + +@torch.no_grad() +def precompute_target(grain, target, ckpt, out_root, marker_channel=None, channel=None, + fluor_csv=None, control=None, n_cells=20, w=1.5, alphas=VIEWER_ALPHAS, + device="cuda", upsize=256, score=True, batch=48, n_workers=8, + load_workers=10, keep_crops=False, accuracy_parquet=None, variant=None, + accuracy_fluor_csv=None, invert_anchors=True): + """Decode + save the α-frame sequence for the first n_cells control cells of one + (marker, target). Batched GPU decode, batched re-encode → per-image classifier + confidence (sigmoid of the control→KD LR logit), threaded WebP save. Writes frames + + meta.json + scores.json. control=None → NTC-anchored; else an A→B anchor class. + load_workers: parallel zarr crop-read workers (the gather dominates runtime). + accuracy_parquet/variant: accuracy-selected cells (both A & B) into a variant modality (phase_topacc).""" + ctx = _setup(grain, target, out_root, device, ckpt=ckpt, marker_channel=marker_channel, + channel=channel, fluor_csv=fluor_csv, control=control, num_workers=load_workers, + return_images=True, accuracy_parquet=accuracy_parquet, variant=variant, + accuracy_fluor_csv=accuracy_fluor_csv) + dev, cfg, slug, _out, embs, labels, fixed_dir, gap, diffae, null_base, real_imgs = ctx + ci = np.flatnonzero(labels == 0) + ncell = min(n_cells, len(ci)) + H, al, A = cfg.crop_size, sorted(alphas), len(sorted(alphas)) + modality = (slugify(marker_channel) if marker_channel else "phase") + (f"_{variant}" if variant else "") + adir = Path(out_root) / _ASSETS / modality / grain / slug + + # per-cell latent: DDIM-invert the source cell (α=0 = the real cell) or a fixed random seed + real_norm = normalize(real_imgs[ci[:ncell]]) # source-cell crops, [-1,1] — inversion + display + conds, xts, keys = [], [], [] + for cell in range(ncell): + z0 = torch.as_tensor(embs[ci[cell]:ci[cell] + 1], dtype=torch.float32, device=dev) + if invert_anchors: + x0c = torch.as_tensor(real_norm[cell:cell + 1], dtype=torch.float32, device=dev) + xT = _ddim_guided(diffae, x0c, z0, null_base, w, cfg, inverse=True) + else: + xT = torch.randn(1, 1, H, H, generator=torch.Generator(device=dev).manual_seed(1234 + cell), device=dev) + for ai, a in enumerate(al): + conds.append(z0 + (a * gap) * fixed_dir); xts.append(xT); keys.append((cell, ai)) + + gen = np.empty((ncell, A, H, H), dtype=np.float32) + for i0 in range(0, len(conds), batch): + cb = torch.cat(conds[i0:i0 + batch], 0) + xb = torch.cat(xts[i0:i0 + batch], 0) + nb = null_base.expand(cb.shape[0], -1) + out = _sample_guided(diffae, xb, cb, nb, w, cfg).cpu().numpy()[:, 0] + for j, (cell, ai) in enumerate(keys[i0:i0 + batch]): + gen[cell, ai] = out[j] + + scores = None + if score: # re-encode every frame → LR class confidence + _, lr_w, lr_b, _ = supervised_direction(embs, labels, cfg) + gemb = embed_crops(gen.reshape(-1, 1, H, H).astype(np.float32), cfg, cache_path=None) + logits = gemb @ lr_w + lr_b + scores = (1.0 / (1.0 + np.exp(-logits))).reshape(ncell, A) + + real = real_norm # actual source-cell crops, [-1,1] for display + pool = ThreadPoolExecutor(max_workers=n_workers) + for cell in range(ncell): + cdir = adir / f"cell{cell}"; cdir.mkdir(parents=True, exist_ok=True) + pool.submit(_save_webp, cdir / "real.webp", real[cell, 0], upsize) # static real cell (α=0 is a recon) + for ai in range(A): + pool.submit(_save_webp, cdir / f"frame_{ai:02d}.webp", gen[cell, ai], upsize) + pool.shutdown(wait=True) + + if scores is not None: + (adir / "scores.json").write_text(json.dumps({"alphas": al, "scores": np.round(scores, 3).tolist()})) + meta = {"grain": grain, "target": target, "modality": modality, "control": control, + "marker_channel": marker_channel, "channel": channel, "slug": slug, + "w": w, "alphas": al, "gap": float(gap), "n_cells": ncell, "has_scores": scores is not None, + "has_real": True, "asset_dir": f"{modality}/{grain}/{slug}"} + (adir / "meta.json").write_text(json.dumps(meta)) + + if not keep_crops: # drop the ~195MB materialized-crop cache the viewer never needs + cp = Path(_out) / "cache" / f"crops_{slug}_{cfg.crop_size}.npz" + if cp.exists(): + cp.unlink() + print(f"[viewer] {modality}/{grain}/{slug}: {ncell}×{A} frames" + (" +scores" if score else "")) + return meta + + +def build_manifest(out_root, dist_map=None, desc_map=None): + """Aggregate every viewer_assets/*/*/*/meta.json into one manifest.json the frontend reads. + dist_map: optional {(modality, grain, slug): mAP} to attach for sorting targets. + desc_map: optional {target_name: description} (gene function / complex members).""" + root = Path(out_root) / _ASSETS + try: + mb = json.loads((root / "_minibinder_meta.json").read_text()) # per-binder cell_score/binder_prob/gene_target + except Exception: + mb = {} + markers = {} + for mj in sorted(root.glob("*/*/*/meta.json")): + m = json.loads(mj.read_text()) + mod = m["modality"] + if mod == "phase_minibinder": # orphan from a cancelled mis-structured run (couldn't rm on shared FS); minibinders live under phase/minibinder + continue + label = m["marker_channel"] or "Phase" # phase = canonical (accuracy-selected as of 2026-07-12 swap) + mk = markers.setdefault(mod, {"modality": mod, "marker_channel": m["marker_channel"], + "label": label, "channel": m["channel"], "targets": []}) + key = (mod, m["grain"], m["slug"]) + adir = m["asset_dir"] + if adir.startswith("viewer_assets/"): # normalize pre-fix-era meta.json + adir = adir[len("viewer_assets/"):] + mk["targets"].append({"grain": m["grain"], "target": m["target"], "slug": m["slug"], + "control": m.get("control"), # None = NTC-anchored; else A→B anchor class + "has_real": m.get("has_real", False), "real_dir": m.get("real_dir"), + "n_cells": m["n_cells"], "asset_dir": adir, "alphas": m["alphas"], + "dist_map": (dist_map or {}).get(key), + "explained_variance": m.get("explained_variance"), # PC grain: % variance + "desc": (desc_map or {}).get(m["target"]), + **({"binder_prob": mb[m["slug"]]["binder_prob"], "gene_target": mb[m["slug"]]["gene_target"], + "phenotype": mb[m["slug"]]["phenotype"], "cell_score": mb[m["slug"]]["cell_score"]} + if m["grain"] == "minibinder" and m["slug"] in mb else {})}) + for mk in markers.values(): + mk["targets"].sort(key=lambda t: (-(t["dist_map"] or -1), t["target"])) + manifest = {"alphas": list(VIEWER_ALPHAS), "w": 2.0, + "markers": sorted(markers.values(), key=lambda x: (x["marker_channel"] or "", x["modality"]))} + out = root / "manifest.json" + out.write_text(json.dumps(manifest, indent=2)) + n_t = sum(len(mk["targets"]) for mk in markers.values()) + print(f"[viewer] manifest: {len(markers)} markers, {n_t} targets -> {out}") + return str(out) diff --git a/src/ops_model/models/attention/diffex/viewer/render_montage_scales.py b/src/ops_model/models/attention/diffex/viewer/render_montage_scales.py new file mode 100644 index 0000000..7775dea --- /dev/null +++ b/src/ops_model/models/attention/diffex/viewer/render_montage_scales.py @@ -0,0 +1,364 @@ +"""Stitch a DeepZoom montage level into a single composite PNG (white bg) with a bottom-left +embedding legend (leiden_r4, big dots, NTC as a dark labelled circle). The montage image and its +baked viewer-style gene names come straight from the built tiles — finer levels give crisper text. + + python -m ops_model.models.attention.diffex.viewer.render_montage_scales --alphas 1-5 --levels 3,4 + +Each level of `_montage/phase_geneKO_phate_cell1_a_tiles/L/` is a level-of-detail montage +(coarse levels show a decimated non-overlapping subset; finer levels fill in more cells at higher res). +""" +from __future__ import annotations + +import argparse +import json +import os + +import numpy as np +import matplotlib +matplotlib.use("Agg") +import matplotlib.pyplot as plt +import matplotlib.patheffects as pe +from matplotlib.patches import Circle, Rectangle +from PIL import Image + +plt.rcParams["pdf.fonttype"] = 42 +Image.MAX_IMAGE_PIXELS = None + +VA = "/hpc/projects/icd.fast.ops/models/diffex/viewer_assets" +OUT_DIR = "/hpc/projects/icd.fast.ops/analysis/figure4_embedding/montage_scales" +HIRES_ROOT = "/hpc/projects/icd.fast.ops/analysis/figure4_embedding/hires_tiles" +COMPOSED_DIR = "/hpc/projects/icd.fast.ops/analysis/figure4_embedding/montage_composed" +UMAP_H5AD = ("/hpc/projects/icd.fast.ops/organelle_attribution/pca_optimized_v0.3/cell_dino/" + "zscore_per_exp/paper_v2/phase_only/fixed_80%/cosine/gene_embedding_pca_optimized.h5ad") +LEIDEN = "leiden_r4" +VIEWER_ALPHAS = [-5.0, -4.0, -3.0, -2.5, -2.0, -1.5, -1.0, -0.5, 0.0, 0.5, 1.0, 1.5, 2.0, 2.5, 3.0, 4.0, 5.0] +# EBI complexes to box in the montage: label -> (complex name in labels.json, color) +MARKS = {"40S": ("40S cytosolic small ribosomal subunit", "#0b6b73"), # dark teal / dark orange for contrast + "60S": ("60S cytosolic large ribosomal subunit", "#b35900")} +CPLX_TRAVERSAL = "40S_cytosolic_small_ribosomal_subunit__to__60S_cytosolic_large_ribosomal_subunit" +MARK_GENE = {"40S": "RPS25", "60S": "RPL37A"} # specific member to box (verified in the complex) + + +def _frame(alpha): + return max(0, min(16, int(round(8 + alpha / 5.0 * 8)))) # 17 frames span alpha -5..+5; 8 = base + + +def _stretch(im, lo_p=3.0, hi_p=97.0): + """Per-crop percentile contrast stretch (latent_lens _normalize_crop uses 1/99; tighter = punchier). + Input/return uint8.""" + a = im.astype(np.float32) + lo, hi = np.percentile(a, lo_p), np.percentile(a, hi_p) + if hi <= lo: + return im + return (np.clip((a - lo) / (hi - lo), 0, 1) * 255).astype(np.uint8) + + +def _composite(tile_dir, level, width, height, ts, bg=255): + lw = max(1, width >> level); lh = max(1, height >> level) + canv = np.full((lh, lw), bg, np.uint8) + d = f"{tile_dir}/L{level}" + n = 0; x0 = y0 = 1 << 30; x1 = y1 = 0 + for f in os.listdir(d): + if not f.endswith(".png"): + continue + col, row = (int(x) for x in f[:-4].split("_")) + t = np.asarray(Image.open(f"{d}/{f}").convert("L")) + py, px = row * ts, col * ts; th, tw = t.shape + canv[py:py + th, px:px + tw] = t[: lh - py, : lw - px] + n += 1 + x0 = min(x0, px); y0 = min(y0, py); x1 = max(x1, px + tw); y1 = max(y1, py + th) + m = ts // 4 # canvas white; tiles keep their black interior (bright points render on black) + bbox = (max(0, x0 - m), max(0, y0 - m), min(lw, x1 + m), min(lh, y1 + m)) + return canv, n, bbox + + +def _leiden_colors(labels): + clusters = sorted({g.get(LEIDEN, "") for g in labels if g.get(LEIDEN, "") != ""}, + key=lambda s: (len(s), s)) + cmap = plt.get_cmap("hsv") # all bright/saturated, no dark clusters (dark is reserved for NTC) + return {c: cmap(i / max(1, len(clusters) - 1)) for i, c in enumerate(clusters)} + + +def _legend(ax, labels, cw, dpi, box=(0.008, 0.008, 0.26, 0.26), dark=False): + lut = _leiden_colors(labels) + ex = np.array([g["nx"] for g in labels]); ey = np.array([g["ny"] for g in labels]) + ec = np.array([lut.get(g.get(LEIDEN, ""), (0.7, 0.7, 0.7, 1)) for g in labels]) + is_ntc = np.array([str(g["g"]).startswith("NTC") for g in labels]) + tcol = "white" if dark else "#111" + iax = ax.inset_axes(list(box)) + if dark: + iax.set_facecolor("white") + dot = (cw / dpi) * 6 + iax.scatter(ex[~is_ntc], ey[~is_ntc], c=ec[~is_ntc], s=dot, edgecolors="none", zorder=2) + if is_ntc.any(): + nx, ny = ex[is_ntc], ey[is_ntc] + iax.scatter(nx, ny, c="#111", s=dot * 1.2, edgecolors="none", zorder=4) + cx, cy = nx.mean(), ny.mean() + rr = 1.7 * np.percentile(np.hypot(nx - cx, ny - cy), 90) + 0.01 + iax.add_patch(Circle((cx, cy), rr, fill=False, ls="--", ec="#111", lw=max(1.2, cw / dpi * 0.15), zorder=5)) + iax.annotate("NTC", (cx, cy - rr * 1.5), ha="center", va="bottom", + fontsize=max(10, cw / dpi * 0.95), weight="bold", color="#111") + iax.set_aspect("equal"); iax.invert_yaxis(); iax.set_xticks([]); iax.set_yticks([]) + for s in iax.spines.values(): + s.set_visible(False) + iax.set_title(f"embedding · {LEIDEN}", fontsize=max(11, cw / dpi * 1.1), weight="bold", color=tcol) + + +def _mark(ax, tile_dir, labels, W, lvl, ts, x0c, y0c, cw, dpi): + """Box the montage tile nearest each marked complex's embedding centroid + label it.""" + lw = W >> lvl + occ = [tuple(int(v) for v in f[:-4].split("_")) + for f in os.listdir(f"{tile_dir}/L{lvl}") if f.endswith(".png")] + for lab, (cname, color) in MARKS.items(): + pts = [(g["nx"] * lw, g["ny"] * lw) for g in labels if g.get("ebi_complex") == cname] + if not pts: + continue + cx = np.mean([p[0] for p in pts]); cy = np.mean([p[1] for p in pts]) + col, row = min(occ, key=lambda t: ((t[0] + 0.5) * ts - cx) ** 2 + ((t[1] + 0.5) * ts - cy) ** 2) + rx, ry = col * ts - x0c, row * ts - y0c + ax.add_patch(Rectangle((rx, ry), ts, ts, fill=False, ec=color, lw=max(3, cw / dpi * 0.5), zorder=6)) + ax.text(rx + ts / 2, ry - 6, lab, ha="center", va="bottom", color=color, + fontsize=max(13, cw / dpi * 1.3), weight="bold", zorder=7) + + +def build_hires(alphas, cell=1, crop=512, embedding="phate", border_field=LEIDEN): + """Rebuild the phase geneKO montage at a larger crop_size (crisper baked text, 64px font at crop 512) + with per-cell leiden borders, to a SEPARATE dir (viewer cache untouched).""" + from .build_umap_montage import montage_from_cache, montage_to_tiles, ZARR_SCRATCH + import shutil + os.makedirs(HIRES_ROOT, exist_ok=True) + for a in alphas: + stem = f"phase_geneKO_{embedding}_cell{cell}_a{a:g}" + scratch = f"{ZARR_SCRATCH}/{stem}_hires.zarr" + _, placed = montage_from_cache(UMAP_H5AD, scratch, cell=cell, alpha=a, modality="phase", grain="geneKO", + tile=crop, px_per_umap=int(round(crop * 5600 / 256)), + embedding=embedding, border_field=border_field) + montage_to_tiles(scratch, UMAP_H5AD, out_dir=f"{HIRES_ROOT}/{stem}_tiles", + placed=set(placed), embedding=embedding) + shutil.rmtree(scratch, ignore_errors=True) + print(f"[hires] built {stem}") + + +def render(alpha, levels, out_dir=None, tiles_root=None): + out_dir = out_dir or OUT_DIR + os.makedirs(out_dir, exist_ok=True) + tile_dir = f"{tiles_root or (VA + '/_montage')}/phase_geneKO_phate_cell1_a{alpha}_tiles" + meta = json.load(open(f"{tile_dir}/tiles.json")) + W, H, ts = meta["width"], meta["height"], meta["tileSize"] + labels = json.load(open(f"{tile_dir}/labels.json")) + for lvl in levels: + if not os.path.isdir(f"{tile_dir}/L{lvl}"): + print(f"[montage] a{alpha} L{lvl}: missing, skip"); continue + canv, n, (x0, y0, x1, y1) = _composite(tile_dir, lvl, W, H, ts) + sub = canv[y0:y1, x0:x1] + ch, cw = sub.shape + dpi = 150 + fig = plt.figure(figsize=(cw / dpi, ch / dpi), dpi=dpi, facecolor="white") + ax = fig.add_axes([0, 0, 1, 1]); ax.imshow(sub, cmap="gray", vmin=0, vmax=255, aspect="auto"); ax.axis("off") + _mark(ax, tile_dir, labels, W, lvl, ts, x0, y0, cw, dpi) + _legend(ax, labels, cw, dpi) + out = f"{out_dir}/montage_a{alpha}_L{lvl}_{cw}x{ch}.png" + fig.savefig(out, dpi=dpi, facecolor="white") + plt.close(fig) + print(f"[montage] a{alpha} L{lvl}: {cw}x{ch} ({n} tiles) -> {out}") + + +def render_traversal(cells=(0, 1, 2, 3, 4, 5), alphas=(-3, 0, 1, 3), out_dir=None): + """One PNG per cell of the 40S->60S complex-level DiffAE traversal (tight, no inter-image gaps).""" + out_dir = out_dir or OUT_DIR + os.makedirs(out_dir, exist_ok=True) + fis = [int(np.argmin([abs(v - a) for v in VIEWER_ALPHAS])) for a in alphas] + for cell in cells: + d = f"{VA}/phase/complex/{CPLX_TRAVERSAL}/cell{cell}" + if not os.path.isdir(d): + continue + n = len(alphas) + fig, axes = plt.subplots(1, n, figsize=(3 * n, 3.35), gridspec_kw={"wspace": 0.03}) + for ax, a, fi in zip(np.atleast_1d(axes), alphas, fis): + p = f"{d}/frame_{fi:02d}.webp" + if os.path.exists(p): + ax.imshow(np.asarray(Image.open(p).convert("L")), cmap="gray", vmin=0, vmax=255) + ax.set_title(f"α = {a:g}", fontsize=13, weight="bold") + ax.set_xticks([]); ax.set_yticks([]) + for s in ax.spines.values(): + s.set_visible(False) + axes[0].set_ylabel("40S", fontsize=15, weight="bold", color=MARKS["40S"][1]) + axes[-1].yaxis.set_label_position("right") + axes[-1].set_ylabel("60S", fontsize=15, weight="bold", color=MARKS["60S"][1], rotation=270, labelpad=18) + fig.suptitle(f"40S → 60S complex traversal · cell {cell}", fontsize=15, weight="bold", y=1.04) + out = f"{out_dir}/traversal_40S_to_60S_cell{cell}.png" + fig.savefig(out, dpi=200, bbox_inches="tight", facecolor="white") + plt.close(fig) + print(f"[traversal] 40S->60S cells {list(cells)} alphas {list(alphas)} -> {out_dir}/traversal_40S_to_60S_cell*.png") + + +def render_ntc_pair(cell=1, ref_gene="FANCC", out_dir=None): + """Separate PNG: the original NTC cell (real) next to its generative reconstruction (α=0, frame_08).""" + out_dir = out_dir or OUT_DIR + os.makedirs(out_dir, exist_ok=True) + real = f"{VA}/phase/_anchors/NTC/cell{cell}/real.webp" + gen = f"{VA}/phase/geneKO/{ref_gene}/cell{cell}/frame_08.webp" # α=0 = unmorphed base reconstruction + fig, axes = plt.subplots(1, 2, figsize=(6.4, 3.6), gridspec_kw={"wspace": 0.04}) + for ax, p, t in [(axes[0], real, "original"), (axes[1], gen, "generative")]: + if os.path.exists(p): + im = np.asarray(Image.open(p).convert("L")) + if t == "original": + im = im[::-1, ::-1] # flip H+V so the original matches the generative orientation + ax.imshow(im, cmap="gray", vmin=0, vmax=255) + ax.set_title(t, fontsize=14, weight="bold") + ax.set_xticks([]); ax.set_yticks([]) + for s in ax.spines.values(): + s.set_visible(False) + fig.suptitle("NTC cell", fontsize=15, weight="bold", y=1.03) + out = f"{out_dir}/ntc_original_vs_generative_cell{cell}.png" + fig.savefig(out, dpi=200, bbox_inches="tight", facecolor="white") + plt.close(fig) + print(f"[ntc] cell{cell} original vs generative -> {out}") + + +def render_composed(alphas, cell=1, level=4, crop=256, ppu=5600, embedding="phate", out_dir=None, + bg="white", stretch=True, marks=True, channel="phase", stained_dir=None, assets="viewer_assets"): + """EXACTLY reproduce the viewer montage layout (same placed-gene set/order, same compute_priority + + compute_canvas_size + assign_cells_to_grid grid-snapped placement, crop 256 / ppu 5600), then overlay + the added features: per-cell leiden borders, crisp gene labels, and one correct 40S/60S member box. + Native 256px crops (dark cells) on white. -> COMPOSED_DIR.""" + import anndata as ad + from latent_lens.grid import assign_cells_to_grid, compute_priority, compute_canvas_size + from .build_umap_montage import _embed_coords, OUT as MOUT, VIEWER_ALPHAS + from ..classifier.config import slugify + + out_dir = out_dir or COMPOSED_DIR + os.makedirs(out_dir, exist_ok=True) + ann = ad.read_h5ad(UMAP_H5AD) + coords_all = _embed_coords(ann, embedding).astype(np.float32) + perts = ann.obs["perturbation"].astype(str).values + leiden = (ann.obs[LEIDEN].astype(str).values if LEIDEN in ann.obs else np.array([""] * len(perts))) + ebi = (ann.obs["ebi_complex"].astype(str).values if "ebi_complex" in ann.obs else np.array([""] * len(perts))) + va = f"{MOUT}/{assets}/phase/geneKO" # phase-frame source (also defines the placed-gene layout for all channels) + al = list(VIEWER_ALPHAS); a0 = int(np.argmin([abs(x) for x in al])) + + # placed-gene set/order EXACTLY as montage_from_cache: real genes with cache, then NTC nodes + real = [i for i, g in enumerate(perts) + if not g.startswith("NTC") and os.path.exists(f"{va}/{slugify(g)}/cell{cell}/frame_{a0:02d}.webp")] + ntc = [i for i, g in enumerate(perts) if g.startswith("NTC")] + order = np.array(real + ntc) + ntc_ref = slugify(perts[real[0]]) if real else None + gcoords = coords_all[order] + priority = compute_priority(gcoords) + canvas_w, canvas_h, _umin, umap_px_0 = compute_canvas_size(gcoords, ppu, crop, 8) + ds = 2 ** level + gh = (canvas_h // ds) // crop; gw = (canvas_w // ds) // crop + selected, positions = assign_cells_to_grid(umap_px_0 / ds, crop, gh, gw, priority) + if len(selected) == 0: + print("[composed] no cells at this level"); return + + clusters = sorted({v for v in leiden if v not in ("", "nan", "None")}, key=lambda s: (len(s), s)) + cmap = plt.get_cmap("hsv") + lut = {c: cmap(i / max(1, len(clusters) - 1)) for i, c in enumerate(clusters)} + axn, ayn = np.ptp(coords_all[:, 0]) or 1, np.ptp(coords_all[:, 1]) or 1 + labels_all = [{"g": perts[i], "nx": float((coords_all[i, 0] - coords_all[:, 0].min()) / axn), + "ny": float((coords_all[i, 1] - coords_all[:, 1].min()) / ayn), LEIDEN: leiden[i]} + for i in range(len(perts))] + + grs = positions[:, 0]; gcs = positions[:, 1]; r0, c0 = grs.min(), gcs.min() + Hc, Wc = int((grs.max() - r0 + 1) * crop), int((gcs.max() - c0 + 1) * crop) + cen = gcoords.mean(0) # embedding centroid (for "arm end" pick) + thin = max(1, crop // 130); mlw = max(9, crop // 16); dpi = 150 + fill = 0 if bg == "black" else 255 + for a in alphas: + ai = int(np.argmin([abs(x - a) for x in al])) + canv = (np.full((Hc, Wc), fill, np.uint8) if channel == "phase" + else np.full((Hc, Wc, 3), 0 if channel == "__allmarkers__" else 255, np.uint8)) # allmarkers: black RGB (fluor merge); single marker: white + magma + info = [] + for sidx, (gr, gc) in zip(selected, positions): + gi = order[sidx]; g = perts[gi]; is_ntc = g.startswith("NTC") + if channel == "phase": + p = f"{va}/{ntc_ref if is_ntc else slugify(g)}/cell{cell}/frame_{(a0 if is_ntc else ai):02d}.webp" + elif channel == "__allmarkers__": # pre-composited RGB 42-marker merge (flat dir) + p = f"{stained_dir}/{'__NTC_c%d' % cell if is_ntc else slugify(g) + '_c%d_a%g' % (cell, a)}.webp" + else: # virtually-stained single marker: per-marker subdir + p = (f"{stained_dir}/{channel}/__NTC_c{cell}.webp" if is_ntc + else f"{stained_dir}/{channel}/{slugify(g)}_c{cell}_a{a:g}.webp") + if not os.path.exists(p): + continue + yy, xx = int((gr - r0) * crop), int((gc - c0) * crop) + if channel == "__allmarkers__": # already RGB (unique hue per marker) → place directly + canv[yy:yy + crop, xx:xx + crop] = np.asarray(Image.open(p).convert("RGB"))[:crop, :crop] + else: + im = np.asarray(Image.open(p).convert("L"))[:crop, :crop] + if stretch: + im = _stretch(im) # match the built montage's punchy contrast + if channel == "phase": + canv[yy:yy + crop, xx:xx + crop] = im + else: # magma-map the fluor tile, composite onto white + canv[yy:yy + crop, xx:xx + crop] = (plt.get_cmap("magma")(im / 255.0)[..., :3] * 255).astype(np.uint8) + info.append((xx, yy, "NTC" if is_ntc else g, leiden[gi], ebi[gi], is_ntc, gi)) + + fig = plt.figure(figsize=(Wc / dpi, Hc / dpi), dpi=dpi, facecolor=bg) + ax = fig.add_axes([0, 0, 1, 1]) + (ax.imshow(canv, cmap="gray", vmin=0, vmax=255, aspect="auto") if channel == "phase" + else ax.imshow(canv, aspect="auto")); ax.axis("off") + fs = crop * 0.135 / dpi * 72 + for xx, yy, gname, lv, _e, is_ntc, _gi in info: # thin per-cell leiden border + matching label + col = "#111" if is_ntc else lut.get(lv, (0.7, 0.7, 0.7, 1)) + ax.add_patch(Rectangle((xx, yy), crop, crop, fill=False, ec=col, lw=thin, alpha=0.5, zorder=4)) + tcol = "white" if is_ntc else tuple(0.5 + 0.5 * c for c in col[:3]) # lighten label for legibility + ax.text(xx + crop * 0.04, yy + crop * 0.04, gname, fontsize=fs, color=tcol, weight="bold", + family="monospace", va="top", ha="left", zorder=5, + path_effects=[pe.withStroke(linewidth=fs * 0.2, foreground="black")]) + for lab, (cname, mcol) in (MARKS.items() if marks else []): # pinned (or farthest-out) member, bold box + m = [c for c in info if c[4] == cname] + if not m: + continue + pinned = [c for c in m if c[2] == MARK_GENE.get(lab)] + it = pinned[0] if pinned else max(m, key=lambda c: float(np.linalg.norm(coords_all[c[6]] - cen))) + xx, yy = it[0], it[1] # border sits just OUTSIDE the tile edge (inner edge at boundary) + ax.add_patch(Rectangle((xx - mlw / 2, yy - mlw / 2), crop + mlw, crop + mlw, fill=False, + ec=mcol, lw=mlw, zorder=6)) + ax.text(xx + crop / 2, yy - mlw, f"ribo{lab}", ha="center", va="bottom", + color=mcol, fontsize=fs * 1.2, weight="bold", zorder=7) + _legend(ax, labels_all, Wc, dpi, box=(0.004, 0.004, 0.2, 0.2), dark=(bg == "black")) + out = f"{out_dir}/composed_a{a:g}_L{level}_n{len(info)}_{Wc}x{Hc}.png" + fig.savefig(out, dpi=dpi, facecolor=bg); plt.close(fig) + print(f"[composed] a{a:g} L{level}: {len(info)} cells, {Wc}x{Hc} -> {out}") + + +def _parse_range(s): + if "-" in s: + lo, hi = (int(x) for x in s.split("-")); return list(range(lo, hi + 1)) + return [int(x) for x in s.split(",")] + + +if __name__ == "__main__": + ap = argparse.ArgumentParser() + ap.add_argument("--alphas", default="1-5") + ap.add_argument("--levels", default="3,4") + ap.add_argument("--out", default=None) + ap.add_argument("--traversal", action="store_true", help="also render the 40S->60S complex traversal panels") + ap.add_argument("--ntc-pair", dest="ntc_pair", action="store_true", help="render the NTC original-vs-generative panel") + ap.add_argument("--cells", type=int, nargs="+", default=[0, 1, 2, 3, 4, 5], help="cells for the traversal panels") + ap.add_argument("--montage", dest="montage", action="store_true", default=True) + ap.add_argument("--no-montage", dest="montage", action="store_false", help="skip montage; only traversal") + ap.add_argument("--hires", action="store_true", help="rebuild montage at crop 512 (crisp text + leiden borders) first") + ap.add_argument("--crop", type=int, default=512) + ap.add_argument("--composed", action="store_true", help="reproduce viewer montage layout + overlaid features") + ap.add_argument("--no-stretch", dest="stretch", action="store_false", help="skip per-crop percentile contrast stretch") + ap.add_argument("--no-marks", dest="marks", action="store_false", help="omit the 40S/60S highlight boxes") + ap.add_argument("--bg", default="white", choices=["white", "black"]) + a = ap.parse_args() + alphas = _parse_range(a.alphas) + if a.composed: + for lv in _parse_range(a.levels): + render_composed(alphas, level=lv, out_dir=a.out, bg=a.bg, stretch=a.stretch, marks=a.marks) + tiles_root = None + if a.hires: + build_hires(alphas, crop=a.crop) + tiles_root = HIRES_ROOT + if a.montage and not a.composed: + for al in alphas: + render(al, _parse_range(a.levels), a.out, tiles_root=tiles_root) + if a.traversal: + render_traversal(cells=a.cells, out_dir=a.out) + if a.ntc_pair: + for c in a.cells: + render_ntc_pair(cell=c, out_dir=a.out) diff --git a/src/ops_model/models/attention/diffex/viewer/score_generated.py b/src/ops_model/models/attention/diffex/viewer/score_generated.py new file mode 100644 index 0000000..b1ce894 --- /dev/null +++ b/src/ops_model/models/attention/diffex/viewer/score_generated.py @@ -0,0 +1,285 @@ +"""v5 SetTransformer accuracy of GENERATED traversal frames — the accuracy-vs-α validation overlay. + +v5 is no-mask/160px, so generated frames feed the classifier directly: + frame → embed_crops (no seg mask) → z-standardize on the NTC control (cached anchor ctrl.npz) + → bag → v5 SetTransformer → softmax → P(target). +Writes scores_v5.json into each traversal dir: {alphas, p_target:[...], top1_target:[...]}. + +geneKO target = the gene's class prob. complex target = sum of member-gene probs (ebionly model, +gene-labeled), matching the complex-centroid definition. A→B alt-anchor = the target (B) prob. +""" +from __future__ import annotations +import glob +import json +import os + +import numpy as np +from PIL import Image + +CACHE = os.environ.get("OPS_DIFFEX_ASSETS", "viewer_assets") +V5_BASE = f"/hpc/projects/icd.fast.ops/models/diffex/{CACHE}/phase" + + +def _emb_frames(cfg, trav, ai, embed_crops): + """embed_crops the α=ai frame of every cell in a traversal (reverse _save_webp → [-1,1], no mask).""" + imgs = [] + for c in range(len(glob.glob(f"{trav}/cell*"))): + f = f"{trav}/cell{c}/frame_{ai:02d}.webp" + if os.path.exists(f): + imgs.append(np.asarray(Image.open(f).convert("L"), np.float32) / 255.0 * 2 - 1) + if not imgs: + return None + return embed_crops(np.stack(imgs)[:, None].astype(np.float32), cfg, cache_path=None) + + +def diagnose(gene="TOMM20", device="cuda"): + """Isolate the generated~0 issue: score REAL cells of `gene` through OUR embed_crops pipeline + (raw, and z-standardized on NTC control), vs Alex's own val embeddings (positive control).""" + import torch + from .set_classifier import load_set_classifier, score_bags, V5_CKPT_ROOT, V5_RUNS + from .precompute import _gather_class + from ..directions.config import DirConfig + m, g2i, c2i = load_set_classifier(run=V5_RUNS[("phase", "geneKO")], device=device, root=V5_CKPT_ROOT) + ci = c2i.get("Phase2D", 0); gi = g2i[gene] + # Alex val (already z-std) — positive control + av = torch.load(f"/hpc/projects/icd.fast.ops/models/alex_lin_attention/v5/paper_v2_phase/val/{gene}.pt", map_location="cpu")["embeddings"].numpy() + p_alex = float(score_bags(m, av[:100][None], channel_idx=ci, device=device)[0][gi]) + # our pipeline on REAL cells + cfg = DirConfig(grain="geneKO", target=gene, device=device) + _, real = _gather_class(cfg, gene, 100) # materialize + embed_crops (our extraction) + z = np.load(f"{V5_BASE}/_anchors/NTC/ctrl.npz"); mu = z["ctrl_embs"].mean(0); sd = z["ctrl_embs"].std(0) + 1e-6 + p_raw = float(score_bags(m, real[None], channel_idx=ci, device=device)[0][gi]) + p_zstd = float(score_bags(m, ((real - mu) / sd)[None], channel_idx=ci, device=device)[0][gi]) + print(f"[diagnose {gene}] Alex-val={p_alex:.3f} | OUR-real raw={p_raw:.3f} zstd-on-NTC={p_zstd:.3f}") + print(f" emb scale: Alex mean={av.mean():.2f} std={av.std():.2f} | our-real mean={real.mean():.2f} std={real.std():.2f} | our-zstd mean={((real-mu)/sd).mean():.2f} std={((real-mu)/sd).std():.2f}") + import json + json.dump({"gene": gene, "p_alex": p_alex, "p_raw": p_raw, "p_zstd": p_zstd}, + open("/hpc/projects/icd.fast.ops/models/diffex/viewer_assets_v5/_diag.json", "w")) + + +def diagnose2(gene="TOMM20", device="cuda"): + """Test the domain-offset fix: standardize GENERATED frames by GENERATED-NTC(α0) stats instead of real NTC. + Reports P(target) vs α for both references + whether generated α0 is recognized as NTC.""" + import torch, json + from .set_classifier import load_set_classifier, score_bags, V5_CKPT_ROOT, V5_RUNS + from ..directions.config import DirConfig + from ..classifier.celldino_features import embed_crops + m, g2i, c2i = load_set_classifier(run=V5_RUNS[("phase", "geneKO")], device=device, root=V5_CKPT_ROOT) + ci = c2i.get("Phase2D", 0); gi = g2i[gene]; ni = g2i.get("NTC") + trav = f"{V5_BASE}/geneKO/{gene}"; alphas = json.load(open(f"{trav}/meta.json"))["alphas"] + embs = [_emb_frames(DirConfig(grain="geneKO", target=gene, device=device), trav, ai, embed_crops) for ai in range(len(alphas))] + z0 = len(alphas) // 2 + zr = np.load(f"{V5_BASE}/_anchors/NTC/ctrl.npz"); mu_r, sd_r = zr["ctrl_embs"].mean(0), zr["ctrl_embs"].std(0) + 1e-6 + mu_g, sd_g = embs[z0].mean(0), embs[z0].std(0) + 1e-6 # generated-NTC (α0) stats + def curve(mu, sd, idx): + return [round(float(score_bags(m, ((e - mu) / sd)[None], channel_idx=ci, device=device)[0][idx]), 3) for e in embs] + res = {"gene": gene, "alphas": alphas, "p_target_realNTC": curve(mu_r, sd_r, gi), + "p_target_genNTC": curve(mu_g, sd_g, gi), + "p_ntc_genNTC": curve(mu_g, sd_g, ni) if ni is not None else None} + json.dump(res, open("/hpc/projects/icd.fast.ops/models/diffex/viewer_assets_v5/_diag2.json", "w")) + print("DIAG2 DONE") + + +def _real_expectation(grain, target): + """Alex real-cell top1_acc by bag size. geneKO → the gene's row; complex → MEAN over the complex's + member-gene rows (grouped by Alex's label_name in the ebionly eval).""" + import csv as _csv + from collections import defaultdict + E = "/hpc/projects/icd.fast.ops/models/alex_lin_attention/v5/phase" + if grain == "geneKO": + return {int(r["n_cells"]): float(r["top1_acc"]) for r in _csv.DictReader(open(f"{E}/eval_phase_e200_pergene_val.csv")) if r["gene_name"] == target} + by = defaultdict(list) + for r in _csv.DictReader(open(f"{E}/eval_phase_ebionly_e200_pergene_val.csv")): + if r["label_name"] == target: + by[int(r["n_cells"])].append(float(r["top1_acc"])) + return {b: float(np.mean(v)) for b, v in by.items()} + + +def bag_experiment(grain="geneKO", target="MICOS13", n_max=200, sizes=(20, 50, 100, 150, 200), n_bags=30, device="cuda"): + """Regenerate `target` (grain geneKO|complex) with n_max cells, then at the peak α sweep bag size → + generated top1_acc + mean P(target), vs Alex's REAL top1_acc-by-bag (gene row, or mean-member for a + complex). Writes _bagexp_.json. Isolated OPS_DIFFEX_ASSETS (bagtest dir set by caller).""" + import json, torch + from .precompute import precompute_marker + from .submit import PHASE_CK + from .set_classifier import load_set_classifier, score_bags, V5_CKPT_ROOT, V5_RUNS + from ..directions.config import DirConfig + from ..classifier.celldino_features import embed_crops + from ..classifier.config import slugify + OUT = "/hpc/projects/icd.fast.ops/models/diffex" + run = V5_RUNS[("phase", "geneKO" if grain == "geneKO" else "complex_ebionly")] + precompute_marker(grain=grain, targets=[target], ckpt=PHASE_CK, out_root=OUT, n_cells=n_max, + score=False, device=device, force=True) + model, cmap, c2i = load_set_classifier(run=run, device=device, root=V5_CKPT_ROOT) + if target not in cmap: + print(f"[skip] {target} not in class map"); return + gi = cmap[target]; ci = c2i.get("Phase2D", 0) + slug = target if grain == "geneKO" else slugify(target) + trav = f"{V5_BASE}/{grain}/{slug}"; alphas = json.load(open(f"{trav}/meta.json"))["alphas"] + cfg = DirConfig(grain=grain, target=(target if grain == "geneKO" else "NTC"), device=device) + embs = [_emb_frames(cfg, trav, ai, embed_crops) for ai in range(len(alphas))] + z0 = len(alphas) // 2; mu = embs[z0].mean(0); sd = embs[z0].std(0) + 1e-6 + N = len(embs[z0]) + fullp = [float(score_bags(model, ((embs[ai] - mu) / sd)[None], channel_idx=ci, device=device)[0][gi]) for ai in range(len(alphas))] + ai_pk = int(np.argmax(fullp)); Epk = (embs[ai_pk] - mu) / sd + rng = np.random.default_rng(0) + res = {"grain": grain, "target": target, "alphas": alphas, "n_generated": N, "peak_alpha": alphas[ai_pk], + "real_expectation": _real_expectation(grain, target), "bag": {}} + for sz in sizes: + if sz > N: + continue + nb = 1 if sz >= N else n_bags + ps, t1 = [], [] + for _ in range(nb): + idx = rng.choice(N, sz, replace=False) + prob = score_bags(model, Epk[idx][None], channel_idx=ci, device=device)[0] + ps.append(float(prob[gi])); t1.append(int(int(np.argmax(prob)) == gi)) + res["bag"][str(sz)] = {"top1_acc": float(np.mean(t1)), "mean_p": float(np.mean(ps)), "n_bags": nb} + print(f"[bag {sz}] {target[:30]} gen top1={np.mean(t1):.2f} | real={res['real_expectation'].get(sz,'-')}") + os.makedirs(f"{OUT}/viewer_assets_v5_bagtest", exist_ok=True) + json.dump(res, open(f"{OUT}/viewer_assets_v5_bagtest/_bagexp_{slug}.json", "w")) + print(f"BAGEXP DONE {target} (peak α={alphas[ai_pk]:+g})") + + +def score_embs_v5(embs, alphas, tgt, model, g2i, ci, run, device="cuda", bag=None): + """Score already-embedded per-α traversal frames with the v5 SetTransformer. embs[ai] = ncell×1024 + CellDINO embs for α=alphas[ai] (or None). Standardize on the α0 (middle) generated frames → removes + the DiffAE domain offset. bag: score a FIXED-size bag (first `bag` cells) even if more are generated, + for cross-approach comparability. Returns scores_v5.json dict, or None if tgt not in class map / no α0.""" + from .set_classifier import score_bags + idxs = [g2i[tgt]] if tgt in g2i else [] + z0 = len(alphas) // 2 + if not idxs or embs[z0] is None: + return None + E = [None if e is None else (e[:bag] if bag else e) for e in embs] # fixed bag → same statistic across approaches + mu = E[z0].mean(0); sd = E[z0].std(0) + 1e-6 + tset = set(idxs) + ptgt, top1, top5, ranks = [], [], [], [] + for emb in E: + if emb is None: + ptgt.append(None); top1.append(None); top5.append(None); ranks.append(None); continue + prob = score_bags(model, ((emb - mu) / sd)[None], channel_idx=ci, device=device)[0] + order = np.argsort(prob)[::-1].tolist() + ptgt.append(float(prob[idxs].sum())) + top1.append(int(order[0] in tset)) + top5.append(int(bool(tset & set(order[:5])))) # target among the 5 most-likely classes + ranks.append(int(min(order.index(i) for i in idxs)) + 1) # 1-indexed rank of the target class + return {"alphas": alphas, "p_target": ptgt, "top1_target": top1, "top5_target": top5, "rank_target": ranks, "run": run, "bag": bag or len(E[z0])} + + +MAIN_V5 = "/hpc/projects/icd.fast.ops/models/diffex/viewer_assets_v5/phase" + + +def bag_scaling(grain, targets, n_max=500, sizes=(20, 50, 100, 200), n_bags=30, device="cuda"): + """Lean bag-scaling pass. For each target: read its PEAK α from the existing main-build scores_v5.json, + generate ONLY {α=0, peak α} × n_max cells (α=0 = standardization reference; peak = the pool), then sweep + bag size sampling n_bags DISTINCT bags per size (resampled from the pool, like Alex's real-cell eval) → + generated top1_acc + mean P(target) vs Alex's real expectation. Writes _bagexp_.json in bagtest.""" + import json, torch + from .precompute import precompute_marker + from .submit import PHASE_CK + from .set_classifier import load_set_classifier, score_bags, V5_CKPT_ROOT, V5_RUNS + from ..directions.config import DirConfig + from ..classifier.celldino_features import embed_crops + from ..classifier.config import slugify + OUT = "/hpc/projects/icd.fast.ops/models/diffex" + run = V5_RUNS[("phase", "geneKO" if grain == "geneKO" else "complex_ebionly")] + model, cmap, c2i = load_set_classifier(run=run, device=device, root=V5_CKPT_ROOT) + ci = c2i.get("Phase2D", 0); rng = np.random.default_rng(0) + for tgt in targets: + try: + if tgt not in cmap: + print(f"[skip] {tgt} not in class map"); continue + gi = cmap[tgt]; slug = tgt if grain == "geneKO" else slugify(tgt) + sc = json.load(open(f"{MAIN_V5}/{grain}/{slug}/scores_v5.json")) # peak α from the main 20-cell build + pk_a = sc["alphas"][int(np.nanargmax([-1 if v is None else v for v in sc["p_target"]]))] + precompute_marker(grain=grain, targets=[tgt], ckpt=PHASE_CK, out_root=OUT, n_cells=n_max, + alphas=[0.0, pk_a], score=False, device=device, force=True) # only α0 + peak α + trav = f"{V5_BASE}/{grain}/{slug}" + cfg = DirConfig(grain=grain, target=(tgt if grain == "geneKO" else "NTC"), device=device) + al = sorted([0.0, pk_a]); pk_i = al.index(pk_a); z0_i = al.index(0.0) + e0 = _emb_frames(cfg, trav, z0_i, embed_crops); ep = _emb_frames(cfg, trav, pk_i, embed_crops) + mu = e0.mean(0); sd = e0.std(0) + 1e-6; E = (ep - mu) / sd; N = len(E) + res = {"grain": grain, "target": tgt, "peak_alpha": pk_a, "n_generated": N, + "real_expectation": _real_expectation(grain, tgt), "bag": {}} + for sz in sizes: + if sz > N: + continue + nb = 1 if sz >= N else n_bags + ps, t1 = [], [] + for _ in range(nb): + idx = rng.choice(N, sz, replace=False) # distinct resampled bag (Alex-style) + prob = score_bags(model, E[idx][None], channel_idx=ci, device=device)[0] + ps.append(float(prob[gi])); t1.append(int(int(np.argmax(prob)) == gi)) + res["bag"][str(sz)] = {"top1_acc": float(np.mean(t1)), "mean_p": float(np.mean(ps)), "n_bags": nb} + os.makedirs(f"{OUT}/viewer_assets_v5_bagtest", exist_ok=True) + json.dump(res, open(f"{OUT}/viewer_assets_v5_bagtest/_bagexp_{slug}.json", "w")) + print(f"[bagscale {tgt[:30]}] peakα={pk_a:+g} N={N} @20={res['bag'].get('20',{}).get('top1_acc')} @200={res['bag'].get('200',{}).get('top1_acc')}") + except Exception as e: + import traceback; print(f"[ERR {tgt}] {repr(e)[:120]}"); traceback.print_exc() + + +def score_targets(grain, targets, device="cuda", run=None, members_map=None, bag=None): + """Score a list of traversals (grain='geneKO'|'complex'). members_map: complex→member genes (complex only). + Writes scores_v5.json per traversal dir. Returns {target: p_target-per-α}.""" + import torch # noqa + from .set_classifier import load_set_classifier, score_bags, V5_CKPT_ROOT, V5_RUNS + from ..directions.config import DirConfig + from ..classifier.celldino_features import embed_crops + from ..classifier.config import slugify + run = run or V5_RUNS[("phase", "geneKO" if grain == "geneKO" else "complex_ebionly")] + model, g2i, c2i = load_set_classifier(run=run, device=device, root=V5_CKPT_ROOT) + ci = c2i.get("Phase2D", 0) + sub = "geneKO" if grain == "geneKO" else "complex" + out = {} + for tgt in targets: + try: # per-target guard: one bad target must not kill the shard + # geneKO → gene class; complex → the model classifies complexes directly (label_to_idx). g2i is the model's class map. + idxs = [g2i[tgt]] if tgt in g2i else [] + trav = f"{V5_BASE}/{sub}/{tgt if grain == 'geneKO' else slugify(tgt)}" + mp = f"{trav}/meta.json" + if not idxs or not os.path.exists(mp): + print(f"[skip] {tgt}: idxs={len(idxs)} meta={os.path.exists(mp)}"); continue + alphas = json.load(open(mp))["alphas"] + cfg = DirConfig(grain=grain, target=(tgt if grain == "geneKO" else "NTC"), device=device) + embs = [_emb_frames(cfg, trav, ai, embed_crops) for ai in range(len(alphas))] + d = score_embs_v5(embs, alphas, tgt, model, g2i, ci, run, device, bag) + if d is None: + print(f"[skip] {tgt}: no α0 frames / not in class map"); continue + json.dump(d, open(f"{trav}/scores_v5.json", "w")) + out[tgt] = d["p_target"]; z0 = len(alphas) // 2 + print(f"[score] {tgt}: a0={d['p_target'][z0]:.3f} a+5={d['p_target'][-1]:.3f} rise={d['p_target'][-1]-d['p_target'][z0]:+.3f}") + except Exception as e: + import traceback; print(f"[ERR] {tgt}: {repr(e)[:150]}"); traceback.print_exc() + return out + + +def score_anchor_traversals(grain, device="cuda"): + """Score the A→B alt-anchor traversals: P(target B) per α via the v5 SetTransformer, standardizing on the + α0 (= anchor-A) frames — same recipe as the NTC traversals, just target=B. Writes scores_v5.json per dir.""" + import glob + from .set_classifier import load_set_classifier, V5_CKPT_ROOT, V5_RUNS + from ..directions.config import DirConfig + from ..classifier.celldino_features import embed_crops + run = V5_RUNS[("phase", "geneKO" if grain == "geneKO" else "complex_ebionly")] + model, g2i, c2i = load_set_classifier(run=run, device=device, root=V5_CKPT_ROOT) + ci = c2i.get("Phase2D", 0) + sub = "geneKO" if grain == "geneKO" else "complex" + out = {} + for trav in sorted(glob.glob(f"{V5_BASE}/{sub}/*__to__*")): + mp = f"{trav}/meta.json" + if not os.path.exists(mp): + continue + m = json.load(open(mp)); B = m["target"]; alphas = m["alphas"] + cfg = DirConfig(grain=grain, target=B, device=device) + embs = [_emb_frames(cfg, trav, ai, embed_crops) for ai in range(len(alphas))] + d = score_embs_v5(embs, alphas, B, model, g2i, ci, run, device) + name = os.path.basename(trav) + if d is None: + print(f"[skip] {name}: B={B} not in class map / no α0"); continue + json.dump(d, open(f"{trav}/scores_v5.json", "w")) + z0 = len(alphas) // 2; p = d["p_target"] + out[name] = p + print(f"[anchor-score] {name}: a0={p[z0]:.3f} peak={max(v for v in p if v is not None):.3f}") + print(f"ANCHOR SCORING DONE {grain}: {len(out)} traversals") + return out diff --git a/src/ops_model/models/attention/diffex/viewer/set_classifier.py b/src/ops_model/models/attention/diffex/viewer/set_classifier.py new file mode 100644 index 0000000..49b0f00 --- /dev/null +++ b/src/ops_model/models/attention/diffex/viewer/set_classifier.py @@ -0,0 +1,127 @@ +"""Alex Lin's cellstate-set-classifier (SetTransformer) — reconstructed from the checkpoint state_dict ++ config (the training code lives in a private repo). A bag of per-cell CellDINO features → class logits. + +Used for the DiffEx viewer's honest 1-of-N score: feed a bag of the N generated cells at a given α → +softmax → P(target class). Report it against the REAL-cell bag score (the achievable ceiling, since +val_acc is only ~0.50 even on real phase cells). + +Input space = MASKED CellDINO ViT-L/16, z-standardized on control (Alex's train_ops_zstdcontrol_cdino). +Checkpoints: /hpc/projects/icd.fast.ops/models/alex_lin_attention/v4/wandb/cellstate_set_classifier// + miwkg1cy=1K phase geneKO, epzvv0m1=EBI phase, hx6q8byj/ggdfggsn/ciw91el9=fluor. +""" +from __future__ import annotations + +import glob + +import numpy as np +import torch +import torch.nn as nn +import torch.nn.functional as F + +CKPT_ROOT = "/hpc/projects/icd.fast.ops/models/alex_lin_attention/v4/wandb/cellstate_set_classifier" +# v5 (paper-v2) checkpoints: no-mask/160px, cosine head. Same SetClassifier arch (loads clean). +V5_CKPT_ROOT = "/hpc/projects/icd.fast.ops/models/alex_lin_attention/v5/checkpoints" +V5_RUNS = {("phase", "geneKO"): "six29oan", ("phase", "complex_ebionly"): "5dbnlgl5", ("phase", "complex_ebifb"): "tcmqqj8z", + ("fluor", "geneKO"): "fgaf9ni2", ("fluor", "complex_ebionly"): "avzy8p7x", ("fluor", "complex_ebifb"): "gexmq7ks"} + + +class MAB(nn.Module): + """Pre-norm multihead cross-attention block (query attends to kv) + feed-forward, both residual.""" + def __init__(self, d, heads, d_ff): + super().__init__() + self.attn = nn.MultiheadAttention(d, heads, batch_first=True) + self.norm_q = nn.LayerNorm(d) + self.norm_ff = nn.LayerNorm(d) + self.ff = nn.Sequential(nn.Linear(d, d_ff), nn.GELU(), nn.Dropout(0.0), nn.Linear(d_ff, d)) + + def forward(self, q, kv): + a, _ = self.attn(self.norm_q(q), kv, kv, need_weights=False) + h = q + a + return h + self.ff(self.norm_ff(h)) + + +class ISAB(nn.Module): + """Induced set-attention block: inducing points attend to X (cross1), then X attends to that (cross2).""" + def __init__(self, d, heads, m, d_ff): + super().__init__() + self.inducing = nn.Parameter(torch.zeros(1, m, d)) + self.cross1 = MAB(d, heads, d_ff) + self.cross2 = MAB(d, heads, d_ff) + + def forward(self, x): + h = self.cross1(self.inducing.expand(x.size(0), -1, -1), x) + return self.cross2(x, h) + + +class PMA(nn.Module): + """Pooling by multihead attention: k seed(s) attend to the set → k pooled vectors.""" + def __init__(self, d, heads, k, d_ff): + super().__init__() + self.seeds = nn.Parameter(torch.zeros(1, k, d)) + self.cross = MAB(d, heads, d_ff) + + def forward(self, x): + return self.cross(self.seeds.expand(x.size(0), -1, -1), x) + + +class Encoder(nn.Module): + def __init__(self, d, heads, m, n_layers, d_ff): + super().__init__() + self.layers = nn.ModuleList([ISAB(d, heads, m, d_ff) for _ in range(n_layers)]) + self.pool = PMA(d, heads, 1, d_ff) + self.final_norm = nn.LayerNorm(d) + + def forward(self, x): + for lyr in self.layers: + x = lyr(x) + return self.final_norm(self.pool(x).squeeze(1)) + + +class CosineHead(nn.Module): + """Cosine classifier: temperature-scaled cosine similarity between the pooled vector and class prototypes.""" + def __init__(self, d, n): + super().__init__() + self.weight = nn.Parameter(torch.zeros(n, d)) + self.log_scale = nn.Parameter(torch.zeros(())) + + def forward(self, x): + return self.log_scale.exp() * (F.normalize(x, dim=-1) @ F.normalize(self.weight, dim=-1).t()) + + +class SetClassifier(nn.Module): + def __init__(self, d=512, heads=4, m=32, n_layers=2, n_classes=1001, n_channels=1, d_ff=2048): + super().__init__() + self.input_proj = nn.Linear(1024, d) + self.channel_embeddings = nn.Embedding(n_channels, d) + self.concat_proj = nn.Linear(2 * d, d) + self.encoder = Encoder(d, heads, m, n_layers, d_ff) + self.head = nn.Sequential(nn.Identity(), CosineHead(d, n_classes)) + + def forward(self, feats, channel_idx): + x = self.input_proj(feats) # (B,N,d) + ce = self.channel_embeddings(channel_idx)[:, None, :].expand(-1, x.size(1), -1) + x = self.concat_proj(torch.cat([x, ce], -1)) # concat channel conditioning + return self.head(self.encoder(x)) # (B, n_classes) + + +def load_set_classifier(run="miwkg1cy", device="cpu", root=CKPT_ROOT): + """Load a checkpoint → (model, class_to_idx, channel_to_idx). Architecture read from the bundled config. + class_to_idx = label_to_idx when present (the EBI-only model classifies 99 COMPLEXES directly), + else gene_to_idx (the geneKO model classifies 1001 genes). root=V5_CKPT_ROOT for the v5 checkpoints.""" + ckpt = torch.load(glob.glob(f"{root}/{run}/**/*.pt", recursive=True)[0], map_location=device, weights_only=False) + mc = ckpt["config"]["model"] if "model" in ckpt.get("config", {}) else ckpt["config"] + cmap = ckpt.get("label_to_idx") or ckpt["gene_to_idx"] # complex model → label_to_idx (99); geneKO → gene_to_idx + m = SetClassifier(d=mc["d_model"], heads=mc["n_heads"], m=mc["n_inducing_cell"], + n_layers=mc["n_layers_cell"], n_classes=len(cmap), + n_channels=len(ckpt["channel_to_idx"]), d_ff=mc.get("d_ff") or 4 * mc["d_model"]) + m.load_state_dict(ckpt["model_state_dict"]) + m.eval().to(device) + return m, cmap, ckpt["channel_to_idx"] + + +@torch.no_grad() +def score_bags(model, feats, channel_idx=0, device="cpu"): + """feats: (B, N, 1024) bags → softmax probabilities (B, n_classes).""" + f = torch.as_tensor(np.asarray(feats), dtype=torch.float32, device=device) + ci = torch.full((f.size(0),), int(channel_idx), dtype=torch.long, device=device) + return F.softmax(model(f, ci), dim=-1).cpu().numpy() diff --git a/src/ops_model/models/attention/diffex/viewer/submit.py b/src/ops_model/models/attention/diffex/viewer/submit.py new file mode 100644 index 0000000..2b99f18 --- /dev/null +++ b/src/ops_model/models/attention/diffex/viewer/submit.py @@ -0,0 +1,256 @@ +"""Build the DiffEx viewer cache — reproducible, version-controlled entrypoint (replaces the +one-off scratchpad drivers). All target selection comes from `catalog.py`. + + python -m ops_model.models.attention.diffex.viewer.submit seed # per-marker NTC traversals + python -m ops_model.models.attention.diffex.viewer.submit anchors --k 5 # A→B anchor pairs + python -m ops_model.models.attention.diffex.viewer.submit manifest # rebuild manifest.json (local) + python -m ops_model.models.attention.diffex.viewer.submit montage --cell 0 --alpha 2 # harvest cache -> UMAP montage zarr +""" +from __future__ import annotations + +import argparse +import os + +_ASSETS = os.environ.get("OPS_DIFFEX_ASSETS", "viewer_assets") # isolated v5 build → viewer_assets_v5 + +from ops_utils.hpc.slurm_batch_utils import submit_parallel_jobs + +from ..classifier.config import slugify +from . import catalog as C +from .build_umap_montage import build_montage_grid, build_montage_web +from .precompute import build_manifest, precompute_marker, precompute_target + +PHASE_CK = f"{C.DD}/phase_v1/diffae_best.pt" +UMAP_H5AD = ("/hpc/projects/icd.fast.ops/organelle_attribution/pca_optimized_v0.3/cell_dino/" + "zscore_per_exp/paper_v2/phase_only/fixed_80%/cosine/gene_embedding_pca_optimized.h5ad") +PHASE_COMPLEXES = ["40S cytosolic small ribosomal subunit", "60S cytosolic large ribosomal subunit", + "DNA-directed RNA polymerase II complex", "Chaperonin-containing T-complex", "SF3B complex"] +FLUOR_EBI_CSV = "/hpc/projects/icd.fast.ops/models/alex_lin_attention/v4/pma_fluorescent_cells_ebi_all.csv" + + +def _gpu(**kw): + return {"slurm_partition": "gpu", "gpus_per_node": 1, "cpus_per_task": 12, "mem_gb": 64, + "slurm_constraint": "[a100_80|h100|h200|6000_blackwell]", **kw} + + +def _job(name, func, kwargs, stage): + return {"name": name, "func": func, "kwargs": kwargs, "metadata": {"stage": stage}} + + +def cmd_seed(args): + """Per-marker NTC traversals: every complete fluorescent marker (top-N genes) + phase geneKO + phase complex.""" + dist = C.dist_matrix(); jobs = [] + for d, mc, ch in C.complete_markers(min_ep=args.min_ep): + rep = C.rep_of(dist, mc) + if not rep or rep not in dist.columns: + continue + if args.map_thr is not None: # full buildout: ALL genes the marker distinguishes >= thr (0 → all ~1000) + sc = dist[rep] + tg = [g for g in sc.index[sc >= args.map_thr] if not str(g).startswith("NTC")] + else: + tg = C.top_genes(dist, rep, args.n) + if tg: + jobs.append(_job(f"pm_{slugify(mc)[:20]}", precompute_marker, + dict(grain="geneKO", targets=tg, marker_channel=mc, channel=ch, + ckpt=f"{C.DD}/{d}/diffae_best.pt", out_root=C.OUT, load_workers=12, + score=not args.no_score, force=getattr(args, "force", False)), "seed")) + if args.map_thr is None: # phase already fully built — only (re)seed with top-N mode + jobs.append(_job("pm_phase_geneKO", precompute_marker, + dict(grain="geneKO", targets=C.top_genes(dist, "Phase", args.n + 4), + ckpt=PHASE_CK, out_root=C.OUT, load_workers=12, force=getattr(args, "force", False)), "seed")) + jobs.append(_job("pm_phase_complex", precompute_marker, + dict(grain="complex", targets=PHASE_COMPLEXES, ckpt=PHASE_CK, out_root=C.OUT, load_workers=12, force=getattr(args, "force", False)), "seed")) + tgt = sum(len(j["kwargs"]["targets"]) for j in jobs) + print(f"seed: {len(jobs)} per-marker jobs, {tgt} total targets" + + (f" (mAP>={args.map_thr} filter)" if args.map_thr is not None else f" (top-{args.n})")) + sync = getattr(args, "sync", False) # --sync: wait, then refresh manifest/attention/montages + sp = _gpu(timeout_min=args.timeout) + if args.parallel is not None: # default None = no concurrency cap (all markers at once) + sp["slurm_array_parallelism"] = args.parallel + submit_parallel_jobs(jobs_to_submit=jobs, experiment="diffex_gifs", + slurm_params=sp, log_dir="diffex_gifs", wait_for_completion=sync, + post_completion_callback=(lambda *_: run_full_sync()) if sync else None) + + +def cmd_anchors(args): + """A→B anchor traversals: all ordered pairs among each marker's top-K classes (phase + fluor).""" + dist = C.dist_matrix(); jobs = [] + markers = C.complete_markers() + if args.markers: + markers = [m for m in markers if m[1] in args.markers or slugify(m[1]) in args.markers] + for d, mc, ch in markers: + rep = C.rep_of(dist, mc) + top = C.top_genes(dist, rep, args.k) if rep else [] + for a in top: + for b in top: + if a == b: + continue + jobs.append(_job(f"ab_{slugify(mc)[:12]}_{a[:6]}_{b[:6]}", precompute_target, + dict(grain="geneKO", target=b, control=a, marker_channel=mc, channel=ch, + ckpt=f"{C.DD}/{d}/diffae_best.pt", out_root=C.OUT, load_workers=12), "anchors")) + print(f"anchors: {len(jobs)} A→B pair jobs across {len(markers)} markers") + submit_parallel_jobs(jobs_to_submit=jobs, experiment="diffex_gifs", + slurm_params=_gpu(timeout_min=45, slurm_array_parallelism=args.parallel), + log_dir="diffex_gifs", wait_for_completion=False) + + +def cmd_manifest(args): + """Rebuild manifest.json in place (dist mAP for sorting + gene/complex descriptions). Local, no SLURM. + Also writes gene_desc.json (ALL genes) so the viewer shows info for genes not yet cached as targets.""" + import json + va = f"{C.OUT}/{_ASSETS}" + dm = C.desc_map() + build_manifest(C.OUT, dist_map=C.dist_map_for_assets(va), desc_map=dm) + open(f"{va}/gene_desc.json", "w").write(json.dumps(dm)) + print(f"[viewer] gene_desc.json: {len(dm)} entries") + + +def cmd_montage(args): + """Per-marker UMAP montage: place each gene's cached α-frame at its gene-UMAP coord. Layout is ALWAYS + the shared phase gene embedding (UMAP_H5AD); only the images swap per marker (modality). ONE SLURM job + per discrete montage (marker × emb × cell × α) for maximal concurrency; content-aware skip at submit + time so only stale montages are even queued. No decode/re-embed — reads the traversal cache. CPU-only.""" + import glob + import os + import shutil + import time + from pathlib import Path + va = f"{C.OUT}/{_ASSETS}" + # sweep orphan montage zarrs (transient intermediates now live in _montage_zarr, outside viewer_assets; + # each job deletes its own, but killed jobs leave them). Also sweep the legacy viewer_assets/_montage + # location for old stragglers. Skip any <10 min old so a concurrently-running build isn't disturbed. + now = time.time(); freed = 0 + for zdir in (f"{C.OUT}/_montage_zarr", f"{va}/_montage"): + for z in glob.glob(f"{zdir}/*.zarr"): + if now - os.path.getmtime(z) > 600: + shutil.rmtree(z, ignore_errors=True); freed += 1 + if freed: + print(f"[montage] swept {freed} orphan zarr intermediates") + mods = ["phase"] # phase + markers that have geneKO traversal frames + for mdir in sorted(glob.glob(f"{va}/*/geneKO")): + mod = Path(mdir).parent.name + if mod != "phase" and any(e.is_dir() for e in os.scandir(mdir)): # cheap: ≥1 gene dir (don't enumerate all frames) + mods.append(mod) + if args.markers: + mods = [m for m in mods if m in args.markers or slugify(m) in args.markers] + jobs = [] # one job PER (marker, emb, cell, α) = a discrete unit + for mod in mods: + gk = f"{va}/{mod}/geneKO" + cache_mtime = os.path.getmtime(gk) if os.path.isdir(gk) else 0 # bumps when a new gene traversal lands + for emb in args.embeddings: + for cell in args.cells: + for a in args.alphas: + oz = f"{va}/_montage/{mod}_geneKO_{emb}_cell{cell}_a{a:g}.zarr" + tj = f"{oz[:-5]}_tiles/tiles.json" + if not args.force and os.path.exists(tj) and os.path.getmtime(tj) >= cache_mtime: + continue # already reflects the current cache → don't even queue it + jobs.append(_job(f"mtg_{mod[:10]}_{emb[:2]}_c{cell}_a{a:g}", build_montage_web, + dict(h5ad=UMAP_H5AD, out_zarr=oz, cell=cell, alpha=a, embedding=emb, modality=mod), "montage")) + print(f"montage: {len(jobs)} discrete jobs across {len(mods)} markers " + f"(≤{len(mods) * len(args.embeddings) * len(args.cells) * len(args.alphas)} combos; skipped up-to-date)") + submit_parallel_jobs(jobs_to_submit=jobs, experiment="diffex_gifs", + slurm_params={"slurm_partition": "cpu", "cpus_per_task": 4, "mem_gb": 24, "timeout_min": 60, + "slurm_array_parallelism": args.parallel}, + log_dir="diffex_gifs", wait_for_completion=getattr(args, "wait", False)) + + +def run_full_sync(): + """Refresh the entire viewer from the CURRENT cache (all incremental): manifest → attention render → + per-marker montages (full cell×α×emb grid). Safe to re-run; only new/missing assets are built. + Top-level (picklable) so it can be a submitit job func or a seed post-completion callback.""" + import argparse + from .build_attention_heads import submit as attn_submit + # content-aware montage skip handles freshness (rebuilds only montages older than the marker's cache) + ns = argparse.Namespace(cells=list(range(20)), alphas=[1., 2., 3., 4., 5.], + embeddings=["umap", "phate"], markers=None, force=False, wait=True, parallel=100) + print("[sync] 1/3 manifest"); cmd_manifest(ns) + print("[sync] 2/3 attention-head render"); attn_submit(parallel=40) # waits + rebuilds index.json + print("[sync] 3/3 montages"); cmd_montage(ns) # waits + print("[sync] viewer refreshed") + return "sync complete" + + +def cmd_sync(args): + """Make the viewer current. Inline by default; with --after , submit a SLURM gate job that + runs the refresh automatically once those (seed) jobs finish (afterany dependency).""" + if getattr(args, "after", None): + ids = ":".join(str(j) for j in args.after) + print(f"[sync] gating on afterany:{ids} → will refresh when the seed build finishes") + submit_parallel_jobs( + jobs_to_submit=[{"name": "viewer_sync", "func": run_full_sync, "kwargs": {}, "metadata": {"stage": "sync"}}], + experiment="diffex_sync", + slurm_params={"slurm_partition": "cpu", "cpus_per_task": 4, "mem_gb": 16, "timeout_min": 600, + "slurm_additional_parameters": {"dependency": f"afterany:{ids}"}}, + log_dir="diffex_sync", wait_for_completion=False) + else: + run_full_sync() + + +def _chunks(lst, n): + return [lst[i:i + n] for i in range(0, len(lst), n)] + + +def cmd_fluor_complex(args): + """NTC-anchored complex traversals (all EBI complexes) for every complete fluorescent marker.""" + cx = C.ebi_complexes() + markers = C.complete_markers() + if args.markers: + markers = [m for m in markers if m[1] in args.markers or slugify(m[1]) in args.markers] + jobs = [_job(f"fcx_{slugify(mc)[:18]}", precompute_marker, + dict(grain="complex", targets=cx, marker_channel=mc, channel=ch, fluor_csv=C.EBI_FLUOR_CSV, + ckpt=f"{C.DD}/{d}/diffae_best.pt", out_root=C.OUT, load_workers=12, batch=args.batch, force=getattr(args, "force", False)), "fluor_complex") + for d, mc, ch in markers] + print(f"fluor-complex: {len(jobs)} markers × {len(cx)} complexes (batch={args.batch}, no constraint/cap)") + sp = {"slurm_partition": "gpu", "gpus_per_node": 1, "cpus_per_task": 12, "mem_gb": 64, "timeout_min": 720} + submit_parallel_jobs(jobs_to_submit=jobs, experiment="diffex_gifs", slurm_params=sp, + log_dir="diffex_gifs", wait_for_completion=False) + + +def cmd_phase_morpho(args): + """Morpho-demo targets end to end: generated overlay masks + full features + top-accuracy store real ref + + cached production-label real-cell images. Keys from MORPHO_TARGETS (e.g. MICOS13 TOMM20 CCT). One SLURM job + per target, run in PARALLEL (each target stages its own per-target zarr → no shared-zarr race). --parallel caps + concurrency (default = all targets at once).""" + from .morpho_pipeline import build_morpho_target, MORPHO_TARGETS + targets = args.targets or list(MORPHO_TARGETS) + jobs = [_job(f"pm_{slugify(t)}", build_morpho_target, dict(key=t, n_cells=args.n_cells), "phase_morpho") for t in targets] + print(f"phase-morpho: {len(jobs)} target(s) in parallel (cap {args.parallel or len(jobs)}) -> {targets}") + sp = {"slurm_partition": "gpu", "gpus_per_node": 1, "cpus_per_task": 12, "mem_gb": 128, "timeout_min": 240, + "slurm_array_parallelism": args.parallel or len(jobs)} + submit_parallel_jobs(jobs_to_submit=jobs, experiment="diffex_gifs", slurm_params=sp, + log_dir="diffex_morpho", wait_for_completion=False) + + +def cmd_phase_full(args): + """Full phase NTC cache on v1: all ~1000 geneKOs + all EBI complexes, chunked across GPU jobs. + No GPU-type constraint and no concurrency cap — batch is shrunk (default 24 → ~28GB peak) so it + fits any GPU (incl. the plentiful 40GB a100 / 48GB l40s|a40|a6000), maximizing availability.""" + genes, cx = C.all_genes(), C.ebi_complexes() + jobs = [_job(f"phg_{i}", precompute_marker, + dict(grain="geneKO", targets=ch, ckpt=PHASE_CK, out_root=C.OUT, load_workers=12, batch=args.batch, force=getattr(args, "force", False)), "phase_full") + for i, ch in enumerate(_chunks(genes, args.chunk_size))] + jobs += [_job(f"phc_{i}", precompute_marker, + dict(grain="complex", targets=ch, ckpt=PHASE_CK, out_root=C.OUT, load_workers=12, batch=args.batch, force=getattr(args, "force", False)), "phase_full") + for i, ch in enumerate(_chunks(cx, args.chunk_size))] + print(f"phase-full: {len(genes)} geneKO + {len(cx)} complex → {len(jobs)} chunked jobs (batch={args.batch}, no constraint/cap)") + sp = {"slurm_partition": "gpu", "gpus_per_node": 1, "cpus_per_task": 12, "mem_gb": 64, "timeout_min": 720} + submit_parallel_jobs(jobs_to_submit=jobs, experiment="diffex_gifs", slurm_params=sp, + log_dir="diffex_gifs", wait_for_completion=False) + + +def main(): + ap = argparse.ArgumentParser(description="Build the DiffEx viewer cache") + sub = ap.add_subparsers(dest="cmd", required=True) + s = sub.add_parser("seed"); s.add_argument("--n", type=int, default=8); s.add_argument("--map-thr", dest="map_thr", type=float, default=None); s.add_argument("--min-ep", dest="min_ep", type=int, default=0, help="min generator epoch to include; default 0 = no epoch gate (diffae_best.pt banks the peak regardless — epoch != quality)"); s.add_argument("--no-score", dest="no_score", action="store_true"); s.add_argument("--parallel", type=int, default=None, help="max concurrent SLURM tasks; default None = no cap"); s.add_argument("--timeout", type=int, default=180, help="per-marker SLURM timeout (min); bump for full ~1000-gene buildouts"); s.add_argument("--sync", action="store_true", help="on completion, auto-refresh manifest + attention + montages"); s.add_argument("--force", action="store_true", help="rebuild existing traversals in place"); s.set_defaults(fn=cmd_seed) + a = sub.add_parser("anchors"); a.add_argument("--k", type=int, default=5); a.add_argument("--markers", nargs="*"); a.add_argument("--parallel", type=int, default=12); a.set_defaults(fn=cmd_anchors) + m = sub.add_parser("manifest"); m.set_defaults(fn=cmd_manifest) + g = sub.add_parser("montage"); g.add_argument("--cells", type=int, nargs="+", default=list(range(20))); g.add_argument("--alphas", type=float, nargs="+", default=[1.0, 2.0, 3.0, 4.0, 5.0]); g.add_argument("--embeddings", nargs="+", default=["umap", "phate"]); g.add_argument("--markers", nargs="+", help="restrict to these markers (raw or slug); default all with geneKO traversals"); g.add_argument("--force", action="store_true", help="rebuild montages even if tiles already exist"); g.add_argument("--parallel", type=int, default=100, help="max concurrent SLURM tasks"); g.set_defaults(fn=cmd_montage) + fc = sub.add_parser("fluor-complex"); fc.add_argument("--markers", nargs="*"); fc.add_argument("--batch", type=int, default=24); fc.add_argument("--force", action="store_true"); fc.set_defaults(fn=cmd_fluor_complex) + pf = sub.add_parser("phase-full"); pf.add_argument("--chunk-size", type=int, default=50); pf.add_argument("--batch", type=int, default=24); pf.add_argument("--force", action="store_true"); pf.set_defaults(fn=cmd_phase_full) + pm = sub.add_parser("phase-morpho"); pm.add_argument("--targets", nargs="*"); pm.add_argument("--n-cells", dest="n_cells", type=int, default=12); pm.add_argument("--parallel", type=int, default=None, help="max concurrent targets (default = all)"); pm.set_defaults(fn=cmd_phase_morpho) + sy = sub.add_parser("sync", help="refresh manifest + attention + montages from the current cache"); sy.add_argument("--after", nargs="+", help="SLURM job IDs to gate on (afterany); refreshes when they finish"); sy.set_defaults(fn=cmd_sync) + args = ap.parse_args(); args.fn(args) + + +if __name__ == "__main__": + main() diff --git a/src/ops_model/models/attention/diffex/viewer/webapp/app.js b/src/ops_model/models/attention/diffex/viewer/webapp/app.js new file mode 100644 index 0000000..eb72e25 --- /dev/null +++ b/src/ops_model/models/attention/diffex/viewer/webapp/app.js @@ -0,0 +1,1888 @@ +// DiffEx traversal viewer — static, reads viewer_assets/manifest.json (or window.MANIFEST_URL). +// α scrubs precomputed WebP frames. One grid: perturbation rows (current + pinned) × cells-per-page. +const ASSET_VER = "v5"; // v4/v5 toggle removed — the viewer always uses the v5 assets +const ASSET_PREFIX = ""; // app is served from inside the v5 assets dir — data is alongside index.html +const MANIFEST_URL = window.MANIFEST_URL || (ASSET_PREFIX + "manifest.json"); +const BASE = MANIFEST_URL.replace(/manifest\.json$/, ""); + +// ── internal (full) vs public (demo subset) feature gating ───────────────── +// The real public deployment shows only Top cells + Traversal (+ How-it-works steps 1–6). +// Env resolution: window.VIEWER_ENV (future Argus-injected env.js) → hostname rule → default internal. +// Add the prod ingress host to PUBLIC_HOSTS once assigned to auto-enable public mode there. +const PUBLIC_HOSTS = []; +const IS_PUBLIC_DEPLOY = (() => { + const v = (window.VIEWER_ENV || "").toLowerCase(); + if (["public", "prod", "production"].includes(v)) return true; + if (["internal", "staging", "rdev", "dev"].includes(v)) return false; + return PUBLIC_HOSTS.includes(location.hostname); +})(); +const PUBLIC_HIDDEN_TABS = new Set(["montage", "pc", "attn"]); // hidden in public; Top cells/Traversal/How-it-works stay +const PUBLIC_HIDDEN_GRAINS = new Set(["minibinder", "pc"]); // hidden from the Type selector in public +const isPublic = () => IS_PUBLIC_DEPLOY || state.publicPreview; // publicPreview = staging-only manual toggle +function setSidePanel(m, init) { // right panel: 'info' (selected perturbation) | 'about' (viewer overview). Re-click active → hide. + const bar = $("sidebar"); + if (!init && state.sidePanel === m) { bar.classList.toggle("hidden"); } // re-click toggles visibility + else { + state.sidePanel = m; + $("side-info-view").style.display = m === "info" ? "" : "none"; + $("side-about-view").style.display = m === "about" ? "" : "none"; + if (!init) bar.classList.remove("hidden"); + } + const vis = !bar.classList.contains("hidden"); // a button is blue ONLY when its panel is actually showing + $("side-info").classList.toggle("active", vis && state.sidePanel === "info"); + $("side-about").classList.toggle("active", vis && state.sidePanel === "about"); + if (!init && typeof saveState === "function") saveState(); +} +const NOCACHE = "?t=" + Date.now(); // per-load cache-bust for the small JSON metadata (manifest/index/labels/…) + // so reloads always get the freshly-rebuilt data; images stay cached +const PAD = (i) => String(i).padStart(2, "0"); +const $ = (id) => document.getElementById(id); + +// universal image levels: rewindow all page images/canvases to [lo,hi] via the #img-levels SVG filter. +// output = clamp((in - lo)/(hi - lo)); lo=0,hi=1 → identity (no-op). Display-time RGB stretch, not raw-data. +function updateImgLevels() { + const lo = state.imgClimLo, hi = state.imgClimHi, d = Math.max(1e-3, hi - lo); + const slope = (1 / d).toFixed(4), icpt = (-lo / d).toFixed(4); + for (const id of ["lvlR", "lvlG", "lvlB"]) { const f = $(id); if (f) { f.setAttribute("slope", slope); f.setAttribute("intercept", icpt); } } +} + +// show/hide features for the current mode (internal vs public). Cosmetic gate — the real +// protection is that the hidden tabs' data isn't in the public S3 bucket. renderMethods() +// independently trims the deck to steps 1–6 when isPublic(). +function applyFeatureGate() { + const pub = isPublic(); + document.querySelectorAll("#tabbar .tab, .about-tabdesc").forEach(b => + b.classList.toggle("feat-hidden", pub && PUBLIC_HIDDEN_TABS.has(b.dataset.tab))); + const grain = $("grain"); // hide minibinder + PC from the Type selector in public + if (grain && grain._seg) grain._seg.querySelectorAll("button").forEach(b => + b.classList.toggle("feat-hidden", pub && PUBLIC_HIDDEN_GRAINS.has(b.dataset.value))); + if (grain && pub && PUBLIC_HIDDEN_GRAINS.has(grain.value)) { // currently on a hidden type → fall back to geneKO + grain.value = "geneKO"; grain.dispatchEvent(new Event("change", { bubbles: true })); + } + if (state.manifest) updateBagUI(); // public hides the anchor-bag selector (forces multi_bag); reflect on toggle too + if (pub && PUBLIC_HIDDEN_TABS.has(state.view)) { // current view just got hidden → fall back to Traversal + const t = document.querySelector('#tabbar .tab[data-tab="traversal"]'); if (t) t.click(); + } else if (state.view === "methods") renderMethods(); // re-trim the deck in place + const tg = $("envtoggle"); + if (tg) { tg.classList.toggle("pub-on", state.publicPreview); + tg.textContent = state.publicPreview ? "👁 Previewing public" : "🔒 Preview public"; } +} + +const state = { + manifest: null, marker: null, markerIdx: null, targets: [], target: null, anchor: "NTC", sidePanel: "info", + cellCount: 8, page: 0, pinned: [], panels: [], alphas: [], + idx: 0, playing: false, playSeq: [], playPos: 0, frameMs: 180, // default 1× (180ms/frame) + scoreMode: "ptarget", showReal: false, scores: {}, scoresV5: {}, groups: [], realAcc20: null, // scoreMode: none|linear|ptarget|rank. scores[dir]=linear per-cell; scoresV5[dir]=v5 set (bag,per-α); realAcc20[dir]=real top1@bag20 + + pausePoints: new Set(), pauseN: -1, // α indices where autoplay dwells (click ticks to toggle) + rangeLo: 0, rangeHi: 0, alphaLimit: 5, // autoplay sweeps only within ±alphaLimit (scrub stays full) + targetSort: "setacc", // perturbation list order default: "setacc" (SetTransformer set-accuracy) | "map" | "alpha" + publicPreview: false, // staging-only: preview the public feature subset without a real prod deploy + altAnchorsOnly: false, // filter perturbation list to those with a non-NTC (A→B) anchor + view: "traversal", // active view: traversal | montage | attn (all driven by browse selection) + attnIndex: null, attnHeadsCache: {}, attnImgCache: {}, // attention-head assets + attnHead: "all", attnNorm: "map", // default: show ALL heads per cell; per-cell (per-tile max) normalization + attnClimLo: 0, attnClimHi: 1, attnAlpha: 0.6, attnImgOpacity: 1, // clim [vmin,vmax] + constant overlay alpha (Ritvik uses 0.6) + cell-image dimming + imgClimLo: 0, imgClimHi: 1, // universal display-time levels stretch on all page images/canvases (0–1 = no-op) + attnPinned: [], // extra perturbations (geneKO) pinned for side-by-side comparison, like traversal +}; + +// inferno colormap (256 RGB triples, flat) — applied client-side so the attention overlay's +// normalization + opacity are live display options (no baked-in variants). +const INFERNO = Uint8ClampedArray.from([0,0,4,1,0,5,1,1,6,1,1,8,2,1,10,2,2,12,2,2,14,3,2,16,4,3,18,4,3,20,5,4,23,6,4,25,7,5,27,8,5,29,9,6,31,10,7,34,11,7,36,12,8,38,13,8,41,14,9,43,16,9,45,17,10,48,18,10,50,20,11,52,21,11,55,22,11,57,24,12,60,25,12,62,27,12,65,28,12,67,30,12,69,31,12,72,33,12,74,35,12,76,36,12,79,38,12,81,40,11,83,41,11,85,43,11,87,45,11,89,47,10,91,49,10,92,50,10,94,52,10,95,54,9,97,56,9,98,57,9,99,59,9,100,61,9,101,62,9,102,64,10,103,66,10,104,68,10,104,69,10,105,71,11,106,73,11,106,74,12,107,76,12,107,77,13,108,79,13,108,81,14,108,82,14,109,84,15,109,85,15,109,87,16,110,89,16,110,90,17,110,92,18,110,93,18,110,95,19,110,97,19,110,98,20,110,100,21,110,101,21,110,103,22,110,105,22,110,106,23,110,108,24,110,109,24,110,111,25,110,113,25,110,114,26,110,116,26,110,117,27,110,119,28,109,120,28,109,122,29,109,124,29,109,125,30,109,127,30,108,128,31,108,130,32,108,132,32,107,133,33,107,135,33,107,136,34,106,138,34,106,140,35,105,141,35,105,143,36,105,144,37,104,146,37,104,147,38,103,149,38,103,151,39,102,152,39,102,154,40,101,155,41,100,157,41,100,159,42,99,160,42,99,162,43,98,163,44,97,165,44,96,166,45,96,168,46,95,169,46,94,171,47,94,173,48,93,174,48,92,176,49,91,177,50,90,179,50,90,180,51,89,182,52,88,183,53,87,185,53,86,186,54,85,188,55,84,189,56,83,191,57,82,192,58,81,193,58,80,195,59,79,196,60,78,198,61,77,199,62,76,200,63,75,202,64,74,203,65,73,204,66,72,206,67,71,207,68,70,208,69,69,210,70,68,211,71,67,212,72,66,213,74,65,215,75,63,216,76,62,217,77,61,218,78,60,219,80,59,221,81,58,222,82,56,223,83,55,224,85,54,225,86,53,226,87,52,227,89,51,228,90,49,229,92,48,230,93,47,231,94,46,232,96,45,233,97,43,234,99,42,235,100,41,235,102,40,236,103,38,237,105,37,238,106,36,239,108,35,239,110,33,240,111,32,241,113,31,241,115,29,242,116,28,243,118,27,243,120,25,244,121,24,245,123,23,245,125,21,246,126,20,246,128,19,247,130,18,247,132,16,248,133,15,248,135,14,248,137,12,249,139,11,249,140,10,249,142,9,250,144,8,250,146,7,250,148,7,251,150,6,251,151,6,251,153,6,251,155,6,251,157,7,252,159,7,252,161,8,252,163,9,252,165,10,252,166,12,252,168,13,252,170,15,252,172,17,252,174,18,252,176,20,252,178,22,252,180,24,251,182,26,251,184,29,251,186,31,251,188,33,251,190,35,250,192,38,250,194,40,250,196,42,250,198,45,249,199,47,249,201,50,249,203,53,248,205,55,248,207,58,247,209,61,247,211,64,246,213,67,246,215,70,245,217,73,245,219,76,244,221,79,244,223,83,244,225,86,243,227,90,243,229,93,242,230,97,242,232,101,242,234,105,241,236,109,241,237,113,241,239,117,241,241,121,242,242,125,242,244,130,243,245,134,243,246,138,244,248,142,245,249,146,246,250,150,248,251,154,249,252,157,250,253,161,252,255,164]); + +const PALETTE = ["#26c6ff", "#ff5252", "#f0a020", "#7ee787", "#c586ff", "#ff9edb", "#5ad1c7", "#ffd166"]; +// traversal frames/scores/anchor can switch NTC anchor pool (v5 only): accuracy = 25 hand-picked (_v5acc/), attention = v5 build (BASE). __to__ alt-anchors have no accpool → stay on BASE. +function travBase(dir) { return BASE; } // v5 attention+accuracy cells are consolidated into one dir under BASE +const SET_MODES = ["ptarget", "rank"]; // v5 SetTransformer per-traversal (bag) score modes → row-header chip +// adaptive score-overlay legend caption per mode (see #scoremode); shown only on Traversal when scoreMode != none +const SCORE_LEGEND = { + linear: "per-cell classifier score (NTC → knockout): 0 → 1", + ptarget: "P(target) for the whole set (bag): 0 → 100%", + rank: "target rank within the set: ≥100 → rank 1", +}; +function updateScoreLegend() { + const show = state.scoreMode !== "none" && (!state.view || state.view === "traversal"); + $("score-legend").style.display = show ? "flex" : "none"; + if (show) $("score-legend-txt").innerHTML = "score overlay (white → red) — " + (SCORE_LEGEND[state.scoreMode] || ""); +} +function setChip(sv, i) { // {txt, bg, fg, showReal} for the selected set-mode at α-index i, or null. Both modes use the one white→red heat. + if (!sv || !SET_MODES.includes(state.scoreMode)) return null; + if (state.scoreMode === "rank") { + const arr = sv.rank_target; if (!arr) return null; + const r = arr[Math.min(i, arr.length - 1)]; if (r == null) return null; + const v = Math.max(0, Math.min(1, 1 - Math.log10(Math.max(1, r)) / 2)); // rank1→1 (deep red), rank10→0.5, rank≥100→0 (white) + return { txt: `rank ${r}`, bg: heat(v), fg: v > 0.55 ? "#fff" : "#111", showReal: false }; // rank vs real-fraction differ in units → no real overlay + } + const arr = sv.p_target; if (!arr) return null; + const x = arr[Math.min(i, arr.length - 1)]; if (x == null) return null; + return { txt: `set-acc ${Math.round(x * 100)}%`, bg: heat(x), fg: x > 0.55 ? "#fff" : "#111", showReal: true }; +} +const frameURL = (dir, cell, i) => `${travBase(dir)}${dir}/cell${cell}/frame_${PAD(i)}.webp`; +const heat = (v) => { // classifier confidence 0→1 as white → deep red (#99000d) + const r = Math.round(255 + (153 - 255) * v), gg = Math.round(255 - 255 * v), b = Math.round(255 + (13 - 255) * v); + return `rgb(${r},${gg},${b})`; +}; +const pertOf = (markerName, t, anchor) => ({ markerName, target: t.target, anchor, slug: t.slug, + asset_dir: t.asset_dir, alphas: t.alphas, n_cells: t.n_cells, has_real: t.has_real, + real_dir: t.real_dir || t.asset_dir, key: markerName + "|" + t.slug }); + +// ---- persist the browse selection + display prefs across page reloads (localStorage; works on static S3) ---- +const LS_KEY = "opsin.state.v1"; +let restoredAlpha = null; // α VALUE restored from localStorage, applied on the first rebuild after a reload/version-switch +function saveState() { + try { + localStorage.setItem(LS_KEY, JSON.stringify({ + marker: markerLabel(state.markerIdx), grain: $("grain").value, + target: state.target ? state.target.target : null, anchor: state.anchor, + cellCount: state.cellCount, page: state.page, tilepx: $("tile-scale").value, + scoreMode: state.scoreMode, showReal: state.showReal, altAnchor: state.altAnchorsOnly, sidePanel: state.sidePanel, + cols: $("colslayout").checked, tcCols: $("tc-cols").checked, speed: $("speed").value, alphaLimit: $("alphalimit").value, view: state.view, + bag: $("m-bag") ? $("m-bag").value : null, + alpha: (state.alphas && state.idx != null && state.idx < state.alphas.length) ? state.alphas[state.idx] : null, // hold α (by value) across reloads + pinned: state.pinned.map(p => ({ target: p.target, anchor: p.anchor })), + })); + } catch (e) { /* private mode / quota — non-fatal */ } +} +function restoreState() { // returns true if a saved snapshot was applied (skips the default selection) + let s; try { s = JSON.parse(localStorage.getItem(LS_KEY) || "null"); } catch (e) { s = null; } + if (!s) return false; + if (s.alpha != null) restoredAlpha = s.alpha; // consumed by the first rebuild → holds α across reload/version-switch + // filters/prefs that affect the target list — set BEFORE marker/target resolution + if (s.grain && s.grain !== "all") $("grain").value = s.grain; // 'all' grain removed → fall back to default geneKO + if (s.bag && $("m-bag")) $("m-bag").value = s.bag; // restore anchor bag before marker resolution (montageCells depends on it) + if (s.cols != null) { $("colslayout").checked = s.cols; $("grid").classList.toggle("cols-layout", s.cols); } + if (s.tcCols != null) { $("tc-cols").checked = s.tcCols; $("tc-view").classList.toggle("cols-layout", s.tcCols); } + if (s.altAnchor != null) { $("altanchor").checked = s.altAnchor; state.altAnchorsOnly = s.altAnchor; } + if (s.scoreMode) { state.scoreMode = s.scoreMode; $("scoremode").value = s.scoreMode; updateScoreLegend(); } + if (s.showReal != null) { $("showreal").checked = s.showReal; state.showReal = s.showReal; } + if (s.sidePanel) setSidePanel(s.sidePanel, true); + if (s.cellCount) { state.cellCount = s.cellCount; $("cellcount").value = s.cellCount; } + if (s.tilepx) { $("tile-scale").value = s.tilepx; document.documentElement.style.setProperty("--tilepx", s.tilepx + "px"); } + if (s.speed) { $("speed").value = s.speed; state.frameMs = +s.speed; } + if (s.alphaLimit) { $("alphalimit").value = s.alphaLimit; state.alphaLimit = +s.alphaLimit; } + if (s.anchor) state.anchor = s.anchor; // populateAnchors keeps it if valid for the target, else resets to NTC + let mi = state.manifest.markers.findIndex(m => (m.label || m.marker_channel || "Phase") === s.marker); + if (mi < 0) mi = 0; + selectMarker(mi); $("markerfilter").value = markerLabel(mi); // refreshTargets → selects first target by default + if (s.target) { const t = state.targets.find(x => x.target === s.target); if (t) { $("filter").value = targetLabel(t); selectTarget(t.slug); } } + if (Array.isArray(s.pinned) && s.pinned.length) { // re-pin same-marker comparisons + const mc = state.marker.marker_channel || "Phase"; + state.pinned = s.pinned.map(pp => { const e = resolveEntry(pp.target, pp.anchor || "NTC"); return e ? pertOf(mc, e, e.control || "NTC") : null; }).filter(Boolean); + renderPinned(); rebuild(); + } + if (s.page) { state.page = s.page; rebuild(); } // restore the cell page (selectTarget reset it to 0) + if (s.view && s.view !== "traversal") { const b = document.querySelector(`.tab[data-tab="${s.view}"]`); if (b) b.click(); } + return true; +} + +async function boot() { + state.manifest = await (await fetch(MANIFEST_URL + NOCACHE)).json(); + state.geneDesc = await fetch(`${BASE}gene_desc.json${NOCACHE}`).then(r => r.ok ? r.json() : {}).catch(() => ({})); // desc for ALL genes (incl un-cached) + state.geneNarr = await fetch(`${BASE}gene_narrative.json${NOCACHE}`).then(r => r.ok ? r.json() : {}).catch(() => ({})); // affinage mechanistic narratives (gene → prose), pre-fetched + state.realAcc20 = await fetch(`${BASE}real_acc20.json${NOCACHE}`).then(r => r.ok ? r.json() : {}).catch(() => ({})); // real-cell top1_acc@bag20 by asset_dir (feasibility ceiling) + state.attnIndex = await fetch(`${BASE}attention_heads/index.json${NOCACHE}`).then(r => r.ok ? r.json() : null).catch(() => null); // {global_max, assets:{modality:{grain:[keys]}}} + mont.rmMap = await fetch(`${BASE}_montage/render_mode.json${NOCACHE}`).then(r => r.ok ? r.json() : {}).catch(() => ({})); // per-marker renderer: tiles (per-marker montage) vs live + ensureSetacc(() => {}); // preload set-accuracy so the default "by SET ACC" ordering is ready + wireCombo("markerfilter", "marker-list", renderMarkerList, () => markerLabel(state.markerIdx)); + wireCombo("filter", "target-list", renderTargetList, () => state.target ? targetLabel(state.target) : ""); + $("tprev").onclick = () => stepTarget(-1); // step through perturbations quickly + $("tnext").onclick = () => stepTarget(1); + $("target-sort").onchange = () => { state.targetSort = $("target-sort").value; + const show = () => { renderTargetList(); $("target-list").classList.remove("hidden"); }; + state.targetSort === "setacc" ? ensureSetacc(show) : show(); }; + $("altanchor").onchange = () => { state.altAnchorsOnly = $("altanchor").checked; refreshTargets(); }; + $("grain").onchange = refreshTargets; + $("cellcount").onchange = () => { state.cellCount = Math.max(1, +$("cellcount").value | 0); state.page = 0; rebuild(); if (state.view === "attn") renderAttn(); if (state.view === "top") renderTop(); }; + $("cprev").onclick = () => { state.page = Math.max(0, state.page - 1); rebuild(); if (state.view === "attn") renderAttn(); if (state.view === "top") renderTop(); }; + $("cnext").onclick = () => { state.page++; rebuild(); if (state.view === "attn") renderAttn(); if (state.view === "top") renderTop(); }; + $("anchor").onchange = () => { state.anchor = $("anchor").value; rebuild(); }; + $("addpanel").onclick = () => { + const set = activeSet(); if (!set.length) return; + const p = set[0]; // the current resolved (marker, anchor→target) + if (!state.pinned.some(q => q.key === p.key)) state.pinned.push(p); + renderPinned(); rebuild(); + }; + $("clearpanels").onclick = () => { state.pinned = []; renderPinned(); rebuild(); }; + $("alpha").oninput = () => showIdx(+$("alpha").value); + $("alpha").onchange = saveState; // persist α on release so reloads / v4↔v5 hold the current traversal position + $("tile-scale").oninput = () => document.documentElement.style.setProperty("--tilepx", $("tile-scale").value + "px"); // traversal image scale + $("tc-scale").oninput = () => document.documentElement.style.setProperty("--tcpx", $("tc-scale").value + "px"); // top-cells image scale + $("colslayout").onchange = () => $("grid").classList.toggle("cols-layout", $("colslayout").checked); // perturbations rows ↔ columns + $("exportgif").onclick = exportGif; + $("play").onclick = togglePlay; + $("scoremode").onchange = () => { + state.scoreMode = $("scoremode").value; + updateScoreLegend(); + saveState(); showIdx(state.idx); + }; + $("speed").onchange = () => { state.frameMs = +$("speed").value; }; + $("showreal").onchange = () => { state.showReal = $("showreal").checked; rebuild(); }; + $("alphalimit").onchange = () => { state.alphaLimit = +$("alphalimit").value; computeRange(); buildPlaySeq(state.alphas); }; + setSidePanel(state.sidePanel, true); // init the right-panel Info/About toggle + $("a-head").onchange = () => { const v = $("a-head").value; state.attnHead = v === "all" ? "all" : +v; renderAttn(); }; + $("a-norm").onchange = () => { state.attnNorm = $("a-norm").value; renderAttn(); }; + $("a-climlo").oninput = () => { // dual-handle clim; keep lo ≤ hi + let lo = +$("a-climlo").value; if (lo > state.attnClimHi) { lo = state.attnClimHi; $("a-climlo").value = lo; } + state.attnClimLo = lo; renderAttn(); + }; + $("a-climhi").oninput = () => { + let hi = +$("a-climhi").value; if (hi < state.attnClimLo) { hi = state.attnClimLo; $("a-climhi").value = hi; } + state.attnClimHi = hi; renderAttn(); + }; + $("a-alpha").oninput = () => { state.attnAlpha = +$("a-alpha").value; renderAttn(); }; + $("a-img").oninput = () => { state.attnImgOpacity = +$("a-img").value; renderAttn(); }; + $("a-reset").onclick = () => { // reset all attention-head display controls to defaults + Object.assign(state, { attnHead: "all", attnNorm: "map", attnClimLo: 0, attnClimHi: 1, attnAlpha: 0.6, attnImgOpacity: 1 }); + $("a-head").value = "all"; $("a-norm").value = "map"; $("a-climlo").value = 0; $("a-climhi").value = 1; + $("a-alpha").value = 0.6; $("a-img").value = 1; + $("a-norm")._segSync?.(); + renderAttn(); + }; + $("a-pin").onclick = () => { // pin the current perturbation for comparison (mirrors traversal pin) + const r = attnCurrentRef(); + if (r && !state.attnPinned.some(p => sameRef(p, r))) { state.attnPinned.push(r); renderAttnPinned(); renderAttn(); } + }; + $("a-pinclear").onclick = () => { state.attnPinned = []; renderAttnPinned(); renderAttn(); }; + $("i-climlo").oninput = () => { // universal image clim (dual-handle; keep lo ≤ hi) + let lo = +$("i-climlo").value; if (lo > state.imgClimHi) { lo = state.imgClimHi; $("i-climlo").value = lo; } + state.imgClimLo = lo; updateImgLevels(); + }; + $("i-climhi").oninput = () => { + let hi = +$("i-climhi").value; if (hi < state.imgClimLo) { hi = state.imgClimLo; $("i-climhi").value = hi; } + state.imgClimHi = hi; updateImgLevels(); + }; + $("i-climreset").onclick = () => { + state.imgClimLo = 0; state.imgClimHi = 1; $("i-climlo").value = 0; $("i-climhi").value = 1; updateImgLevels(); + }; + document.querySelectorAll(".tab").forEach(b => b.onclick = () => { // view switcher (all views share the browse selection) + const view = b.dataset.tab; state.view = view; + document.querySelectorAll(".tab").forEach(x => x.classList.toggle("active", x === b)); + document.querySelectorAll(".tabpane").forEach(p => p.classList.toggle("hidden", p.id !== "tab-" + view)); + $("stage").classList.toggle("montage-active", view === "montage"); + $("stage").classList.toggle("attn-active", view === "attn"); + $("stage").classList.toggle("pc-active", view === "pc"); + $("stage").classList.toggle("top-active", view === "top"); + $("stage").classList.toggle("methods-active", view === "methods"); + updateScoreLegend(); // score legend is traversal-only + adaptive to the selected overlay mode + if (view === "montage") { updateVsUI(); if (mont.renderMode === "live") liveLoad(); else { ensureMontage(); focusMontageOnSelection(); } } + if (view === "attn") renderAttn(); + if (view === "pc") loadPC(); + if (view === "top") loadTop(); + if (view === "methods") renderMethods(); + }); + fillCellDropdown(); // per-modality cell count (phase 45, markers 20); refilled on marker change + fillOverlayMarkers(); // load the 42 VS marker names into the Overlay dropdown + $("m-cell").value = 1; // default NTC cell = 1 + $("m-alpha").value = "5"; // force default α=5 (exaggerated); overrides any browser-restored form value + const LIVE = () => mont.renderMode === "live"; + $("m-render").onchange = setRenderMode; + $("m-emb").onchange = () => LIVE() ? liveLoad() : loadMontage(); + $("m-alpha").onchange = () => LIVE() ? liveRefresh() : loadMontage(); + const ALPHA_MEANING = { "1": "α = 1 (centroid)", "2": "α = 2", "3": "α = 3", "4": "α = 4", "5": "α = 5 (exaggerated)" }; + const updateAlphaRead = () => { $("m-alpha-read").textContent = ALPHA_MEANING[$("m-alpha").value] || `α = ${$("m-alpha").value}`; }; + $("m-alpha").oninput = updateAlphaRead; // live label while dragging (montage only reloads on release via onchange) + updateAlphaRead(); + $("m-cell").onchange = () => LIVE() ? liveRefresh() : loadMontage(); + $("m-vs").onclick = () => { // Virtual staining toggle: tiles renderer, cell scrubbable over cells that have stained tiles (α scrubbable) + $("m-vs").classList.toggle("active"); + const vs = vsMode(); $("m-cell").disabled = false; + if (vs) { $("m-render").value = "tiles"; setRenderMode(); ovlSel = "__allmarkers__"; $("m-ovsel").value = overlayLabel(); } // default overlay = all markers on entering VS + else { ovlSel = "off"; $("m-ovsel").value = "off"; $("m-phaseoff").classList.remove("active"); } // leaving VS clears overlay + fillCellDropdown(); // VS → only cells with stained tiles; non-VS → full montage range + updateVsUI(); + loadMontage(); + }; + wireCombo("m-ovsel", "m-ovsel-list", renderOverlayList, overlayLabel); // searchable Overlay picker + $("m-phaseoff").onclick = () => { $("m-phaseoff").classList.toggle("active"); loadMontage(); }; // hide/show phase base + $("m-mode").onchange = () => { setMode($("m-mode").value); if (LIVE()) liveDraw(); }; + $("m-imgalpha").oninput = () => { mont.imgAlpha = +$("m-imgalpha").value; LIVE() ? liveDraw() : applyLayers(); }; + $("m-ovlalpha").oninput = () => { mont.ovlAlpha = +$("m-ovlalpha").value; applyLayers(); }; // stained-overlay opacity over phase + $("m-ptalpha").oninput = () => { mont.ptAlpha = +$("m-ptalpha").value; LIVE() ? liveDraw() : drawOverlay(); }; + $("m-tilesize").oninput = () => { mont.tileSize = +$("m-tilesize").value; liveDraw(); }; + $("m-detail").oninput = () => { + mont.detail = +$("m-detail").value; + if (mont.osd) { // apply to EVERY layer (phase base + VS overlay) so they stay in sync + mont.osd.minPixelRatio = mont.detail; + for (let i = 0; i < mont.osd.world.getItemCount(); i++) mont.osd.world.getItemAt(i).minPixelRatio = mont.detail; + mont.osd.forceRedraw(); + } + }; + wireCombo("m-color-search", "m-color-list", renderColorList, colorLabel); + $("m-cmap").onchange = () => { mont.cmapName = $("m-cmap").value; if (isFeatField()) { renderLegend(); drawOverlay(); if (LIVE()) liveDraw(); } }; + $("pc-tfidf").onchange = () => { pc.tfidf = $("pc-tfidf").checked; if (pc.cur) showPC(pc.cur); }; + const pcSetMode = (m) => { pc.mode = m; $("pc-mode-onto").classList.toggle("active", m === "onto"); $("pc-mode-feat").classList.toggle("active", m === "feat"); $("pc-feat-opts").style.display = m === "feat" ? "" : "none"; if (pc.cur) showPC(pc.cur); }; + $("pc-mode-onto").onclick = () => pcSetMode("onto"); + $("pc-mode-feat").onclick = () => pcSetMode("feat"); + $("pc-dedup").onchange = () => { pc.dedup = $("pc-dedup").checked; buildPCList(); if (pc.cur) showPC(pc.cur); }; + $("pc-norm").onchange = () => { pc.norm = $("pc-norm").checked; buildPCList(); if (pc.cur) showPC(pc.cur); }; + $("pc-sort").onchange = () => { pc.sort = $("pc-sort").value; buildPCList(); }; + $("tc-pin").onclick = () => { const g = state.target && state.target.target; // pin current gene in the current mode + if (g && !tc.pinned.some(p => p.gene === g && p.mode === tc.mode)) tc.pinned.push({ gene: g, mode: tc.mode }); renderTopPins(); renderTop(); }; + $("tc-pinclear").onclick = () => { tc.pinned = []; renderTopPins(); renderTop(); }; + $("tc-cols").onchange = () => $("tc-view").classList.toggle("cols-layout", $("tc-cols").checked); // top cells rows ↔ columns + $("tc-mask").onchange = () => { tc.mask = $("tc-mask").checked; $("tc-view").classList.toggle("masked", tc.mask); saveState(); }; // blue seg overlay on/off + $("tc-inorm").onchange = () => { tc.inorm = $("tc-inorm").checked; renderTop(); saveState(); }; // marker-global vs per-cell intensity (fluor) + $("tc-acc").onchange = () => { tc.showAcc = $("tc-acc").checked; ensureSetacc(renderTop); saveState(); }; // per-group set-accuracy chip + $("tc-accbin").onchange = () => { tc.accBin = +$("tc-accbin").value; renderTop(); saveState(); }; // classifier bag size the accuracy is measured at + $("m-labels").onchange = () => { mont.showLabels = $("m-labels").checked; drawOverlay(); }; + const onSetacc = () => { // off / geneKO / complex × P(target)|rank — per-tile v5 set-score at the montage α (v5 cache only) + mont.setaccMode = $("m-setacc").value; + mont.setaccMetric = $("m-setacc-metric").value; + const done = () => { + if (mont.setaccMode === "complex") { // color by EBI-complex, chip on the nearest-member dot + if (mont.prevField == null) mont.prevField = mont.field; + computeCxAnchors(); mont.field = "ebi_complex"; $("m-color-search").value = colorLabel(); setField(); + } else { // geneKO / off — restore the prior coloring if we'd switched it + if (mont.prevField != null) { mont.field = mont.prevField; mont.prevField = null; $("m-color-search").value = colorLabel(); setField(); } + else drawOverlay(); + } + }; + const need = mont.setaccMode === "geneKO" && mont.setaccMetric === "ptarget" && !mont.setacc ? ["setacc.json", j => mont.setacc = j] + : mont.setaccMode === "geneKO" && mont.setaccMetric === "rank" && !mont.setaccRank ? ["setacc_rank.json", j => mont.setaccRank = j] + : mont.setaccMode === "complex" && mont.setaccMetric === "ptarget" && !mont.setaccCx ? ["setacc_complex.json", j => mont.setaccCx = j] + : mont.setaccMode === "complex" && mont.setaccMetric === "rank" && !mont.setaccCxRank ? ["setacc_complex_rank.json", j => mont.setaccCxRank = j] : null; + if (need) fetch(`${BASE}_montage/${need[0]}${NOCACHE}`).then(r => r.ok ? r.json() : null).then(j => { need[1](j); done(); }).catch(done); + else done(); + }; + $("m-setacc").onchange = onSetacc; $("m-setacc-metric").onchange = onSetacc; + let _saveTimer; // persist selection/prefs after any change settles (debounced; snapshot reads live state) + const scheduleSave = () => { clearTimeout(_saveTimer); _saveTimer = setTimeout(saveState, 250); }; + document.addEventListener("change", scheduleSave); document.addEventListener("click", scheduleSave); + setupBagUI(); // anchor-bag selector (multi_bag default) before first marker selection + if (!restoreState()) { // restore last session, else default = phase marker + HSPA5 + selectMarker(0); $("markerfilter").value = markerLabel(0); + const defT = state.targets.find(t => t.target === "HSPA5"); + if (defT) { $("filter").value = targetLabel(defT); selectTarget(defT.slug); } + } + document.querySelectorAll("select[data-seg]").forEach(segmentize); // small dropdowns → segmented pills + document.querySelectorAll("label.chk").forEach(toggleize); // checkboxes → off/on segmented switches + if (IS_PUBLIC_DEPLOY) $("envtoggle").style.display = "none"; // the manual toggle is a staging-only preview aid + else $("envtoggle").onclick = () => { state.publicPreview = !state.publicPreview; applyFeatureGate(); }; + applyFeatureGate(); + updateScoreLegend(); // init the adaptive score-overlay legend (fresh loads default to ptarget) + requestAnimationFrame(() => $("loading").classList.add("gone")); // reveal the app only after first full render (no traversal flash) +} +// Turn a checkbox into an [off | ] segmented switch (like Ontology/Features). The native checkbox +// stays in the DOM (hidden) so existing .checked/.onchange logic is untouched; buttons mirror + drive it. +function toggleize(label) { + const inp = label.querySelector('input[type=checkbox]'); if (!inp || inp._tog) return; + const onLbl = inp.dataset.on || label.textContent.trim(); + const g = document.createElement("div"); g.className = "seg-group tog"; g.title = label.textContent.trim(); + const off = document.createElement("button"), on = document.createElement("button"); + off.type = on.type = "button"; off.className = on.className = "seg"; off.textContent = "off"; on.textContent = onLbl; + const sync = () => { off.classList.toggle("active", !inp.checked); on.classList.toggle("active", inp.checked); }; + off.onclick = () => { if (!inp.checked) return; inp.checked = false; inp.dispatchEvent(new Event("change", { bubbles: true })); sync(); }; + on.onclick = () => { if (inp.checked) return; inp.checked = true; inp.dispatchEvent(new Event("change", { bubbles: true })); sync(); }; + g.append(off, on); inp._tog = g; inp._togSync = sync; + label.style.display = "none"; label.after(g); sync(); +} +// Skin a + + + +
+ + +
+
+ +
+
+ + +
+ + +
+
+
+ +
+ + + +
+
+ +
+ +
+
+ + +
+ +
+
+ + + + +
+ + +
+
+
+
+
α=1 · true centroid
+
+
anti-phenotypeAnchor (α=0)phenotype
+
+
+
+ +
+ +
+
+ α = 0.0 +
+
+ + +
+ +
+
+
+
+
+
+
+
+ +
+
+ +
+
+ + + + + + diff --git a/src/ops_model/models/attention/diffex/viewer/webapp/methods.js b/src/ops_model/models/attention/diffex/viewer/webapp/methods.js new file mode 100644 index 0000000..97990e2 --- /dev/null +++ b/src/ops_model/models/attention/diffex/viewer/webapp/methods.js @@ -0,0 +1,406 @@ +// "How it works" tab — a slide deck of the ML methods behind the viewer, for a general audience but rigorous. +// Each slide: an animated SVG, a short body, a "why it matters" line, explicit paper links, and a "Key terms" glossary. +// Animations are CSS keyframes (style.css .mth-*) driven by group transforms + opacity so they loop while shown. + +const MTH_C = { acc: "#26c6ff", ntc: "#8b949e", ko: "#f0a020", grn: "#3fb950", pur: "#bc8cff", yel: "#d29922", fg: "#e6e8ec" }; + +// ---- SVG builders ---- +const _cellBlob = (cx, cy, r, fill, cls = "") => + `` + + `` + + ``; +const _bars = (x, y, vals, w, gap, col, cls = "mth-rise") => vals.map((v, i) => + ``).join(""); +const _arrow = (x1, x2, y, col = "#26c6ff") => + `` + + ``; +const _box = (x, y, w, h, l1, l2 = "") => + `` + + `${l1}` + + (l2 ? `${l2}` : ""); +const _lbl = (x, y, t, col = "#8b949e", sz = 11) => `${t}`; +const _mix = (a, b, t) => { const p = h => [1, 3, 5].map(i => parseInt(h.slice(i, i + 2), 16)); const A = p(a), B = p(b); return `rgb(${A.map((v, i) => Math.round(v + (B[i] - v) * t)).join(",")})`; }; // hex→hex lerp +const _patchGrid = (x, y, s, n, hot) => Array.from({ length: n * n }, (_, i) => { + const gx = x + (i % n) * s, gy = y + Math.floor(i / n) * s, on = hot.includes(i); + return ``; +}).join(""); +// shared phenotype-cell presets (used by Embedding + The Screen): membrane aspect (ar), nucleus fraction (nf), organelle offsets (frac of s) +const MTH_PHENO = [ + { ar: 1.0, nf: 0.24, org: [[-0.27, -0.33]] }, + { ar: 1.27, nf: 0.38, org: [[0.27, 0.33]] }, + { ar: 0.82, nf: 0.5, org: [[0.27, -0.4], [-0.27, 0.4]] }, + { ar: 0.72, nf: 0.3, org: [[-0.3, -0.18], [0.32, 0.26]] }, +]; +const _pheno = (cx, cy, s, p, fill) => `` + + p.org.map(o => { const ox = cx + o[0] * s, oy = cy + o[1] * s; return ``; }).join(""); +// low-res black&white "static" — a grid of grayscale pixels (deterministic per seed), like the paper schematic +const _pixnoise = (x, y, w, h, cols, rows, seed, alpha) => { const cw = w / cols, ch = h / rows; let s = ""; + for (let r = 0; r < rows; r++) for (let c = 0; c < cols; c++) { const k = r * cols + c, v = Math.abs(Math.sin((k + 1) * 12.9898 + seed * 78.233)) * 43758.5453, g = Math.floor((v - Math.floor(v)) * 255); + s += ``; } + return s; }; + +const MTH_REFS = { + "The screen": [["Optical pooled screens · Feldman 2019", "https://doi.org/10.1016/j.cell.2019.09.016"], ["Funk et al. · Cell 2022", "https://www.sciencedirect.com/science/article/pii/S0092867422013599"], ["Liu et al. · bioRxiv 2026", "https://www.biorxiv.org/content/10.64898/2026.06.01.728087v1"]], + "Embedding": [["Cell-DINO · Moutakanni · PLOS Comput Biol 2025", "https://journals.plos.org/ploscompbiol/article?id=10.1371/journal.pcbi.1013828"], ["DINO · Caron 2021", "https://arxiv.org/abs/2104.14294"], ["DINOv2 · Oquab 2023", "https://arxiv.org/abs/2304.07193"]], + "Classifier": [["Set Transformer · Lee 2019", "https://arxiv.org/abs/1810.00825"], ["Attention-based multiple-instance learning · Ilse 2018", "https://arxiv.org/abs/1802.04712"]], + "Top cells": [["Explaining by removing · Covert 2021", "https://arxiv.org/abs/2011.14878"], ["SHAP · Lundberg & Lee 2017", "https://arxiv.org/abs/1705.07874"]], + "Diffusion": [["DDPM · Ho 2020", "https://arxiv.org/abs/2006.11239"], ["Diffusion autoencoders · Preechakul 2022", "https://arxiv.org/abs/2111.15640"], ["DDIM · Song 2020", "https://arxiv.org/abs/2010.02502"]], + "Traversal": [["DiffEx · Bourou 2025", "https://arxiv.org/abs/2502.09663"], ["Diffusion autoencoders · Preechakul 2022", "https://arxiv.org/abs/2111.15640"], ["Classifier-free guidance · Ho & Salimans 2022", "https://arxiv.org/abs/2207.12598"]], + "DDIM": [["DDIM · Song 2020", "https://arxiv.org/abs/2010.02502"]], + "Attention heads": [["DINO attention · Caron 2021", "https://arxiv.org/abs/2104.14294"], ["Attention is all you need · Vaswani 2017", "https://arxiv.org/abs/1706.03762"]], + "Montage": [["UMAP · McInnes 2018", "https://arxiv.org/abs/1802.03426"], ["PHATE · Moon 2019", "https://doi.org/10.1038/s41587-019-0336-3"]], + "Virtual staining": [["In silico labeling · Christiansen 2018", "https://doi.org/10.1016/j.cell.2018.03.040"], ["Diffusion autoencoders · Preechakul 2022", "https://arxiv.org/abs/2111.15640"]], + "mRNA phenotypes": [["CROP-seq · Datlinger 2017", "https://doi.org/10.1038/nmeth.4177"], ["Perturbation autoencoder (CPA) · Lotfollahi 2023", "https://doi.org/10.15252/msb.202211517"]], +}; + +// longer "Learn more" paragraph per slide (keyed by nav) — unpacks the concept for readers who want depth +const MTH_MORE = { + "The screen": "A traditional CRISPR screen gives one number per perturbation — did cells grow, did a reporter switch on. An optical pooled screen instead pools cells carrying many different perturbations into one dish and images them together in place, so a single experiment reads out thousands of knockouts side by side. The trick: each cell also carries a short DNA barcode unique to its CRISPR guide, and a round of in-situ sequencing lights that barcode up letter-by-letter right in the microscope — so we can match every cell's image to the exact geneKO. One imaging run yields millions of (perturbation, picture) pairs. The hard part is what comes next: those cells are a mix of thousands of perturbations, each present in many copies that vary enormously for reasons unrelated to the knockout — cell cycle, size, local density, position on the plate — and most knockouts nudge the phenotype only slightly. So the central problem is signal-in-noise: for every perturbation, which cells and which features actually capture its distinctive effect, rather than the technique's inherent variability? The classifier, rankings, and generative traversals in this viewer are complementary tools for answering exactly that.", + "Embedding": "A raw microscope image is hundreds of thousands of pixels — too many, and too noisy, to compare directly. We need a compact summary that keeps what's biologically meaningful (shape, texture, organelle layout) and drops what isn't (exact position, lighting). CellDINO is a vision transformer trained by self-supervision: shown only images, never perturbation labels, it learns embeddings where two crops of the same cell agree and different cells differ. The result is a 1,024-number fingerprint per cell whose distances track real morphological similarity — which is what lets us cluster phenotypes, rank perturbations, and steer the generative model later.", + "Classifier": "A single knocked-out cell is often ambiguous — cells vary a lot even with no perturbation — so the signal lives in the distribution of a perturbation's cells. We therefore classify a whole bag at once (multiple-instance learning). The SetTransformer tags each cell's fingerprint with its imaging channel, uses attention so the cells in a bag can inform one another, then pools them into a single bag-vector with attention (PMA) rather than a plain average, letting it weight the informative cells. Because attention and pooling ignore order, shuffling the bag can't change the answer; because it trained across bag sizes it can score anywhere from 10 to thousands of cells. The output is a probability over 1,000 perturbations (+NTC), or over 99 protein complexes when we ask about pathways instead of single geneKOs.", + "Top cells": "Once the classifier works, we ask which individual cells it actually relies on, borrowing an idea from model explanation called \"explaining by removing\": a cell is important if taking it out of a bag makes the classifier less sure of the right answer. Concretely we take the probability of the correct class with the cell in the bag minus the probability without it, and average that marginal contribution over many random bags and bag sizes (a Monte-Carlo estimate). The highest-scoring cells are the clearest examples of a perturbation's phenotype — these top-predictive cells are exactly what the Top Cells tab shows and what the diffusion traversals are anchored to. Up-weighting them also sharpens the perturbation-level distinctiveness score (mAP).", + "Diffusion": "The purpose of the diffusion model here is to be a decoder for CellDINO space: given a cell's CellDINO vector, produce the image that vector describes. That is what makes the fingerprint editable — nudge the vector and the decoded cell changes. It learns this by reversing a corruption process: during training we add Gaussian noise to a real cell over and over until it is pure static, and a network learns to predict — from a noisy image plus the cell's CellDINO vector — the noise that was added. To generate, we start from static and repeatedly subtract the predicted noise (steered by the vector) until a realistic cell condenses out. By default the starting noise is random, so a given vector decodes to a representative cell, not one particular real cell — and a morph from a generic cell is hard to trust. DDIM (Denoising Diffusion Implicit Models) fixes this. The \"implicit\" part means it denoises along a single fixed trajectory — a deterministic ODE rather than a random Markov walk — so the same seed always yields the same cell. Because it is deterministic it can be integrated in reverse: from a real cell's pixels, run the trajectory backwards to recover the exact noise seed that regenerates it. We invert under the same guidance used for generation (\"guided inversion\"), so decoding at α = 0 reproduces the original cell almost perfectly (pixel correlation ≈ 0.99). Every traversal therefore starts anchored to a genuine cell.", + "Traversal": "The semantic code lives in a space where nearby points are similar-looking cells and directions correspond to consistent visual changes. To build a knockout's \"movie\" we average the codes of control (NTC) cells and of that knockout's cells; the vector between them is the direction that turns control-looking into knockout-looking. Starting from one cell's code we step along it — α = 0 is the cell itself, α = 1 applies the full average shift, larger α exaggerates subtle effects — decoding an image at each step. Because the noise seed is held fixed throughout, only the phenotype moves: you're watching the same cell change, not a slideshow of different cells.", + "DDIM": "By default the noise seed is random, so decoding a code gives a representative cell, not any particular real one — and a traversal from a generic cell is hard to trust. DDIM fixes this two ways. First it makes generation deterministic: the same seed always yields the same image (a diffusion \"ODE\", not a random walk). Second, being deterministic, it can run in reverse — starting from a real cell's pixels and integrating backwards to recover the exact noise seed that regenerates it. We invert under the same guidance used to generate (\"guided inversion\"), so decoding at α = 0 reproduces the original cell almost perfectly (pixel correlation ≈ 0.99). Every traversal here therefore begins anchored to a genuine cell, and the changes you see are real counterfactuals, not artifacts of a random start.", + "Attention heads": "The vision transformer doesn't read pixels one at a time — it breaks the image into a grid of patches and, in each attention \"head\", decides how much every patch should influence its summary of the cell. Reading those weights back out gives a heat-map showing where the model concentrated. Different heads specialize on different structures, and the viewer lets you step through them and compare a perturbation against its control. The payoff is interpretability: instead of only telling you that a perturbation has a distinctive phenotype, the map shows you where the model is looking — e.g. a head that consistently lights up mitochondria for a mitochondrial geneKO — which you can check against the biology.", + "Montage": "A single traversal shows one perturbation's effect; the Montage tab shows all of them at once — and controls for cell-to-cell variation by using the same anchor cell throughout. We morph that one cell toward each of the ~1,000 knockouts, giving 1,000 counterfactual images of the same starting cell. Each image is then positioned by its perturbation's coordinates on a UMAP or PHATE embedding of the perturbation-level phenotype space, so knockouts that produce similar morphologies land near one another. LatentLens stitches the crops into a continuous, zoomable atlas — pan and zoom to compare neighborhoods, spot phenotype clusters, and see where a perturbation of interest falls relative to the whole library.", + "Virtual staining": "Fluorescent markers reveal specific structures but cost extra dyes, channels, and imaging, and you can only stain a few at once. Label-free phase imaging is cheap and gentle but hard to read. Virtual staining bridges them: we reuse the diffusion autoencoder and condition it on a phase image in two complementary ways. The semantic path sends the phase image through a frozen Cell-DINO ViT to a pooled code that, together with a learned marker id, globally modulates the U-Net (FiLM) — telling it what to draw. The spatial path concatenates the raw phase image as an extra U-Net input channel, so generation keeps the input's pixel layout and the predicted marker stays registered to the real cell (this is what lifts fidelity from ~0.13 to ~0.78 Pearson). One model, trained on paired (phase, marker) crops, covers all 42 live markers; switching the marker id renders any channel, and applied to a traversal it gives a full multi-channel phenotype for every perturbation from a single grayscale image.", + "mRNA phenotypes": "So far every direction has been morphological — defined in the CellDINO image-fingerprint space. But the same knockout library was also profiled by CROP-seq, which reads each cell's transcriptome (its gene-expression profile) instead of its picture, giving every perturbation a transcriptional signature (its pseudobulk expression shift versus NTC). How would we drive the diffusion decoder from that? Two designs. The light one reuses everything: fit a map — linear first, then a small MLP — from a perturbation's transcriptional shift to its shift in CellDINO space, then feed that predicted CellDINO direction into the existing decoder. No retraining, and the map's R² directly measures how much of morphology is even predictable from transcriptome. The full one conditions the diffusion model on the transcriptome directly, in the style of a compositional perturbation autoencoder (CPA): a single per-perturbation embedding drives both a transcriptome decoder (reconstructing the CROP-seq profile) and the image decoder, so the two modalities are tied through one shared latent and any transcriptional state — even an unseen combination — can be rendered. Either way the payoff is the same: a transcriptome↔morphology divergence map that flags the perturbations that reshape the transcriptome but barely change the image, or vice versa — decoupling that is invisible to either assay alone.", +}; + +const METHODS_SLIDES = [ + { + nav: "The screen", kicker: "THE QUESTION", title: "Thousands of knockouts, mixed in one noisy dish", + svg: () => { const COL = [MTH_C.grn, MTH_C.acc, MTH_C.ko, MTH_C.pur, MTH_C.yel]; const cells = []; + for (let r = 0; r < 6; r++) for (let c = 0; c < 5; c++) { const x = 30 + c * 30 + ((r * 37) % 9 - 4), y = 40 + r * 24 + ((c * 53) % 9 - 4); + if (((x - 92) / 74) ** 2 + ((y - 100) / 80) ** 2 <= 0.92) cells.push([x, y, (r * 2 + c) % 5]); } + return ` + + ${cells.map(([x, y, ci]) => { const r = ci === 1 ? 8.75 : 7; + return `${_cellBlob(x, y, r, COL[ci])}`; }).join("")} + ${_lbl(92, 194, "well — pooled knockouts, all mixed")} + ${_arrow(176, 200, 100)} + ${[0, 1, 2, 3, 4, 5, 6].map(i => ``).join("")} + debarcode + ${_arrow(242, 270, 100)} + + one knockout's cells + ${[[308, 88], [370, 88], [308, 126], [370, 126]].map((p, i) => `${_pheno(p[0], p[1], 15, MTH_PHENO[i], MTH_C.acc)}`).join("")} + many phenotypes —which is real? + `; }, + body: "In a pooled optical CRISPR screen, thousands of geneKOs are mixed in one dish and imaged together; each cell's DNA barcode, sequenced in place, names the geneKO inside it. This yields millions of (perturbation, image) pairs — but each knockout's real effect is subtle and buried in enormous cell-to-cell variation.", + why: "The core question this whole viewer answers: for each perturbation, which change in the cell captures its true phenotype — separated from the noise of the technique's scale and heterogeneity? Everything that follows is one answer.", + defs: [["Pooled optical screen", "imaging a mixed population where every cell has a different geneKO, all together."], + ["Barcode", "a short DNA tag, read out in situ, that identifies which CRISPR guide (perturbation) is in each cell."], + ["Perturbation", "the genetic change applied to a cell — here a CRISPR knockout; NTC = non-targeting control (no geneKO)."]] + }, + { + nav: "Embedding", kicker: "REPRESENT", title: "Turning a cell image into numbers (a ViT)", + svg: () => { const A = MTH_C.acc, P = MTH_C.pur, gx = 14, gy = 72, gs = 15, gn = 4; + // three discrete phenotype states — cell (shared MTH_PHENO presets), active patches, tokens, and the 1024-d vector all switch together + const HOTC = [[2, 5, 9], [1, 6, 10, 13], [0, 4, 7, 11, 14]]; + const HOTT = [[0, 2, 4], [1, 4, 5], [0, 3, 4]]; + const VEC = [[58, 26, 72, 40, 54, 20, 66, 34], [30, 64, 20, 80, 44, 60, 28, 72], [70, 18, 50, 34, 78, 26, 42, 60]]; + const lattice = Array.from({ length: gn + 1 }, (_, k) => ``).join(""); + const hotC = hot => hot.map(i => ``).join(""); + const tokens = hot => [0, 1, 2, 3, 4, 5].map(i => ``).join(""); + const stateG = s => `${_pheno(44, 100, 30, MTH_PHENO[s], A)}${hotC(HOTC[s])}${lattice}${tokens(HOTT[s])}${_bars(352, 136, VEC[s], 11, 4, P, "")}`; + return ` + ${[0, 1, 2].map(stateG).join("")} + ${_lbl(44, 150, "image crop")} + patchify + ${_arrow(82, 128, 100)} + ${_lbl(143, 178, "patch tokens")} + ${_arrow(158, 200, 100)} + ${_box(204, 74, 98, 52, "Transformer", "self-attention")} + ${_arrow(306, 346, 100)} + ${_lbl(400, 156, "1024-d feature")} + `; }, + body: "A vision transformer (CellDINO) cuts the image into a grid of patches, turns each into a token, and lets them attend to one another; the tokens are pooled into one 1,024-d feature vector — a compact fingerprint of the cell's morphology. It's trained self-supervised (no labels), so similar cells get similar vectors.", + why: "Turning each cell into a comparable vector is what makes morphology measurable — the basis for clustering, ranking, and steering the generative model.", + defs: [["Patch / token", "the small square pieces the image is cut into; each becomes one input token to the transformer."], + ["Vision transformer (ViT)", "a network that relates all patch-tokens with attention, rather than scanning with convolutions."], + ["Attention", "a weighted lookup between tokens: each patch emits a query and a key, their match sets a weight, and the patch is updated as a weighted sum of the others' values. High weight = \"this patch is relevant to me,\" so the update pulls in information from wherever in the cell matters most."], + ["Self-attention", "attention run within one image — queries, keys, and values all come from the same set of patch-tokens, so every patch is refined by the whole cell's context at once. That lets the ViT link distant structures (e.g. a nucleus and a far-off organelle) in a single step, which convolutions can't do locally."], + ["Self-supervised (DINO)", "trained on images alone — no perturbation labels — so it learns general-purpose morphology features."], + ["Feature vector / embedding", "the pooled 1,024 numbers summarizing the cell; distances between vectors track visual similarity."]] + }, + { + nav: "Classifier", kicker: "THE MODEL", title: "A classifier that reads a whole group of cells", + svg: () => { const X = [118, 205, 290, 392], L0 = [46, 82, 118, 154], LH = [64, 100, 136], OY = [38, 72, 106, 140, 174], OC = [MTH_C.grn, MTH_C.acc, MTH_C.ko, MTH_C.pur, MTH_C.yel]; + const edges = (xa, ya, xb, yb, d) => ya.map(a => yb.map(b => ``).join("")).join(""); + const set = (col, by, hl) => `${[[16, by], [30, by], [16, by + 14], [30, by + 14]].map(c => _cellBlob(c[0], c[1], 5.5, col)).join("")}`; + return ` + ${[MTH_C.grn, MTH_C.acc, MTH_C.ko, MTH_C.pur].map((col, s) => set(col, 30 + s * 40, col === MTH_C.acc)).join("")} + CellDINO vectors(one per cell; a set = one perturbation) + ${_arrow(46, 108, 100, MTH_C.acc)} + ${edges(X[0], L0, X[1], LH, 0.25)}${edges(X[1], LH, X[2], LH, 0.75)}${LH.map(a => OY.map((b, j) => ``).join("")).join("")} + ${L0.map(y => ``).join("")} + ${LH.map(y => ``).join("")} + ${LH.map(y => ``).join("")} + ${OY.map((y, i) => ``).join("")} + ${_lbl(118, 18, "encode + channel-embed", "#8b949e", 7.5)} + ${_lbl(247, 18, "ISAB ×2 · inducing-point attention", "#8b949e", 7.5)} + ${_lbl(392, 18, "PMA pool → cosine", "#8b949e", 7.5)} + ${_lbl(392, 196, "class scores (perturbations)")} + ← predicted + `; }, + body: "Because single cells are noisy, we show the model a whole group of a perturbation's cells at once — each as its CellDINO feature vector (from the Embedding tab), not the raw image — and ask it to name the perturbation. It weighs the cells against each other, pools them into one verdict, and outputs a probability over the 1,000 perturbations (or 99 protein complexes). The architecture is a SetTransformer — the how is below.", + why: "Trained on random bags of 100 cells yet able to score any bag size (10–5,000), it reads the population phenotype and is invariant to how the cells are ordered.", + defs: [["Multiple-instance learning", "classify a whole bag from one shared label, without labeling individual cells."], + ["Bag / permutation-invariant", "an unordered set of a perturbation's cells; shuffling them can't change the prediction."], + ["ISAB (inducing-point attention)", "self-attention routed through 32 learned reference points, so cost grows linearly and scales to tens of thousands of cells."], + ["Channel embedding", "a learned tag marking each cell's imaging channel (phase, or a given fluorescent marker)."], + ["PMA + cosine classifier", "pooling-by-attention collapses the bag to one vector; a cosine classifier turns it into class probabilities."]] + }, + { + nav: "Top cells", kicker: "EXPLAINING BY REMOVING", title: "Which cells carry the phenotype?", + svg: () => ` + + ${_cellBlob(34, 82, 9, MTH_C.acc)}${_cellBlob(58, 82, 9, MTH_C.acc)}${_cellBlob(34, 116, 9, MTH_C.acc)} + ${_cellBlob(58, 116, 12.5, MTH_C.ko)} + ${_lbl(47, 170, "bag · cell x")} + ${_arrow(84, 140, 100)} + ${[64, 100, 136].flatMap(a => [64, 100, 136].map(b => ``)).join("")} + ${[64, 100, 136].map(y => ``).join("")} + ${[64, 100, 136].map(y => ``).join("")} + ${_lbl(171, 28, "classifier", "#8b949e", 8)} + ${_arrow(204, 240, 100)} + accuracy score = P(perturbation X) + + + + + Δ = score(x) + ${_lbl(285, 48, "with x", MTH_C.ko, 9)}${_lbl(285, 156, "without x", "#8b949e", 9)} + ${_arrow(334, 368, 100)} + ${(() => { let cx = 372; return [0, 1, 2, 3, 4].map(i => { const top = i === 4; const r = 4.5 + i * 1.1; if (i > 0) cx += (4.5 + (i - 1) * 1.1) + r + 5; return ``; }).join(""); })()} + ${_lbl(404, 128, "higher rank →")} + `, + body: "To find a perturbation's most telling cells, we score each cell by how much it helps the classifier: the drop in predicted probability when the cell is removed from a bag (\"explaining by removing\"), averaged over many bag sizes and random partners. The top-scoring top-predictive cells carry the phenotypic signature — the cells the viewer anchors its traversals to and shows in Top Cells. Re-weighting the perturbation-level mAP by these cells sharpens the distinctiveness ranking.", + why: "It picks, per perturbation, the handful of cells that most define its phenotype — the exemplars every traversal starts from.", + defs: [["Explaining by removing", "gauge a cell's importance by how much the prediction drops when you take it out of the bag."], + ["Marginal contribution", "score(x) = p(class | bag with x) − p(class | bag without x), averaged over many bags (sizes 1–500)."], + ["Top-predictive cells", "the highest-scoring cells for a class; used as traversal anchors and in the Top Cells tab."], + ["Distinctiveness (mAP)", "the perturbation-level separability score (Fig. 2), recomputed with these top cells up-weighted."]] + }, + { + nav: "Diffusion", kicker: "DECODE & INVERT", title: "Turning a CellDINO vector back into a cell", + sections: [ + { text: "We condition a diffusion model on a cell's CellDINO vector (the same fingerprint from the Embedding tab). Starting from noise it removes a little at each step, steered by that vector, until it produces the exact cell the vector describes — so the model becomes a decoder for CellDINO space.", + svg: () => ` + ${_bars(10, 80, [22, 11, 30, 15, 24, 13], 7, 3, MTH_C.pur, "mth-jit")} + CellDINO z + ${_arrow(76, 96, 56, MTH_C.pur)} + ${[0, 1, 2, 3].map(i => { const x = 98 + i * 90, op = i / 3; + return `${_pixnoise(x + 4, 30, 58, 58, 12, 12, i + 1, (1 - op * 0.92).toFixed(2))}${_cellBlob(x + 33, 59, 21, MTH_C.acc)}`; }).join("")} + + ${_lbl(268, 16, "z conditions the denoising — z paints its cell →", MTH_C.acc, 10)} + ${_lbl(131, 108, "noise")}${_lbl(401, 108, "the cell z describes")} + ` }, + { text: "DDIM = Denoising Diffusion Implicit Models. \"Implicit\" means it denoises along one fixed, deterministic path (an ODE) instead of a random walk — so the same seed always gives the same cell. Being deterministic, it can also run in reverse: from a real cell's pixels it recovers the exact noise seed that regenerates it, so decoding at α = 0 reproduces that cell (r ≈ 0.99).", + svg: () => ` + ${_cellBlob(70, 90, 32, MTH_C.acc)}${_lbl(70, 146, "real cell")} + + ${_pixnoise(185, 69, 50, 42, 12, 10, 7, 0.95)} + ${_lbl(210, 146, "its exact seed x_T")} + ${_cellBlob(350, 90, 32, MTH_C.acc)}${_lbl(350, 146, "generated cell (r≈0.99)")} + ${_arrow(112, 176, 78, MTH_C.pur)}${_lbl(144, 60, "invert ↩", MTH_C.pur, 10)} + ${_arrow(244, 310, 78, MTH_C.acc)}${_lbl(277, 60, "generate →", MTH_C.acc, 10)} + ` } + ], + body: "Why a diffusion model? We want to turn a CellDINO vector back into an image — a generative decoder conditioned on the fingerprint — so we can then nudge the vector and watch the cell change (that's the Traversal). It learns this by reversing noise: corrupt a real cell to static, then learn to undo it step by step, steered by that cell's CellDINO vector. DDIM makes the reverse deterministic, so it also runs backwards to recover a real cell's exact seed and anchor the edit.", + why: "It makes CellDINO space visual and editable: any fingerprint — real or shifted — becomes a cell you can see, which is what turns a number (a distinctiveness score, a direction) into a watchable phenotype.", + defs: [["Decoder for CellDINO", "the diffusion model is trained to reconstruct a cell from its CellDINO vector, so it maps the fingerprint space back to images — the inverse of the Embedding step."], + ["Conditioning on z", "the vector steers every denoising step; change the vector (e.g. along a knockout direction) and the decoded cell changes to match."], + ["Forward / reverse process", "forward adds Gaussian noise to a real cell until it is pure noise (x_T); the reverse network predicts and removes that noise, guided by z, to recover a cell (x_0)."], + ["DDIM, word by word", "Denoising (removes noise) · Diffusion (the noise process) · Implicit (it follows one fixed, non-random path — an ODE — rather than a random walk) · Models. Upshot: deterministic sampling (same seed → same cell), which makes it both invertible and much faster (fewer steps)."], + ["Inversion (encoding)", "running DDIM backwards to recover the exact noise seed of a specific real cell, so a traversal can begin from it (guided inversion → α = 0 reconstructs it, r ≈ 0.99)."]] + }, + { + nav: "Traversal", kicker: "COUNTERFACTUAL", title: "\"How would this cell look if we applied a given perturbation?\"", + svg: () => ` + ${_cellBlob(210, 86, 42, MTH_C.acc)} + + + + ${_lbl(70, 180, "NTC · α0")}${_lbl(350, 180, "knockout · α+", MTH_C.ko)} + ${(() => { const zA = [8, 4, 12, 6, 10, 5], zB = [11, 7, 8, 9, 7, 8]; return "" + zA.map((a, i) => ``).join("") + ""; })()}${_lbl(211, 40, "z (CellDINO vector)", MTH_C.pur, 9)} + `, + body: "We take the cell's semantic code and slide it along the NTC → knockout direction (the average difference between control and knockout codes), decoding each step. α = 0 is the start; α = 1 applies the full knockout shift; beyond exaggerates it.", + why: "It renders the phenotype a perturbation induces as a smooth, watchable transformation of one cell.", + defs: [["Semantic direction", "in the identity-code space (the \"semantic code\", see Diffusion), the vector pointing from control (NTC) toward a knockout — a mean difference."], + ["α (alpha)", "how far we push along that direction — 0 = start, 1 = full shift, |α|>1 = extrapolation."], + ["Counterfactual", "a generated \"what this cell would look like if…\" image — not an observed one."], + ["Classifier-free guidance (w)", "a strength knob controlling how firmly the code steers the generated image."]] + }, + { + nav: "DDIM", kicker: "ANCHOR TO A REAL CELL", title: "DDIM — running the model backwards", + svg: () => ` + ${_cellBlob(70, 92, 34, MTH_C.acc)}${_lbl(70, 150, "real cell")} + ${Array.from({ length: 22 }, (_, i) => ``).join("")} + ${_lbl(210, 150, "its exact seed x_T")} + ${_cellBlob(350, 92, 34, MTH_C.acc)}${_lbl(350, 150, "generated cell (r≈0.99)")} + ${_arrow(110, 178, 78, MTH_C.pur)}${_lbl(144, 68, "invert ↩", MTH_C.pur, 10)} + ${_arrow(242, 312, 78, MTH_C.acc)}${_lbl(277, 68, "generate →", MTH_C.acc, 10)} + `, + body: "DDIM (Denoising Diffusion Implicit Models) makes generation deterministic — same seed, same cell, every time. Because it's deterministic it can also run backwards: given a real cell it recovers the exact noise seed that regenerates it. So the traversal starts from a true cell (α = 0 reconstructs it, pixel r ≈ 0.99).", + why: "The morph is anchored to a real cell's identity — its size, texture, and context are preserved while only the phenotype moves.", + defs: [["DDIM", "Denoising Diffusion Implicit Models — a deterministic way to sample a diffusion model (same training, no randomness at generation)."], + ["Deterministic", "same input always gives the same output — no dice-rolling — which is what makes it reversible."], + ["Inversion (encoding)", "running the model backwards to find the exact noise seed of a specific real image."], + ["Guided inversion", "inverting under the same guidance used for generation, so α = 0 reconstructs the cell faithfully."]] + }, + { + nav: "Attention heads", kicker: "INTERPRET", title: "Where does the model look?", + svg: () => { const cell = (cx, hot, lbl) => ` + + + + ${[[cx + 16, 80, 25, "M -11 0 q 5 -8 10 -1 q 6 7 12 -1"], [cx + 24, 100, -40, "M -10 1 q 6 -7 11 0 q 4 6 10 -3"], [cx + 16, 116, 60, "M -12 -1 q 4 7 9 1 q 6 -6 12 1"], [cx + 30, 120, -12, "M -9 0 q 7 -6 12 1 q 3 6 9 -2"]].map(([x, y, rot, d]) => ``).join("")} + ${_patchGrid(cx - 42, 56, 14, 6, hot)} + ${_lbl(cx, 186, lbl)}`; + return ` + ${cell(112, [16, 17, 22, 23], "head A — tubules")} + ${cell(330, [7, 8, 13, 14], "head B — nucleoli")} + `; }, + body: "The vision transformer splits the cell into patches, and each attention head weights which patches it focuses on. Drawn as a heat-map, a head shows which structures the model attends to for a given cell — for example one that concentrates on mitochondria.", + why: "It makes the model's focus visible and checkable against known biology, rather than leaving the phenotype call as a black box.", + defs: [["Patch / token", "the small square pieces a vision transformer breaks the image into."], + ["Attention head", "one of several parallel attention 'spotlights'; each weights which patches to focus on, and different heads specialize."], + ["Attention map", "a head's per-patch weights, drawn as a heat-map — where the model concentrates (an attention pattern, suggestive of but not a formal attribution of the decision)."]] + }, + { + nav: "Montage", kicker: "THE MAP", title: "Every knockout, from one cell, on one map", + svg: () => { const cx = 150, cy = 98, B = [[MTH_C.acc, -2.35], [MTH_C.ko, -0.75], [MTH_C.grn, 0.55], [MTH_C.pur, 2.0]]; + // each arm = a different phenotype axis: 0 acc grows · 1 ko elongates · 2 grn (bottom) shrinks · 3 pur nucleus grows + // tt = 0 at the first cell (≈ the original NTC cell) → 1 at the arm tip (full phenotype); gradual divergence + const mcell = (x, y, b, i, col, ang) => { const tt = (i - 1) / 4, lp = (a, z) => a + (z - a) * tt, hue = _mix(MTH_C.ntc, col, tt), dark = _mix(col, "#000000", .3); + let rx, ry, nr; + if (b === 0) { rx = ry = 5 + i * 1.2; nr = rx * 0.3; } // acc — grows (tuned to arm spacing so cells don't overlap) + else if (b === 1) { rx = lp(11, 15.5); ry = lp(11, 7); nr = ry * 0.34; } // ko — elongates (never smaller) + else if (b === 3) { rx = ry = 11; nr = lp(3.2, 8.6); } // pur — nucleus grows + else { rx = ry = lp(11, 5); nr = rx * 0.3; } // grn (bottom) — shrinks + const ox = x + rx * 0.32, oy = y + ry * 0.3; + return ``; }; + let out = ""; + B.forEach(([col, a0], b) => { for (let i = 1; i < 6; i++) { const r = 14 + i * 26, a = a0 + i * 0.15, x = cx + r * Math.cos(a), y = cy + r * Math.sin(a) * 0.66; out += `${mcell(x, y, b, i, col, (a * 57).toFixed(0))}`; } }); + return ` + ${out} + ${_cellBlob(cx, cy, 11, MTH_C.ntc)} + ${_lbl(cx, cy + 26, "Single")}${_lbl(cx, cy + 38, "Control")}${_lbl(cx, cy + 50, "Cell")} + ${_lbl(96, 14, "Cluster A", MTH_C.acc, 10)}${_lbl(332, 100, "Cluster B", MTH_C.ko, 10)} + + ${_lbl(392, 36, "each tile = a gene-KO", "#8b949e", 9)} + ${_lbl(230, 194, "each arm = a different phenotype axis (size · shape · nucleus); distance = how different", "#8b949e", 9.5)} + `; }, + body: "The Montage tab takes a single anchor cell, traverses it toward every one of the ~1,000 perturbations, and drops each morphed cell at that perturbation's spot on a perturbation-similarity map (UMAP/PHATE). LatentLens tiles thousands of these crops into one zoomable montage.", + why: "It turns 1,000 separate what-ifs into a single navigable landscape — perturbations with similar phenotypes cluster together, visible at a glance.", + defs: [["Perturbation embedding (UMAP / PHATE)", "a 2-D map where each point is a perturbation, placed so phenotypically similar knockouts sit close together."], + ["Anchor cell", "the one real cell (see DDIM) whose counterfactual we render for every perturbation, so the comparison is apples-to-apples."], + ["LatentLens", "the tiling engine that lays thousands of image crops onto the map as a smooth, zoomable montage."]] + }, + { + nav: "Virtual staining", kicker: "CROSS-CHANNEL", title: "Predicting fluorescent stains from phase", + svg: () => { const M = [["mitochondria", MTH_C.ko], ["ER", "#2ca089"], ["nucleus", MTH_C.acc], ["actin", MTH_C.pur], ["lysosome", MTH_C.yel]]; + const comp = (cx, cy, col, kind, s) => kind === "mitochondria" ? [["M -7 0 q 3 -5 6 -1 q 4 5 8 -1", -3, -3, 25], ["M -6 1 q 4 -5 7 0 q 3 4 7 -2", 2, 1, -40], ["M -8 -1 q 3 5 6 1 q 4 -4 8 1", -1, 4, 60]].map(([d, dx, dy, rot]) => ``).join("") + : kind === "ER" ? [0, 1, 2, 3].map(k => ``).join("") + : kind === "nucleus" ? `` + : kind === "actin" ? [0, 1, 2].map(k => ``).join("") + : [0, 1, 2, 3, 4].map(k => ``).join(""); + return ` + ${_cellBlob(42, 100, 28, MTH_C.ntc)} + ${M.map((m, i) => `${comp(42, 100, m[1], m[0], 1.5)}`).join("")} + ${_lbl(42, 146, "phase cell")} + ${_arrow(78, 130, 100)} + + diffusion staining + model + semantic (what) + + spatial (where) + ${_arrow(264, 300, 100)} + ${M.map((m, i) => { const y = 34 + i * 33; return `${comp(328, y, m[1], m[0], 1)}${m[0]}`; }).join("")} + ${_lbl(356, 197, "same cell, different stain \u2014 pick any of 42")} + `; }, + body: "We reuse the diffusion autoencoder as a virtual-staining model, conditioning it on a phase image in two ways: a semantic path (phase → frozen Cell-DINO ViT → a pooled code that FiLM-conditions the U-Net, with a marker id — the what) and a spatial path (the raw phase pixels concatenated into the U-Net input — keeping the layout, so the output stays pixel-registered). Switching the marker id renders any of 42 fluorescent channels from the same phase cell.", + why: "The spatial conditioning is what makes it faithful — predicted markers line up with the real cell's structures (Pearson lifts ~0.13 → ~0.78). One model covers all 42 live markers from a single label-free image; run on a traversal it yields a full multi-channel phenotype per perturbation.", + defs: [["Virtual staining", "predicting fluorescent-marker images from a label-free phase image."], + ["Semantic (FiLM) conditioning", "phase → frozen Cell-DINO ViT → a pooled code that globally steers the U-Net (with a marker id) — the \"what\" to render, carrying no spatial layout."], + ["Spatial conditioning", "the raw phase image concatenated as an extra U-Net input channel, so the prediction keeps the input's layout and stays pixel-registered (fidelity lift ~0.13 → ~0.78)."], + ["FiLM", "feature-wise modulation — how the semantic code + marker id scale the U-Net's features to select what to generate."], + ["Marker id", "a learned token selecting which of the 42 fluorescent channels to render from the same phase cell."]] + }, + { + nav: "mRNA phenotypes", kicker: "THE NEXT DIRECTION", title: "From transcription to morphology", + svg: () => { const bases = ["A", "U", "G", "C", "A", "G", "U", "C"], bcol = { A: MTH_C.grn, U: MTH_C.yel, G: MTH_C.acc, C: MTH_C.pur }; + const sx = 12, sw = 84, n = bases.length, step = sw / (n - 1), sy = 52; + const pos = bases.map((b, i) => [sx + i * step, sy + (i % 2 ? -8 : 8)]); + const backbone = "M " + pos.map(p => `${p[0].toFixed(1)} ${p[1].toFixed(1)}`).join(" L "); + const beads = bases.map((b, i) => `${b}`).join(""); + return ` + interleukin mRNA \u2191 + ${beads} + + + ${_lbl(54, 170, "expression")} + ${_arrow(98, 250, 100)} + ${_lbl(176, 86, "predict", "#8b949e", 9)} + ${_cellBlob(348, 100, 42, MTH_C.acc)} + ${_lbl(348, 164, "morphology follows")} + `; }, + body: "The same library was also profiled by CROP-seq (single-cell RNA), so each perturbation has a transcriptional signature too. Next we condition the diffusion decoder on transcriptional change instead of a morphological direction — either by mapping a perturbation's expression shift onto a CellDINO direction (reusing this decoder), or by feeding the transcriptome straight into the model as a learned perturbation vector added in its latent space. Then we ask how a cell's shape should follow its gene-expression state \u2014 e.g. as interleukin expression rises, watch the morphology change.", + why: "Bridging transcriptome and image lets us see which perturbations decouple the two (loud in RNA, silent in shape — or the reverse) — the core question of the transcriptional project.", + defs: [["CROP-seq", "a pooled CRISPR screen read out by single-cell RNA sequencing (the transcriptome), on the same knockout library."], + ["Transcriptional signature", "a perturbation's pseudobulk mRNA shift vs NTC (or a learned scRNA latent, or pathway module scores) — the vector that stands in for the CellDINO direction."], + ["Reuse map (light design)", "fit transcriptional-shift → CellDINO-shift (linear, then a small MLP), then feed that predicted direction into the existing diffusion decoder — no retraining; its R² measures how predictable morphology is from transcriptome."], + ["Learned perturbation vector", "represent each perturbation\u2019s effect as a vector added in the model\u2019s latent space; a transcriptome decoder and the image decoder share it, so any transcriptional state can be rendered. (This is the idea behind a \u201ccompositional perturbation autoencoder,\u201d CPA.)"]] + }, +]; + +// display order (narrative arc): data → represent → classify+interpret → generate+arrange → cross-channel → next +// core methods (numbered 1..N, count stops at Montage) then "Extra" add-on slides (unnumbered) +const MTH_ORDER = ["The screen", "Embedding", "Classifier", "Top cells", "Diffusion", "Traversal", "Montage", + "Attention heads", "Virtual staining", "mRNA phenotypes"]; // last 3 are Extras; DDIM folded into Diffusion +const MTH_EXTRA = new Set(["Attention heads", "Virtual staining", "mRNA phenotypes"]); +const MTH_DECK = MTH_ORDER.map(n => METHODS_SLIDES.find(s => s.nav === n)).filter(Boolean); +// public deck = steps 1–6 (Screen…Traversal); internal = full deck. isPublic() from app.js (loaded first). +const _mthDeck = () => (typeof isPublic === "function" && isPublic()) ? MTH_DECK.slice(0, 6) : MTH_DECK; + +let _mthIdx = 0; +function renderMethods() { + const deck = _mthDeck(), coreN = deck.filter(s => !MTH_EXTRA.has(s.nav)).length; + if (_mthIdx >= deck.length) _mthIdx = deck.length - 1; + const rail = document.getElementById("tab-methods"); + if (rail && rail.querySelectorAll(".mth-railitem").length !== deck.length) { // (re)build when the deck size changes (e.g. public toggle) + if (!rail.dataset.inited) { const sv = +localStorage.getItem("opsin.mth"); if (sv >= 0 && sv < deck.length) _mthIdx = sv; rail.dataset.inited = "1"; } // restore last-viewed once + rail.innerHTML = `
A visual tour of the methods behind this viewer — click through, ← → to navigate.
+
${deck.map((s, i) => { const x = MTH_EXTRA.has(s.nav); + return ``; }).join("")}
`; + } + const s = deck[_mthIdx], view = document.getElementById("methods-view"); + if (!view) return; + const refs = MTH_REFS[s.nav] || []; + view.innerHTML = `
+
${s.kicker} · ${MTH_EXTRA.has(s.nav) ? "Extra" : (_mthIdx + 1) + " / " + coreN}
+

${s.title}

+ ${s.sections ? s.sections.map(sec => `
${sec.cap ? `
${sec.cap}
` : ""}
${sec.svg()}
${sec.text ? `
${sec.text}
` : ""}
`).join("") : `
${s.svg()}
`} +
${s.body}
+ ${MTH_MORE[s.nav] ? `
Learn more

${MTH_MORE[s.nav]}

` : ""} +
Why it matters — ${s.why}
+ ${refs.length ? `
📄 ${refs.map(r => `${r[0]}`).join("  ·  ")}
` : ""} + ${(s.defs || []).length ? `
Key terms
${s.defs.map(d => `
${d[0]}
${d[1]}
`).join("")}
` : ""} +
+ +
${deck.map((_, i) => ``).join("")}
+ +
+
`; + document.querySelectorAll(".mth-railitem").forEach((b, i) => b.classList.toggle("on", i === _mthIdx)); + const st = document.getElementById("stage"); if (st) st.scrollTop = 0; // new slide → back to top (don't inherit the previous slide's scroll) +} +function methodsGo(i) { _mthIdx = Math.max(0, Math.min(_mthDeck().length - 1, i)); try { localStorage.setItem("opsin.mth", _mthIdx); } catch (e) { } renderMethods(); } +function methodsStep(d) { methodsGo(_mthIdx + d); } +document.addEventListener("keydown", (e) => { + const active = document.querySelector(".tab.active"); + if (active && active.dataset.tab === "methods") { if (e.key === "ArrowRight") methodsStep(1); if (e.key === "ArrowLeft") methodsStep(-1); } +}); diff --git a/src/ops_model/models/attention/diffex/viewer/webapp/morpho_demo.html b/src/ops_model/models/attention/diffex/viewer/webapp/morpho_demo.html new file mode 100644 index 0000000..b1fa7c9 --- /dev/null +++ b/src/ops_model/models/attention/diffex/viewer/webapp/morpho_demo.html @@ -0,0 +1,333 @@ + + + + +DiffEx — Morphometrics demo + + + +
+

Morphometrics — real organelle_profiler seg on generated cells

+ + + + + + + + + +
+
+
+
+
Blobs = real org-seg, colored by the selected feature (fixed clim, shared with the reference cells). Scrub α → watch the phenotype amplify.
+
+ +
+
+
+
Real cells (same org-seg + colormap)
+
+
+
+ + + diff --git a/src/ops_model/models/attention/diffex/viewer/webapp/openseadragon.min.js b/src/ops_model/models/attention/diffex/viewer/webapp/openseadragon.min.js new file mode 100644 index 0000000..b2778cf --- /dev/null +++ b/src/ops_model/models/attention/diffex/viewer/webapp/openseadragon.min.js @@ -0,0 +1,9 @@ +//! openseadragon 4.1.0 +//! Built on 2023-05-25 +//! Git commit: v4.1.0-0-8849681 +//! http://openseadragon.github.io +//! License: http://openseadragon.github.io/license/ + + +function OpenSeadragon(e){return new OpenSeadragon.Viewer(e)}!function(n){n.version={versionStr:"4.1.0",major:parseInt("4",10),minor:parseInt("1",10),revision:parseInt("0",10)};var t={"[object Boolean]":"boolean","[object Number]":"number","[object String]":"string","[object Function]":"function","[object AsyncFunction]":"function","[object Promise]":"promise","[object Array]":"array","[object Date]":"date","[object RegExp]":"regexp","[object Object]":"object"},i=Object.prototype.toString,o=Object.prototype.hasOwnProperty;n.isFunction=function(e){return"function"===n.type(e)};n.isArray=Array.isArray||function(e){return"array"===n.type(e)};n.isWindow=function(e){return e&&"object"==typeof e&&"setInterval"in e};n.type=function(e){return null==e?String(e):t[i.call(e)]||"object"};n.isPlainObject=function(e){if(!e||"object"!==OpenSeadragon.type(e)||e.nodeType||n.isWindow(e))return!1;if(e.constructor&&!o.call(e,"constructor")&&!o.call(e.constructor.prototype,"isPrototypeOf"))return!1;var t;for(var i in e)t=i;return void 0===t||o.call(e,t)};n.isEmptyObject=function(e){for(var t in e)return!1;return!0};n.freezeObject=function(e){Object.freeze?n.freezeObject=Object.freeze:n.freezeObject=function(e){return e};return n.freezeObject(e)};n.supportsCanvas=(e=document.createElement("canvas"),!(!n.isFunction(e.getContext)||!e.getContext("2d")));var e;n.isCanvasTainted=function(e){var t=!1;try{e.getContext("2d").getImageData(0,0,1,1)}catch(e){t=!0}return t};n.supportsAddEventListener=!(!document.documentElement.addEventListener||!document.addEventListener);n.supportsRemoveEventListener=!(!document.documentElement.removeEventListener||!document.removeEventListener);n.supportsEventListenerOptions=function(){var t=0;if(n.supportsAddEventListener)try{var e={get capture(){t++;return!1},get once(){t++;return!1},get passive(){t++;return!1}};window.addEventListener("test",null,e);window.removeEventListener("test",null,e)}catch(e){t=0}return 3<=t}();n.getCurrentPixelDensityRatio=function(){if(n.supportsCanvas){var e=document.createElement("canvas").getContext("2d");var t=window.devicePixelRatio||1;e=e.webkitBackingStorePixelRatio||e.mozBackingStorePixelRatio||e.msBackingStorePixelRatio||e.oBackingStorePixelRatio||e.backingStorePixelRatio||1;return Math.max(t,1)/e}return 1};n.pixelDensityRatio=n.getCurrentPixelDensityRatio()}(OpenSeadragon);!function(u){u.extend=function(){var e,t,i,n,o,r=arguments[0]||{},s=arguments.length,a=!1,l=1;if("boolean"==typeof r){a=r;r=arguments[1]||{};l=2}"object"==typeof r||OpenSeadragon.isFunction(r)||(r={});if(s===l){r=this;--l}for(;l=i.x&&t.x=i.y},getMousePosition:function(e){if("number"==typeof e.pageX)u.getMousePosition=function(e){var t=new u.Point;t.x=e.pageX;t.y=e.pageY;return t};else{if("number"!=typeof e.clientX)throw new Error("Unknown event mouse position, no known technique.");u.getMousePosition=function(e){var t=new u.Point;t.x=e.clientX+document.body.scrollLeft+document.documentElement.scrollLeft;t.y=e.clientY+document.body.scrollTop+document.documentElement.scrollTop;return t}}return u.getMousePosition(e)},getPageScroll:function(){var e=document.documentElement||{},t=document.body||{};if("number"==typeof window.pageXOffset)u.getPageScroll=function(){return new u.Point(window.pageXOffset,window.pageYOffset)};else if(t.scrollLeft||t.scrollTop)u.getPageScroll=function(){return new u.Point(document.body.scrollLeft,document.body.scrollTop)};else{if(!e.scrollLeft&&!e.scrollTop)return new u.Point(0,0);u.getPageScroll=function(){return new u.Point(document.documentElement.scrollLeft,document.documentElement.scrollTop)}}return u.getPageScroll()},setPageScroll:function(e){if(void 0!==window.scrollTo)u.setPageScroll=function(e){window.scrollTo(e.x,e.y)};else{var t=u.getPageScroll();if(t.x===e.x&&t.y===e.y)return;document.body.scrollLeft=e.x;document.body.scrollTop=e.y;var i=u.getPageScroll();if(i.x!==t.x&&i.y!==t.y){u.setPageScroll=function(e){document.body.scrollLeft=e.x;document.body.scrollTop=e.y};return}document.documentElement.scrollLeft=e.x;document.documentElement.scrollTop=e.y;if((i=u.getPageScroll()).x!==t.x&&i.y!==t.y){u.setPageScroll=function(e){document.documentElement.scrollLeft=e.x;document.documentElement.scrollTop=e.y};return}u.setPageScroll=function(e){}}u.setPageScroll(e)},getWindowSize:function(){var e=document.documentElement||{},t=document.body||{};if("number"==typeof window.innerWidth)u.getWindowSize=function(){return new u.Point(window.innerWidth,window.innerHeight)};else if(e.clientWidth||e.clientHeight)u.getWindowSize=function(){return new u.Point(document.documentElement.clientWidth,document.documentElement.clientHeight)};else{if(!t.clientWidth&&!t.clientHeight)throw new Error("Unknown window size, no known technique.");u.getWindowSize=function(){return new u.Point(document.body.clientWidth,document.body.clientHeight)}}return u.getWindowSize()},makeCenteredNode:function(e){e=u.getElement(e);var t=[u.makeNeutralElement("div"),u.makeNeutralElement("div"),u.makeNeutralElement("div")];u.extend(t[0].style,{display:"table",height:"100%",width:"100%"});u.extend(t[1].style,{display:"table-row"});u.extend(t[2].style,{display:"table-cell",verticalAlign:"middle",textAlign:"center"});t[0].appendChild(t[1]);t[1].appendChild(t[2]);t[2].appendChild(e);return t[0]},makeNeutralElement:function(e){var t=document.createElement(e),e=t.style;e.background="transparent none";e.border="none";e.margin="0px";e.padding="0px";e.position="static";return t},now:function(){Date.now?u.now=Date.now:u.now=function(){return(new Date).getTime()};return u.now()},makeTransparentImage:function(e){var t=u.makeNeutralElement("img");t.src=e;return t},setElementOpacity:function(e,t,i){e=u.getElement(e);i&&!u.Browser.alpha&&(t=Math.round(t));if(u.Browser.opacity)e.style.opacity=t<1?t:"";else if(t<1){t=Math.round(100*t);e.style.filter="alpha(opacity="+t+")"}else e.style.filter=""},setElementTouchActionNone:function(e){void 0!==(e=u.getElement(e)).style.touchAction?e.style.touchAction="none":void 0!==e.style.msTouchAction&&(e.style.msTouchAction="none")},setElementPointerEvents:function(e,t){void 0!==(e=u.getElement(e)).style&&void 0!==e.style.pointerEvents&&(e.style.pointerEvents=t)},setElementPointerEventsNone:function(e){u.setElementPointerEvents(e,"none")},addClass:function(e,t){(e=u.getElement(e)).className?-1===(" "+e.className+" ").indexOf(" "+t+" ")&&(e.className+=" "+t):e.className=t},indexOf:function(e,t,i){Array.prototype.indexOf?this.indexOf=function(e,t,i){return e.indexOf(t,i)}:this.indexOf=function(e,t,i){var n,o,i=i||0;if(!e)throw new TypeError;if(0===(o=e.length)||o<=i)return-1;for(n=i=i<0?o-Math.abs(i):i;nt.touches.length-r&&c.console.warn("Tracked touch contact count doesn't match event.touches.length");var a={originalEvent:t,eventType:"pointerdown",pointerType:"touch",isEmulated:!1};B(e,a);for(n=0;n\s*$/))n=m.parseXml(n);else if(n.match(/^\s*[{[].*[}\]]\s*$/))try{var e=m.parseJSON(n);n=e}catch(e){}function l(e,t){if(e.ready)r(e);else{e.addHandler("ready",function(){r(e)});e.addHandler("open-failed",function(e){s({message:e.message,source:t})})}}setTimeout(function(){if("string"===m.type(n))(n=new m.TileSource({url:n,crossOriginPolicy:(void 0!==o.crossOriginPolicy?o:i).crossOriginPolicy,ajaxWithCredentials:i.ajaxWithCredentials,ajaxHeaders:o.ajaxHeaders||i.ajaxHeaders,splitHashDataForPost:i.splitHashDataForPost,useCanvas:i.useCanvas,success:function(e){r(e.tileSource)}})).addHandler("open-failed",function(e){s(e)});else if(m.isPlainObject(n)||n.nodeType){void 0!==n.crossOriginPolicy||void 0===o.crossOriginPolicy&&void 0===i.crossOriginPolicy||(n.crossOriginPolicy=(void 0!==o.crossOriginPolicy?o:i).crossOriginPolicy);void 0===n.ajaxWithCredentials&&(n.ajaxWithCredentials=i.ajaxWithCredentials);void 0===n.useCanvas&&(n.useCanvas=i.useCanvas);if(m.isFunction(n.getTileUrl)){var e=new m.TileSource(n);e.getTileUrl=n.getTileUrl;r(e)}else{var t=m.TileSource.determineType(a,n);if(t){e=t.prototype.configure.apply(a,[n]);l(new t(e),n)}else s({message:"Unable to load TileSource",source:n})}}else l(n,n)})}(this,i.tileSource,i,function(e){o.tileSource=e;s()},function(e){e.options=i;t(e);s()})}function s(){var e,t;for(;n._loadQueue.length&&(e=n._loadQueue[0]).tileSource;){n._loadQueue.splice(0,1);if(e.options.replace){var i=n.world.getIndexOfItem(e.options.replaceItem);-1!==i&&(e.options.index=i);n.world.removeItem(e.options.replaceItem)}t=new m.TiledImage({viewer:n,source:e.tileSource,viewport:n.viewport,drawer:n.drawer,tileCache:n.tileCache,imageLoader:n.imageLoader,x:e.options.x,y:e.options.y,width:e.options.width,height:e.options.height,fitBounds:e.options.fitBounds,fitBoundsPlacement:e.options.fitBoundsPlacement,clip:e.options.clip,placeholderFillStyle:e.options.placeholderFillStyle,opacity:e.options.opacity,preload:e.options.preload,degrees:e.options.degrees,flipped:e.options.flipped,compositeOperation:e.options.compositeOperation,springStiffness:n.springStiffness,animationTime:n.animationTime,minZoomImageRatio:n.minZoomImageRatio,wrapHorizontal:n.wrapHorizontal,wrapVertical:n.wrapVertical,immediateRender:n.immediateRender,blendTime:n.blendTime,alwaysBlend:n.alwaysBlend,minPixelRatio:n.minPixelRatio,smoothTileEdgesMinZoom:n.smoothTileEdgesMinZoom,iOSDevice:n.iOSDevice,crossOriginPolicy:e.options.crossOriginPolicy,ajaxWithCredentials:e.options.ajaxWithCredentials,loadTilesWithAjax:e.options.loadTilesWithAjax,ajaxHeaders:e.options.ajaxHeaders,debugMode:n.debugMode,subPixelRoundingForTransparency:n.subPixelRoundingForTransparency});n.collectionMode&&n.world.setAutoRefigureSizes(!1);if(n.navigator){i=m.extend({},e.options,{replace:!1,originalTiledImage:t,tileSource:e.tileSource});n.navigator.addTiledImage(i)}n.world.addItem(t,{index:e.options.index});0===n._loadQueue.length&&r(e);1!==n.world.getItemCount()||n.preserveViewport||n.viewport.goHome(!0);e.options.success&&e.options.success({item:t})}}},addSimpleImage:function(e){m.console.assert(e,"[Viewer.addSimpleImage] options is required");m.console.assert(e.url,"[Viewer.addSimpleImage] options.url is required");e=m.extend({},e,{tileSource:{type:"image",url:e.url}});delete e.url;this.addTiledImage(e)},addLayer:function(t){var i=this;m.console.error("[Viewer.addLayer] this function is deprecated; use Viewer.addTiledImage() instead.");var e=m.extend({},t,{success:function(e){i.raiseEvent("add-layer",{options:t,drawer:e.item})},error:function(e){i.raiseEvent("add-layer-failed",e)}});this.addTiledImage(e);return this},getLayerAtLevel:function(e){m.console.error("[Viewer.getLayerAtLevel] this function is deprecated; use World.getItemAt() instead.");return this.world.getItemAt(e)},getLevelOfLayer:function(e){m.console.error("[Viewer.getLevelOfLayer] this function is deprecated; use World.getIndexOfItem() instead.");return this.world.getIndexOfItem(e)},getLayersCount:function(){m.console.error("[Viewer.getLayersCount] this function is deprecated; use World.getItemCount() instead.");return this.world.getItemCount()},setLayerLevel:function(e,t){m.console.error("[Viewer.setLayerLevel] this function is deprecated; use World.setItemIndex() instead.");return this.world.setItemIndex(e,t)},removeLayer:function(e){m.console.error("[Viewer.removeLayer] this function is deprecated; use World.removeItem() instead.");return this.world.removeItem(e)},forceRedraw:function(){c[this.hash].forceRedraw=!0;return this},forceResize:function(){c[this.hash].needsResize=!0;c[this.hash].forceResize=!0},bindSequenceControls:function(){var e=m.delegate(this,v),t=m.delegate(this,f),i=m.delegate(this,this.goToNextPage),n=m.delegate(this,this.goToPreviousPage),o=this.navImages,r=!0;if(this.showSequenceControl){(this.previousButton||this.nextButton)&&(r=!1);this.previousButton=new m.Button({element:this.previousButton?m.getElement(this.previousButton):null,clickTimeThreshold:this.clickTimeThreshold,clickDistThreshold:this.clickDistThreshold,tooltip:m.getString("Tooltips.PreviousPage"),srcRest:B(this.prefixUrl,o.previous.REST),srcGroup:B(this.prefixUrl,o.previous.GROUP),srcHover:B(this.prefixUrl,o.previous.HOVER),srcDown:B(this.prefixUrl,o.previous.DOWN),onRelease:n,onFocus:e,onBlur:t});this.nextButton=new m.Button({element:this.nextButton?m.getElement(this.nextButton):null,clickTimeThreshold:this.clickTimeThreshold,clickDistThreshold:this.clickDistThreshold,tooltip:m.getString("Tooltips.NextPage"),srcRest:B(this.prefixUrl,o.next.REST),srcGroup:B(this.prefixUrl,o.next.GROUP),srcHover:B(this.prefixUrl,o.next.HOVER),srcDown:B(this.prefixUrl,o.next.DOWN),onRelease:i,onFocus:e,onBlur:t});this.navPrevNextWrap||this.previousButton.disable();this.tileSources&&this.tileSources.length||this.nextButton.disable();if(r){this.paging=new m.ButtonGroup({buttons:[this.previousButton,this.nextButton],clickTimeThreshold:this.clickTimeThreshold,clickDistThreshold:this.clickDistThreshold});this.pagingControl=this.paging.element;this.toolbar?this.toolbar.addControl(this.pagingControl,{anchor:m.ControlAnchor.BOTTOM_RIGHT}):this.addControl(this.pagingControl,{anchor:this.sequenceControlAnchor||m.ControlAnchor.TOP_LEFT})}}return this},bindStandardControls:function(){var e=m.delegate(this,L),t=m.delegate(this,M),i=m.delegate(this,N),n=m.delegate(this,F),o=m.delegate(this,A),r=m.delegate(this,U),s=m.delegate(this,V),a=m.delegate(this,j),l=m.delegate(this,G),h=m.delegate(this,q),c=m.delegate(this,v),u=m.delegate(this,f),d=this.navImages,p=[],g=!0;if(this.showNavigationControl){(this.zoomInButton||this.zoomOutButton||this.homeButton||this.fullPageButton||this.rotateLeftButton||this.rotateRightButton||this.flipButton)&&(g=!1);if(this.showZoomControl){p.push(this.zoomInButton=new m.Button({element:this.zoomInButton?m.getElement(this.zoomInButton):null,clickTimeThreshold:this.clickTimeThreshold,clickDistThreshold:this.clickDistThreshold,tooltip:m.getString("Tooltips.ZoomIn"),srcRest:B(this.prefixUrl,d.zoomIn.REST),srcGroup:B(this.prefixUrl,d.zoomIn.GROUP),srcHover:B(this.prefixUrl,d.zoomIn.HOVER),srcDown:B(this.prefixUrl,d.zoomIn.DOWN),onPress:e,onRelease:t,onClick:i,onEnter:e,onExit:t,onFocus:c,onBlur:u}));p.push(this.zoomOutButton=new m.Button({element:this.zoomOutButton?m.getElement(this.zoomOutButton):null,clickTimeThreshold:this.clickTimeThreshold,clickDistThreshold:this.clickDistThreshold,tooltip:m.getString("Tooltips.ZoomOut"),srcRest:B(this.prefixUrl,d.zoomOut.REST),srcGroup:B(this.prefixUrl,d.zoomOut.GROUP),srcHover:B(this.prefixUrl,d.zoomOut.HOVER),srcDown:B(this.prefixUrl,d.zoomOut.DOWN),onPress:n,onRelease:t,onClick:o,onEnter:n,onExit:t,onFocus:c,onBlur:u}))}this.showHomeControl&&p.push(this.homeButton=new m.Button({element:this.homeButton?m.getElement(this.homeButton):null,clickTimeThreshold:this.clickTimeThreshold,clickDistThreshold:this.clickDistThreshold,tooltip:m.getString("Tooltips.Home"),srcRest:B(this.prefixUrl,d.home.REST),srcGroup:B(this.prefixUrl,d.home.GROUP),srcHover:B(this.prefixUrl,d.home.HOVER),srcDown:B(this.prefixUrl,d.home.DOWN),onRelease:r,onFocus:c,onBlur:u}));this.showFullPageControl&&p.push(this.fullPageButton=new m.Button({element:this.fullPageButton?m.getElement(this.fullPageButton):null,clickTimeThreshold:this.clickTimeThreshold,clickDistThreshold:this.clickDistThreshold,tooltip:m.getString("Tooltips.FullPage"),srcRest:B(this.prefixUrl,d.fullpage.REST),srcGroup:B(this.prefixUrl,d.fullpage.GROUP),srcHover:B(this.prefixUrl,d.fullpage.HOVER),srcDown:B(this.prefixUrl,d.fullpage.DOWN),onRelease:s,onFocus:c,onBlur:u}));if(this.showRotationControl){p.push(this.rotateLeftButton=new m.Button({element:this.rotateLeftButton?m.getElement(this.rotateLeftButton):null,clickTimeThreshold:this.clickTimeThreshold,clickDistThreshold:this.clickDistThreshold,tooltip:m.getString("Tooltips.RotateLeft"),srcRest:B(this.prefixUrl,d.rotateleft.REST),srcGroup:B(this.prefixUrl,d.rotateleft.GROUP),srcHover:B(this.prefixUrl,d.rotateleft.HOVER),srcDown:B(this.prefixUrl,d.rotateleft.DOWN),onRelease:a,onFocus:c,onBlur:u}));p.push(this.rotateRightButton=new m.Button({element:this.rotateRightButton?m.getElement(this.rotateRightButton):null,clickTimeThreshold:this.clickTimeThreshold,clickDistThreshold:this.clickDistThreshold,tooltip:m.getString("Tooltips.RotateRight"),srcRest:B(this.prefixUrl,d.rotateright.REST),srcGroup:B(this.prefixUrl,d.rotateright.GROUP),srcHover:B(this.prefixUrl,d.rotateright.HOVER),srcDown:B(this.prefixUrl,d.rotateright.DOWN),onRelease:l,onFocus:c,onBlur:u}))}this.showFlipControl&&p.push(this.flipButton=new m.Button({element:this.flipButton?m.getElement(this.flipButton):null,clickTimeThreshold:this.clickTimeThreshold,clickDistThreshold:this.clickDistThreshold,tooltip:m.getString("Tooltips.Flip"),srcRest:B(this.prefixUrl,d.flip.REST),srcGroup:B(this.prefixUrl,d.flip.GROUP),srcHover:B(this.prefixUrl,d.flip.HOVER),srcDown:B(this.prefixUrl,d.flip.DOWN),onRelease:h,onFocus:c,onBlur:u}));if(g){this.buttonGroup=new m.ButtonGroup({buttons:p,clickTimeThreshold:this.clickTimeThreshold,clickDistThreshold:this.clickDistThreshold});this.navControl=this.buttonGroup.element;this.addHandler("open",m.delegate(this,W));(this.toolbar||this).addControl(this.navControl,{anchor:this.navigationControlAnchor||m.ControlAnchor.TOP_LEFT})}else this.customButtons=p}return this},currentPage:function(){return this._sequenceIndex},goToPage:function(e){if(this.tileSources&&0<=e&&e=this.tileSources.length&&(e=0);this.goToPage(e)},isAnimating:function(){return c[this.hash].animating}});function r(e){e=m.getElement(e);return new m.Point(0===e.clientWidth?1:e.clientWidth,0===e.clientHeight?1:e.clientHeight)}function h(e,t){if(t instanceof m.Overlay)return t;var i=null;if(t.element)i=m.getElement(t.element);else{var n=t.id||"openseadragon-overlay-"+Math.floor(1e7*Math.random());(i=m.getElement(t.id))||((i=document.createElement("a")).href="#/overlay/"+n);i.id=n;m.addClass(i,t.className||"openseadragon-overlay")}var o=t.location;var r=t.width;var s=t.height;if(!o){n=t.x;var a=t.y;if(void 0!==t.px){e=e.viewport.imageToViewportRectangle(new m.Rect(t.px,t.py,r||0,s||0));n=e.x;a=e.y;r=void 0!==r?e.width:void 0;s=void 0!==s?e.height:void 0}o=new m.Point(n,a)}a=t.placement;a&&"string"===m.type(a)&&(a=m.Placement[t.placement.toUpperCase()]);return new m.Overlay({element:i,location:o,placement:a,onDraw:t.onDraw,checkResize:t.checkResize,width:r,height:s,rotationMode:t.rotationMode})}function s(e,t){var i;for(i=e.length-1;0<=i;i--)if(e[i].element===t)return i;return-1}function a(e,t){return m.requestAnimationFrame(function(){t(e)})}function l(e){m.requestAnimationFrame(function(){!function(e){var t,i,n;if(e.controlsShouldFade){t=m.now();t=t-e.controlsFadeBeginTime;i=1-t/e.controlsFadeLength;i=Math.min(1,i);i=Math.max(0,i);for(n=e.controls.length-1;0<=n;n--)e.controls[n].autoFade&&e.controls[n].setOpacity(i);0=t.flickMinSpeed){var n=0;this.panHorizontal&&(n=t.flickMomentum*e.speed*Math.cos(e.direction));i=0;this.panVertical&&(i=t.flickMomentum*e.speed*Math.sin(e.direction));e=this.viewport.pixelFromPoint(this.viewport.getCenter(!0));i=this.viewport.pointFromPixel(new m.Point(e.x-n,e.y-i));this.viewport.panTo(i,!1)}this.viewport.applyConstraints()}t.dblClickDragToZoom&&!0===c[this.hash].draggingToZoom&&(c[this.hash].draggingToZoom=!1)}function S(e){this.raiseEvent("canvas-enter",{tracker:e.eventSource,pointerType:e.pointerType,position:e.position,buttons:e.buttons,pointers:e.pointers,insideElementPressed:e.insideElementPressed,buttonDownAny:e.buttonDownAny,originalEvent:e.originalEvent})}function E(e){this.raiseEvent("canvas-exit",{tracker:e.eventSource,pointerType:e.pointerType,position:e.position,buttons:e.buttons,pointers:e.pointers,insideElementPressed:e.insideElementPressed,buttonDownAny:e.buttonDownAny,originalEvent:e.originalEvent})}function P(e){this.raiseEvent("canvas-press",{tracker:e.eventSource,pointerType:e.pointerType,position:e.position,insideElementPressed:e.insideElementPressed,insideElementReleased:e.insideElementReleased,originalEvent:e.originalEvent});if(this.gestureSettingsByDeviceType(e.pointerType).dblClickDragToZoom){var t=c[this.hash].lastClickTime;e=m.now();if(null!==t){e-tthis.minScrollDeltaTime){this._lastScrollTime=n;t={tracker:e.eventSource,position:e.position,scroll:e.scroll,shift:e.shift,originalEvent:e.originalEvent,preventDefaultAction:!1,preventDefault:!0};this.raiseEvent("canvas-scroll",t);if(!t.preventDefaultAction&&this.viewport){this.viewport.flipped&&(e.position.x=this.viewport.getContainerSize().x-e.position.x);if((i=this.gestureSettingsByDeviceType(e.pointerType)).scrollToZoom){n=Math.pow(this.zoomPerScroll,e.scroll);this.viewport.zoomBy(n,i.zoomToRefPoint?this.viewport.pointFromPixel(e.position,!0):null);this.viewport.applyConstraints()}}e.preventDefault=t.preventDefault}else e.preventDefault=!0}function k(e){c[this.hash].mouseInside=!0;n(this);this.raiseEvent("container-enter",{tracker:e.eventSource,pointerType:e.pointerType,position:e.position,buttons:e.buttons,pointers:e.pointers,insideElementPressed:e.insideElementPressed,buttonDownAny:e.buttonDownAny,originalEvent:e.originalEvent})}function H(e){if(e.pointers<1){c[this.hash].mouseInside=!1;c[this.hash].animating||u(this)}this.raiseEvent("container-exit",{tracker:e.eventSource,pointerType:e.pointerType,position:e.position,buttons:e.buttons,pointers:e.pointers,insideElementPressed:e.insideElementPressed,buttonDownAny:e.buttonDownAny,originalEvent:e.originalEvent})}function z(e){!function(e){if(!e._opening&&c[e.hash]){if(e.autoResize||c[e.hash].forceResize){if(e._autoResizePolling){i=r(e.container);var t=c[e.hash].prevContainerSize;i.equals(t)||(c[e.hash].needsResize=!0)}c[e.hash].needsResize&&function(e,t){var i=e.viewport;var n=i.getZoom();var o=i.getCenter();i.resize(t,e.preserveImageSizeOnResize);i.panTo(o,!0);var r;if(e.preserveImageSizeOnResize)r=c[e.hash].prevContainerSize.x/t.x;else{var s=new m.Point(0,0);o=new m.Point(c[e.hash].prevContainerSize.x,c[e.hash].prevContainerSize.y).distanceTo(s);s=new m.Point(t.x,t.y).distanceTo(s);r=s/o*c[e.hash].prevContainerSize.x/t.x}i.zoomTo(n*r,null,!0);c[e.hash].prevContainerSize=t;c[e.hash].forceRedraw=!0;c[e.hash].needsResize=!1;c[e.hash].forceResize=!1}(e,i||r(e.container))}t=e.viewport.update();var i=e.world.update()||t;t&&e.raiseEvent("viewport-change");e.referenceStrip&&(i=e.referenceStrip.update(e.viewport)||i);t=c[e.hash].animating;if(!t&&i){e.raiseEvent("animation-start");n(e)}t=t&&!i;t&&(c[e.hash].animating=!1);if(i||t||c[e.hash].forceRedraw||e.world.needsDraw()){!function(e){e.imageLoader.clear();e.drawer.clear();e.world.draw();e.raiseEvent("update-viewport",{})}(e);e._drawOverlays();e.navigator&&e.navigator.update(e.viewport);c[e.hash].forceRedraw=!1;i&&e.raiseEvent("animation")}if(t){e.raiseEvent("animation-finish");c[e.hash].mouseInside||u(e)}c[e.hash].animating=i}}(e);e.isOpen()?e._updateRequestId=a(e,z):e._updateRequestId=!1}function B(e,t){return e?e+t:t}function L(){c[this.hash].lastZoomTime=m.now();c[this.hash].zoomFactor=this.zoomPerSecond;c[this.hash].zooming=!0;i(this)}function F(){c[this.hash].lastZoomTime=m.now();c[this.hash].zoomFactor=1/this.zoomPerSecond;c[this.hash].zooming=!0;i(this)}function M(){c[this.hash].zooming=!1}function i(e){m.requestAnimationFrame(m.delegate(e,t))}function t(){var e,t;if(c[this.hash].zooming&&this.viewport){t=(e=m.now())-c[this.hash].lastZoomTime;t=Math.pow(c[this.hash].zoomFactor,t/1e3);this.viewport.zoomBy(t);this.viewport.applyConstraints();c[this.hash].lastZoomTime=e;i(this)}}function N(){if(this.viewport){c[this.hash].zooming=!1;this.viewport.zoomBy(+this.zoomPerClick);this.viewport.applyConstraints()}}function A(){if(this.viewport){c[this.hash].zooming=!1;this.viewport.zoomBy(1/this.zoomPerClick);this.viewport.applyConstraints()}}function W(){if(this.buttonGroup){this.buttonGroup.emulateEnter();this.buttonGroup.emulateLeave()}}function U(){this.viewport&&this.viewport.goHome()}function V(){this.isFullPage()&&!m.isFullScreen()?this.setFullPage(!1):this.setFullScreen(!this.isFullPage());this.buttonGroup&&this.buttonGroup.emulateLeave();this.fullPageButton.element.focus();this.viewport&&this.viewport.applyConstraints()}function j(){if(this.viewport){var e=this.viewport.getRotation();this.viewport.flipped?e+=this.rotationIncrement:e-=this.rotationIncrement;this.viewport.setRotation(e)}}function G(){if(this.viewport){var e=this.viewport.getRotation();this.viewport.flipped?e-=this.rotationIncrement:e+=this.rotationIncrement;this.viewport.setRotation(e)}}function q(){this.viewport.toggleFlip()}}(OpenSeadragon);!function(r){r.Navigator=function(i){var e,t=i.viewer,n=this;if(i.element||i.id){if(i.element){i.id&&r.console.warn("Given option.id for Navigator was ignored since option.element was provided and is being used instead.");i.element.id?i.id=i.element.id:i.id="navigator-"+r.now();this.element=i.element}else this.element=document.getElementById(i.id);i.controlOptions={anchor:r.ControlAnchor.NONE,attachToViewer:!1,autoFade:!1}}else{i.id="navigator-"+r.now();this.element=r.makeNeutralElement("div");i.controlOptions={anchor:r.ControlAnchor.TOP_RIGHT,attachToViewer:!0,autoFade:i.autoFade};if(i.position)if("BOTTOM_RIGHT"===i.position)i.controlOptions.anchor=r.ControlAnchor.BOTTOM_RIGHT;else if("BOTTOM_LEFT"===i.position)i.controlOptions.anchor=r.ControlAnchor.BOTTOM_LEFT;else if("TOP_RIGHT"===i.position)i.controlOptions.anchor=r.ControlAnchor.TOP_RIGHT;else if("TOP_LEFT"===i.position)i.controlOptions.anchor=r.ControlAnchor.TOP_LEFT;else if("ABSOLUTE"===i.position){i.controlOptions.anchor=r.ControlAnchor.ABSOLUTE;i.controlOptions.top=i.top;i.controlOptions.left=i.left;i.controlOptions.height=i.height;i.controlOptions.width=i.width}}this.element.id=i.id;this.element.className+=" navigator";(i=r.extend(!0,{sizeRatio:r.DEFAULT_SETTINGS.navigatorSizeRatio},i,{element:this.element,tabIndex:-1,showNavigator:!1,mouseNavEnabled:!1,showNavigationControl:!1,showSequenceControl:!1,immediateRender:!0,blendTime:0,animationTime:i.animationTime,autoResize:!1,minZoomImageRatio:1,background:i.background,opacity:i.opacity,borderColor:i.borderColor,displayRegionColor:i.displayRegionColor})).minPixelRatio=this.minPixelRatio=t.minPixelRatio;r.setElementTouchActionNone(this.element);this.borderWidth=2;this.fudge=new r.Point(1,1);this.totalBorderWidths=new r.Point(2*this.borderWidth,2*this.borderWidth).minus(this.fudge);i.controlOptions.anchor!==r.ControlAnchor.NONE&&function(e,t){e.margin="0px";e.border=t+"px solid "+i.borderColor;e.padding="0px";e.background=i.background;e.opacity=i.opacity;e.overflow="hidden"}(this.element.style,this.borderWidth);this.displayRegion=r.makeNeutralElement("div");this.displayRegion.id=this.element.id+"-displayregion";this.displayRegion.className="displayregion";!function(e,t){e.position="relative";e.top="0px";e.left="0px";e.fontSize="0px";e.overflow="hidden";e.border=t+"px solid "+i.displayRegionColor;e.margin="0px";e.padding="0px";e.background="transparent";e.float="left";e.cssFloat="left";e.styleFloat="left";e.zIndex=999999999;e.cursor="default";e.boxSizing="content-box"}(this.displayRegion.style,this.borderWidth);r.setElementPointerEventsNone(this.displayRegion);r.setElementTouchActionNone(this.displayRegion);this.displayRegionContainer=r.makeNeutralElement("div");this.displayRegionContainer.id=this.element.id+"-displayregioncontainer";this.displayRegionContainer.className="displayregioncontainer";this.displayRegionContainer.style.width="100%";this.displayRegionContainer.style.height="100%";r.setElementPointerEventsNone(this.displayRegionContainer);r.setElementTouchActionNone(this.displayRegionContainer);t.addControl(this.element,i.controlOptions);this._resizeWithViewer=i.controlOptions.anchor!==r.ControlAnchor.ABSOLUTE&&i.controlOptions.anchor!==r.ControlAnchor.NONE;if(i.width&&i.height){this.setWidth(i.width);this.setHeight(i.height)}else if(this._resizeWithViewer){e=r.getElementSize(t.element);this.element.style.height=Math.round(e.y*i.sizeRatio)+"px";this.element.style.width=Math.round(e.x*i.sizeRatio)+"px";this.oldViewerSize=e;e=r.getElementSize(this.element);this.elementArea=e.x*e.y}this.oldContainerSize=new r.Point(0,0);r.Viewer.apply(this,[i]);this.displayRegionContainer.appendChild(this.displayRegion);this.element.getElementsByTagName("div")[0].appendChild(this.displayRegionContainer);function o(e,t){c(n.displayRegionContainer,e);c(n.displayRegion,-e);n.viewport.setRotation(e,t)}if(i.navigatorRotate){o(i.viewer.viewport?i.viewer.viewport.getRotation():i.viewer.degrees||0,!0);i.viewer.addHandler("rotate",function(e){o(e.degrees,e.immediately)})}this.innerTracker.destroy();this.innerTracker=new r.MouseTracker({userData:"Navigator.innerTracker",element:this.element,dragHandler:r.delegate(this,a),clickHandler:r.delegate(this,s),releaseHandler:r.delegate(this,l),scrollHandler:r.delegate(this,h),preProcessEventHandler:function(e){"wheel"===e.eventType&&(e.preventDefault=!0)}});this.outerTracker.userData="Navigator.outerTracker";r.setElementPointerEventsNone(this.canvas);r.setElementPointerEventsNone(this.container);this.addHandler("reset-size",function(){n.viewport&&n.viewport.goHome(!0)});t.world.addHandler("item-index-change",function(t){window.setTimeout(function(){var e=n.world.getItemAt(t.previousIndex);n.world.setItemIndex(e,t.newIndex)},1)});t.world.addHandler("remove-item",function(e){e=e.item;e=n._getMatchingItem(e);e&&n.world.removeItem(e)});this.update(t.viewport)};r.extend(r.Navigator.prototype,r.EventSource.prototype,r.Viewer.prototype,{updateSize:function(){if(this.viewport){var e=new r.Point(0===this.container.clientWidth?1:this.container.clientWidth,0===this.container.clientHeight?1:this.container.clientHeight);if(!e.equals(this.oldContainerSize)){this.viewport.resize(e,!0);this.viewport.goHome(!0);this.oldContainerSize=e;this.drawer.clear();this.world.draw()}}},setWidth:function(e){this.width=e;this.element.style.width="number"==typeof e?e+"px":e;this._resizeWithViewer=!1;this.updateSize()},setHeight:function(e){this.height=e;this.element.style.height="number"==typeof e?e+"px":e;this._resizeWithViewer=!1;this.updateSize()},setFlip:function(e){this.viewport.setFlip(e);this.setDisplayTransform(this.viewer.viewport.getFlip()?"scale(-1,1)":"scale(1,1)");return this},setDisplayTransform:function(e){i(this.displayRegion,e);i(this.canvas,e);i(this.element,e)},update:function(e){var t,i;t=r.getElementSize(this.viewer.element);if(this._resizeWithViewer&&t.x&&t.y&&!t.equals(this.oldViewerSize)){this.oldViewerSize=t;if(this.maintainSizeRatio||!this.elementArea){i=t.x*this.sizeRatio;o=t.y*this.sizeRatio}else{i=Math.sqrt(this.elementArea*(t.x/t.y));o=this.elementArea/i}this.element.style.width=Math.round(i)+"px";this.element.style.height=Math.round(o)+"px";this.elementArea||(this.elementArea=i*o);this.updateSize()}if(e&&this.viewport){i=e.getBoundsNoRotate(!0);o=this.viewport.pixelFromPointNoRotate(i.getTopLeft(),!1);i=this.viewport.pixelFromPointNoRotate(i.getBottomRight(),!1).minus(this.totalBorderWidths);if(!this.navigatorRotate){var n=e.getRotation(!0);c(this.displayRegion,-n)}e=this.displayRegion.style;e.display=this.world.getItemCount()?"block":"none";e.top=o.y.toFixed(2)+"px";e.left=o.x.toFixed(2)+"px";n=i.x-o.x;var o=i.y-o.y;e.width=Math.round(Math.max(n,0))+"px";e.height=Math.round(Math.max(o,0))+"px"}},addTiledImage:function(e){var n=this;var o=e.originalTiledImage;delete e.original;e=r.extend({},e,{success:function(e){var t=e.item;t._originalForNavigator=o;n._matchBounds(t,o,!0);n._matchOpacity(t,o);n._matchCompositeOperation(t,o);function i(){n._matchBounds(t,o)}o.addHandler("bounds-change",i);o.addHandler("clip-change",i);o.addHandler("opacity-change",function(){n._matchOpacity(t,o)});o.addHandler("composite-operation-change",function(){n._matchCompositeOperation(t,o)})}});return r.Viewer.prototype.addTiledImage.apply(this,[e])},destroy:function(){return r.Viewer.prototype.destroy.apply(this)},_getMatchingItem:function(e){var t=this.world.getItemCount();var i;for(var n=0;n=1/this.aspectRatio-1e-15&&(n=this.getNumTiles(e).y-1);return new h.Point(i,n)},getTileBounds:function(e,t,i,n){var o=this.dimensions.times(this.getLevelScale(e)),r=this.getTileWidth(e),s=this.getTileHeight(e),a=0===t?0:r*t-this.tileOverlap,e=0===i?0:s*i-this.tileOverlap,t=r+(0===t?1:2)*this.tileOverlap,s=s+(0===i?1:2)*this.tileOverlap,i=1/o.x;t=Math.min(t,o.x-a);s=Math.min(s,o.y-e);return n?new h.Rect(0,0,t,s):new h.Rect(a*i,e*i,t*i,s*i)},getImageInfo:function(n){var t,i,e,o,r,s=this;n&&-1<(r=(o=(e=n.split("/"))[e.length-1]).lastIndexOf("."))&&(e[e.length-1]=o.slice(0,r));var a=null;if(this.splitHashDataForPost){var l=n.indexOf("#");if(-1!==l){a=n.substring(l+1);n=n.substr(0,l)}}t=function(e){"string"==typeof e&&(e=h.parseXml(e));var t=h.TileSource.determineType(s,e,n);if(t){void 0===(i=t.prototype.configure.apply(s,[e,n,a])).ajaxWithCredentials&&(i.ajaxWithCredentials=s.ajaxWithCredentials);i=new t(i);s.ready=!0;s.raiseEvent("ready",{tileSource:i})}else s.raiseEvent("open-failed",{message:"Unable to load TileSource",source:n})};if(n.match(/\.js$/)){l=n.split("/").pop().replace(".js","");h.jsonp({url:n,async:!1,callbackName:l,callback:t})}else h.makeAjaxRequest({url:n,postData:a,withCredentials:this.ajaxWithCredentials,headers:this.ajaxHeaders,success:function(e){e=function(t){var e,i,n=t.responseText,o=t.status;{if(!t)throw new Error(h.getString("Errors.Security"));if(200!==t.status&&0!==t.status){o=t.status;e=404===o?"Not Found":t.statusText;throw new Error(h.getString("Errors.Status",o,e))}}if(n.match(/^\s*<.*/))try{i=t.responseXML&&t.responseXML.documentElement?t.responseXML:h.parseXml(n)}catch(e){i=t.responseText}else if(n.match(/\s*[{[].*/))try{i=h.parseJSON(n)}catch(e){i=n}else i=n;return i}(e);t(e)},error:function(e,t){var i;try{i="HTTP "+e.status+" attempting to load TileSource: "+n}catch(e){i=(void 0!==t&&t.toString?t.toString():"Unknown error")+" attempting to load TileSource: "+n}h.console.error(i);s.raiseEvent("open-failed",{message:i,source:n,postData:a})}})},supports:function(e,t){return!1},configure:function(e,t,i){throw new Error("Method not implemented.")},getTileUrl:function(e,t,i){throw new Error("Method not implemented.")},getTilePostData:function(e,t,i){return null},getTileAjaxHeaders:function(e,t,i){return{}},getTileHashKey:function(e,t,i,n,o,r){function s(e){return o?e+"+"+JSON.stringify(o):e}return s("string"!=typeof n?e+"/"+t+"_"+i:n)},tileExists:function(e,t,i){var n=this.getNumTiles(e);return e>=this.minLevel&&e<=this.maxLevel&&0<=t&&0<=i&&tthis.maxLevel)return!1;if(!h||!h.length)return!0;for(l=h.length-1;0<=l;l--)if(!(e<(n=h[l]).minLevel||e>n.maxLevel)){a=this.getLevelScale(e);o=n.x*a;r=n.y*a;s=o+n.width*a;a=r+n.height*a;o=Math.floor(o/this._tileWidth);r=Math.floor(r/this._tileWidth);s=Math.ceil(s/this._tileWidth);a=Math.ceil(a/this._tileWidth);if(o<=t&&t=this.minLevel&&e<=this.maxLevel?this.levels[e].width/this.levels[this.maxLevel].width:t}return h.TileSource.prototype.getLevelScale.call(this,e)},getNumTiles:function(e){if(this.emulateLegacyImagePyramid)return this.getLevelScale(e)?new h.Point(1,1):new h.Point(0,0);if(this.levelSizes){var t=this.levelSizes[e];var i=Math.ceil(t.width/this.getTileWidth(e)),t=Math.ceil(t.height/this.getTileHeight(e));return new h.Point(i,t)}return h.TileSource.prototype.getNumTiles.call(this,e)},getTileAtPoint:function(e,t){if(this.emulateLegacyImagePyramid)return new h.Point(0,0);if(this.levelSizes){var i=0<=t.x&&t.x<=1&&0<=t.y&&t.y<=1/this.aspectRatio;h.console.assert(i,"[TileSource.getTileAtPoint] must be called with a valid point.");var n=this.levelSizes[e].width;i=t.x*n;n=t.y*n;i=Math.floor(i/this.getTileWidth(e));n=Math.floor(n/this.getTileHeight(e));1<=t.x&&(i=this.getNumTiles(e).x-1);t.y>=1/this.aspectRatio-1e-15&&(n=this.getNumTiles(e).y-1);return new h.Point(i,n)}return h.TileSource.prototype.getTileAtPoint.call(this,e,t)},getTileUrl:function(e,t,i){if(this.emulateLegacyImagePyramid){var n=null;return n=0=this.minLevel&&e<=this.maxLevel?this.levels[e].url:n}var o,r,s,a,l,h,c,u,d=Math.pow(.5,this.maxLevel-e);if(this.levelSizes){o=this.levelSizes[e].width;r=this.levelSizes[e].height}else{o=Math.ceil(this.width*d);r=Math.ceil(this.height*d)}c=this.getTileWidth(e);u=this.getTileHeight(e);a=Math.round(c/d);l=Math.round(u/d);n=1===this.version?"native."+this.tileFormat:"default."+this.tileFormat;if(oe.tileSize||parseInt(t.y,10)>e.tileSize;){t.x=Math.floor(t.x/2);t.y=Math.floor(t.y/2);e.imageSizes.push({x:t.x,y:t.y});e.gridSize.push(this._getGridSize(t.x,t.y,e.tileSize))}e.imageSizes.reverse();e.gridSize.reverse();e.minLevel=0;e.maxLevel=e.gridSize.length-1;OpenSeadragon.TileSource.apply(this,[e])};e.extend(e.ZoomifyTileSource.prototype,e.TileSource.prototype,{_getGridSize:function(e,t,i){return{x:Math.ceil(e/i),y:Math.ceil(t/i)}},_calculateAbsoluteTileNumber:function(e,t,i){var n=0;var o={};for(var r=0;r");return n.sort(function(e,t){return e.height-t.height})}(t.levels);if(0=this.minLevel&&e<=this.maxLevel?this.levels[e].width/this.levels[this.maxLevel].width:t},getNumTiles:function(e){return this.getLevelScale(e)?new a.Point(1,1):new a.Point(0,0)},getTileUrl:function(e,t,i){var n=null;return n=0=this.minLevel&&e<=this.maxLevel?this.levels[e].url:n}})}(OpenSeadragon);!function(a){a.ImageTileSource=function(e){e=a.extend({buildPyramid:!0,crossOriginPolicy:!1,ajaxWithCredentials:!1,useCanvas:!0},e);a.TileSource.apply(this,[e])};a.extend(a.ImageTileSource.prototype,a.TileSource.prototype,{supports:function(e,t){return e.type&&"image"===e.type},configure:function(e,t,i){return e},getImageInfo:function(e){var t=this._image=new Image;var i=this;this.crossOriginPolicy&&(t.crossOrigin=this.crossOriginPolicy);this.ajaxWithCredentials&&(t.useCredentials=this.ajaxWithCredentials);a.addEvent(t,"load",function(){i.width=t.naturalWidth;i.height=t.naturalHeight;i.aspectRatio=i.width/i.height;i.dimensions=new a.Point(i.width,i.height);i._tileWidth=i.width;i._tileHeight=i.height;i.tileOverlap=0;i.minLevel=0;i.levels=i._buildLevels();i.maxLevel=i.levels.length-1;i.ready=!0;i.raiseEvent("ready",{tileSource:i})});a.addEvent(t,"error",function(){i.raiseEvent("open-failed",{message:"Error loading image at "+e,source:e})});t.src=e},getLevelScale:function(e){var t=NaN;return t=e>=this.minLevel&&e<=this.maxLevel?this.levels[e].width/this.levels[this.maxLevel].width:t},getNumTiles:function(e){return this.getLevelScale(e)?new a.Point(1,1):new a.Point(0,0)},getTileUrl:function(e,t,i){var n=null;return n=e>=this.minLevel&&e<=this.maxLevel?this.levels[e].url:n},getContext2D:function(e,t,i){var n=null;return n=e>=this.minLevel&&e<=this.maxLevel?this.levels[e].context2D:n},destroy:function(){this._freeupCanvasMemory()},_buildLevels:function(){var e=[{url:this._image.src,width:this._image.naturalWidth,height:this._image.naturalHeight}];if(!this.buildPyramid||!a.supportsCanvas||!this.useCanvas){delete this._image;return e}var t=this._image.naturalWidth;var i=this._image.naturalHeight;var n=document.createElement("canvas");var o=n.getContext("2d");n.width=t;n.height=i;o.drawImage(this._image,0,0,t,i);e[0].context2D=o;delete this._image;if(a.isCanvasTainted(n))return e;for(;2<=t&&2<=i;){t=Math.floor(t/2);i=Math.floor(i/2);var r=document.createElement("canvas");var s=r.getContext("2d");r.width=t;r.height=i;s.drawImage(n,0,0,t,i);e.splice(0,0,{context2D:s,width:t,height:i});n=r;o=s}return e},_freeupCanvasMemory:function(){for(var e=0;e=i.ButtonState.GROUP&&e.currentState===i.ButtonState.REST){!function(e){e.shouldFade=!1;e.imgGroup&&i.setElementOpacity(e.imgGroup,1,!0)}(e);e.currentState=i.ButtonState.GROUP}if(t>=i.ButtonState.HOVER&&e.currentState===i.ButtonState.GROUP){e.imgHover&&(e.imgHover.style.visibility="");e.currentState=i.ButtonState.HOVER}if(t>=i.ButtonState.DOWN&&e.currentState===i.ButtonState.HOVER){e.imgDown&&(e.imgDown.style.visibility="");e.currentState=i.ButtonState.DOWN}}}function r(e,t){if(!e.element.disabled){if(t<=i.ButtonState.HOVER&&e.currentState===i.ButtonState.DOWN){e.imgDown&&(e.imgDown.style.visibility="hidden");e.currentState=i.ButtonState.HOVER}if(t<=i.ButtonState.GROUP&&e.currentState===i.ButtonState.HOVER){e.imgHover&&(e.imgHover.style.visibility="hidden");e.currentState=i.ButtonState.GROUP}if(t<=i.ButtonState.REST&&e.currentState===i.ButtonState.GROUP){!function(e){e.shouldFade=!0;e.fadeBeginTime=i.now()+e.fadeDelay;window.setTimeout(function(){n(e)},e.fadeDelay)}(e);e.currentState=i.ButtonState.REST}}}}(OpenSeadragon);!function(o){o.ButtonGroup=function(e){o.extend(!0,this,{buttons:[],clickTimeThreshold:o.DEFAULT_SETTINGS.clickTimeThreshold,clickDistThreshold:o.DEFAULT_SETTINGS.clickDistThreshold,labelText:""},e);var t,i=this.buttons.concat([]),n=this;this.element=e.element||o.makeNeutralElement("div");if(!e.group){this.element.style.display="inline-block";for(t=0;tu&&(u=m.x);m.yp&&(p=m.y)}return new v.Rect(c,d,u-c,p-d)},_getSegments:function(){var e=this.getTopLeft();var t=this.getTopRight();var i=this.getBottomLeft();var n=this.getBottomRight();return[[e,t],[t,n],[n,i],[i,e]]},rotate:function(e,t){if(0===(e=v.positiveModulo(e,360)))return this.clone();t=t||this.getCenter();var i=this.getTopLeft().rotate(e,t);e=this.getTopRight().rotate(e,t).minus(i);e=e.apply(function(e){return Math.abs(e)<1e-15?0:e});t=Math.atan(e.y/e.x);e.x<0?t+=Math.PI:e.y<0&&(t+=2*Math.PI);return new v.Rect(i.x,i.y,this.width,this.height,t/Math.PI*180)},getBoundingBox:function(){if(0===this.degrees)return this.clone();var e=this.getTopLeft();var t=this.getTopRight();var i=this.getBottomLeft();var n=this.getBottomRight();var o=Math.min(e.x,t.x,i.x,n.x);var r=Math.max(e.x,t.x,i.x,n.x);var s=Math.min(e.y,t.y,i.y,n.y);n=Math.max(e.y,t.y,i.y,n.y);return new v.Rect(o,s,r-o,n-s)},getIntegerBoundingBox:function(){var e=this.getBoundingBox();var t=Math.floor(e.x);var i=Math.floor(e.y);var n=Math.ceil(e.width+e.x-t);e=Math.ceil(e.height+e.y-i);return new v.Rect(t,i,n,e)},containsPoint:function(e,t){t=t||0;var i=this.getTopLeft();var n=this.getTopRight();var o=this.getBottomLeft();var r=n.minus(i);var s=o.minus(i);return(e.x-i.x)*r.x+(e.y-i.y)*r.y>=-t&&(e.x-n.x)*r.x+(e.y-n.y)*r.y<=t&&(e.x-i.x)*s.x+(e.y-i.y)*s.y>=-t&&(e.x-o.x)*s.x+(e.y-o.y)*s.y<=t},toString:function(){return"["+Math.round(100*this.x)/100+", "+Math.round(100*this.y)/100+", "+Math.round(100*this.width)/100+"x"+Math.round(100*this.height)/100+", "+Math.round(100*this.degrees)/100+"deg]"}}}(OpenSeadragon);!function(h){var s={};h.ReferenceStrip=function(e){var t,i,n,o=e.viewer,r=h.getElementSize(o.element);if(!e.id){e.id="referencestrip-"+h.now();this.element=h.makeNeutralElement("div");this.element.id=e.id;this.element.className="referencestrip"}e=h.extend(!0,{sizeRatio:h.DEFAULT_SETTINGS.referenceStripSizeRatio,position:h.DEFAULT_SETTINGS.referenceStripPosition,scroll:h.DEFAULT_SETTINGS.referenceStripScroll,clickTimeThreshold:h.DEFAULT_SETTINGS.clickTimeThreshold},e,{element:this.element});h.extend(this,e);s[this.id]={animating:!1};this.minPixelRatio=this.viewer.minPixelRatio;this.element.tabIndex=0;(i=this.element.style).marginTop="0px";i.marginRight="0px";i.marginBottom="0px";i.marginLeft="0px";i.left="0px";i.bottom="0px";i.border="0px";i.background="#000";i.position="relative";h.setElementTouchActionNone(this.element);h.setElementOpacity(this.element,.8);this.viewer=o;this.tracker=new h.MouseTracker({userData:"ReferenceStrip.tracker",element:this.element,clickHandler:h.delegate(this,a),dragHandler:h.delegate(this,l),scrollHandler:h.delegate(this,c),enterHandler:h.delegate(this,d),leaveHandler:h.delegate(this,p),keyDownHandler:h.delegate(this,g),keyHandler:h.delegate(this,m),preProcessEventHandler:function(e){"wheel"===e.eventType&&(e.preventDefault=!0)}});if(e.width&&e.height){this.element.style.width=e.width+"px";this.element.style.height=e.height+"px";o.addControl(this.element,{anchor:h.ControlAnchor.BOTTOM_LEFT})}else if("horizontal"===e.scroll){this.element.style.width=r.x*e.sizeRatio*o.tileSources.length+12*o.tileSources.length+"px";this.element.style.height=r.y*e.sizeRatio+"px";o.addControl(this.element,{anchor:h.ControlAnchor.BOTTOM_LEFT})}else{this.element.style.height=r.y*e.sizeRatio*o.tileSources.length+12*o.tileSources.length+"px";this.element.style.width=r.x*e.sizeRatio+"px";o.addControl(this.element,{anchor:h.ControlAnchor.TOP_LEFT})}this.panelWidth=r.x*this.sizeRatio+8;this.panelHeight=r.y*this.sizeRatio+8;this.panels=[];this.miniViewers={};for(n=0;ns+n.x-this.panelWidth){t=Math.min(t,o-n.x);this.element.style.marginLeft=-t+"px";u(this,n.x,-t)}else if(ta+n.y-this.panelHeight){t=Math.min(t,r-n.y);this.element.style.marginTop=-t+"px";u(this,n.y,-t)}else if(t-(n-r.x)){this.element.style.marginLeft=t+2*e.delta.x+"px";u(this,r.x,t+2*e.delta.x)}}else if(-e.delta.x<0&&t<0){this.element.style.marginLeft=t+2*e.delta.x+"px";u(this,r.x,t+2*e.delta.x)}}else if(0<-e.delta.y){if(i>-(o-r.y)){this.element.style.marginTop=i+2*e.delta.y+"px";u(this,r.y,i+2*e.delta.y)}}else if(-e.delta.y<0&&i<0){this.element.style.marginTop=i+2*e.delta.y+"px";u(this,r.y,i+2*e.delta.y)}}}function c(e){if(this.element){var t=Number(this.element.style.marginLeft.replace("px","")),i=Number(this.element.style.marginTop.replace("px","")),n=Number(this.element.style.width.replace("px","")),o=Number(this.element.style.height.replace("px","")),r=h.getElementSize(this.viewer.canvas);if("horizontal"===this.scroll){if(0-(n-r.x)){this.element.style.marginLeft=t-60*e.scroll+"px";u(this,r.x,t-60*e.scroll)}}else if(e.scroll<0&&t<0){this.element.style.marginLeft=t-60*e.scroll+"px";u(this,r.x,t-60*e.scroll)}}else if(e.scroll<0){if(i>r.y-o){this.element.style.marginTop=i+60*e.scroll+"px";u(this,r.y,i+60*e.scroll)}}else if(0=this.target.time?t:e+(t-e)*(n=this.springStiffness,i=(this.current.time-this.start.time)/(this.target.time-this.start.time),(1-Math.exp(n*-i))/(1-Math.exp(-n)));var i;var n=this.current.value;this._exponential?this.current.value=Math.exp(i):this.current.value=i;return n!==this.current.value},isAtTargetValue:function(){return this.current.value===this.target.value}}}(OpenSeadragon);!function(n){n.ImageJob=function(e){n.extend(!0,this,{timeout:n.DEFAULT_SETTINGS.timeout,jobId:null,tries:0},e);this.data=null;this.userData={};this.errorMsg=null};n.ImageJob.prototype={start:function(){this.tries++;var e=this;var t=this.abort;this.jobId=window.setTimeout(function(){e.finish(null,null,"Image load exceeded timeout ("+e.timeout+" ms)")},this.timeout);this.abort=function(){e.source.downloadTileAbort(e);"function"==typeof t&&t()};this.source.downloadTileStart(this)},finish:function(e,t,i){this.data=e;this.request=t;this.errorMsg=i;this.jobId&&window.clearTimeout(this.jobId);this.callback(this)}};n.ImageLoader=function(e){n.extend(!0,this,{jobLimit:n.DEFAULT_SETTINGS.imageLoaderLimit,timeout:n.DEFAULT_SETTINGS.timeout,jobQueue:[],failedTiles:[],jobsInProgress:0},e)};n.ImageLoader.prototype={addJob:function(t){if(!t.source){n.console.error("ImageLoader.prototype.addJob() requires [options.source]. TileSource since new API defines how images are fetched. Creating a dummy TileSource.");var e=n.TileSource.prototype;t.source={downloadTileStart:e.downloadTileStart,downloadTileAbort:e.downloadTileAbort}}var i=this,e={src:t.src,tile:t.tile||{},source:t.source,loadWithAjax:t.loadWithAjax,ajaxHeaders:t.loadWithAjax?t.ajaxHeaders:null,crossOriginPolicy:t.crossOriginPolicy,ajaxWithCredentials:t.ajaxWithCredentials,postData:t.postData,callback:function(e){!function(e,t,i){""!==t.errorMsg&&(null===t.data||void 0===t.data)&&t.tries<1+e.tileRetryMax&&e.failedTiles.push(t);var n;e.jobsInProgress--;if((!e.jobLimit||e.jobsInProgressthis.canvas.width&&(r.width=this.canvas.width-r.x);if(r.y<0){r.height+=r.y;r.y=0}r.y+r.height>this.canvas.height&&(r.height=this.canvas.height-r.y);this.context.drawImage(this.sketchCanvas,r.x,r.y,r.width,r.height,r.x,r.y,r.width,r.height)}else{t=o.scale||1;e=(i=o.translate)instanceof a.Point?i:new a.Point(0,0);n=0;r=0;if(i){o=this.sketchCanvas.width-this.canvas.width;i=this.sketchCanvas.height-this.canvas.height;n=Math.round(o/2);r=Math.round(i/2)}this.context.drawImage(this.sketchCanvas,e.x-n*t,e.y-r*t,(this.canvas.width+2*n)*t,(this.canvas.height+2*r)*t,-n,-r,this.canvas.width+2*n,this.canvas.height+2*r)}this.context.restore()}},drawDebugInfo:function(e,t,i,n){if(this.useCanvas){var o=this.viewer.world.getIndexOfItem(n)%this.debugGridColor.length;var r=this.context;r.save();r.lineWidth=2*a.pixelDensityRatio;r.font="small-caps bold "+13*a.pixelDensityRatio+"px arial";r.strokeStyle=this.debugGridColor[o];r.fillStyle=this.debugGridColor[o];this.viewport.getRotation(!0)%360!=0&&this._offsetForRotation({degrees:this.viewport.getRotation(!0)});n.getRotation(!0)%360!=0&&this._offsetForRotation({degrees:n.getRotation(!0),point:n.viewport.pixelFromPointNoRotate(n._getRotationPoint(!0),!0)});n.viewport.getRotation(!0)%360==0&&n.getRotation(!0)%360==0&&n._drawer.viewer.viewport.getFlip()&&n._drawer._flip();r.strokeRect(e.position.x*a.pixelDensityRatio,e.position.y*a.pixelDensityRatio,e.size.x*a.pixelDensityRatio,e.size.y*a.pixelDensityRatio);var s=(e.position.x+e.size.x/2)*a.pixelDensityRatio;o=(e.position.y+e.size.y/2)*a.pixelDensityRatio;r.translate(s,o);r.rotate(Math.PI/180*-this.viewport.getRotation(!0));r.translate(-s,-o);if(0===e.x&&0===e.y){r.fillText("Zoom: "+this.viewport.getZoom(),e.position.x*a.pixelDensityRatio,(e.position.y-30)*a.pixelDensityRatio);r.fillText("Pan: "+this.viewport.getBounds().toString(),e.position.x*a.pixelDensityRatio,(e.position.y-20)*a.pixelDensityRatio)}r.fillText("Level: "+e.level,(e.position.x+10)*a.pixelDensityRatio,(e.position.y+20)*a.pixelDensityRatio);r.fillText("Column: "+e.x,(e.position.x+10)*a.pixelDensityRatio,(e.position.y+30)*a.pixelDensityRatio);r.fillText("Row: "+e.y,(e.position.x+10)*a.pixelDensityRatio,(e.position.y+40)*a.pixelDensityRatio);r.fillText("Order: "+i+" of "+t,(e.position.x+10)*a.pixelDensityRatio,(e.position.y+50)*a.pixelDensityRatio);r.fillText("Size: "+e.size.toString(),(e.position.x+10)*a.pixelDensityRatio,(e.position.y+60)*a.pixelDensityRatio);r.fillText("Position: "+e.position.toString(),(e.position.x+10)*a.pixelDensityRatio,(e.position.y+70)*a.pixelDensityRatio);this.viewport.getRotation(!0)%360!=0&&this._restoreRotationChanges();n.getRotation(!0)%360!=0&&this._restoreRotationChanges();n.viewport.getRotation(!0)%360==0&&n.getRotation(!0)%360==0&&n._drawer.viewer.viewport.getFlip()&&n._drawer._flip();r.restore()}},debugRect:function(e){if(this.useCanvas){var t=this.context;t.save();t.lineWidth=2*a.pixelDensityRatio;t.strokeStyle=this.debugGridColor[0];t.fillStyle=this.debugGridColor[0];t.strokeRect(e.x*a.pixelDensityRatio,e.y*a.pixelDensityRatio,e.width*a.pixelDensityRatio,e.height*a.pixelDensityRatio);t.restore()}},setImageSmoothingEnabled:function(e){if(this.useCanvas){this._imageSmoothingEnabled=e;this._updateImageSmoothingEnabled(this.context);this.viewer.forceRedraw()}},_updateImageSmoothingEnabled:function(e){e.msImageSmoothingEnabled=this._imageSmoothingEnabled;e.imageSmoothingEnabled=this._imageSmoothingEnabled},getCanvasSize:function(e){e=this._getContext(e).canvas;return new a.Point(e.width,e.height)},getCanvasCenter:function(){return new a.Point(this.canvas.width/2,this.canvas.height/2)},_offsetForRotation:function(e){var t=e.point?e.point.times(a.pixelDensityRatio):this.getCanvasCenter();var i=this._getContext(e.useSketch);i.save();i.translate(t.x,t.y);if(this.viewer.viewport.flipped){i.rotate(Math.PI/180*-e.degrees);i.scale(-1,1)}else i.rotate(Math.PI/180*e.degrees);i.translate(-t.x,-t.y)},_flip:function(e){var t=(e=e||{}).point?e.point.times(a.pixelDensityRatio):this.getCanvasCenter();e=this._getContext(e.useSketch);e.translate(t.x,0);e.scale(-1,1);e.translate(-t.x,0)},_restoreRotationChanges:function(e){this._getContext(e).restore()},_calculateCanvasSize:function(){var e=a.pixelDensityRatio;var t=this.viewport.getContainerSize();return{x:Math.round(t.x*e),y:Math.round(t.y*e)}},_calculateSketchCanvasSize:function(){var e=this._calculateCanvasSize();if(0===this.viewport.getRotation())return e;e=Math.ceil(Math.sqrt(e.x*e.x+e.y*e.y));return{x:e,y:e}}}}(OpenSeadragon);!function(h){h.Viewport=function(e){var t=arguments;if((e=t.length&&t[0]instanceof h.Point?{containerSize:t[0],contentSize:t[1],config:t[2]}:e).config){h.extend(!0,e,e.config);delete e.config}this._margins=h.extend({left:0,top:0,right:0,bottom:0},e.margins||{});delete e.margins;e.initialDegrees=e.degrees;delete e.degrees;h.extend(!0,this,{containerSize:null,contentSize:null,zoomPoint:null,rotationPivot:null,viewer:null,springStiffness:h.DEFAULT_SETTINGS.springStiffness,animationTime:h.DEFAULT_SETTINGS.animationTime,minZoomImageRatio:h.DEFAULT_SETTINGS.minZoomImageRatio,maxZoomPixelRatio:h.DEFAULT_SETTINGS.maxZoomPixelRatio,visibilityRatio:h.DEFAULT_SETTINGS.visibilityRatio,wrapHorizontal:h.DEFAULT_SETTINGS.wrapHorizontal,wrapVertical:h.DEFAULT_SETTINGS.wrapVertical,defaultZoomLevel:h.DEFAULT_SETTINGS.defaultZoomLevel,minZoomLevel:h.DEFAULT_SETTINGS.minZoomLevel,maxZoomLevel:h.DEFAULT_SETTINGS.maxZoomLevel,initialDegrees:h.DEFAULT_SETTINGS.degrees,flipped:h.DEFAULT_SETTINGS.flipped,homeFillsViewer:h.DEFAULT_SETTINGS.homeFillsViewer,silenceMultiImageWarnings:h.DEFAULT_SETTINGS.silenceMultiImageWarnings},e);this._updateContainerInnerSize();this.centerSpringX=new h.Spring({initial:0,springStiffness:this.springStiffness,animationTime:this.animationTime});this.centerSpringY=new h.Spring({initial:0,springStiffness:this.springStiffness,animationTime:this.animationTime});this.zoomSpring=new h.Spring({exponential:!0,initial:1,springStiffness:this.springStiffness,animationTime:this.animationTime});this.degreesSpring=new h.Spring({initial:e.initialDegrees,springStiffness:this.springStiffness,animationTime:this.animationTime});this._oldCenterX=this.centerSpringX.current.value;this._oldCenterY=this.centerSpringY.current.value;this._oldZoom=this.zoomSpring.current.value;this._oldDegrees=this.degreesSpring.current.value;this._setContentBounds(new h.Rect(0,0,1,1),1);this.goHome(!0);this.update()};h.Viewport.prototype={get degrees(){h.console.warn("Accessing [Viewport.degrees] is deprecated. Use viewport.getRotation instead.");return this.getRotation()},set degrees(e){h.console.warn("Setting [Viewport.degrees] is deprecated. Use viewport.rotateTo, viewport.rotateBy, or viewport.setRotation instead.");this.rotateTo(e)},resetContentSize:function(e){h.console.assert(e,"[Viewport.resetContentSize] contentSize is required");h.console.assert(e instanceof h.Point,"[Viewport.resetContentSize] contentSize must be an OpenSeadragon.Point");h.console.assert(0i.width?this.visibilityRatio*i.width:this.visibilityRatio*t.width;r=i.x-r+a;s=s-t.x-a;if(a>i.width){t.x+=(r+s)/2;n=!0}else if(s<0){t.x+=s;n=!0}else if(0i.height?this.visibilityRatio*i.height:this.visibilityRatio*t.height;l=i.y-l+r;s=s-t.y-r;if(r>i.height){t.y+=(l+s)/2;o=!0}else if(s<0){t.y+=s;o=!0}else if(0=o?s.height=s.width/o:s.width=s.height*o;s.x=r.x-s.width/2;s.y=r.y-s.height/2;var a=1/s.width;if(i){this.panTo(r,!0);this.zoomTo(a,null,!0);n&&this.applyConstraints(!0);return this}var l=this.getCenter(!0);t=this.getZoom(!0);this.panTo(l,!0);this.zoomTo(t,null,!0);e=this.getBounds();o=this.getZoom();if(0===o||Math.abs(a/o-1)<1e-8){this.zoomTo(a,null,!0);this.panTo(r,i);n&&this.applyConstraints(!1);return this}if(n){this.panTo(r,!1);a=this._applyZoomConstraints(a);this.zoomTo(a,null,!1);r=this.getConstrainedBounds();this.panTo(l,!0);this.zoomTo(t,null,!0);this.fitBounds(r)}else{o=s.rotate(-this.getRotation()).getTopLeft().times(a).minus(e.getTopLeft().times(o)).divide(a-o);this.zoomTo(a,o,i)}return this},fitBounds:function(e,t){return this._fitBounds(e,{immediately:t,constraints:!1})},fitBoundsWithConstraints:function(e,t){return this._fitBounds(e,{immediately:t,constraints:!0})},fitVertically:function(e){var t=new h.Rect(this._contentBounds.x+this._contentBounds.width/2,this._contentBounds.y,0,this._contentBounds.height);return this.fitBounds(t,e)},fitHorizontally:function(e){var t=new h.Rect(this._contentBounds.x,this._contentBounds.y+this._contentBounds.height/2,this._contentBounds.width,0);return this.fitBounds(t,e)},getConstrainedBounds:function(e){e=this.getBounds(e);return this._applyBoundaryConstraints(e)},panBy:function(e,t){var i=new h.Point(this.centerSpringX.target.value,this.centerSpringY.target.value);return this.panTo(i.plus(e),t)},panTo:function(e,t){if(t){this.centerSpringX.resetTo(e.x);this.centerSpringY.resetTo(e.y)}else{this.centerSpringX.springTo(e.x);this.centerSpringY.springTo(e.y)}this.viewer&&this.viewer.raiseEvent("pan",{center:e,immediately:t});return this},zoomBy:function(e,t,i){return this.zoomTo(this.zoomSpring.target.value*e,t,i)},zoomTo:function(e,t,i){var n=this;this.zoomPoint=t instanceof h.Point&&!isNaN(t.x)&&!isNaN(t.y)?t:null;i?this._adjustCenterSpringsForZoomPoint(function(){n.zoomSpring.resetTo(e)}):this.zoomSpring.springTo(e);this.viewer&&this.viewer.raiseEvent("zoom",{zoom:e,refPoint:t,immediately:i});return this},setRotation:function(e,t){return this.rotateTo(e,null,t)},getRotation:function(e){return(e?this.degreesSpring.current:this.degreesSpring.target).value},setRotationWithPivot:function(e,t,i){return this.rotateTo(e,t,i)},rotateTo:function(e,t,i){if(!this.viewer||!this.viewer.drawer.canRotate())return this;if(this.degreesSpring.target.value===e&&this.degreesSpring.isAtTargetValue())return this;this.rotationPivot=t instanceof h.Point&&!isNaN(t.x)&&!isNaN(t.y)?t:null;if(i)if(this.rotationPivot){if(!(e-this._oldDegrees)){this.rotationPivot=null;return this}this._rotateAboutPivot(e)}else this.degreesSpring.resetTo(e);else{var n=h.positiveModulo(this.degreesSpring.current.value,360);var o=h.positiveModulo(e,360);t=o-n;180o){r=this._clip.x/this._clip.height*e.height;s=this._clip.y/this._clip.height*e.height}else{r=this._clip.x/this._clip.width*e.width;s=this._clip.y/this._clip.width*e.width}}if(e.getAspectRatio()>o){var l=e.height/t;t=0;n.isHorizontallyCentered?t=(e.width-e.height*o)/2:n.isRight&&(t=e.width-e.height*o);this.setPosition(new y.Point(e.x-r+t,e.y-s),i);this.setHeight(l,i)}else{l=e.width/a;a=0;n.isVerticallyCentered?a=(e.height-e.width/o)/2:n.isBottom&&(a=e.height-e.width/o);this.setPosition(new y.Point(e.x-r,e.y-s+a),i);this.setWidth(l,i)}},getClip:function(){return this._clip?this._clip.clone():null},setClip:function(e){y.console.assert(!e||e instanceof y.Rect,"[TiledImage.setClip] newClip must be an OpenSeadragon.Rect or null");e instanceof y.Rect?this._clip=e.clone():this._clip=null;this._needsDraw=!0;this.raiseEvent("clip-change")},getFlip:function(){return!!this.flipped},setFlip:function(e){this.flipped=!!e;this._needsDraw=!0;this._raiseBoundsChange()},getOpacity:function(){return this.opacity},setOpacity:function(e){if(e!==this.opacity){this.opacity=e;this._needsDraw=!0;this.raiseEvent("opacity-change",{opacity:this.opacity})}},getPreload:function(){return this._preload},setPreload:function(e){this._preload=!!e;this._needsDraw=!0},getRotation:function(e){return(e?this._degreesSpring.current:this._degreesSpring.target).value},setRotation:function(e,t){if(this._degreesSpring.target.value!==e||!this._degreesSpring.isAtTargetValue()){t?this._degreesSpring.resetTo(e):this._degreesSpring.springTo(e);this._needsDraw=!0;this._raiseBoundsChange()}},_getRotationPoint:function(e){return this.getBoundsNoRotate(e).getCenter()},getCompositeOperation:function(){return this.compositeOperation},setCompositeOperation:function(e){if(e!==this.compositeOperation){this.compositeOperation=e;this._needsDraw=!0;this.raiseEvent("composite-operation-change",{compositeOperation:this.compositeOperation})}},setAjaxHeaders:function(e,t){if(y.isPlainObject(e=null===e?{}:e)){this._ownAjaxHeaders=e;this._updateAjaxHeaders(t)}else console.error("[TiledImage.setAjaxHeaders] Ignoring invalid headers, must be a plain object")},_updateAjaxHeaders:function(e){void 0===e&&(e=!0);y.isPlainObject(this.viewer.ajaxHeaders)?this.ajaxHeaders=y.extend({},this.viewer.ajaxHeaders,this._ownAjaxHeaders):this.ajaxHeaders=this._ownAjaxHeaders;if(e){var t,i;for(var n in this.tilesMatrix){t=this.source.getNumTiles(n);for(var o in this.tilesMatrix[n]){i=(t.x+o%t.x)%t.x;for(var r in this.tilesMatrix[n][o]){s=(t.y+r%t.y)%t.y;(r=this.tilesMatrix[n][o][r]).loadWithAjax=this.loadTilesWithAjax;if(r.loadWithAjax){var s=this.source.getTileAjaxHeaders(n,i,s);r.ajaxHeaders=y.extend({},this.ajaxHeaders,s)}else r.ajaxHeaders=null}}}for(var a=0;a=this.minPixelRatio)r=l=!0;else if(!r)continue;var c=e.deltaPixelsFromPointsNoRotate(this.source.getPixelRatio(a),!1).x*this._scaleSpring.current.value;var u=e.deltaPixelsFromPointsNoRotate(this.source.getPixelRatio(Math.max(this.source.getClosestLevel(),0)),!1).x*this._scaleSpring.current.value;u=this.immediateRender?1:u;h=Math.min(1,(h-.5)/.5);c=u/Math.abs(u-c);o=this._updateLevel(r,l,a,h,c,t,s,o);if(this._providesCoverage(this.coverage,a))break}this._drawTiles(this.lastDrawn);if(o&&!o.context2D){this._loadTile(o,s);this._needsDraw=!0;this._setFullyLoaded(!1)}else this._setFullyLoaded(0===this._tilesLoading)},_getCornerTiles:function(e,t,i){var n;var o;if(this.wrapHorizontal){n=y.positiveModulo(t.x,1);o=y.positiveModulo(i.x,1)}else{n=Math.max(0,t.x);o=Math.min(1,i.x)}var r=1/this.source.aspectRatio;if(this.wrapVertical){s=y.positiveModulo(t.y,r);a=y.positiveModulo(i.y,r)}else{s=Math.max(0,t.y);a=Math.min(r,i.y)}var s=this.source.getTileAtPoint(e,new y.Point(n,s));var a=this.source.getTileAtPoint(e,new y.Point(o,a));e=this.source.getNumTiles(e);if(this.wrapHorizontal){s.x+=e.x*Math.floor(t.x);a.x+=e.x*Math.floor(i.x)}if(this.wrapVertical){s.y+=e.y*Math.floor(t.y/r);a.y+=e.y*Math.floor(i.y/r)}return{topLeft:s,bottomRight:a}},_updateLevel:function(e,t,i,n,o,r,s,a){var l=r.getBoundingBox().getTopLeft();var h=r.getBoundingBox().getBottomRight();this.viewer&&this.viewer.raiseEvent("update-level",{tiledImage:this,havedrawn:e,level:i,opacity:n,visibility:o,drawArea:r,topleft:l,bottomright:h,currenttime:s,best:a});this._resetCoverage(this.coverage,i);this._resetCoverage(this.loadingCoverage,i);h=this._getCornerTiles(i,l,h);var c=h.topLeft;var u=h.bottomRight;var d=this.source.getNumTiles(i);var p=this.viewport.pixelFromPoint(this.viewport.getCenter());if(this.getFlip()){u.x+=1;this.wrapHorizontal||(u.x=Math.min(u.x,d.x-1))}for(var g=c.x;g<=u.x;g++)for(var m=c.y;m<=u.y;m++){if(this.getFlip()){var v=(d.x+g%d.x)%d.x;v=g+d.x-v-v-1}else v=g;null!==r.intersection(this.getTileBounds(i,v,m))&&(a=this._updateTile(t,e,v,m,i,n,o,p,d,s,a))}return a},_updateTile:function(e,t,i,n,o,r,s,a,l,h,c){var u=this._getTile(i,n,o,h,l,this._worldWidthCurrent,this._worldHeightCurrent),l=t;this.viewer&&this.viewer.raiseEvent("update-tile",{tiledImage:this,tile:u});this._setCoverage(this.coverage,o,i,n,!1);t=u.loaded||u.loading||this._isCovered(this.loadingCoverage,o,i,n);this._setCoverage(this.loadingCoverage,o,i,n,t);if(!u.exists)return c;e&&!l&&(this._isCovered(this.coverage,o,i,n)?this._setCoverage(this.coverage,o,i,n,!0):l=!0);if(!l)return c;this._positionTile(u,this.source.tileOverlap,this.viewport,a,s);if(!u.loaded)if(u.context2D)this._setTileLoaded(u);else{s=this._tileCache.getImageRecord(u.cacheKey);s&&this._setTileLoaded(u,s.getData())}u.loaded?this._blendTile(u,i,n,o,r,h)&&(this._needsDraw=!0):u.loading?this._tilesLoading++:t||(c=this._compareTiles(c,u));return c},_getTile:function(e,t,i,n,o,r,s){var a,l,h,c,u,d,p,g,m,v=this.tilesMatrix,f=this.source;v[i]||(v[i]={});v[i][e]||(v[i][e]={});if(!v[i][e][t]||!v[i][e][t].flipped!=!this.flipped){a=(o.x+e%o.x)%o.x;l=(o.y+t%o.y)%o.y;h=this.getTileBounds(i,e,t);c=f.getTileBounds(i,a,l,!0);u=f.tileExists(i,a,l);d=f.getTileUrl(i,a,l);m=f.getTilePostData(i,a,l);if(this.loadTilesWithAjax){p=f.getTileAjaxHeaders(i,a,l);y.isPlainObject(this.ajaxHeaders)&&(p=y.extend({},this.ajaxHeaders,p))}else p=null;g=f.getContext2D?f.getContext2D(i,a,l):void 0;m=new y.Tile(i,e,t,h,u,d,g,this.loadTilesWithAjax,p,c,m,f.getTileHashKey(i,a,l,d,p,m));this.getFlip()?0==a&&(m.isRightMost=!0):a==o.x-1&&(m.isRightMost=!0);l==o.y-1&&(m.isBottomMost=!0);m.flipped=this.flipped;v[i][e][t]=m}(m=v[i][e][t]).lastTouchTime=n;return m},_loadTile:function(n,o){var r=this;n.loading=!0;this._imageLoader.addJob({src:n.getUrl(),tile:n,source:this.source,postData:n.postData,loadWithAjax:n.loadWithAjax,ajaxHeaders:n.ajaxHeaders,crossOriginPolicy:this.crossOriginPolicy,ajaxWithCredentials:this.ajaxWithCredentials,callback:function(e,t,i){r._onTileLoad(n,o,e,t,i)},abort:function(){n.loading=!1}})},_onTileLoad:function(t,e,i,n,o){if(i){t.exists=!0;if(ee.visibility||t.visibility===e.visibility&&t.squaredDistancethis.smoothTileEdgesMinZoom&&!this.iOSDevice&&this.getRotation(!0)%360==0&&y.supportsCanvas&&this.viewer.useCanvas){i=!0;n=t.getScaleForEdgeSmoothing();o=t.getTranslationForEdgeSmoothing(n,this._drawer.getCanvasSize(!1),this._drawer.getCanvasSize(!0))}var a;if(i){if(!n){a=this.viewport.viewportToViewerElementRectangle(this.getClippedBounds(!0)).getIntegerBoundingBox();this._drawer.viewer.viewport.getFlip()&&(this.viewport.getRotation(!0)%360==0&&this.getRotation(!0)%360==0||(a.x=this._drawer.viewer.container.clientWidth-(a.x+a.width)));a=a.times(y.pixelDensityRatio)}this._drawer._clear(!0,a)}if(!n){this.viewport.getRotation(!0)%360!=0&&this._drawer._offsetForRotation({degrees:this.viewport.getRotation(!0),useSketch:i});this.getRotation(!0)%360!=0&&this._drawer._offsetForRotation({degrees:this.getRotation(!0),point:this.viewport.pixelFromPointNoRotate(this._getRotationPoint(!0),!0),useSketch:i});this.viewport.getRotation(!0)%360==0&&this.getRotation(!0)%360==0&&this._drawer.viewer.viewport.getFlip()&&this._drawer._flip()}r=!1;if(this._clip){this._drawer.saveContext(i);s=this.imageToViewportRectangle(this._clip,!0);s=s.rotate(-this.getRotation(!0),this._getRotationPoint(!0));s=this._drawer.viewportToDrawerRectangle(s);n&&(s=s.times(n));o&&(s=s.translate(o));this._drawer.setClip(s,i);r=!0}if(this._croppingPolygons){var l=this;this._drawer.saveContext(i);try{var h=this._croppingPolygons.map(function(e){return e.map(function(e){e=l.imageToViewportCoordinates(e.x,e.y,!0).rotate(-l.getRotation(!0),l._getRotationPoint(!0));e=l._drawer.viewportCoordToDrawerCoord(e);n&&(e=e.times(n));return e=o?e.plus(o):e})});this._drawer.clipWithPolygons(h,i)}catch(e){y.console.error(e)}r=!0}if(this.placeholderFillStyle&&!1===this._hasOpaqueTile){h=this._drawer.viewportToDrawerRectangle(this.getBounds(!0));n&&(h=h.times(n));o&&(h=h.translate(o));var c=null;c="function"==typeof this.placeholderFillStyle?this.placeholderFillStyle(this,this._drawer.context):this.placeholderFillStyle;this._drawer.drawRectangle(h,c,i)}c=function(e){if("number"==typeof e)return m(e);if(!e||!y.Browser)return p;var t=e[y.Browser.vendor];g(t)&&(t=e["*"]);return m(t)}(this.subPixelRoundingForTransparency);var u=!1;c===y.SUBPIXEL_ROUNDING_OCCURRENCES.ALWAYS?u=!0:c===y.SUBPIXEL_ROUNDING_OCCURRENCES.ONLY_AT_REST&&(u=!(this.viewer&&this.viewer.isAnimating()));for(var d=e.length-1;0<=d;d--){t=e[d];this._drawer.drawTile(t,this._drawingHandler,i,n,o,u,this.source);t.beingDrawn=!0;this.viewer&&this.viewer.raiseEvent("tile-drawn",{tiledImage:this,tile:t})}r&&this._drawer.restoreContext(i);if(!n){this.getRotation(!0)%360!=0&&this._drawer._restoreRotationChanges(i);this.viewport.getRotation(!0)%360!=0&&this._drawer._restoreRotationChanges(i)}if(i){if(n){this.viewport.getRotation(!0)%360!=0&&this._drawer._offsetForRotation({degrees:this.viewport.getRotation(!0),useSketch:!1});this.getRotation(!0)%360!=0&&this._drawer._offsetForRotation({degrees:this.getRotation(!0),point:this.viewport.pixelFromPointNoRotate(this._getRotationPoint(!0),!0),useSketch:!1})}this._drawer.blendSketch({opacity:this.opacity,scale:n,translate:o,compositeOperation:this.compositeOperation,bounds:a});if(n){this.getRotation(!0)%360!=0&&this._drawer._restoreRotationChanges(!1);this.viewport.getRotation(!0)%360!=0&&this._drawer._restoreRotationChanges(!1)}}n||this.viewport.getRotation(!0)%360==0&&this.getRotation(!0)%360==0&&this._drawer.viewer.viewport.getFlip()&&this._drawer._flip();this._drawDebugInfo(e)}},_drawDebugInfo:function(e){if(this.debugMode)for(var t=e.length-1;0<=t;t--){var i=e[t];try{this._drawer.drawDebugInfo(i,e.length,t,this)}catch(e){y.console.error(e)}}},_providesCoverage:function(e,t,i,n){var o,r,s,a;if(!e[t])return!1;if(void 0!==i&&void 0!==n)return void 0===e[t][i]||void 0===e[t][i][n]||!0===e[t][i][n];for(s in o=e[t])if(Object.prototype.hasOwnProperty.call(o,s))for(a in r=o[s])if(Object.prototype.hasOwnProperty.call(r,a)&&!r[a])return!1;return!0},_isCovered:function(e,t,i,n){return void 0===i||void 0===n?this._providesCoverage(e,t+1):this._providesCoverage(e,t+1,2*i,2*n)&&this._providesCoverage(e,t+1,2*i,2*n+1)&&this._providesCoverage(e,t+1,2*i+1,2*n)&&this._providesCoverage(e,t+1,2*i+1,2*n+1)},_setCoverage:function(e,t,i,n,o){if(e[t]){e[t][i]||(e[t][i]={});e[t][i][n]=o}else y.console.warn("Setting coverage for a tile before its level's coverage has been reset: %s",t)},_resetCoverage:function(e,t){e[t]={}}});var p=y.SUBPIXEL_ROUNDING_OCCURRENCES.NEVER;function g(e){return e!==y.SUBPIXEL_ROUNDING_OCCURRENCES.ALWAYS&&e!==y.SUBPIXEL_ROUNDING_OCCURRENCES.ONLY_AT_REST&&e!==y.SUBPIXEL_ROUNDING_OCCURRENCES.NEVER}function m(e){return g(e)?p:e}}(OpenSeadragon);!function(g){function m(e){g.console.assert(e,"[TileCache.cacheTile] options is required");g.console.assert(e.tile,"[TileCache.cacheTile] options.tile is required");g.console.assert(e.tiledImage,"[TileCache.cacheTile] options.tiledImage is required");this.tile=e.tile;this.tiledImage=e.tiledImage}function v(e){g.console.assert(e,"[ImageRecord] options is required");g.console.assert(e.data,"[ImageRecord] options.data is required");this._tiles=[];e.create.apply(null,[this,e.data,e.ownerTile]);this._destroyImplementation=e.destroy.bind(null,this);this.getImage=e.getImage.bind(null,this);this.getData=e.getData.bind(null,this);this.getRenderedContext=e.getRenderedContext.bind(null,this)}v.prototype={destroy:function(){this._destroyImplementation();this._tiles=null},addTile:function(e){g.console.assert(e,"[ImageRecord.addTile] tile is required");this._tiles.push(e)},removeTile:function(e){for(var t=0;tthis._maxImageCacheCount){var o=null;var r=-1;var s=null;var a,l,h,c,u,d;for(var p=this._tilesLoaded.length-1;0<=p;p--)if(!((a=(d=this._tilesLoaded[p]).tile).level<=t||a.beingDrawn))if(o){c=a.lastTouchTime;l=o.lastTouchTime;u=a.level;h=o.level;if(c=this._items.length)throw new Error("Index bigger than number of layers.");if(t!==i&&-1!==i){this._items.splice(i,1);this._items.splice(t,0,e);this._needsDraw=!0;this.raiseEvent("item-index-change",{item:e,previousIndex:i,newIndex:t})}},removeItem:function(e){g.console.assert(e,"[World.removeItem] item is required");var t=g.indexOf(this._items,e);if(-1!==t){e.removeHandler("bounds-change",this._delegatedFigureSizes);e.removeHandler("clip-change",this._delegatedFigureSizes);e.destroy();this._items.splice(t,1);this._figureSizes();this._needsDraw=!0;this._raiseRemoveItem(e)}},removeAll:function(){this.viewer._cancelPendingImages();var e;var t;for(t=0;td.height?r:r*(d.width/d.height))*(d.height/d.width);d=new g.Point(l+(r-u)/2,h+(r-d)/2);c.setPosition(d,t);c.setWidth(u,t);"horizontal"===i?l+=s:h+=s}this.setAutoRefigureSizes(!0)},_figureSizes:function(){var e=this._homeBounds?this._homeBounds.clone():null;var t=this._contentSize?this._contentSize.clone():null;var i=this._contentFactor||0;if(this._items.length){var n=this._items[0];var o=n.getBounds();this._contentFactor=n.getContentSize().x/o.width;var r=n.getClippedBounds().getBoundingBox();var s=r.x;var a=r.y;var l=r.x+r.width;var h=r.y+r.height;for(var c=1;c \ No newline at end of file diff --git a/src/ops_model/models/attention/diffex/viewer/webapp/style.css b/src/ops_model/models/attention/diffex/viewer/webapp/style.css new file mode 100644 index 0000000..3b8dcae --- /dev/null +++ b/src/ops_model/models/attention/diffex/viewer/webapp/style.css @@ -0,0 +1,521 @@ +:root{--bg:#0d0f13;--panel:#161a21;--line:#2a2f3a;--fg:#e6e8ec;--mut:#8b93a1; + --neg:#f0a020;--mid:#26c6ff;--pos:#ff5252;--accent:#26c6ff;} +*{box-sizing:border-box} +body{margin:0;background:var(--bg);color:var(--fg);display:flex;flex-direction:column;height:100vh;overflow:hidden; + font:14px/1.4 -apple-system,Segoe UI,Roboto,Helvetica,Arial,sans-serif} +header{padding:8px 20px;border-bottom:1px solid var(--line);display:flex; + align-items:center;justify-content:space-between;gap:16px} +header h1{margin:0;font-size:18px;font-weight:650} +header .sub{color:var(--mut);font-size:12px} +.toggle{background:var(--panel);color:var(--fg);border:1px solid var(--line);border-radius:6px; + padding:8px 12px;font-size:13px;cursor:pointer;white-space:nowrap} +.toggle:hover{border-color:var(--accent)} +.hdr-right{display:flex;align-items:center;gap:14px} +.asset-ver{display:inline-flex;align-items:center;gap:5px} +.asset-ver .av-label{font-size:11px;text-transform:uppercase;letter-spacing:.06em;color:var(--mut)} +.verbtn{background:var(--panel);color:var(--mut);border:1px solid var(--line);border-radius:6px; + padding:4px 11px;font-size:12px;font-weight:600;cursor:pointer} +.verbtn:hover{border-color:var(--accent)} +.verbtn.active{background:var(--accent);color:var(--bg);border-color:var(--accent)} +#sidebar{width:340px;flex:none;border-left:1px solid var(--line);padding:18px;overflow:auto} +#sidebar.hidden{display:none} +.side-hd{font-size:11px;text-transform:uppercase;letter-spacing:.06em;color:var(--mut); + border-bottom:1px solid var(--line);padding-bottom:8px;margin-bottom:12px} +#info-title{font-size:20px;font-weight:650;line-height:1.25} +#info-sub{color:var(--mut);font-size:12px;font-family:ui-monospace,monospace;margin:4px 0 14px} +#info-body{font-size:13px;line-height:1.55} +#info-body .sec{margin-top:12px} +#info-body .sec-lbl{font-size:11px;text-transform:uppercase;letter-spacing:.04em;color:var(--accent);margin-bottom:3px} +#info-body .empty{color:var(--mut);font-style:italic} +#info-body a{color:var(--accent);text-decoration:none} +#info-body a:hover{text-decoration:underline} +#app{display:flex;flex:1;min-height:0} +/* top crossbar: the core selection controls, laid out horizontally */ +#topbar{display:flex;gap:18px;padding:12px 20px;border-bottom:1px solid var(--line); + background:var(--panel);align-items:flex-start;flex-wrap:wrap;flex:none} +.tb-field{display:flex;flex-direction:column;gap:6px;min-width:160px} +.tb-field:has(#markerfilter){min-width:240px;max-width:320px} /* Marker search: ~50% wider than default, capped so it doesn't squish the rest */ +.tb-field:has(#grain){min-width:0} /* Type: sizes to its pills — adapts to count (2 public / 4 internal) + label length */ +.tb-field:has(#grain) .seg{flex:0 0 auto;padding:4px 15px} /* pills hug their (variable-length) labels with roomy padding */ +.tb-field.tb-grow{flex:0.6;min-width:170px;max-width:340px} +.tb-field>label{font-size:11px;text-transform:uppercase;letter-spacing:.04em;color:var(--mut)} +.tb-field .hint{text-transform:none;letter-spacing:0;font-size:11px;opacity:.7} +#topbar select,#topbar input{width:100%} +#topbar select.mini{width:auto;display:inline-block;padding:2px 4px;font-size:11px;margin-left:4px} +#topbar .row{display:flex;gap:6px;align-items:center} +#topbar .row #cellcount{width:56px} +#topbar .row button{flex:none;padding:7px 11px} +#topbar .row .combo{flex:1;min-width:0} +/* combobox: input shows the selection; focus reveals a scrollable (~10-row) filtered list */ +.combo{position:relative} +.combo-list{position:absolute;top:100%;left:0;right:0;z-index:60;margin-top:2px; + background:var(--panel);border:1px solid var(--accent);border-radius:6px; + max-height:264px;overflow:auto;box-shadow:0 10px 24px rgba(0,0,0,.55)} +.combo-list.hidden{display:none} +.combo-item{padding:6px 9px;font:13px/1.3 ui-monospace,Menlo,Consolas,monospace;cursor:pointer; + white-space:nowrap;overflow:hidden;text-overflow:ellipsis} +.combo-item:hover{background:var(--line)} +.combo-item.sel{color:var(--accent)} +.combo-empty{padding:8px 9px;color:var(--mut);font-style:italic;font-size:12px} +#controls{width:290px;padding:16px;border-right:1px solid var(--line); + display:flex;flex-direction:column;gap:14px;overflow:hidden} /* pane scrolls, tabs stay put */ +/* full-width view-tab bar under the crossbar — connected "page" tabs: the active tab takes the page + surface + an accent highlight and its bottom edge dissolves into the content below (not a button) */ +#tabbar{display:flex;gap:3px;flex-wrap:wrap;align-items:flex-end;flex:none;padding:10px 18px 0; + background:linear-gradient(180deg,rgba(22,27,34,.6),rgba(13,15,19,.35)); + backdrop-filter:blur(14px) saturate(1.3);-webkit-backdrop-filter:blur(14px) saturate(1.3); + border-bottom:1px solid var(--line);position:relative;z-index:6} +#tabbar .tab{appearance:none;background:transparent;border:1px solid transparent;border-radius:12px 12px 0 0; + color:var(--mut);padding:11px 22px;font:600 14px ui-sans-serif,system-ui,sans-serif;white-space:nowrap; + cursor:pointer;margin-bottom:-1px;transition:color .15s,background .15s} +#tabbar .tab:hover:not(.active){color:var(--fg);background:rgba(255,255,255,.045)} +#tabbar .tab.active{color:var(--fg);background:var(--bg);border-color:var(--line);border-top:2px solid var(--accent); + border-bottom-color:var(--bg);box-shadow:0 -3px 14px rgba(0,0,0,.32)} +.tabpane{display:flex;flex-direction:column;gap:14px;flex:1;min-height:0;overflow-y:auto;padding-right:4px} +.tabpane.hidden{display:none} +#controls label{display:flex;flex-direction:column;gap:6px;font-size:12px;color:var(--mut); + text-transform:uppercase;letter-spacing:.04em} +#controls .hint{text-transform:none;letter-spacing:0;font-size:11px;opacity:.7} +.desc{font-size:11px;line-height:1.45;color:var(--mut);background:var(--panel); + border:1px solid var(--line);border-radius:6px;padding:8px;max-height:150px;overflow:auto} +.desc:empty{display:none} +select,input[type=text],input[type=number]{background:var(--panel);color:var(--fg); + border:1px solid var(--line);border-radius:6px;padding:8px;font-size:13px;outline:none} +select:focus,input:focus{border-color:var(--accent)} +#target{font-family:ui-monospace,Menlo,Consolas,monospace} +.grp{display:flex;flex-direction:column;gap:8px;border-top:1px solid var(--line);padding-top:12px} +.grp .lbl{font-size:12px;color:var(--mut);text-transform:uppercase;letter-spacing:.04em} +.row{display:flex;gap:8px} +button{background:var(--panel);color:var(--fg);border:1px solid var(--line);border-radius:6px; + padding:7px 10px;font-size:13px;cursor:pointer} +button:hover{border-color:var(--accent)} +.hidden{display:none!important} +#panellist,#a-panellist{list-style:none;margin:6px 0 0;padding:0;font-size:12px;font-family:ui-monospace,monospace} +#panellist li,#a-panellist li{display:flex;justify-content:space-between;gap:6px;padding:4px 6px;background:var(--panel); + border:1px solid var(--line);border-radius:5px;margin-bottom:4px} +#panellist li button,#a-panellist li button{padding:0 6px;font-size:12px} + +#stage{flex:1;display:flex;flex-direction:column;align-items:center;gap:14px;padding:18px;overflow:auto} +#heatbar{width:min(80vw,720px);margin-top:22px;box-sizing:border-box;padding:0 118px 0 52px} /* inset to the α-slider track: play 38+gap14 left, α-read 104+gap14 right */ +#heat-track{position:relative} /* tick/real % map to this (= slider width) so the scrub thumb + colorbar tick move in lockstep */ +#heat-real{position:absolute;top:-19px;transform:translateX(-50%);text-align:center;pointer-events:none} +#heat-real span{font-size:10px;color:#fff;background:rgba(0,0,0,.65);padding:1px 5px;border-radius:4px;white-space:nowrap} +#heat-real::after{content:"";display:block;width:2px;height:16px;background:#fff;opacity:.75;margin:2px auto 0} +#heat-grad{height:14px;border-radius:7px; + background:linear-gradient(90deg,var(--neg) 0%,var(--mid) 50%,var(--pos) 100%)} +#heat-tick{position:absolute;top:-3px;width:3px;height:20px;background:#fff;border-radius:2px; + box-shadow:0 0 4px #000;transition:left .05s linear} +.heat-lbls{display:flex;justify-content:space-between;font-size:11px;margin-top:3px} +.heat-lbls .neg{color:var(--neg)}.heat-lbls .mid{color:var(--mid)}.heat-lbls .pos{color:var(--pos)} +#scrub{display:flex;align-items:center;gap:14px;width:min(80vw,720px)} +.slider-col{flex:1;display:flex;flex-direction:column;gap:2px} +#scrub #alpha{width:100%;accent-color:var(--accent)} +#ticks{position:relative;height:20px} +#ticks .rangespan{position:absolute;top:2px;height:4px;background:var(--accent);opacity:.25;border-radius:2px} +.tick{position:absolute;top:0;width:18px;height:20px;transform:translateX(-50%);cursor:pointer;z-index:1} +.tick::before{content:"";position:absolute;left:50%;top:0;width:2px;height:8px;background:var(--mut); + opacity:.5;transform:translateX(-50%);border-radius:1px} +.tick:hover::before{opacity:1;background:var(--fg);height:13px} +.tick.on::before{height:13px;width:3px;opacity:1;background:var(--accent)} +.tick.on::after{content:"●";position:absolute;left:50%;top:12px;transform:translateX(-50%);font-size:8px;color:var(--accent)} +.tickhint{font-size:11px;color:var(--mut);width:min(80vw,720px)} +#play{width:38px;height:34px;font-size:15px} +#alpha-read{font-family:ui-monospace,monospace;min-width:104px;text-align:right;color:var(--mut)} +.playctl{display:flex;gap:8px;width:min(80vw,720px);justify-content:flex-end} +.playctl .seg{flex:0 0 auto;padding:4px 11px} /* speed pills size to their text (0.25×/0.5× fully visible); group grows leftward since playctl is right-justified */ + +#grid{display:flex;flex-direction:column;gap:20px;width:100%;max-width:1100px} +#grid.cols-layout{display:flex;flex-direction:row;flex-wrap:nowrap;align-items:flex-start;max-width:none;width:100%;overflow-x:auto} /* perturbations ALWAYS side-by-side; scroll if many */ +#grid.cols-layout .group{flex:0 0 auto;width:var(--tilepx,170px)} /* each column fixed to tile width (title can't force it wide) */ +#grid.cols-layout .group-cells{grid-template-columns:var(--tilepx,170px)} /* one vertical column of cells */ +#grid.cols-layout .group-hd{display:flex;flex-wrap:wrap;align-items:center;justify-content:center;text-align:center;white-space:normal;word-break:break-word} /* same font as rows layout; wrap+center within the column */ +#grid.cols-layout .setacc{display:block;margin:3px auto 0} /* score chip drops below the title so it can't stretch the narrow column */ +.group{display:flex;flex-direction:column;gap:8px} +.group-hd{font:650 20px ui-monospace,Menlo,monospace} /* single header per group, ~2× old caption */ +.rowlbl{font-size:11px;color:var(--mut);text-transform:uppercase;letter-spacing:.04em} +.group-cells{display:grid;gap:10px;grid-template-columns:repeat(auto-fill,minmax(var(--tilepx,170px),1fr))} /* image-scale slider drives --tilepx: smaller = more per row + more rows */ +.panel{position:relative} +.panel img{width:100%;aspect-ratio:1;background:#000;border:1px solid var(--line);border-radius:6px;display:block} +.panel.lead img{border-left-width:4px} /* colour bar only on the group's left-most image */ +.setacc{margin-left:10px;font:800 17px ui-monospace,monospace;padding:2px 8px;border-radius:6px;vertical-align:middle;box-shadow:0 0 3px rgba(0,0,0,.4)} +.badge{position:absolute;top:6px;right:6px;font:700 13px ui-monospace,monospace; + border-radius:5px;padding:2px 7px;box-shadow:0 0 3px rgba(0,0,0,.5)} +.chk{flex-direction:row!important;align-items:center;gap:8px;text-transform:none!important; + color:var(--fg)!important;font-size:13px!important} +.rng{flex-direction:row!important;align-items:center;gap:3px;text-transform:none!important; + color:var(--mut)!important;font-size:11px!important} +.rng select{padding:4px 6px;font-size:12px} +#score-legend{display:flex;align-items:center;gap:8px;font-size:11px;color:var(--mut);align-self:flex-end} +#score-legend .bar{width:130px;height:12px;border-radius:6px;border:1px solid var(--line); + background:linear-gradient(90deg,#fff,#99000d)} +#ticktip{position:fixed;pointer-events:none;z-index:100;display:none;white-space:nowrap; + background:rgba(0,0,0,.88);color:#fff;font-size:11px;padding:3px 7px;border-radius:4px} +.panel .cap{font-size:11px;color:var(--mut);font-family:ui-monospace,monospace;text-align:center; + white-space:nowrap;overflow:hidden;text-overflow:ellipsis} +#meta{color:var(--mut);font-size:12px;font-family:ui-monospace,monospace;text-align:center} + +/* Montage tab: OpenSeadragon fills the stage; points overlay + hover tooltip */ +#montage-view{display:none;position:relative;width:100%;height:100%} +#stage.montage-active{align-items:stretch;padding:0;overflow:hidden} +#stage.montage-active>*{display:none} +#stage.montage-active>#montage-view{display:block} +#osd{width:100%;height:100%;background:#000} +#m-overlay{position:absolute;inset:0;pointer-events:none;z-index:5} +#m-live{position:absolute;inset:0;width:100%;height:100%;background:#000;display:none;cursor:grab} +#m-live.drag{cursor:grabbing} +#controls .live-only{display:none} /* renderer-specific controls in the Embedding tab (id beats #controls label) */ +#tab-montage.liverender .tiles-only{display:none} +#tab-montage.liverender .live-only{display:flex} +#montage-view.live #osd,#montage-view.live #m-overlay{display:none} +#montage-view.live #m-live{display:block} +#m-tip{position:fixed;pointer-events:none;z-index:100;display:none;white-space:nowrap; + background:rgba(0,0,0,.9);color:var(--accent);font:600 13px ui-monospace,monospace; + padding:3px 8px;border-radius:4px;border:1px solid var(--line)} +#m-legend{margin-top:10px;max-height:260px;overflow:auto;font-size:11px;font-family:ui-monospace,monospace} +#m-legend .leg-hd{text-transform:uppercase;letter-spacing:.04em;color:var(--accent);margin-bottom:5px} +#m-legend .leg-i{display:flex;align-items:center;gap:6px;padding:1px 0;color:var(--mut); + white-space:nowrap;overflow:hidden;text-overflow:ellipsis} +#m-legend .leg-i.more{font-style:italic;opacity:.7} +#m-legend .leg-grad{height:14px;border-radius:4px;border:1px solid var(--line);margin:2px 0 3px} +#m-legend .leg-cont{display:flex;justify-content:space-between;color:var(--mut);font-family:ui-monospace,monospace} +.combo-hd{padding:5px 9px 2px;font-size:10px;text-transform:uppercase;letter-spacing:.05em;color:var(--accent)} +.ci-sub{color:var(--mut);font-size:11px} +#m-legend .sw{width:11px;height:11px;border-radius:2px;flex:none} + +/* dual-handle range: two overlaid range inputs sharing one track (clim min + max) */ +.dual{position:relative;height:22px} +.dual::before{content:"";position:absolute;left:0;right:0;top:50%;transform:translateY(-50%);height:4px;background:var(--line);border-radius:2px} +.dual input[type=range]{position:absolute;left:0;top:0;width:100%;height:22px;margin:0;background:none; + pointer-events:none;-webkit-appearance:none;appearance:none} +.dual input[type=range]::-webkit-slider-thumb{pointer-events:auto;-webkit-appearance:none;appearance:none; + width:14px;height:14px;border-radius:50%;background:var(--accent);border:2px solid var(--bg);cursor:pointer;box-shadow:0 0 3px #000} +.dual input[type=range]::-moz-range-thumb{pointer-events:auto;width:14px;height:14px;border-radius:50%; + background:var(--accent);border:2px solid var(--bg);cursor:pointer} +.dual input[type=range]::-webkit-slider-runnable-track{background:none} +.dual input[type=range]::-moz-range-track{background:none} + +/* Attention-head view: grid of real phenotype crops with per-head inferno attribution overlay */ +#attn-view{display:none;width:100%;max-width:1100px;flex-direction:column;gap:12px} +#stage.attn-active{justify-content:flex-start} +#stage.attn-active>*{display:none} +#stage.attn-active>#attn-view{display:flex} +#attn-head-lbl{font:650 16px ui-monospace,Menlo,monospace;color:var(--fg);line-height:1.35} +#attn-head-lbl .sub{display:block;font:400 12px ui-monospace,monospace;color:var(--mut);margin-top:3px} +#attn-grid{display:flex;flex-direction:column;gap:20px} +.agroup{display:flex;flex-direction:column;gap:8px} +.agroup-hd{font:650 15px ui-monospace,Menlo,monospace} /* colored per perturbation */ +.arow{display:flex;gap:10px;align-items:stretch} +.arow-lbl{flex:0 0 92px;border-left:4px solid;padding-left:8px;font:11px/1.3 ui-monospace,monospace; + color:var(--mut);display:flex;align-items:center} +.arow-cells{flex:1;display:grid;gap:8px;grid-template-columns:repeat(var(--acols,4),minmax(0,1fr))} +.acell{position:relative} +.acell canvas{width:100%;aspect-ratio:1;background:#000;border:1px solid var(--line);border-radius:6px;display:block} +.acell .cap{font-size:10px;color:var(--mut);font-family:ui-monospace,monospace;text-align:center;margin-top:2px} +#attn-view .empty{color:var(--mut);font-style:italic;padding:24px;text-align:center} + +/* PCs tab: PC list (left pane) + strip explorer (stage) */ +.pc-list{list-style:none;margin:6px 0 0;padding:0;max-height:calc(100vh - 260px);overflow:auto} +.pc-item{display:flex;align-items:center;gap:6px;padding:4px 6px;border-radius:4px;cursor:pointer; + font:12px ui-monospace,monospace;color:var(--mut)} +.pc-item:hover{background:var(--panel)} +.pc-item.active{background:var(--panel);color:var(--fg)} +.pc-item .pc-lbl{width:44px;flex:none} +.pc-item .pc-bar{flex:1;height:8px;background:var(--line);border-radius:4px;overflow:hidden} +.pc-item .pc-bar span{display:block;height:100%;background:var(--accent)} +.pc-item .pc-pct{width:40px;text-align:right;flex:none} +#pc-view{display:none;width:100%;max-width:1400px} +#stage.pc-active{justify-content:flex-start} +#stage.pc-active>*{display:none} +#stage.pc-active>#pc-view{display:block} +#tc-view{display:none;width:100%;max-width:1400px} +#stage.top-active{justify-content:flex-start} +#stage.top-active>*{display:none} +#stage.top-active>#tc-view{display:block} +.tc-row{margin-bottom:18px} +.tc-hd{font:650 15px ui-monospace,monospace;margin-bottom:6px;display:flex;align-items:center;gap:8px} +.tc-hd .tc-n{margin-left:auto;font:400 11px ui-monospace,monospace;color:var(--mut)} +.tc-hd .tc-acc{font:700 13px ui-monospace,monospace;padding:2px 8px;border-radius:6px} +.tc-strip{gap:4px;flex-wrap:wrap;background:var(--panel);border:1px solid var(--line);border-radius:8px;padding:8px} +#stage.top-active>#tc-view.cols-layout{display:flex;flex-direction:row;flex-wrap:nowrap;align-items:flex-start;gap:12px;overflow-x:auto} /* perturbations side-by-side; selector must out-specify the '#stage.top-active>#tc-view{display:block}' show rule */ +#tc-view.cols-layout .tc-row{flex:0 0 auto;margin-bottom:0;width:calc(var(--tcpx,150px) + 18px)} /* fixed-width column (cell + strip padding) so long titles can't stretch it */ +#tc-view.cols-layout .tc-strip{flex-direction:column;flex-wrap:nowrap} +#tc-view.cols-layout .tc-hd{white-space:normal;text-align:center;justify-content:center;word-break:break-word;flex-wrap:wrap} +.tc-cell{position:relative;flex:none} +#tc-view .pc-cell{width:var(--tcpx,150px);height:var(--tcpx,150px)} /* Image-scale slider (--tcpx) drives top-cell size; scoped so the PC tab stays 150px */ +.tc-rank{position:absolute;top:3px;left:3px;font:700 11px ui-monospace,monospace;color:#fff;background:rgba(0,0,0,.6);border-radius:4px;padding:0 5px} +.tc-ov{position:absolute;top:0;left:0;width:var(--tcpx,150px);height:var(--tcpx,150px);border-radius:3px;pointer-events:none;display:none} +#tc-view.masked .tc-ov{display:block} +.pc-head{display:flex;align-items:center;gap:14px;margin-bottom:12px} +.pc-head h2{margin:0;font:650 20px ui-monospace,monospace} +.pc-head .row{margin-left:auto} +.pc-sort-row{display:flex;align-items:center;gap:6px;font-size:11px;color:var(--mut);text-transform:none;letter-spacing:0} +.pc-sort-row select{flex:1;padding:4px 6px;font-size:12px} +.pc-strip{background:var(--panel);border:1px solid var(--line);border-radius:8px;padding:10px;overflow-x:auto} +.pc-strip-row{display:flex;gap:3px;margin-bottom:3px} +.pc-cell{width:150px;height:150px;flex:none;border-radius:3px;background:#000;cursor:pointer;object-fit:cover} +.pc-cell:hover{outline:2px solid var(--accent)} +.pc-cell.ph{background:#111;cursor:default} +.pc-axis{display:flex;justify-content:space-between;font-size:11px;color:var(--mut);margin-top:4px;padding:0 2px} +.pc-genes{display:grid;grid-template-columns:1fr 1fr;gap:20px;margin-top:16px} +.pc-genes .sec-lbl{font-size:11px;text-transform:uppercase;letter-spacing:.04em;color:var(--accent);margin-bottom:6px} +.chips{display:flex;flex-wrap:wrap;gap:5px} +.chip{background:var(--panel);border:1px solid var(--line);border-radius:12px;padding:2px 9px;font-size:12px; + font-family:ui-monospace,monospace;cursor:pointer} +.chip:hover{border-color:var(--accent)} +.chip small{color:var(--mut)} +.chip.pos{border-color:#2e6}.chip.neg{border-color:#e55} +.pc-overlay{position:fixed;inset:0;background:rgba(0,0,0,.7);z-index:100;display:flex;align-items:center;justify-content:center} +.pc-ov-card{background:var(--panel);border:1px solid var(--line);border-radius:12px;padding:18px;width:min(90%,640px);max-height:80vh;overflow:auto} +.pc-ov-hd{display:flex;justify-content:space-between;align-items:center;font:650 16px ui-monospace,monospace;margin-bottom:8px} +.pc-ov-hd button{background:none;border:none;color:var(--mut);font-size:18px;cursor:pointer} +.ov-row{display:flex;align-items:center;gap:8px;padding:3px 4px;border-radius:4px;cursor:pointer;font:12px ui-monospace,monospace} +.ov-row:hover{background:var(--line)} +.ov-pc{width:46px;text-align:right;color:var(--mut)} +.ov-track{flex:1;height:14px;background:var(--bg);border-radius:3px;position:relative} +.ov-track span{position:absolute;height:100%;border-radius:3px} +.ov-val{width:52px;text-align:right} +/* speedrichr-style enrichment bars (per library, ordered by adj p-value) */ +.enr-lib{margin-top:12px} +.enr-hd{font-size:11px;text-transform:uppercase;letter-spacing:.04em;color:var(--mut);margin-bottom:4px} +.enr-row{display:flex;align-items:center;gap:8px;font-size:11px;margin-bottom:2px;border-radius:4px;padding:1px 3px;transition:background .12s} +.enr-track{flex:1;position:relative;height:16px;background:var(--bg);border-radius:4px;overflow:hidden;min-width:0;transition:height .12s} +.enr-bar{position:absolute;left:0;top:0;bottom:0;background-color:var(--accent);opacity:.34;transition:opacity .12s} +.enr-term{position:absolute;left:7px;right:6px;top:0;line-height:16px;white-space:nowrap;overflow:hidden;text-overflow:ellipsis;color:var(--fg);transition:font-size .12s,line-height .12s} +.enr-n{flex:none;width:118px;text-align:right;color:var(--mut);font-family:ui-monospace,monospace;transition:font-size .12s,color .12s} +.enr-bar.neg{background-color:#e55} +.enr-row:hover{background:rgba(255,255,255,.06)} +.enr-row:hover .enr-track{height:24px} +.enr-row:hover .enr-bar{opacity:.6} +.enr-row:hover .enr-term{font-size:13px;line-height:24px;font-weight:600} +.enr-row:hover .enr-n{font-size:12.5px;color:var(--fg)} +.enr-warn{color:#e0a94a;text-transform:none;letter-spacing:0;font-weight:400} +/* reusable segmented control (Ontology/Features + all small dropdowns) */ +.pc-mode,.seg-group{display:flex;gap:0;margin:6px 0 2px} +.seg{flex:1;min-width:0;background:var(--bg);border:1px solid var(--line);color:var(--mut);padding:4px 8px;font-size:12px;cursor:pointer;white-space:nowrap;overflow:hidden;text-overflow:ellipsis;text-align:center;text-transform:none;letter-spacing:0} +.seg+.seg{border-left:none} +.seg:first-child{border-radius:5px 0 0 5px} +.seg:last-child{border-radius:0 5px 5px 0} +.seg:hover:not(.active){color:var(--fg);border-color:var(--accent)} +.seg.active{background:var(--accent);color:#fff;border-color:var(--accent)} +/* inline segmented (e.g. the "list order" mini control sitting after a label) */ +.seg-group.mini{display:inline-flex;margin:0 0 0 6px;vertical-align:middle} +.seg-group.mini .seg{flex:none;padding:2px 12px;font-size:11px} +/* checkboxes → [off | feature] segmented switches (see toggleize); off-active is neutral, on-active is accent */ +.seg-group.tog{margin:8px 0 2px} +.seg-group.tog .seg:first-child.active{background:var(--line);color:var(--fg);border-color:var(--line)} +.pc-comp{grid-column:1/-1;margin-top:18px} +.comp-block{margin-top:10px} +.comp-bar{display:flex;height:30px;border-radius:9px;overflow:hidden; + background:rgba(255,255,255,.03);border:1px solid rgba(255,255,255,.1);box-shadow:inset 0 1px 2px rgba(0,0,0,.35)} +/* frosted glass: translucent color fill (see-through) + backdrop blur + hairline divider — no gloss */ +.comp-seg{position:relative;height:100%;min-width:0;display:flex;align-items:center;justify-content:center; + backdrop-filter:blur(7px) saturate(1.1);-webkit-backdrop-filter:blur(7px) saturate(1.1); + box-shadow:inset -1px 0 0 rgba(255,255,255,.14);transition:filter .12s,background-color .12s} +.comp-seg:last-child{box-shadow:none} +.comp-seg:hover{filter:brightness(1.22)} +.comp-lbl{padding:0 8px;font:650 10.5px/1 ui-sans-serif,-apple-system,BlinkMacSystemFont,"Segoe UI",sans-serif; + text-transform:uppercase;letter-spacing:.08em;color:rgba(255,255,255,.97);white-space:nowrap;overflow:hidden; + text-overflow:ellipsis;text-shadow:0 1px 3px rgba(0,0,0,.9);transition:font-size .12s,letter-spacing .12s} +.comp-seg:hover .comp-lbl{font-size:12.5px;letter-spacing:.1em} +/* glassy pop-out inset (draggable, stackable; viewer still visible behind) */ +.popout{position:fixed;z-index:200;width:450px;container-type:inline-size;background:rgba(22,27,34,.55);backdrop-filter:blur(10px); + -webkit-backdrop-filter:blur(10px);border:1px solid rgba(255,255,255,.18);border-radius:10px; + box-shadow:0 14px 44px rgba(0,0,0,.55);overflow:hidden;resize:horizontal;min-width:160px;max-width:90vw} +.po-bar{display:flex;justify-content:space-between;align-items:center;gap:8px;padding:6px 10px;cursor:move; + font:600 clamp(11px,4cqw,18px) ui-monospace,monospace;background:rgba(255,255,255,.07);border-bottom:1px solid rgba(255,255,255,.1)} +.po-bar span{overflow:hidden;text-overflow:ellipsis;white-space:nowrap} +.po-bar button{background:none;border:none;color:var(--fg);font-size:16px;cursor:pointer;line-height:1} +.popout img{width:100%;display:block;background:#000} +.po-img{position:relative;width:100%} +.popout img.po-ov{position:absolute;top:0;left:0;width:100%;height:100%;background:transparent!important;pointer-events:none} +.po-body{padding:8px 10px;font-size:clamp(11px,4.2cqw,20px);line-height:1.5;color:var(--fg);overflow-wrap:break-word} +.po-body .po-sub{color:var(--mut);font-family:ui-monospace,monospace} +.po-body a{color:var(--accent);text-decoration:none} + +/* ---- "How it works" methods tab ---- */ +#methods-view{display:none;width:100%;max-width:860px} +#stage.methods-active{justify-content:flex-start} +#stage.methods-active>*{display:none} +#stage.methods-active>#methods-view{display:block} +.mth-rail{display:flex;flex-direction:column;gap:4px;margin-top:10px} +.mth-railitem{display:flex;align-items:center;gap:9px;text-align:left;background:none;border:none;color:var(--mut); + padding:7px 9px;border-radius:7px;cursor:pointer;font-size:13px;border-left:3px solid transparent;width:100%} +.mth-railitem:hover{background:rgba(255,255,255,.05);color:var(--fg)} +.mth-railitem.on{background:rgba(88,166,255,.12);color:var(--fg);border-left-color:var(--accent)} +.mth-num{display:inline-flex;align-items:center;justify-content:center;width:20px;height:20px;border-radius:50%; + background:rgba(255,255,255,.08);font-size:11px;flex:0 0 auto} +.mth-railitem.on .mth-num{background:var(--accent);color:#08111f} +.mth-railitem-x{font-style:italic} +.mth-railitem-x .mth-num{background:rgba(255,255,255,.04);color:var(--mut);font-weight:600} +.mth-card{background:rgba(22,27,34,.5);border:1px solid rgba(255,255,255,.1);border-radius:14px;padding:22px 26px 18px} +.mth-kicker{font:600 11px ui-monospace,monospace;letter-spacing:.14em;color:var(--accent);text-transform:uppercase} +.mth-title{margin:6px 0 4px;font-size:24px;font-weight:700;color:var(--fg)} +.mth-stage{display:flex;justify-content:center;padding:8px 0 14px} +.mth-stage svg.mth{width:100%;max-width:520px;height:auto} +.mth-body{font-size:15px;line-height:1.6;color:var(--fg);max-width:660px} +.mth-body i{color:var(--mut)} +.mth-why{margin-top:12px;font-size:13.5px;line-height:1.55;color:var(--mut);border-left:3px solid var(--accent);padding-left:11px} +.mth-why b{color:var(--accent)} +.mth-navbar{display:flex;align-items:center;justify-content:space-between;margin-top:18px} +.mth-navbar button{background:rgba(255,255,255,.07);border:1px solid rgba(255,255,255,.14);color:var(--fg); + padding:6px 14px;border-radius:8px;cursor:pointer;font-size:13px} +.mth-navbar button:disabled{opacity:.35;cursor:default} +.mth-dots{display:flex;gap:7px} +.mth-dot{width:9px;height:9px;border-radius:50%;background:rgba(255,255,255,.2);cursor:pointer} +.mth-dot.on{background:var(--accent)} +/* SVG animations (group transforms + opacity; per-class transform-origin) */ +.mth-pulse,.mth-breathe,.mth-rise,.mth-collapse,.mth-morph,.mth-slide,.mth-glow,.mth-sweep,.mth-hi,.mth-jit,.mth-soft,.mth-branch,.mth-cyc{transform-box:fill-box} +@keyframes mthPulse{0%,100%{opacity:.35}50%{opacity:1}} +.mth-pulse{transform-origin:center;animation:mthPulse 2.6s ease-in-out infinite} +@keyframes mthBreathe{0%,100%{transform:scale(1)}50%{transform:scale(1.07)}} +.mth-breathe{transform-origin:center;animation:mthBreathe 3s ease-in-out infinite} +@keyframes mthRise{0%{transform:scaleY(.05)}55%,100%{transform:scaleY(1)}} +.mth-rise{transform-origin:bottom;animation:mthRise 2.6s ease-out infinite} +@keyframes mthCollapse{0%{transform:scale(1.5);opacity:.85}62%{transform:scale(.18);opacity:0}100%{opacity:0}} +.mth-collapse{transform-origin:center;animation:mthCollapse 3.6s ease-in infinite} +@keyframes mthEmerge{0%,45%{opacity:0}72%,100%{opacity:1}} +.mth-emerge{animation:mthEmerge 3.6s ease-out infinite} +@keyframes mthMorph{0%,100%{transform:scale(1,1)}50%{transform:scale(1.28,.74)}} +.mth-morph{transform-origin:center;animation:mthMorph 3.2s ease-in-out infinite} +/* Traversal: z-bars morph A(NTC)→B(knockout)→A in lockstep with the α dot + cell (all 3.2s) */ +@keyframes mthZmorph{0%,100%{transform:scaleY(1)}50%{transform:scaleY(var(--sy,1))}} +.mth-zmorph{transform-box:fill-box;transform-origin:bottom;animation:mthZmorph 3.2s ease-in-out infinite} +@keyframes mthSlide{0%,100%{transform:translateX(0)}50%{transform:translateX(280px)}} +.mth-slide{animation:mthSlide 3.2s ease-in-out infinite} +@keyframes mthGlow{0%,100%{transform:scale(.7);opacity:.35}50%{transform:scale(1.15);opacity:.9}} +.mth-glow{transform-origin:center;animation:mthGlow 2.8s ease-in-out infinite} +@keyframes mthCycA{0%,100%{opacity:1}50%{opacity:.12}} +.mth-cycA{animation:mthCycA 3s ease-in-out infinite} +@keyframes mthCycB{0%,100%{opacity:.12}50%{opacity:1}} +.mth-cycB{animation:mthCycB 3s ease-in-out infinite} +.mth-refs{margin-top:12px;font-size:12px;color:var(--mut)} +.mth-refs a{color:var(--accent);text-decoration:none} +.mth-refs a:hover{text-decoration:underline} +.mth-defs{margin-top:12px;border-top:1px solid rgba(255,255,255,.08);padding-top:8px} +.mth-defs>summary{cursor:pointer;font-size:12px;color:var(--accent);font-weight:600;list-style:none} +.mth-defs>summary::-webkit-details-marker{display:none} +.mth-defs>summary::before{content:"▸ ";color:var(--mut)} +.mth-defs[open]>summary::before{content:"▾ "} +.mth-defs dl{margin:9px 0 0} +.mth-defs dl>div{margin-bottom:7px} +.mth-defs dt{font-size:12.5px;font-weight:600;color:var(--fg)} +.mth-defs dd{margin:1px 0 0;font-size:12.5px;line-height:1.5;color:var(--mut)} +@keyframes mthSweep{0%{transform:translateX(0);opacity:0}12%{opacity:.9}88%{transform:translateX(var(--sw,300px));opacity:.9}100%{transform:translateX(var(--sw,300px));opacity:0}} +.mth-sweep{animation:mthSweep 4s ease-in-out infinite} +/* methods tab — larger text + wider layout + Learn-more dropdown (feedback) */ +#methods-view{max-width:1180px} +#stage.methods-active{padding:16px 30px} +.mth-title{font-size:30px} +.mth-kicker{font-size:14px} +.mth-body{font-size:19px;max-width:940px} +.mth-why{font-size:17px} +.mth-refs{font-size:15px} +.mth-stage svg.mth{max-width:680px} +.mth-defs>summary{font-size:15px} +.mth-defs dt{font-size:15.5px} +.mth-defs dd{font-size:15.5px} +.mth-navbar button{font-size:16px} +.mth-railitem{font-size:14px} +.mth-more{margin-top:10px;border-top:1px solid rgba(255,255,255,.06);padding-top:8px} +.mth-more>summary{cursor:pointer;color:var(--accent);font-weight:600;font-size:15px;list-style:none} +.mth-more>summary::-webkit-details-marker{display:none} +.mth-more>summary::before{content:"▸ ";color:var(--mut)} +.mth-more[open]>summary::before{content:"▾ "} +.mth-more p{margin:8px 0 0;font-size:16px;line-height:1.65;color:var(--fg)} +.mth-sec{margin-bottom:4px} +.mth-seccap{font-size:16px;color:var(--fg);font-weight:700;margin:10px 0 4px;text-align:center} /* match deck typography (not the bright accent-blue) */ +.mth-sec .mth-stage{padding:4px 0 8px} +.mth-sec .mth-stage svg.mth{max-width:620px} +.mth-sectext{font-size:19px;line-height:1.6;color:var(--fg);text-align:left;max-width:940px;margin:4px 0 10px} /* same as .mth-body */ +@keyframes mthHi{0%,100%{transform:scale(1);opacity:.85}50%{transform:scale(1.16);opacity:1}} +.mth-hi{transform-origin:center;animation:mthHi 3.2s ease-in-out infinite} +@keyframes mthJit{0%,100%{transform:scaleY(.9)}50%{transform:scaleY(1.05)}} +.mth-jit{transform-origin:bottom;animation:mthJit 2.6s ease-in-out infinite} + +@keyframes mthSoft{0%,100%{opacity:.5}50%{opacity:.78}} +.mth-soft{animation:mthSoft 3.6s ease-in-out infinite backwards} +@keyframes mthBranch{0%,100%{opacity:1}50%{opacity:.4}} +.mth-branch{animation:mthBranch 3.7s ease-in-out infinite both} +@keyframes mthCyc{0%,10%{transform:scale(1);opacity:.55}16%,24%{transform:scale(1.18);opacity:1}32%,100%{transform:scale(1);opacity:.55}} +.mth-cyc{transform-origin:center;animation:mthCyc 5s ease-in-out infinite} +@keyframes mthCycIn{0%,10%{opacity:0}17%,25%{opacity:.95}33%,100%{opacity:0}} +.mth-cycin{animation:mthCycIn 5s ease-in-out infinite both} + +@keyframes mthGrow{0%,100%{transform:scaleY(.25)}50%{transform:scaleY(1)}} +.mth-grow{transform-box:fill-box;transform-origin:bottom;animation:mthGrow 3s ease-in-out infinite} +@keyframes mthExpress{0%,100%{opacity:.15}50%{opacity:1}} +.mth-express{animation:mthExpress 3s ease-in-out infinite} +@keyframes mthDrop{0%,100%{transform:scaleY(1)}50%{transform:scaleY(.06)}} +.mth-drop{transform-box:fill-box;transform-origin:bottom;animation:mthDrop 3s ease-in-out infinite} +@keyframes mthRemove{0%,100%{opacity:1}50%{opacity:.12}} +.mth-remove{animation:mthRemove 3s ease-in-out infinite} +@keyframes mthShrink{0%,100%{transform:scale(1)}50%{transform:scale(.75)}} +.mth-shrink{transform-box:fill-box;transform-origin:center;animation:mthShrink 3s ease-in-out infinite} +/* Embedding: 3 discrete states switch in lockstep (synchronized) with snappy transitions (stochastic, not a wave) */ +@keyframes mthSt0{0%,28%{opacity:1}34%,94%{opacity:0}100%{opacity:1}} +@keyframes mthSt1{0%,28%{opacity:0}34%,61%{opacity:1}67%,100%{opacity:0}} +@keyframes mthSt2{0%,61%{opacity:0}67%,94%{opacity:1}100%{opacity:0}} +.mth-st0,.mth-st1,.mth-st2{animation-duration:4.2s;animation-timing-function:ease;animation-iteration-count:infinite} +.mth-st0{animation-name:mthSt0}.mth-st1{animation-name:mthSt1}.mth-st2{animation-name:mthSt2} +@keyframes mthFlowEdge{0%{opacity:.2}10%{opacity:.95}26%{opacity:.2}100%{opacity:.2}} +.mth-flow-edge{animation:mthFlowEdge 3s ease-in-out infinite} +@keyframes mthFlowNode{0%,100%{transform:scale(1)}10%{transform:scale(1.4)}28%{transform:scale(1)}} +.mth-flow-node{transform-box:fill-box;transform-origin:center;animation:mthFlowNode 3s ease-in-out infinite} +/* predicted class node: pops large + bright when the signal reaches the last layer (delay 1.25s, same 3s cycle as the flow) */ +@keyframes mthPredict{0%,100%{transform:scale(1);opacity:.4}15%{transform:scale(1.6);opacity:1}55%{transform:scale(1.2);opacity:.9}} +.mth-predict{transform-box:fill-box;transform-origin:center;animation:mthPredict 3s ease-in-out infinite} +.brand{display:flex;align-items:center;gap:9px} +.opsin-eyes{height:28px;width:auto;display:block} +.biohub-logo{display:inline-flex;align-items:center;padding-left:10px;border-left:1px solid var(--line);margin-left:2px;opacity:.9;transition:opacity .15s} +.biohub-logo img{height:26px;width:auto;display:block} +.biohub-logo:hover{opacity:1} +/* full-screen loading veil — hides the initial traversal flash until first render */ +#loading{position:fixed;inset:0;z-index:200;display:flex;align-items:center;justify-content:center;background:var(--bg);transition:opacity .4s ease} +#loading.gone{opacity:0;pointer-events:none} +#loading .load-box{display:flex;flex-direction:column;align-items:center;gap:12px} +#loading .load-box img{width:56px;height:56px;animation:loadpulse 1.3s ease-in-out infinite} +#loading .load-t{font:700 20px ui-sans-serif,system-ui,sans-serif;color:var(--fg);letter-spacing:.06em} +#loading .load-s{font:400 12px ui-sans-serif,system-ui,sans-serif;color:var(--mut);letter-spacing:.08em} +@keyframes loadpulse{0%,100%{opacity:.55;transform:scale(1)}50%{opacity:1;transform:scale(1.09)}} + +/* universal image clim — display-time levels stretch applied to every panel's images/canvases */ +#stage img, #stage canvas, .popout img { filter: url(#img-levels); } /* .popout is fixed/outside #stage → include it */ +#imgclim-field { min-width: 150px; } +#imgclim-field .dual { align-self: center; } +#imgclim-field #i-climreset { padding: 2px 7px; } + +/* affinage mechanistic narrative — collapsible prose in the Info panel */ +#info-body .narr{color:var(--fg);line-height:1.55} +#info-body .narr-more{display:block;margin-top:6px;background:none;border:none;padding:0; + color:var(--accent);font:inherit;font-size:12px;cursor:pointer} +#info-body .narr-more:hover{text-decoration:underline} + +/* per-member affinage narrative on the complex page — collapsed by default */ +#info-body .cx-narr{margin:2px 0 2px 12px} +#info-body .cx-narr>summary{cursor:pointer;color:var(--accent);font-size:11px;list-style:none} +#info-body .cx-narr>summary::-webkit-details-marker{display:none} +#info-body .cx-narr>summary::before{content:"▸ "} +#info-body .cx-narr[open]>summary::before{content:"▾ "} +#info-body .cx-narr .narr{font-size:12px;color:var(--fg);line-height:1.5;margin:4px 0 6px 12px} + +/* internal-vs-public feature gate + staging-only preview toggle */ +.feat-hidden{display:none !important} +.envtoggle{margin-right:12px;border-color:#d29922;color:#e3b341} +.envtoggle:hover{border-color:#e3b341} +.envtoggle.pub-on{background:rgba(210,153,34,.18);border-color:#e3b341;color:#f2cc60} + +/* About panel — branded hero + tighter, styled sections */ +.about-hero{display:flex;flex-direction:column;align-items:center;text-align:center;gap:3px;padding:4px 0 14px} +.about-hero .about-eyes{width:54px;height:54px;margin-bottom:2px} +.about-name{font:700 22px ui-sans-serif,system-ui,sans-serif;letter-spacing:.07em;color:var(--fg)} +.about-tag{font-size:10.5px;color:var(--mut);letter-spacing:.03em;max-width:210px;line-height:1.35} +.about-lead{font-size:12.5px;line-height:1.55;color:var(--mut);padding-bottom:13px;border-bottom:1px solid var(--line);margin-bottom:14px} +.about-lead b{color:var(--fg)} +.about-sec{margin-bottom:15px} +.about-cap{font-size:11px;text-transform:uppercase;letter-spacing:.05em;color:var(--accent);margin-bottom:7px} +.about-tabs{list-style:none;margin:0;padding:0;display:flex;flex-direction:column;gap:7px} +.about-tabs li{font-size:12px;line-height:1.45;color:var(--mut);padding-left:10px;border-left:2px solid var(--line)} +.about-tabs li b{color:var(--fg)} +.about-foot{display:flex;align-items:center;gap:10px;padding-top:13px;border-top:1px solid var(--line)} +.about-foot img{height:20px;opacity:.92} +.about-foot span{font-size:11px;color:var(--mut);line-height:1.35} From eccd236aa2da85f3d22285b6ac3337aacf145b61 Mon Sep 17 00:00:00 2001 From: Gav Sturm Date: Tue, 11 Aug 2026 09:25:52 -0700 Subject: [PATCH 02/13] rename models/attention -> models/interpretability Whole-dir rename; 52 intra-tree module-path refs updated. 0 external ops_model importers (blast radius nil). Core diffex subpackages import-verified. --- .../{attention => interpretability}/RUNBOOK.md | 0 .../atlas/attention_accuracy_umap_animation.py | 0 .../atlas/attention_atlas.py | 0 .../atlas/attention_atlas_shap.py | 0 .../atlas/low_attention_phase_atlas.py | 0 .../atlas/make_scale_bar.py | 0 .../atlas/marker_selection_distribution.py | 0 .../atlas/plot_eval_accuracy_curves.py | 0 .../diffex/CROPSEQ_TO_MORPHOLOGY.md | 0 .../{attention => interpretability}/diffex/PLAN.md | 0 .../diffex/README.md | 0 .../diffex/classifier/README.md | 0 .../diffex/classifier/__init__.py | 0 .../diffex/classifier/aggregate.py | 2 +- .../diffex/classifier/celldino_features.py | 0 .../diffex/classifier/config.py | 0 .../diffex/classifier/data.py | 0 .../diffex/classifier/models.py | 0 .../diffex/classifier/run.py | 4 ++-- .../diffex/classifier/submit.py | 6 +++--- .../diffex/classifier/train.py | 0 .../diffex/diffae/__init__.py | 0 .../diffex/diffae/config.py | 0 .../diffex/diffae/data.py | 0 .../diffex/diffae/diagnose_conditioning.py | 2 +- .../diffex/diffae/model.py | 0 .../diffex/diffae/plot_metrics.py | 0 .../diffex/diffae/recon.py | 0 .../diffex/diffae/run.py | 2 +- .../diffex/diffae/submit.py | 2 +- .../diffex/diffae/train.py | 0 .../diffex/diffae/virtstain_eval.py | 2 +- .../diffex/diffae/virtstain_multi.py | 2 +- .../diffex/directions/__init__.py | 0 .../diffex/directions/batch.py | 2 +- .../diffex/directions/config.py | 0 .../diffex/directions/data.py | 0 .../diffex/directions/flow.py | 0 .../diffex/directions/grid.py | 0 .../diffex/directions/losses.py | 0 .../diffex/directions/make_gifs.py | 0 .../diffex/directions/model.py | 0 .../diffex/directions/proto_ddim_anchors.py | 2 +- .../diffex/directions/rank.py | 0 .../diffex/directions/run.py | 4 ++-- .../diffex/directions/submit.py | 2 +- .../diffex/directions/train_directions.py | 0 .../diffex/directions/traverse.py | 0 .../diffex/figures/METHODS_final.txt | 0 .../diffex/figures/METHODS_traversal_montage.md | 0 .../diffex/figures/METHODS_traversal_montage.txt | 0 .../diffex/figures/_setacc_common.py | 10 +++++----- .../diffex/figures/_setacc_phase.py | 0 .../diffex/figures/auto_pick_and_plot.py | 4 ++-- .../diffex/figures/cis_golgi_alternatives.py | 0 .../diffex/figures/debug_setacc_top100.py | 2 +- .../diffex/figures/ebi_peripheral_droplets.py | 2 +- .../diffex/figures/figure4_morpho_traversal.py | 0 .../diffex/figures/figure4_morpho_violin.py | 4 ++-- .../diffex/figures/figure4_setacc_panel.py | 0 .../diffex/figures/figure4_setacc_panel_fluorB.py | 0 .../diffex/figures/figure4_setacc_panel_newpheno.py | 0 .../diffex/figures/figure4_setacc_panel_phase.py | 0 .../diffex/figures/figure_ebi_morpho_violin.py | 0 .../diffex/figures/figure_multirank_ebi_grid.py | 0 .../diffex/figures/fluor_panel_montages.py | 0 .../diffex/figures/fluor_shap_montages.py | 2 +- .../figures/gen_validation/bag_sweep_plots.py | 2 +- .../figures/gen_validation/bag_sweep_score.py | 6 +++--- .../figures/gen_validation/centroid_bagsweep.py | 4 ++-- .../figures/gen_validation/centroid_halves.py | 4 ++-- .../gen_validation/centroid_pooled_bagsweep.py | 4 ++-- .../figures/gen_validation/control_halves_zscore.py | 4 ++-- .../diffex/figures/gen_validation/embcheck.py | 6 +++--- .../figures/gen_validation/embedding_diagnostics.py | 0 .../gen_validation/figure4_bagsize_reachfrac.py | 0 .../gen_validation/figure4_v5_accuracy_summary.py | 4 ++-- .../figures/gen_validation/gen_alpha_embedding.py | 0 .../figures/gen_validation/gen_embed_refit.py | 4 ++-- .../figures/gen_validation/gen_phate_passthrough.py | 0 .../figures/gen_validation/gen_real_centroid.py | 12 ++++++------ .../figures/gen_validation/gen_real_distinct.py | 0 .../figures/gen_validation/ntc_inverse_gap.py | 10 +++++----- .../figures/gen_validation/patch_cache_real.py | 6 +++--- .../figures/gen_validation/publish_multibag_page.py | 0 .../diffex/figures/gen_validation/rank_summary.py | 0 .../figures/gen_validation/st_halves_score.py | 6 +++--- .../figures/gen_validation/std_anchor_test.py | 0 .../figures/gen_validation/stepabl_compare.py | 0 .../figures/gen_validation/valid200_alphastep.py | 0 .../figures/gen_validation/valid200_cache_build.py | 0 .../figures/gen_validation/valid200_capcheck.py | 0 .../figures/gen_validation/valid200_map_compare.py | 2 +- .../figures/gen_validation/valid200_metrics.py | 0 .../diffex/figures/nc_ratio.py | 0 .../diffex/figures/ntc_anchor_compare.py | 0 .../diffex/figures/phase_montages.py | 0 .../diffex/figures/phase_multibag_montages.py | 0 .../diffex/figures/phase_sample_montages.py | 0 .../diffex/figures/rab_candidate_montages.py | 2 +- .../diffex/figures/rebuild_traversals_n100.py | 12 ++++++------ .../diffex/figures/traversal_montage_schematic.py | 0 .../diffex/figures/virtual_staining_schematic.py | 0 .../diffex/kyle_pcs/build_static_explorer.py | 0 .../diffex/kyle_pcs/compute_pc_strips.py | 0 .../diffex/viewer/__init__.py | 0 .../diffex/viewer/_altanchor_build.py | 0 .../diffex/viewer/_anchortest.py | 0 .../diffex/viewer/_build_stepablation.py | 0 .../diffex/viewer/_build_v5_inverted.py | 2 +- .../diffex/viewer/_build_v5_montages.py | 2 +- .../diffex/viewer/_build_valid200.py | 6 +++--- .../diffex/viewer/_consolidate_cells.py | 0 .../diffex/viewer/_fluor_complex_build.py | 0 .../diffex/viewer/_fluor_topcells.py | 0 .../diffex/viewer/_fluor_v5_build.py | 0 .../diffex/viewer/_migrate_v4_to_v5.py | 0 .../diffex/viewer/_phase_vs.py | 0 .../diffex/viewer/_rebuild_v5.py | 0 .../diffex/viewer/_rescore_rank.py | 2 +- .../diffex/viewer/_score_v4.py | 0 .../diffex/viewer/_v4acc_test.py | 0 .../diffex/viewer/_verify_pt_space.py | 0 .../diffex/viewer/_verify_score_bridge.py | 0 .../diffex/viewer/altanchor_pairs.json | 0 .../diffex/viewer/anchor_cells.py | 0 .../diffex/viewer/build_attention_heads.py | 8 ++++---- .../diffex/viewer/build_complex_ebi_map.py | 0 .../diffex/viewer/build_fluor_shap_rankings.py | 4 ++-- .../diffex/viewer/build_montage_features.py | 2 +- .../diffex/viewer/build_pc_crops_masked.py | 4 ++-- .../diffex/viewer/build_pc_features.py | 2 +- .../diffex/viewer/build_pc_walks.py | 4 ++-- .../diffex/viewer/build_pcs.py | 4 ++-- .../diffex/viewer/build_pcs_marker.py | 2 +- .../diffex/viewer/build_phase_shap_rankings.py | 4 ++-- .../diffex/viewer/build_phate_figure.py | 2 +- .../diffex/viewer/build_setacc_bins.py | 0 .../diffex/viewer/build_setacc_bymarker.py | 0 .../diffex/viewer/build_top_cells.py | 4 ++-- .../diffex/viewer/build_umap_montage.py | 0 .../diffex/viewer/catalog.py | 0 .../diffex/viewer/deploy/README.md | 0 .../diffex/viewer/marker_leaves.py | 0 .../diffex/viewer/mimic_alex_embed.py | 0 .../diffex/viewer/morpho_pipeline.py | 0 .../diffex/viewer/morphometrics.py | 0 .../diffex/viewer/nway_clf.py | 0 .../diffex/viewer/phenotype_cells.py | 0 .../diffex/viewer/precompute.py | 0 .../diffex/viewer/render_montage_scales.py | 2 +- .../diffex/viewer/score_generated.py | 0 .../diffex/viewer/set_classifier.py | 0 .../diffex/viewer/submit.py | 8 ++++---- .../diffex/viewer/webapp/app.js | 0 .../diffex/viewer/webapp/biohub-mark.png | Bin .../diffex/viewer/webapp/biohub-wordmark.png | Bin .../diffex/viewer/webapp/build_gene_narratives.py | 0 .../diffex/viewer/webapp/gif.js | 0 .../diffex/viewer/webapp/gif.worker.js | 0 .../diffex/viewer/webapp/index.html | 0 .../diffex/viewer/webapp/methods.js | 0 .../diffex/viewer/webapp/morpho_demo.html | 0 .../diffex/viewer/webapp/openseadragon.min.js | 0 .../diffex/viewer/webapp/opsin-eyes.svg | 0 .../diffex/viewer/webapp/style.css | 0 .../embedding/generate_ko_violin_plots.py | 0 .../embedding/regen_umap_gav.py | 0 .../embedding/regen_umap_html.py | 0 .../embedding/run_all_atlases.py | 0 .../embedding/top_attention_embed_and_score.py | 0 .../shap/analyze_chad_variants.py | 0 .../shap/generate_shap_captions_combined.py | 0 .../shap/ko_shap_features.py | 0 .../shap/merge_shap_shards.py | 0 .../shap/ntc_attention_compare.py | 0 .../shap/ntc_pick_cells.py | 0 .../shap/ntc_shap_features.py | 0 .../shap/run_all_shap.py | 0 .../shap/run_shap_pipeline.py | 0 .../shap/shap_approach_compare.py | 0 .../titration/decay/map_attention_decay.py | 0 .../titration/decay/phate_peak_groups.py | 0 .../titration/decay/plot_3way_summary_bars.py | 0 .../decay/plot_all_cells_correction_bars.py | 0 .../expansion/count_genes_above_threshold.py | 4 ++-- .../expansion/map_attention_expansion_v4.py | 0 .../expansion/plot_sgrna_coverage_sweep.py | 0 .../titration/expansion/run_percentile_sweep.py | 0 .../weighted_aggregation/_v4_attn_worker.py | 0 .../weighted_aggregation/analyze_v3_acc_bins.py | 0 .../weighted_aggregation/plot_v4_attn_comparison.py | 0 .../run_v3_pipeline_on_v4_attn_weighted.py | 0 .../run_v3_pipeline_on_v4_features.py | 0 194 files changed, 105 insertions(+), 105 deletions(-) rename src/ops_model/models/{attention => interpretability}/RUNBOOK.md (100%) rename src/ops_model/models/{attention => interpretability}/atlas/attention_accuracy_umap_animation.py (100%) rename src/ops_model/models/{attention => interpretability}/atlas/attention_atlas.py (100%) rename src/ops_model/models/{attention => interpretability}/atlas/attention_atlas_shap.py (100%) rename src/ops_model/models/{attention => interpretability}/atlas/low_attention_phase_atlas.py (100%) rename src/ops_model/models/{attention => interpretability}/atlas/make_scale_bar.py (100%) rename src/ops_model/models/{attention => interpretability}/atlas/marker_selection_distribution.py (100%) rename src/ops_model/models/{attention => interpretability}/atlas/plot_eval_accuracy_curves.py (100%) rename src/ops_model/models/{attention => interpretability}/diffex/CROPSEQ_TO_MORPHOLOGY.md (100%) rename src/ops_model/models/{attention => interpretability}/diffex/PLAN.md (100%) rename src/ops_model/models/{attention => interpretability}/diffex/README.md (100%) rename src/ops_model/models/{attention => interpretability}/diffex/classifier/README.md (100%) rename src/ops_model/models/{attention => interpretability}/diffex/classifier/__init__.py (100%) rename src/ops_model/models/{attention => interpretability}/diffex/classifier/aggregate.py (97%) rename src/ops_model/models/{attention => interpretability}/diffex/classifier/celldino_features.py (100%) rename src/ops_model/models/{attention => interpretability}/diffex/classifier/config.py (100%) rename src/ops_model/models/{attention => interpretability}/diffex/classifier/data.py (100%) rename src/ops_model/models/{attention => interpretability}/diffex/classifier/models.py (100%) rename src/ops_model/models/{attention => interpretability}/diffex/classifier/run.py (95%) rename src/ops_model/models/{attention => interpretability}/diffex/classifier/submit.py (92%) rename src/ops_model/models/{attention => interpretability}/diffex/classifier/train.py (100%) rename src/ops_model/models/{attention => interpretability}/diffex/diffae/__init__.py (100%) rename src/ops_model/models/{attention => interpretability}/diffex/diffae/config.py (100%) rename src/ops_model/models/{attention => interpretability}/diffex/diffae/data.py (100%) rename src/ops_model/models/{attention => interpretability}/diffex/diffae/diagnose_conditioning.py (98%) rename src/ops_model/models/{attention => interpretability}/diffex/diffae/model.py (100%) rename src/ops_model/models/{attention => interpretability}/diffex/diffae/plot_metrics.py (100%) rename src/ops_model/models/{attention => interpretability}/diffex/diffae/recon.py (100%) rename src/ops_model/models/{attention => interpretability}/diffex/diffae/run.py (97%) rename src/ops_model/models/{attention => interpretability}/diffex/diffae/submit.py (98%) rename src/ops_model/models/{attention => interpretability}/diffex/diffae/train.py (100%) rename src/ops_model/models/{attention => interpretability}/diffex/diffae/virtstain_eval.py (98%) rename src/ops_model/models/{attention => interpretability}/diffex/diffae/virtstain_multi.py (99%) rename src/ops_model/models/{attention => interpretability}/diffex/directions/__init__.py (100%) rename src/ops_model/models/{attention => interpretability}/diffex/directions/batch.py (99%) rename src/ops_model/models/{attention => interpretability}/diffex/directions/config.py (100%) rename src/ops_model/models/{attention => interpretability}/diffex/directions/data.py (100%) rename src/ops_model/models/{attention => interpretability}/diffex/directions/flow.py (100%) rename src/ops_model/models/{attention => interpretability}/diffex/directions/grid.py (100%) rename src/ops_model/models/{attention => interpretability}/diffex/directions/losses.py (100%) rename src/ops_model/models/{attention => interpretability}/diffex/directions/make_gifs.py (100%) rename src/ops_model/models/{attention => interpretability}/diffex/directions/model.py (100%) rename src/ops_model/models/{attention => interpretability}/diffex/directions/proto_ddim_anchors.py (99%) rename src/ops_model/models/{attention => interpretability}/diffex/directions/rank.py (100%) rename src/ops_model/models/{attention => interpretability}/diffex/directions/run.py (96%) rename src/ops_model/models/{attention => interpretability}/diffex/directions/submit.py (96%) rename src/ops_model/models/{attention => interpretability}/diffex/directions/train_directions.py (100%) rename src/ops_model/models/{attention => interpretability}/diffex/directions/traverse.py (100%) rename src/ops_model/models/{attention => interpretability}/diffex/figures/METHODS_final.txt (100%) rename src/ops_model/models/{attention => interpretability}/diffex/figures/METHODS_traversal_montage.md (100%) rename src/ops_model/models/{attention => interpretability}/diffex/figures/METHODS_traversal_montage.txt (100%) rename src/ops_model/models/{attention => interpretability}/diffex/figures/_setacc_common.py (94%) rename src/ops_model/models/{attention => interpretability}/diffex/figures/_setacc_phase.py (100%) rename src/ops_model/models/{attention => interpretability}/diffex/figures/auto_pick_and_plot.py (96%) rename src/ops_model/models/{attention => interpretability}/diffex/figures/cis_golgi_alternatives.py (100%) rename src/ops_model/models/{attention => interpretability}/diffex/figures/debug_setacc_top100.py (97%) rename src/ops_model/models/{attention => interpretability}/diffex/figures/ebi_peripheral_droplets.py (98%) rename src/ops_model/models/{attention => interpretability}/diffex/figures/figure4_morpho_traversal.py (100%) rename src/ops_model/models/{attention => interpretability}/diffex/figures/figure4_morpho_violin.py (97%) rename src/ops_model/models/{attention => interpretability}/diffex/figures/figure4_setacc_panel.py (100%) rename src/ops_model/models/{attention => interpretability}/diffex/figures/figure4_setacc_panel_fluorB.py (100%) rename src/ops_model/models/{attention => interpretability}/diffex/figures/figure4_setacc_panel_newpheno.py (100%) rename src/ops_model/models/{attention => interpretability}/diffex/figures/figure4_setacc_panel_phase.py (100%) rename src/ops_model/models/{attention => interpretability}/diffex/figures/figure_ebi_morpho_violin.py (100%) rename src/ops_model/models/{attention => interpretability}/diffex/figures/figure_multirank_ebi_grid.py (100%) rename src/ops_model/models/{attention => interpretability}/diffex/figures/fluor_panel_montages.py (100%) rename src/ops_model/models/{attention => interpretability}/diffex/figures/fluor_shap_montages.py (98%) rename src/ops_model/models/{attention => interpretability}/diffex/figures/gen_validation/bag_sweep_plots.py (99%) rename src/ops_model/models/{attention => interpretability}/diffex/figures/gen_validation/bag_sweep_score.py (90%) rename src/ops_model/models/{attention => interpretability}/diffex/figures/gen_validation/centroid_bagsweep.py (96%) rename src/ops_model/models/{attention => interpretability}/diffex/figures/gen_validation/centroid_halves.py (95%) rename src/ops_model/models/{attention => interpretability}/diffex/figures/gen_validation/centroid_pooled_bagsweep.py (97%) rename src/ops_model/models/{attention => interpretability}/diffex/figures/gen_validation/control_halves_zscore.py (96%) rename src/ops_model/models/{attention => interpretability}/diffex/figures/gen_validation/embcheck.py (86%) rename src/ops_model/models/{attention => interpretability}/diffex/figures/gen_validation/embedding_diagnostics.py (100%) rename src/ops_model/models/{attention => interpretability}/diffex/figures/gen_validation/figure4_bagsize_reachfrac.py (100%) rename src/ops_model/models/{attention => interpretability}/diffex/figures/gen_validation/figure4_v5_accuracy_summary.py (97%) rename src/ops_model/models/{attention => interpretability}/diffex/figures/gen_validation/gen_alpha_embedding.py (100%) rename src/ops_model/models/{attention => interpretability}/diffex/figures/gen_validation/gen_embed_refit.py (99%) rename src/ops_model/models/{attention => interpretability}/diffex/figures/gen_validation/gen_phate_passthrough.py (100%) rename src/ops_model/models/{attention => interpretability}/diffex/figures/gen_validation/gen_real_centroid.py (97%) rename src/ops_model/models/{attention => interpretability}/diffex/figures/gen_validation/gen_real_distinct.py (100%) rename src/ops_model/models/{attention => interpretability}/diffex/figures/gen_validation/ntc_inverse_gap.py (95%) rename src/ops_model/models/{attention => interpretability}/diffex/figures/gen_validation/patch_cache_real.py (90%) rename src/ops_model/models/{attention => interpretability}/diffex/figures/gen_validation/publish_multibag_page.py (100%) rename src/ops_model/models/{attention => interpretability}/diffex/figures/gen_validation/rank_summary.py (100%) rename src/ops_model/models/{attention => interpretability}/diffex/figures/gen_validation/st_halves_score.py (90%) rename src/ops_model/models/{attention => interpretability}/diffex/figures/gen_validation/std_anchor_test.py (100%) rename src/ops_model/models/{attention => interpretability}/diffex/figures/gen_validation/stepabl_compare.py (100%) rename src/ops_model/models/{attention => interpretability}/diffex/figures/gen_validation/valid200_alphastep.py (100%) rename src/ops_model/models/{attention => interpretability}/diffex/figures/gen_validation/valid200_cache_build.py (100%) rename src/ops_model/models/{attention => interpretability}/diffex/figures/gen_validation/valid200_capcheck.py (100%) rename src/ops_model/models/{attention => interpretability}/diffex/figures/gen_validation/valid200_map_compare.py (98%) rename src/ops_model/models/{attention => interpretability}/diffex/figures/gen_validation/valid200_metrics.py (100%) rename src/ops_model/models/{attention => interpretability}/diffex/figures/nc_ratio.py (100%) rename src/ops_model/models/{attention => interpretability}/diffex/figures/ntc_anchor_compare.py (100%) rename src/ops_model/models/{attention => interpretability}/diffex/figures/phase_montages.py (100%) rename src/ops_model/models/{attention => interpretability}/diffex/figures/phase_multibag_montages.py (100%) rename src/ops_model/models/{attention => interpretability}/diffex/figures/phase_sample_montages.py (100%) rename src/ops_model/models/{attention => interpretability}/diffex/figures/rab_candidate_montages.py (92%) rename src/ops_model/models/{attention => interpretability}/diffex/figures/rebuild_traversals_n100.py (95%) rename src/ops_model/models/{attention => interpretability}/diffex/figures/traversal_montage_schematic.py (100%) rename src/ops_model/models/{attention => interpretability}/diffex/figures/virtual_staining_schematic.py (100%) rename src/ops_model/models/{attention => interpretability}/diffex/kyle_pcs/build_static_explorer.py (100%) rename src/ops_model/models/{attention => interpretability}/diffex/kyle_pcs/compute_pc_strips.py (100%) rename src/ops_model/models/{attention => interpretability}/diffex/viewer/__init__.py (100%) rename src/ops_model/models/{attention => interpretability}/diffex/viewer/_altanchor_build.py (100%) rename src/ops_model/models/{attention => interpretability}/diffex/viewer/_anchortest.py (100%) rename src/ops_model/models/{attention => interpretability}/diffex/viewer/_build_stepablation.py (100%) rename src/ops_model/models/{attention => interpretability}/diffex/viewer/_build_v5_inverted.py (99%) rename src/ops_model/models/{attention => interpretability}/diffex/viewer/_build_v5_montages.py (98%) rename src/ops_model/models/{attention => interpretability}/diffex/viewer/_build_valid200.py (95%) rename src/ops_model/models/{attention => interpretability}/diffex/viewer/_consolidate_cells.py (100%) rename src/ops_model/models/{attention => interpretability}/diffex/viewer/_fluor_complex_build.py (100%) rename src/ops_model/models/{attention => interpretability}/diffex/viewer/_fluor_topcells.py (100%) rename src/ops_model/models/{attention => interpretability}/diffex/viewer/_fluor_v5_build.py (100%) rename src/ops_model/models/{attention => interpretability}/diffex/viewer/_migrate_v4_to_v5.py (100%) rename src/ops_model/models/{attention => interpretability}/diffex/viewer/_phase_vs.py (100%) rename src/ops_model/models/{attention => interpretability}/diffex/viewer/_rebuild_v5.py (100%) rename src/ops_model/models/{attention => interpretability}/diffex/viewer/_rescore_rank.py (97%) rename src/ops_model/models/{attention => interpretability}/diffex/viewer/_score_v4.py (100%) rename src/ops_model/models/{attention => interpretability}/diffex/viewer/_v4acc_test.py (100%) rename src/ops_model/models/{attention => interpretability}/diffex/viewer/_verify_pt_space.py (100%) rename src/ops_model/models/{attention => interpretability}/diffex/viewer/_verify_score_bridge.py (100%) rename src/ops_model/models/{attention => interpretability}/diffex/viewer/altanchor_pairs.json (100%) rename src/ops_model/models/{attention => interpretability}/diffex/viewer/anchor_cells.py (100%) rename src/ops_model/models/{attention => interpretability}/diffex/viewer/build_attention_heads.py (95%) rename src/ops_model/models/{attention => interpretability}/diffex/viewer/build_complex_ebi_map.py (100%) rename src/ops_model/models/{attention => interpretability}/diffex/viewer/build_fluor_shap_rankings.py (96%) rename src/ops_model/models/{attention => interpretability}/diffex/viewer/build_montage_features.py (96%) rename src/ops_model/models/{attention => interpretability}/diffex/viewer/build_pc_crops_masked.py (97%) rename src/ops_model/models/{attention => interpretability}/diffex/viewer/build_pc_features.py (99%) rename src/ops_model/models/{attention => interpretability}/diffex/viewer/build_pc_walks.py (97%) rename src/ops_model/models/{attention => interpretability}/diffex/viewer/build_pcs.py (97%) rename src/ops_model/models/{attention => interpretability}/diffex/viewer/build_pcs_marker.py (99%) rename src/ops_model/models/{attention => interpretability}/diffex/viewer/build_phase_shap_rankings.py (95%) rename src/ops_model/models/{attention => interpretability}/diffex/viewer/build_phate_figure.py (99%) rename src/ops_model/models/{attention => interpretability}/diffex/viewer/build_setacc_bins.py (100%) rename src/ops_model/models/{attention => interpretability}/diffex/viewer/build_setacc_bymarker.py (100%) rename src/ops_model/models/{attention => interpretability}/diffex/viewer/build_top_cells.py (97%) rename src/ops_model/models/{attention => interpretability}/diffex/viewer/build_umap_montage.py (100%) rename src/ops_model/models/{attention => interpretability}/diffex/viewer/catalog.py (100%) rename src/ops_model/models/{attention => interpretability}/diffex/viewer/deploy/README.md (100%) rename src/ops_model/models/{attention => interpretability}/diffex/viewer/marker_leaves.py (100%) rename src/ops_model/models/{attention => interpretability}/diffex/viewer/mimic_alex_embed.py (100%) rename src/ops_model/models/{attention => interpretability}/diffex/viewer/morpho_pipeline.py (100%) rename src/ops_model/models/{attention => interpretability}/diffex/viewer/morphometrics.py (100%) rename src/ops_model/models/{attention => interpretability}/diffex/viewer/nway_clf.py (100%) rename src/ops_model/models/{attention => interpretability}/diffex/viewer/phenotype_cells.py (100%) rename src/ops_model/models/{attention => interpretability}/diffex/viewer/precompute.py (100%) rename src/ops_model/models/{attention => interpretability}/diffex/viewer/render_montage_scales.py (99%) rename src/ops_model/models/{attention => interpretability}/diffex/viewer/score_generated.py (100%) rename src/ops_model/models/{attention => interpretability}/diffex/viewer/set_classifier.py (100%) rename src/ops_model/models/{attention => interpretability}/diffex/viewer/submit.py (97%) rename src/ops_model/models/{attention => interpretability}/diffex/viewer/webapp/app.js (100%) rename src/ops_model/models/{attention => interpretability}/diffex/viewer/webapp/biohub-mark.png (100%) rename src/ops_model/models/{attention => interpretability}/diffex/viewer/webapp/biohub-wordmark.png (100%) rename src/ops_model/models/{attention => interpretability}/diffex/viewer/webapp/build_gene_narratives.py (100%) rename src/ops_model/models/{attention => interpretability}/diffex/viewer/webapp/gif.js (100%) rename src/ops_model/models/{attention => interpretability}/diffex/viewer/webapp/gif.worker.js (100%) rename src/ops_model/models/{attention => interpretability}/diffex/viewer/webapp/index.html (100%) rename src/ops_model/models/{attention => interpretability}/diffex/viewer/webapp/methods.js (100%) rename src/ops_model/models/{attention => interpretability}/diffex/viewer/webapp/morpho_demo.html (100%) rename src/ops_model/models/{attention => interpretability}/diffex/viewer/webapp/openseadragon.min.js (100%) rename src/ops_model/models/{attention => interpretability}/diffex/viewer/webapp/opsin-eyes.svg (100%) rename src/ops_model/models/{attention => interpretability}/diffex/viewer/webapp/style.css (100%) rename src/ops_model/models/{attention => interpretability}/embedding/generate_ko_violin_plots.py (100%) rename src/ops_model/models/{attention => interpretability}/embedding/regen_umap_gav.py (100%) rename src/ops_model/models/{attention => interpretability}/embedding/regen_umap_html.py (100%) rename src/ops_model/models/{attention => interpretability}/embedding/run_all_atlases.py (100%) rename src/ops_model/models/{attention => interpretability}/embedding/top_attention_embed_and_score.py (100%) rename src/ops_model/models/{attention => interpretability}/shap/analyze_chad_variants.py (100%) rename src/ops_model/models/{attention => interpretability}/shap/generate_shap_captions_combined.py (100%) rename src/ops_model/models/{attention => interpretability}/shap/ko_shap_features.py (100%) rename src/ops_model/models/{attention => interpretability}/shap/merge_shap_shards.py (100%) rename src/ops_model/models/{attention => interpretability}/shap/ntc_attention_compare.py (100%) rename src/ops_model/models/{attention => interpretability}/shap/ntc_pick_cells.py (100%) rename src/ops_model/models/{attention => interpretability}/shap/ntc_shap_features.py (100%) rename src/ops_model/models/{attention => interpretability}/shap/run_all_shap.py (100%) rename src/ops_model/models/{attention => interpretability}/shap/run_shap_pipeline.py (100%) rename src/ops_model/models/{attention => interpretability}/shap/shap_approach_compare.py (100%) rename src/ops_model/models/{attention => interpretability}/titration/decay/map_attention_decay.py (100%) rename src/ops_model/models/{attention => interpretability}/titration/decay/phate_peak_groups.py (100%) rename src/ops_model/models/{attention => interpretability}/titration/decay/plot_3way_summary_bars.py (100%) rename src/ops_model/models/{attention => interpretability}/titration/decay/plot_all_cells_correction_bars.py (100%) rename src/ops_model/models/{attention => interpretability}/titration/expansion/count_genes_above_threshold.py (98%) rename src/ops_model/models/{attention => interpretability}/titration/expansion/map_attention_expansion_v4.py (100%) rename src/ops_model/models/{attention => interpretability}/titration/expansion/plot_sgrna_coverage_sweep.py (100%) rename src/ops_model/models/{attention => interpretability}/titration/expansion/run_percentile_sweep.py (100%) rename src/ops_model/models/{attention => interpretability}/weighted_aggregation/_v4_attn_worker.py (100%) rename src/ops_model/models/{attention => interpretability}/weighted_aggregation/analyze_v3_acc_bins.py (100%) rename src/ops_model/models/{attention => interpretability}/weighted_aggregation/plot_v4_attn_comparison.py (100%) rename src/ops_model/models/{attention => interpretability}/weighted_aggregation/run_v3_pipeline_on_v4_attn_weighted.py (100%) rename src/ops_model/models/{attention => interpretability}/weighted_aggregation/run_v3_pipeline_on_v4_features.py (100%) diff --git a/src/ops_model/models/attention/RUNBOOK.md b/src/ops_model/models/interpretability/RUNBOOK.md similarity index 100% rename from src/ops_model/models/attention/RUNBOOK.md rename to src/ops_model/models/interpretability/RUNBOOK.md diff --git a/src/ops_model/models/attention/atlas/attention_accuracy_umap_animation.py b/src/ops_model/models/interpretability/atlas/attention_accuracy_umap_animation.py similarity index 100% rename from src/ops_model/models/attention/atlas/attention_accuracy_umap_animation.py rename to src/ops_model/models/interpretability/atlas/attention_accuracy_umap_animation.py diff --git a/src/ops_model/models/attention/atlas/attention_atlas.py b/src/ops_model/models/interpretability/atlas/attention_atlas.py similarity index 100% rename from src/ops_model/models/attention/atlas/attention_atlas.py rename to src/ops_model/models/interpretability/atlas/attention_atlas.py diff --git a/src/ops_model/models/attention/atlas/attention_atlas_shap.py b/src/ops_model/models/interpretability/atlas/attention_atlas_shap.py similarity index 100% rename from src/ops_model/models/attention/atlas/attention_atlas_shap.py rename to src/ops_model/models/interpretability/atlas/attention_atlas_shap.py diff --git a/src/ops_model/models/attention/atlas/low_attention_phase_atlas.py b/src/ops_model/models/interpretability/atlas/low_attention_phase_atlas.py similarity index 100% rename from src/ops_model/models/attention/atlas/low_attention_phase_atlas.py rename to src/ops_model/models/interpretability/atlas/low_attention_phase_atlas.py diff --git a/src/ops_model/models/attention/atlas/make_scale_bar.py b/src/ops_model/models/interpretability/atlas/make_scale_bar.py similarity index 100% rename from src/ops_model/models/attention/atlas/make_scale_bar.py rename to src/ops_model/models/interpretability/atlas/make_scale_bar.py diff --git a/src/ops_model/models/attention/atlas/marker_selection_distribution.py b/src/ops_model/models/interpretability/atlas/marker_selection_distribution.py similarity index 100% rename from src/ops_model/models/attention/atlas/marker_selection_distribution.py rename to src/ops_model/models/interpretability/atlas/marker_selection_distribution.py diff --git a/src/ops_model/models/attention/atlas/plot_eval_accuracy_curves.py b/src/ops_model/models/interpretability/atlas/plot_eval_accuracy_curves.py similarity index 100% rename from src/ops_model/models/attention/atlas/plot_eval_accuracy_curves.py rename to src/ops_model/models/interpretability/atlas/plot_eval_accuracy_curves.py diff --git a/src/ops_model/models/attention/diffex/CROPSEQ_TO_MORPHOLOGY.md b/src/ops_model/models/interpretability/diffex/CROPSEQ_TO_MORPHOLOGY.md similarity index 100% rename from src/ops_model/models/attention/diffex/CROPSEQ_TO_MORPHOLOGY.md rename to src/ops_model/models/interpretability/diffex/CROPSEQ_TO_MORPHOLOGY.md diff --git a/src/ops_model/models/attention/diffex/PLAN.md b/src/ops_model/models/interpretability/diffex/PLAN.md similarity index 100% rename from src/ops_model/models/attention/diffex/PLAN.md rename to src/ops_model/models/interpretability/diffex/PLAN.md diff --git a/src/ops_model/models/attention/diffex/README.md b/src/ops_model/models/interpretability/diffex/README.md similarity index 100% rename from src/ops_model/models/attention/diffex/README.md rename to src/ops_model/models/interpretability/diffex/README.md diff --git a/src/ops_model/models/attention/diffex/classifier/README.md b/src/ops_model/models/interpretability/diffex/classifier/README.md similarity index 100% rename from src/ops_model/models/attention/diffex/classifier/README.md rename to src/ops_model/models/interpretability/diffex/classifier/README.md diff --git a/src/ops_model/models/attention/diffex/classifier/__init__.py b/src/ops_model/models/interpretability/diffex/classifier/__init__.py similarity index 100% rename from src/ops_model/models/attention/diffex/classifier/__init__.py rename to src/ops_model/models/interpretability/diffex/classifier/__init__.py diff --git a/src/ops_model/models/attention/diffex/classifier/aggregate.py b/src/ops_model/models/interpretability/diffex/classifier/aggregate.py similarity index 97% rename from src/ops_model/models/attention/diffex/classifier/aggregate.py rename to src/ops_model/models/interpretability/diffex/classifier/aggregate.py index 68fc0a6..2577bc0 100644 --- a/src/ops_model/models/attention/diffex/classifier/aggregate.py +++ b/src/ops_model/models/interpretability/diffex/classifier/aggregate.py @@ -1,6 +1,6 @@ """Aggregate per-class classifier metrics into a ranked table. - python -m ops_model.models.attention.diffex.classifier.aggregate --grain complex + python -m ops_model.models.interpretability.diffex.classifier.aggregate --grain complex Collects every ///metrics_.json into one CSV ranked by test AUROC (how cleanly/distinctly each class's top-attention cells classify), plus diff --git a/src/ops_model/models/attention/diffex/classifier/celldino_features.py b/src/ops_model/models/interpretability/diffex/classifier/celldino_features.py similarity index 100% rename from src/ops_model/models/attention/diffex/classifier/celldino_features.py rename to src/ops_model/models/interpretability/diffex/classifier/celldino_features.py diff --git a/src/ops_model/models/attention/diffex/classifier/config.py b/src/ops_model/models/interpretability/diffex/classifier/config.py similarity index 100% rename from src/ops_model/models/attention/diffex/classifier/config.py rename to src/ops_model/models/interpretability/diffex/classifier/config.py diff --git a/src/ops_model/models/attention/diffex/classifier/data.py b/src/ops_model/models/interpretability/diffex/classifier/data.py similarity index 100% rename from src/ops_model/models/attention/diffex/classifier/data.py rename to src/ops_model/models/interpretability/diffex/classifier/data.py diff --git a/src/ops_model/models/attention/diffex/classifier/models.py b/src/ops_model/models/interpretability/diffex/classifier/models.py similarity index 100% rename from src/ops_model/models/attention/diffex/classifier/models.py rename to src/ops_model/models/interpretability/diffex/classifier/models.py diff --git a/src/ops_model/models/attention/diffex/classifier/run.py b/src/ops_model/models/interpretability/diffex/classifier/run.py similarity index 95% rename from src/ops_model/models/attention/diffex/classifier/run.py rename to src/ops_model/models/interpretability/diffex/classifier/run.py index 8459f89..afe55ef 100644 --- a/src/ops_model/models/attention/diffex/classifier/run.py +++ b/src/ops_model/models/interpretability/diffex/classifier/run.py @@ -1,7 +1,7 @@ """Orchestrator for the DiffEx single-cell classifier PoC. - python -m ops_model.models.attention.diffex.classifier.run --model B --gene HSPA5 - python -m ops_model.models.attention.diffex.classifier.run --model C --gene HSPA5 + python -m ops_model.models.interpretability.diffex.classifier.run --model B --gene HSPA5 + python -m ops_model.models.interpretability.diffex.classifier.run --model C --gene HSPA5 Shared steps: build cell table -> materialize phase crops (cached) -> split. Then B trains a ResNet on crops; C embeds the crops with CellDINO and trains an diff --git a/src/ops_model/models/attention/diffex/classifier/submit.py b/src/ops_model/models/interpretability/diffex/classifier/submit.py similarity index 92% rename from src/ops_model/models/attention/diffex/classifier/submit.py rename to src/ops_model/models/interpretability/diffex/classifier/submit.py index 3853193..ccf6c72 100644 --- a/src/ops_model/models/attention/diffex/classifier/submit.py +++ b/src/ops_model/models/interpretability/diffex/classifier/submit.py @@ -1,13 +1,13 @@ """Submit the classifier sweep to SLURM (GPU) via submit_parallel_jobs. # one gene, both models - python -m ops_model.models.attention.diffex.classifier.submit --gene HSPA5 --models B C + python -m ops_model.models.interpretability.diffex.classifier.submit --gene HSPA5 --models B C # all 98 EBI complexes, model C - python -m ops_model.models.attention.diffex.classifier.submit --grain complex --all-classes --models C + python -m ops_model.models.interpretability.diffex.classifier.submit --grain complex --all-classes --models C # specific classes - python -m ops_model.models.attention.diffex.classifier.submit --grain complex \ + python -m ops_model.models.interpretability.diffex.classifier.submit --grain complex \ --classes "19S proteasome regulatory complex" "Commander complex" --models C One GPU job per (class, model). Outputs under ///. diff --git a/src/ops_model/models/attention/diffex/classifier/train.py b/src/ops_model/models/interpretability/diffex/classifier/train.py similarity index 100% rename from src/ops_model/models/attention/diffex/classifier/train.py rename to src/ops_model/models/interpretability/diffex/classifier/train.py diff --git a/src/ops_model/models/attention/diffex/diffae/__init__.py b/src/ops_model/models/interpretability/diffex/diffae/__init__.py similarity index 100% rename from src/ops_model/models/attention/diffex/diffae/__init__.py rename to src/ops_model/models/interpretability/diffex/diffae/__init__.py diff --git a/src/ops_model/models/attention/diffex/diffae/config.py b/src/ops_model/models/interpretability/diffex/diffae/config.py similarity index 100% rename from src/ops_model/models/attention/diffex/diffae/config.py rename to src/ops_model/models/interpretability/diffex/diffae/config.py diff --git a/src/ops_model/models/attention/diffex/diffae/data.py b/src/ops_model/models/interpretability/diffex/diffae/data.py similarity index 100% rename from src/ops_model/models/attention/diffex/diffae/data.py rename to src/ops_model/models/interpretability/diffex/diffae/data.py diff --git a/src/ops_model/models/attention/diffex/diffae/diagnose_conditioning.py b/src/ops_model/models/interpretability/diffex/diffae/diagnose_conditioning.py similarity index 98% rename from src/ops_model/models/attention/diffex/diffae/diagnose_conditioning.py rename to src/ops_model/models/interpretability/diffex/diffae/diagnose_conditioning.py index 32b32a8..87f9714 100644 --- a/src/ops_model/models/attention/diffex/diffae/diagnose_conditioning.py +++ b/src/ops_model/models/interpretability/diffex/diffae/diagnose_conditioning.py @@ -6,7 +6,7 @@ MSE between two DIFFERENT noises — if embedding-driven change ≪ noise-driven change, the embedding has weak control (the bug we suspect). - python -m ops_model.models.attention.diffex.diffae.diagnose_conditioning + python -m ops_model.models.interpretability.diffex.diffae.diagnose_conditioning """ from __future__ import annotations diff --git a/src/ops_model/models/attention/diffex/diffae/model.py b/src/ops_model/models/interpretability/diffex/diffae/model.py similarity index 100% rename from src/ops_model/models/attention/diffex/diffae/model.py rename to src/ops_model/models/interpretability/diffex/diffae/model.py diff --git a/src/ops_model/models/attention/diffex/diffae/plot_metrics.py b/src/ops_model/models/interpretability/diffex/diffae/plot_metrics.py similarity index 100% rename from src/ops_model/models/attention/diffex/diffae/plot_metrics.py rename to src/ops_model/models/interpretability/diffex/diffae/plot_metrics.py diff --git a/src/ops_model/models/attention/diffex/diffae/recon.py b/src/ops_model/models/interpretability/diffex/diffae/recon.py similarity index 100% rename from src/ops_model/models/attention/diffex/diffae/recon.py rename to src/ops_model/models/interpretability/diffex/diffae/recon.py diff --git a/src/ops_model/models/attention/diffex/diffae/run.py b/src/ops_model/models/interpretability/diffex/diffae/run.py similarity index 97% rename from src/ops_model/models/attention/diffex/diffae/run.py rename to src/ops_model/models/interpretability/diffex/diffae/run.py index 61d913a..e78a0cb 100644 --- a/src/ops_model/models/attention/diffex/diffae/run.py +++ b/src/ops_model/models/interpretability/diffex/diffae/run.py @@ -1,6 +1,6 @@ """Orchestrator for the DiffAE generator stage. - python -m ops_model.models.attention.diffex.diffae.run + python -m ops_model.models.interpretability.diffex.diffae.run Steps: sample broad phase crops (cached) -> normalize -> train DiffAE (joint encoder + conditional UNet) -> periodic + final reconstruction gate. Writes diff --git a/src/ops_model/models/attention/diffex/diffae/submit.py b/src/ops_model/models/interpretability/diffex/diffae/submit.py similarity index 98% rename from src/ops_model/models/attention/diffex/diffae/submit.py rename to src/ops_model/models/interpretability/diffex/diffae/submit.py index ce3bea2..d3fe3c3 100644 --- a/src/ops_model/models/attention/diffex/diffae/submit.py +++ b/src/ops_model/models/interpretability/diffex/diffae/submit.py @@ -1,6 +1,6 @@ """Submit the DiffAE training to SLURM (1 GPU, longer wall clock). - python -m ops_model.models.attention.diffex.diffae.submit + python -m ops_model.models.interpretability.diffex.diffae.submit =============================== RUNBOOK =============================== Checkpoints (root /hpc/projects/icd.fast.ops/models/diffex/diffae/) and their diff --git a/src/ops_model/models/attention/diffex/diffae/train.py b/src/ops_model/models/interpretability/diffex/diffae/train.py similarity index 100% rename from src/ops_model/models/attention/diffex/diffae/train.py rename to src/ops_model/models/interpretability/diffex/diffae/train.py diff --git a/src/ops_model/models/attention/diffex/diffae/virtstain_eval.py b/src/ops_model/models/interpretability/diffex/diffae/virtstain_eval.py similarity index 98% rename from src/ops_model/models/attention/diffex/diffae/virtstain_eval.py rename to src/ops_model/models/interpretability/diffex/diffae/virtstain_eval.py index 0800deb..c08df1a 100644 --- a/src/ops_model/models/attention/diffex/diffae/virtstain_eval.py +++ b/src/ops_model/models/interpretability/diffex/diffae/virtstain_eval.py @@ -2,7 +2,7 @@ embedding on a HELD-OUT set of cells (fresh seed → disjoint from training), then report Pearson(pred, real) and save a `phase | predicted | real` montage. - python -m ops_model.models.attention.diffex.diffae.virtstain_eval \ + python -m ops_model.models.interpretability.diffex.diffae.virtstain_eval \ --out-dir /hpc/projects/icd.fast.ops/analysis/virtual_staining/chromalive561_from_phase \ --marker-channel "mitochondria_ChromaLIVE 561 excitation" --channel mCherry --cond-channel Phase2D """ diff --git a/src/ops_model/models/attention/diffex/diffae/virtstain_multi.py b/src/ops_model/models/interpretability/diffex/diffae/virtstain_multi.py similarity index 99% rename from src/ops_model/models/attention/diffex/diffae/virtstain_multi.py rename to src/ops_model/models/interpretability/diffex/diffae/virtstain_multi.py index b7c6eb7..90d731e 100644 --- a/src/ops_model/models/attention/diffex/diffae/virtstain_multi.py +++ b/src/ops_model/models/interpretability/diffex/diffae/virtstain_multi.py @@ -3,7 +3,7 @@ (each marker's own exps); the phase image is concatenated into the UNet (registered stain) and the marker id selects which channel to render. - python -m ops_model.models.attention.diffex.diffae.virtstain_multi --submit --cap 2500 --epochs 120 + python -m ops_model.models.interpretability.diffex.diffae.virtstain_multi --submit --cap 2500 --epochs 120 """ from __future__ import annotations diff --git a/src/ops_model/models/attention/diffex/directions/__init__.py b/src/ops_model/models/interpretability/diffex/directions/__init__.py similarity index 100% rename from src/ops_model/models/attention/diffex/directions/__init__.py rename to src/ops_model/models/interpretability/diffex/directions/__init__.py diff --git a/src/ops_model/models/attention/diffex/directions/batch.py b/src/ops_model/models/interpretability/diffex/directions/batch.py similarity index 99% rename from src/ops_model/models/attention/diffex/directions/batch.py rename to src/ops_model/models/interpretability/diffex/directions/batch.py index 714185b..eea65fc 100644 --- a/src/ops_model/models/attention/diffex/directions/batch.py +++ b/src/ops_model/models/interpretability/diffex/directions/batch.py @@ -4,7 +4,7 @@ submits one GPU job per target. Each job: run_directions at w=5 only → per-cell strips + scores, then a GIF for the auto-picked best-Δscore cell. - python -m ops_model.models.attention.diffex.directions.batch \ + python -m ops_model.models.interpretability.diffex.directions.batch \ --genes-csv --complex-csv \ --n-genes 50 --n-complex 20 """ diff --git a/src/ops_model/models/attention/diffex/directions/config.py b/src/ops_model/models/interpretability/diffex/directions/config.py similarity index 100% rename from src/ops_model/models/attention/diffex/directions/config.py rename to src/ops_model/models/interpretability/diffex/directions/config.py diff --git a/src/ops_model/models/attention/diffex/directions/data.py b/src/ops_model/models/interpretability/diffex/directions/data.py similarity index 100% rename from src/ops_model/models/attention/diffex/directions/data.py rename to src/ops_model/models/interpretability/diffex/directions/data.py diff --git a/src/ops_model/models/attention/diffex/directions/flow.py b/src/ops_model/models/interpretability/diffex/directions/flow.py similarity index 100% rename from src/ops_model/models/attention/diffex/directions/flow.py rename to src/ops_model/models/interpretability/diffex/directions/flow.py diff --git a/src/ops_model/models/attention/diffex/directions/grid.py b/src/ops_model/models/interpretability/diffex/directions/grid.py similarity index 100% rename from src/ops_model/models/attention/diffex/directions/grid.py rename to src/ops_model/models/interpretability/diffex/directions/grid.py diff --git a/src/ops_model/models/attention/diffex/directions/losses.py b/src/ops_model/models/interpretability/diffex/directions/losses.py similarity index 100% rename from src/ops_model/models/attention/diffex/directions/losses.py rename to src/ops_model/models/interpretability/diffex/directions/losses.py diff --git a/src/ops_model/models/attention/diffex/directions/make_gifs.py b/src/ops_model/models/interpretability/diffex/directions/make_gifs.py similarity index 100% rename from src/ops_model/models/attention/diffex/directions/make_gifs.py rename to src/ops_model/models/interpretability/diffex/directions/make_gifs.py diff --git a/src/ops_model/models/attention/diffex/directions/model.py b/src/ops_model/models/interpretability/diffex/directions/model.py similarity index 100% rename from src/ops_model/models/attention/diffex/directions/model.py rename to src/ops_model/models/interpretability/diffex/directions/model.py diff --git a/src/ops_model/models/attention/diffex/directions/proto_ddim_anchors.py b/src/ops_model/models/interpretability/diffex/directions/proto_ddim_anchors.py similarity index 99% rename from src/ops_model/models/attention/diffex/directions/proto_ddim_anchors.py rename to src/ops_model/models/interpretability/diffex/directions/proto_ddim_anchors.py index 131ed36..fef1cc8 100644 --- a/src/ops_model/models/attention/diffex/directions/proto_ddim_anchors.py +++ b/src/ops_model/models/interpretability/diffex/directions/proto_ddim_anchors.py @@ -9,7 +9,7 @@ Same NTC anchor cells v5 uses (top-rank NTC), KIF23 + POLR1B, α 0->5. Everything reused from the existing traversal stack; the only new step is _ddim(..., inverse=True) to get xT. - python -m ops_model.models.attention.diffex.directions.proto_ddim_anchors --submit + python -m ops_model.models.interpretability.diffex.directions.proto_ddim_anchors --submit """ from __future__ import annotations diff --git a/src/ops_model/models/attention/diffex/directions/rank.py b/src/ops_model/models/interpretability/diffex/directions/rank.py similarity index 100% rename from src/ops_model/models/attention/diffex/directions/rank.py rename to src/ops_model/models/interpretability/diffex/directions/rank.py diff --git a/src/ops_model/models/attention/diffex/directions/run.py b/src/ops_model/models/interpretability/diffex/directions/run.py similarity index 96% rename from src/ops_model/models/attention/diffex/directions/run.py rename to src/ops_model/models/interpretability/diffex/directions/run.py index 3f9a37b..5971d42 100644 --- a/src/ops_model/models/attention/diffex/directions/run.py +++ b/src/ops_model/models/interpretability/diffex/directions/run.py @@ -1,7 +1,7 @@ """Orchestrator for Stage 3 (directions → ranking → traversal). - python -m ops_model.models.attention.diffex.directions.run --target HSPA5 - python -m ops_model.models.attention.diffex.directions.run --grain complex \ + python -m ops_model.models.interpretability.diffex.directions.run --target HSPA5 + python -m ops_model.models.interpretability.diffex.directions.run --grain complex \ --target "Chaperonin-containing T-complex" Steps: gather target+control crops/embeddings → train K direction MLPs (unsupervised) diff --git a/src/ops_model/models/attention/diffex/directions/submit.py b/src/ops_model/models/interpretability/diffex/directions/submit.py similarity index 96% rename from src/ops_model/models/attention/diffex/directions/submit.py rename to src/ops_model/models/interpretability/diffex/directions/submit.py index 43e7ef0..d131efd 100644 --- a/src/ops_model/models/attention/diffex/directions/submit.py +++ b/src/ops_model/models/interpretability/diffex/directions/submit.py @@ -1,6 +1,6 @@ """Submit Stage 3 (directions + traversal) to SLURM (1 GPU). - python -m ops_model.models.attention.diffex.directions.submit --target HSPA5 + python -m ops_model.models.interpretability.diffex.directions.submit --target HSPA5 """ from __future__ import annotations diff --git a/src/ops_model/models/attention/diffex/directions/train_directions.py b/src/ops_model/models/interpretability/diffex/directions/train_directions.py similarity index 100% rename from src/ops_model/models/attention/diffex/directions/train_directions.py rename to src/ops_model/models/interpretability/diffex/directions/train_directions.py diff --git a/src/ops_model/models/attention/diffex/directions/traverse.py b/src/ops_model/models/interpretability/diffex/directions/traverse.py similarity index 100% rename from src/ops_model/models/attention/diffex/directions/traverse.py rename to src/ops_model/models/interpretability/diffex/directions/traverse.py diff --git a/src/ops_model/models/attention/diffex/figures/METHODS_final.txt b/src/ops_model/models/interpretability/diffex/figures/METHODS_final.txt similarity index 100% rename from src/ops_model/models/attention/diffex/figures/METHODS_final.txt rename to src/ops_model/models/interpretability/diffex/figures/METHODS_final.txt diff --git a/src/ops_model/models/attention/diffex/figures/METHODS_traversal_montage.md b/src/ops_model/models/interpretability/diffex/figures/METHODS_traversal_montage.md similarity index 100% rename from src/ops_model/models/attention/diffex/figures/METHODS_traversal_montage.md rename to src/ops_model/models/interpretability/diffex/figures/METHODS_traversal_montage.md diff --git a/src/ops_model/models/attention/diffex/figures/METHODS_traversal_montage.txt b/src/ops_model/models/interpretability/diffex/figures/METHODS_traversal_montage.txt similarity index 100% rename from src/ops_model/models/attention/diffex/figures/METHODS_traversal_montage.txt rename to src/ops_model/models/interpretability/diffex/figures/METHODS_traversal_montage.txt diff --git a/src/ops_model/models/attention/diffex/figures/_setacc_common.py b/src/ops_model/models/interpretability/diffex/figures/_setacc_common.py similarity index 94% rename from src/ops_model/models/attention/diffex/figures/_setacc_common.py rename to src/ops_model/models/interpretability/diffex/figures/_setacc_common.py index ad739d3..702475d 100644 --- a/src/ops_model/models/attention/diffex/figures/_setacc_common.py +++ b/src/ops_model/models/interpretability/diffex/figures/_setacc_common.py @@ -10,11 +10,11 @@ import pandas as pd import zarr -from ops_model.models.attention.diffex.classifier.config import slugify -from ops_model.models.attention.diffex.classifier.data import make_labels_df, materialize_crops -from ops_model.models.attention.diffex.directions.config import DirConfig -from ops_model.models.attention.diffex.viewer._fluor_topcells import _overlay_rgba -from ops_model.models.attention.diffex.viewer.build_pc_crops_masked import BASE, CROP_SIZE, _crop, _zarr_patch +from ops_model.models.interpretability.diffex.classifier.config import slugify +from ops_model.models.interpretability.diffex.classifier.data import make_labels_df, materialize_crops +from ops_model.models.interpretability.diffex.directions.config import DirConfig +from ops_model.models.interpretability.diffex.viewer._fluor_topcells import _overlay_rgba +from ops_model.models.interpretability.diffex.viewer.build_pc_crops_masked import BASE, CROP_SIZE, _crop, _zarr_patch OUT = "/hpc/projects/icd.fast.ops/analysis/figure4_setacc_panel" RANK_BASE = "/hpc/projects/icd.fast.ops/models/diffex/viewer_assets_v5/_rankings/fluor_shap" diff --git a/src/ops_model/models/attention/diffex/figures/_setacc_phase.py b/src/ops_model/models/interpretability/diffex/figures/_setacc_phase.py similarity index 100% rename from src/ops_model/models/attention/diffex/figures/_setacc_phase.py rename to src/ops_model/models/interpretability/diffex/figures/_setacc_phase.py diff --git a/src/ops_model/models/attention/diffex/figures/auto_pick_and_plot.py b/src/ops_model/models/interpretability/diffex/figures/auto_pick_and_plot.py similarity index 96% rename from src/ops_model/models/attention/diffex/figures/auto_pick_and_plot.py rename to src/ops_model/models/interpretability/diffex/figures/auto_pick_and_plot.py index 8323212..234004e 100644 --- a/src/ops_model/models/attention/diffex/figures/auto_pick_and_plot.py +++ b/src/ops_model/models/interpretability/diffex/figures/auto_pick_and_plot.py @@ -10,8 +10,8 @@ sys.path.insert(0, os.path.dirname(os.path.abspath(__file__))) os.environ.setdefault("OPS_DIFFEX_ASSETS", "viewer_assets_v5") -from ops_model.models.attention.diffex.viewer.morpho_pipeline import MORPHO_TARGETS -from ops_model.models.attention.diffex.classifier.config import slugify +from ops_model.models.interpretability.diffex.viewer.morpho_pipeline import MORPHO_TARGETS +from ops_model.models.interpretability.diffex.classifier.config import slugify VA = "/hpc/projects/icd.fast.ops/models/diffex/viewer_assets_v5/_morphometrics" BAD = re.compile("moment|hu_|inertia|eigval|intensity|haralick|zernike|glcm|orientation|centroid|_timing") diff --git a/src/ops_model/models/attention/diffex/figures/cis_golgi_alternatives.py b/src/ops_model/models/interpretability/diffex/figures/cis_golgi_alternatives.py similarity index 100% rename from src/ops_model/models/attention/diffex/figures/cis_golgi_alternatives.py rename to src/ops_model/models/interpretability/diffex/figures/cis_golgi_alternatives.py diff --git a/src/ops_model/models/attention/diffex/figures/debug_setacc_top100.py b/src/ops_model/models/interpretability/diffex/figures/debug_setacc_top100.py similarity index 97% rename from src/ops_model/models/attention/diffex/figures/debug_setacc_top100.py rename to src/ops_model/models/interpretability/diffex/figures/debug_setacc_top100.py index 46a44ae..80f8e6a 100644 --- a/src/ops_model/models/attention/diffex/figures/debug_setacc_top100.py +++ b/src/ops_model/models/interpretability/diffex/figures/debug_setacc_top100.py @@ -15,7 +15,7 @@ import matplotlib.pyplot as plt import numpy as np -from ops_model.models.attention.diffex.classifier.config import slugify +from ops_model.models.interpretability.diffex.classifier.config import slugify from _setacc_common import (COMPLEX_COLS, GENE_COLS, OUT, CROP_SIZE, materialize_class, seg_crop, composite) plt.rcParams["pdf.fonttype"] = 42 diff --git a/src/ops_model/models/attention/diffex/figures/ebi_peripheral_droplets.py b/src/ops_model/models/interpretability/diffex/figures/ebi_peripheral_droplets.py similarity index 98% rename from src/ops_model/models/attention/diffex/figures/ebi_peripheral_droplets.py rename to src/ops_model/models/interpretability/diffex/figures/ebi_peripheral_droplets.py index 25c7bcd..2f671c1 100644 --- a/src/ops_model/models/attention/diffex/figures/ebi_peripheral_droplets.py +++ b/src/ops_model/models/interpretability/diffex/figures/ebi_peripheral_droplets.py @@ -41,7 +41,7 @@ from figure_ebi_morpho_violin import draw_violin from figure_multirank_ebi_grid import CACHE, OUT, ebi_rows, top_rows -from ops_model.models.attention.diffex.viewer.build_pc_crops_masked import BASE, _crop, _zarr_patch +from ops_model.models.interpretability.diffex.viewer.build_pc_crops_masked import BASE, _crop, _zarr_patch from organelle_profiler.feature_extraction.localization_features import compute_localization_features diff --git a/src/ops_model/models/attention/diffex/figures/figure4_morpho_traversal.py b/src/ops_model/models/interpretability/diffex/figures/figure4_morpho_traversal.py similarity index 100% rename from src/ops_model/models/attention/diffex/figures/figure4_morpho_traversal.py rename to src/ops_model/models/interpretability/diffex/figures/figure4_morpho_traversal.py diff --git a/src/ops_model/models/attention/diffex/figures/figure4_morpho_violin.py b/src/ops_model/models/interpretability/diffex/figures/figure4_morpho_violin.py similarity index 97% rename from src/ops_model/models/attention/diffex/figures/figure4_morpho_violin.py rename to src/ops_model/models/interpretability/diffex/figures/figure4_morpho_violin.py index 97dd87b..fc6aafb 100644 --- a/src/ops_model/models/attention/diffex/figures/figure4_morpho_violin.py +++ b/src/ops_model/models/interpretability/diffex/figures/figure4_morpho_violin.py @@ -17,8 +17,8 @@ import pandas as pd from figure4_morpho_traversal import FIGURES, VA, image_panels, render_images -from ops_model.models.attention.diffex.viewer.morpho_pipeline import MORPHO_TARGETS, real_percell -from ops_model.models.attention.diffex.classifier.config import slugify +from ops_model.models.interpretability.diffex.viewer.morpho_pipeline import MORPHO_TARGETS, real_percell +from ops_model.models.interpretability.diffex.classifier.config import slugify plt.rcParams["pdf.fonttype"] = 42 plt.rcParams["svg.fonttype"] = "none" diff --git a/src/ops_model/models/attention/diffex/figures/figure4_setacc_panel.py b/src/ops_model/models/interpretability/diffex/figures/figure4_setacc_panel.py similarity index 100% rename from src/ops_model/models/attention/diffex/figures/figure4_setacc_panel.py rename to src/ops_model/models/interpretability/diffex/figures/figure4_setacc_panel.py diff --git a/src/ops_model/models/attention/diffex/figures/figure4_setacc_panel_fluorB.py b/src/ops_model/models/interpretability/diffex/figures/figure4_setacc_panel_fluorB.py similarity index 100% rename from src/ops_model/models/attention/diffex/figures/figure4_setacc_panel_fluorB.py rename to src/ops_model/models/interpretability/diffex/figures/figure4_setacc_panel_fluorB.py diff --git a/src/ops_model/models/attention/diffex/figures/figure4_setacc_panel_newpheno.py b/src/ops_model/models/interpretability/diffex/figures/figure4_setacc_panel_newpheno.py similarity index 100% rename from src/ops_model/models/attention/diffex/figures/figure4_setacc_panel_newpheno.py rename to src/ops_model/models/interpretability/diffex/figures/figure4_setacc_panel_newpheno.py diff --git a/src/ops_model/models/attention/diffex/figures/figure4_setacc_panel_phase.py b/src/ops_model/models/interpretability/diffex/figures/figure4_setacc_panel_phase.py similarity index 100% rename from src/ops_model/models/attention/diffex/figures/figure4_setacc_panel_phase.py rename to src/ops_model/models/interpretability/diffex/figures/figure4_setacc_panel_phase.py diff --git a/src/ops_model/models/attention/diffex/figures/figure_ebi_morpho_violin.py b/src/ops_model/models/interpretability/diffex/figures/figure_ebi_morpho_violin.py similarity index 100% rename from src/ops_model/models/attention/diffex/figures/figure_ebi_morpho_violin.py rename to src/ops_model/models/interpretability/diffex/figures/figure_ebi_morpho_violin.py diff --git a/src/ops_model/models/attention/diffex/figures/figure_multirank_ebi_grid.py b/src/ops_model/models/interpretability/diffex/figures/figure_multirank_ebi_grid.py similarity index 100% rename from src/ops_model/models/attention/diffex/figures/figure_multirank_ebi_grid.py rename to src/ops_model/models/interpretability/diffex/figures/figure_multirank_ebi_grid.py diff --git a/src/ops_model/models/attention/diffex/figures/fluor_panel_montages.py b/src/ops_model/models/interpretability/diffex/figures/fluor_panel_montages.py similarity index 100% rename from src/ops_model/models/attention/diffex/figures/fluor_panel_montages.py rename to src/ops_model/models/interpretability/diffex/figures/fluor_panel_montages.py diff --git a/src/ops_model/models/attention/diffex/figures/fluor_shap_montages.py b/src/ops_model/models/interpretability/diffex/figures/fluor_shap_montages.py similarity index 98% rename from src/ops_model/models/attention/diffex/figures/fluor_shap_montages.py rename to src/ops_model/models/interpretability/diffex/figures/fluor_shap_montages.py index 7048fee..3fe62eb 100644 --- a/src/ops_model/models/attention/diffex/figures/fluor_shap_montages.py +++ b/src/ops_model/models/interpretability/diffex/figures/fluor_shap_montages.py @@ -15,7 +15,7 @@ import numpy as np import pandas as pd -from ops_model.models.attention.diffex.classifier.config import slugify +from ops_model.models.interpretability.diffex.classifier.config import slugify from _setacc_common import _materialize, seg_crop, composite, CROP_SIZE MR = "/hpc/projects/icd.fast.ops/models/alex_lin_attention/v5/multi_rank/shap_screen/shap_screen_fluor_all.compact.parquet" diff --git a/src/ops_model/models/attention/diffex/figures/gen_validation/bag_sweep_plots.py b/src/ops_model/models/interpretability/diffex/figures/gen_validation/bag_sweep_plots.py similarity index 99% rename from src/ops_model/models/attention/diffex/figures/gen_validation/bag_sweep_plots.py rename to src/ops_model/models/interpretability/diffex/figures/gen_validation/bag_sweep_plots.py index daf3d7e..f1e77d8 100644 --- a/src/ops_model/models/attention/diffex/figures/gen_validation/bag_sweep_plots.py +++ b/src/ops_model/models/interpretability/diffex/figures/gen_validation/bag_sweep_plots.py @@ -29,7 +29,7 @@ CENT_SUF = "_perbag" if CENT_STD == "perbag" else "" COL = {b: cm.viridis(i / (len(BAGS) - 1)) for i, b in enumerate(BAGS)} COLK = {k: cm.viridis(i / (len(KS) - 1)) for i, k in enumerate(KS)} -from ops_model.models.attention.diffex.classifier.config import slugify +from ops_model.models.interpretability.diffex.classifier.config import slugify REAL = json.load(open(f"{B}/viewer_assets_v5/real_acc20.json")) diff --git a/src/ops_model/models/attention/diffex/figures/gen_validation/bag_sweep_score.py b/src/ops_model/models/interpretability/diffex/figures/gen_validation/bag_sweep_score.py similarity index 90% rename from src/ops_model/models/attention/diffex/figures/gen_validation/bag_sweep_score.py rename to src/ops_model/models/interpretability/diffex/figures/gen_validation/bag_sweep_score.py index a7a6e50..773bd66 100644 --- a/src/ops_model/models/attention/diffex/figures/gen_validation/bag_sweep_score.py +++ b/src/ops_model/models/interpretability/diffex/figures/gen_validation/bag_sweep_score.py @@ -16,11 +16,11 @@ def score_shard(genes): import torch - from ops_model.models.attention.diffex.viewer.score_generated import score_embs_v5 - from ops_model.models.attention.diffex.viewer.set_classifier import load_set_classifier, V5_CKPT_ROOT, V5_RUNS + from ops_model.models.interpretability.diffex.viewer.score_generated import score_embs_v5 + from ops_model.models.interpretability.diffex.viewer.set_classifier import load_set_classifier, V5_CKPT_ROOT, V5_RUNS os.makedirs(OUT, exist_ok=True) dev = "cuda" if torch.cuda.is_available() else "cpu" - from ops_model.models.attention.diffex.classifier.config import slugify + from ops_model.models.interpretability.diffex.classifier.config import slugify run = V5_RUNS[("phase", "geneKO" if GRAIN == "geneKO" else "complex_ebionly")] m, g2i, c2i = load_set_classifier(run=run, device=dev, root=V5_CKPT_ROOT) ci = c2i.get("Phase2D", 0) diff --git a/src/ops_model/models/attention/diffex/figures/gen_validation/centroid_bagsweep.py b/src/ops_model/models/interpretability/diffex/figures/gen_validation/centroid_bagsweep.py similarity index 96% rename from src/ops_model/models/attention/diffex/figures/gen_validation/centroid_bagsweep.py rename to src/ops_model/models/interpretability/diffex/figures/gen_validation/centroid_bagsweep.py index 6f5df48..2041eff 100644 --- a/src/ops_model/models/attention/diffex/figures/gen_validation/centroid_bagsweep.py +++ b/src/ops_model/models/interpretability/diffex/figures/gen_validation/centroid_bagsweep.py @@ -21,7 +21,7 @@ def _cz(): - from ops_model.models.attention.diffex.classifier.config import slugify + from ops_model.models.interpretability.diffex.classifier.config import slugify d = np.load(f"{CENTD}/{GRAIN}_centroids.npz", allow_pickle=True) names = list(d["names"]); cidx = {slugify(str(c)): i for i, c in enumerate(names)} cz = (d["cents"] - d["mu"]) / d["sd"]; cz = cz / (np.linalg.norm(cz, axis=1, keepdims=True) + 1e-9) @@ -42,7 +42,7 @@ def compute_mu(): def score_shard(genes): - from ops_model.models.attention.diffex.classifier.config import slugify + from ops_model.models.interpretability.diffex.classifier.config import slugify os.makedirs(PART, exist_ok=True) cz, cidx = _cz() if STD == "global": diff --git a/src/ops_model/models/attention/diffex/figures/gen_validation/centroid_halves.py b/src/ops_model/models/interpretability/diffex/figures/gen_validation/centroid_halves.py similarity index 95% rename from src/ops_model/models/attention/diffex/figures/gen_validation/centroid_halves.py rename to src/ops_model/models/interpretability/diffex/figures/gen_validation/centroid_halves.py index 92d73a2..d13cd03 100644 --- a/src/ops_model/models/attention/diffex/figures/gen_validation/centroid_halves.py +++ b/src/ops_model/models/interpretability/diffex/figures/gen_validation/centroid_halves.py @@ -14,7 +14,7 @@ def _cz(): - from ops_model.models.attention.diffex.classifier.config import slugify + from ops_model.models.interpretability.diffex.classifier.config import slugify d = np.load(f"{CENTD}/{GRAIN}_centroids.npz", allow_pickle=True) names = list(d["names"]); cidx = {slugify(str(c)): i for i, c in enumerate(names)} cz = (d["cents"] - d["mu"]) / d["sd"]; cz = cz / (np.linalg.norm(cz, axis=1, keepdims=True) + 1e-9) @@ -29,7 +29,7 @@ def _top1(vecs, cz, ti): def score_shard(genes): - from ops_model.models.attention.diffex.classifier.config import slugify + from ops_model.models.interpretability.diffex.classifier.config import slugify os.makedirs(PART, exist_ok=True) cz, cidx = _cz(); mg = np.load(MU); mu_g, sd_g = mg["mu"], mg["sd"] first, second = {}, {} diff --git a/src/ops_model/models/attention/diffex/figures/gen_validation/centroid_pooled_bagsweep.py b/src/ops_model/models/interpretability/diffex/figures/gen_validation/centroid_pooled_bagsweep.py similarity index 97% rename from src/ops_model/models/attention/diffex/figures/gen_validation/centroid_pooled_bagsweep.py rename to src/ops_model/models/interpretability/diffex/figures/gen_validation/centroid_pooled_bagsweep.py index 1772cc0..36cd9ec 100644 --- a/src/ops_model/models/attention/diffex/figures/gen_validation/centroid_pooled_bagsweep.py +++ b/src/ops_model/models/interpretability/diffex/figures/gen_validation/centroid_pooled_bagsweep.py @@ -23,7 +23,7 @@ def _cz(): - from ops_model.models.attention.diffex.classifier.config import slugify + from ops_model.models.interpretability.diffex.classifier.config import slugify d = np.load(f"{CENTD}/{GRAIN}_centroids.npz", allow_pickle=True) names = list(d["names"]); cidx = {slugify(str(c)): i for i, c in enumerate(names)} cz = (d["cents"] - d["mu"]) / d["sd"]; cz = cz / (np.linalg.norm(cz, axis=1, keepdims=True) + 1e-9) @@ -38,7 +38,7 @@ def _pooled(vecs, cz, ti): def score_shard(genes): - from ops_model.models.attention.diffex.classifier.config import slugify + from ops_model.models.interpretability.diffex.classifier.config import slugify os.makedirs(PART, exist_ok=True) cz, cidx, mu_r, sd_r = _cz() if STD == "global": diff --git a/src/ops_model/models/attention/diffex/figures/gen_validation/control_halves_zscore.py b/src/ops_model/models/interpretability/diffex/figures/gen_validation/control_halves_zscore.py similarity index 96% rename from src/ops_model/models/attention/diffex/figures/gen_validation/control_halves_zscore.py rename to src/ops_model/models/interpretability/diffex/figures/gen_validation/control_halves_zscore.py index 75e5921..b8b9785 100644 --- a/src/ops_model/models/attention/diffex/figures/gen_validation/control_halves_zscore.py +++ b/src/ops_model/models/interpretability/diffex/figures/gen_validation/control_halves_zscore.py @@ -19,7 +19,7 @@ def _cz(): - from ops_model.models.attention.diffex.classifier.config import slugify + from ops_model.models.interpretability.diffex.classifier.config import slugify d = np.load(f"{CENTD}/{GRAIN}_centroids.npz", allow_pickle=True) names = list(d["names"]); cidx = {slugify(str(c)): i for i, c in enumerate(names)} cz = (d["cents"] - d["mu"]) / d["sd"]; cz = cz / (np.linalg.norm(cz, axis=1, keepdims=True) + 1e-9) @@ -35,7 +35,7 @@ def _t1(vecs, cz, ti): def score_shard(genes): - from ops_model.models.attention.diffex.classifier.config import slugify + from ops_model.models.interpretability.diffex.classifier.config import slugify os.makedirs(PART, exist_ok=True) cz, cidx = _cz(); mg = np.load(MU); mu_g, sd_g = mg["mu"], mg["sd"] res = {s: {h: {} for h in ("first", "second")} for s in ("global", "perbag")} diff --git a/src/ops_model/models/attention/diffex/figures/gen_validation/embcheck.py b/src/ops_model/models/interpretability/diffex/figures/gen_validation/embcheck.py similarity index 86% rename from src/ops_model/models/attention/diffex/figures/gen_validation/embcheck.py rename to src/ops_model/models/interpretability/diffex/figures/gen_validation/embcheck.py index 825e412..16d251e 100644 --- a/src/ops_model/models/attention/diffex/figures/gen_validation/embcheck.py +++ b/src/ops_model/models/interpretability/diffex/figures/gen_validation/embcheck.py @@ -8,9 +8,9 @@ def check(gene="AACS", ai=6): - from ops_model.models.attention.diffex.viewer.score_generated import _emb_frames - from ops_model.models.attention.diffex.classifier.celldino_features import embed_crops - from ops_model.models.attention.diffex.directions.config import DirConfig + from ops_model.models.interpretability.diffex.viewer.score_generated import _emb_frames + from ops_model.models.interpretability.diffex.classifier.celldino_features import embed_crops + from ops_model.models.interpretability.diffex.directions.config import DirConfig trav = f"{B}/viewer_assets_valid200/phase/geneKO/{gene}" cfg = DirConfig(grain="geneKO", target=gene, device="cuda") embB = np.asarray(_emb_frames(cfg, trav, ai, embed_crops), np.float32) # original path diff --git a/src/ops_model/models/attention/diffex/figures/gen_validation/embedding_diagnostics.py b/src/ops_model/models/interpretability/diffex/figures/gen_validation/embedding_diagnostics.py similarity index 100% rename from src/ops_model/models/attention/diffex/figures/gen_validation/embedding_diagnostics.py rename to src/ops_model/models/interpretability/diffex/figures/gen_validation/embedding_diagnostics.py diff --git a/src/ops_model/models/attention/diffex/figures/gen_validation/figure4_bagsize_reachfrac.py b/src/ops_model/models/interpretability/diffex/figures/gen_validation/figure4_bagsize_reachfrac.py similarity index 100% rename from src/ops_model/models/attention/diffex/figures/gen_validation/figure4_bagsize_reachfrac.py rename to src/ops_model/models/interpretability/diffex/figures/gen_validation/figure4_bagsize_reachfrac.py diff --git a/src/ops_model/models/attention/diffex/figures/gen_validation/figure4_v5_accuracy_summary.py b/src/ops_model/models/interpretability/diffex/figures/gen_validation/figure4_v5_accuracy_summary.py similarity index 97% rename from src/ops_model/models/attention/diffex/figures/gen_validation/figure4_v5_accuracy_summary.py rename to src/ops_model/models/interpretability/diffex/figures/gen_validation/figure4_v5_accuracy_summary.py index b7eb101..796b604 100644 --- a/src/ops_model/models/attention/diffex/figures/gen_validation/figure4_v5_accuracy_summary.py +++ b/src/ops_model/models/interpretability/diffex/figures/gen_validation/figure4_v5_accuracy_summary.py @@ -34,7 +34,7 @@ def _complex_allow(thr=0.9): label_name in the ebionly eval CSV (his reported members), then mean of those member accuracies.""" import csv as _csv from collections import defaultdict - from ops_model.models.attention.diffex.classifier.config import slugify + from ops_model.models.interpretability.diffex.classifier.config import slugify by = defaultdict(list) for r in _csv.DictReader(open(f"{EVAL}/eval_phase_ebionly_e200_pergene_val.csv")): if int(r["n_cells"]) == 20: @@ -48,7 +48,7 @@ def _real_map(sub): return _real_acc20("eval_phase_e200_pergene_val.csv") import csv as _csv from collections import defaultdict - from ops_model.models.attention.diffex.classifier.config import slugify + from ops_model.models.interpretability.diffex.classifier.config import slugify by = defaultdict(list) for r in _csv.DictReader(open(f"{EVAL}/eval_phase_ebionly_e200_pergene_val.csv")): if int(r["n_cells"]) == 20: diff --git a/src/ops_model/models/attention/diffex/figures/gen_validation/gen_alpha_embedding.py b/src/ops_model/models/interpretability/diffex/figures/gen_validation/gen_alpha_embedding.py similarity index 100% rename from src/ops_model/models/attention/diffex/figures/gen_validation/gen_alpha_embedding.py rename to src/ops_model/models/interpretability/diffex/figures/gen_validation/gen_alpha_embedding.py diff --git a/src/ops_model/models/attention/diffex/figures/gen_validation/gen_embed_refit.py b/src/ops_model/models/interpretability/diffex/figures/gen_validation/gen_embed_refit.py similarity index 99% rename from src/ops_model/models/attention/diffex/figures/gen_validation/gen_embed_refit.py rename to src/ops_model/models/interpretability/diffex/figures/gen_validation/gen_embed_refit.py index a023c5d..5e45314 100644 --- a/src/ops_model/models/attention/diffex/figures/gen_validation/gen_embed_refit.py +++ b/src/ops_model/models/interpretability/diffex/figures/gen_validation/gen_embed_refit.py @@ -155,8 +155,8 @@ def webp_vs_float_v5(n=120): import os, glob from PIL import Image from scipy.spatial.distance import cdist - from ops_model.models.attention.diffex.classifier.celldino_features import embed_crops - from ops_model.models.attention.diffex.directions.config import DirConfig + from ops_model.models.interpretability.diffex.classifier.celldino_features import embed_crops + from ops_model.models.interpretability.diffex.directions.config import DirConfig a, comp, mean = gp._load_embedding() Xr = np.asarray(a.obsm["X_pca"], np.float64); idx = {nm: i for i, nm in enumerate(a.obs_names)} genes = [g for g in sorted(os.listdir(V5WEBP)) if g in idx diff --git a/src/ops_model/models/attention/diffex/figures/gen_validation/gen_phate_passthrough.py b/src/ops_model/models/interpretability/diffex/figures/gen_validation/gen_phate_passthrough.py similarity index 100% rename from src/ops_model/models/attention/diffex/figures/gen_validation/gen_phate_passthrough.py rename to src/ops_model/models/interpretability/diffex/figures/gen_validation/gen_phate_passthrough.py diff --git a/src/ops_model/models/attention/diffex/figures/gen_validation/gen_real_centroid.py b/src/ops_model/models/interpretability/diffex/figures/gen_validation/gen_real_centroid.py similarity index 97% rename from src/ops_model/models/attention/diffex/figures/gen_validation/gen_real_centroid.py rename to src/ops_model/models/interpretability/diffex/figures/gen_validation/gen_real_centroid.py index 16c5bcc..116c93a 100644 --- a/src/ops_model/models/attention/diffex/figures/gen_validation/gen_real_centroid.py +++ b/src/ops_model/models/interpretability/diffex/figures/gen_validation/gen_real_centroid.py @@ -24,9 +24,9 @@ def _classes(grain): def embed_centroids(grain, classes): - from ops_model.models.attention.diffex.viewer.precompute import _gather_class - from ops_model.models.attention.diffex.directions.config import DirConfig - from ops_model.models.attention.diffex.classifier.config import slugify + from ops_model.models.interpretability.diffex.viewer.precompute import _gather_class + from ops_model.models.interpretability.diffex.directions.config import DirConfig + from ops_model.models.interpretability.diffex.classifier.config import slugify cfg = DirConfig(grain=grain, target=classes[0], device="cuda"); cfg.num_workers = 12 cents, S, SS, n = {}, np.zeros(1024), np.zeros(1024), 0 embs, lbl = [], [] @@ -61,7 +61,7 @@ def merge(grain): def score(grain, cache=None, out=None, cap=None): - from ops_model.models.attention.diffex.classifier.config import slugify + from ops_model.models.interpretability.diffex.classifier.config import slugify CACHE_ = cache or CACHE; OUT_ = out or OUT # explicit args (SLURM-safe) override module globals cap = cap if cap is not None else (int(os.environ.get("GRC_GEN_CAP", "0")) or None) # subsample gen bag to first `cap` cells/class d = np.load(f"{OUT_}/{grain}_centroids.npz", allow_pickle=True) @@ -107,7 +107,7 @@ def score(grain, cache=None, out=None, cap=None): def ceiling(grain): """Real cells (cached embed_crops) → nearest faithful centroid: per-class real mAP/top1/top5 (the ceiling).""" - from ops_model.models.attention.diffex.classifier.config import slugify + from ops_model.models.interpretability.diffex.classifier.config import slugify d = np.load(f"{OUT}/{grain}_centroids.npz", allow_pickle=True) names = list(d["names"]); cidx = {slugify(str(c)): i for i, c in enumerate(names)} cz = (d["cents"] - d["mu"]) / d["sd"]; cz = cz / (np.linalg.norm(cz, axis=1, keepdims=True) + 1e-9) @@ -136,7 +136,7 @@ def plot(min_dist=None, acc_thr=None, fname="centroid_topk", overlay=None, overl """3 scores (mAP, top-1, top-5) vs α from the faithful centroids. min_dist: restrict to classes whose real distinctiveness/EBI mAP@20 > min_dist. acc_thr: restrict to real top1_acc>acc_thr @bag20 (SetTransformer subset). overlay: {grain: scored_dir} → dashed second line (e.g. the 200-cell bag) on that grain's axis.""" - from ops_model.models.attention.diffex.classifier.config import slugify + from ops_model.models.interpretability.diffex.classifier.config import slugify import matplotlib matplotlib.use("Agg") import matplotlib.pyplot as plt diff --git a/src/ops_model/models/attention/diffex/figures/gen_validation/gen_real_distinct.py b/src/ops_model/models/interpretability/diffex/figures/gen_validation/gen_real_distinct.py similarity index 100% rename from src/ops_model/models/attention/diffex/figures/gen_validation/gen_real_distinct.py rename to src/ops_model/models/interpretability/diffex/figures/gen_validation/gen_real_distinct.py diff --git a/src/ops_model/models/attention/diffex/figures/gen_validation/ntc_inverse_gap.py b/src/ops_model/models/interpretability/diffex/figures/gen_validation/ntc_inverse_gap.py similarity index 95% rename from src/ops_model/models/attention/diffex/figures/gen_validation/ntc_inverse_gap.py rename to src/ops_model/models/interpretability/diffex/figures/gen_validation/ntc_inverse_gap.py index 6b0ecaf..7938830 100644 --- a/src/ops_model/models/attention/diffex/figures/gen_validation/ntc_inverse_gap.py +++ b/src/ops_model/models/interpretability/diffex/figures/gen_validation/ntc_inverse_gap.py @@ -41,8 +41,8 @@ def _first_done(channel=None): def run(channel=None, gene=None): import torch # noqa - from ops_model.models.attention.diffex.classifier.celldino_features import embed_crops - from ops_model.models.attention.diffex.directions.config import DirConfig + from ops_model.models.interpretability.diffex.classifier.celldino_features import embed_crops + from ops_model.models.interpretability.diffex.directions.config import DirConfig os.makedirs(OUT, exist_ok=True) ch, g = _first_done(channel) gene = gene or g @@ -122,9 +122,9 @@ def webp_ab(): (A) float straight into CellDINO, (B) round-tripped through the traversal's 8-bit _save_webp path — and compare CellDINO cosine. Answers: does saving as a proper (float/zarr) image bring generated closer to real?""" import torch, tempfile # noqa - from ops_model.models.attention.diffex.classifier.celldino_features import embed_crops - from ops_model.models.attention.diffex.directions.config import DirConfig - from ops_model.models.attention.diffex.viewer.precompute import _save_webp + from ops_model.models.interpretability.diffex.classifier.celldino_features import embed_crops + from ops_model.models.interpretability.diffex.directions.config import DirConfig + from ops_model.models.interpretability.diffex.viewer.precompute import _save_webp os.makedirs(OUT, exist_ok=True) d = np.load(CTRL, allow_pickle=True) imgs = d["anchor_imgs"].astype(np.float32) # (45,1,160,160) float [-1,1] diff --git a/src/ops_model/models/attention/diffex/figures/gen_validation/patch_cache_real.py b/src/ops_model/models/interpretability/diffex/figures/gen_validation/patch_cache_real.py similarity index 90% rename from src/ops_model/models/attention/diffex/figures/gen_validation/patch_cache_real.py rename to src/ops_model/models/interpretability/diffex/figures/gen_validation/patch_cache_real.py index ad1a0ed..b27ce0e 100644 --- a/src/ops_model/models/attention/diffex/figures/gen_validation/patch_cache_real.py +++ b/src/ops_model/models/interpretability/diffex/figures/gen_validation/patch_cache_real.py @@ -20,9 +20,9 @@ def _drop42(): def run(): import pandas as pd - from ops_model.models.attention.diffex.viewer.precompute import _gather_class - from ops_model.models.attention.diffex.directions.config import DirConfig - from ops_model.models.attention.diffex.classifier.config import slugify + from ops_model.models.interpretability.diffex.viewer.precompute import _gather_class + from ops_model.models.interpretability.diffex.directions.config import DirConfig + from ops_model.models.interpretability.diffex.classifier.config import slugify genes = _drop42() orig = {slugify(str(x)): str(x) for x in pd.read_parquet(RANKP, columns=["gene"])["gene"].unique()} # slug→ranking name cfg = DirConfig(grain="geneKO", target=genes[0], device="cuda"); cfg.num_workers = 12 diff --git a/src/ops_model/models/attention/diffex/figures/gen_validation/publish_multibag_page.py b/src/ops_model/models/interpretability/diffex/figures/gen_validation/publish_multibag_page.py similarity index 100% rename from src/ops_model/models/attention/diffex/figures/gen_validation/publish_multibag_page.py rename to src/ops_model/models/interpretability/diffex/figures/gen_validation/publish_multibag_page.py diff --git a/src/ops_model/models/attention/diffex/figures/gen_validation/rank_summary.py b/src/ops_model/models/interpretability/diffex/figures/gen_validation/rank_summary.py similarity index 100% rename from src/ops_model/models/attention/diffex/figures/gen_validation/rank_summary.py rename to src/ops_model/models/interpretability/diffex/figures/gen_validation/rank_summary.py diff --git a/src/ops_model/models/attention/diffex/figures/gen_validation/st_halves_score.py b/src/ops_model/models/interpretability/diffex/figures/gen_validation/st_halves_score.py similarity index 90% rename from src/ops_model/models/attention/diffex/figures/gen_validation/st_halves_score.py rename to src/ops_model/models/interpretability/diffex/figures/gen_validation/st_halves_score.py index 98dfa28..6d62ac9 100644 --- a/src/ops_model/models/attention/diffex/figures/gen_validation/st_halves_score.py +++ b/src/ops_model/models/interpretability/diffex/figures/gen_validation/st_halves_score.py @@ -12,9 +12,9 @@ def score_shard(genes): import torch - from ops_model.models.attention.diffex.viewer.score_generated import score_embs_v5 - from ops_model.models.attention.diffex.viewer.set_classifier import load_set_classifier, V5_CKPT_ROOT, V5_RUNS - from ops_model.models.attention.diffex.classifier.config import slugify + from ops_model.models.interpretability.diffex.viewer.score_generated import score_embs_v5 + from ops_model.models.interpretability.diffex.viewer.set_classifier import load_set_classifier, V5_CKPT_ROOT, V5_RUNS + from ops_model.models.interpretability.diffex.classifier.config import slugify os.makedirs(OUT, exist_ok=True) dev = "cuda" if torch.cuda.is_available() else "cpu" run = V5_RUNS[("phase", "geneKO" if GRAIN == "geneKO" else "complex_ebionly")] diff --git a/src/ops_model/models/attention/diffex/figures/gen_validation/std_anchor_test.py b/src/ops_model/models/interpretability/diffex/figures/gen_validation/std_anchor_test.py similarity index 100% rename from src/ops_model/models/attention/diffex/figures/gen_validation/std_anchor_test.py rename to src/ops_model/models/interpretability/diffex/figures/gen_validation/std_anchor_test.py diff --git a/src/ops_model/models/attention/diffex/figures/gen_validation/stepabl_compare.py b/src/ops_model/models/interpretability/diffex/figures/gen_validation/stepabl_compare.py similarity index 100% rename from src/ops_model/models/attention/diffex/figures/gen_validation/stepabl_compare.py rename to src/ops_model/models/interpretability/diffex/figures/gen_validation/stepabl_compare.py diff --git a/src/ops_model/models/attention/diffex/figures/gen_validation/valid200_alphastep.py b/src/ops_model/models/interpretability/diffex/figures/gen_validation/valid200_alphastep.py similarity index 100% rename from src/ops_model/models/attention/diffex/figures/gen_validation/valid200_alphastep.py rename to src/ops_model/models/interpretability/diffex/figures/gen_validation/valid200_alphastep.py diff --git a/src/ops_model/models/attention/diffex/figures/gen_validation/valid200_cache_build.py b/src/ops_model/models/interpretability/diffex/figures/gen_validation/valid200_cache_build.py similarity index 100% rename from src/ops_model/models/attention/diffex/figures/gen_validation/valid200_cache_build.py rename to src/ops_model/models/interpretability/diffex/figures/gen_validation/valid200_cache_build.py diff --git a/src/ops_model/models/attention/diffex/figures/gen_validation/valid200_capcheck.py b/src/ops_model/models/interpretability/diffex/figures/gen_validation/valid200_capcheck.py similarity index 100% rename from src/ops_model/models/attention/diffex/figures/gen_validation/valid200_capcheck.py rename to src/ops_model/models/interpretability/diffex/figures/gen_validation/valid200_capcheck.py diff --git a/src/ops_model/models/attention/diffex/figures/gen_validation/valid200_map_compare.py b/src/ops_model/models/interpretability/diffex/figures/gen_validation/valid200_map_compare.py similarity index 98% rename from src/ops_model/models/attention/diffex/figures/gen_validation/valid200_map_compare.py rename to src/ops_model/models/interpretability/diffex/figures/gen_validation/valid200_map_compare.py index aea1814..67ad8a8 100644 --- a/src/ops_model/models/attention/diffex/figures/gen_validation/valid200_map_compare.py +++ b/src/ops_model/models/interpretability/diffex/figures/gen_validation/valid200_map_compare.py @@ -9,7 +9,7 @@ """ import json, glob, os import numpy as np -from ops_model.models.attention.diffex.classifier.config import slugify +from ops_model.models.interpretability.diffex.classifier.config import slugify CV = "/hpc/projects/icd.fast.ops/analysis/figure4_traversals" OUT = f"{CV}/valid200_metrics" diff --git a/src/ops_model/models/attention/diffex/figures/gen_validation/valid200_metrics.py b/src/ops_model/models/interpretability/diffex/figures/gen_validation/valid200_metrics.py similarity index 100% rename from src/ops_model/models/attention/diffex/figures/gen_validation/valid200_metrics.py rename to src/ops_model/models/interpretability/diffex/figures/gen_validation/valid200_metrics.py diff --git a/src/ops_model/models/attention/diffex/figures/nc_ratio.py b/src/ops_model/models/interpretability/diffex/figures/nc_ratio.py similarity index 100% rename from src/ops_model/models/attention/diffex/figures/nc_ratio.py rename to src/ops_model/models/interpretability/diffex/figures/nc_ratio.py diff --git a/src/ops_model/models/attention/diffex/figures/ntc_anchor_compare.py b/src/ops_model/models/interpretability/diffex/figures/ntc_anchor_compare.py similarity index 100% rename from src/ops_model/models/attention/diffex/figures/ntc_anchor_compare.py rename to src/ops_model/models/interpretability/diffex/figures/ntc_anchor_compare.py diff --git a/src/ops_model/models/attention/diffex/figures/phase_montages.py b/src/ops_model/models/interpretability/diffex/figures/phase_montages.py similarity index 100% rename from src/ops_model/models/attention/diffex/figures/phase_montages.py rename to src/ops_model/models/interpretability/diffex/figures/phase_montages.py diff --git a/src/ops_model/models/attention/diffex/figures/phase_multibag_montages.py b/src/ops_model/models/interpretability/diffex/figures/phase_multibag_montages.py similarity index 100% rename from src/ops_model/models/attention/diffex/figures/phase_multibag_montages.py rename to src/ops_model/models/interpretability/diffex/figures/phase_multibag_montages.py diff --git a/src/ops_model/models/attention/diffex/figures/phase_sample_montages.py b/src/ops_model/models/interpretability/diffex/figures/phase_sample_montages.py similarity index 100% rename from src/ops_model/models/attention/diffex/figures/phase_sample_montages.py rename to src/ops_model/models/interpretability/diffex/figures/phase_sample_montages.py diff --git a/src/ops_model/models/attention/diffex/figures/rab_candidate_montages.py b/src/ops_model/models/interpretability/diffex/figures/rab_candidate_montages.py similarity index 92% rename from src/ops_model/models/attention/diffex/figures/rab_candidate_montages.py rename to src/ops_model/models/interpretability/diffex/figures/rab_candidate_montages.py index fc9bdb4..0286e20 100644 --- a/src/ops_model/models/attention/diffex/figures/rab_candidate_montages.py +++ b/src/ops_model/models/interpretability/diffex/figures/rab_candidate_montages.py @@ -2,7 +2,7 @@ cells can be picked. Complexes only have top-30 cells.""" from cis_golgi_alternatives import CANDS from debug_setacc_top100 import montage -from ops_model.models.attention.diffex.classifier.config import slugify +from ops_model.models.interpretability.diffex.classifier.config import slugify ntc_done = set() for c in CANDS: diff --git a/src/ops_model/models/attention/diffex/figures/rebuild_traversals_n100.py b/src/ops_model/models/interpretability/diffex/figures/rebuild_traversals_n100.py similarity index 95% rename from src/ops_model/models/attention/diffex/figures/rebuild_traversals_n100.py rename to src/ops_model/models/interpretability/diffex/figures/rebuild_traversals_n100.py index ab2677c..59d273e 100644 --- a/src/ops_model/models/attention/diffex/figures/rebuild_traversals_n100.py +++ b/src/ops_model/models/interpretability/diffex/figures/rebuild_traversals_n100.py @@ -13,8 +13,8 @@ import sys from pathlib import Path -from ops_model.models.attention.diffex.viewer import catalog as C -from ops_model.models.attention.diffex.classifier.config import slugify +from ops_model.models.interpretability.diffex.viewer import catalog as C +from ops_model.models.interpretability.diffex.classifier.config import slugify ASSETS = "viewer_assets_v5" RANK = f"{C.OUT}/{ASSETS}/_rankings/fluor" @@ -81,7 +81,7 @@ def _clear_anchor(modality): def rebuild_marker(modality): os.environ["OPS_DIFFEX_ASSETS"] = ASSETS - from ops_model.models.attention.diffex.viewer import precompute as P + from ops_model.models.interpretability.diffex.viewer import precompute as P P._ASSETS = ASSETS _clear_anchor(modality) tg = TARGETS[modality] @@ -105,7 +105,7 @@ def rebuild_marker(modality): def gen_phase(targets): """Generate specific phase geneKO targets at n=100 (reuses the existing 100-cell phase anchor cache).""" os.environ["OPS_DIFFEX_ASSETS"] = ASSETS - from ops_model.models.attention.diffex.viewer import precompute as P + from ops_model.models.interpretability.diffex.viewer import precompute as P P._ASSETS = ASSETS _clear_anchor("phase") # idempotent: keeps the 100-cell cache P.precompute_marker(grain="geneKO", targets=targets, ckpt=PHASE_CK, out_root=C.OUT, @@ -115,7 +115,7 @@ def gen_phase(targets): def gen_phase_complex(targets): """Generate phase COMPLEX targets at n=100 (full complex names; reuses the 100-cell phase anchor cache).""" os.environ["OPS_DIFFEX_ASSETS"] = ASSETS - from ops_model.models.attention.diffex.viewer import precompute as P + from ops_model.models.interpretability.diffex.viewer import precompute as P P._ASSETS = ASSETS _clear_anchor("phase") P.precompute_marker(grain="complex", targets=targets, ckpt=PHASE_CK, out_root=C.OUT, @@ -139,7 +139,7 @@ def gen_phase_chunk(targets, cell_range): cached direction (both ckpt-independent), each writing its own cell{c} dirs → parallelizes the per-cell DDIM inversion across GPUs instead of one long serial job. v5 scoring skipped (whole-target only).""" os.environ["OPS_DIFFEX_ASSETS"] = ASSETS - from ops_model.models.attention.diffex.viewer import precompute as P + from ops_model.models.interpretability.diffex.viewer import precompute as P P._ASSETS = ASSETS _clear_anchor("phase") # idempotent: 100-anchor cache kept P.precompute_marker(grain="geneKO", targets=targets, ckpt=PHASE_CK, out_root=C.OUT, diff --git a/src/ops_model/models/attention/diffex/figures/traversal_montage_schematic.py b/src/ops_model/models/interpretability/diffex/figures/traversal_montage_schematic.py similarity index 100% rename from src/ops_model/models/attention/diffex/figures/traversal_montage_schematic.py rename to src/ops_model/models/interpretability/diffex/figures/traversal_montage_schematic.py diff --git a/src/ops_model/models/attention/diffex/figures/virtual_staining_schematic.py b/src/ops_model/models/interpretability/diffex/figures/virtual_staining_schematic.py similarity index 100% rename from src/ops_model/models/attention/diffex/figures/virtual_staining_schematic.py rename to src/ops_model/models/interpretability/diffex/figures/virtual_staining_schematic.py diff --git a/src/ops_model/models/attention/diffex/kyle_pcs/build_static_explorer.py b/src/ops_model/models/interpretability/diffex/kyle_pcs/build_static_explorer.py similarity index 100% rename from src/ops_model/models/attention/diffex/kyle_pcs/build_static_explorer.py rename to src/ops_model/models/interpretability/diffex/kyle_pcs/build_static_explorer.py diff --git a/src/ops_model/models/attention/diffex/kyle_pcs/compute_pc_strips.py b/src/ops_model/models/interpretability/diffex/kyle_pcs/compute_pc_strips.py similarity index 100% rename from src/ops_model/models/attention/diffex/kyle_pcs/compute_pc_strips.py rename to src/ops_model/models/interpretability/diffex/kyle_pcs/compute_pc_strips.py diff --git a/src/ops_model/models/attention/diffex/viewer/__init__.py b/src/ops_model/models/interpretability/diffex/viewer/__init__.py similarity index 100% rename from src/ops_model/models/attention/diffex/viewer/__init__.py rename to src/ops_model/models/interpretability/diffex/viewer/__init__.py diff --git a/src/ops_model/models/attention/diffex/viewer/_altanchor_build.py b/src/ops_model/models/interpretability/diffex/viewer/_altanchor_build.py similarity index 100% rename from src/ops_model/models/attention/diffex/viewer/_altanchor_build.py rename to src/ops_model/models/interpretability/diffex/viewer/_altanchor_build.py diff --git a/src/ops_model/models/attention/diffex/viewer/_anchortest.py b/src/ops_model/models/interpretability/diffex/viewer/_anchortest.py similarity index 100% rename from src/ops_model/models/attention/diffex/viewer/_anchortest.py rename to src/ops_model/models/interpretability/diffex/viewer/_anchortest.py diff --git a/src/ops_model/models/attention/diffex/viewer/_build_stepablation.py b/src/ops_model/models/interpretability/diffex/viewer/_build_stepablation.py similarity index 100% rename from src/ops_model/models/attention/diffex/viewer/_build_stepablation.py rename to src/ops_model/models/interpretability/diffex/viewer/_build_stepablation.py diff --git a/src/ops_model/models/attention/diffex/viewer/_build_v5_inverted.py b/src/ops_model/models/interpretability/diffex/viewer/_build_v5_inverted.py similarity index 99% rename from src/ops_model/models/attention/diffex/viewer/_build_v5_inverted.py rename to src/ops_model/models/interpretability/diffex/viewer/_build_v5_inverted.py index adfb096..ba09d31 100644 --- a/src/ops_model/models/attention/diffex/viewer/_build_v5_inverted.py +++ b/src/ops_model/models/interpretability/diffex/viewer/_build_v5_inverted.py @@ -9,7 +9,7 @@ complex) whose per-cell rankings exist. That is Alex Lin's top1_acc>0.5@100-cell distinctiveness filter (see _fluor_v5_build.py) — NOT missing data (cells exist for all 1000×55). Lower-acc combos need Alex to gen more. - python -m ops_model.models.attention.diffex.viewer._build_v5_inverted markers + python -m ops_model.models.interpretability.diffex.viewer._build_v5_inverted markers """ import json import os diff --git a/src/ops_model/models/attention/diffex/viewer/_build_v5_montages.py b/src/ops_model/models/interpretability/diffex/viewer/_build_v5_montages.py similarity index 98% rename from src/ops_model/models/attention/diffex/viewer/_build_v5_montages.py rename to src/ops_model/models/interpretability/diffex/viewer/_build_v5_montages.py index 259e8c6..fb7311b 100644 --- a/src/ops_model/models/attention/diffex/viewer/_build_v5_montages.py +++ b/src/ops_model/models/interpretability/diffex/viewer/_build_v5_montages.py @@ -5,7 +5,7 @@ phase embedding. Reads the merged inverted frames from viewer_assets_v5//geneKO and writes tiles to viewer_assets_v5/_montage/. Only builds markers whose geneKO is 100% present in viewer_assets_v5. - python -m ops_model.models.attention.diffex.viewer._build_v5_montages + python -m ops_model.models.interpretability.diffex.viewer._build_v5_montages """ import glob import os diff --git a/src/ops_model/models/attention/diffex/viewer/_build_valid200.py b/src/ops_model/models/interpretability/diffex/viewer/_build_valid200.py similarity index 95% rename from src/ops_model/models/attention/diffex/viewer/_build_valid200.py rename to src/ops_model/models/interpretability/diffex/viewer/_build_valid200.py index 6a2a232..01858da 100644 --- a/src/ops_model/models/attention/diffex/viewer/_build_valid200.py +++ b/src/ops_model/models/interpretability/diffex/viewer/_build_valid200.py @@ -8,9 +8,9 @@ Reuses the v5 per-class directions (d_vec + gap are the PRODUCTION values we are validating) via a symlinked _directions tree, so no direction re-fit — only the 200-anchor inversion + 200×7 decodes per target. - python -m ops_model.models.attention.diffex.viewer._build_valid200 anchor # 1 GPU: build the 200-cell NTC anchor - python -m ops_model.models.attention.diffex.viewer._build_valid200 submit # shard all 1000 geneKO (after anchor) - python -m ops_model.models.attention.diffex.viewer._build_valid200 all # anchor job -> shards (afterok dep) + python -m ops_model.models.interpretability.diffex.viewer._build_valid200 anchor # 1 GPU: build the 200-cell NTC anchor + python -m ops_model.models.interpretability.diffex.viewer._build_valid200 submit # shard all 1000 geneKO (after anchor) + python -m ops_model.models.interpretability.diffex.viewer._build_valid200 all # anchor job -> shards (afterok dep) """ import os import sys diff --git a/src/ops_model/models/attention/diffex/viewer/_consolidate_cells.py b/src/ops_model/models/interpretability/diffex/viewer/_consolidate_cells.py similarity index 100% rename from src/ops_model/models/attention/diffex/viewer/_consolidate_cells.py rename to src/ops_model/models/interpretability/diffex/viewer/_consolidate_cells.py diff --git a/src/ops_model/models/attention/diffex/viewer/_fluor_complex_build.py b/src/ops_model/models/interpretability/diffex/viewer/_fluor_complex_build.py similarity index 100% rename from src/ops_model/models/attention/diffex/viewer/_fluor_complex_build.py rename to src/ops_model/models/interpretability/diffex/viewer/_fluor_complex_build.py diff --git a/src/ops_model/models/attention/diffex/viewer/_fluor_topcells.py b/src/ops_model/models/interpretability/diffex/viewer/_fluor_topcells.py similarity index 100% rename from src/ops_model/models/attention/diffex/viewer/_fluor_topcells.py rename to src/ops_model/models/interpretability/diffex/viewer/_fluor_topcells.py diff --git a/src/ops_model/models/attention/diffex/viewer/_fluor_v5_build.py b/src/ops_model/models/interpretability/diffex/viewer/_fluor_v5_build.py similarity index 100% rename from src/ops_model/models/attention/diffex/viewer/_fluor_v5_build.py rename to src/ops_model/models/interpretability/diffex/viewer/_fluor_v5_build.py diff --git a/src/ops_model/models/attention/diffex/viewer/_migrate_v4_to_v5.py b/src/ops_model/models/interpretability/diffex/viewer/_migrate_v4_to_v5.py similarity index 100% rename from src/ops_model/models/attention/diffex/viewer/_migrate_v4_to_v5.py rename to src/ops_model/models/interpretability/diffex/viewer/_migrate_v4_to_v5.py diff --git a/src/ops_model/models/attention/diffex/viewer/_phase_vs.py b/src/ops_model/models/interpretability/diffex/viewer/_phase_vs.py similarity index 100% rename from src/ops_model/models/attention/diffex/viewer/_phase_vs.py rename to src/ops_model/models/interpretability/diffex/viewer/_phase_vs.py diff --git a/src/ops_model/models/attention/diffex/viewer/_rebuild_v5.py b/src/ops_model/models/interpretability/diffex/viewer/_rebuild_v5.py similarity index 100% rename from src/ops_model/models/attention/diffex/viewer/_rebuild_v5.py rename to src/ops_model/models/interpretability/diffex/viewer/_rebuild_v5.py diff --git a/src/ops_model/models/attention/diffex/viewer/_rescore_rank.py b/src/ops_model/models/interpretability/diffex/viewer/_rescore_rank.py similarity index 97% rename from src/ops_model/models/attention/diffex/viewer/_rescore_rank.py rename to src/ops_model/models/interpretability/diffex/viewer/_rescore_rank.py index 0d4603b..79bbf8d 100644 --- a/src/ops_model/models/attention/diffex/viewer/_rescore_rank.py +++ b/src/ops_model/models/interpretability/diffex/viewer/_rescore_rank.py @@ -11,7 +11,7 @@ def _retarget(assets): """Point the score module at `assets` and return the module (V5_BASE is read at call time).""" - import ops_model.models.attention.diffex.viewer.score_generated as SG + import ops_model.models.interpretability.diffex.viewer.score_generated as SG SG.V5_BASE = f"{BASE}/{assets}/phase" return SG diff --git a/src/ops_model/models/attention/diffex/viewer/_score_v4.py b/src/ops_model/models/interpretability/diffex/viewer/_score_v4.py similarity index 100% rename from src/ops_model/models/attention/diffex/viewer/_score_v4.py rename to src/ops_model/models/interpretability/diffex/viewer/_score_v4.py diff --git a/src/ops_model/models/attention/diffex/viewer/_v4acc_test.py b/src/ops_model/models/interpretability/diffex/viewer/_v4acc_test.py similarity index 100% rename from src/ops_model/models/attention/diffex/viewer/_v4acc_test.py rename to src/ops_model/models/interpretability/diffex/viewer/_v4acc_test.py diff --git a/src/ops_model/models/attention/diffex/viewer/_verify_pt_space.py b/src/ops_model/models/interpretability/diffex/viewer/_verify_pt_space.py similarity index 100% rename from src/ops_model/models/attention/diffex/viewer/_verify_pt_space.py rename to src/ops_model/models/interpretability/diffex/viewer/_verify_pt_space.py diff --git a/src/ops_model/models/attention/diffex/viewer/_verify_score_bridge.py b/src/ops_model/models/interpretability/diffex/viewer/_verify_score_bridge.py similarity index 100% rename from src/ops_model/models/attention/diffex/viewer/_verify_score_bridge.py rename to src/ops_model/models/interpretability/diffex/viewer/_verify_score_bridge.py diff --git a/src/ops_model/models/attention/diffex/viewer/altanchor_pairs.json b/src/ops_model/models/interpretability/diffex/viewer/altanchor_pairs.json similarity index 100% rename from src/ops_model/models/attention/diffex/viewer/altanchor_pairs.json rename to src/ops_model/models/interpretability/diffex/viewer/altanchor_pairs.json diff --git a/src/ops_model/models/attention/diffex/viewer/anchor_cells.py b/src/ops_model/models/interpretability/diffex/viewer/anchor_cells.py similarity index 100% rename from src/ops_model/models/attention/diffex/viewer/anchor_cells.py rename to src/ops_model/models/interpretability/diffex/viewer/anchor_cells.py diff --git a/src/ops_model/models/attention/diffex/viewer/build_attention_heads.py b/src/ops_model/models/interpretability/diffex/viewer/build_attention_heads.py similarity index 95% rename from src/ops_model/models/attention/diffex/viewer/build_attention_heads.py rename to src/ops_model/models/interpretability/diffex/viewer/build_attention_heads.py index 9ee8308..6bfedf9 100644 --- a/src/ops_model/models/attention/diffex/viewer/build_attention_heads.py +++ b/src/ops_model/models/interpretability/diffex/viewer/build_attention_heads.py @@ -18,10 +18,10 @@ {AH}/index.json {global_max, assets:{:{:[keys]}}} where modality = "phase" | slugify(marker_channel), grain = geneKO|complex, key = gene | complex-slug. - python -m ops_model.models.attention.diffex.viewer.build_attention_heads render # SLURM (all trees) - python -m ops_model.models.attention.diffex.viewer.build_attention_heads render --local # serial, no SLURM - python -m ops_model.models.attention.diffex.viewer.build_attention_heads render --dry-run - python -m ops_model.models.attention.diffex.viewer.build_attention_heads index # (re)aggregate index.json + python -m ops_model.models.interpretability.diffex.viewer.build_attention_heads render # SLURM (all trees) + python -m ops_model.models.interpretability.diffex.viewer.build_attention_heads render --local # serial, no SLURM + python -m ops_model.models.interpretability.diffex.viewer.build_attention_heads render --dry-run + python -m ops_model.models.interpretability.diffex.viewer.build_attention_heads index # (re)aggregate index.json """ from __future__ import annotations diff --git a/src/ops_model/models/attention/diffex/viewer/build_complex_ebi_map.py b/src/ops_model/models/interpretability/diffex/viewer/build_complex_ebi_map.py similarity index 100% rename from src/ops_model/models/attention/diffex/viewer/build_complex_ebi_map.py rename to src/ops_model/models/interpretability/diffex/viewer/build_complex_ebi_map.py diff --git a/src/ops_model/models/attention/diffex/viewer/build_fluor_shap_rankings.py b/src/ops_model/models/interpretability/diffex/viewer/build_fluor_shap_rankings.py similarity index 96% rename from src/ops_model/models/attention/diffex/viewer/build_fluor_shap_rankings.py rename to src/ops_model/models/interpretability/diffex/viewer/build_fluor_shap_rankings.py index a8a8203..1e806cb 100644 --- a/src/ops_model/models/attention/diffex/viewer/build_fluor_shap_rankings.py +++ b/src/ops_model/models/interpretability/diffex/viewer/build_fluor_shap_rankings.py @@ -5,8 +5,8 @@ new CSV cols: gene, channel_name, rank, shap, ..., experiment, well, x_pheno, y_pheno, segmentation_id old schema: channel_name, gene, rank, pma_attention, experiment, well, x_pheno, y_pheno, segmentation, rank_type - python -m ops_model.models.attention.diffex.viewer.build_fluor_shap_rankings # local (needs ~64GB) - python -m ops_model.models.attention.diffex.viewer.build_fluor_shap_rankings --submit # SLURM cpu, mem 96 + python -m ops_model.models.interpretability.diffex.viewer.build_fluor_shap_rankings # local (needs ~64GB) + python -m ops_model.models.interpretability.diffex.viewer.build_fluor_shap_rankings --submit # SLURM cpu, mem 96 """ from __future__ import annotations diff --git a/src/ops_model/models/attention/diffex/viewer/build_montage_features.py b/src/ops_model/models/interpretability/diffex/viewer/build_montage_features.py similarity index 96% rename from src/ops_model/models/attention/diffex/viewer/build_montage_features.py rename to src/ops_model/models/interpretability/diffex/viewer/build_montage_features.py index 94c2d15..521d742 100644 --- a/src/ops_model/models/attention/diffex/viewer/build_montage_features.py +++ b/src/ops_model/models/interpretability/diffex/viewer/build_montage_features.py @@ -7,7 +7,7 @@ viewer_assets/montage_features.json {"features": [base names], "range": {feat: [lo, hi]}, "values": {gene: [0..1 per feature | null]}} - python -m ops_model.models.attention.diffex.viewer.build_montage_features + python -m ops_model.models.interpretability.diffex.viewer.build_montage_features """ from __future__ import annotations diff --git a/src/ops_model/models/attention/diffex/viewer/build_pc_crops_masked.py b/src/ops_model/models/interpretability/diffex/viewer/build_pc_crops_masked.py similarity index 97% rename from src/ops_model/models/attention/diffex/viewer/build_pc_crops_masked.py rename to src/ops_model/models/interpretability/diffex/viewer/build_pc_crops_masked.py index 4eac1dc..3869f87 100644 --- a/src/ops_model/models/attention/diffex/viewer/build_pc_crops_masked.py +++ b/src/ops_model/models/interpretability/diffex/viewer/build_pc_crops_masked.py @@ -10,8 +10,8 @@ {BASE}/{exp}/3-assembly/phenotyping_v3.zarr/{row}/{col}/0/0 image [1,C,1,Y,X], Phase2D=ch0 {BASE}/{exp}/3-assembly/phenotyping_v3.zarr/{row}/{col}/0/labels/cell_seg/0 int32 labels - python -m ops_model.models.attention.diffex.viewer.build_pc_crops_masked --sample 24 # preview - python -m ops_model.models.attention.diffex.viewer.build_pc_crops_masked # full (overwrites crops/) + python -m ops_model.models.interpretability.diffex.viewer.build_pc_crops_masked --sample 24 # preview + python -m ops_model.models.interpretability.diffex.viewer.build_pc_crops_masked # full (overwrites crops/) """ from __future__ import annotations diff --git a/src/ops_model/models/attention/diffex/viewer/build_pc_features.py b/src/ops_model/models/interpretability/diffex/viewer/build_pc_features.py similarity index 99% rename from src/ops_model/models/attention/diffex/viewer/build_pc_features.py rename to src/ops_model/models/interpretability/diffex/viewer/build_pc_features.py index d530975..5846745 100644 --- a/src/ops_model/models/attention/diffex/viewer/build_pc_features.py +++ b/src/ops_model/models/interpretability/diffex/viewer/build_pc_features.py @@ -14,7 +14,7 @@ to make "+corr features" line up with the strip's high bins. Compositions are unsigned (|r| / tf-idf) so they need no flip. - python -m ops_model.models.attention.diffex.viewer.build_pc_features + python -m ops_model.models.interpretability.diffex.viewer.build_pc_features """ from __future__ import annotations diff --git a/src/ops_model/models/attention/diffex/viewer/build_pc_walks.py b/src/ops_model/models/interpretability/diffex/viewer/build_pc_walks.py similarity index 97% rename from src/ops_model/models/attention/diffex/viewer/build_pc_walks.py rename to src/ops_model/models/interpretability/diffex/viewer/build_pc_walks.py index ed7a50c..f55b529 100644 --- a/src/ops_model/models/attention/diffex/viewer/build_pc_walks.py +++ b/src/ops_model/models/interpretability/diffex/viewer/build_pc_walks.py @@ -8,8 +8,8 @@ the PC score std; v_p is the (z-scored) eigenvector, mapped back to raw CellDINO space by the mean per-exp sd. Output: one composite figure per marker (rows = PCs, cols = α). - python -m ops_model.models.attention.diffex.viewer.build_pc_walks --markers "Mitochondria_TOMM20" - python -m ops_model.models.attention.diffex.viewer.build_pc_walks --all # SLURM, every marker + python -m ops_model.models.interpretability.diffex.viewer.build_pc_walks --markers "Mitochondria_TOMM20" + python -m ops_model.models.interpretability.diffex.viewer.build_pc_walks --all # SLURM, every marker """ from __future__ import annotations diff --git a/src/ops_model/models/attention/diffex/viewer/build_pcs.py b/src/ops_model/models/interpretability/diffex/viewer/build_pcs.py similarity index 97% rename from src/ops_model/models/attention/diffex/viewer/build_pcs.py rename to src/ops_model/models/interpretability/diffex/viewer/build_pcs.py index 63cc84a..7b46d2b 100644 --- a/src/ops_model/models/attention/diffex/viewer/build_pcs.py +++ b/src/ops_model/models/interpretability/diffex/viewer/build_pcs.py @@ -9,8 +9,8 @@ the raw artifacts dir is gone). If Kyle regenerates artifacts, re-run his build_static_explorer.py then point --html at the fresh output. - python -m ops_model.models.attention.diffex.viewer.build_pcs - python -m ops_model.models.attention.diffex.viewer.build_pcs --html /path/to/pc_explorer_static.html + python -m ops_model.models.interpretability.diffex.viewer.build_pcs + python -m ops_model.models.interpretability.diffex.viewer.build_pcs --html /path/to/pc_explorer_static.html """ from __future__ import annotations diff --git a/src/ops_model/models/attention/diffex/viewer/build_pcs_marker.py b/src/ops_model/models/interpretability/diffex/viewer/build_pcs_marker.py similarity index 99% rename from src/ops_model/models/attention/diffex/viewer/build_pcs_marker.py rename to src/ops_model/models/interpretability/diffex/viewer/build_pcs_marker.py index a99a427..20b7256 100644 --- a/src/ops_model/models/attention/diffex/viewer/build_pcs_marker.py +++ b/src/ops_model/models/interpretability/diffex/viewer/build_pcs_marker.py @@ -9,7 +9,7 @@ viewer_assets/pcs/markers//index.json (same schema as the phase pcs/index.json) viewer_assets/pcs/markers//crops/pc###_bin##_row#.png - python -m ops_model.models.attention.diffex.viewer.build_pcs_marker --marker "autophagosome_MAP1LC3B" + python -m ops_model.models.interpretability.diffex.viewer.build_pcs_marker --marker "autophagosome_MAP1LC3B" """ from __future__ import annotations diff --git a/src/ops_model/models/attention/diffex/viewer/build_phase_shap_rankings.py b/src/ops_model/models/interpretability/diffex/viewer/build_phase_shap_rankings.py similarity index 95% rename from src/ops_model/models/attention/diffex/viewer/build_phase_shap_rankings.py rename to src/ops_model/models/interpretability/diffex/viewer/build_phase_shap_rankings.py index f649d94..d55aaee 100644 --- a/src/ops_model/models/attention/diffex/viewer/build_phase_shap_rankings.py +++ b/src/ops_model/models/interpretability/diffex/viewer/build_phase_shap_rankings.py @@ -6,8 +6,8 @@ pma_attention, rank, rank_type) complex → pma_shap_phase_complex.parquet (adds predicted_class=complex, gene=member gene; EBI-pooled) - python -m ops_model.models.attention.diffex.viewer.build_phase_shap_rankings --geneko --submit - python -m ops_model.models.attention.diffex.viewer.build_phase_shap_rankings --complex --submit + python -m ops_model.models.interpretability.diffex.viewer.build_phase_shap_rankings --geneko --submit + python -m ops_model.models.interpretability.diffex.viewer.build_phase_shap_rankings --complex --submit """ from __future__ import annotations diff --git a/src/ops_model/models/attention/diffex/viewer/build_phate_figure.py b/src/ops_model/models/interpretability/diffex/viewer/build_phate_figure.py similarity index 99% rename from src/ops_model/models/attention/diffex/viewer/build_phate_figure.py rename to src/ops_model/models/interpretability/diffex/viewer/build_phate_figure.py index b3ac3ee..648ffd1 100644 --- a/src/ops_model/models/attention/diffex/viewer/build_phate_figure.py +++ b/src/ops_model/models/interpretability/diffex/viewer/build_phate_figure.py @@ -5,7 +5,7 @@ Each panel: the same PHATE scatter (grey), that panel's groups colored + leader-labelled with the single-cell generated morph (NTC cell1 → group, alpha=+5). NTC original shown top-left of panel E. - python -m ops_model.models.attention.diffex.viewer.build_phate_figure + python -m ops_model.models.interpretability.diffex.viewer.build_phate_figure """ from __future__ import annotations diff --git a/src/ops_model/models/attention/diffex/viewer/build_setacc_bins.py b/src/ops_model/models/interpretability/diffex/viewer/build_setacc_bins.py similarity index 100% rename from src/ops_model/models/attention/diffex/viewer/build_setacc_bins.py rename to src/ops_model/models/interpretability/diffex/viewer/build_setacc_bins.py diff --git a/src/ops_model/models/attention/diffex/viewer/build_setacc_bymarker.py b/src/ops_model/models/interpretability/diffex/viewer/build_setacc_bymarker.py similarity index 100% rename from src/ops_model/models/attention/diffex/viewer/build_setacc_bymarker.py rename to src/ops_model/models/interpretability/diffex/viewer/build_setacc_bymarker.py diff --git a/src/ops_model/models/attention/diffex/viewer/build_top_cells.py b/src/ops_model/models/interpretability/diffex/viewer/build_top_cells.py similarity index 97% rename from src/ops_model/models/attention/diffex/viewer/build_top_cells.py rename to src/ops_model/models/interpretability/diffex/viewer/build_top_cells.py index bed2423..843fefd 100644 --- a/src/ops_model/models/attention/diffex/viewer/build_top_cells.py +++ b/src/ops_model/models/interpretability/diffex/viewer/build_top_cells.py @@ -7,8 +7,8 @@ viewer_assets_v5/top_cells/index.json {"top_n", "genes"|"complexes": {CLASS: {"accuracy": [rec...]}}} viewer_assets_v5/top_cells/crops/.png - python -m ops_model.models.attention.diffex.viewer.build_top_cells geneKO # SLURM crop shards + finalize - python -m ops_model.models.attention.diffex.viewer.build_top_cells complex --finalize # rebuild index only + python -m ops_model.models.interpretability.diffex.viewer.build_top_cells geneKO # SLURM crop shards + finalize + python -m ops_model.models.interpretability.diffex.viewer.build_top_cells complex --finalize # rebuild index only """ from __future__ import annotations diff --git a/src/ops_model/models/attention/diffex/viewer/build_umap_montage.py b/src/ops_model/models/interpretability/diffex/viewer/build_umap_montage.py similarity index 100% rename from src/ops_model/models/attention/diffex/viewer/build_umap_montage.py rename to src/ops_model/models/interpretability/diffex/viewer/build_umap_montage.py diff --git a/src/ops_model/models/attention/diffex/viewer/catalog.py b/src/ops_model/models/interpretability/diffex/viewer/catalog.py similarity index 100% rename from src/ops_model/models/attention/diffex/viewer/catalog.py rename to src/ops_model/models/interpretability/diffex/viewer/catalog.py diff --git a/src/ops_model/models/attention/diffex/viewer/deploy/README.md b/src/ops_model/models/interpretability/diffex/viewer/deploy/README.md similarity index 100% rename from src/ops_model/models/attention/diffex/viewer/deploy/README.md rename to src/ops_model/models/interpretability/diffex/viewer/deploy/README.md diff --git a/src/ops_model/models/attention/diffex/viewer/marker_leaves.py b/src/ops_model/models/interpretability/diffex/viewer/marker_leaves.py similarity index 100% rename from src/ops_model/models/attention/diffex/viewer/marker_leaves.py rename to src/ops_model/models/interpretability/diffex/viewer/marker_leaves.py diff --git a/src/ops_model/models/attention/diffex/viewer/mimic_alex_embed.py b/src/ops_model/models/interpretability/diffex/viewer/mimic_alex_embed.py similarity index 100% rename from src/ops_model/models/attention/diffex/viewer/mimic_alex_embed.py rename to src/ops_model/models/interpretability/diffex/viewer/mimic_alex_embed.py diff --git a/src/ops_model/models/attention/diffex/viewer/morpho_pipeline.py b/src/ops_model/models/interpretability/diffex/viewer/morpho_pipeline.py similarity index 100% rename from src/ops_model/models/attention/diffex/viewer/morpho_pipeline.py rename to src/ops_model/models/interpretability/diffex/viewer/morpho_pipeline.py diff --git a/src/ops_model/models/attention/diffex/viewer/morphometrics.py b/src/ops_model/models/interpretability/diffex/viewer/morphometrics.py similarity index 100% rename from src/ops_model/models/attention/diffex/viewer/morphometrics.py rename to src/ops_model/models/interpretability/diffex/viewer/morphometrics.py diff --git a/src/ops_model/models/attention/diffex/viewer/nway_clf.py b/src/ops_model/models/interpretability/diffex/viewer/nway_clf.py similarity index 100% rename from src/ops_model/models/attention/diffex/viewer/nway_clf.py rename to src/ops_model/models/interpretability/diffex/viewer/nway_clf.py diff --git a/src/ops_model/models/attention/diffex/viewer/phenotype_cells.py b/src/ops_model/models/interpretability/diffex/viewer/phenotype_cells.py similarity index 100% rename from src/ops_model/models/attention/diffex/viewer/phenotype_cells.py rename to src/ops_model/models/interpretability/diffex/viewer/phenotype_cells.py diff --git a/src/ops_model/models/attention/diffex/viewer/precompute.py b/src/ops_model/models/interpretability/diffex/viewer/precompute.py similarity index 100% rename from src/ops_model/models/attention/diffex/viewer/precompute.py rename to src/ops_model/models/interpretability/diffex/viewer/precompute.py diff --git a/src/ops_model/models/attention/diffex/viewer/render_montage_scales.py b/src/ops_model/models/interpretability/diffex/viewer/render_montage_scales.py similarity index 99% rename from src/ops_model/models/attention/diffex/viewer/render_montage_scales.py rename to src/ops_model/models/interpretability/diffex/viewer/render_montage_scales.py index 7775dea..eb708ad 100644 --- a/src/ops_model/models/attention/diffex/viewer/render_montage_scales.py +++ b/src/ops_model/models/interpretability/diffex/viewer/render_montage_scales.py @@ -2,7 +2,7 @@ embedding legend (leiden_r4, big dots, NTC as a dark labelled circle). The montage image and its baked viewer-style gene names come straight from the built tiles — finer levels give crisper text. - python -m ops_model.models.attention.diffex.viewer.render_montage_scales --alphas 1-5 --levels 3,4 + python -m ops_model.models.interpretability.diffex.viewer.render_montage_scales --alphas 1-5 --levels 3,4 Each level of `_montage/phase_geneKO_phate_cell1_a_tiles/L/` is a level-of-detail montage (coarse levels show a decimated non-overlapping subset; finer levels fill in more cells at higher res). diff --git a/src/ops_model/models/attention/diffex/viewer/score_generated.py b/src/ops_model/models/interpretability/diffex/viewer/score_generated.py similarity index 100% rename from src/ops_model/models/attention/diffex/viewer/score_generated.py rename to src/ops_model/models/interpretability/diffex/viewer/score_generated.py diff --git a/src/ops_model/models/attention/diffex/viewer/set_classifier.py b/src/ops_model/models/interpretability/diffex/viewer/set_classifier.py similarity index 100% rename from src/ops_model/models/attention/diffex/viewer/set_classifier.py rename to src/ops_model/models/interpretability/diffex/viewer/set_classifier.py diff --git a/src/ops_model/models/attention/diffex/viewer/submit.py b/src/ops_model/models/interpretability/diffex/viewer/submit.py similarity index 97% rename from src/ops_model/models/attention/diffex/viewer/submit.py rename to src/ops_model/models/interpretability/diffex/viewer/submit.py index 2b99f18..699eac2 100644 --- a/src/ops_model/models/attention/diffex/viewer/submit.py +++ b/src/ops_model/models/interpretability/diffex/viewer/submit.py @@ -1,10 +1,10 @@ """Build the DiffEx viewer cache — reproducible, version-controlled entrypoint (replaces the one-off scratchpad drivers). All target selection comes from `catalog.py`. - python -m ops_model.models.attention.diffex.viewer.submit seed # per-marker NTC traversals - python -m ops_model.models.attention.diffex.viewer.submit anchors --k 5 # A→B anchor pairs - python -m ops_model.models.attention.diffex.viewer.submit manifest # rebuild manifest.json (local) - python -m ops_model.models.attention.diffex.viewer.submit montage --cell 0 --alpha 2 # harvest cache -> UMAP montage zarr + python -m ops_model.models.interpretability.diffex.viewer.submit seed # per-marker NTC traversals + python -m ops_model.models.interpretability.diffex.viewer.submit anchors --k 5 # A→B anchor pairs + python -m ops_model.models.interpretability.diffex.viewer.submit manifest # rebuild manifest.json (local) + python -m ops_model.models.interpretability.diffex.viewer.submit montage --cell 0 --alpha 2 # harvest cache -> UMAP montage zarr """ from __future__ import annotations diff --git a/src/ops_model/models/attention/diffex/viewer/webapp/app.js b/src/ops_model/models/interpretability/diffex/viewer/webapp/app.js similarity index 100% rename from src/ops_model/models/attention/diffex/viewer/webapp/app.js rename to src/ops_model/models/interpretability/diffex/viewer/webapp/app.js diff --git a/src/ops_model/models/attention/diffex/viewer/webapp/biohub-mark.png b/src/ops_model/models/interpretability/diffex/viewer/webapp/biohub-mark.png similarity index 100% rename from src/ops_model/models/attention/diffex/viewer/webapp/biohub-mark.png rename to src/ops_model/models/interpretability/diffex/viewer/webapp/biohub-mark.png diff --git a/src/ops_model/models/attention/diffex/viewer/webapp/biohub-wordmark.png b/src/ops_model/models/interpretability/diffex/viewer/webapp/biohub-wordmark.png similarity index 100% rename from src/ops_model/models/attention/diffex/viewer/webapp/biohub-wordmark.png rename to src/ops_model/models/interpretability/diffex/viewer/webapp/biohub-wordmark.png diff --git a/src/ops_model/models/attention/diffex/viewer/webapp/build_gene_narratives.py b/src/ops_model/models/interpretability/diffex/viewer/webapp/build_gene_narratives.py similarity index 100% rename from src/ops_model/models/attention/diffex/viewer/webapp/build_gene_narratives.py rename to src/ops_model/models/interpretability/diffex/viewer/webapp/build_gene_narratives.py diff --git a/src/ops_model/models/attention/diffex/viewer/webapp/gif.js b/src/ops_model/models/interpretability/diffex/viewer/webapp/gif.js similarity index 100% rename from src/ops_model/models/attention/diffex/viewer/webapp/gif.js rename to src/ops_model/models/interpretability/diffex/viewer/webapp/gif.js diff --git a/src/ops_model/models/attention/diffex/viewer/webapp/gif.worker.js b/src/ops_model/models/interpretability/diffex/viewer/webapp/gif.worker.js similarity index 100% rename from src/ops_model/models/attention/diffex/viewer/webapp/gif.worker.js rename to src/ops_model/models/interpretability/diffex/viewer/webapp/gif.worker.js diff --git a/src/ops_model/models/attention/diffex/viewer/webapp/index.html b/src/ops_model/models/interpretability/diffex/viewer/webapp/index.html similarity index 100% rename from src/ops_model/models/attention/diffex/viewer/webapp/index.html rename to src/ops_model/models/interpretability/diffex/viewer/webapp/index.html diff --git a/src/ops_model/models/attention/diffex/viewer/webapp/methods.js b/src/ops_model/models/interpretability/diffex/viewer/webapp/methods.js similarity index 100% rename from src/ops_model/models/attention/diffex/viewer/webapp/methods.js rename to src/ops_model/models/interpretability/diffex/viewer/webapp/methods.js diff --git a/src/ops_model/models/attention/diffex/viewer/webapp/morpho_demo.html b/src/ops_model/models/interpretability/diffex/viewer/webapp/morpho_demo.html similarity index 100% rename from src/ops_model/models/attention/diffex/viewer/webapp/morpho_demo.html rename to src/ops_model/models/interpretability/diffex/viewer/webapp/morpho_demo.html diff --git a/src/ops_model/models/attention/diffex/viewer/webapp/openseadragon.min.js b/src/ops_model/models/interpretability/diffex/viewer/webapp/openseadragon.min.js similarity index 100% rename from src/ops_model/models/attention/diffex/viewer/webapp/openseadragon.min.js rename to src/ops_model/models/interpretability/diffex/viewer/webapp/openseadragon.min.js diff --git a/src/ops_model/models/attention/diffex/viewer/webapp/opsin-eyes.svg b/src/ops_model/models/interpretability/diffex/viewer/webapp/opsin-eyes.svg similarity index 100% rename from src/ops_model/models/attention/diffex/viewer/webapp/opsin-eyes.svg rename to src/ops_model/models/interpretability/diffex/viewer/webapp/opsin-eyes.svg diff --git a/src/ops_model/models/attention/diffex/viewer/webapp/style.css b/src/ops_model/models/interpretability/diffex/viewer/webapp/style.css similarity index 100% rename from src/ops_model/models/attention/diffex/viewer/webapp/style.css rename to src/ops_model/models/interpretability/diffex/viewer/webapp/style.css diff --git a/src/ops_model/models/attention/embedding/generate_ko_violin_plots.py b/src/ops_model/models/interpretability/embedding/generate_ko_violin_plots.py similarity index 100% rename from src/ops_model/models/attention/embedding/generate_ko_violin_plots.py rename to src/ops_model/models/interpretability/embedding/generate_ko_violin_plots.py diff --git a/src/ops_model/models/attention/embedding/regen_umap_gav.py b/src/ops_model/models/interpretability/embedding/regen_umap_gav.py similarity index 100% rename from src/ops_model/models/attention/embedding/regen_umap_gav.py rename to src/ops_model/models/interpretability/embedding/regen_umap_gav.py diff --git a/src/ops_model/models/attention/embedding/regen_umap_html.py b/src/ops_model/models/interpretability/embedding/regen_umap_html.py similarity index 100% rename from src/ops_model/models/attention/embedding/regen_umap_html.py rename to src/ops_model/models/interpretability/embedding/regen_umap_html.py diff --git a/src/ops_model/models/attention/embedding/run_all_atlases.py b/src/ops_model/models/interpretability/embedding/run_all_atlases.py similarity index 100% rename from src/ops_model/models/attention/embedding/run_all_atlases.py rename to src/ops_model/models/interpretability/embedding/run_all_atlases.py diff --git a/src/ops_model/models/attention/embedding/top_attention_embed_and_score.py b/src/ops_model/models/interpretability/embedding/top_attention_embed_and_score.py similarity index 100% rename from src/ops_model/models/attention/embedding/top_attention_embed_and_score.py rename to src/ops_model/models/interpretability/embedding/top_attention_embed_and_score.py diff --git a/src/ops_model/models/attention/shap/analyze_chad_variants.py b/src/ops_model/models/interpretability/shap/analyze_chad_variants.py similarity index 100% rename from src/ops_model/models/attention/shap/analyze_chad_variants.py rename to src/ops_model/models/interpretability/shap/analyze_chad_variants.py diff --git a/src/ops_model/models/attention/shap/generate_shap_captions_combined.py b/src/ops_model/models/interpretability/shap/generate_shap_captions_combined.py similarity index 100% rename from src/ops_model/models/attention/shap/generate_shap_captions_combined.py rename to src/ops_model/models/interpretability/shap/generate_shap_captions_combined.py diff --git a/src/ops_model/models/attention/shap/ko_shap_features.py b/src/ops_model/models/interpretability/shap/ko_shap_features.py similarity index 100% rename from src/ops_model/models/attention/shap/ko_shap_features.py rename to src/ops_model/models/interpretability/shap/ko_shap_features.py diff --git a/src/ops_model/models/attention/shap/merge_shap_shards.py b/src/ops_model/models/interpretability/shap/merge_shap_shards.py similarity index 100% rename from src/ops_model/models/attention/shap/merge_shap_shards.py rename to src/ops_model/models/interpretability/shap/merge_shap_shards.py diff --git a/src/ops_model/models/attention/shap/ntc_attention_compare.py b/src/ops_model/models/interpretability/shap/ntc_attention_compare.py similarity index 100% rename from src/ops_model/models/attention/shap/ntc_attention_compare.py rename to src/ops_model/models/interpretability/shap/ntc_attention_compare.py diff --git a/src/ops_model/models/attention/shap/ntc_pick_cells.py b/src/ops_model/models/interpretability/shap/ntc_pick_cells.py similarity index 100% rename from src/ops_model/models/attention/shap/ntc_pick_cells.py rename to src/ops_model/models/interpretability/shap/ntc_pick_cells.py diff --git a/src/ops_model/models/attention/shap/ntc_shap_features.py b/src/ops_model/models/interpretability/shap/ntc_shap_features.py similarity index 100% rename from src/ops_model/models/attention/shap/ntc_shap_features.py rename to src/ops_model/models/interpretability/shap/ntc_shap_features.py diff --git a/src/ops_model/models/attention/shap/run_all_shap.py b/src/ops_model/models/interpretability/shap/run_all_shap.py similarity index 100% rename from src/ops_model/models/attention/shap/run_all_shap.py rename to src/ops_model/models/interpretability/shap/run_all_shap.py diff --git a/src/ops_model/models/attention/shap/run_shap_pipeline.py b/src/ops_model/models/interpretability/shap/run_shap_pipeline.py similarity index 100% rename from src/ops_model/models/attention/shap/run_shap_pipeline.py rename to src/ops_model/models/interpretability/shap/run_shap_pipeline.py diff --git a/src/ops_model/models/attention/shap/shap_approach_compare.py b/src/ops_model/models/interpretability/shap/shap_approach_compare.py similarity index 100% rename from src/ops_model/models/attention/shap/shap_approach_compare.py rename to src/ops_model/models/interpretability/shap/shap_approach_compare.py diff --git a/src/ops_model/models/attention/titration/decay/map_attention_decay.py b/src/ops_model/models/interpretability/titration/decay/map_attention_decay.py similarity index 100% rename from src/ops_model/models/attention/titration/decay/map_attention_decay.py rename to src/ops_model/models/interpretability/titration/decay/map_attention_decay.py diff --git a/src/ops_model/models/attention/titration/decay/phate_peak_groups.py b/src/ops_model/models/interpretability/titration/decay/phate_peak_groups.py similarity index 100% rename from src/ops_model/models/attention/titration/decay/phate_peak_groups.py rename to src/ops_model/models/interpretability/titration/decay/phate_peak_groups.py diff --git a/src/ops_model/models/attention/titration/decay/plot_3way_summary_bars.py b/src/ops_model/models/interpretability/titration/decay/plot_3way_summary_bars.py similarity index 100% rename from src/ops_model/models/attention/titration/decay/plot_3way_summary_bars.py rename to src/ops_model/models/interpretability/titration/decay/plot_3way_summary_bars.py diff --git a/src/ops_model/models/attention/titration/decay/plot_all_cells_correction_bars.py b/src/ops_model/models/interpretability/titration/decay/plot_all_cells_correction_bars.py similarity index 100% rename from src/ops_model/models/attention/titration/decay/plot_all_cells_correction_bars.py rename to src/ops_model/models/interpretability/titration/decay/plot_all_cells_correction_bars.py diff --git a/src/ops_model/models/attention/titration/expansion/count_genes_above_threshold.py b/src/ops_model/models/interpretability/titration/expansion/count_genes_above_threshold.py similarity index 98% rename from src/ops_model/models/attention/titration/expansion/count_genes_above_threshold.py rename to src/ops_model/models/interpretability/titration/expansion/count_genes_above_threshold.py index 4faf30a..f402319 100644 --- a/src/ops_model/models/attention/titration/expansion/count_genes_above_threshold.py +++ b/src/ops_model/models/interpretability/titration/expansion/count_genes_above_threshold.py @@ -17,10 +17,10 @@ Usage:: # Submit one SLURM task per K (9 tasks; ~5 min wall once they land) - uv run python -m ops_model.models.attention.titration.expansion.count_genes_above_threshold --slurm + uv run python -m ops_model.models.interpretability.titration.expansion.count_genes_above_threshold --slurm # Replot from cached per-gene CSVs (no SLURM) - uv run python -m ops_model.models.attention.titration.expansion.count_genes_above_threshold --replot + uv run python -m ops_model.models.interpretability.titration.expansion.count_genes_above_threshold --replot """ from __future__ import annotations diff --git a/src/ops_model/models/attention/titration/expansion/map_attention_expansion_v4.py b/src/ops_model/models/interpretability/titration/expansion/map_attention_expansion_v4.py similarity index 100% rename from src/ops_model/models/attention/titration/expansion/map_attention_expansion_v4.py rename to src/ops_model/models/interpretability/titration/expansion/map_attention_expansion_v4.py diff --git a/src/ops_model/models/attention/titration/expansion/plot_sgrna_coverage_sweep.py b/src/ops_model/models/interpretability/titration/expansion/plot_sgrna_coverage_sweep.py similarity index 100% rename from src/ops_model/models/attention/titration/expansion/plot_sgrna_coverage_sweep.py rename to src/ops_model/models/interpretability/titration/expansion/plot_sgrna_coverage_sweep.py diff --git a/src/ops_model/models/attention/titration/expansion/run_percentile_sweep.py b/src/ops_model/models/interpretability/titration/expansion/run_percentile_sweep.py similarity index 100% rename from src/ops_model/models/attention/titration/expansion/run_percentile_sweep.py rename to src/ops_model/models/interpretability/titration/expansion/run_percentile_sweep.py diff --git a/src/ops_model/models/attention/weighted_aggregation/_v4_attn_worker.py b/src/ops_model/models/interpretability/weighted_aggregation/_v4_attn_worker.py similarity index 100% rename from src/ops_model/models/attention/weighted_aggregation/_v4_attn_worker.py rename to src/ops_model/models/interpretability/weighted_aggregation/_v4_attn_worker.py diff --git a/src/ops_model/models/attention/weighted_aggregation/analyze_v3_acc_bins.py b/src/ops_model/models/interpretability/weighted_aggregation/analyze_v3_acc_bins.py similarity index 100% rename from src/ops_model/models/attention/weighted_aggregation/analyze_v3_acc_bins.py rename to src/ops_model/models/interpretability/weighted_aggregation/analyze_v3_acc_bins.py diff --git a/src/ops_model/models/attention/weighted_aggregation/plot_v4_attn_comparison.py b/src/ops_model/models/interpretability/weighted_aggregation/plot_v4_attn_comparison.py similarity index 100% rename from src/ops_model/models/attention/weighted_aggregation/plot_v4_attn_comparison.py rename to src/ops_model/models/interpretability/weighted_aggregation/plot_v4_attn_comparison.py diff --git a/src/ops_model/models/attention/weighted_aggregation/run_v3_pipeline_on_v4_attn_weighted.py b/src/ops_model/models/interpretability/weighted_aggregation/run_v3_pipeline_on_v4_attn_weighted.py similarity index 100% rename from src/ops_model/models/attention/weighted_aggregation/run_v3_pipeline_on_v4_attn_weighted.py rename to src/ops_model/models/interpretability/weighted_aggregation/run_v3_pipeline_on_v4_attn_weighted.py diff --git a/src/ops_model/models/attention/weighted_aggregation/run_v3_pipeline_on_v4_features.py b/src/ops_model/models/interpretability/weighted_aggregation/run_v3_pipeline_on_v4_features.py similarity index 100% rename from src/ops_model/models/attention/weighted_aggregation/run_v3_pipeline_on_v4_features.py rename to src/ops_model/models/interpretability/weighted_aggregation/run_v3_pipeline_on_v4_features.py From 4f875f9e0c69f10af7b3eee9a508386e2960208d Mon Sep 17 00:00:00 2001 From: Gav Sturm Date: Tue, 11 Aug 2026 09:28:37 -0700 Subject: [PATCH 03/13] rename diffex -> diffae, inner diffae stage -> generator - interpretability/diffex/ -> interpretability/diffae/ (the DiffEx->DiffAE pipeline) - inner diffae/ (stage-2 diffusion generator) -> generator/ (resolves the name clash) - patched module paths in .py (5 inner-refs, 46 parent-refs, 17 relative ..diffae->..generator) and doc run-commands; preserved the DiffEx method name, DiffAE class, /models/diffex/ output paths, and diffex-viewer infra names. Core pipeline import-verified (directions -> ..generator wiring intact). --- .../{diffex => diffae}/CROPSEQ_TO_MORPHOLOGY.md | 0 .../interpretability/{diffex => diffae}/PLAN.md | 4 ++-- .../interpretability/{diffex => diffae}/README.md | 12 ++++++------ .../{diffex => diffae}/classifier/README.md | 0 .../{diffex => diffae}/classifier/__init__.py | 0 .../{diffex => diffae}/classifier/aggregate.py | 2 +- .../classifier/celldino_features.py | 0 .../{diffex => diffae}/classifier/config.py | 0 .../{diffex => diffae}/classifier/data.py | 0 .../{diffex => diffae}/classifier/models.py | 0 .../{diffex => diffae}/classifier/run.py | 4 ++-- .../{diffex => diffae}/classifier/submit.py | 6 +++--- .../{diffex => diffae}/classifier/train.py | 0 .../{diffex => diffae}/directions/__init__.py | 0 .../{diffex => diffae}/directions/batch.py | 2 +- .../{diffex => diffae}/directions/config.py | 0 .../{diffex => diffae}/directions/data.py | 0 .../{diffex => diffae}/directions/flow.py | 0 .../{diffex => diffae}/directions/grid.py | 0 .../{diffex => diffae}/directions/losses.py | 0 .../{diffex => diffae}/directions/make_gifs.py | 0 .../{diffex => diffae}/directions/model.py | 0 .../directions/proto_ddim_anchors.py | 4 ++-- .../{diffex => diffae}/directions/rank.py | 0 .../{diffex => diffae}/directions/run.py | 6 +++--- .../{diffex => diffae}/directions/submit.py | 2 +- .../directions/train_directions.py | 0 .../{diffex => diffae}/directions/traverse.py | 4 ++-- .../{diffex => diffae}/figures/METHODS_final.txt | 0 .../figures/METHODS_traversal_montage.md | 0 .../figures/METHODS_traversal_montage.txt | 0 .../{diffex => diffae}/figures/_setacc_common.py | 10 +++++----- .../{diffex => diffae}/figures/_setacc_phase.py | 0 .../figures/auto_pick_and_plot.py | 4 ++-- .../figures/cis_golgi_alternatives.py | 0 .../figures/debug_setacc_top100.py | 2 +- .../figures/ebi_peripheral_droplets.py | 2 +- .../figures/figure4_morpho_traversal.py | 0 .../figures/figure4_morpho_violin.py | 4 ++-- .../figures/figure4_setacc_panel.py | 0 .../figures/figure4_setacc_panel_fluorB.py | 0 .../figures/figure4_setacc_panel_newpheno.py | 0 .../figures/figure4_setacc_panel_phase.py | 0 .../figures/figure_ebi_morpho_violin.py | 0 .../figures/figure_multirank_ebi_grid.py | 0 .../figures/fluor_panel_montages.py | 0 .../figures/fluor_shap_montages.py | 2 +- .../figures/gen_validation/bag_sweep_plots.py | 2 +- .../figures/gen_validation/bag_sweep_score.py | 6 +++--- .../figures/gen_validation/centroid_bagsweep.py | 4 ++-- .../figures/gen_validation/centroid_halves.py | 4 ++-- .../gen_validation/centroid_pooled_bagsweep.py | 4 ++-- .../figures/gen_validation/control_halves_zscore.py | 4 ++-- .../figures/gen_validation/embcheck.py | 6 +++--- .../figures/gen_validation/embedding_diagnostics.py | 0 .../gen_validation/figure4_bagsize_reachfrac.py | 0 .../gen_validation/figure4_v5_accuracy_summary.py | 4 ++-- .../figures/gen_validation/gen_alpha_embedding.py | 0 .../figures/gen_validation/gen_embed_refit.py | 4 ++-- .../figures/gen_validation/gen_phate_passthrough.py | 0 .../figures/gen_validation/gen_real_centroid.py | 12 ++++++------ .../figures/gen_validation/gen_real_distinct.py | 0 .../figures/gen_validation/ntc_inverse_gap.py | 10 +++++----- .../figures/gen_validation/patch_cache_real.py | 6 +++--- .../figures/gen_validation/publish_multibag_page.py | 0 .../figures/gen_validation/rank_summary.py | 0 .../figures/gen_validation/st_halves_score.py | 6 +++--- .../figures/gen_validation/std_anchor_test.py | 0 .../figures/gen_validation/stepabl_compare.py | 0 .../figures/gen_validation/valid200_alphastep.py | 0 .../figures/gen_validation/valid200_cache_build.py | 0 .../figures/gen_validation/valid200_capcheck.py | 0 .../figures/gen_validation/valid200_map_compare.py | 2 +- .../figures/gen_validation/valid200_metrics.py | 0 .../{diffex => diffae}/figures/nc_ratio.py | 0 .../figures/ntc_anchor_compare.py | 0 .../{diffex => diffae}/figures/phase_montages.py | 0 .../figures/phase_multibag_montages.py | 0 .../figures/phase_sample_montages.py | 0 .../figures/rab_candidate_montages.py | 2 +- .../figures/rebuild_traversals_n100.py | 12 ++++++------ .../figures/traversal_montage_schematic.py | 0 .../figures/virtual_staining_schematic.py | 0 .../{diffex/diffae => diffae/generator}/__init__.py | 0 .../{diffex/diffae => diffae/generator}/config.py | 0 .../{diffex/diffae => diffae/generator}/data.py | 0 .../generator}/diagnose_conditioning.py | 2 +- .../{diffex/diffae => diffae/generator}/model.py | 0 .../diffae => diffae/generator}/plot_metrics.py | 0 .../{diffex/diffae => diffae/generator}/recon.py | 0 .../{diffex/diffae => diffae/generator}/run.py | 2 +- .../{diffex/diffae => diffae/generator}/submit.py | 2 +- .../{diffex/diffae => diffae/generator}/train.py | 0 .../diffae => diffae/generator}/virtstain_eval.py | 2 +- .../diffae => diffae/generator}/virtstain_multi.py | 2 +- .../kyle_pcs/build_static_explorer.py | 0 .../kyle_pcs/compute_pc_strips.py | 0 .../{diffex => diffae}/viewer/__init__.py | 0 .../{diffex => diffae}/viewer/_altanchor_build.py | 0 .../{diffex => diffae}/viewer/_anchortest.py | 0 .../viewer/_build_stepablation.py | 0 .../{diffex => diffae}/viewer/_build_v5_inverted.py | 10 +++++----- .../{diffex => diffae}/viewer/_build_v5_montages.py | 2 +- .../{diffex => diffae}/viewer/_build_valid200.py | 8 ++++---- .../{diffex => diffae}/viewer/_consolidate_cells.py | 0 .../viewer/_fluor_complex_build.py | 0 .../{diffex => diffae}/viewer/_fluor_topcells.py | 2 +- .../{diffex => diffae}/viewer/_fluor_v5_build.py | 0 .../{diffex => diffae}/viewer/_migrate_v4_to_v5.py | 0 .../{diffex => diffae}/viewer/_phase_vs.py | 6 +++--- .../{diffex => diffae}/viewer/_rebuild_v5.py | 2 +- .../{diffex => diffae}/viewer/_rescore_rank.py | 2 +- .../{diffex => diffae}/viewer/_score_v4.py | 0 .../{diffex => diffae}/viewer/_v4acc_test.py | 0 .../{diffex => diffae}/viewer/_verify_pt_space.py | 0 .../viewer/_verify_score_bridge.py | 0 .../{diffex => diffae}/viewer/altanchor_pairs.json | 0 .../{diffex => diffae}/viewer/anchor_cells.py | 0 .../viewer/build_attention_heads.py | 8 ++++---- .../viewer/build_complex_ebi_map.py | 0 .../viewer/build_fluor_shap_rankings.py | 4 ++-- .../viewer/build_montage_features.py | 2 +- .../viewer/build_pc_crops_masked.py | 4 ++-- .../{diffex => diffae}/viewer/build_pc_features.py | 2 +- .../{diffex => diffae}/viewer/build_pc_walks.py | 4 ++-- .../{diffex => diffae}/viewer/build_pcs.py | 4 ++-- .../{diffex => diffae}/viewer/build_pcs_marker.py | 2 +- .../viewer/build_phase_shap_rankings.py | 4 ++-- .../{diffex => diffae}/viewer/build_phate_figure.py | 2 +- .../{diffex => diffae}/viewer/build_setacc_bins.py | 0 .../viewer/build_setacc_bymarker.py | 0 .../{diffex => diffae}/viewer/build_top_cells.py | 4 ++-- .../{diffex => diffae}/viewer/build_umap_montage.py | 0 .../{diffex => diffae}/viewer/catalog.py | 0 .../{diffex => diffae}/viewer/deploy/README.md | 0 .../{diffex => diffae}/viewer/marker_leaves.py | 0 .../{diffex => diffae}/viewer/mimic_alex_embed.py | 0 .../{diffex => diffae}/viewer/morpho_pipeline.py | 4 ++-- .../{diffex => diffae}/viewer/morphometrics.py | 0 .../{diffex => diffae}/viewer/nway_clf.py | 0 .../{diffex => diffae}/viewer/phenotype_cells.py | 0 .../{diffex => diffae}/viewer/precompute.py | 2 +- .../viewer/render_montage_scales.py | 2 +- .../{diffex => diffae}/viewer/score_generated.py | 0 .../{diffex => diffae}/viewer/set_classifier.py | 0 .../{diffex => diffae}/viewer/submit.py | 8 ++++---- .../{diffex => diffae}/viewer/webapp/app.js | 0 .../viewer/webapp/biohub-mark.png | Bin .../viewer/webapp/biohub-wordmark.png | Bin .../viewer/webapp/build_gene_narratives.py | 0 .../{diffex => diffae}/viewer/webapp/gif.js | 0 .../{diffex => diffae}/viewer/webapp/gif.worker.js | 0 .../{diffex => diffae}/viewer/webapp/index.html | 0 .../{diffex => diffae}/viewer/webapp/methods.js | 0 .../viewer/webapp/morpho_demo.html | 0 .../viewer/webapp/openseadragon.min.js | 0 .../{diffex => diffae}/viewer/webapp/opsin-eyes.svg | 0 .../{diffex => diffae}/viewer/webapp/style.css | 0 158 files changed, 128 insertions(+), 128 deletions(-) rename src/ops_model/models/interpretability/{diffex => diffae}/CROPSEQ_TO_MORPHOLOGY.md (100%) rename src/ops_model/models/interpretability/{diffex => diffae}/PLAN.md (99%) rename src/ops_model/models/interpretability/{diffex => diffae}/README.md (72%) rename src/ops_model/models/interpretability/{diffex => diffae}/classifier/README.md (100%) rename src/ops_model/models/interpretability/{diffex => diffae}/classifier/__init__.py (100%) rename src/ops_model/models/interpretability/{diffex => diffae}/classifier/aggregate.py (98%) rename src/ops_model/models/interpretability/{diffex => diffae}/classifier/celldino_features.py (100%) rename src/ops_model/models/interpretability/{diffex => diffae}/classifier/config.py (100%) rename src/ops_model/models/interpretability/{diffex => diffae}/classifier/data.py (100%) rename src/ops_model/models/interpretability/{diffex => diffae}/classifier/models.py (100%) rename src/ops_model/models/interpretability/{diffex => diffae}/classifier/run.py (96%) rename src/ops_model/models/interpretability/{diffex => diffae}/classifier/submit.py (95%) rename src/ops_model/models/interpretability/{diffex => diffae}/classifier/train.py (100%) rename src/ops_model/models/interpretability/{diffex => diffae}/directions/__init__.py (100%) rename src/ops_model/models/interpretability/{diffex => diffae}/directions/batch.py (99%) rename src/ops_model/models/interpretability/{diffex => diffae}/directions/config.py (100%) rename src/ops_model/models/interpretability/{diffex => diffae}/directions/data.py (100%) rename src/ops_model/models/interpretability/{diffex => diffae}/directions/flow.py (100%) rename src/ops_model/models/interpretability/{diffex => diffae}/directions/grid.py (100%) rename src/ops_model/models/interpretability/{diffex => diffae}/directions/losses.py (100%) rename src/ops_model/models/interpretability/{diffex => diffae}/directions/make_gifs.py (100%) rename src/ops_model/models/interpretability/{diffex => diffae}/directions/model.py (100%) rename src/ops_model/models/interpretability/{diffex => diffae}/directions/proto_ddim_anchors.py (98%) rename src/ops_model/models/interpretability/{diffex => diffae}/directions/rank.py (100%) rename src/ops_model/models/interpretability/{diffex => diffae}/directions/run.py (96%) rename src/ops_model/models/interpretability/{diffex => diffae}/directions/submit.py (97%) rename src/ops_model/models/interpretability/{diffex => diffae}/directions/train_directions.py (100%) rename src/ops_model/models/interpretability/{diffex => diffae}/directions/traverse.py (99%) rename src/ops_model/models/interpretability/{diffex => diffae}/figures/METHODS_final.txt (100%) rename src/ops_model/models/interpretability/{diffex => diffae}/figures/METHODS_traversal_montage.md (100%) rename src/ops_model/models/interpretability/{diffex => diffae}/figures/METHODS_traversal_montage.txt (100%) rename src/ops_model/models/interpretability/{diffex => diffae}/figures/_setacc_common.py (96%) rename src/ops_model/models/interpretability/{diffex => diffae}/figures/_setacc_phase.py (100%) rename src/ops_model/models/interpretability/{diffex => diffae}/figures/auto_pick_and_plot.py (97%) rename src/ops_model/models/interpretability/{diffex => diffae}/figures/cis_golgi_alternatives.py (100%) rename src/ops_model/models/interpretability/{diffex => diffae}/figures/debug_setacc_top100.py (98%) rename src/ops_model/models/interpretability/{diffex => diffae}/figures/ebi_peripheral_droplets.py (99%) rename src/ops_model/models/interpretability/{diffex => diffae}/figures/figure4_morpho_traversal.py (100%) rename src/ops_model/models/interpretability/{diffex => diffae}/figures/figure4_morpho_violin.py (98%) rename src/ops_model/models/interpretability/{diffex => diffae}/figures/figure4_setacc_panel.py (100%) rename src/ops_model/models/interpretability/{diffex => diffae}/figures/figure4_setacc_panel_fluorB.py (100%) rename src/ops_model/models/interpretability/{diffex => diffae}/figures/figure4_setacc_panel_newpheno.py (100%) rename src/ops_model/models/interpretability/{diffex => diffae}/figures/figure4_setacc_panel_phase.py (100%) rename src/ops_model/models/interpretability/{diffex => diffae}/figures/figure_ebi_morpho_violin.py (100%) rename src/ops_model/models/interpretability/{diffex => diffae}/figures/figure_multirank_ebi_grid.py (100%) rename src/ops_model/models/interpretability/{diffex => diffae}/figures/fluor_panel_montages.py (100%) rename src/ops_model/models/interpretability/{diffex => diffae}/figures/fluor_shap_montages.py (99%) rename src/ops_model/models/interpretability/{diffex => diffae}/figures/gen_validation/bag_sweep_plots.py (99%) rename src/ops_model/models/interpretability/{diffex => diffae}/figures/gen_validation/bag_sweep_score.py (93%) rename src/ops_model/models/interpretability/{diffex => diffae}/figures/gen_validation/centroid_bagsweep.py (97%) rename src/ops_model/models/interpretability/{diffex => diffae}/figures/gen_validation/centroid_halves.py (96%) rename src/ops_model/models/interpretability/{diffex => diffae}/figures/gen_validation/centroid_pooled_bagsweep.py (97%) rename src/ops_model/models/interpretability/{diffex => diffae}/figures/gen_validation/control_halves_zscore.py (97%) rename src/ops_model/models/interpretability/{diffex => diffae}/figures/gen_validation/embcheck.py (90%) rename src/ops_model/models/interpretability/{diffex => diffae}/figures/gen_validation/embedding_diagnostics.py (100%) rename src/ops_model/models/interpretability/{diffex => diffae}/figures/gen_validation/figure4_bagsize_reachfrac.py (100%) rename src/ops_model/models/interpretability/{diffex => diffae}/figures/gen_validation/figure4_v5_accuracy_summary.py (98%) rename src/ops_model/models/interpretability/{diffex => diffae}/figures/gen_validation/gen_alpha_embedding.py (100%) rename src/ops_model/models/interpretability/{diffex => diffae}/figures/gen_validation/gen_embed_refit.py (99%) rename src/ops_model/models/interpretability/{diffex => diffae}/figures/gen_validation/gen_phate_passthrough.py (100%) rename src/ops_model/models/interpretability/{diffex => diffae}/figures/gen_validation/gen_real_centroid.py (98%) rename src/ops_model/models/interpretability/{diffex => diffae}/figures/gen_validation/gen_real_distinct.py (100%) rename src/ops_model/models/interpretability/{diffex => diffae}/figures/gen_validation/ntc_inverse_gap.py (96%) rename src/ops_model/models/interpretability/{diffex => diffae}/figures/gen_validation/patch_cache_real.py (93%) rename src/ops_model/models/interpretability/{diffex => diffae}/figures/gen_validation/publish_multibag_page.py (100%) rename src/ops_model/models/interpretability/{diffex => diffae}/figures/gen_validation/rank_summary.py (100%) rename src/ops_model/models/interpretability/{diffex => diffae}/figures/gen_validation/st_halves_score.py (93%) rename src/ops_model/models/interpretability/{diffex => diffae}/figures/gen_validation/std_anchor_test.py (100%) rename src/ops_model/models/interpretability/{diffex => diffae}/figures/gen_validation/stepabl_compare.py (100%) rename src/ops_model/models/interpretability/{diffex => diffae}/figures/gen_validation/valid200_alphastep.py (100%) rename src/ops_model/models/interpretability/{diffex => diffae}/figures/gen_validation/valid200_cache_build.py (100%) rename src/ops_model/models/interpretability/{diffex => diffae}/figures/gen_validation/valid200_capcheck.py (100%) rename src/ops_model/models/interpretability/{diffex => diffae}/figures/gen_validation/valid200_map_compare.py (98%) rename src/ops_model/models/interpretability/{diffex => diffae}/figures/gen_validation/valid200_metrics.py (100%) rename src/ops_model/models/interpretability/{diffex => diffae}/figures/nc_ratio.py (100%) rename src/ops_model/models/interpretability/{diffex => diffae}/figures/ntc_anchor_compare.py (100%) rename src/ops_model/models/interpretability/{diffex => diffae}/figures/phase_montages.py (100%) rename src/ops_model/models/interpretability/{diffex => diffae}/figures/phase_multibag_montages.py (100%) rename src/ops_model/models/interpretability/{diffex => diffae}/figures/phase_sample_montages.py (100%) rename src/ops_model/models/interpretability/{diffex => diffae}/figures/rab_candidate_montages.py (93%) rename src/ops_model/models/interpretability/{diffex => diffae}/figures/rebuild_traversals_n100.py (96%) rename src/ops_model/models/interpretability/{diffex => diffae}/figures/traversal_montage_schematic.py (100%) rename src/ops_model/models/interpretability/{diffex => diffae}/figures/virtual_staining_schematic.py (100%) rename src/ops_model/models/interpretability/{diffex/diffae => diffae/generator}/__init__.py (100%) rename src/ops_model/models/interpretability/{diffex/diffae => diffae/generator}/config.py (100%) rename src/ops_model/models/interpretability/{diffex/diffae => diffae/generator}/data.py (100%) rename src/ops_model/models/interpretability/{diffex/diffae => diffae/generator}/diagnose_conditioning.py (98%) rename src/ops_model/models/interpretability/{diffex/diffae => diffae/generator}/model.py (100%) rename src/ops_model/models/interpretability/{diffex/diffae => diffae/generator}/plot_metrics.py (100%) rename src/ops_model/models/interpretability/{diffex/diffae => diffae/generator}/recon.py (100%) rename src/ops_model/models/interpretability/{diffex/diffae => diffae/generator}/run.py (97%) rename src/ops_model/models/interpretability/{diffex/diffae => diffae/generator}/submit.py (98%) rename src/ops_model/models/interpretability/{diffex/diffae => diffae/generator}/train.py (100%) rename src/ops_model/models/interpretability/{diffex/diffae => diffae/generator}/virtstain_eval.py (98%) rename src/ops_model/models/interpretability/{diffex/diffae => diffae/generator}/virtstain_multi.py (99%) rename src/ops_model/models/interpretability/{diffex => diffae}/kyle_pcs/build_static_explorer.py (100%) rename src/ops_model/models/interpretability/{diffex => diffae}/kyle_pcs/compute_pc_strips.py (100%) rename src/ops_model/models/interpretability/{diffex => diffae}/viewer/__init__.py (100%) rename src/ops_model/models/interpretability/{diffex => diffae}/viewer/_altanchor_build.py (100%) rename src/ops_model/models/interpretability/{diffex => diffae}/viewer/_anchortest.py (100%) rename src/ops_model/models/interpretability/{diffex => diffae}/viewer/_build_stepablation.py (100%) rename src/ops_model/models/interpretability/{diffex => diffae}/viewer/_build_v5_inverted.py (99%) rename src/ops_model/models/interpretability/{diffex => diffae}/viewer/_build_v5_montages.py (98%) rename src/ops_model/models/interpretability/{diffex => diffae}/viewer/_build_valid200.py (97%) rename src/ops_model/models/interpretability/{diffex => diffae}/viewer/_consolidate_cells.py (100%) rename src/ops_model/models/interpretability/{diffex => diffae}/viewer/_fluor_complex_build.py (100%) rename src/ops_model/models/interpretability/{diffex => diffae}/viewer/_fluor_topcells.py (99%) rename src/ops_model/models/interpretability/{diffex => diffae}/viewer/_fluor_v5_build.py (100%) rename src/ops_model/models/interpretability/{diffex => diffae}/viewer/_migrate_v4_to_v5.py (100%) rename src/ops_model/models/interpretability/{diffex => diffae}/viewer/_phase_vs.py (99%) rename src/ops_model/models/interpretability/{diffex => diffae}/viewer/_rebuild_v5.py (98%) rename src/ops_model/models/interpretability/{diffex => diffae}/viewer/_rescore_rank.py (97%) rename src/ops_model/models/interpretability/{diffex => diffae}/viewer/_score_v4.py (100%) rename src/ops_model/models/interpretability/{diffex => diffae}/viewer/_v4acc_test.py (100%) rename src/ops_model/models/interpretability/{diffex => diffae}/viewer/_verify_pt_space.py (100%) rename src/ops_model/models/interpretability/{diffex => diffae}/viewer/_verify_score_bridge.py (100%) rename src/ops_model/models/interpretability/{diffex => diffae}/viewer/altanchor_pairs.json (100%) rename src/ops_model/models/interpretability/{diffex => diffae}/viewer/anchor_cells.py (100%) rename src/ops_model/models/interpretability/{diffex => diffae}/viewer/build_attention_heads.py (97%) rename src/ops_model/models/interpretability/{diffex => diffae}/viewer/build_complex_ebi_map.py (100%) rename src/ops_model/models/interpretability/{diffex => diffae}/viewer/build_fluor_shap_rankings.py (97%) rename src/ops_model/models/interpretability/{diffex => diffae}/viewer/build_montage_features.py (97%) rename src/ops_model/models/interpretability/{diffex => diffae}/viewer/build_pc_crops_masked.py (98%) rename src/ops_model/models/interpretability/{diffex => diffae}/viewer/build_pc_features.py (99%) rename src/ops_model/models/interpretability/{diffex => diffae}/viewer/build_pc_walks.py (98%) rename src/ops_model/models/interpretability/{diffex => diffae}/viewer/build_pcs.py (98%) rename src/ops_model/models/interpretability/{diffex => diffae}/viewer/build_pcs_marker.py (99%) rename src/ops_model/models/interpretability/{diffex => diffae}/viewer/build_phase_shap_rankings.py (97%) rename src/ops_model/models/interpretability/{diffex => diffae}/viewer/build_phate_figure.py (99%) rename src/ops_model/models/interpretability/{diffex => diffae}/viewer/build_setacc_bins.py (100%) rename src/ops_model/models/interpretability/{diffex => diffae}/viewer/build_setacc_bymarker.py (100%) rename src/ops_model/models/interpretability/{diffex => diffae}/viewer/build_top_cells.py (98%) rename src/ops_model/models/interpretability/{diffex => diffae}/viewer/build_umap_montage.py (100%) rename src/ops_model/models/interpretability/{diffex => diffae}/viewer/catalog.py (100%) rename src/ops_model/models/interpretability/{diffex => diffae}/viewer/deploy/README.md (100%) rename src/ops_model/models/interpretability/{diffex => diffae}/viewer/marker_leaves.py (100%) rename src/ops_model/models/interpretability/{diffex => diffae}/viewer/mimic_alex_embed.py (100%) rename src/ops_model/models/interpretability/{diffex => diffae}/viewer/morpho_pipeline.py (99%) rename src/ops_model/models/interpretability/{diffex => diffae}/viewer/morphometrics.py (100%) rename src/ops_model/models/interpretability/{diffex => diffae}/viewer/nway_clf.py (100%) rename src/ops_model/models/interpretability/{diffex => diffae}/viewer/phenotype_cells.py (100%) rename src/ops_model/models/interpretability/{diffex => diffae}/viewer/precompute.py (99%) rename src/ops_model/models/interpretability/{diffex => diffae}/viewer/render_montage_scales.py (99%) rename src/ops_model/models/interpretability/{diffex => diffae}/viewer/score_generated.py (100%) rename src/ops_model/models/interpretability/{diffex => diffae}/viewer/set_classifier.py (100%) rename src/ops_model/models/interpretability/{diffex => diffae}/viewer/submit.py (98%) rename src/ops_model/models/interpretability/{diffex => diffae}/viewer/webapp/app.js (100%) rename src/ops_model/models/interpretability/{diffex => diffae}/viewer/webapp/biohub-mark.png (100%) rename src/ops_model/models/interpretability/{diffex => diffae}/viewer/webapp/biohub-wordmark.png (100%) rename src/ops_model/models/interpretability/{diffex => diffae}/viewer/webapp/build_gene_narratives.py (100%) rename src/ops_model/models/interpretability/{diffex => diffae}/viewer/webapp/gif.js (100%) rename src/ops_model/models/interpretability/{diffex => diffae}/viewer/webapp/gif.worker.js (100%) rename src/ops_model/models/interpretability/{diffex => diffae}/viewer/webapp/index.html (100%) rename src/ops_model/models/interpretability/{diffex => diffae}/viewer/webapp/methods.js (100%) rename src/ops_model/models/interpretability/{diffex => diffae}/viewer/webapp/morpho_demo.html (100%) rename src/ops_model/models/interpretability/{diffex => diffae}/viewer/webapp/openseadragon.min.js (100%) rename src/ops_model/models/interpretability/{diffex => diffae}/viewer/webapp/opsin-eyes.svg (100%) rename src/ops_model/models/interpretability/{diffex => diffae}/viewer/webapp/style.css (100%) diff --git a/src/ops_model/models/interpretability/diffex/CROPSEQ_TO_MORPHOLOGY.md b/src/ops_model/models/interpretability/diffae/CROPSEQ_TO_MORPHOLOGY.md similarity index 100% rename from src/ops_model/models/interpretability/diffex/CROPSEQ_TO_MORPHOLOGY.md rename to src/ops_model/models/interpretability/diffae/CROPSEQ_TO_MORPHOLOGY.md diff --git a/src/ops_model/models/interpretability/diffex/PLAN.md b/src/ops_model/models/interpretability/diffae/PLAN.md similarity index 99% rename from src/ops_model/models/interpretability/diffex/PLAN.md rename to src/ops_model/models/interpretability/diffae/PLAN.md index 26e7b49..e5cafd5 100644 --- a/src/ops_model/models/interpretability/diffex/PLAN.md +++ b/src/ops_model/models/interpretability/diffae/PLAN.md @@ -333,7 +333,7 @@ and the CellDINO encoder are all local. ## Build log ### 2026-06-16 — classifier B/C package built -Package: `ops_model/models/attention/diffex/classifier/` (config, data, models, +Package: `ops_model/models/interpretability/diffae/classifier/` (config, data, models, celldino_features, train, run, submit, README). Locked params: binary HSPA5-vs-rest, negatives = other genes' top-5 (distinct), 1000/class, 160×160 phase crops (no mask), 3-way train/val/test split grouped by experiment (val=selection, test=clean reported AUROC; @@ -344,7 +344,7 @@ train+val AUROC logged per epoch to watch over/under-fit). Outputs under - **Verified:** full B pipeline end-to-end on CPU (tiny config) — filtered parquet read (no OOM), store resolution, crops materialized non-degenerate, train→AUROC→artifacts. SLURM submitter dry-run OK (2 GPU jobs). (Crop cache key includes mask state so masked/unmasked don't collide.) -- **Next (GPU):** `python -m ops_model.models.attention.diffex.classifier.submit --gene HSPA5` +- **Next (GPU):** `python -m ops_model.models.interpretability.diffae.classifier.submit --gene HSPA5` → compare B vs C held-out AUROC, pick the DiffEx target. C needs GPU (CellDINO ViT-L). ### 2026-06-16 — HSPA5 PoC results (job 34280479, experiment-grouped split, 1000/class) diff --git a/src/ops_model/models/interpretability/diffex/README.md b/src/ops_model/models/interpretability/diffae/README.md similarity index 72% rename from src/ops_model/models/interpretability/diffex/README.md rename to src/ops_model/models/interpretability/diffae/README.md index 1f72e6c..19307cc 100644 --- a/src/ops_model/models/interpretability/diffex/README.md +++ b/src/ops_model/models/interpretability/diffae/README.md @@ -13,22 +13,22 @@ See [PLAN.md](PLAN.md) for the design rationale and the full running log. | stage | package | what it does | |---|---|---| | 1 | [`classifier/`](classifier/) | per-class single-cell classifier on **top-attention cells** — the model whose decision DiffEx explains / that ranks directions. B = ResNet on phase crops; **C = MLP on CellDINO features** (chosen). | -| 2 | [`diffae/`](diffae/) | **conditional diffusion** generator: UNet that generates a cell image conditioned on its CellDINO embedding (conditioning dropout + EMA + CFG). | +| 2 | [`generator/`](generator/) | **conditional diffusion** generator (the DiffAE): UNet that generates a cell image conditioned on its CellDINO embedding (conditioning dropout + EMA + CFG). | | 3 | [`directions/`](directions/) | **contrastive direction discovery** (InfoNCE + decorrelation, unsupervised) → rank directions by a control-vs-target classifier → **CFG traversal** α∈[−,+] → DDIM-sample a counterfactual strip + Δ-pixel heatmap, verified by re-encoded score. | ## Run order (each stage has `run.py` for local + `submit.py` for SLURM) ```bash # Stage 1 — classifier (per gene/complex, or sweep --all-classes) -python -m ops_model.models.attention.diffex.classifier.submit --grain complex --all-classes --models C -python -m ops_model.models.attention.diffex.classifier.aggregate --grain complex --model C +python -m ops_model.models.interpretability.diffae.classifier.submit --grain complex --all-classes --models C +python -m ops_model.models.interpretability.diffae.classifier.aggregate --grain complex --model C # Stage 2 — train the conditional DiffAE (resume-able; gate = embedding/noise ratio) -python -m ops_model.models.attention.diffex.diffae.submit --epochs 120 --batch-size 48 -python -m ops_model.models.attention.diffex.diffae.diagnose_conditioning # conditioning-strength check +python -m ops_model.models.interpretability.diffae.generator.submit --epochs 120 --batch-size 48 +python -m ops_model.models.interpretability.diffae.generator.diagnose_conditioning # conditioning-strength check # Stage 3 — directions + counterfactual traversal for a target -python -m ops_model.models.attention.diffex.directions.submit --grain geneKO --target HSPA5 +python -m ops_model.models.interpretability.diffae.directions.submit --grain geneKO --target HSPA5 ``` Outputs: `/hpc/projects/icd.fast.ops/models/diffex/{,diffae,directions}/…`. diff --git a/src/ops_model/models/interpretability/diffex/classifier/README.md b/src/ops_model/models/interpretability/diffae/classifier/README.md similarity index 100% rename from src/ops_model/models/interpretability/diffex/classifier/README.md rename to src/ops_model/models/interpretability/diffae/classifier/README.md diff --git a/src/ops_model/models/interpretability/diffex/classifier/__init__.py b/src/ops_model/models/interpretability/diffae/classifier/__init__.py similarity index 100% rename from src/ops_model/models/interpretability/diffex/classifier/__init__.py rename to src/ops_model/models/interpretability/diffae/classifier/__init__.py diff --git a/src/ops_model/models/interpretability/diffex/classifier/aggregate.py b/src/ops_model/models/interpretability/diffae/classifier/aggregate.py similarity index 98% rename from src/ops_model/models/interpretability/diffex/classifier/aggregate.py rename to src/ops_model/models/interpretability/diffae/classifier/aggregate.py index 2577bc0..c5ee963 100644 --- a/src/ops_model/models/interpretability/diffex/classifier/aggregate.py +++ b/src/ops_model/models/interpretability/diffae/classifier/aggregate.py @@ -1,6 +1,6 @@ """Aggregate per-class classifier metrics into a ranked table. - python -m ops_model.models.interpretability.diffex.classifier.aggregate --grain complex + python -m ops_model.models.interpretability.diffae.classifier.aggregate --grain complex Collects every ///metrics_.json into one CSV ranked by test AUROC (how cleanly/distinctly each class's top-attention cells classify), plus diff --git a/src/ops_model/models/interpretability/diffex/classifier/celldino_features.py b/src/ops_model/models/interpretability/diffae/classifier/celldino_features.py similarity index 100% rename from src/ops_model/models/interpretability/diffex/classifier/celldino_features.py rename to src/ops_model/models/interpretability/diffae/classifier/celldino_features.py diff --git a/src/ops_model/models/interpretability/diffex/classifier/config.py b/src/ops_model/models/interpretability/diffae/classifier/config.py similarity index 100% rename from src/ops_model/models/interpretability/diffex/classifier/config.py rename to src/ops_model/models/interpretability/diffae/classifier/config.py diff --git a/src/ops_model/models/interpretability/diffex/classifier/data.py b/src/ops_model/models/interpretability/diffae/classifier/data.py similarity index 100% rename from src/ops_model/models/interpretability/diffex/classifier/data.py rename to src/ops_model/models/interpretability/diffae/classifier/data.py diff --git a/src/ops_model/models/interpretability/diffex/classifier/models.py b/src/ops_model/models/interpretability/diffae/classifier/models.py similarity index 100% rename from src/ops_model/models/interpretability/diffex/classifier/models.py rename to src/ops_model/models/interpretability/diffae/classifier/models.py diff --git a/src/ops_model/models/interpretability/diffex/classifier/run.py b/src/ops_model/models/interpretability/diffae/classifier/run.py similarity index 96% rename from src/ops_model/models/interpretability/diffex/classifier/run.py rename to src/ops_model/models/interpretability/diffae/classifier/run.py index afe55ef..370b5ac 100644 --- a/src/ops_model/models/interpretability/diffex/classifier/run.py +++ b/src/ops_model/models/interpretability/diffae/classifier/run.py @@ -1,7 +1,7 @@ """Orchestrator for the DiffEx single-cell classifier PoC. - python -m ops_model.models.interpretability.diffex.classifier.run --model B --gene HSPA5 - python -m ops_model.models.interpretability.diffex.classifier.run --model C --gene HSPA5 + python -m ops_model.models.interpretability.diffae.classifier.run --model B --gene HSPA5 + python -m ops_model.models.interpretability.diffae.classifier.run --model C --gene HSPA5 Shared steps: build cell table -> materialize phase crops (cached) -> split. Then B trains a ResNet on crops; C embeds the crops with CellDINO and trains an diff --git a/src/ops_model/models/interpretability/diffex/classifier/submit.py b/src/ops_model/models/interpretability/diffae/classifier/submit.py similarity index 95% rename from src/ops_model/models/interpretability/diffex/classifier/submit.py rename to src/ops_model/models/interpretability/diffae/classifier/submit.py index ccf6c72..a785bf0 100644 --- a/src/ops_model/models/interpretability/diffex/classifier/submit.py +++ b/src/ops_model/models/interpretability/diffae/classifier/submit.py @@ -1,13 +1,13 @@ """Submit the classifier sweep to SLURM (GPU) via submit_parallel_jobs. # one gene, both models - python -m ops_model.models.interpretability.diffex.classifier.submit --gene HSPA5 --models B C + python -m ops_model.models.interpretability.diffae.classifier.submit --gene HSPA5 --models B C # all 98 EBI complexes, model C - python -m ops_model.models.interpretability.diffex.classifier.submit --grain complex --all-classes --models C + python -m ops_model.models.interpretability.diffae.classifier.submit --grain complex --all-classes --models C # specific classes - python -m ops_model.models.interpretability.diffex.classifier.submit --grain complex \ + python -m ops_model.models.interpretability.diffae.classifier.submit --grain complex \ --classes "19S proteasome regulatory complex" "Commander complex" --models C One GPU job per (class, model). Outputs under ///. diff --git a/src/ops_model/models/interpretability/diffex/classifier/train.py b/src/ops_model/models/interpretability/diffae/classifier/train.py similarity index 100% rename from src/ops_model/models/interpretability/diffex/classifier/train.py rename to src/ops_model/models/interpretability/diffae/classifier/train.py diff --git a/src/ops_model/models/interpretability/diffex/directions/__init__.py b/src/ops_model/models/interpretability/diffae/directions/__init__.py similarity index 100% rename from src/ops_model/models/interpretability/diffex/directions/__init__.py rename to src/ops_model/models/interpretability/diffae/directions/__init__.py diff --git a/src/ops_model/models/interpretability/diffex/directions/batch.py b/src/ops_model/models/interpretability/diffae/directions/batch.py similarity index 99% rename from src/ops_model/models/interpretability/diffex/directions/batch.py rename to src/ops_model/models/interpretability/diffae/directions/batch.py index eea65fc..69581aa 100644 --- a/src/ops_model/models/interpretability/diffex/directions/batch.py +++ b/src/ops_model/models/interpretability/diffae/directions/batch.py @@ -4,7 +4,7 @@ submits one GPU job per target. Each job: run_directions at w=5 only → per-cell strips + scores, then a GIF for the auto-picked best-Δscore cell. - python -m ops_model.models.interpretability.diffex.directions.batch \ + python -m ops_model.models.interpretability.diffae.directions.batch \ --genes-csv --complex-csv \ --n-genes 50 --n-complex 20 """ diff --git a/src/ops_model/models/interpretability/diffex/directions/config.py b/src/ops_model/models/interpretability/diffae/directions/config.py similarity index 100% rename from src/ops_model/models/interpretability/diffex/directions/config.py rename to src/ops_model/models/interpretability/diffae/directions/config.py diff --git a/src/ops_model/models/interpretability/diffex/directions/data.py b/src/ops_model/models/interpretability/diffae/directions/data.py similarity index 100% rename from src/ops_model/models/interpretability/diffex/directions/data.py rename to src/ops_model/models/interpretability/diffae/directions/data.py diff --git a/src/ops_model/models/interpretability/diffex/directions/flow.py b/src/ops_model/models/interpretability/diffae/directions/flow.py similarity index 100% rename from src/ops_model/models/interpretability/diffex/directions/flow.py rename to src/ops_model/models/interpretability/diffae/directions/flow.py diff --git a/src/ops_model/models/interpretability/diffex/directions/grid.py b/src/ops_model/models/interpretability/diffae/directions/grid.py similarity index 100% rename from src/ops_model/models/interpretability/diffex/directions/grid.py rename to src/ops_model/models/interpretability/diffae/directions/grid.py diff --git a/src/ops_model/models/interpretability/diffex/directions/losses.py b/src/ops_model/models/interpretability/diffae/directions/losses.py similarity index 100% rename from src/ops_model/models/interpretability/diffex/directions/losses.py rename to src/ops_model/models/interpretability/diffae/directions/losses.py diff --git a/src/ops_model/models/interpretability/diffex/directions/make_gifs.py b/src/ops_model/models/interpretability/diffae/directions/make_gifs.py similarity index 100% rename from src/ops_model/models/interpretability/diffex/directions/make_gifs.py rename to src/ops_model/models/interpretability/diffae/directions/make_gifs.py diff --git a/src/ops_model/models/interpretability/diffex/directions/model.py b/src/ops_model/models/interpretability/diffae/directions/model.py similarity index 100% rename from src/ops_model/models/interpretability/diffex/directions/model.py rename to src/ops_model/models/interpretability/diffae/directions/model.py diff --git a/src/ops_model/models/interpretability/diffex/directions/proto_ddim_anchors.py b/src/ops_model/models/interpretability/diffae/directions/proto_ddim_anchors.py similarity index 98% rename from src/ops_model/models/interpretability/diffex/directions/proto_ddim_anchors.py rename to src/ops_model/models/interpretability/diffae/directions/proto_ddim_anchors.py index fef1cc8..581f939 100644 --- a/src/ops_model/models/interpretability/diffex/directions/proto_ddim_anchors.py +++ b/src/ops_model/models/interpretability/diffae/directions/proto_ddim_anchors.py @@ -9,7 +9,7 @@ Same NTC anchor cells v5 uses (top-rank NTC), KIF23 + POLR1B, α 0->5. Everything reused from the existing traversal stack; the only new step is _ddim(..., inverse=True) to get xT. - python -m ops_model.models.interpretability.diffex.directions.proto_ddim_anchors --submit + python -m ops_model.models.interpretability.diffae.directions.proto_ddim_anchors --submit """ from __future__ import annotations @@ -21,7 +21,7 @@ import torch from ..classifier.config import slugify -from ..diffae.data import normalize +from ..generator.data import normalize from .config import DirConfig from .rank import supervised_direction from .traverse import _ddim_guided, _sample_guided, load_diffae diff --git a/src/ops_model/models/interpretability/diffex/directions/rank.py b/src/ops_model/models/interpretability/diffae/directions/rank.py similarity index 100% rename from src/ops_model/models/interpretability/diffex/directions/rank.py rename to src/ops_model/models/interpretability/diffae/directions/rank.py diff --git a/src/ops_model/models/interpretability/diffex/directions/run.py b/src/ops_model/models/interpretability/diffae/directions/run.py similarity index 96% rename from src/ops_model/models/interpretability/diffex/directions/run.py rename to src/ops_model/models/interpretability/diffae/directions/run.py index 5971d42..92df0d7 100644 --- a/src/ops_model/models/interpretability/diffex/directions/run.py +++ b/src/ops_model/models/interpretability/diffae/directions/run.py @@ -1,7 +1,7 @@ """Orchestrator for Stage 3 (directions → ranking → traversal). - python -m ops_model.models.interpretability.diffex.directions.run --target HSPA5 - python -m ops_model.models.interpretability.diffex.directions.run --grain complex \ + python -m ops_model.models.interpretability.diffae.directions.run --target HSPA5 + python -m ops_model.models.interpretability.diffae.directions.run --grain complex \ --target "Chaperonin-containing T-complex" Steps: gather target+control crops/embeddings → train K direction MLPs (unsupervised) @@ -18,7 +18,7 @@ import torch from ..classifier.config import DEFAULT_OUT_ROOT, GRAINS, slugify -from ..diffae.data import normalize +from ..generator.data import normalize from .config import DirConfig from .data import gather from .model import DirectionBank diff --git a/src/ops_model/models/interpretability/diffex/directions/submit.py b/src/ops_model/models/interpretability/diffae/directions/submit.py similarity index 97% rename from src/ops_model/models/interpretability/diffex/directions/submit.py rename to src/ops_model/models/interpretability/diffae/directions/submit.py index d131efd..8c49241 100644 --- a/src/ops_model/models/interpretability/diffex/directions/submit.py +++ b/src/ops_model/models/interpretability/diffae/directions/submit.py @@ -1,6 +1,6 @@ """Submit Stage 3 (directions + traversal) to SLURM (1 GPU). - python -m ops_model.models.interpretability.diffex.directions.submit --target HSPA5 + python -m ops_model.models.interpretability.diffae.directions.submit --target HSPA5 """ from __future__ import annotations diff --git a/src/ops_model/models/interpretability/diffex/directions/train_directions.py b/src/ops_model/models/interpretability/diffae/directions/train_directions.py similarity index 100% rename from src/ops_model/models/interpretability/diffex/directions/train_directions.py rename to src/ops_model/models/interpretability/diffae/directions/train_directions.py diff --git a/src/ops_model/models/interpretability/diffex/directions/traverse.py b/src/ops_model/models/interpretability/diffae/directions/traverse.py similarity index 99% rename from src/ops_model/models/interpretability/diffex/directions/traverse.py rename to src/ops_model/models/interpretability/diffae/directions/traverse.py index edd2840..816c305 100644 --- a/src/ops_model/models/interpretability/diffex/directions/traverse.py +++ b/src/ops_model/models/interpretability/diffae/directions/traverse.py @@ -14,8 +14,8 @@ from ..classifier.celldino_features import embed_crops from ..classifier.config import slugify -from ..diffae.config import DiffAEConfig -from ..diffae.model import DiffAE +from ..generator.config import DiffAEConfig +from ..generator.model import DiffAE def load_diffae(cfg, dev): diff --git a/src/ops_model/models/interpretability/diffex/figures/METHODS_final.txt b/src/ops_model/models/interpretability/diffae/figures/METHODS_final.txt similarity index 100% rename from src/ops_model/models/interpretability/diffex/figures/METHODS_final.txt rename to src/ops_model/models/interpretability/diffae/figures/METHODS_final.txt diff --git a/src/ops_model/models/interpretability/diffex/figures/METHODS_traversal_montage.md b/src/ops_model/models/interpretability/diffae/figures/METHODS_traversal_montage.md similarity index 100% rename from src/ops_model/models/interpretability/diffex/figures/METHODS_traversal_montage.md rename to src/ops_model/models/interpretability/diffae/figures/METHODS_traversal_montage.md diff --git a/src/ops_model/models/interpretability/diffex/figures/METHODS_traversal_montage.txt b/src/ops_model/models/interpretability/diffae/figures/METHODS_traversal_montage.txt similarity index 100% rename from src/ops_model/models/interpretability/diffex/figures/METHODS_traversal_montage.txt rename to src/ops_model/models/interpretability/diffae/figures/METHODS_traversal_montage.txt diff --git a/src/ops_model/models/interpretability/diffex/figures/_setacc_common.py b/src/ops_model/models/interpretability/diffae/figures/_setacc_common.py similarity index 96% rename from src/ops_model/models/interpretability/diffex/figures/_setacc_common.py rename to src/ops_model/models/interpretability/diffae/figures/_setacc_common.py index 702475d..a0f1d15 100644 --- a/src/ops_model/models/interpretability/diffex/figures/_setacc_common.py +++ b/src/ops_model/models/interpretability/diffae/figures/_setacc_common.py @@ -10,11 +10,11 @@ import pandas as pd import zarr -from ops_model.models.interpretability.diffex.classifier.config import slugify -from ops_model.models.interpretability.diffex.classifier.data import make_labels_df, materialize_crops -from ops_model.models.interpretability.diffex.directions.config import DirConfig -from ops_model.models.interpretability.diffex.viewer._fluor_topcells import _overlay_rgba -from ops_model.models.interpretability.diffex.viewer.build_pc_crops_masked import BASE, CROP_SIZE, _crop, _zarr_patch +from ops_model.models.interpretability.diffae.classifier.config import slugify +from ops_model.models.interpretability.diffae.classifier.data import make_labels_df, materialize_crops +from ops_model.models.interpretability.diffae.directions.config import DirConfig +from ops_model.models.interpretability.diffae.viewer._fluor_topcells import _overlay_rgba +from ops_model.models.interpretability.diffae.viewer.build_pc_crops_masked import BASE, CROP_SIZE, _crop, _zarr_patch OUT = "/hpc/projects/icd.fast.ops/analysis/figure4_setacc_panel" RANK_BASE = "/hpc/projects/icd.fast.ops/models/diffex/viewer_assets_v5/_rankings/fluor_shap" diff --git a/src/ops_model/models/interpretability/diffex/figures/_setacc_phase.py b/src/ops_model/models/interpretability/diffae/figures/_setacc_phase.py similarity index 100% rename from src/ops_model/models/interpretability/diffex/figures/_setacc_phase.py rename to src/ops_model/models/interpretability/diffae/figures/_setacc_phase.py diff --git a/src/ops_model/models/interpretability/diffex/figures/auto_pick_and_plot.py b/src/ops_model/models/interpretability/diffae/figures/auto_pick_and_plot.py similarity index 97% rename from src/ops_model/models/interpretability/diffex/figures/auto_pick_and_plot.py rename to src/ops_model/models/interpretability/diffae/figures/auto_pick_and_plot.py index 234004e..a9f5e52 100644 --- a/src/ops_model/models/interpretability/diffex/figures/auto_pick_and_plot.py +++ b/src/ops_model/models/interpretability/diffae/figures/auto_pick_and_plot.py @@ -10,8 +10,8 @@ sys.path.insert(0, os.path.dirname(os.path.abspath(__file__))) os.environ.setdefault("OPS_DIFFEX_ASSETS", "viewer_assets_v5") -from ops_model.models.interpretability.diffex.viewer.morpho_pipeline import MORPHO_TARGETS -from ops_model.models.interpretability.diffex.classifier.config import slugify +from ops_model.models.interpretability.diffae.viewer.morpho_pipeline import MORPHO_TARGETS +from ops_model.models.interpretability.diffae.classifier.config import slugify VA = "/hpc/projects/icd.fast.ops/models/diffex/viewer_assets_v5/_morphometrics" BAD = re.compile("moment|hu_|inertia|eigval|intensity|haralick|zernike|glcm|orientation|centroid|_timing") diff --git a/src/ops_model/models/interpretability/diffex/figures/cis_golgi_alternatives.py b/src/ops_model/models/interpretability/diffae/figures/cis_golgi_alternatives.py similarity index 100% rename from src/ops_model/models/interpretability/diffex/figures/cis_golgi_alternatives.py rename to src/ops_model/models/interpretability/diffae/figures/cis_golgi_alternatives.py diff --git a/src/ops_model/models/interpretability/diffex/figures/debug_setacc_top100.py b/src/ops_model/models/interpretability/diffae/figures/debug_setacc_top100.py similarity index 98% rename from src/ops_model/models/interpretability/diffex/figures/debug_setacc_top100.py rename to src/ops_model/models/interpretability/diffae/figures/debug_setacc_top100.py index 80f8e6a..d9dfce0 100644 --- a/src/ops_model/models/interpretability/diffex/figures/debug_setacc_top100.py +++ b/src/ops_model/models/interpretability/diffae/figures/debug_setacc_top100.py @@ -15,7 +15,7 @@ import matplotlib.pyplot as plt import numpy as np -from ops_model.models.interpretability.diffex.classifier.config import slugify +from ops_model.models.interpretability.diffae.classifier.config import slugify from _setacc_common import (COMPLEX_COLS, GENE_COLS, OUT, CROP_SIZE, materialize_class, seg_crop, composite) plt.rcParams["pdf.fonttype"] = 42 diff --git a/src/ops_model/models/interpretability/diffex/figures/ebi_peripheral_droplets.py b/src/ops_model/models/interpretability/diffae/figures/ebi_peripheral_droplets.py similarity index 99% rename from src/ops_model/models/interpretability/diffex/figures/ebi_peripheral_droplets.py rename to src/ops_model/models/interpretability/diffae/figures/ebi_peripheral_droplets.py index 2f671c1..2eb9bb8 100644 --- a/src/ops_model/models/interpretability/diffex/figures/ebi_peripheral_droplets.py +++ b/src/ops_model/models/interpretability/diffae/figures/ebi_peripheral_droplets.py @@ -41,7 +41,7 @@ from figure_ebi_morpho_violin import draw_violin from figure_multirank_ebi_grid import CACHE, OUT, ebi_rows, top_rows -from ops_model.models.interpretability.diffex.viewer.build_pc_crops_masked import BASE, _crop, _zarr_patch +from ops_model.models.interpretability.diffae.viewer.build_pc_crops_masked import BASE, _crop, _zarr_patch from organelle_profiler.feature_extraction.localization_features import compute_localization_features diff --git a/src/ops_model/models/interpretability/diffex/figures/figure4_morpho_traversal.py b/src/ops_model/models/interpretability/diffae/figures/figure4_morpho_traversal.py similarity index 100% rename from src/ops_model/models/interpretability/diffex/figures/figure4_morpho_traversal.py rename to src/ops_model/models/interpretability/diffae/figures/figure4_morpho_traversal.py diff --git a/src/ops_model/models/interpretability/diffex/figures/figure4_morpho_violin.py b/src/ops_model/models/interpretability/diffae/figures/figure4_morpho_violin.py similarity index 98% rename from src/ops_model/models/interpretability/diffex/figures/figure4_morpho_violin.py rename to src/ops_model/models/interpretability/diffae/figures/figure4_morpho_violin.py index fc6aafb..036e31e 100644 --- a/src/ops_model/models/interpretability/diffex/figures/figure4_morpho_violin.py +++ b/src/ops_model/models/interpretability/diffae/figures/figure4_morpho_violin.py @@ -17,8 +17,8 @@ import pandas as pd from figure4_morpho_traversal import FIGURES, VA, image_panels, render_images -from ops_model.models.interpretability.diffex.viewer.morpho_pipeline import MORPHO_TARGETS, real_percell -from ops_model.models.interpretability.diffex.classifier.config import slugify +from ops_model.models.interpretability.diffae.viewer.morpho_pipeline import MORPHO_TARGETS, real_percell +from ops_model.models.interpretability.diffae.classifier.config import slugify plt.rcParams["pdf.fonttype"] = 42 plt.rcParams["svg.fonttype"] = "none" diff --git a/src/ops_model/models/interpretability/diffex/figures/figure4_setacc_panel.py b/src/ops_model/models/interpretability/diffae/figures/figure4_setacc_panel.py similarity index 100% rename from src/ops_model/models/interpretability/diffex/figures/figure4_setacc_panel.py rename to src/ops_model/models/interpretability/diffae/figures/figure4_setacc_panel.py diff --git a/src/ops_model/models/interpretability/diffex/figures/figure4_setacc_panel_fluorB.py b/src/ops_model/models/interpretability/diffae/figures/figure4_setacc_panel_fluorB.py similarity index 100% rename from src/ops_model/models/interpretability/diffex/figures/figure4_setacc_panel_fluorB.py rename to src/ops_model/models/interpretability/diffae/figures/figure4_setacc_panel_fluorB.py diff --git a/src/ops_model/models/interpretability/diffex/figures/figure4_setacc_panel_newpheno.py b/src/ops_model/models/interpretability/diffae/figures/figure4_setacc_panel_newpheno.py similarity index 100% rename from src/ops_model/models/interpretability/diffex/figures/figure4_setacc_panel_newpheno.py rename to src/ops_model/models/interpretability/diffae/figures/figure4_setacc_panel_newpheno.py diff --git a/src/ops_model/models/interpretability/diffex/figures/figure4_setacc_panel_phase.py b/src/ops_model/models/interpretability/diffae/figures/figure4_setacc_panel_phase.py similarity index 100% rename from src/ops_model/models/interpretability/diffex/figures/figure4_setacc_panel_phase.py rename to src/ops_model/models/interpretability/diffae/figures/figure4_setacc_panel_phase.py diff --git a/src/ops_model/models/interpretability/diffex/figures/figure_ebi_morpho_violin.py b/src/ops_model/models/interpretability/diffae/figures/figure_ebi_morpho_violin.py similarity index 100% rename from src/ops_model/models/interpretability/diffex/figures/figure_ebi_morpho_violin.py rename to src/ops_model/models/interpretability/diffae/figures/figure_ebi_morpho_violin.py diff --git a/src/ops_model/models/interpretability/diffex/figures/figure_multirank_ebi_grid.py b/src/ops_model/models/interpretability/diffae/figures/figure_multirank_ebi_grid.py similarity index 100% rename from src/ops_model/models/interpretability/diffex/figures/figure_multirank_ebi_grid.py rename to src/ops_model/models/interpretability/diffae/figures/figure_multirank_ebi_grid.py diff --git a/src/ops_model/models/interpretability/diffex/figures/fluor_panel_montages.py b/src/ops_model/models/interpretability/diffae/figures/fluor_panel_montages.py similarity index 100% rename from src/ops_model/models/interpretability/diffex/figures/fluor_panel_montages.py rename to src/ops_model/models/interpretability/diffae/figures/fluor_panel_montages.py diff --git a/src/ops_model/models/interpretability/diffex/figures/fluor_shap_montages.py b/src/ops_model/models/interpretability/diffae/figures/fluor_shap_montages.py similarity index 99% rename from src/ops_model/models/interpretability/diffex/figures/fluor_shap_montages.py rename to src/ops_model/models/interpretability/diffae/figures/fluor_shap_montages.py index 3fe62eb..4aad69a 100644 --- a/src/ops_model/models/interpretability/diffex/figures/fluor_shap_montages.py +++ b/src/ops_model/models/interpretability/diffae/figures/fluor_shap_montages.py @@ -15,7 +15,7 @@ import numpy as np import pandas as pd -from ops_model.models.interpretability.diffex.classifier.config import slugify +from ops_model.models.interpretability.diffae.classifier.config import slugify from _setacc_common import _materialize, seg_crop, composite, CROP_SIZE MR = "/hpc/projects/icd.fast.ops/models/alex_lin_attention/v5/multi_rank/shap_screen/shap_screen_fluor_all.compact.parquet" diff --git a/src/ops_model/models/interpretability/diffex/figures/gen_validation/bag_sweep_plots.py b/src/ops_model/models/interpretability/diffae/figures/gen_validation/bag_sweep_plots.py similarity index 99% rename from src/ops_model/models/interpretability/diffex/figures/gen_validation/bag_sweep_plots.py rename to src/ops_model/models/interpretability/diffae/figures/gen_validation/bag_sweep_plots.py index f1e77d8..68fa82b 100644 --- a/src/ops_model/models/interpretability/diffex/figures/gen_validation/bag_sweep_plots.py +++ b/src/ops_model/models/interpretability/diffae/figures/gen_validation/bag_sweep_plots.py @@ -29,7 +29,7 @@ CENT_SUF = "_perbag" if CENT_STD == "perbag" else "" COL = {b: cm.viridis(i / (len(BAGS) - 1)) for i, b in enumerate(BAGS)} COLK = {k: cm.viridis(i / (len(KS) - 1)) for i, k in enumerate(KS)} -from ops_model.models.interpretability.diffex.classifier.config import slugify +from ops_model.models.interpretability.diffae.classifier.config import slugify REAL = json.load(open(f"{B}/viewer_assets_v5/real_acc20.json")) diff --git a/src/ops_model/models/interpretability/diffex/figures/gen_validation/bag_sweep_score.py b/src/ops_model/models/interpretability/diffae/figures/gen_validation/bag_sweep_score.py similarity index 93% rename from src/ops_model/models/interpretability/diffex/figures/gen_validation/bag_sweep_score.py rename to src/ops_model/models/interpretability/diffae/figures/gen_validation/bag_sweep_score.py index 773bd66..82e6cdb 100644 --- a/src/ops_model/models/interpretability/diffex/figures/gen_validation/bag_sweep_score.py +++ b/src/ops_model/models/interpretability/diffae/figures/gen_validation/bag_sweep_score.py @@ -16,11 +16,11 @@ def score_shard(genes): import torch - from ops_model.models.interpretability.diffex.viewer.score_generated import score_embs_v5 - from ops_model.models.interpretability.diffex.viewer.set_classifier import load_set_classifier, V5_CKPT_ROOT, V5_RUNS + from ops_model.models.interpretability.diffae.viewer.score_generated import score_embs_v5 + from ops_model.models.interpretability.diffae.viewer.set_classifier import load_set_classifier, V5_CKPT_ROOT, V5_RUNS os.makedirs(OUT, exist_ok=True) dev = "cuda" if torch.cuda.is_available() else "cpu" - from ops_model.models.interpretability.diffex.classifier.config import slugify + from ops_model.models.interpretability.diffae.classifier.config import slugify run = V5_RUNS[("phase", "geneKO" if GRAIN == "geneKO" else "complex_ebionly")] m, g2i, c2i = load_set_classifier(run=run, device=dev, root=V5_CKPT_ROOT) ci = c2i.get("Phase2D", 0) diff --git a/src/ops_model/models/interpretability/diffex/figures/gen_validation/centroid_bagsweep.py b/src/ops_model/models/interpretability/diffae/figures/gen_validation/centroid_bagsweep.py similarity index 97% rename from src/ops_model/models/interpretability/diffex/figures/gen_validation/centroid_bagsweep.py rename to src/ops_model/models/interpretability/diffae/figures/gen_validation/centroid_bagsweep.py index 2041eff..43fffc2 100644 --- a/src/ops_model/models/interpretability/diffex/figures/gen_validation/centroid_bagsweep.py +++ b/src/ops_model/models/interpretability/diffae/figures/gen_validation/centroid_bagsweep.py @@ -21,7 +21,7 @@ def _cz(): - from ops_model.models.interpretability.diffex.classifier.config import slugify + from ops_model.models.interpretability.diffae.classifier.config import slugify d = np.load(f"{CENTD}/{GRAIN}_centroids.npz", allow_pickle=True) names = list(d["names"]); cidx = {slugify(str(c)): i for i, c in enumerate(names)} cz = (d["cents"] - d["mu"]) / d["sd"]; cz = cz / (np.linalg.norm(cz, axis=1, keepdims=True) + 1e-9) @@ -42,7 +42,7 @@ def compute_mu(): def score_shard(genes): - from ops_model.models.interpretability.diffex.classifier.config import slugify + from ops_model.models.interpretability.diffae.classifier.config import slugify os.makedirs(PART, exist_ok=True) cz, cidx = _cz() if STD == "global": diff --git a/src/ops_model/models/interpretability/diffex/figures/gen_validation/centroid_halves.py b/src/ops_model/models/interpretability/diffae/figures/gen_validation/centroid_halves.py similarity index 96% rename from src/ops_model/models/interpretability/diffex/figures/gen_validation/centroid_halves.py rename to src/ops_model/models/interpretability/diffae/figures/gen_validation/centroid_halves.py index d13cd03..d8096c0 100644 --- a/src/ops_model/models/interpretability/diffex/figures/gen_validation/centroid_halves.py +++ b/src/ops_model/models/interpretability/diffae/figures/gen_validation/centroid_halves.py @@ -14,7 +14,7 @@ def _cz(): - from ops_model.models.interpretability.diffex.classifier.config import slugify + from ops_model.models.interpretability.diffae.classifier.config import slugify d = np.load(f"{CENTD}/{GRAIN}_centroids.npz", allow_pickle=True) names = list(d["names"]); cidx = {slugify(str(c)): i for i, c in enumerate(names)} cz = (d["cents"] - d["mu"]) / d["sd"]; cz = cz / (np.linalg.norm(cz, axis=1, keepdims=True) + 1e-9) @@ -29,7 +29,7 @@ def _top1(vecs, cz, ti): def score_shard(genes): - from ops_model.models.interpretability.diffex.classifier.config import slugify + from ops_model.models.interpretability.diffae.classifier.config import slugify os.makedirs(PART, exist_ok=True) cz, cidx = _cz(); mg = np.load(MU); mu_g, sd_g = mg["mu"], mg["sd"] first, second = {}, {} diff --git a/src/ops_model/models/interpretability/diffex/figures/gen_validation/centroid_pooled_bagsweep.py b/src/ops_model/models/interpretability/diffae/figures/gen_validation/centroid_pooled_bagsweep.py similarity index 97% rename from src/ops_model/models/interpretability/diffex/figures/gen_validation/centroid_pooled_bagsweep.py rename to src/ops_model/models/interpretability/diffae/figures/gen_validation/centroid_pooled_bagsweep.py index 36cd9ec..4d4c20d 100644 --- a/src/ops_model/models/interpretability/diffex/figures/gen_validation/centroid_pooled_bagsweep.py +++ b/src/ops_model/models/interpretability/diffae/figures/gen_validation/centroid_pooled_bagsweep.py @@ -23,7 +23,7 @@ def _cz(): - from ops_model.models.interpretability.diffex.classifier.config import slugify + from ops_model.models.interpretability.diffae.classifier.config import slugify d = np.load(f"{CENTD}/{GRAIN}_centroids.npz", allow_pickle=True) names = list(d["names"]); cidx = {slugify(str(c)): i for i, c in enumerate(names)} cz = (d["cents"] - d["mu"]) / d["sd"]; cz = cz / (np.linalg.norm(cz, axis=1, keepdims=True) + 1e-9) @@ -38,7 +38,7 @@ def _pooled(vecs, cz, ti): def score_shard(genes): - from ops_model.models.interpretability.diffex.classifier.config import slugify + from ops_model.models.interpretability.diffae.classifier.config import slugify os.makedirs(PART, exist_ok=True) cz, cidx, mu_r, sd_r = _cz() if STD == "global": diff --git a/src/ops_model/models/interpretability/diffex/figures/gen_validation/control_halves_zscore.py b/src/ops_model/models/interpretability/diffae/figures/gen_validation/control_halves_zscore.py similarity index 97% rename from src/ops_model/models/interpretability/diffex/figures/gen_validation/control_halves_zscore.py rename to src/ops_model/models/interpretability/diffae/figures/gen_validation/control_halves_zscore.py index b8b9785..3a94d54 100644 --- a/src/ops_model/models/interpretability/diffex/figures/gen_validation/control_halves_zscore.py +++ b/src/ops_model/models/interpretability/diffae/figures/gen_validation/control_halves_zscore.py @@ -19,7 +19,7 @@ def _cz(): - from ops_model.models.interpretability.diffex.classifier.config import slugify + from ops_model.models.interpretability.diffae.classifier.config import slugify d = np.load(f"{CENTD}/{GRAIN}_centroids.npz", allow_pickle=True) names = list(d["names"]); cidx = {slugify(str(c)): i for i, c in enumerate(names)} cz = (d["cents"] - d["mu"]) / d["sd"]; cz = cz / (np.linalg.norm(cz, axis=1, keepdims=True) + 1e-9) @@ -35,7 +35,7 @@ def _t1(vecs, cz, ti): def score_shard(genes): - from ops_model.models.interpretability.diffex.classifier.config import slugify + from ops_model.models.interpretability.diffae.classifier.config import slugify os.makedirs(PART, exist_ok=True) cz, cidx = _cz(); mg = np.load(MU); mu_g, sd_g = mg["mu"], mg["sd"] res = {s: {h: {} for h in ("first", "second")} for s in ("global", "perbag")} diff --git a/src/ops_model/models/interpretability/diffex/figures/gen_validation/embcheck.py b/src/ops_model/models/interpretability/diffae/figures/gen_validation/embcheck.py similarity index 90% rename from src/ops_model/models/interpretability/diffex/figures/gen_validation/embcheck.py rename to src/ops_model/models/interpretability/diffae/figures/gen_validation/embcheck.py index 16d251e..4e4868c 100644 --- a/src/ops_model/models/interpretability/diffex/figures/gen_validation/embcheck.py +++ b/src/ops_model/models/interpretability/diffae/figures/gen_validation/embcheck.py @@ -8,9 +8,9 @@ def check(gene="AACS", ai=6): - from ops_model.models.interpretability.diffex.viewer.score_generated import _emb_frames - from ops_model.models.interpretability.diffex.classifier.celldino_features import embed_crops - from ops_model.models.interpretability.diffex.directions.config import DirConfig + from ops_model.models.interpretability.diffae.viewer.score_generated import _emb_frames + from ops_model.models.interpretability.diffae.classifier.celldino_features import embed_crops + from ops_model.models.interpretability.diffae.directions.config import DirConfig trav = f"{B}/viewer_assets_valid200/phase/geneKO/{gene}" cfg = DirConfig(grain="geneKO", target=gene, device="cuda") embB = np.asarray(_emb_frames(cfg, trav, ai, embed_crops), np.float32) # original path diff --git a/src/ops_model/models/interpretability/diffex/figures/gen_validation/embedding_diagnostics.py b/src/ops_model/models/interpretability/diffae/figures/gen_validation/embedding_diagnostics.py similarity index 100% rename from src/ops_model/models/interpretability/diffex/figures/gen_validation/embedding_diagnostics.py rename to src/ops_model/models/interpretability/diffae/figures/gen_validation/embedding_diagnostics.py diff --git a/src/ops_model/models/interpretability/diffex/figures/gen_validation/figure4_bagsize_reachfrac.py b/src/ops_model/models/interpretability/diffae/figures/gen_validation/figure4_bagsize_reachfrac.py similarity index 100% rename from src/ops_model/models/interpretability/diffex/figures/gen_validation/figure4_bagsize_reachfrac.py rename to src/ops_model/models/interpretability/diffae/figures/gen_validation/figure4_bagsize_reachfrac.py diff --git a/src/ops_model/models/interpretability/diffex/figures/gen_validation/figure4_v5_accuracy_summary.py b/src/ops_model/models/interpretability/diffae/figures/gen_validation/figure4_v5_accuracy_summary.py similarity index 98% rename from src/ops_model/models/interpretability/diffex/figures/gen_validation/figure4_v5_accuracy_summary.py rename to src/ops_model/models/interpretability/diffae/figures/gen_validation/figure4_v5_accuracy_summary.py index 796b604..8233566 100644 --- a/src/ops_model/models/interpretability/diffex/figures/gen_validation/figure4_v5_accuracy_summary.py +++ b/src/ops_model/models/interpretability/diffae/figures/gen_validation/figure4_v5_accuracy_summary.py @@ -34,7 +34,7 @@ def _complex_allow(thr=0.9): label_name in the ebionly eval CSV (his reported members), then mean of those member accuracies.""" import csv as _csv from collections import defaultdict - from ops_model.models.interpretability.diffex.classifier.config import slugify + from ops_model.models.interpretability.diffae.classifier.config import slugify by = defaultdict(list) for r in _csv.DictReader(open(f"{EVAL}/eval_phase_ebionly_e200_pergene_val.csv")): if int(r["n_cells"]) == 20: @@ -48,7 +48,7 @@ def _real_map(sub): return _real_acc20("eval_phase_e200_pergene_val.csv") import csv as _csv from collections import defaultdict - from ops_model.models.interpretability.diffex.classifier.config import slugify + from ops_model.models.interpretability.diffae.classifier.config import slugify by = defaultdict(list) for r in _csv.DictReader(open(f"{EVAL}/eval_phase_ebionly_e200_pergene_val.csv")): if int(r["n_cells"]) == 20: diff --git a/src/ops_model/models/interpretability/diffex/figures/gen_validation/gen_alpha_embedding.py b/src/ops_model/models/interpretability/diffae/figures/gen_validation/gen_alpha_embedding.py similarity index 100% rename from src/ops_model/models/interpretability/diffex/figures/gen_validation/gen_alpha_embedding.py rename to src/ops_model/models/interpretability/diffae/figures/gen_validation/gen_alpha_embedding.py diff --git a/src/ops_model/models/interpretability/diffex/figures/gen_validation/gen_embed_refit.py b/src/ops_model/models/interpretability/diffae/figures/gen_validation/gen_embed_refit.py similarity index 99% rename from src/ops_model/models/interpretability/diffex/figures/gen_validation/gen_embed_refit.py rename to src/ops_model/models/interpretability/diffae/figures/gen_validation/gen_embed_refit.py index 5e45314..c6752c7 100644 --- a/src/ops_model/models/interpretability/diffex/figures/gen_validation/gen_embed_refit.py +++ b/src/ops_model/models/interpretability/diffae/figures/gen_validation/gen_embed_refit.py @@ -155,8 +155,8 @@ def webp_vs_float_v5(n=120): import os, glob from PIL import Image from scipy.spatial.distance import cdist - from ops_model.models.interpretability.diffex.classifier.celldino_features import embed_crops - from ops_model.models.interpretability.diffex.directions.config import DirConfig + from ops_model.models.interpretability.diffae.classifier.celldino_features import embed_crops + from ops_model.models.interpretability.diffae.directions.config import DirConfig a, comp, mean = gp._load_embedding() Xr = np.asarray(a.obsm["X_pca"], np.float64); idx = {nm: i for i, nm in enumerate(a.obs_names)} genes = [g for g in sorted(os.listdir(V5WEBP)) if g in idx diff --git a/src/ops_model/models/interpretability/diffex/figures/gen_validation/gen_phate_passthrough.py b/src/ops_model/models/interpretability/diffae/figures/gen_validation/gen_phate_passthrough.py similarity index 100% rename from src/ops_model/models/interpretability/diffex/figures/gen_validation/gen_phate_passthrough.py rename to src/ops_model/models/interpretability/diffae/figures/gen_validation/gen_phate_passthrough.py diff --git a/src/ops_model/models/interpretability/diffex/figures/gen_validation/gen_real_centroid.py b/src/ops_model/models/interpretability/diffae/figures/gen_validation/gen_real_centroid.py similarity index 98% rename from src/ops_model/models/interpretability/diffex/figures/gen_validation/gen_real_centroid.py rename to src/ops_model/models/interpretability/diffae/figures/gen_validation/gen_real_centroid.py index 116c93a..0ca9896 100644 --- a/src/ops_model/models/interpretability/diffex/figures/gen_validation/gen_real_centroid.py +++ b/src/ops_model/models/interpretability/diffae/figures/gen_validation/gen_real_centroid.py @@ -24,9 +24,9 @@ def _classes(grain): def embed_centroids(grain, classes): - from ops_model.models.interpretability.diffex.viewer.precompute import _gather_class - from ops_model.models.interpretability.diffex.directions.config import DirConfig - from ops_model.models.interpretability.diffex.classifier.config import slugify + from ops_model.models.interpretability.diffae.viewer.precompute import _gather_class + from ops_model.models.interpretability.diffae.directions.config import DirConfig + from ops_model.models.interpretability.diffae.classifier.config import slugify cfg = DirConfig(grain=grain, target=classes[0], device="cuda"); cfg.num_workers = 12 cents, S, SS, n = {}, np.zeros(1024), np.zeros(1024), 0 embs, lbl = [], [] @@ -61,7 +61,7 @@ def merge(grain): def score(grain, cache=None, out=None, cap=None): - from ops_model.models.interpretability.diffex.classifier.config import slugify + from ops_model.models.interpretability.diffae.classifier.config import slugify CACHE_ = cache or CACHE; OUT_ = out or OUT # explicit args (SLURM-safe) override module globals cap = cap if cap is not None else (int(os.environ.get("GRC_GEN_CAP", "0")) or None) # subsample gen bag to first `cap` cells/class d = np.load(f"{OUT_}/{grain}_centroids.npz", allow_pickle=True) @@ -107,7 +107,7 @@ def score(grain, cache=None, out=None, cap=None): def ceiling(grain): """Real cells (cached embed_crops) → nearest faithful centroid: per-class real mAP/top1/top5 (the ceiling).""" - from ops_model.models.interpretability.diffex.classifier.config import slugify + from ops_model.models.interpretability.diffae.classifier.config import slugify d = np.load(f"{OUT}/{grain}_centroids.npz", allow_pickle=True) names = list(d["names"]); cidx = {slugify(str(c)): i for i, c in enumerate(names)} cz = (d["cents"] - d["mu"]) / d["sd"]; cz = cz / (np.linalg.norm(cz, axis=1, keepdims=True) + 1e-9) @@ -136,7 +136,7 @@ def plot(min_dist=None, acc_thr=None, fname="centroid_topk", overlay=None, overl """3 scores (mAP, top-1, top-5) vs α from the faithful centroids. min_dist: restrict to classes whose real distinctiveness/EBI mAP@20 > min_dist. acc_thr: restrict to real top1_acc>acc_thr @bag20 (SetTransformer subset). overlay: {grain: scored_dir} → dashed second line (e.g. the 200-cell bag) on that grain's axis.""" - from ops_model.models.interpretability.diffex.classifier.config import slugify + from ops_model.models.interpretability.diffae.classifier.config import slugify import matplotlib matplotlib.use("Agg") import matplotlib.pyplot as plt diff --git a/src/ops_model/models/interpretability/diffex/figures/gen_validation/gen_real_distinct.py b/src/ops_model/models/interpretability/diffae/figures/gen_validation/gen_real_distinct.py similarity index 100% rename from src/ops_model/models/interpretability/diffex/figures/gen_validation/gen_real_distinct.py rename to src/ops_model/models/interpretability/diffae/figures/gen_validation/gen_real_distinct.py diff --git a/src/ops_model/models/interpretability/diffex/figures/gen_validation/ntc_inverse_gap.py b/src/ops_model/models/interpretability/diffae/figures/gen_validation/ntc_inverse_gap.py similarity index 96% rename from src/ops_model/models/interpretability/diffex/figures/gen_validation/ntc_inverse_gap.py rename to src/ops_model/models/interpretability/diffae/figures/gen_validation/ntc_inverse_gap.py index 7938830..9f67deb 100644 --- a/src/ops_model/models/interpretability/diffex/figures/gen_validation/ntc_inverse_gap.py +++ b/src/ops_model/models/interpretability/diffae/figures/gen_validation/ntc_inverse_gap.py @@ -41,8 +41,8 @@ def _first_done(channel=None): def run(channel=None, gene=None): import torch # noqa - from ops_model.models.interpretability.diffex.classifier.celldino_features import embed_crops - from ops_model.models.interpretability.diffex.directions.config import DirConfig + from ops_model.models.interpretability.diffae.classifier.celldino_features import embed_crops + from ops_model.models.interpretability.diffae.directions.config import DirConfig os.makedirs(OUT, exist_ok=True) ch, g = _first_done(channel) gene = gene or g @@ -122,9 +122,9 @@ def webp_ab(): (A) float straight into CellDINO, (B) round-tripped through the traversal's 8-bit _save_webp path — and compare CellDINO cosine. Answers: does saving as a proper (float/zarr) image bring generated closer to real?""" import torch, tempfile # noqa - from ops_model.models.interpretability.diffex.classifier.celldino_features import embed_crops - from ops_model.models.interpretability.diffex.directions.config import DirConfig - from ops_model.models.interpretability.diffex.viewer.precompute import _save_webp + from ops_model.models.interpretability.diffae.classifier.celldino_features import embed_crops + from ops_model.models.interpretability.diffae.directions.config import DirConfig + from ops_model.models.interpretability.diffae.viewer.precompute import _save_webp os.makedirs(OUT, exist_ok=True) d = np.load(CTRL, allow_pickle=True) imgs = d["anchor_imgs"].astype(np.float32) # (45,1,160,160) float [-1,1] diff --git a/src/ops_model/models/interpretability/diffex/figures/gen_validation/patch_cache_real.py b/src/ops_model/models/interpretability/diffae/figures/gen_validation/patch_cache_real.py similarity index 93% rename from src/ops_model/models/interpretability/diffex/figures/gen_validation/patch_cache_real.py rename to src/ops_model/models/interpretability/diffae/figures/gen_validation/patch_cache_real.py index b27ce0e..3122e5b 100644 --- a/src/ops_model/models/interpretability/diffex/figures/gen_validation/patch_cache_real.py +++ b/src/ops_model/models/interpretability/diffae/figures/gen_validation/patch_cache_real.py @@ -20,9 +20,9 @@ def _drop42(): def run(): import pandas as pd - from ops_model.models.interpretability.diffex.viewer.precompute import _gather_class - from ops_model.models.interpretability.diffex.directions.config import DirConfig - from ops_model.models.interpretability.diffex.classifier.config import slugify + from ops_model.models.interpretability.diffae.viewer.precompute import _gather_class + from ops_model.models.interpretability.diffae.directions.config import DirConfig + from ops_model.models.interpretability.diffae.classifier.config import slugify genes = _drop42() orig = {slugify(str(x)): str(x) for x in pd.read_parquet(RANKP, columns=["gene"])["gene"].unique()} # slug→ranking name cfg = DirConfig(grain="geneKO", target=genes[0], device="cuda"); cfg.num_workers = 12 diff --git a/src/ops_model/models/interpretability/diffex/figures/gen_validation/publish_multibag_page.py b/src/ops_model/models/interpretability/diffae/figures/gen_validation/publish_multibag_page.py similarity index 100% rename from src/ops_model/models/interpretability/diffex/figures/gen_validation/publish_multibag_page.py rename to src/ops_model/models/interpretability/diffae/figures/gen_validation/publish_multibag_page.py diff --git a/src/ops_model/models/interpretability/diffex/figures/gen_validation/rank_summary.py b/src/ops_model/models/interpretability/diffae/figures/gen_validation/rank_summary.py similarity index 100% rename from src/ops_model/models/interpretability/diffex/figures/gen_validation/rank_summary.py rename to src/ops_model/models/interpretability/diffae/figures/gen_validation/rank_summary.py diff --git a/src/ops_model/models/interpretability/diffex/figures/gen_validation/st_halves_score.py b/src/ops_model/models/interpretability/diffae/figures/gen_validation/st_halves_score.py similarity index 93% rename from src/ops_model/models/interpretability/diffex/figures/gen_validation/st_halves_score.py rename to src/ops_model/models/interpretability/diffae/figures/gen_validation/st_halves_score.py index 6d62ac9..ed8e5b2 100644 --- a/src/ops_model/models/interpretability/diffex/figures/gen_validation/st_halves_score.py +++ b/src/ops_model/models/interpretability/diffae/figures/gen_validation/st_halves_score.py @@ -12,9 +12,9 @@ def score_shard(genes): import torch - from ops_model.models.interpretability.diffex.viewer.score_generated import score_embs_v5 - from ops_model.models.interpretability.diffex.viewer.set_classifier import load_set_classifier, V5_CKPT_ROOT, V5_RUNS - from ops_model.models.interpretability.diffex.classifier.config import slugify + from ops_model.models.interpretability.diffae.viewer.score_generated import score_embs_v5 + from ops_model.models.interpretability.diffae.viewer.set_classifier import load_set_classifier, V5_CKPT_ROOT, V5_RUNS + from ops_model.models.interpretability.diffae.classifier.config import slugify os.makedirs(OUT, exist_ok=True) dev = "cuda" if torch.cuda.is_available() else "cpu" run = V5_RUNS[("phase", "geneKO" if GRAIN == "geneKO" else "complex_ebionly")] diff --git a/src/ops_model/models/interpretability/diffex/figures/gen_validation/std_anchor_test.py b/src/ops_model/models/interpretability/diffae/figures/gen_validation/std_anchor_test.py similarity index 100% rename from src/ops_model/models/interpretability/diffex/figures/gen_validation/std_anchor_test.py rename to src/ops_model/models/interpretability/diffae/figures/gen_validation/std_anchor_test.py diff --git a/src/ops_model/models/interpretability/diffex/figures/gen_validation/stepabl_compare.py b/src/ops_model/models/interpretability/diffae/figures/gen_validation/stepabl_compare.py similarity index 100% rename from src/ops_model/models/interpretability/diffex/figures/gen_validation/stepabl_compare.py rename to src/ops_model/models/interpretability/diffae/figures/gen_validation/stepabl_compare.py diff --git a/src/ops_model/models/interpretability/diffex/figures/gen_validation/valid200_alphastep.py b/src/ops_model/models/interpretability/diffae/figures/gen_validation/valid200_alphastep.py similarity index 100% rename from src/ops_model/models/interpretability/diffex/figures/gen_validation/valid200_alphastep.py rename to src/ops_model/models/interpretability/diffae/figures/gen_validation/valid200_alphastep.py diff --git a/src/ops_model/models/interpretability/diffex/figures/gen_validation/valid200_cache_build.py b/src/ops_model/models/interpretability/diffae/figures/gen_validation/valid200_cache_build.py similarity index 100% rename from src/ops_model/models/interpretability/diffex/figures/gen_validation/valid200_cache_build.py rename to src/ops_model/models/interpretability/diffae/figures/gen_validation/valid200_cache_build.py diff --git a/src/ops_model/models/interpretability/diffex/figures/gen_validation/valid200_capcheck.py b/src/ops_model/models/interpretability/diffae/figures/gen_validation/valid200_capcheck.py similarity index 100% rename from src/ops_model/models/interpretability/diffex/figures/gen_validation/valid200_capcheck.py rename to src/ops_model/models/interpretability/diffae/figures/gen_validation/valid200_capcheck.py diff --git a/src/ops_model/models/interpretability/diffex/figures/gen_validation/valid200_map_compare.py b/src/ops_model/models/interpretability/diffae/figures/gen_validation/valid200_map_compare.py similarity index 98% rename from src/ops_model/models/interpretability/diffex/figures/gen_validation/valid200_map_compare.py rename to src/ops_model/models/interpretability/diffae/figures/gen_validation/valid200_map_compare.py index 67ad8a8..34384ce 100644 --- a/src/ops_model/models/interpretability/diffex/figures/gen_validation/valid200_map_compare.py +++ b/src/ops_model/models/interpretability/diffae/figures/gen_validation/valid200_map_compare.py @@ -9,7 +9,7 @@ """ import json, glob, os import numpy as np -from ops_model.models.interpretability.diffex.classifier.config import slugify +from ops_model.models.interpretability.diffae.classifier.config import slugify CV = "/hpc/projects/icd.fast.ops/analysis/figure4_traversals" OUT = f"{CV}/valid200_metrics" diff --git a/src/ops_model/models/interpretability/diffex/figures/gen_validation/valid200_metrics.py b/src/ops_model/models/interpretability/diffae/figures/gen_validation/valid200_metrics.py similarity index 100% rename from src/ops_model/models/interpretability/diffex/figures/gen_validation/valid200_metrics.py rename to src/ops_model/models/interpretability/diffae/figures/gen_validation/valid200_metrics.py diff --git a/src/ops_model/models/interpretability/diffex/figures/nc_ratio.py b/src/ops_model/models/interpretability/diffae/figures/nc_ratio.py similarity index 100% rename from src/ops_model/models/interpretability/diffex/figures/nc_ratio.py rename to src/ops_model/models/interpretability/diffae/figures/nc_ratio.py diff --git a/src/ops_model/models/interpretability/diffex/figures/ntc_anchor_compare.py b/src/ops_model/models/interpretability/diffae/figures/ntc_anchor_compare.py similarity index 100% rename from src/ops_model/models/interpretability/diffex/figures/ntc_anchor_compare.py rename to src/ops_model/models/interpretability/diffae/figures/ntc_anchor_compare.py diff --git a/src/ops_model/models/interpretability/diffex/figures/phase_montages.py b/src/ops_model/models/interpretability/diffae/figures/phase_montages.py similarity index 100% rename from src/ops_model/models/interpretability/diffex/figures/phase_montages.py rename to src/ops_model/models/interpretability/diffae/figures/phase_montages.py diff --git a/src/ops_model/models/interpretability/diffex/figures/phase_multibag_montages.py b/src/ops_model/models/interpretability/diffae/figures/phase_multibag_montages.py similarity index 100% rename from src/ops_model/models/interpretability/diffex/figures/phase_multibag_montages.py rename to src/ops_model/models/interpretability/diffae/figures/phase_multibag_montages.py diff --git a/src/ops_model/models/interpretability/diffex/figures/phase_sample_montages.py b/src/ops_model/models/interpretability/diffae/figures/phase_sample_montages.py similarity index 100% rename from src/ops_model/models/interpretability/diffex/figures/phase_sample_montages.py rename to src/ops_model/models/interpretability/diffae/figures/phase_sample_montages.py diff --git a/src/ops_model/models/interpretability/diffex/figures/rab_candidate_montages.py b/src/ops_model/models/interpretability/diffae/figures/rab_candidate_montages.py similarity index 93% rename from src/ops_model/models/interpretability/diffex/figures/rab_candidate_montages.py rename to src/ops_model/models/interpretability/diffae/figures/rab_candidate_montages.py index 0286e20..e71a542 100644 --- a/src/ops_model/models/interpretability/diffex/figures/rab_candidate_montages.py +++ b/src/ops_model/models/interpretability/diffae/figures/rab_candidate_montages.py @@ -2,7 +2,7 @@ cells can be picked. Complexes only have top-30 cells.""" from cis_golgi_alternatives import CANDS from debug_setacc_top100 import montage -from ops_model.models.interpretability.diffex.classifier.config import slugify +from ops_model.models.interpretability.diffae.classifier.config import slugify ntc_done = set() for c in CANDS: diff --git a/src/ops_model/models/interpretability/diffex/figures/rebuild_traversals_n100.py b/src/ops_model/models/interpretability/diffae/figures/rebuild_traversals_n100.py similarity index 96% rename from src/ops_model/models/interpretability/diffex/figures/rebuild_traversals_n100.py rename to src/ops_model/models/interpretability/diffae/figures/rebuild_traversals_n100.py index 59d273e..c74f8fe 100644 --- a/src/ops_model/models/interpretability/diffex/figures/rebuild_traversals_n100.py +++ b/src/ops_model/models/interpretability/diffae/figures/rebuild_traversals_n100.py @@ -13,8 +13,8 @@ import sys from pathlib import Path -from ops_model.models.interpretability.diffex.viewer import catalog as C -from ops_model.models.interpretability.diffex.classifier.config import slugify +from ops_model.models.interpretability.diffae.viewer import catalog as C +from ops_model.models.interpretability.diffae.classifier.config import slugify ASSETS = "viewer_assets_v5" RANK = f"{C.OUT}/{ASSETS}/_rankings/fluor" @@ -81,7 +81,7 @@ def _clear_anchor(modality): def rebuild_marker(modality): os.environ["OPS_DIFFEX_ASSETS"] = ASSETS - from ops_model.models.interpretability.diffex.viewer import precompute as P + from ops_model.models.interpretability.diffae.viewer import precompute as P P._ASSETS = ASSETS _clear_anchor(modality) tg = TARGETS[modality] @@ -105,7 +105,7 @@ def rebuild_marker(modality): def gen_phase(targets): """Generate specific phase geneKO targets at n=100 (reuses the existing 100-cell phase anchor cache).""" os.environ["OPS_DIFFEX_ASSETS"] = ASSETS - from ops_model.models.interpretability.diffex.viewer import precompute as P + from ops_model.models.interpretability.diffae.viewer import precompute as P P._ASSETS = ASSETS _clear_anchor("phase") # idempotent: keeps the 100-cell cache P.precompute_marker(grain="geneKO", targets=targets, ckpt=PHASE_CK, out_root=C.OUT, @@ -115,7 +115,7 @@ def gen_phase(targets): def gen_phase_complex(targets): """Generate phase COMPLEX targets at n=100 (full complex names; reuses the 100-cell phase anchor cache).""" os.environ["OPS_DIFFEX_ASSETS"] = ASSETS - from ops_model.models.interpretability.diffex.viewer import precompute as P + from ops_model.models.interpretability.diffae.viewer import precompute as P P._ASSETS = ASSETS _clear_anchor("phase") P.precompute_marker(grain="complex", targets=targets, ckpt=PHASE_CK, out_root=C.OUT, @@ -139,7 +139,7 @@ def gen_phase_chunk(targets, cell_range): cached direction (both ckpt-independent), each writing its own cell{c} dirs → parallelizes the per-cell DDIM inversion across GPUs instead of one long serial job. v5 scoring skipped (whole-target only).""" os.environ["OPS_DIFFEX_ASSETS"] = ASSETS - from ops_model.models.interpretability.diffex.viewer import precompute as P + from ops_model.models.interpretability.diffae.viewer import precompute as P P._ASSETS = ASSETS _clear_anchor("phase") # idempotent: 100-anchor cache kept P.precompute_marker(grain="geneKO", targets=targets, ckpt=PHASE_CK, out_root=C.OUT, diff --git a/src/ops_model/models/interpretability/diffex/figures/traversal_montage_schematic.py b/src/ops_model/models/interpretability/diffae/figures/traversal_montage_schematic.py similarity index 100% rename from src/ops_model/models/interpretability/diffex/figures/traversal_montage_schematic.py rename to src/ops_model/models/interpretability/diffae/figures/traversal_montage_schematic.py diff --git a/src/ops_model/models/interpretability/diffex/figures/virtual_staining_schematic.py b/src/ops_model/models/interpretability/diffae/figures/virtual_staining_schematic.py similarity index 100% rename from src/ops_model/models/interpretability/diffex/figures/virtual_staining_schematic.py rename to src/ops_model/models/interpretability/diffae/figures/virtual_staining_schematic.py diff --git a/src/ops_model/models/interpretability/diffex/diffae/__init__.py b/src/ops_model/models/interpretability/diffae/generator/__init__.py similarity index 100% rename from src/ops_model/models/interpretability/diffex/diffae/__init__.py rename to src/ops_model/models/interpretability/diffae/generator/__init__.py diff --git a/src/ops_model/models/interpretability/diffex/diffae/config.py b/src/ops_model/models/interpretability/diffae/generator/config.py similarity index 100% rename from src/ops_model/models/interpretability/diffex/diffae/config.py rename to src/ops_model/models/interpretability/diffae/generator/config.py diff --git a/src/ops_model/models/interpretability/diffex/diffae/data.py b/src/ops_model/models/interpretability/diffae/generator/data.py similarity index 100% rename from src/ops_model/models/interpretability/diffex/diffae/data.py rename to src/ops_model/models/interpretability/diffae/generator/data.py diff --git a/src/ops_model/models/interpretability/diffex/diffae/diagnose_conditioning.py b/src/ops_model/models/interpretability/diffae/generator/diagnose_conditioning.py similarity index 98% rename from src/ops_model/models/interpretability/diffex/diffae/diagnose_conditioning.py rename to src/ops_model/models/interpretability/diffae/generator/diagnose_conditioning.py index 87f9714..c192839 100644 --- a/src/ops_model/models/interpretability/diffex/diffae/diagnose_conditioning.py +++ b/src/ops_model/models/interpretability/diffae/generator/diagnose_conditioning.py @@ -6,7 +6,7 @@ MSE between two DIFFERENT noises — if embedding-driven change ≪ noise-driven change, the embedding has weak control (the bug we suspect). - python -m ops_model.models.interpretability.diffex.diffae.diagnose_conditioning + python -m ops_model.models.interpretability.diffae.generator.diagnose_conditioning """ from __future__ import annotations diff --git a/src/ops_model/models/interpretability/diffex/diffae/model.py b/src/ops_model/models/interpretability/diffae/generator/model.py similarity index 100% rename from src/ops_model/models/interpretability/diffex/diffae/model.py rename to src/ops_model/models/interpretability/diffae/generator/model.py diff --git a/src/ops_model/models/interpretability/diffex/diffae/plot_metrics.py b/src/ops_model/models/interpretability/diffae/generator/plot_metrics.py similarity index 100% rename from src/ops_model/models/interpretability/diffex/diffae/plot_metrics.py rename to src/ops_model/models/interpretability/diffae/generator/plot_metrics.py diff --git a/src/ops_model/models/interpretability/diffex/diffae/recon.py b/src/ops_model/models/interpretability/diffae/generator/recon.py similarity index 100% rename from src/ops_model/models/interpretability/diffex/diffae/recon.py rename to src/ops_model/models/interpretability/diffae/generator/recon.py diff --git a/src/ops_model/models/interpretability/diffex/diffae/run.py b/src/ops_model/models/interpretability/diffae/generator/run.py similarity index 97% rename from src/ops_model/models/interpretability/diffex/diffae/run.py rename to src/ops_model/models/interpretability/diffae/generator/run.py index e78a0cb..50e6694 100644 --- a/src/ops_model/models/interpretability/diffex/diffae/run.py +++ b/src/ops_model/models/interpretability/diffae/generator/run.py @@ -1,6 +1,6 @@ """Orchestrator for the DiffAE generator stage. - python -m ops_model.models.interpretability.diffex.diffae.run + python -m ops_model.models.interpretability.diffae.generator.run Steps: sample broad phase crops (cached) -> normalize -> train DiffAE (joint encoder + conditional UNet) -> periodic + final reconstruction gate. Writes diff --git a/src/ops_model/models/interpretability/diffex/diffae/submit.py b/src/ops_model/models/interpretability/diffae/generator/submit.py similarity index 98% rename from src/ops_model/models/interpretability/diffex/diffae/submit.py rename to src/ops_model/models/interpretability/diffae/generator/submit.py index d3fe3c3..fcd0925 100644 --- a/src/ops_model/models/interpretability/diffex/diffae/submit.py +++ b/src/ops_model/models/interpretability/diffae/generator/submit.py @@ -1,6 +1,6 @@ """Submit the DiffAE training to SLURM (1 GPU, longer wall clock). - python -m ops_model.models.interpretability.diffex.diffae.submit + python -m ops_model.models.interpretability.diffae.generator.submit =============================== RUNBOOK =============================== Checkpoints (root /hpc/projects/icd.fast.ops/models/diffex/diffae/) and their diff --git a/src/ops_model/models/interpretability/diffex/diffae/train.py b/src/ops_model/models/interpretability/diffae/generator/train.py similarity index 100% rename from src/ops_model/models/interpretability/diffex/diffae/train.py rename to src/ops_model/models/interpretability/diffae/generator/train.py diff --git a/src/ops_model/models/interpretability/diffex/diffae/virtstain_eval.py b/src/ops_model/models/interpretability/diffae/generator/virtstain_eval.py similarity index 98% rename from src/ops_model/models/interpretability/diffex/diffae/virtstain_eval.py rename to src/ops_model/models/interpretability/diffae/generator/virtstain_eval.py index c08df1a..f819d56 100644 --- a/src/ops_model/models/interpretability/diffex/diffae/virtstain_eval.py +++ b/src/ops_model/models/interpretability/diffae/generator/virtstain_eval.py @@ -2,7 +2,7 @@ embedding on a HELD-OUT set of cells (fresh seed → disjoint from training), then report Pearson(pred, real) and save a `phase | predicted | real` montage. - python -m ops_model.models.interpretability.diffex.diffae.virtstain_eval \ + python -m ops_model.models.interpretability.diffae.generator.virtstain_eval \ --out-dir /hpc/projects/icd.fast.ops/analysis/virtual_staining/chromalive561_from_phase \ --marker-channel "mitochondria_ChromaLIVE 561 excitation" --channel mCherry --cond-channel Phase2D """ diff --git a/src/ops_model/models/interpretability/diffex/diffae/virtstain_multi.py b/src/ops_model/models/interpretability/diffae/generator/virtstain_multi.py similarity index 99% rename from src/ops_model/models/interpretability/diffex/diffae/virtstain_multi.py rename to src/ops_model/models/interpretability/diffae/generator/virtstain_multi.py index 90d731e..d984689 100644 --- a/src/ops_model/models/interpretability/diffex/diffae/virtstain_multi.py +++ b/src/ops_model/models/interpretability/diffae/generator/virtstain_multi.py @@ -3,7 +3,7 @@ (each marker's own exps); the phase image is concatenated into the UNet (registered stain) and the marker id selects which channel to render. - python -m ops_model.models.interpretability.diffex.diffae.virtstain_multi --submit --cap 2500 --epochs 120 + python -m ops_model.models.interpretability.diffae.generator.virtstain_multi --submit --cap 2500 --epochs 120 """ from __future__ import annotations diff --git a/src/ops_model/models/interpretability/diffex/kyle_pcs/build_static_explorer.py b/src/ops_model/models/interpretability/diffae/kyle_pcs/build_static_explorer.py similarity index 100% rename from src/ops_model/models/interpretability/diffex/kyle_pcs/build_static_explorer.py rename to src/ops_model/models/interpretability/diffae/kyle_pcs/build_static_explorer.py diff --git a/src/ops_model/models/interpretability/diffex/kyle_pcs/compute_pc_strips.py b/src/ops_model/models/interpretability/diffae/kyle_pcs/compute_pc_strips.py similarity index 100% rename from src/ops_model/models/interpretability/diffex/kyle_pcs/compute_pc_strips.py rename to src/ops_model/models/interpretability/diffae/kyle_pcs/compute_pc_strips.py diff --git a/src/ops_model/models/interpretability/diffex/viewer/__init__.py b/src/ops_model/models/interpretability/diffae/viewer/__init__.py similarity index 100% rename from src/ops_model/models/interpretability/diffex/viewer/__init__.py rename to src/ops_model/models/interpretability/diffae/viewer/__init__.py diff --git a/src/ops_model/models/interpretability/diffex/viewer/_altanchor_build.py b/src/ops_model/models/interpretability/diffae/viewer/_altanchor_build.py similarity index 100% rename from src/ops_model/models/interpretability/diffex/viewer/_altanchor_build.py rename to src/ops_model/models/interpretability/diffae/viewer/_altanchor_build.py diff --git a/src/ops_model/models/interpretability/diffex/viewer/_anchortest.py b/src/ops_model/models/interpretability/diffae/viewer/_anchortest.py similarity index 100% rename from src/ops_model/models/interpretability/diffex/viewer/_anchortest.py rename to src/ops_model/models/interpretability/diffae/viewer/_anchortest.py diff --git a/src/ops_model/models/interpretability/diffex/viewer/_build_stepablation.py b/src/ops_model/models/interpretability/diffae/viewer/_build_stepablation.py similarity index 100% rename from src/ops_model/models/interpretability/diffex/viewer/_build_stepablation.py rename to src/ops_model/models/interpretability/diffae/viewer/_build_stepablation.py diff --git a/src/ops_model/models/interpretability/diffex/viewer/_build_v5_inverted.py b/src/ops_model/models/interpretability/diffae/viewer/_build_v5_inverted.py similarity index 99% rename from src/ops_model/models/interpretability/diffex/viewer/_build_v5_inverted.py rename to src/ops_model/models/interpretability/diffae/viewer/_build_v5_inverted.py index ba09d31..904425e 100644 --- a/src/ops_model/models/interpretability/diffex/viewer/_build_v5_inverted.py +++ b/src/ops_model/models/interpretability/diffae/viewer/_build_v5_inverted.py @@ -9,7 +9,7 @@ complex) whose per-cell rankings exist. That is Alex Lin's top1_acc>0.5@100-cell distinctiveness filter (see _fluor_v5_build.py) — NOT missing data (cells exist for all 1000×55). Lower-acc combos need Alex to gen more. - python -m ops_model.models.interpretability.diffex.viewer._build_v5_inverted markers + python -m ops_model.models.interpretability.diffae.viewer._build_v5_inverted markers """ import json import os @@ -94,7 +94,7 @@ def build_phase_anchor_200(): from pathlib import Path from concurrent.futures import ThreadPoolExecutor from .precompute import _gather_class, _save_webp - from ..diffae.data import normalize + from ..generator.data import normalize from ..directions.config import DirConfig _use_v5() rd = Path(C.OUT) / _V5 / "phase" / "_anchors" / "NTC" @@ -136,7 +136,7 @@ def build_phase_anchor_multirank(hi=400): import numpy as np from concurrent.futures import ThreadPoolExecutor from .precompute import _gather_class, _save_webp - from ..diffae.data import normalize + from ..generator.data import normalize from ..directions.config import DirConfig _use_v5() rd = Path(C.OUT) / _V5 / "phase" / "_anchors" / "NTC" @@ -253,7 +253,7 @@ def build_phase_anchor(): from pathlib import Path from concurrent.futures import ThreadPoolExecutor from .precompute import _gather_class, _save_webp - from ..diffae.data import normalize + from ..generator.data import normalize from ..directions.config import DirConfig _use_v5() cfg = DirConfig(grain="geneKO", target="NTC", control="NTC", device="cuda") @@ -432,7 +432,7 @@ def prebuild_marker_anchor(d, marker_channel, channel): import numpy as np, pandas as pd from concurrent.futures import ThreadPoolExecutor from .precompute import _gather_class, _save_webp - from ..diffae.data import normalize + from ..generator.data import normalize from ..directions.config import DirConfig frp = f"{FRP_DIR}/{slugify(marker_channel)}.parquet" if not os.path.exists(frp): diff --git a/src/ops_model/models/interpretability/diffex/viewer/_build_v5_montages.py b/src/ops_model/models/interpretability/diffae/viewer/_build_v5_montages.py similarity index 98% rename from src/ops_model/models/interpretability/diffex/viewer/_build_v5_montages.py rename to src/ops_model/models/interpretability/diffae/viewer/_build_v5_montages.py index fb7311b..6afad71 100644 --- a/src/ops_model/models/interpretability/diffex/viewer/_build_v5_montages.py +++ b/src/ops_model/models/interpretability/diffae/viewer/_build_v5_montages.py @@ -5,7 +5,7 @@ phase embedding. Reads the merged inverted frames from viewer_assets_v5//geneKO and writes tiles to viewer_assets_v5/_montage/. Only builds markers whose geneKO is 100% present in viewer_assets_v5. - python -m ops_model.models.interpretability.diffex.viewer._build_v5_montages + python -m ops_model.models.interpretability.diffae.viewer._build_v5_montages """ import glob import os diff --git a/src/ops_model/models/interpretability/diffex/viewer/_build_valid200.py b/src/ops_model/models/interpretability/diffae/viewer/_build_valid200.py similarity index 97% rename from src/ops_model/models/interpretability/diffex/viewer/_build_valid200.py rename to src/ops_model/models/interpretability/diffae/viewer/_build_valid200.py index 01858da..22e6555 100644 --- a/src/ops_model/models/interpretability/diffex/viewer/_build_valid200.py +++ b/src/ops_model/models/interpretability/diffae/viewer/_build_valid200.py @@ -8,9 +8,9 @@ Reuses the v5 per-class directions (d_vec + gap are the PRODUCTION values we are validating) via a symlinked _directions tree, so no direction re-fit — only the 200-anchor inversion + 200×7 decodes per target. - python -m ops_model.models.interpretability.diffex.viewer._build_valid200 anchor # 1 GPU: build the 200-cell NTC anchor - python -m ops_model.models.interpretability.diffex.viewer._build_valid200 submit # shard all 1000 geneKO (after anchor) - python -m ops_model.models.interpretability.diffex.viewer._build_valid200 all # anchor job -> shards (afterok dep) + python -m ops_model.models.interpretability.diffae.viewer._build_valid200 anchor # 1 GPU: build the 200-cell NTC anchor + python -m ops_model.models.interpretability.diffae.viewer._build_valid200 submit # shard all 1000 geneKO (after anchor) + python -m ops_model.models.interpretability.diffae.viewer._build_valid200 all # anchor job -> shards (afterok dep) """ import os import sys @@ -65,7 +65,7 @@ def build_anchor(): import numpy as np from concurrent.futures import ThreadPoolExecutor from .precompute import _gather_class, _save_webp - from ..diffae.data import normalize + from ..generator.data import normalize from ..directions.config import DirConfig _use_valid() setup_dirs() diff --git a/src/ops_model/models/interpretability/diffex/viewer/_consolidate_cells.py b/src/ops_model/models/interpretability/diffae/viewer/_consolidate_cells.py similarity index 100% rename from src/ops_model/models/interpretability/diffex/viewer/_consolidate_cells.py rename to src/ops_model/models/interpretability/diffae/viewer/_consolidate_cells.py diff --git a/src/ops_model/models/interpretability/diffex/viewer/_fluor_complex_build.py b/src/ops_model/models/interpretability/diffae/viewer/_fluor_complex_build.py similarity index 100% rename from src/ops_model/models/interpretability/diffex/viewer/_fluor_complex_build.py rename to src/ops_model/models/interpretability/diffae/viewer/_fluor_complex_build.py diff --git a/src/ops_model/models/interpretability/diffex/viewer/_fluor_topcells.py b/src/ops_model/models/interpretability/diffae/viewer/_fluor_topcells.py similarity index 99% rename from src/ops_model/models/interpretability/diffex/viewer/_fluor_topcells.py rename to src/ops_model/models/interpretability/diffae/viewer/_fluor_topcells.py index 464beea..0afc9a9 100644 --- a/src/ops_model/models/interpretability/diffex/viewer/_fluor_topcells.py +++ b/src/ops_model/models/interpretability/diffae/viewer/_fluor_topcells.py @@ -14,7 +14,7 @@ from ..classifier.config import slugify, GRAINS from ..classifier.data import make_labels_df, materialize_crops from ..directions.config import DirConfig -from ..diffae.data import normalize +from ..generator.data import normalize from .precompute import _save_webp from .build_pc_crops_masked import BASE, CROP_SIZE, MASK_DILATION, OVERLAY_RGB, OVERLAY_ALPHA, _crop, _zarr_patch diff --git a/src/ops_model/models/interpretability/diffex/viewer/_fluor_v5_build.py b/src/ops_model/models/interpretability/diffae/viewer/_fluor_v5_build.py similarity index 100% rename from src/ops_model/models/interpretability/diffex/viewer/_fluor_v5_build.py rename to src/ops_model/models/interpretability/diffae/viewer/_fluor_v5_build.py diff --git a/src/ops_model/models/interpretability/diffex/viewer/_migrate_v4_to_v5.py b/src/ops_model/models/interpretability/diffae/viewer/_migrate_v4_to_v5.py similarity index 100% rename from src/ops_model/models/interpretability/diffex/viewer/_migrate_v4_to_v5.py rename to src/ops_model/models/interpretability/diffae/viewer/_migrate_v4_to_v5.py diff --git a/src/ops_model/models/interpretability/diffex/viewer/_phase_vs.py b/src/ops_model/models/interpretability/diffae/viewer/_phase_vs.py similarity index 99% rename from src/ops_model/models/interpretability/diffex/viewer/_phase_vs.py rename to src/ops_model/models/interpretability/diffae/viewer/_phase_vs.py index 446979c..e216753 100644 --- a/src/ops_model/models/interpretability/diffex/viewer/_phase_vs.py +++ b/src/ops_model/models/interpretability/diffae/viewer/_phase_vs.py @@ -27,8 +27,8 @@ def load_vs(dev): - from ..diffae.config import DiffAEConfig - from ..diffae.model import DiffAE + from ..generator.config import DiffAEConfig + from ..generator.model import DiffAE markers = json.load(open(f"{VS_OUT}/markers.json")) cfg = DiffAEConfig(spatial_cond=True, n_markers=len(markers), device="cuda", epochs=1) ema = DiffAE(cfg).to(dev).eval() @@ -48,7 +48,7 @@ def _load_phase(gene, cell, ai, H): @torch.no_grad() def stain(ema, markers, cfg, dev, phase_np, seed=0): """phase_np (1,1,H,H) in [-1,1] → {marker_idx: pred (H,H)} for all markers (fixed xT seed).""" - from ..diffae.virtstain_multi import _sample_marker + from ..generator.virtstain_multi import _sample_marker from ..classifier.celldino_features import embed_crops H = cfg.crop_size emb = torch.as_tensor(embed_crops(phase_np, cfg), dtype=torch.float32, device=dev) diff --git a/src/ops_model/models/interpretability/diffex/viewer/_rebuild_v5.py b/src/ops_model/models/interpretability/diffae/viewer/_rebuild_v5.py similarity index 98% rename from src/ops_model/models/interpretability/diffex/viewer/_rebuild_v5.py rename to src/ops_model/models/interpretability/diffae/viewer/_rebuild_v5.py index 4353051..05ef44f 100644 --- a/src/ops_model/models/interpretability/diffex/viewer/_rebuild_v5.py +++ b/src/ops_model/models/interpretability/diffae/viewer/_rebuild_v5.py @@ -33,7 +33,7 @@ def build_accpool_anchor(): from pathlib import Path from concurrent.futures import ThreadPoolExecutor from .precompute import _gather_class, _ASSETS, _save_webp - from ..diffae.data import normalize + from ..generator.data import normalize from ..directions.config import DirConfig sel = pd.read_csv(SEL25) parq = pd.DataFrame({"gene": "NTC", "experiment": sel.experiment, "well": sel.well, "segmentation": sel.segmentation, diff --git a/src/ops_model/models/interpretability/diffex/viewer/_rescore_rank.py b/src/ops_model/models/interpretability/diffae/viewer/_rescore_rank.py similarity index 97% rename from src/ops_model/models/interpretability/diffex/viewer/_rescore_rank.py rename to src/ops_model/models/interpretability/diffae/viewer/_rescore_rank.py index 79bbf8d..e59758e 100644 --- a/src/ops_model/models/interpretability/diffex/viewer/_rescore_rank.py +++ b/src/ops_model/models/interpretability/diffae/viewer/_rescore_rank.py @@ -11,7 +11,7 @@ def _retarget(assets): """Point the score module at `assets` and return the module (V5_BASE is read at call time).""" - import ops_model.models.interpretability.diffex.viewer.score_generated as SG + import ops_model.models.interpretability.diffae.viewer.score_generated as SG SG.V5_BASE = f"{BASE}/{assets}/phase" return SG diff --git a/src/ops_model/models/interpretability/diffex/viewer/_score_v4.py b/src/ops_model/models/interpretability/diffae/viewer/_score_v4.py similarity index 100% rename from src/ops_model/models/interpretability/diffex/viewer/_score_v4.py rename to src/ops_model/models/interpretability/diffae/viewer/_score_v4.py diff --git a/src/ops_model/models/interpretability/diffex/viewer/_v4acc_test.py b/src/ops_model/models/interpretability/diffae/viewer/_v4acc_test.py similarity index 100% rename from src/ops_model/models/interpretability/diffex/viewer/_v4acc_test.py rename to src/ops_model/models/interpretability/diffae/viewer/_v4acc_test.py diff --git a/src/ops_model/models/interpretability/diffex/viewer/_verify_pt_space.py b/src/ops_model/models/interpretability/diffae/viewer/_verify_pt_space.py similarity index 100% rename from src/ops_model/models/interpretability/diffex/viewer/_verify_pt_space.py rename to src/ops_model/models/interpretability/diffae/viewer/_verify_pt_space.py diff --git a/src/ops_model/models/interpretability/diffex/viewer/_verify_score_bridge.py b/src/ops_model/models/interpretability/diffae/viewer/_verify_score_bridge.py similarity index 100% rename from src/ops_model/models/interpretability/diffex/viewer/_verify_score_bridge.py rename to src/ops_model/models/interpretability/diffae/viewer/_verify_score_bridge.py diff --git a/src/ops_model/models/interpretability/diffex/viewer/altanchor_pairs.json b/src/ops_model/models/interpretability/diffae/viewer/altanchor_pairs.json similarity index 100% rename from src/ops_model/models/interpretability/diffex/viewer/altanchor_pairs.json rename to src/ops_model/models/interpretability/diffae/viewer/altanchor_pairs.json diff --git a/src/ops_model/models/interpretability/diffex/viewer/anchor_cells.py b/src/ops_model/models/interpretability/diffae/viewer/anchor_cells.py similarity index 100% rename from src/ops_model/models/interpretability/diffex/viewer/anchor_cells.py rename to src/ops_model/models/interpretability/diffae/viewer/anchor_cells.py diff --git a/src/ops_model/models/interpretability/diffex/viewer/build_attention_heads.py b/src/ops_model/models/interpretability/diffae/viewer/build_attention_heads.py similarity index 97% rename from src/ops_model/models/interpretability/diffex/viewer/build_attention_heads.py rename to src/ops_model/models/interpretability/diffae/viewer/build_attention_heads.py index 6bfedf9..7c406b1 100644 --- a/src/ops_model/models/interpretability/diffex/viewer/build_attention_heads.py +++ b/src/ops_model/models/interpretability/diffae/viewer/build_attention_heads.py @@ -18,10 +18,10 @@ {AH}/index.json {global_max, assets:{:{:[keys]}}} where modality = "phase" | slugify(marker_channel), grain = geneKO|complex, key = gene | complex-slug. - python -m ops_model.models.interpretability.diffex.viewer.build_attention_heads render # SLURM (all trees) - python -m ops_model.models.interpretability.diffex.viewer.build_attention_heads render --local # serial, no SLURM - python -m ops_model.models.interpretability.diffex.viewer.build_attention_heads render --dry-run - python -m ops_model.models.interpretability.diffex.viewer.build_attention_heads index # (re)aggregate index.json + python -m ops_model.models.interpretability.diffae.viewer.build_attention_heads render # SLURM (all trees) + python -m ops_model.models.interpretability.diffae.viewer.build_attention_heads render --local # serial, no SLURM + python -m ops_model.models.interpretability.diffae.viewer.build_attention_heads render --dry-run + python -m ops_model.models.interpretability.diffae.viewer.build_attention_heads index # (re)aggregate index.json """ from __future__ import annotations diff --git a/src/ops_model/models/interpretability/diffex/viewer/build_complex_ebi_map.py b/src/ops_model/models/interpretability/diffae/viewer/build_complex_ebi_map.py similarity index 100% rename from src/ops_model/models/interpretability/diffex/viewer/build_complex_ebi_map.py rename to src/ops_model/models/interpretability/diffae/viewer/build_complex_ebi_map.py diff --git a/src/ops_model/models/interpretability/diffex/viewer/build_fluor_shap_rankings.py b/src/ops_model/models/interpretability/diffae/viewer/build_fluor_shap_rankings.py similarity index 97% rename from src/ops_model/models/interpretability/diffex/viewer/build_fluor_shap_rankings.py rename to src/ops_model/models/interpretability/diffae/viewer/build_fluor_shap_rankings.py index 1e806cb..1d2b39f 100644 --- a/src/ops_model/models/interpretability/diffex/viewer/build_fluor_shap_rankings.py +++ b/src/ops_model/models/interpretability/diffae/viewer/build_fluor_shap_rankings.py @@ -5,8 +5,8 @@ new CSV cols: gene, channel_name, rank, shap, ..., experiment, well, x_pheno, y_pheno, segmentation_id old schema: channel_name, gene, rank, pma_attention, experiment, well, x_pheno, y_pheno, segmentation, rank_type - python -m ops_model.models.interpretability.diffex.viewer.build_fluor_shap_rankings # local (needs ~64GB) - python -m ops_model.models.interpretability.diffex.viewer.build_fluor_shap_rankings --submit # SLURM cpu, mem 96 + python -m ops_model.models.interpretability.diffae.viewer.build_fluor_shap_rankings # local (needs ~64GB) + python -m ops_model.models.interpretability.diffae.viewer.build_fluor_shap_rankings --submit # SLURM cpu, mem 96 """ from __future__ import annotations diff --git a/src/ops_model/models/interpretability/diffex/viewer/build_montage_features.py b/src/ops_model/models/interpretability/diffae/viewer/build_montage_features.py similarity index 97% rename from src/ops_model/models/interpretability/diffex/viewer/build_montage_features.py rename to src/ops_model/models/interpretability/diffae/viewer/build_montage_features.py index 521d742..06894b7 100644 --- a/src/ops_model/models/interpretability/diffex/viewer/build_montage_features.py +++ b/src/ops_model/models/interpretability/diffae/viewer/build_montage_features.py @@ -7,7 +7,7 @@ viewer_assets/montage_features.json {"features": [base names], "range": {feat: [lo, hi]}, "values": {gene: [0..1 per feature | null]}} - python -m ops_model.models.interpretability.diffex.viewer.build_montage_features + python -m ops_model.models.interpretability.diffae.viewer.build_montage_features """ from __future__ import annotations diff --git a/src/ops_model/models/interpretability/diffex/viewer/build_pc_crops_masked.py b/src/ops_model/models/interpretability/diffae/viewer/build_pc_crops_masked.py similarity index 98% rename from src/ops_model/models/interpretability/diffex/viewer/build_pc_crops_masked.py rename to src/ops_model/models/interpretability/diffae/viewer/build_pc_crops_masked.py index 3869f87..3938deb 100644 --- a/src/ops_model/models/interpretability/diffex/viewer/build_pc_crops_masked.py +++ b/src/ops_model/models/interpretability/diffae/viewer/build_pc_crops_masked.py @@ -10,8 +10,8 @@ {BASE}/{exp}/3-assembly/phenotyping_v3.zarr/{row}/{col}/0/0 image [1,C,1,Y,X], Phase2D=ch0 {BASE}/{exp}/3-assembly/phenotyping_v3.zarr/{row}/{col}/0/labels/cell_seg/0 int32 labels - python -m ops_model.models.interpretability.diffex.viewer.build_pc_crops_masked --sample 24 # preview - python -m ops_model.models.interpretability.diffex.viewer.build_pc_crops_masked # full (overwrites crops/) + python -m ops_model.models.interpretability.diffae.viewer.build_pc_crops_masked --sample 24 # preview + python -m ops_model.models.interpretability.diffae.viewer.build_pc_crops_masked # full (overwrites crops/) """ from __future__ import annotations diff --git a/src/ops_model/models/interpretability/diffex/viewer/build_pc_features.py b/src/ops_model/models/interpretability/diffae/viewer/build_pc_features.py similarity index 99% rename from src/ops_model/models/interpretability/diffex/viewer/build_pc_features.py rename to src/ops_model/models/interpretability/diffae/viewer/build_pc_features.py index 5846745..c976812 100644 --- a/src/ops_model/models/interpretability/diffex/viewer/build_pc_features.py +++ b/src/ops_model/models/interpretability/diffae/viewer/build_pc_features.py @@ -14,7 +14,7 @@ to make "+corr features" line up with the strip's high bins. Compositions are unsigned (|r| / tf-idf) so they need no flip. - python -m ops_model.models.interpretability.diffex.viewer.build_pc_features + python -m ops_model.models.interpretability.diffae.viewer.build_pc_features """ from __future__ import annotations diff --git a/src/ops_model/models/interpretability/diffex/viewer/build_pc_walks.py b/src/ops_model/models/interpretability/diffae/viewer/build_pc_walks.py similarity index 98% rename from src/ops_model/models/interpretability/diffex/viewer/build_pc_walks.py rename to src/ops_model/models/interpretability/diffae/viewer/build_pc_walks.py index f55b529..5d7063d 100644 --- a/src/ops_model/models/interpretability/diffex/viewer/build_pc_walks.py +++ b/src/ops_model/models/interpretability/diffae/viewer/build_pc_walks.py @@ -8,8 +8,8 @@ the PC score std; v_p is the (z-scored) eigenvector, mapped back to raw CellDINO space by the mean per-exp sd. Output: one composite figure per marker (rows = PCs, cols = α). - python -m ops_model.models.interpretability.diffex.viewer.build_pc_walks --markers "Mitochondria_TOMM20" - python -m ops_model.models.interpretability.diffex.viewer.build_pc_walks --all # SLURM, every marker + python -m ops_model.models.interpretability.diffae.viewer.build_pc_walks --markers "Mitochondria_TOMM20" + python -m ops_model.models.interpretability.diffae.viewer.build_pc_walks --all # SLURM, every marker """ from __future__ import annotations diff --git a/src/ops_model/models/interpretability/diffex/viewer/build_pcs.py b/src/ops_model/models/interpretability/diffae/viewer/build_pcs.py similarity index 98% rename from src/ops_model/models/interpretability/diffex/viewer/build_pcs.py rename to src/ops_model/models/interpretability/diffae/viewer/build_pcs.py index 7b46d2b..36e90b6 100644 --- a/src/ops_model/models/interpretability/diffex/viewer/build_pcs.py +++ b/src/ops_model/models/interpretability/diffae/viewer/build_pcs.py @@ -9,8 +9,8 @@ the raw artifacts dir is gone). If Kyle regenerates artifacts, re-run his build_static_explorer.py then point --html at the fresh output. - python -m ops_model.models.interpretability.diffex.viewer.build_pcs - python -m ops_model.models.interpretability.diffex.viewer.build_pcs --html /path/to/pc_explorer_static.html + python -m ops_model.models.interpretability.diffae.viewer.build_pcs + python -m ops_model.models.interpretability.diffae.viewer.build_pcs --html /path/to/pc_explorer_static.html """ from __future__ import annotations diff --git a/src/ops_model/models/interpretability/diffex/viewer/build_pcs_marker.py b/src/ops_model/models/interpretability/diffae/viewer/build_pcs_marker.py similarity index 99% rename from src/ops_model/models/interpretability/diffex/viewer/build_pcs_marker.py rename to src/ops_model/models/interpretability/diffae/viewer/build_pcs_marker.py index 20b7256..62c969c 100644 --- a/src/ops_model/models/interpretability/diffex/viewer/build_pcs_marker.py +++ b/src/ops_model/models/interpretability/diffae/viewer/build_pcs_marker.py @@ -9,7 +9,7 @@ viewer_assets/pcs/markers//index.json (same schema as the phase pcs/index.json) viewer_assets/pcs/markers//crops/pc###_bin##_row#.png - python -m ops_model.models.interpretability.diffex.viewer.build_pcs_marker --marker "autophagosome_MAP1LC3B" + python -m ops_model.models.interpretability.diffae.viewer.build_pcs_marker --marker "autophagosome_MAP1LC3B" """ from __future__ import annotations diff --git a/src/ops_model/models/interpretability/diffex/viewer/build_phase_shap_rankings.py b/src/ops_model/models/interpretability/diffae/viewer/build_phase_shap_rankings.py similarity index 97% rename from src/ops_model/models/interpretability/diffex/viewer/build_phase_shap_rankings.py rename to src/ops_model/models/interpretability/diffae/viewer/build_phase_shap_rankings.py index d55aaee..acff805 100644 --- a/src/ops_model/models/interpretability/diffex/viewer/build_phase_shap_rankings.py +++ b/src/ops_model/models/interpretability/diffae/viewer/build_phase_shap_rankings.py @@ -6,8 +6,8 @@ pma_attention, rank, rank_type) complex → pma_shap_phase_complex.parquet (adds predicted_class=complex, gene=member gene; EBI-pooled) - python -m ops_model.models.interpretability.diffex.viewer.build_phase_shap_rankings --geneko --submit - python -m ops_model.models.interpretability.diffex.viewer.build_phase_shap_rankings --complex --submit + python -m ops_model.models.interpretability.diffae.viewer.build_phase_shap_rankings --geneko --submit + python -m ops_model.models.interpretability.diffae.viewer.build_phase_shap_rankings --complex --submit """ from __future__ import annotations diff --git a/src/ops_model/models/interpretability/diffex/viewer/build_phate_figure.py b/src/ops_model/models/interpretability/diffae/viewer/build_phate_figure.py similarity index 99% rename from src/ops_model/models/interpretability/diffex/viewer/build_phate_figure.py rename to src/ops_model/models/interpretability/diffae/viewer/build_phate_figure.py index 648ffd1..0ae77f3 100644 --- a/src/ops_model/models/interpretability/diffex/viewer/build_phate_figure.py +++ b/src/ops_model/models/interpretability/diffae/viewer/build_phate_figure.py @@ -5,7 +5,7 @@ Each panel: the same PHATE scatter (grey), that panel's groups colored + leader-labelled with the single-cell generated morph (NTC cell1 → group, alpha=+5). NTC original shown top-left of panel E. - python -m ops_model.models.interpretability.diffex.viewer.build_phate_figure + python -m ops_model.models.interpretability.diffae.viewer.build_phate_figure """ from __future__ import annotations diff --git a/src/ops_model/models/interpretability/diffex/viewer/build_setacc_bins.py b/src/ops_model/models/interpretability/diffae/viewer/build_setacc_bins.py similarity index 100% rename from src/ops_model/models/interpretability/diffex/viewer/build_setacc_bins.py rename to src/ops_model/models/interpretability/diffae/viewer/build_setacc_bins.py diff --git a/src/ops_model/models/interpretability/diffex/viewer/build_setacc_bymarker.py b/src/ops_model/models/interpretability/diffae/viewer/build_setacc_bymarker.py similarity index 100% rename from src/ops_model/models/interpretability/diffex/viewer/build_setacc_bymarker.py rename to src/ops_model/models/interpretability/diffae/viewer/build_setacc_bymarker.py diff --git a/src/ops_model/models/interpretability/diffex/viewer/build_top_cells.py b/src/ops_model/models/interpretability/diffae/viewer/build_top_cells.py similarity index 98% rename from src/ops_model/models/interpretability/diffex/viewer/build_top_cells.py rename to src/ops_model/models/interpretability/diffae/viewer/build_top_cells.py index 843fefd..c61df19 100644 --- a/src/ops_model/models/interpretability/diffex/viewer/build_top_cells.py +++ b/src/ops_model/models/interpretability/diffae/viewer/build_top_cells.py @@ -7,8 +7,8 @@ viewer_assets_v5/top_cells/index.json {"top_n", "genes"|"complexes": {CLASS: {"accuracy": [rec...]}}} viewer_assets_v5/top_cells/crops/.png - python -m ops_model.models.interpretability.diffex.viewer.build_top_cells geneKO # SLURM crop shards + finalize - python -m ops_model.models.interpretability.diffex.viewer.build_top_cells complex --finalize # rebuild index only + python -m ops_model.models.interpretability.diffae.viewer.build_top_cells geneKO # SLURM crop shards + finalize + python -m ops_model.models.interpretability.diffae.viewer.build_top_cells complex --finalize # rebuild index only """ from __future__ import annotations diff --git a/src/ops_model/models/interpretability/diffex/viewer/build_umap_montage.py b/src/ops_model/models/interpretability/diffae/viewer/build_umap_montage.py similarity index 100% rename from src/ops_model/models/interpretability/diffex/viewer/build_umap_montage.py rename to src/ops_model/models/interpretability/diffae/viewer/build_umap_montage.py diff --git a/src/ops_model/models/interpretability/diffex/viewer/catalog.py b/src/ops_model/models/interpretability/diffae/viewer/catalog.py similarity index 100% rename from src/ops_model/models/interpretability/diffex/viewer/catalog.py rename to src/ops_model/models/interpretability/diffae/viewer/catalog.py diff --git a/src/ops_model/models/interpretability/diffex/viewer/deploy/README.md b/src/ops_model/models/interpretability/diffae/viewer/deploy/README.md similarity index 100% rename from src/ops_model/models/interpretability/diffex/viewer/deploy/README.md rename to src/ops_model/models/interpretability/diffae/viewer/deploy/README.md diff --git a/src/ops_model/models/interpretability/diffex/viewer/marker_leaves.py b/src/ops_model/models/interpretability/diffae/viewer/marker_leaves.py similarity index 100% rename from src/ops_model/models/interpretability/diffex/viewer/marker_leaves.py rename to src/ops_model/models/interpretability/diffae/viewer/marker_leaves.py diff --git a/src/ops_model/models/interpretability/diffex/viewer/mimic_alex_embed.py b/src/ops_model/models/interpretability/diffae/viewer/mimic_alex_embed.py similarity index 100% rename from src/ops_model/models/interpretability/diffex/viewer/mimic_alex_embed.py rename to src/ops_model/models/interpretability/diffae/viewer/mimic_alex_embed.py diff --git a/src/ops_model/models/interpretability/diffex/viewer/morpho_pipeline.py b/src/ops_model/models/interpretability/diffae/viewer/morpho_pipeline.py similarity index 99% rename from src/ops_model/models/interpretability/diffex/viewer/morpho_pipeline.py rename to src/ops_model/models/interpretability/diffae/viewer/morpho_pipeline.py index 512d946..e5addf4 100644 --- a/src/ops_model/models/interpretability/diffex/viewer/morpho_pipeline.py +++ b/src/ops_model/models/interpretability/diffae/viewer/morpho_pipeline.py @@ -100,8 +100,8 @@ def _vs_h2b_nucleus_npz(marker_dir, target, grain, out_npz, n_cells, force=False from PIL import Image if os.path.exists(out_npz) and not force: print(f"[vs-nuc] cache {out_npz}"); return out_npz - from ..diffae.config import DiffAEConfig - from ..diffae.model import DiffAE + from ..generator.config import DiffAEConfig + from ..generator.model import DiffAE from ..classifier.celldino_features import embed_crops from diffusers import DDIMScheduler from cellpose import models as cpm diff --git a/src/ops_model/models/interpretability/diffex/viewer/morphometrics.py b/src/ops_model/models/interpretability/diffae/viewer/morphometrics.py similarity index 100% rename from src/ops_model/models/interpretability/diffex/viewer/morphometrics.py rename to src/ops_model/models/interpretability/diffae/viewer/morphometrics.py diff --git a/src/ops_model/models/interpretability/diffex/viewer/nway_clf.py b/src/ops_model/models/interpretability/diffae/viewer/nway_clf.py similarity index 100% rename from src/ops_model/models/interpretability/diffex/viewer/nway_clf.py rename to src/ops_model/models/interpretability/diffae/viewer/nway_clf.py diff --git a/src/ops_model/models/interpretability/diffex/viewer/phenotype_cells.py b/src/ops_model/models/interpretability/diffae/viewer/phenotype_cells.py similarity index 100% rename from src/ops_model/models/interpretability/diffex/viewer/phenotype_cells.py rename to src/ops_model/models/interpretability/diffae/viewer/phenotype_cells.py diff --git a/src/ops_model/models/interpretability/diffex/viewer/precompute.py b/src/ops_model/models/interpretability/diffae/viewer/precompute.py similarity index 99% rename from src/ops_model/models/interpretability/diffex/viewer/precompute.py rename to src/ops_model/models/interpretability/diffae/viewer/precompute.py index 4e89a8f..4b37968 100644 --- a/src/ops_model/models/interpretability/diffex/viewer/precompute.py +++ b/src/ops_model/models/interpretability/diffae/viewer/precompute.py @@ -34,7 +34,7 @@ from ..classifier.celldino_features import embed_crops from ..classifier.config import GRAINS, slugify from ..classifier.data import _BASE_COLS, make_labels_df, materialize_crops -from ..diffae.data import normalize +from ..generator.data import normalize from ..directions.config import DirConfig from ..directions.data import _top_cells from ..directions.make_gifs import _pair_slug, _setup, _sample_guided diff --git a/src/ops_model/models/interpretability/diffex/viewer/render_montage_scales.py b/src/ops_model/models/interpretability/diffae/viewer/render_montage_scales.py similarity index 99% rename from src/ops_model/models/interpretability/diffex/viewer/render_montage_scales.py rename to src/ops_model/models/interpretability/diffae/viewer/render_montage_scales.py index eb708ad..eb725c8 100644 --- a/src/ops_model/models/interpretability/diffex/viewer/render_montage_scales.py +++ b/src/ops_model/models/interpretability/diffae/viewer/render_montage_scales.py @@ -2,7 +2,7 @@ embedding legend (leiden_r4, big dots, NTC as a dark labelled circle). The montage image and its baked viewer-style gene names come straight from the built tiles — finer levels give crisper text. - python -m ops_model.models.interpretability.diffex.viewer.render_montage_scales --alphas 1-5 --levels 3,4 + python -m ops_model.models.interpretability.diffae.viewer.render_montage_scales --alphas 1-5 --levels 3,4 Each level of `_montage/phase_geneKO_phate_cell1_a_tiles/L/` is a level-of-detail montage (coarse levels show a decimated non-overlapping subset; finer levels fill in more cells at higher res). diff --git a/src/ops_model/models/interpretability/diffex/viewer/score_generated.py b/src/ops_model/models/interpretability/diffae/viewer/score_generated.py similarity index 100% rename from src/ops_model/models/interpretability/diffex/viewer/score_generated.py rename to src/ops_model/models/interpretability/diffae/viewer/score_generated.py diff --git a/src/ops_model/models/interpretability/diffex/viewer/set_classifier.py b/src/ops_model/models/interpretability/diffae/viewer/set_classifier.py similarity index 100% rename from src/ops_model/models/interpretability/diffex/viewer/set_classifier.py rename to src/ops_model/models/interpretability/diffae/viewer/set_classifier.py diff --git a/src/ops_model/models/interpretability/diffex/viewer/submit.py b/src/ops_model/models/interpretability/diffae/viewer/submit.py similarity index 98% rename from src/ops_model/models/interpretability/diffex/viewer/submit.py rename to src/ops_model/models/interpretability/diffae/viewer/submit.py index 699eac2..eee7314 100644 --- a/src/ops_model/models/interpretability/diffex/viewer/submit.py +++ b/src/ops_model/models/interpretability/diffae/viewer/submit.py @@ -1,10 +1,10 @@ """Build the DiffEx viewer cache — reproducible, version-controlled entrypoint (replaces the one-off scratchpad drivers). All target selection comes from `catalog.py`. - python -m ops_model.models.interpretability.diffex.viewer.submit seed # per-marker NTC traversals - python -m ops_model.models.interpretability.diffex.viewer.submit anchors --k 5 # A→B anchor pairs - python -m ops_model.models.interpretability.diffex.viewer.submit manifest # rebuild manifest.json (local) - python -m ops_model.models.interpretability.diffex.viewer.submit montage --cell 0 --alpha 2 # harvest cache -> UMAP montage zarr + python -m ops_model.models.interpretability.diffae.viewer.submit seed # per-marker NTC traversals + python -m ops_model.models.interpretability.diffae.viewer.submit anchors --k 5 # A→B anchor pairs + python -m ops_model.models.interpretability.diffae.viewer.submit manifest # rebuild manifest.json (local) + python -m ops_model.models.interpretability.diffae.viewer.submit montage --cell 0 --alpha 2 # harvest cache -> UMAP montage zarr """ from __future__ import annotations diff --git a/src/ops_model/models/interpretability/diffex/viewer/webapp/app.js b/src/ops_model/models/interpretability/diffae/viewer/webapp/app.js similarity index 100% rename from src/ops_model/models/interpretability/diffex/viewer/webapp/app.js rename to src/ops_model/models/interpretability/diffae/viewer/webapp/app.js diff --git a/src/ops_model/models/interpretability/diffex/viewer/webapp/biohub-mark.png b/src/ops_model/models/interpretability/diffae/viewer/webapp/biohub-mark.png similarity index 100% rename from src/ops_model/models/interpretability/diffex/viewer/webapp/biohub-mark.png rename to src/ops_model/models/interpretability/diffae/viewer/webapp/biohub-mark.png diff --git a/src/ops_model/models/interpretability/diffex/viewer/webapp/biohub-wordmark.png b/src/ops_model/models/interpretability/diffae/viewer/webapp/biohub-wordmark.png similarity index 100% rename from src/ops_model/models/interpretability/diffex/viewer/webapp/biohub-wordmark.png rename to src/ops_model/models/interpretability/diffae/viewer/webapp/biohub-wordmark.png diff --git a/src/ops_model/models/interpretability/diffex/viewer/webapp/build_gene_narratives.py b/src/ops_model/models/interpretability/diffae/viewer/webapp/build_gene_narratives.py similarity index 100% rename from src/ops_model/models/interpretability/diffex/viewer/webapp/build_gene_narratives.py rename to src/ops_model/models/interpretability/diffae/viewer/webapp/build_gene_narratives.py diff --git a/src/ops_model/models/interpretability/diffex/viewer/webapp/gif.js b/src/ops_model/models/interpretability/diffae/viewer/webapp/gif.js similarity index 100% rename from src/ops_model/models/interpretability/diffex/viewer/webapp/gif.js rename to src/ops_model/models/interpretability/diffae/viewer/webapp/gif.js diff --git a/src/ops_model/models/interpretability/diffex/viewer/webapp/gif.worker.js b/src/ops_model/models/interpretability/diffae/viewer/webapp/gif.worker.js similarity index 100% rename from src/ops_model/models/interpretability/diffex/viewer/webapp/gif.worker.js rename to src/ops_model/models/interpretability/diffae/viewer/webapp/gif.worker.js diff --git a/src/ops_model/models/interpretability/diffex/viewer/webapp/index.html b/src/ops_model/models/interpretability/diffae/viewer/webapp/index.html similarity index 100% rename from src/ops_model/models/interpretability/diffex/viewer/webapp/index.html rename to src/ops_model/models/interpretability/diffae/viewer/webapp/index.html diff --git a/src/ops_model/models/interpretability/diffex/viewer/webapp/methods.js b/src/ops_model/models/interpretability/diffae/viewer/webapp/methods.js similarity index 100% rename from src/ops_model/models/interpretability/diffex/viewer/webapp/methods.js rename to src/ops_model/models/interpretability/diffae/viewer/webapp/methods.js diff --git a/src/ops_model/models/interpretability/diffex/viewer/webapp/morpho_demo.html b/src/ops_model/models/interpretability/diffae/viewer/webapp/morpho_demo.html similarity index 100% rename from src/ops_model/models/interpretability/diffex/viewer/webapp/morpho_demo.html rename to src/ops_model/models/interpretability/diffae/viewer/webapp/morpho_demo.html diff --git a/src/ops_model/models/interpretability/diffex/viewer/webapp/openseadragon.min.js b/src/ops_model/models/interpretability/diffae/viewer/webapp/openseadragon.min.js similarity index 100% rename from src/ops_model/models/interpretability/diffex/viewer/webapp/openseadragon.min.js rename to src/ops_model/models/interpretability/diffae/viewer/webapp/openseadragon.min.js diff --git a/src/ops_model/models/interpretability/diffex/viewer/webapp/opsin-eyes.svg b/src/ops_model/models/interpretability/diffae/viewer/webapp/opsin-eyes.svg similarity index 100% rename from src/ops_model/models/interpretability/diffex/viewer/webapp/opsin-eyes.svg rename to src/ops_model/models/interpretability/diffae/viewer/webapp/opsin-eyes.svg diff --git a/src/ops_model/models/interpretability/diffex/viewer/webapp/style.css b/src/ops_model/models/interpretability/diffae/viewer/webapp/style.css similarity index 100% rename from src/ops_model/models/interpretability/diffex/viewer/webapp/style.css rename to src/ops_model/models/interpretability/diffae/viewer/webapp/style.css From bbf77773553ab3f5363a8f3e898154ee913edee5 Mon Sep 17 00:00:00 2001 From: Gav Sturm Date: Tue, 11 Aug 2026 09:30:48 -0700 Subject: [PATCH 04/13] quarantine uncoupled analysis dirs into interpretability/_internal/ MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Moved (0 core dependencies, 0 residual refs): atlas, shap, titration, embedding, weighted_aggregation, diffae/kyle_pcs, and the SHAP RUNBOOK.md. Fixed titration's one docstring -m path. viewer + figures NOT yet moved — they're coupled (see next). --- .../models/interpretability/{ => _internal}/RUNBOOK.md | 0 .../atlas/attention_accuracy_umap_animation.py | 0 .../interpretability/{ => _internal}/atlas/attention_atlas.py | 0 .../{ => _internal}/atlas/attention_atlas_shap.py | 0 .../{ => _internal}/atlas/low_attention_phase_atlas.py | 0 .../interpretability/{ => _internal}/atlas/make_scale_bar.py | 0 .../{ => _internal}/atlas/marker_selection_distribution.py | 0 .../{ => _internal}/atlas/plot_eval_accuracy_curves.py | 0 .../{ => _internal}/embedding/generate_ko_violin_plots.py | 0 .../{ => _internal}/embedding/regen_umap_gav.py | 0 .../{ => _internal}/embedding/regen_umap_html.py | 0 .../{ => _internal}/embedding/run_all_atlases.py | 0 .../embedding/top_attention_embed_and_score.py | 0 .../{diffae => _internal}/kyle_pcs/build_static_explorer.py | 0 .../{diffae => _internal}/kyle_pcs/compute_pc_strips.py | 0 .../{ => _internal}/shap/analyze_chad_variants.py | 0 .../{ => _internal}/shap/generate_shap_captions_combined.py | 0 .../interpretability/{ => _internal}/shap/ko_shap_features.py | 0 .../{ => _internal}/shap/merge_shap_shards.py | 0 .../{ => _internal}/shap/ntc_attention_compare.py | 0 .../interpretability/{ => _internal}/shap/ntc_pick_cells.py | 0 .../{ => _internal}/shap/ntc_shap_features.py | 0 .../interpretability/{ => _internal}/shap/run_all_shap.py | 0 .../{ => _internal}/shap/run_shap_pipeline.py | 0 .../{ => _internal}/shap/shap_approach_compare.py | 0 .../{ => _internal}/titration/decay/map_attention_decay.py | 0 .../{ => _internal}/titration/decay/phate_peak_groups.py | 0 .../{ => _internal}/titration/decay/plot_3way_summary_bars.py | 0 .../titration/decay/plot_all_cells_correction_bars.py | 0 .../titration/expansion/count_genes_above_threshold.py | 4 ++-- .../titration/expansion/map_attention_expansion_v4.py | 0 .../titration/expansion/plot_sgrna_coverage_sweep.py | 0 .../titration/expansion/run_percentile_sweep.py | 0 .../{ => _internal}/weighted_aggregation/_v4_attn_worker.py | 0 .../weighted_aggregation/analyze_v3_acc_bins.py | 0 .../weighted_aggregation/plot_v4_attn_comparison.py | 0 .../run_v3_pipeline_on_v4_attn_weighted.py | 0 .../weighted_aggregation/run_v3_pipeline_on_v4_features.py | 0 38 files changed, 2 insertions(+), 2 deletions(-) rename src/ops_model/models/interpretability/{ => _internal}/RUNBOOK.md (100%) rename src/ops_model/models/interpretability/{ => _internal}/atlas/attention_accuracy_umap_animation.py (100%) rename src/ops_model/models/interpretability/{ => _internal}/atlas/attention_atlas.py (100%) rename src/ops_model/models/interpretability/{ => _internal}/atlas/attention_atlas_shap.py (100%) rename src/ops_model/models/interpretability/{ => _internal}/atlas/low_attention_phase_atlas.py (100%) rename src/ops_model/models/interpretability/{ => _internal}/atlas/make_scale_bar.py (100%) rename src/ops_model/models/interpretability/{ => _internal}/atlas/marker_selection_distribution.py (100%) rename src/ops_model/models/interpretability/{ => _internal}/atlas/plot_eval_accuracy_curves.py (100%) rename src/ops_model/models/interpretability/{ => _internal}/embedding/generate_ko_violin_plots.py (100%) rename src/ops_model/models/interpretability/{ => _internal}/embedding/regen_umap_gav.py (100%) rename src/ops_model/models/interpretability/{ => _internal}/embedding/regen_umap_html.py (100%) rename src/ops_model/models/interpretability/{ => _internal}/embedding/run_all_atlases.py (100%) rename src/ops_model/models/interpretability/{ => _internal}/embedding/top_attention_embed_and_score.py (100%) rename src/ops_model/models/interpretability/{diffae => _internal}/kyle_pcs/build_static_explorer.py (100%) rename src/ops_model/models/interpretability/{diffae => _internal}/kyle_pcs/compute_pc_strips.py (100%) rename src/ops_model/models/interpretability/{ => _internal}/shap/analyze_chad_variants.py (100%) rename src/ops_model/models/interpretability/{ => _internal}/shap/generate_shap_captions_combined.py (100%) rename src/ops_model/models/interpretability/{ => _internal}/shap/ko_shap_features.py (100%) rename src/ops_model/models/interpretability/{ => _internal}/shap/merge_shap_shards.py (100%) rename src/ops_model/models/interpretability/{ => _internal}/shap/ntc_attention_compare.py (100%) rename src/ops_model/models/interpretability/{ => _internal}/shap/ntc_pick_cells.py (100%) rename src/ops_model/models/interpretability/{ => _internal}/shap/ntc_shap_features.py (100%) rename src/ops_model/models/interpretability/{ => _internal}/shap/run_all_shap.py (100%) rename src/ops_model/models/interpretability/{ => _internal}/shap/run_shap_pipeline.py (100%) rename src/ops_model/models/interpretability/{ => _internal}/shap/shap_approach_compare.py (100%) rename src/ops_model/models/interpretability/{ => _internal}/titration/decay/map_attention_decay.py (100%) rename src/ops_model/models/interpretability/{ => _internal}/titration/decay/phate_peak_groups.py (100%) rename src/ops_model/models/interpretability/{ => _internal}/titration/decay/plot_3way_summary_bars.py (100%) rename src/ops_model/models/interpretability/{ => _internal}/titration/decay/plot_all_cells_correction_bars.py (100%) rename src/ops_model/models/interpretability/{ => _internal}/titration/expansion/count_genes_above_threshold.py (98%) rename src/ops_model/models/interpretability/{ => _internal}/titration/expansion/map_attention_expansion_v4.py (100%) rename src/ops_model/models/interpretability/{ => _internal}/titration/expansion/plot_sgrna_coverage_sweep.py (100%) rename src/ops_model/models/interpretability/{ => _internal}/titration/expansion/run_percentile_sweep.py (100%) rename src/ops_model/models/interpretability/{ => _internal}/weighted_aggregation/_v4_attn_worker.py (100%) rename src/ops_model/models/interpretability/{ => _internal}/weighted_aggregation/analyze_v3_acc_bins.py (100%) rename src/ops_model/models/interpretability/{ => _internal}/weighted_aggregation/plot_v4_attn_comparison.py (100%) rename src/ops_model/models/interpretability/{ => _internal}/weighted_aggregation/run_v3_pipeline_on_v4_attn_weighted.py (100%) rename src/ops_model/models/interpretability/{ => _internal}/weighted_aggregation/run_v3_pipeline_on_v4_features.py (100%) diff --git a/src/ops_model/models/interpretability/RUNBOOK.md b/src/ops_model/models/interpretability/_internal/RUNBOOK.md similarity index 100% rename from src/ops_model/models/interpretability/RUNBOOK.md rename to src/ops_model/models/interpretability/_internal/RUNBOOK.md diff --git a/src/ops_model/models/interpretability/atlas/attention_accuracy_umap_animation.py b/src/ops_model/models/interpretability/_internal/atlas/attention_accuracy_umap_animation.py similarity index 100% rename from src/ops_model/models/interpretability/atlas/attention_accuracy_umap_animation.py rename to src/ops_model/models/interpretability/_internal/atlas/attention_accuracy_umap_animation.py diff --git a/src/ops_model/models/interpretability/atlas/attention_atlas.py b/src/ops_model/models/interpretability/_internal/atlas/attention_atlas.py similarity index 100% rename from src/ops_model/models/interpretability/atlas/attention_atlas.py rename to src/ops_model/models/interpretability/_internal/atlas/attention_atlas.py diff --git a/src/ops_model/models/interpretability/atlas/attention_atlas_shap.py b/src/ops_model/models/interpretability/_internal/atlas/attention_atlas_shap.py similarity index 100% rename from src/ops_model/models/interpretability/atlas/attention_atlas_shap.py rename to src/ops_model/models/interpretability/_internal/atlas/attention_atlas_shap.py diff --git a/src/ops_model/models/interpretability/atlas/low_attention_phase_atlas.py b/src/ops_model/models/interpretability/_internal/atlas/low_attention_phase_atlas.py similarity index 100% rename from src/ops_model/models/interpretability/atlas/low_attention_phase_atlas.py rename to src/ops_model/models/interpretability/_internal/atlas/low_attention_phase_atlas.py diff --git a/src/ops_model/models/interpretability/atlas/make_scale_bar.py b/src/ops_model/models/interpretability/_internal/atlas/make_scale_bar.py similarity index 100% rename from src/ops_model/models/interpretability/atlas/make_scale_bar.py rename to src/ops_model/models/interpretability/_internal/atlas/make_scale_bar.py diff --git a/src/ops_model/models/interpretability/atlas/marker_selection_distribution.py b/src/ops_model/models/interpretability/_internal/atlas/marker_selection_distribution.py similarity index 100% rename from src/ops_model/models/interpretability/atlas/marker_selection_distribution.py rename to src/ops_model/models/interpretability/_internal/atlas/marker_selection_distribution.py diff --git a/src/ops_model/models/interpretability/atlas/plot_eval_accuracy_curves.py b/src/ops_model/models/interpretability/_internal/atlas/plot_eval_accuracy_curves.py similarity index 100% rename from src/ops_model/models/interpretability/atlas/plot_eval_accuracy_curves.py rename to src/ops_model/models/interpretability/_internal/atlas/plot_eval_accuracy_curves.py diff --git a/src/ops_model/models/interpretability/embedding/generate_ko_violin_plots.py b/src/ops_model/models/interpretability/_internal/embedding/generate_ko_violin_plots.py similarity index 100% rename from src/ops_model/models/interpretability/embedding/generate_ko_violin_plots.py rename to src/ops_model/models/interpretability/_internal/embedding/generate_ko_violin_plots.py diff --git a/src/ops_model/models/interpretability/embedding/regen_umap_gav.py b/src/ops_model/models/interpretability/_internal/embedding/regen_umap_gav.py similarity index 100% rename from src/ops_model/models/interpretability/embedding/regen_umap_gav.py rename to src/ops_model/models/interpretability/_internal/embedding/regen_umap_gav.py diff --git a/src/ops_model/models/interpretability/embedding/regen_umap_html.py b/src/ops_model/models/interpretability/_internal/embedding/regen_umap_html.py similarity index 100% rename from src/ops_model/models/interpretability/embedding/regen_umap_html.py rename to src/ops_model/models/interpretability/_internal/embedding/regen_umap_html.py diff --git a/src/ops_model/models/interpretability/embedding/run_all_atlases.py b/src/ops_model/models/interpretability/_internal/embedding/run_all_atlases.py similarity index 100% rename from src/ops_model/models/interpretability/embedding/run_all_atlases.py rename to src/ops_model/models/interpretability/_internal/embedding/run_all_atlases.py diff --git a/src/ops_model/models/interpretability/embedding/top_attention_embed_and_score.py b/src/ops_model/models/interpretability/_internal/embedding/top_attention_embed_and_score.py similarity index 100% rename from src/ops_model/models/interpretability/embedding/top_attention_embed_and_score.py rename to src/ops_model/models/interpretability/_internal/embedding/top_attention_embed_and_score.py diff --git a/src/ops_model/models/interpretability/diffae/kyle_pcs/build_static_explorer.py b/src/ops_model/models/interpretability/_internal/kyle_pcs/build_static_explorer.py similarity index 100% rename from src/ops_model/models/interpretability/diffae/kyle_pcs/build_static_explorer.py rename to src/ops_model/models/interpretability/_internal/kyle_pcs/build_static_explorer.py diff --git a/src/ops_model/models/interpretability/diffae/kyle_pcs/compute_pc_strips.py b/src/ops_model/models/interpretability/_internal/kyle_pcs/compute_pc_strips.py similarity index 100% rename from src/ops_model/models/interpretability/diffae/kyle_pcs/compute_pc_strips.py rename to src/ops_model/models/interpretability/_internal/kyle_pcs/compute_pc_strips.py diff --git a/src/ops_model/models/interpretability/shap/analyze_chad_variants.py b/src/ops_model/models/interpretability/_internal/shap/analyze_chad_variants.py similarity index 100% rename from src/ops_model/models/interpretability/shap/analyze_chad_variants.py rename to src/ops_model/models/interpretability/_internal/shap/analyze_chad_variants.py diff --git a/src/ops_model/models/interpretability/shap/generate_shap_captions_combined.py b/src/ops_model/models/interpretability/_internal/shap/generate_shap_captions_combined.py similarity index 100% rename from src/ops_model/models/interpretability/shap/generate_shap_captions_combined.py rename to src/ops_model/models/interpretability/_internal/shap/generate_shap_captions_combined.py diff --git a/src/ops_model/models/interpretability/shap/ko_shap_features.py b/src/ops_model/models/interpretability/_internal/shap/ko_shap_features.py similarity index 100% rename from src/ops_model/models/interpretability/shap/ko_shap_features.py rename to src/ops_model/models/interpretability/_internal/shap/ko_shap_features.py diff --git a/src/ops_model/models/interpretability/shap/merge_shap_shards.py b/src/ops_model/models/interpretability/_internal/shap/merge_shap_shards.py similarity index 100% rename from src/ops_model/models/interpretability/shap/merge_shap_shards.py rename to src/ops_model/models/interpretability/_internal/shap/merge_shap_shards.py diff --git a/src/ops_model/models/interpretability/shap/ntc_attention_compare.py b/src/ops_model/models/interpretability/_internal/shap/ntc_attention_compare.py similarity index 100% rename from src/ops_model/models/interpretability/shap/ntc_attention_compare.py rename to src/ops_model/models/interpretability/_internal/shap/ntc_attention_compare.py diff --git a/src/ops_model/models/interpretability/shap/ntc_pick_cells.py b/src/ops_model/models/interpretability/_internal/shap/ntc_pick_cells.py similarity index 100% rename from src/ops_model/models/interpretability/shap/ntc_pick_cells.py rename to src/ops_model/models/interpretability/_internal/shap/ntc_pick_cells.py diff --git a/src/ops_model/models/interpretability/shap/ntc_shap_features.py b/src/ops_model/models/interpretability/_internal/shap/ntc_shap_features.py similarity index 100% rename from src/ops_model/models/interpretability/shap/ntc_shap_features.py rename to src/ops_model/models/interpretability/_internal/shap/ntc_shap_features.py diff --git a/src/ops_model/models/interpretability/shap/run_all_shap.py b/src/ops_model/models/interpretability/_internal/shap/run_all_shap.py similarity index 100% rename from src/ops_model/models/interpretability/shap/run_all_shap.py rename to src/ops_model/models/interpretability/_internal/shap/run_all_shap.py diff --git a/src/ops_model/models/interpretability/shap/run_shap_pipeline.py b/src/ops_model/models/interpretability/_internal/shap/run_shap_pipeline.py similarity index 100% rename from src/ops_model/models/interpretability/shap/run_shap_pipeline.py rename to src/ops_model/models/interpretability/_internal/shap/run_shap_pipeline.py diff --git a/src/ops_model/models/interpretability/shap/shap_approach_compare.py b/src/ops_model/models/interpretability/_internal/shap/shap_approach_compare.py similarity index 100% rename from src/ops_model/models/interpretability/shap/shap_approach_compare.py rename to src/ops_model/models/interpretability/_internal/shap/shap_approach_compare.py diff --git a/src/ops_model/models/interpretability/titration/decay/map_attention_decay.py b/src/ops_model/models/interpretability/_internal/titration/decay/map_attention_decay.py similarity index 100% rename from src/ops_model/models/interpretability/titration/decay/map_attention_decay.py rename to src/ops_model/models/interpretability/_internal/titration/decay/map_attention_decay.py diff --git a/src/ops_model/models/interpretability/titration/decay/phate_peak_groups.py b/src/ops_model/models/interpretability/_internal/titration/decay/phate_peak_groups.py similarity index 100% rename from src/ops_model/models/interpretability/titration/decay/phate_peak_groups.py rename to src/ops_model/models/interpretability/_internal/titration/decay/phate_peak_groups.py diff --git a/src/ops_model/models/interpretability/titration/decay/plot_3way_summary_bars.py b/src/ops_model/models/interpretability/_internal/titration/decay/plot_3way_summary_bars.py similarity index 100% rename from src/ops_model/models/interpretability/titration/decay/plot_3way_summary_bars.py rename to src/ops_model/models/interpretability/_internal/titration/decay/plot_3way_summary_bars.py diff --git a/src/ops_model/models/interpretability/titration/decay/plot_all_cells_correction_bars.py b/src/ops_model/models/interpretability/_internal/titration/decay/plot_all_cells_correction_bars.py similarity index 100% rename from src/ops_model/models/interpretability/titration/decay/plot_all_cells_correction_bars.py rename to src/ops_model/models/interpretability/_internal/titration/decay/plot_all_cells_correction_bars.py diff --git a/src/ops_model/models/interpretability/titration/expansion/count_genes_above_threshold.py b/src/ops_model/models/interpretability/_internal/titration/expansion/count_genes_above_threshold.py similarity index 98% rename from src/ops_model/models/interpretability/titration/expansion/count_genes_above_threshold.py rename to src/ops_model/models/interpretability/_internal/titration/expansion/count_genes_above_threshold.py index f402319..e6e1b45 100644 --- a/src/ops_model/models/interpretability/titration/expansion/count_genes_above_threshold.py +++ b/src/ops_model/models/interpretability/_internal/titration/expansion/count_genes_above_threshold.py @@ -17,10 +17,10 @@ Usage:: # Submit one SLURM task per K (9 tasks; ~5 min wall once they land) - uv run python -m ops_model.models.interpretability.titration.expansion.count_genes_above_threshold --slurm + uv run python -m ops_model.models.interpretability._internal.titration.expansion.count_genes_above_threshold --slurm # Replot from cached per-gene CSVs (no SLURM) - uv run python -m ops_model.models.interpretability.titration.expansion.count_genes_above_threshold --replot + uv run python -m ops_model.models.interpretability._internal.titration.expansion.count_genes_above_threshold --replot """ from __future__ import annotations diff --git a/src/ops_model/models/interpretability/titration/expansion/map_attention_expansion_v4.py b/src/ops_model/models/interpretability/_internal/titration/expansion/map_attention_expansion_v4.py similarity index 100% rename from src/ops_model/models/interpretability/titration/expansion/map_attention_expansion_v4.py rename to src/ops_model/models/interpretability/_internal/titration/expansion/map_attention_expansion_v4.py diff --git a/src/ops_model/models/interpretability/titration/expansion/plot_sgrna_coverage_sweep.py b/src/ops_model/models/interpretability/_internal/titration/expansion/plot_sgrna_coverage_sweep.py similarity index 100% rename from src/ops_model/models/interpretability/titration/expansion/plot_sgrna_coverage_sweep.py rename to src/ops_model/models/interpretability/_internal/titration/expansion/plot_sgrna_coverage_sweep.py diff --git a/src/ops_model/models/interpretability/titration/expansion/run_percentile_sweep.py b/src/ops_model/models/interpretability/_internal/titration/expansion/run_percentile_sweep.py similarity index 100% rename from src/ops_model/models/interpretability/titration/expansion/run_percentile_sweep.py rename to src/ops_model/models/interpretability/_internal/titration/expansion/run_percentile_sweep.py diff --git a/src/ops_model/models/interpretability/weighted_aggregation/_v4_attn_worker.py b/src/ops_model/models/interpretability/_internal/weighted_aggregation/_v4_attn_worker.py similarity index 100% rename from src/ops_model/models/interpretability/weighted_aggregation/_v4_attn_worker.py rename to src/ops_model/models/interpretability/_internal/weighted_aggregation/_v4_attn_worker.py diff --git a/src/ops_model/models/interpretability/weighted_aggregation/analyze_v3_acc_bins.py b/src/ops_model/models/interpretability/_internal/weighted_aggregation/analyze_v3_acc_bins.py similarity index 100% rename from src/ops_model/models/interpretability/weighted_aggregation/analyze_v3_acc_bins.py rename to src/ops_model/models/interpretability/_internal/weighted_aggregation/analyze_v3_acc_bins.py diff --git a/src/ops_model/models/interpretability/weighted_aggregation/plot_v4_attn_comparison.py b/src/ops_model/models/interpretability/_internal/weighted_aggregation/plot_v4_attn_comparison.py similarity index 100% rename from src/ops_model/models/interpretability/weighted_aggregation/plot_v4_attn_comparison.py rename to src/ops_model/models/interpretability/_internal/weighted_aggregation/plot_v4_attn_comparison.py diff --git a/src/ops_model/models/interpretability/weighted_aggregation/run_v3_pipeline_on_v4_attn_weighted.py b/src/ops_model/models/interpretability/_internal/weighted_aggregation/run_v3_pipeline_on_v4_attn_weighted.py similarity index 100% rename from src/ops_model/models/interpretability/weighted_aggregation/run_v3_pipeline_on_v4_attn_weighted.py rename to src/ops_model/models/interpretability/_internal/weighted_aggregation/run_v3_pipeline_on_v4_attn_weighted.py diff --git a/src/ops_model/models/interpretability/weighted_aggregation/run_v3_pipeline_on_v4_features.py b/src/ops_model/models/interpretability/_internal/weighted_aggregation/run_v3_pipeline_on_v4_features.py similarity index 100% rename from src/ops_model/models/interpretability/weighted_aggregation/run_v3_pipeline_on_v4_features.py rename to src/ops_model/models/interpretability/_internal/weighted_aggregation/run_v3_pipeline_on_v4_features.py From 7e5e3c5df0a62b3ada7d75be7f50ef7d97dec879 Mon Sep 17 00:00:00 2001 From: Gav Sturm Date: Tue, 11 Aug 2026 09:41:24 -0700 Subject: [PATCH 05/13] split viewer (option C): promote 5 shared modules to diffae/traversal/, viewer -> _internal Core figures depend on 5 viewer modules (catalog, precompute, morpho_pipeline, _fluor_topcells, build_pc_crops_masked) that form a closed set over the core stages. Promoted them to diffae/traversal/ (same depth -> their ..classifier/etc. relatives survive). Remaining ~43 viewer files (webapp/deploy/one-offs/diagnostics) -> _internal/viewer/, with 83 cross-pkg relatives + 24 relatives-to-the-5 rewritten to absolute imports. Core stages' 2 lazy ..viewer.<5> refs -> ..traversal. Import-verified: core stages, traversal (all 5), and figures/_setacc_common. --- .../{diffae => _internal}/viewer/__init__.py | 0 .../viewer/_altanchor_build.py | 2 +- .../viewer/_anchortest.py | 2 +- .../viewer/_build_stepablation.py | 2 +- .../viewer/_build_v5_inverted.py | 30 ++++++++-------- .../viewer/_build_v5_montages.py | 4 +-- .../viewer/_build_valid200.py | 16 ++++----- .../viewer/_consolidate_cells.py | 0 .../viewer/_fluor_complex_build.py | 4 +-- .../viewer/_fluor_v5_build.py | 2 +- .../viewer/_migrate_v4_to_v5.py | 0 .../{diffae => _internal}/viewer/_phase_vs.py | 10 +++--- .../viewer/_rebuild_v5.py | 8 ++--- .../viewer/_rescore_rank.py | 2 +- .../{diffae => _internal}/viewer/_score_v4.py | 6 ++-- .../viewer/_v4acc_test.py | 4 +-- .../viewer/_verify_pt_space.py | 8 ++--- .../viewer/_verify_score_bridge.py | 4 +-- .../viewer/altanchor_pairs.json | 0 .../viewer/anchor_cells.py | 4 +-- .../viewer/build_attention_heads.py | 10 +++--- .../viewer/build_complex_ebi_map.py | 0 .../viewer/build_fluor_shap_rankings.py | 6 ++-- .../viewer/build_montage_features.py | 2 +- .../viewer/build_pc_features.py | 2 +- .../viewer/build_pc_walks.py | 10 +++--- .../{diffae => _internal}/viewer/build_pcs.py | 4 +-- .../viewer/build_pcs_marker.py | 4 +-- .../viewer/build_phase_shap_rankings.py | 4 +-- .../viewer/build_phate_figure.py | 2 +- .../viewer/build_setacc_bins.py | 0 .../viewer/build_setacc_bymarker.py | 0 .../viewer/build_top_cells.py | 6 ++-- .../viewer/build_umap_montage.py | 4 +-- .../viewer/deploy/README.md | 0 .../viewer/marker_leaves.py | 0 .../viewer/mimic_alex_embed.py | 0 .../viewer/morphometrics.py | 0 .../{diffae => _internal}/viewer/nway_clf.py | 8 ++--- .../viewer/phenotype_cells.py | 4 +-- .../viewer/render_montage_scales.py | 4 +-- .../viewer/score_generated.py | 34 +++++++++--------- .../viewer/set_classifier.py | 0 .../{diffae => _internal}/viewer/submit.py | 14 ++++---- .../viewer/webapp/app.js | 0 .../viewer/webapp/biohub-mark.png | Bin .../viewer/webapp/biohub-wordmark.png | Bin .../viewer/webapp/build_gene_narratives.py | 0 .../viewer/webapp/gif.js | 0 .../viewer/webapp/gif.worker.js | 0 .../viewer/webapp/index.html | 0 .../viewer/webapp/methods.js | 0 .../viewer/webapp/morpho_demo.html | 0 .../viewer/webapp/openseadragon.min.js | 0 .../viewer/webapp/opsin-eyes.svg | 0 .../viewer/webapp/style.css | 0 .../diffae/directions/proto_ddim_anchors.py | 2 +- .../diffae/figures/_setacc_common.py | 4 +-- .../diffae/figures/auto_pick_and_plot.py | 2 +- .../diffae/figures/ebi_peripheral_droplets.py | 2 +- .../diffae/figures/figure4_morpho_violin.py | 2 +- .../figures/gen_validation/bag_sweep_score.py | 4 +-- .../diffae/figures/gen_validation/embcheck.py | 2 +- .../gen_validation/gen_real_centroid.py | 2 +- .../figures/gen_validation/ntc_inverse_gap.py | 2 +- .../gen_validation/patch_cache_real.py | 2 +- .../figures/gen_validation/st_halves_score.py | 4 +-- .../diffae/figures/rebuild_traversals_n100.py | 10 +++--- .../diffae/generator/virtstain_multi.py | 2 +- .../diffae/traversal/__init__.py | 3 ++ .../{viewer => traversal}/_fluor_topcells.py | 0 .../build_pc_crops_masked.py | 4 +-- .../diffae/{viewer => traversal}/catalog.py | 0 .../{viewer => traversal}/morpho_pipeline.py | 0 .../{viewer => traversal}/precompute.py | 0 75 files changed, 138 insertions(+), 135 deletions(-) rename src/ops_model/models/interpretability/{diffae => _internal}/viewer/__init__.py (100%) rename src/ops_model/models/interpretability/{diffae => _internal}/viewer/_altanchor_build.py (98%) rename src/ops_model/models/interpretability/{diffae => _internal}/viewer/_anchortest.py (92%) rename src/ops_model/models/interpretability/{diffae => _internal}/viewer/_build_stepablation.py (97%) rename src/ops_model/models/interpretability/{diffae => _internal}/viewer/_build_v5_inverted.py (96%) rename src/ops_model/models/interpretability/{diffae => _internal}/viewer/_build_v5_montages.py (97%) rename src/ops_model/models/interpretability/{diffae => _internal}/viewer/_build_valid200.py (90%) rename src/ops_model/models/interpretability/{diffae => _internal}/viewer/_consolidate_cells.py (100%) rename src/ops_model/models/interpretability/{diffae => _internal}/viewer/_fluor_complex_build.py (97%) rename src/ops_model/models/interpretability/{diffae => _internal}/viewer/_fluor_v5_build.py (99%) rename src/ops_model/models/interpretability/{diffae => _internal}/viewer/_migrate_v4_to_v5.py (100%) rename src/ops_model/models/interpretability/{diffae => _internal}/viewer/_phase_vs.py (98%) rename src/ops_model/models/interpretability/{diffae => _internal}/viewer/_rebuild_v5.py (90%) rename src/ops_model/models/interpretability/{diffae => _internal}/viewer/_rescore_rank.py (96%) rename src/ops_model/models/interpretability/{diffae => _internal}/viewer/_score_v4.py (87%) rename src/ops_model/models/interpretability/{diffae => _internal}/viewer/_v4acc_test.py (83%) rename src/ops_model/models/interpretability/{diffae => _internal}/viewer/_verify_pt_space.py (87%) rename src/ops_model/models/interpretability/{diffae => _internal}/viewer/_verify_score_bridge.py (89%) rename src/ops_model/models/interpretability/{diffae => _internal}/viewer/altanchor_pairs.json (100%) rename src/ops_model/models/interpretability/{diffae => _internal}/viewer/anchor_cells.py (95%) rename src/ops_model/models/interpretability/{diffae => _internal}/viewer/build_attention_heads.py (94%) rename src/ops_model/models/interpretability/{diffae => _internal}/viewer/build_complex_ebi_map.py (100%) rename src/ops_model/models/interpretability/{diffae => _internal}/viewer/build_fluor_shap_rankings.py (94%) rename src/ops_model/models/interpretability/{diffae => _internal}/viewer/build_montage_features.py (96%) rename src/ops_model/models/interpretability/{diffae => _internal}/viewer/build_pc_features.py (99%) rename src/ops_model/models/interpretability/{diffae => _internal}/viewer/build_pc_walks.py (93%) rename src/ops_model/models/interpretability/{diffae => _internal}/viewer/build_pcs.py (97%) rename src/ops_model/models/interpretability/{diffae => _internal}/viewer/build_pcs_marker.py (98%) rename src/ops_model/models/interpretability/{diffae => _internal}/viewer/build_phase_shap_rankings.py (95%) rename src/ops_model/models/interpretability/{diffae => _internal}/viewer/build_phate_figure.py (99%) rename src/ops_model/models/interpretability/{diffae => _internal}/viewer/build_setacc_bins.py (100%) rename src/ops_model/models/interpretability/{diffae => _internal}/viewer/build_setacc_bymarker.py (100%) rename src/ops_model/models/interpretability/{diffae => _internal}/viewer/build_top_cells.py (95%) rename src/ops_model/models/interpretability/{diffae => _internal}/viewer/build_umap_montage.py (98%) rename src/ops_model/models/interpretability/{diffae => _internal}/viewer/deploy/README.md (100%) rename src/ops_model/models/interpretability/{diffae => _internal}/viewer/marker_leaves.py (100%) rename src/ops_model/models/interpretability/{diffae => _internal}/viewer/mimic_alex_embed.py (100%) rename src/ops_model/models/interpretability/{diffae => _internal}/viewer/morphometrics.py (100%) rename src/ops_model/models/interpretability/{diffae => _internal}/viewer/nway_clf.py (93%) rename src/ops_model/models/interpretability/{diffae => _internal}/viewer/phenotype_cells.py (96%) rename src/ops_model/models/interpretability/{diffae => _internal}/viewer/render_montage_scales.py (99%) rename src/ops_model/models/interpretability/{diffae => _internal}/viewer/score_generated.py (91%) rename src/ops_model/models/interpretability/{diffae => _internal}/viewer/set_classifier.py (100%) rename src/ops_model/models/interpretability/{diffae => _internal}/viewer/submit.py (95%) rename src/ops_model/models/interpretability/{diffae => _internal}/viewer/webapp/app.js (100%) rename src/ops_model/models/interpretability/{diffae => _internal}/viewer/webapp/biohub-mark.png (100%) rename src/ops_model/models/interpretability/{diffae => _internal}/viewer/webapp/biohub-wordmark.png (100%) rename src/ops_model/models/interpretability/{diffae => _internal}/viewer/webapp/build_gene_narratives.py (100%) rename src/ops_model/models/interpretability/{diffae => _internal}/viewer/webapp/gif.js (100%) rename src/ops_model/models/interpretability/{diffae => _internal}/viewer/webapp/gif.worker.js (100%) rename src/ops_model/models/interpretability/{diffae => _internal}/viewer/webapp/index.html (100%) rename src/ops_model/models/interpretability/{diffae => _internal}/viewer/webapp/methods.js (100%) rename src/ops_model/models/interpretability/{diffae => _internal}/viewer/webapp/morpho_demo.html (100%) rename src/ops_model/models/interpretability/{diffae => _internal}/viewer/webapp/openseadragon.min.js (100%) rename src/ops_model/models/interpretability/{diffae => _internal}/viewer/webapp/opsin-eyes.svg (100%) rename src/ops_model/models/interpretability/{diffae => _internal}/viewer/webapp/style.css (100%) create mode 100644 src/ops_model/models/interpretability/diffae/traversal/__init__.py rename src/ops_model/models/interpretability/diffae/{viewer => traversal}/_fluor_topcells.py (100%) rename src/ops_model/models/interpretability/diffae/{viewer => traversal}/build_pc_crops_masked.py (97%) rename src/ops_model/models/interpretability/diffae/{viewer => traversal}/catalog.py (100%) rename src/ops_model/models/interpretability/diffae/{viewer => traversal}/morpho_pipeline.py (100%) rename src/ops_model/models/interpretability/diffae/{viewer => traversal}/precompute.py (100%) diff --git a/src/ops_model/models/interpretability/diffae/viewer/__init__.py b/src/ops_model/models/interpretability/_internal/viewer/__init__.py similarity index 100% rename from src/ops_model/models/interpretability/diffae/viewer/__init__.py rename to src/ops_model/models/interpretability/_internal/viewer/__init__.py diff --git a/src/ops_model/models/interpretability/diffae/viewer/_altanchor_build.py b/src/ops_model/models/interpretability/_internal/viewer/_altanchor_build.py similarity index 98% rename from src/ops_model/models/interpretability/diffae/viewer/_altanchor_build.py rename to src/ops_model/models/interpretability/_internal/viewer/_altanchor_build.py index 4a2402d..e6fe5d5 100644 --- a/src/ops_model/models/interpretability/diffae/viewer/_altanchor_build.py +++ b/src/ops_model/models/interpretability/_internal/viewer/_altanchor_build.py @@ -7,7 +7,7 @@ """ import os, json, glob from . import catalog as C -from ..classifier.config import slugify +from ops_model.models.interpretability.diffae.classifier.config import slugify ASSETS = "viewer_assets_v5" ROOT = C.OUT diff --git a/src/ops_model/models/interpretability/diffae/viewer/_anchortest.py b/src/ops_model/models/interpretability/_internal/viewer/_anchortest.py similarity index 92% rename from src/ops_model/models/interpretability/diffae/viewer/_anchortest.py rename to src/ops_model/models/interpretability/_internal/viewer/_anchortest.py index 8bbd073..28865d8 100644 --- a/src/ops_model/models/interpretability/diffae/viewer/_anchortest.py +++ b/src/ops_model/models/interpretability/_internal/viewer/_anchortest.py @@ -7,7 +7,7 @@ (NTC) gather reads the v4 parquet while the KD gather uses the v5 accuracy parquet override. Output is isolated under viewer_assets_v5_anchortest/. Delete this module after the diagnosis.""" from . import catalog as C -from .precompute import precompute_marker +from ops_model.models.interpretability.diffae.traversal.precompute import precompute_marker V5_GENEKO = "/hpc/projects/icd.fast.ops/models/diffex/viewer_assets_v5/_rankings/pma_v5_phase_geneKO.parquet" PHASE_CK = f"{C.DD}/phase_v1/diffae_best.pt" diff --git a/src/ops_model/models/interpretability/diffae/viewer/_build_stepablation.py b/src/ops_model/models/interpretability/_internal/viewer/_build_stepablation.py similarity index 97% rename from src/ops_model/models/interpretability/diffae/viewer/_build_stepablation.py rename to src/ops_model/models/interpretability/_internal/viewer/_build_stepablation.py index f1bc3b4..e5fb670 100644 --- a/src/ops_model/models/interpretability/diffae/viewer/_build_stepablation.py +++ b/src/ops_model/models/interpretability/_internal/viewer/_build_stepablation.py @@ -10,7 +10,7 @@ import numpy as np from . import catalog as C -from .precompute import precompute_marker +from ops_model.models.interpretability.diffae.traversal.precompute import precompute_marker STEPS = int(os.environ.get("VAL_STEPS", "50")) INVERT = int(os.environ.get("VAL_INVERT", "1")) # 0 = random-xT (non-inverted, the old scheme) diff --git a/src/ops_model/models/interpretability/diffae/viewer/_build_v5_inverted.py b/src/ops_model/models/interpretability/_internal/viewer/_build_v5_inverted.py similarity index 96% rename from src/ops_model/models/interpretability/diffae/viewer/_build_v5_inverted.py rename to src/ops_model/models/interpretability/_internal/viewer/_build_v5_inverted.py index 904425e..6a7bba9 100644 --- a/src/ops_model/models/interpretability/diffae/viewer/_build_v5_inverted.py +++ b/src/ops_model/models/interpretability/_internal/viewer/_build_v5_inverted.py @@ -9,7 +9,7 @@ complex) whose per-cell rankings exist. That is Alex Lin's top1_acc>0.5@100-cell distinctiveness filter (see _fluor_v5_build.py) — NOT missing data (cells exist for all 1000×55). Lower-acc combos need Alex to gen more. - python -m ops_model.models.interpretability.diffae.viewer._build_v5_inverted markers + python -m ops_model.models.interpretability._internal.viewer._build_v5_inverted markers """ import json import os @@ -17,8 +17,8 @@ from pathlib import Path from . import catalog as C -from ..classifier.config import slugify -from .precompute import precompute_marker +from ops_model.models.interpretability.diffae.classifier.config import slugify +from ops_model.models.interpretability.diffae.traversal.precompute import precompute_marker FRP_DIR = f"{C.OUT}/viewer_assets_v5/_rankings/fluor_shap/geneKO" # NEW shap_screen rankings (robust bin-size top-acc); old at _rankings/fluor/geneKO_OLD_qualifying backup # OUTPUT tree: a FRESH dir so force=False gives skip-done resume (timeouts harmless) without clobbering the old @@ -93,9 +93,9 @@ def build_phase_anchor_200(): import numpy as np, pandas as pd from pathlib import Path from concurrent.futures import ThreadPoolExecutor - from .precompute import _gather_class, _save_webp - from ..generator.data import normalize - from ..directions.config import DirConfig + from ops_model.models.interpretability.diffae.traversal.precompute import _gather_class, _save_webp + from ops_model.models.interpretability.diffae.generator.data import normalize + from ops_model.models.interpretability.diffae.directions.config import DirConfig _use_v5() rd = Path(C.OUT) / _V5 / "phase" / "_anchors" / "NTC" z = dict(np.load(rd / "ctrl.npz")) @@ -135,9 +135,9 @@ def build_phase_anchor_multirank(hi=400): pma_v5 ranking — here we gather NTC by the multirank so cells 200-399 are the clean multi_bag anchors.""" import numpy as np from concurrent.futures import ThreadPoolExecutor - from .precompute import _gather_class, _save_webp - from ..generator.data import normalize - from ..directions.config import DirConfig + from ops_model.models.interpretability.diffae.traversal.precompute import _gather_class, _save_webp + from ops_model.models.interpretability.diffae.generator.data import normalize + from ops_model.models.interpretability.diffae.directions.config import DirConfig _use_v5() rd = Path(C.OUT) / _V5 / "phase" / "_anchors" / "NTC" z = dict(np.load(rd / "ctrl.npz")) @@ -252,9 +252,9 @@ def build_phase_anchor(): import numpy as np, pandas as pd from pathlib import Path from concurrent.futures import ThreadPoolExecutor - from .precompute import _gather_class, _save_webp - from ..generator.data import normalize - from ..directions.config import DirConfig + from ops_model.models.interpretability.diffae.traversal.precompute import _gather_class, _save_webp + from ops_model.models.interpretability.diffae.generator.data import normalize + from ops_model.models.interpretability.diffae.directions.config import DirConfig _use_v5() cfg = DirConfig(grain="geneKO", target="NTC", control="NTC", device="cuda") a_imgs, a_emb = _gather_class(cfg, "NTC", 20) # 20 attention @@ -431,9 +431,9 @@ def prebuild_marker_anchor(d, marker_channel, channel): directions) + real.webp. No reliance on drop/fresh-gather-branch — this IS the anchor build.""" import numpy as np, pandas as pd from concurrent.futures import ThreadPoolExecutor - from .precompute import _gather_class, _save_webp - from ..generator.data import normalize - from ..directions.config import DirConfig + from ops_model.models.interpretability.diffae.traversal.precompute import _gather_class, _save_webp + from ops_model.models.interpretability.diffae.generator.data import normalize + from ops_model.models.interpretability.diffae.directions.config import DirConfig frp = f"{FRP_DIR}/{slugify(marker_channel)}.parquet" if not os.path.exists(frp): return f"skip {marker_channel}" diff --git a/src/ops_model/models/interpretability/diffae/viewer/_build_v5_montages.py b/src/ops_model/models/interpretability/_internal/viewer/_build_v5_montages.py similarity index 97% rename from src/ops_model/models/interpretability/diffae/viewer/_build_v5_montages.py rename to src/ops_model/models/interpretability/_internal/viewer/_build_v5_montages.py index 6afad71..290880d 100644 --- a/src/ops_model/models/interpretability/diffae/viewer/_build_v5_montages.py +++ b/src/ops_model/models/interpretability/_internal/viewer/_build_v5_montages.py @@ -5,7 +5,7 @@ phase embedding. Reads the merged inverted frames from viewer_assets_v5//geneKO and writes tiles to viewer_assets_v5/_montage/. Only builds markers whose geneKO is 100% present in viewer_assets_v5. - python -m ops_model.models.interpretability.diffae.viewer._build_v5_montages + python -m ops_model.models.interpretability._internal.viewer._build_v5_montages """ import glob import os @@ -13,7 +13,7 @@ from . import catalog as C from . import marker_leaves as ML from .build_umap_montage import OUT -from ..classifier.config import slugify +from ops_model.models.interpretability.diffae.classifier.config import slugify V5 = "viewer_assets_v5" CELLS = list(range(20)) diff --git a/src/ops_model/models/interpretability/diffae/viewer/_build_valid200.py b/src/ops_model/models/interpretability/_internal/viewer/_build_valid200.py similarity index 90% rename from src/ops_model/models/interpretability/diffae/viewer/_build_valid200.py rename to src/ops_model/models/interpretability/_internal/viewer/_build_valid200.py index 22e6555..bdebc51 100644 --- a/src/ops_model/models/interpretability/diffae/viewer/_build_valid200.py +++ b/src/ops_model/models/interpretability/_internal/viewer/_build_valid200.py @@ -8,17 +8,17 @@ Reuses the v5 per-class directions (d_vec + gap are the PRODUCTION values we are validating) via a symlinked _directions tree, so no direction re-fit — only the 200-anchor inversion + 200×7 decodes per target. - python -m ops_model.models.interpretability.diffae.viewer._build_valid200 anchor # 1 GPU: build the 200-cell NTC anchor - python -m ops_model.models.interpretability.diffae.viewer._build_valid200 submit # shard all 1000 geneKO (after anchor) - python -m ops_model.models.interpretability.diffae.viewer._build_valid200 all # anchor job -> shards (afterok dep) + python -m ops_model.models.interpretability._internal.viewer._build_valid200 anchor # 1 GPU: build the 200-cell NTC anchor + python -m ops_model.models.interpretability._internal.viewer._build_valid200 submit # shard all 1000 geneKO (after anchor) + python -m ops_model.models.interpretability._internal.viewer._build_valid200 all # anchor job -> shards (afterok dep) """ import os import sys from pathlib import Path -from ..classifier.config import slugify +from ops_model.models.interpretability.diffae.classifier.config import slugify from . import catalog as C -from .precompute import precompute_marker +from ops_model.models.interpretability.diffae.traversal.precompute import precompute_marker W = float(os.environ.get("VALID200_W", "1.5")) # CFG guidance weight (baseline recovery used w=2.0) _VALID = "viewer_assets_valid200" if W == 1.5 else f"viewer_assets_valid200_w{W:g}" @@ -64,9 +64,9 @@ def build_anchor(): (needed for faithful DDIM inversion) + real.webp, once.""" import numpy as np from concurrent.futures import ThreadPoolExecutor - from .precompute import _gather_class, _save_webp - from ..generator.data import normalize - from ..directions.config import DirConfig + from ops_model.models.interpretability.diffae.traversal.precompute import _gather_class, _save_webp + from ops_model.models.interpretability.diffae.generator.data import normalize + from ops_model.models.interpretability.diffae.directions.config import DirConfig _use_valid() setup_dirs() cfg = DirConfig(grain="geneKO", target="NTC", control="NTC", device="cuda") diff --git a/src/ops_model/models/interpretability/diffae/viewer/_consolidate_cells.py b/src/ops_model/models/interpretability/_internal/viewer/_consolidate_cells.py similarity index 100% rename from src/ops_model/models/interpretability/diffae/viewer/_consolidate_cells.py rename to src/ops_model/models/interpretability/_internal/viewer/_consolidate_cells.py diff --git a/src/ops_model/models/interpretability/diffae/viewer/_fluor_complex_build.py b/src/ops_model/models/interpretability/_internal/viewer/_fluor_complex_build.py similarity index 97% rename from src/ops_model/models/interpretability/diffae/viewer/_fluor_complex_build.py rename to src/ops_model/models/interpretability/_internal/viewer/_fluor_complex_build.py index 52e23a4..d5fc370 100644 --- a/src/ops_model/models/interpretability/diffae/viewer/_fluor_complex_build.py +++ b/src/ops_model/models/interpretability/_internal/viewer/_fluor_complex_build.py @@ -9,8 +9,8 @@ import os, re, json import pandas as pd from . import catalog as C -from ..classifier.config import slugify -from ._fluor_topcells import crop_marker_shard, TOP_N, OUT +from ops_model.models.interpretability.diffae.classifier.config import slugify +from ops_model.models.interpretability.diffae.traversal._fluor_topcells import crop_marker_shard, TOP_N, OUT F = "/hpc/projects/icd.fast.ops/models/alex_lin_attention/v5/fluorescence" EBI = f"{F}/misc/gene_marker_ebi_complexqual.compact.parquet" # per-GENE EBI cells (all 55 channels) diff --git a/src/ops_model/models/interpretability/diffae/viewer/_fluor_v5_build.py b/src/ops_model/models/interpretability/_internal/viewer/_fluor_v5_build.py similarity index 99% rename from src/ops_model/models/interpretability/diffae/viewer/_fluor_v5_build.py rename to src/ops_model/models/interpretability/_internal/viewer/_fluor_v5_build.py index 257ec84..4dd046f 100644 --- a/src/ops_model/models/interpretability/diffae/viewer/_fluor_v5_build.py +++ b/src/ops_model/models/interpretability/_internal/viewer/_fluor_v5_build.py @@ -7,7 +7,7 @@ import os, re, json import pandas as pd from . import catalog as C -from ..classifier.config import slugify +from ops_model.models.interpretability.diffae.classifier.config import slugify ASSETS = "viewer_assets_v5" F = "/hpc/projects/icd.fast.ops/models/alex_lin_attention/v5/fluorescence" diff --git a/src/ops_model/models/interpretability/diffae/viewer/_migrate_v4_to_v5.py b/src/ops_model/models/interpretability/_internal/viewer/_migrate_v4_to_v5.py similarity index 100% rename from src/ops_model/models/interpretability/diffae/viewer/_migrate_v4_to_v5.py rename to src/ops_model/models/interpretability/_internal/viewer/_migrate_v4_to_v5.py diff --git a/src/ops_model/models/interpretability/diffae/viewer/_phase_vs.py b/src/ops_model/models/interpretability/_internal/viewer/_phase_vs.py similarity index 98% rename from src/ops_model/models/interpretability/diffae/viewer/_phase_vs.py rename to src/ops_model/models/interpretability/_internal/viewer/_phase_vs.py index e216753..f73948e 100644 --- a/src/ops_model/models/interpretability/diffae/viewer/_phase_vs.py +++ b/src/ops_model/models/interpretability/_internal/viewer/_phase_vs.py @@ -18,7 +18,7 @@ import torch from PIL import Image -from ..classifier.config import slugify +from ops_model.models.interpretability.diffae.classifier.config import slugify VS_OUT = "/hpc/projects/icd.fast.ops/analysis/virtual_staining/multi_marker" V5 = "/hpc/projects/icd.fast.ops/models/diffex/viewer_assets_v5" @@ -27,8 +27,8 @@ def load_vs(dev): - from ..generator.config import DiffAEConfig - from ..generator.model import DiffAE + from ops_model.models.interpretability.diffae.generator.config import DiffAEConfig + from ops_model.models.interpretability.diffae.generator.model import DiffAE markers = json.load(open(f"{VS_OUT}/markers.json")) cfg = DiffAEConfig(spatial_cond=True, n_markers=len(markers), device="cuda", epochs=1) ema = DiffAE(cfg).to(dev).eval() @@ -48,8 +48,8 @@ def _load_phase(gene, cell, ai, H): @torch.no_grad() def stain(ema, markers, cfg, dev, phase_np, seed=0): """phase_np (1,1,H,H) in [-1,1] → {marker_idx: pred (H,H)} for all markers (fixed xT seed).""" - from ..generator.virtstain_multi import _sample_marker - from ..classifier.celldino_features import embed_crops + from ops_model.models.interpretability.diffae.generator.virtstain_multi import _sample_marker + from ops_model.models.interpretability.diffae.classifier.celldino_features import embed_crops H = cfg.crop_size emb = torch.as_tensor(embed_crops(phase_np, cfg), dtype=torch.float32, device=dev) ci = torch.as_tensor(phase_np, dtype=torch.float32, device=dev) diff --git a/src/ops_model/models/interpretability/diffae/viewer/_rebuild_v5.py b/src/ops_model/models/interpretability/_internal/viewer/_rebuild_v5.py similarity index 90% rename from src/ops_model/models/interpretability/diffae/viewer/_rebuild_v5.py rename to src/ops_model/models/interpretability/_internal/viewer/_rebuild_v5.py index 05ef44f..cbb669f 100644 --- a/src/ops_model/models/interpretability/diffae/viewer/_rebuild_v5.py +++ b/src/ops_model/models/interpretability/_internal/viewer/_rebuild_v5.py @@ -3,7 +3,7 @@ scoring the v5 SetTransformer inline (reuses gemb → no separate re-decode/re-embed pass). force=True recomputes directions (v5_KD − v4_NTC) and frames. Removable after the rebuild.""" from . import catalog as C -from .precompute import precompute_marker +from ops_model.models.interpretability.diffae.traversal.precompute import precompute_marker V5G = "/hpc/projects/icd.fast.ops/models/diffex/viewer_assets_v5/_rankings/pma_v5_phase_geneKO.parquet" V5C = "/hpc/projects/icd.fast.ops/models/diffex/viewer_assets_v5/_rankings/pma_v5_phase_complex.parquet" @@ -32,9 +32,9 @@ def build_accpool_anchor(): import numpy as np, pandas as pd from pathlib import Path from concurrent.futures import ThreadPoolExecutor - from .precompute import _gather_class, _ASSETS, _save_webp - from ..generator.data import normalize - from ..directions.config import DirConfig + from ops_model.models.interpretability.diffae.traversal.precompute import _gather_class, _ASSETS, _save_webp + from ops_model.models.interpretability.diffae.generator.data import normalize + from ops_model.models.interpretability.diffae.directions.config import DirConfig sel = pd.read_csv(SEL25) parq = pd.DataFrame({"gene": "NTC", "experiment": sel.experiment, "well": sel.well, "segmentation": sel.segmentation, "x_pheno": sel.x_pheno, "y_pheno": sel.y_pheno, "pma_attention": sel.pma_attention, diff --git a/src/ops_model/models/interpretability/diffae/viewer/_rescore_rank.py b/src/ops_model/models/interpretability/_internal/viewer/_rescore_rank.py similarity index 96% rename from src/ops_model/models/interpretability/diffae/viewer/_rescore_rank.py rename to src/ops_model/models/interpretability/_internal/viewer/_rescore_rank.py index e59758e..6676e4e 100644 --- a/src/ops_model/models/interpretability/diffae/viewer/_rescore_rank.py +++ b/src/ops_model/models/interpretability/_internal/viewer/_rescore_rank.py @@ -11,7 +11,7 @@ def _retarget(assets): """Point the score module at `assets` and return the module (V5_BASE is read at call time).""" - import ops_model.models.interpretability.diffae.viewer.score_generated as SG + import ops_model.models.interpretability._internal.viewer.score_generated as SG SG.V5_BASE = f"{BASE}/{assets}/phase" return SG diff --git a/src/ops_model/models/interpretability/diffae/viewer/_score_v4.py b/src/ops_model/models/interpretability/_internal/viewer/_score_v4.py similarity index 87% rename from src/ops_model/models/interpretability/diffae/viewer/_score_v4.py rename to src/ops_model/models/interpretability/_internal/viewer/_score_v4.py index 4b12ed8..e9fd8c7 100644 --- a/src/ops_model/models/interpretability/diffae/viewer/_score_v4.py +++ b/src/ops_model/models/interpretability/_internal/viewer/_score_v4.py @@ -10,9 +10,9 @@ def score_v4_shard(grain, targets, out_json): from .set_classifier import load_set_classifier, V5_CKPT_ROOT, V5_RUNS from .score_generated import _emb_frames, score_embs_v5 - from ..directions.config import DirConfig - from ..classifier.celldino_features import embed_crops - from ..classifier.config import slugify + from ops_model.models.interpretability.diffae.directions.config import DirConfig + from ops_model.models.interpretability.diffae.classifier.celldino_features import embed_crops + from ops_model.models.interpretability.diffae.classifier.config import slugify run = V5_RUNS[("phase", "geneKO" if grain == "geneKO" else "complex_ebionly")] model, g2i, ci_map = load_set_classifier(run=run, device="cuda", root=V5_CKPT_ROOT) ci = ci_map.get("Phase2D", 0) diff --git a/src/ops_model/models/interpretability/diffae/viewer/_v4acc_test.py b/src/ops_model/models/interpretability/_internal/viewer/_v4acc_test.py similarity index 83% rename from src/ops_model/models/interpretability/diffae/viewer/_v4acc_test.py rename to src/ops_model/models/interpretability/_internal/viewer/_v4acc_test.py index e6ab223..9896d49 100644 --- a/src/ops_model/models/interpretability/diffae/viewer/_v4acc_test.py +++ b/src/ops_model/models/interpretability/_internal/viewer/_v4acc_test.py @@ -10,8 +10,8 @@ def run(): - from ..classifier.config import GRAINS + from ops_model.models.interpretability.diffae.classifier.config import GRAINS GRAINS["complex"]["parquet"] = V4ACC # both anchor (A) and target (B) from v4-accuracy cells - from .precompute import precompute_anchors_marker + from ops_model.models.interpretability.diffae.traversal.precompute import precompute_anchors_marker return precompute_anchors_marker(grain="complex", classes=[C40, C60], ckpt=f"{C.DD}/phase_v1/diffae_best.pt", out_root=C.OUT) diff --git a/src/ops_model/models/interpretability/diffae/viewer/_verify_pt_space.py b/src/ops_model/models/interpretability/_internal/viewer/_verify_pt_space.py similarity index 87% rename from src/ops_model/models/interpretability/diffae/viewer/_verify_pt_space.py rename to src/ops_model/models/interpretability/_internal/viewer/_verify_pt_space.py index 3188184..0e8d1ce 100644 --- a/src/ops_model/models/interpretability/diffae/viewer/_verify_pt_space.py +++ b/src/ops_model/models/interpretability/_internal/viewer/_verify_pt_space.py @@ -2,10 +2,10 @@ space)? Embed top cells via the current gather, match to .pt by segmentation_id, compare. Run as a one-off SLURM job; prints cosine + gap agreement. Delete after.""" import torch, numpy as np -from .precompute import _gather_class -from ..directions.config import DirConfig -from ..directions.data import _top_cells -from ..classifier.config import GRAINS +from ops_model.models.interpretability.diffae.traversal.precompute import _gather_class +from ops_model.models.interpretability.diffae.directions.config import DirConfig +from ops_model.models.interpretability.diffae.directions.data import _top_cells +from ops_model.models.interpretability.diffae.classifier.config import GRAINS PT = "/hpc/projects/icd.fast.ops/models/alex_lin_attention/v4/train_ops_zstdcontrol_cdino_v2" diff --git a/src/ops_model/models/interpretability/diffae/viewer/_verify_score_bridge.py b/src/ops_model/models/interpretability/_internal/viewer/_verify_score_bridge.py similarity index 89% rename from src/ops_model/models/interpretability/diffae/viewer/_verify_score_bridge.py rename to src/ops_model/models/interpretability/_internal/viewer/_verify_score_bridge.py index 6ba459f..44133e8 100644 --- a/src/ops_model/models/interpretability/diffae/viewer/_verify_score_bridge.py +++ b/src/ops_model/models/interpretability/_internal/viewer/_verify_score_bridge.py @@ -3,8 +3,8 @@ Run as a GPU job; prints P(KIF11) for raw vs zstd-control embed_crops bags. Delete after.""" import numpy as np -from ..directions.config import DirConfig -from .precompute import _gather_class +from ops_model.models.interpretability.diffae.directions.config import DirConfig +from ops_model.models.interpretability.diffae.traversal.precompute import _gather_class from .set_classifier import load_set_classifier, score_bags diff --git a/src/ops_model/models/interpretability/diffae/viewer/altanchor_pairs.json b/src/ops_model/models/interpretability/_internal/viewer/altanchor_pairs.json similarity index 100% rename from src/ops_model/models/interpretability/diffae/viewer/altanchor_pairs.json rename to src/ops_model/models/interpretability/_internal/viewer/altanchor_pairs.json diff --git a/src/ops_model/models/interpretability/diffae/viewer/anchor_cells.py b/src/ops_model/models/interpretability/_internal/viewer/anchor_cells.py similarity index 95% rename from src/ops_model/models/interpretability/diffae/viewer/anchor_cells.py rename to src/ops_model/models/interpretability/_internal/viewer/anchor_cells.py index b82a746..5d3c2e3 100644 --- a/src/ops_model/models/interpretability/diffae/viewer/anchor_cells.py +++ b/src/ops_model/models/interpretability/_internal/viewer/anchor_cells.py @@ -10,8 +10,8 @@ import pandas as pd -from ..classifier.config import PMA_PHASE_EBI -from ..classifier.data import _BASE_COLS +from ops_model.models.interpretability.diffae.classifier.config import PMA_PHASE_EBI +from ops_model.models.interpretability.diffae.classifier.data import _BASE_COLS from . import catalog as C OUT_CSV = f"{C.OUT}/viewer_assets/anchor_cells_for_attention.csv" diff --git a/src/ops_model/models/interpretability/diffae/viewer/build_attention_heads.py b/src/ops_model/models/interpretability/_internal/viewer/build_attention_heads.py similarity index 94% rename from src/ops_model/models/interpretability/diffae/viewer/build_attention_heads.py rename to src/ops_model/models/interpretability/_internal/viewer/build_attention_heads.py index 7c406b1..c6cd9d6 100644 --- a/src/ops_model/models/interpretability/diffae/viewer/build_attention_heads.py +++ b/src/ops_model/models/interpretability/_internal/viewer/build_attention_heads.py @@ -18,10 +18,10 @@ {AH}/index.json {global_max, assets:{:{:[keys]}}} where modality = "phase" | slugify(marker_channel), grain = geneKO|complex, key = gene | complex-slug. - python -m ops_model.models.interpretability.diffae.viewer.build_attention_heads render # SLURM (all trees) - python -m ops_model.models.interpretability.diffae.viewer.build_attention_heads render --local # serial, no SLURM - python -m ops_model.models.interpretability.diffae.viewer.build_attention_heads render --dry-run - python -m ops_model.models.interpretability.diffae.viewer.build_attention_heads index # (re)aggregate index.json + python -m ops_model.models.interpretability._internal.viewer.build_attention_heads render # SLURM (all trees) + python -m ops_model.models.interpretability._internal.viewer.build_attention_heads render --local # serial, no SLURM + python -m ops_model.models.interpretability._internal.viewer.build_attention_heads render --dry-run + python -m ops_model.models.interpretability._internal.viewer.build_attention_heads index # (re)aggregate index.json """ from __future__ import annotations @@ -36,7 +36,7 @@ from PIL import Image from scipy.ndimage import gaussian_filter -from ..classifier.config import slugify +from ops_model.models.interpretability.diffae.classifier.config import slugify from . import catalog as C AH_ROOT = f"{C.OUT}/viewer_assets/attention_heads" diff --git a/src/ops_model/models/interpretability/diffae/viewer/build_complex_ebi_map.py b/src/ops_model/models/interpretability/_internal/viewer/build_complex_ebi_map.py similarity index 100% rename from src/ops_model/models/interpretability/diffae/viewer/build_complex_ebi_map.py rename to src/ops_model/models/interpretability/_internal/viewer/build_complex_ebi_map.py diff --git a/src/ops_model/models/interpretability/diffae/viewer/build_fluor_shap_rankings.py b/src/ops_model/models/interpretability/_internal/viewer/build_fluor_shap_rankings.py similarity index 94% rename from src/ops_model/models/interpretability/diffae/viewer/build_fluor_shap_rankings.py rename to src/ops_model/models/interpretability/_internal/viewer/build_fluor_shap_rankings.py index 1d2b39f..496f7dd 100644 --- a/src/ops_model/models/interpretability/diffae/viewer/build_fluor_shap_rankings.py +++ b/src/ops_model/models/interpretability/_internal/viewer/build_fluor_shap_rankings.py @@ -5,8 +5,8 @@ new CSV cols: gene, channel_name, rank, shap, ..., experiment, well, x_pheno, y_pheno, segmentation_id old schema: channel_name, gene, rank, pma_attention, experiment, well, x_pheno, y_pheno, segmentation, rank_type - python -m ops_model.models.interpretability.diffae.viewer.build_fluor_shap_rankings # local (needs ~64GB) - python -m ops_model.models.interpretability.diffae.viewer.build_fluor_shap_rankings --submit # SLURM cpu, mem 96 + python -m ops_model.models.interpretability._internal.viewer.build_fluor_shap_rankings # local (needs ~64GB) + python -m ops_model.models.interpretability._internal.viewer.build_fluor_shap_rankings --submit # SLURM cpu, mem 96 """ from __future__ import annotations @@ -15,7 +15,7 @@ import pandas as pd -from ..classifier.config import slugify +from ops_model.models.interpretability.diffae.classifier.config import slugify CSV = ("/hpc/projects/icd.fast.ops/models/alex_lin_attention/v5/multi_rank/" "shap_screen/shap_screen_fluor_all.csv") diff --git a/src/ops_model/models/interpretability/diffae/viewer/build_montage_features.py b/src/ops_model/models/interpretability/_internal/viewer/build_montage_features.py similarity index 96% rename from src/ops_model/models/interpretability/diffae/viewer/build_montage_features.py rename to src/ops_model/models/interpretability/_internal/viewer/build_montage_features.py index 06894b7..b1fafaa 100644 --- a/src/ops_model/models/interpretability/diffae/viewer/build_montage_features.py +++ b/src/ops_model/models/interpretability/_internal/viewer/build_montage_features.py @@ -7,7 +7,7 @@ viewer_assets/montage_features.json {"features": [base names], "range": {feat: [lo, hi]}, "values": {gene: [0..1 per feature | null]}} - python -m ops_model.models.interpretability.diffae.viewer.build_montage_features + python -m ops_model.models.interpretability._internal.viewer.build_montage_features """ from __future__ import annotations diff --git a/src/ops_model/models/interpretability/diffae/viewer/build_pc_features.py b/src/ops_model/models/interpretability/_internal/viewer/build_pc_features.py similarity index 99% rename from src/ops_model/models/interpretability/diffae/viewer/build_pc_features.py rename to src/ops_model/models/interpretability/_internal/viewer/build_pc_features.py index c976812..dee29f8 100644 --- a/src/ops_model/models/interpretability/diffae/viewer/build_pc_features.py +++ b/src/ops_model/models/interpretability/_internal/viewer/build_pc_features.py @@ -14,7 +14,7 @@ to make "+corr features" line up with the strip's high bins. Compositions are unsigned (|r| / tf-idf) so they need no flip. - python -m ops_model.models.interpretability.diffae.viewer.build_pc_features + python -m ops_model.models.interpretability._internal.viewer.build_pc_features """ from __future__ import annotations diff --git a/src/ops_model/models/interpretability/diffae/viewer/build_pc_walks.py b/src/ops_model/models/interpretability/_internal/viewer/build_pc_walks.py similarity index 93% rename from src/ops_model/models/interpretability/diffae/viewer/build_pc_walks.py rename to src/ops_model/models/interpretability/_internal/viewer/build_pc_walks.py index 5d7063d..322bbc2 100644 --- a/src/ops_model/models/interpretability/diffae/viewer/build_pc_walks.py +++ b/src/ops_model/models/interpretability/_internal/viewer/build_pc_walks.py @@ -8,8 +8,8 @@ the PC score std; v_p is the (z-scored) eigenvector, mapped back to raw CellDINO space by the mean per-exp sd. Output: one composite figure per marker (rows = PCs, cols = α). - python -m ops_model.models.interpretability.diffae.viewer.build_pc_walks --markers "Mitochondria_TOMM20" - python -m ops_model.models.interpretability.diffae.viewer.build_pc_walks --all # SLURM, every marker + python -m ops_model.models.interpretability._internal.viewer.build_pc_walks --markers "Mitochondria_TOMM20" + python -m ops_model.models.interpretability._internal.viewer.build_pc_walks --all # SLURM, every marker """ from __future__ import annotations @@ -34,8 +34,8 @@ def pc_walks_marker(marker_channel, channel, ckpt, out_root, n_pcs=20, n_cells=1 from concurrent.futures import ThreadPoolExecutor from sklearn.decomposition import PCA from .build_pcs_marker import _marker_meta, _fp, _load, FIT_N - from .precompute import DirConfig, load_diffae, _sample_guided, _save_webp, VIEWER_ALPHAS - from ..classifier.config import slugify + from ops_model.models.interpretability.diffae.traversal.precompute import DirConfig, load_diffae, _sample_guided, _save_webp, VIEWER_ALPHAS + from ops_model.models.interpretability.diffae.classifier.config import slugify dev = torch.device(device if torch.cuda.is_available() else "cpu") modality = slugify(marker_channel) if marker_channel else "phase" @@ -120,7 +120,7 @@ def pc_walks_marker(marker_channel, channel, ckpt, out_root, n_pcs=20, n_cells=1 def _marker_jobs(n_pcs, n_cells, force): from . import catalog as C - from ..classifier.config import slugify + from ops_model.models.interpretability.diffae.classifier.config import slugify jobs = [] for d, mc, ch in C.complete_markers(): jobs.append({"name": f"pcw_{slugify(mc)[:18]}", "func": pc_walks_marker, diff --git a/src/ops_model/models/interpretability/diffae/viewer/build_pcs.py b/src/ops_model/models/interpretability/_internal/viewer/build_pcs.py similarity index 97% rename from src/ops_model/models/interpretability/diffae/viewer/build_pcs.py rename to src/ops_model/models/interpretability/_internal/viewer/build_pcs.py index 36e90b6..ff4aa06 100644 --- a/src/ops_model/models/interpretability/diffae/viewer/build_pcs.py +++ b/src/ops_model/models/interpretability/_internal/viewer/build_pcs.py @@ -9,8 +9,8 @@ the raw artifacts dir is gone). If Kyle regenerates artifacts, re-run his build_static_explorer.py then point --html at the fresh output. - python -m ops_model.models.interpretability.diffae.viewer.build_pcs - python -m ops_model.models.interpretability.diffae.viewer.build_pcs --html /path/to/pc_explorer_static.html + python -m ops_model.models.interpretability._internal.viewer.build_pcs + python -m ops_model.models.interpretability._internal.viewer.build_pcs --html /path/to/pc_explorer_static.html """ from __future__ import annotations diff --git a/src/ops_model/models/interpretability/diffae/viewer/build_pcs_marker.py b/src/ops_model/models/interpretability/_internal/viewer/build_pcs_marker.py similarity index 98% rename from src/ops_model/models/interpretability/diffae/viewer/build_pcs_marker.py rename to src/ops_model/models/interpretability/_internal/viewer/build_pcs_marker.py index 62c969c..68b15eb 100644 --- a/src/ops_model/models/interpretability/diffae/viewer/build_pcs_marker.py +++ b/src/ops_model/models/interpretability/_internal/viewer/build_pcs_marker.py @@ -9,7 +9,7 @@ viewer_assets/pcs/markers//index.json (same schema as the phase pcs/index.json) viewer_assets/pcs/markers//crops/pc###_bin##_row#.png - python -m ops_model.models.interpretability.diffae.viewer.build_pcs_marker --marker "autophagosome_MAP1LC3B" + python -m ops_model.models.interpretability._internal.viewer.build_pcs_marker --marker "autophagosome_MAP1LC3B" """ from __future__ import annotations @@ -22,7 +22,7 @@ from . import catalog as C from . import marker_leaves as ML -from .build_pc_crops_masked import CROP_SIZE, _crop, _is_blank, _render, _zarr_patch +from ops_model.models.interpretability.diffae.traversal.build_pc_crops_masked import CROP_SIZE, _crop, _is_blank, _render, _zarr_patch FOPS = "/hpc/projects/intracellular_dashboard/fast_ops" PCS_OUT = f"{C.OUT}/viewer_assets/pcs/markers" diff --git a/src/ops_model/models/interpretability/diffae/viewer/build_phase_shap_rankings.py b/src/ops_model/models/interpretability/_internal/viewer/build_phase_shap_rankings.py similarity index 95% rename from src/ops_model/models/interpretability/diffae/viewer/build_phase_shap_rankings.py rename to src/ops_model/models/interpretability/_internal/viewer/build_phase_shap_rankings.py index acff805..760ac48 100644 --- a/src/ops_model/models/interpretability/diffae/viewer/build_phase_shap_rankings.py +++ b/src/ops_model/models/interpretability/_internal/viewer/build_phase_shap_rankings.py @@ -6,8 +6,8 @@ pma_attention, rank, rank_type) complex → pma_shap_phase_complex.parquet (adds predicted_class=complex, gene=member gene; EBI-pooled) - python -m ops_model.models.interpretability.diffae.viewer.build_phase_shap_rankings --geneko --submit - python -m ops_model.models.interpretability.diffae.viewer.build_phase_shap_rankings --complex --submit + python -m ops_model.models.interpretability._internal.viewer.build_phase_shap_rankings --geneko --submit + python -m ops_model.models.interpretability._internal.viewer.build_phase_shap_rankings --complex --submit """ from __future__ import annotations diff --git a/src/ops_model/models/interpretability/diffae/viewer/build_phate_figure.py b/src/ops_model/models/interpretability/_internal/viewer/build_phate_figure.py similarity index 99% rename from src/ops_model/models/interpretability/diffae/viewer/build_phate_figure.py rename to src/ops_model/models/interpretability/_internal/viewer/build_phate_figure.py index 0ae77f3..0e99bab 100644 --- a/src/ops_model/models/interpretability/diffae/viewer/build_phate_figure.py +++ b/src/ops_model/models/interpretability/_internal/viewer/build_phate_figure.py @@ -5,7 +5,7 @@ Each panel: the same PHATE scatter (grey), that panel's groups colored + leader-labelled with the single-cell generated morph (NTC cell1 → group, alpha=+5). NTC original shown top-left of panel E. - python -m ops_model.models.interpretability.diffae.viewer.build_phate_figure + python -m ops_model.models.interpretability._internal.viewer.build_phate_figure """ from __future__ import annotations diff --git a/src/ops_model/models/interpretability/diffae/viewer/build_setacc_bins.py b/src/ops_model/models/interpretability/_internal/viewer/build_setacc_bins.py similarity index 100% rename from src/ops_model/models/interpretability/diffae/viewer/build_setacc_bins.py rename to src/ops_model/models/interpretability/_internal/viewer/build_setacc_bins.py diff --git a/src/ops_model/models/interpretability/diffae/viewer/build_setacc_bymarker.py b/src/ops_model/models/interpretability/_internal/viewer/build_setacc_bymarker.py similarity index 100% rename from src/ops_model/models/interpretability/diffae/viewer/build_setacc_bymarker.py rename to src/ops_model/models/interpretability/_internal/viewer/build_setacc_bymarker.py diff --git a/src/ops_model/models/interpretability/diffae/viewer/build_top_cells.py b/src/ops_model/models/interpretability/_internal/viewer/build_top_cells.py similarity index 95% rename from src/ops_model/models/interpretability/diffae/viewer/build_top_cells.py rename to src/ops_model/models/interpretability/_internal/viewer/build_top_cells.py index c61df19..c99c04e 100644 --- a/src/ops_model/models/interpretability/diffae/viewer/build_top_cells.py +++ b/src/ops_model/models/interpretability/_internal/viewer/build_top_cells.py @@ -7,8 +7,8 @@ viewer_assets_v5/top_cells/index.json {"top_n", "genes"|"complexes": {CLASS: {"accuracy": [rec...]}}} viewer_assets_v5/top_cells/crops/.png - python -m ops_model.models.interpretability.diffae.viewer.build_top_cells geneKO # SLURM crop shards + finalize - python -m ops_model.models.interpretability.diffae.viewer.build_top_cells complex --finalize # rebuild index only + python -m ops_model.models.interpretability._internal.viewer.build_top_cells geneKO # SLURM crop shards + finalize + python -m ops_model.models.interpretability._internal.viewer.build_top_cells complex --finalize # rebuild index only """ from __future__ import annotations @@ -17,7 +17,7 @@ import os from . import catalog as C -from .build_pc_crops_masked import BASE, CROP_SIZE, PHASE_CHANNEL, _crop, _is_blank, _render, _render_gray, _overlay_rgba, _zarr_patch +from ops_model.models.interpretability.diffae.traversal.build_pc_crops_masked import BASE, CROP_SIZE, PHASE_CHANNEL, _crop, _is_blank, _render, _render_gray, _overlay_rgba, _zarr_patch TOP_N = 40 diff --git a/src/ops_model/models/interpretability/diffae/viewer/build_umap_montage.py b/src/ops_model/models/interpretability/_internal/viewer/build_umap_montage.py similarity index 98% rename from src/ops_model/models/interpretability/diffae/viewer/build_umap_montage.py rename to src/ops_model/models/interpretability/_internal/viewer/build_umap_montage.py index 8c886c9..f677803 100644 --- a/src/ops_model/models/interpretability/diffae/viewer/build_umap_montage.py +++ b/src/ops_model/models/interpretability/_internal/viewer/build_umap_montage.py @@ -17,8 +17,8 @@ from latent_lens import MontageConfig, build_montage -from ..classifier.config import slugify -from .precompute import VIEWER_ALPHAS +from ops_model.models.interpretability.diffae.classifier.config import slugify +from ops_model.models.interpretability.diffae.traversal.precompute import VIEWER_ALPHAS OUT = "/hpc/projects/icd.fast.ops/models/diffex" import os diff --git a/src/ops_model/models/interpretability/diffae/viewer/deploy/README.md b/src/ops_model/models/interpretability/_internal/viewer/deploy/README.md similarity index 100% rename from src/ops_model/models/interpretability/diffae/viewer/deploy/README.md rename to src/ops_model/models/interpretability/_internal/viewer/deploy/README.md diff --git a/src/ops_model/models/interpretability/diffae/viewer/marker_leaves.py b/src/ops_model/models/interpretability/_internal/viewer/marker_leaves.py similarity index 100% rename from src/ops_model/models/interpretability/diffae/viewer/marker_leaves.py rename to src/ops_model/models/interpretability/_internal/viewer/marker_leaves.py diff --git a/src/ops_model/models/interpretability/diffae/viewer/mimic_alex_embed.py b/src/ops_model/models/interpretability/_internal/viewer/mimic_alex_embed.py similarity index 100% rename from src/ops_model/models/interpretability/diffae/viewer/mimic_alex_embed.py rename to src/ops_model/models/interpretability/_internal/viewer/mimic_alex_embed.py diff --git a/src/ops_model/models/interpretability/diffae/viewer/morphometrics.py b/src/ops_model/models/interpretability/_internal/viewer/morphometrics.py similarity index 100% rename from src/ops_model/models/interpretability/diffae/viewer/morphometrics.py rename to src/ops_model/models/interpretability/_internal/viewer/morphometrics.py diff --git a/src/ops_model/models/interpretability/diffae/viewer/nway_clf.py b/src/ops_model/models/interpretability/_internal/viewer/nway_clf.py similarity index 93% rename from src/ops_model/models/interpretability/diffae/viewer/nway_clf.py rename to src/ops_model/models/interpretability/_internal/viewer/nway_clf.py index 1963bf9..734742a 100644 --- a/src/ops_model/models/interpretability/diffae/viewer/nway_clf.py +++ b/src/ops_model/models/interpretability/_internal/viewer/nway_clf.py @@ -21,10 +21,10 @@ import torch.nn as nn from torch.utils.data import DataLoader, TensorDataset -from ..classifier.celldino_features import embed_crops -from ..classifier.config import Config, GRAINS, slugify -from ..classifier.data import _BASE_COLS, make_labels_df, materialize_crops -from ..classifier.models import MLPHead +from ops_model.models.interpretability.diffae.classifier.celldino_features import embed_crops +from ops_model.models.interpretability.diffae.classifier.config import Config, GRAINS, slugify +from ops_model.models.interpretability.diffae.classifier.data import _BASE_COLS, make_labels_df, materialize_crops +from ops_model.models.interpretability.diffae.classifier.models import MLPHead def _all_class_table(cfg, marker_channel, fluor_csv, n_per_class): diff --git a/src/ops_model/models/interpretability/diffae/viewer/phenotype_cells.py b/src/ops_model/models/interpretability/_internal/viewer/phenotype_cells.py similarity index 96% rename from src/ops_model/models/interpretability/diffae/viewer/phenotype_cells.py rename to src/ops_model/models/interpretability/_internal/viewer/phenotype_cells.py index 3f5e66c..c967d54 100644 --- a/src/ops_model/models/interpretability/diffae/viewer/phenotype_cells.py +++ b/src/ops_model/models/interpretability/_internal/viewer/phenotype_cells.py @@ -16,8 +16,8 @@ import pandas as pd import pyarrow.parquet as pq -from ..classifier.config import PMA_PHASE_EBI, PMA_PHASE_GENEKO -from ..classifier.data import _BASE_COLS +from ops_model.models.interpretability.diffae.classifier.config import PMA_PHASE_EBI, PMA_PHASE_GENEKO +from ops_model.models.interpretability.diffae.classifier.data import _BASE_COLS from . import catalog as C _V4 = os.path.dirname(C.EBI_FLUOR_CSV) diff --git a/src/ops_model/models/interpretability/diffae/viewer/render_montage_scales.py b/src/ops_model/models/interpretability/_internal/viewer/render_montage_scales.py similarity index 99% rename from src/ops_model/models/interpretability/diffae/viewer/render_montage_scales.py rename to src/ops_model/models/interpretability/_internal/viewer/render_montage_scales.py index eb725c8..27d519e 100644 --- a/src/ops_model/models/interpretability/diffae/viewer/render_montage_scales.py +++ b/src/ops_model/models/interpretability/_internal/viewer/render_montage_scales.py @@ -2,7 +2,7 @@ embedding legend (leiden_r4, big dots, NTC as a dark labelled circle). The montage image and its baked viewer-style gene names come straight from the built tiles — finer levels give crisper text. - python -m ops_model.models.interpretability.diffae.viewer.render_montage_scales --alphas 1-5 --levels 3,4 + python -m ops_model.models.interpretability._internal.viewer.render_montage_scales --alphas 1-5 --levels 3,4 Each level of `_montage/phase_geneKO_phate_cell1_a_tiles/L/` is a level-of-detail montage (coarse levels show a decimated non-overlapping subset; finer levels fill in more cells at higher res). @@ -225,7 +225,7 @@ def render_composed(alphas, cell=1, level=4, crop=256, ppu=5600, embedding="phat import anndata as ad from latent_lens.grid import assign_cells_to_grid, compute_priority, compute_canvas_size from .build_umap_montage import _embed_coords, OUT as MOUT, VIEWER_ALPHAS - from ..classifier.config import slugify + from ops_model.models.interpretability.diffae.classifier.config import slugify out_dir = out_dir or COMPOSED_DIR os.makedirs(out_dir, exist_ok=True) diff --git a/src/ops_model/models/interpretability/diffae/viewer/score_generated.py b/src/ops_model/models/interpretability/_internal/viewer/score_generated.py similarity index 91% rename from src/ops_model/models/interpretability/diffae/viewer/score_generated.py rename to src/ops_model/models/interpretability/_internal/viewer/score_generated.py index b1ce894..20384e2 100644 --- a/src/ops_model/models/interpretability/diffae/viewer/score_generated.py +++ b/src/ops_model/models/interpretability/_internal/viewer/score_generated.py @@ -37,8 +37,8 @@ def diagnose(gene="TOMM20", device="cuda"): (raw, and z-standardized on NTC control), vs Alex's own val embeddings (positive control).""" import torch from .set_classifier import load_set_classifier, score_bags, V5_CKPT_ROOT, V5_RUNS - from .precompute import _gather_class - from ..directions.config import DirConfig + from ops_model.models.interpretability.diffae.traversal.precompute import _gather_class + from ops_model.models.interpretability.diffae.directions.config import DirConfig m, g2i, c2i = load_set_classifier(run=V5_RUNS[("phase", "geneKO")], device=device, root=V5_CKPT_ROOT) ci = c2i.get("Phase2D", 0); gi = g2i[gene] # Alex val (already z-std) — positive control @@ -62,8 +62,8 @@ def diagnose2(gene="TOMM20", device="cuda"): Reports P(target) vs α for both references + whether generated α0 is recognized as NTC.""" import torch, json from .set_classifier import load_set_classifier, score_bags, V5_CKPT_ROOT, V5_RUNS - from ..directions.config import DirConfig - from ..classifier.celldino_features import embed_crops + from ops_model.models.interpretability.diffae.directions.config import DirConfig + from ops_model.models.interpretability.diffae.classifier.celldino_features import embed_crops m, g2i, c2i = load_set_classifier(run=V5_RUNS[("phase", "geneKO")], device=device, root=V5_CKPT_ROOT) ci = c2i.get("Phase2D", 0); gi = g2i[gene]; ni = g2i.get("NTC") trav = f"{V5_BASE}/geneKO/{gene}"; alphas = json.load(open(f"{trav}/meta.json"))["alphas"] @@ -100,12 +100,12 @@ def bag_experiment(grain="geneKO", target="MICOS13", n_max=200, sizes=(20, 50, 1 generated top1_acc + mean P(target), vs Alex's REAL top1_acc-by-bag (gene row, or mean-member for a complex). Writes _bagexp_.json. Isolated OPS_DIFFEX_ASSETS (bagtest dir set by caller).""" import json, torch - from .precompute import precompute_marker + from ops_model.models.interpretability.diffae.traversal.precompute import precompute_marker from .submit import PHASE_CK from .set_classifier import load_set_classifier, score_bags, V5_CKPT_ROOT, V5_RUNS - from ..directions.config import DirConfig - from ..classifier.celldino_features import embed_crops - from ..classifier.config import slugify + from ops_model.models.interpretability.diffae.directions.config import DirConfig + from ops_model.models.interpretability.diffae.classifier.celldino_features import embed_crops + from ops_model.models.interpretability.diffae.classifier.config import slugify OUT = "/hpc/projects/icd.fast.ops/models/diffex" run = V5_RUNS[("phase", "geneKO" if grain == "geneKO" else "complex_ebionly")] precompute_marker(grain=grain, targets=[target], ckpt=PHASE_CK, out_root=OUT, n_cells=n_max, @@ -176,12 +176,12 @@ def bag_scaling(grain, targets, n_max=500, sizes=(20, 50, 100, 200), n_bags=30, bag size sampling n_bags DISTINCT bags per size (resampled from the pool, like Alex's real-cell eval) → generated top1_acc + mean P(target) vs Alex's real expectation. Writes _bagexp_.json in bagtest.""" import json, torch - from .precompute import precompute_marker + from ops_model.models.interpretability.diffae.traversal.precompute import precompute_marker from .submit import PHASE_CK from .set_classifier import load_set_classifier, score_bags, V5_CKPT_ROOT, V5_RUNS - from ..directions.config import DirConfig - from ..classifier.celldino_features import embed_crops - from ..classifier.config import slugify + from ops_model.models.interpretability.diffae.directions.config import DirConfig + from ops_model.models.interpretability.diffae.classifier.celldino_features import embed_crops + from ops_model.models.interpretability.diffae.classifier.config import slugify OUT = "/hpc/projects/icd.fast.ops/models/diffex" run = V5_RUNS[("phase", "geneKO" if grain == "geneKO" else "complex_ebionly")] model, cmap, c2i = load_set_classifier(run=run, device=device, root=V5_CKPT_ROOT) @@ -224,9 +224,9 @@ def score_targets(grain, targets, device="cuda", run=None, members_map=None, bag Writes scores_v5.json per traversal dir. Returns {target: p_target-per-α}.""" import torch # noqa from .set_classifier import load_set_classifier, score_bags, V5_CKPT_ROOT, V5_RUNS - from ..directions.config import DirConfig - from ..classifier.celldino_features import embed_crops - from ..classifier.config import slugify + from ops_model.models.interpretability.diffae.directions.config import DirConfig + from ops_model.models.interpretability.diffae.classifier.celldino_features import embed_crops + from ops_model.models.interpretability.diffae.classifier.config import slugify run = run or V5_RUNS[("phase", "geneKO" if grain == "geneKO" else "complex_ebionly")] model, g2i, c2i = load_set_classifier(run=run, device=device, root=V5_CKPT_ROOT) ci = c2i.get("Phase2D", 0) @@ -259,8 +259,8 @@ def score_anchor_traversals(grain, device="cuda"): α0 (= anchor-A) frames — same recipe as the NTC traversals, just target=B. Writes scores_v5.json per dir.""" import glob from .set_classifier import load_set_classifier, V5_CKPT_ROOT, V5_RUNS - from ..directions.config import DirConfig - from ..classifier.celldino_features import embed_crops + from ops_model.models.interpretability.diffae.directions.config import DirConfig + from ops_model.models.interpretability.diffae.classifier.celldino_features import embed_crops run = V5_RUNS[("phase", "geneKO" if grain == "geneKO" else "complex_ebionly")] model, g2i, c2i = load_set_classifier(run=run, device=device, root=V5_CKPT_ROOT) ci = c2i.get("Phase2D", 0) diff --git a/src/ops_model/models/interpretability/diffae/viewer/set_classifier.py b/src/ops_model/models/interpretability/_internal/viewer/set_classifier.py similarity index 100% rename from src/ops_model/models/interpretability/diffae/viewer/set_classifier.py rename to src/ops_model/models/interpretability/_internal/viewer/set_classifier.py diff --git a/src/ops_model/models/interpretability/diffae/viewer/submit.py b/src/ops_model/models/interpretability/_internal/viewer/submit.py similarity index 95% rename from src/ops_model/models/interpretability/diffae/viewer/submit.py rename to src/ops_model/models/interpretability/_internal/viewer/submit.py index eee7314..18d3244 100644 --- a/src/ops_model/models/interpretability/diffae/viewer/submit.py +++ b/src/ops_model/models/interpretability/_internal/viewer/submit.py @@ -1,10 +1,10 @@ """Build the DiffEx viewer cache — reproducible, version-controlled entrypoint (replaces the one-off scratchpad drivers). All target selection comes from `catalog.py`. - python -m ops_model.models.interpretability.diffae.viewer.submit seed # per-marker NTC traversals - python -m ops_model.models.interpretability.diffae.viewer.submit anchors --k 5 # A→B anchor pairs - python -m ops_model.models.interpretability.diffae.viewer.submit manifest # rebuild manifest.json (local) - python -m ops_model.models.interpretability.diffae.viewer.submit montage --cell 0 --alpha 2 # harvest cache -> UMAP montage zarr + python -m ops_model.models.interpretability._internal.viewer.submit seed # per-marker NTC traversals + python -m ops_model.models.interpretability._internal.viewer.submit anchors --k 5 # A→B anchor pairs + python -m ops_model.models.interpretability._internal.viewer.submit manifest # rebuild manifest.json (local) + python -m ops_model.models.interpretability._internal.viewer.submit montage --cell 0 --alpha 2 # harvest cache -> UMAP montage zarr """ from __future__ import annotations @@ -15,10 +15,10 @@ from ops_utils.hpc.slurm_batch_utils import submit_parallel_jobs -from ..classifier.config import slugify +from ops_model.models.interpretability.diffae.classifier.config import slugify from . import catalog as C from .build_umap_montage import build_montage_grid, build_montage_web -from .precompute import build_manifest, precompute_marker, precompute_target +from ops_model.models.interpretability.diffae.traversal.precompute import build_manifest, precompute_marker, precompute_target PHASE_CK = f"{C.DD}/phase_v1/diffae_best.pt" UMAP_H5AD = ("/hpc/projects/icd.fast.ops/organelle_attribution/pca_optimized_v0.3/cell_dino/" @@ -211,7 +211,7 @@ def cmd_phase_morpho(args): cached production-label real-cell images. Keys from MORPHO_TARGETS (e.g. MICOS13 TOMM20 CCT). One SLURM job per target, run in PARALLEL (each target stages its own per-target zarr → no shared-zarr race). --parallel caps concurrency (default = all targets at once).""" - from .morpho_pipeline import build_morpho_target, MORPHO_TARGETS + from ops_model.models.interpretability.diffae.traversal.morpho_pipeline import build_morpho_target, MORPHO_TARGETS targets = args.targets or list(MORPHO_TARGETS) jobs = [_job(f"pm_{slugify(t)}", build_morpho_target, dict(key=t, n_cells=args.n_cells), "phase_morpho") for t in targets] print(f"phase-morpho: {len(jobs)} target(s) in parallel (cap {args.parallel or len(jobs)}) -> {targets}") diff --git a/src/ops_model/models/interpretability/diffae/viewer/webapp/app.js b/src/ops_model/models/interpretability/_internal/viewer/webapp/app.js similarity index 100% rename from src/ops_model/models/interpretability/diffae/viewer/webapp/app.js rename to src/ops_model/models/interpretability/_internal/viewer/webapp/app.js diff --git a/src/ops_model/models/interpretability/diffae/viewer/webapp/biohub-mark.png b/src/ops_model/models/interpretability/_internal/viewer/webapp/biohub-mark.png similarity index 100% rename from src/ops_model/models/interpretability/diffae/viewer/webapp/biohub-mark.png rename to src/ops_model/models/interpretability/_internal/viewer/webapp/biohub-mark.png diff --git a/src/ops_model/models/interpretability/diffae/viewer/webapp/biohub-wordmark.png b/src/ops_model/models/interpretability/_internal/viewer/webapp/biohub-wordmark.png similarity index 100% rename from src/ops_model/models/interpretability/diffae/viewer/webapp/biohub-wordmark.png rename to src/ops_model/models/interpretability/_internal/viewer/webapp/biohub-wordmark.png diff --git a/src/ops_model/models/interpretability/diffae/viewer/webapp/build_gene_narratives.py b/src/ops_model/models/interpretability/_internal/viewer/webapp/build_gene_narratives.py similarity index 100% rename from src/ops_model/models/interpretability/diffae/viewer/webapp/build_gene_narratives.py rename to src/ops_model/models/interpretability/_internal/viewer/webapp/build_gene_narratives.py diff --git a/src/ops_model/models/interpretability/diffae/viewer/webapp/gif.js b/src/ops_model/models/interpretability/_internal/viewer/webapp/gif.js similarity index 100% rename from src/ops_model/models/interpretability/diffae/viewer/webapp/gif.js rename to src/ops_model/models/interpretability/_internal/viewer/webapp/gif.js diff --git a/src/ops_model/models/interpretability/diffae/viewer/webapp/gif.worker.js b/src/ops_model/models/interpretability/_internal/viewer/webapp/gif.worker.js similarity index 100% rename from src/ops_model/models/interpretability/diffae/viewer/webapp/gif.worker.js rename to src/ops_model/models/interpretability/_internal/viewer/webapp/gif.worker.js diff --git a/src/ops_model/models/interpretability/diffae/viewer/webapp/index.html b/src/ops_model/models/interpretability/_internal/viewer/webapp/index.html similarity index 100% rename from src/ops_model/models/interpretability/diffae/viewer/webapp/index.html rename to src/ops_model/models/interpretability/_internal/viewer/webapp/index.html diff --git a/src/ops_model/models/interpretability/diffae/viewer/webapp/methods.js b/src/ops_model/models/interpretability/_internal/viewer/webapp/methods.js similarity index 100% rename from src/ops_model/models/interpretability/diffae/viewer/webapp/methods.js rename to src/ops_model/models/interpretability/_internal/viewer/webapp/methods.js diff --git a/src/ops_model/models/interpretability/diffae/viewer/webapp/morpho_demo.html b/src/ops_model/models/interpretability/_internal/viewer/webapp/morpho_demo.html similarity index 100% rename from src/ops_model/models/interpretability/diffae/viewer/webapp/morpho_demo.html rename to src/ops_model/models/interpretability/_internal/viewer/webapp/morpho_demo.html diff --git a/src/ops_model/models/interpretability/diffae/viewer/webapp/openseadragon.min.js b/src/ops_model/models/interpretability/_internal/viewer/webapp/openseadragon.min.js similarity index 100% rename from src/ops_model/models/interpretability/diffae/viewer/webapp/openseadragon.min.js rename to src/ops_model/models/interpretability/_internal/viewer/webapp/openseadragon.min.js diff --git a/src/ops_model/models/interpretability/diffae/viewer/webapp/opsin-eyes.svg b/src/ops_model/models/interpretability/_internal/viewer/webapp/opsin-eyes.svg similarity index 100% rename from src/ops_model/models/interpretability/diffae/viewer/webapp/opsin-eyes.svg rename to src/ops_model/models/interpretability/_internal/viewer/webapp/opsin-eyes.svg diff --git a/src/ops_model/models/interpretability/diffae/viewer/webapp/style.css b/src/ops_model/models/interpretability/_internal/viewer/webapp/style.css similarity index 100% rename from src/ops_model/models/interpretability/diffae/viewer/webapp/style.css rename to src/ops_model/models/interpretability/_internal/viewer/webapp/style.css diff --git a/src/ops_model/models/interpretability/diffae/directions/proto_ddim_anchors.py b/src/ops_model/models/interpretability/diffae/directions/proto_ddim_anchors.py index 581f939..5b99ee1 100644 --- a/src/ops_model/models/interpretability/diffae/directions/proto_ddim_anchors.py +++ b/src/ops_model/models/interpretability/diffae/directions/proto_ddim_anchors.py @@ -62,7 +62,7 @@ def _direction(cfg, gene, modality, ctrl_embs, mu_ctrl, gather, dev): @torch.no_grad() def run(specs, n_cells=4, alphas=(0, 1, 2, 3, 4, 5), ws=(1.0, 1.5, 2.0, 3.0), out_name="ddim_anchors", device="cuda"): - from ..viewer.precompute import _gather_class # gather NTC anchors (imgs + CellDINO embs) + from ..traversal.precompute import _gather_class # gather NTC anchors (imgs + CellDINO embs) dev = torch.device(device if torch.cuda.is_available() else "cpu") out = Path(ANALYSIS) / out_name; out.mkdir(parents=True, exist_ok=True) import matplotlib diff --git a/src/ops_model/models/interpretability/diffae/figures/_setacc_common.py b/src/ops_model/models/interpretability/diffae/figures/_setacc_common.py index a0f1d15..b7f3560 100644 --- a/src/ops_model/models/interpretability/diffae/figures/_setacc_common.py +++ b/src/ops_model/models/interpretability/diffae/figures/_setacc_common.py @@ -13,8 +13,8 @@ from ops_model.models.interpretability.diffae.classifier.config import slugify from ops_model.models.interpretability.diffae.classifier.data import make_labels_df, materialize_crops from ops_model.models.interpretability.diffae.directions.config import DirConfig -from ops_model.models.interpretability.diffae.viewer._fluor_topcells import _overlay_rgba -from ops_model.models.interpretability.diffae.viewer.build_pc_crops_masked import BASE, CROP_SIZE, _crop, _zarr_patch +from ops_model.models.interpretability.diffae.traversal._fluor_topcells import _overlay_rgba +from ops_model.models.interpretability.diffae.traversal.build_pc_crops_masked import BASE, CROP_SIZE, _crop, _zarr_patch OUT = "/hpc/projects/icd.fast.ops/analysis/figure4_setacc_panel" RANK_BASE = "/hpc/projects/icd.fast.ops/models/diffex/viewer_assets_v5/_rankings/fluor_shap" diff --git a/src/ops_model/models/interpretability/diffae/figures/auto_pick_and_plot.py b/src/ops_model/models/interpretability/diffae/figures/auto_pick_and_plot.py index a9f5e52..4f2b704 100644 --- a/src/ops_model/models/interpretability/diffae/figures/auto_pick_and_plot.py +++ b/src/ops_model/models/interpretability/diffae/figures/auto_pick_and_plot.py @@ -10,7 +10,7 @@ sys.path.insert(0, os.path.dirname(os.path.abspath(__file__))) os.environ.setdefault("OPS_DIFFEX_ASSETS", "viewer_assets_v5") -from ops_model.models.interpretability.diffae.viewer.morpho_pipeline import MORPHO_TARGETS +from ops_model.models.interpretability.diffae.traversal.morpho_pipeline import MORPHO_TARGETS from ops_model.models.interpretability.diffae.classifier.config import slugify VA = "/hpc/projects/icd.fast.ops/models/diffex/viewer_assets_v5/_morphometrics" diff --git a/src/ops_model/models/interpretability/diffae/figures/ebi_peripheral_droplets.py b/src/ops_model/models/interpretability/diffae/figures/ebi_peripheral_droplets.py index 2eb9bb8..b2f709f 100644 --- a/src/ops_model/models/interpretability/diffae/figures/ebi_peripheral_droplets.py +++ b/src/ops_model/models/interpretability/diffae/figures/ebi_peripheral_droplets.py @@ -41,7 +41,7 @@ from figure_ebi_morpho_violin import draw_violin from figure_multirank_ebi_grid import CACHE, OUT, ebi_rows, top_rows -from ops_model.models.interpretability.diffae.viewer.build_pc_crops_masked import BASE, _crop, _zarr_patch +from ops_model.models.interpretability.diffae.traversal.build_pc_crops_masked import BASE, _crop, _zarr_patch from organelle_profiler.feature_extraction.localization_features import compute_localization_features diff --git a/src/ops_model/models/interpretability/diffae/figures/figure4_morpho_violin.py b/src/ops_model/models/interpretability/diffae/figures/figure4_morpho_violin.py index 036e31e..a569244 100644 --- a/src/ops_model/models/interpretability/diffae/figures/figure4_morpho_violin.py +++ b/src/ops_model/models/interpretability/diffae/figures/figure4_morpho_violin.py @@ -17,7 +17,7 @@ import pandas as pd from figure4_morpho_traversal import FIGURES, VA, image_panels, render_images -from ops_model.models.interpretability.diffae.viewer.morpho_pipeline import MORPHO_TARGETS, real_percell +from ops_model.models.interpretability.diffae.traversal.morpho_pipeline import MORPHO_TARGETS, real_percell from ops_model.models.interpretability.diffae.classifier.config import slugify plt.rcParams["pdf.fonttype"] = 42 diff --git a/src/ops_model/models/interpretability/diffae/figures/gen_validation/bag_sweep_score.py b/src/ops_model/models/interpretability/diffae/figures/gen_validation/bag_sweep_score.py index 82e6cdb..9a231d5 100644 --- a/src/ops_model/models/interpretability/diffae/figures/gen_validation/bag_sweep_score.py +++ b/src/ops_model/models/interpretability/diffae/figures/gen_validation/bag_sweep_score.py @@ -16,8 +16,8 @@ def score_shard(genes): import torch - from ops_model.models.interpretability.diffae.viewer.score_generated import score_embs_v5 - from ops_model.models.interpretability.diffae.viewer.set_classifier import load_set_classifier, V5_CKPT_ROOT, V5_RUNS + from ops_model.models.interpretability._internal.viewer.score_generated import score_embs_v5 + from ops_model.models.interpretability._internal.viewer.set_classifier import load_set_classifier, V5_CKPT_ROOT, V5_RUNS os.makedirs(OUT, exist_ok=True) dev = "cuda" if torch.cuda.is_available() else "cpu" from ops_model.models.interpretability.diffae.classifier.config import slugify diff --git a/src/ops_model/models/interpretability/diffae/figures/gen_validation/embcheck.py b/src/ops_model/models/interpretability/diffae/figures/gen_validation/embcheck.py index 4e4868c..087c6e0 100644 --- a/src/ops_model/models/interpretability/diffae/figures/gen_validation/embcheck.py +++ b/src/ops_model/models/interpretability/diffae/figures/gen_validation/embcheck.py @@ -8,7 +8,7 @@ def check(gene="AACS", ai=6): - from ops_model.models.interpretability.diffae.viewer.score_generated import _emb_frames + from ops_model.models.interpretability._internal.viewer.score_generated import _emb_frames from ops_model.models.interpretability.diffae.classifier.celldino_features import embed_crops from ops_model.models.interpretability.diffae.directions.config import DirConfig trav = f"{B}/viewer_assets_valid200/phase/geneKO/{gene}" diff --git a/src/ops_model/models/interpretability/diffae/figures/gen_validation/gen_real_centroid.py b/src/ops_model/models/interpretability/diffae/figures/gen_validation/gen_real_centroid.py index 0ca9896..b232fe5 100644 --- a/src/ops_model/models/interpretability/diffae/figures/gen_validation/gen_real_centroid.py +++ b/src/ops_model/models/interpretability/diffae/figures/gen_validation/gen_real_centroid.py @@ -24,7 +24,7 @@ def _classes(grain): def embed_centroids(grain, classes): - from ops_model.models.interpretability.diffae.viewer.precompute import _gather_class + from ops_model.models.interpretability.diffae.traversal.precompute import _gather_class from ops_model.models.interpretability.diffae.directions.config import DirConfig from ops_model.models.interpretability.diffae.classifier.config import slugify cfg = DirConfig(grain=grain, target=classes[0], device="cuda"); cfg.num_workers = 12 diff --git a/src/ops_model/models/interpretability/diffae/figures/gen_validation/ntc_inverse_gap.py b/src/ops_model/models/interpretability/diffae/figures/gen_validation/ntc_inverse_gap.py index 9f67deb..db7e9da 100644 --- a/src/ops_model/models/interpretability/diffae/figures/gen_validation/ntc_inverse_gap.py +++ b/src/ops_model/models/interpretability/diffae/figures/gen_validation/ntc_inverse_gap.py @@ -124,7 +124,7 @@ def webp_ab(): import torch, tempfile # noqa from ops_model.models.interpretability.diffae.classifier.celldino_features import embed_crops from ops_model.models.interpretability.diffae.directions.config import DirConfig - from ops_model.models.interpretability.diffae.viewer.precompute import _save_webp + from ops_model.models.interpretability.diffae.traversal.precompute import _save_webp os.makedirs(OUT, exist_ok=True) d = np.load(CTRL, allow_pickle=True) imgs = d["anchor_imgs"].astype(np.float32) # (45,1,160,160) float [-1,1] diff --git a/src/ops_model/models/interpretability/diffae/figures/gen_validation/patch_cache_real.py b/src/ops_model/models/interpretability/diffae/figures/gen_validation/patch_cache_real.py index 3122e5b..6897ff7 100644 --- a/src/ops_model/models/interpretability/diffae/figures/gen_validation/patch_cache_real.py +++ b/src/ops_model/models/interpretability/diffae/figures/gen_validation/patch_cache_real.py @@ -20,7 +20,7 @@ def _drop42(): def run(): import pandas as pd - from ops_model.models.interpretability.diffae.viewer.precompute import _gather_class + from ops_model.models.interpretability.diffae.traversal.precompute import _gather_class from ops_model.models.interpretability.diffae.directions.config import DirConfig from ops_model.models.interpretability.diffae.classifier.config import slugify genes = _drop42() diff --git a/src/ops_model/models/interpretability/diffae/figures/gen_validation/st_halves_score.py b/src/ops_model/models/interpretability/diffae/figures/gen_validation/st_halves_score.py index ed8e5b2..3e41527 100644 --- a/src/ops_model/models/interpretability/diffae/figures/gen_validation/st_halves_score.py +++ b/src/ops_model/models/interpretability/diffae/figures/gen_validation/st_halves_score.py @@ -12,8 +12,8 @@ def score_shard(genes): import torch - from ops_model.models.interpretability.diffae.viewer.score_generated import score_embs_v5 - from ops_model.models.interpretability.diffae.viewer.set_classifier import load_set_classifier, V5_CKPT_ROOT, V5_RUNS + from ops_model.models.interpretability._internal.viewer.score_generated import score_embs_v5 + from ops_model.models.interpretability._internal.viewer.set_classifier import load_set_classifier, V5_CKPT_ROOT, V5_RUNS from ops_model.models.interpretability.diffae.classifier.config import slugify os.makedirs(OUT, exist_ok=True) dev = "cuda" if torch.cuda.is_available() else "cpu" diff --git a/src/ops_model/models/interpretability/diffae/figures/rebuild_traversals_n100.py b/src/ops_model/models/interpretability/diffae/figures/rebuild_traversals_n100.py index c74f8fe..2c46962 100644 --- a/src/ops_model/models/interpretability/diffae/figures/rebuild_traversals_n100.py +++ b/src/ops_model/models/interpretability/diffae/figures/rebuild_traversals_n100.py @@ -13,7 +13,7 @@ import sys from pathlib import Path -from ops_model.models.interpretability.diffae.viewer import catalog as C +from ops_model.models.interpretability.diffae.traversal import catalog as C from ops_model.models.interpretability.diffae.classifier.config import slugify ASSETS = "viewer_assets_v5" @@ -81,7 +81,7 @@ def _clear_anchor(modality): def rebuild_marker(modality): os.environ["OPS_DIFFEX_ASSETS"] = ASSETS - from ops_model.models.interpretability.diffae.viewer import precompute as P + from ops_model.models.interpretability.diffae.traversal import precompute as P P._ASSETS = ASSETS _clear_anchor(modality) tg = TARGETS[modality] @@ -105,7 +105,7 @@ def rebuild_marker(modality): def gen_phase(targets): """Generate specific phase geneKO targets at n=100 (reuses the existing 100-cell phase anchor cache).""" os.environ["OPS_DIFFEX_ASSETS"] = ASSETS - from ops_model.models.interpretability.diffae.viewer import precompute as P + from ops_model.models.interpretability.diffae.traversal import precompute as P P._ASSETS = ASSETS _clear_anchor("phase") # idempotent: keeps the 100-cell cache P.precompute_marker(grain="geneKO", targets=targets, ckpt=PHASE_CK, out_root=C.OUT, @@ -115,7 +115,7 @@ def gen_phase(targets): def gen_phase_complex(targets): """Generate phase COMPLEX targets at n=100 (full complex names; reuses the 100-cell phase anchor cache).""" os.environ["OPS_DIFFEX_ASSETS"] = ASSETS - from ops_model.models.interpretability.diffae.viewer import precompute as P + from ops_model.models.interpretability.diffae.traversal import precompute as P P._ASSETS = ASSETS _clear_anchor("phase") P.precompute_marker(grain="complex", targets=targets, ckpt=PHASE_CK, out_root=C.OUT, @@ -139,7 +139,7 @@ def gen_phase_chunk(targets, cell_range): cached direction (both ckpt-independent), each writing its own cell{c} dirs → parallelizes the per-cell DDIM inversion across GPUs instead of one long serial job. v5 scoring skipped (whole-target only).""" os.environ["OPS_DIFFEX_ASSETS"] = ASSETS - from ops_model.models.interpretability.diffae.viewer import precompute as P + from ops_model.models.interpretability.diffae.traversal import precompute as P P._ASSETS = ASSETS _clear_anchor("phase") # idempotent: 100-anchor cache kept P.precompute_marker(grain="geneKO", targets=targets, ckpt=PHASE_CK, out_root=C.OUT, diff --git a/src/ops_model/models/interpretability/diffae/generator/virtstain_multi.py b/src/ops_model/models/interpretability/diffae/generator/virtstain_multi.py index d984689..5be487e 100644 --- a/src/ops_model/models/interpretability/diffae/generator/virtstain_multi.py +++ b/src/ops_model/models/interpretability/diffae/generator/virtstain_multi.py @@ -33,7 +33,7 @@ def markers_list(): CP1_/CP2_/4i_R*) are stained AFTER the live phase acquisition, so their cells have moved/changed and the phase→marker registration is broken — that misalignment poisons the spatial conditioning, so they are excluded. Live channels (GFP/mCherry/Cy5/farred) are imaged concurrently with phase → registered.""" - from ..viewer import catalog as C + from ..traversal import catalog as C seen, out = set(), [] for d, mc, ch in C.complete_markers(): if ch.startswith(("CP", "4i")): diff --git a/src/ops_model/models/interpretability/diffae/traversal/__init__.py b/src/ops_model/models/interpretability/diffae/traversal/__init__.py new file mode 100644 index 0000000..c56afc8 --- /dev/null +++ b/src/ops_model/models/interpretability/diffae/traversal/__init__.py @@ -0,0 +1,3 @@ +"""Shared traversal/asset computation used by the paper figures and the (internal) viewer: +catalog selection, per-cell precompute, morphometrics, fluor top-cells, PC-strip crops. +Depends only on the core stages (classifier/generator/directions).""" diff --git a/src/ops_model/models/interpretability/diffae/viewer/_fluor_topcells.py b/src/ops_model/models/interpretability/diffae/traversal/_fluor_topcells.py similarity index 100% rename from src/ops_model/models/interpretability/diffae/viewer/_fluor_topcells.py rename to src/ops_model/models/interpretability/diffae/traversal/_fluor_topcells.py diff --git a/src/ops_model/models/interpretability/diffae/viewer/build_pc_crops_masked.py b/src/ops_model/models/interpretability/diffae/traversal/build_pc_crops_masked.py similarity index 97% rename from src/ops_model/models/interpretability/diffae/viewer/build_pc_crops_masked.py rename to src/ops_model/models/interpretability/diffae/traversal/build_pc_crops_masked.py index 3938deb..9172482 100644 --- a/src/ops_model/models/interpretability/diffae/viewer/build_pc_crops_masked.py +++ b/src/ops_model/models/interpretability/diffae/traversal/build_pc_crops_masked.py @@ -10,8 +10,8 @@ {BASE}/{exp}/3-assembly/phenotyping_v3.zarr/{row}/{col}/0/0 image [1,C,1,Y,X], Phase2D=ch0 {BASE}/{exp}/3-assembly/phenotyping_v3.zarr/{row}/{col}/0/labels/cell_seg/0 int32 labels - python -m ops_model.models.interpretability.diffae.viewer.build_pc_crops_masked --sample 24 # preview - python -m ops_model.models.interpretability.diffae.viewer.build_pc_crops_masked # full (overwrites crops/) + python -m ops_model.models.interpretability.diffae.traversal.build_pc_crops_masked --sample 24 # preview + python -m ops_model.models.interpretability.diffae.traversal.build_pc_crops_masked # full (overwrites crops/) """ from __future__ import annotations diff --git a/src/ops_model/models/interpretability/diffae/viewer/catalog.py b/src/ops_model/models/interpretability/diffae/traversal/catalog.py similarity index 100% rename from src/ops_model/models/interpretability/diffae/viewer/catalog.py rename to src/ops_model/models/interpretability/diffae/traversal/catalog.py diff --git a/src/ops_model/models/interpretability/diffae/viewer/morpho_pipeline.py b/src/ops_model/models/interpretability/diffae/traversal/morpho_pipeline.py similarity index 100% rename from src/ops_model/models/interpretability/diffae/viewer/morpho_pipeline.py rename to src/ops_model/models/interpretability/diffae/traversal/morpho_pipeline.py diff --git a/src/ops_model/models/interpretability/diffae/viewer/precompute.py b/src/ops_model/models/interpretability/diffae/traversal/precompute.py similarity index 100% rename from src/ops_model/models/interpretability/diffae/viewer/precompute.py rename to src/ops_model/models/interpretability/diffae/traversal/precompute.py From f822ba76ff3828aeb6cd11d1b335eb2566d77b54 Mon Sep 17 00:00:00 2001 From: Gav Sturm Date: Tue, 11 Aug 2026 09:42:34 -0700 Subject: [PATCH 06/13] move figures/gen_validation (27 diagnostic scripts) -> _internal/gen_validation MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Diagnostic/validation probes (bag sweeps, centroid recovery, valid200) — not published figures. Absolute imports only, so clean move; self-path -m refs repointed. NOTE: 3 figure scripts (mtor_lyso_loc_traversal, mtor_lysosome_localization, real_gen_measure) are untracked WIP on diffex-interpretability -> not vendored. --- .../figures => _internal}/gen_validation/bag_sweep_plots.py | 0 .../figures => _internal}/gen_validation/bag_sweep_score.py | 0 .../figures => _internal}/gen_validation/centroid_bagsweep.py | 0 .../figures => _internal}/gen_validation/centroid_halves.py | 0 .../gen_validation/centroid_pooled_bagsweep.py | 0 .../figures => _internal}/gen_validation/control_halves_zscore.py | 0 .../{diffae/figures => _internal}/gen_validation/embcheck.py | 0 .../figures => _internal}/gen_validation/embedding_diagnostics.py | 0 .../gen_validation/figure4_bagsize_reachfrac.py | 0 .../gen_validation/figure4_v5_accuracy_summary.py | 0 .../figures => _internal}/gen_validation/gen_alpha_embedding.py | 0 .../figures => _internal}/gen_validation/gen_embed_refit.py | 0 .../figures => _internal}/gen_validation/gen_phate_passthrough.py | 0 .../figures => _internal}/gen_validation/gen_real_centroid.py | 0 .../figures => _internal}/gen_validation/gen_real_distinct.py | 0 .../figures => _internal}/gen_validation/ntc_inverse_gap.py | 0 .../figures => _internal}/gen_validation/patch_cache_real.py | 0 .../figures => _internal}/gen_validation/publish_multibag_page.py | 0 .../{diffae/figures => _internal}/gen_validation/rank_summary.py | 0 .../figures => _internal}/gen_validation/st_halves_score.py | 0 .../figures => _internal}/gen_validation/std_anchor_test.py | 0 .../figures => _internal}/gen_validation/stepabl_compare.py | 0 .../figures => _internal}/gen_validation/valid200_alphastep.py | 0 .../figures => _internal}/gen_validation/valid200_cache_build.py | 0 .../figures => _internal}/gen_validation/valid200_capcheck.py | 0 .../figures => _internal}/gen_validation/valid200_map_compare.py | 0 .../figures => _internal}/gen_validation/valid200_metrics.py | 0 27 files changed, 0 insertions(+), 0 deletions(-) rename src/ops_model/models/interpretability/{diffae/figures => _internal}/gen_validation/bag_sweep_plots.py (100%) rename src/ops_model/models/interpretability/{diffae/figures => _internal}/gen_validation/bag_sweep_score.py (100%) rename src/ops_model/models/interpretability/{diffae/figures => _internal}/gen_validation/centroid_bagsweep.py (100%) rename src/ops_model/models/interpretability/{diffae/figures => _internal}/gen_validation/centroid_halves.py (100%) rename src/ops_model/models/interpretability/{diffae/figures => _internal}/gen_validation/centroid_pooled_bagsweep.py (100%) rename src/ops_model/models/interpretability/{diffae/figures => _internal}/gen_validation/control_halves_zscore.py (100%) rename src/ops_model/models/interpretability/{diffae/figures => _internal}/gen_validation/embcheck.py (100%) rename src/ops_model/models/interpretability/{diffae/figures => _internal}/gen_validation/embedding_diagnostics.py (100%) rename src/ops_model/models/interpretability/{diffae/figures => _internal}/gen_validation/figure4_bagsize_reachfrac.py (100%) rename src/ops_model/models/interpretability/{diffae/figures => _internal}/gen_validation/figure4_v5_accuracy_summary.py (100%) rename src/ops_model/models/interpretability/{diffae/figures => _internal}/gen_validation/gen_alpha_embedding.py (100%) rename src/ops_model/models/interpretability/{diffae/figures => _internal}/gen_validation/gen_embed_refit.py (100%) rename src/ops_model/models/interpretability/{diffae/figures => _internal}/gen_validation/gen_phate_passthrough.py (100%) rename src/ops_model/models/interpretability/{diffae/figures => _internal}/gen_validation/gen_real_centroid.py (100%) rename src/ops_model/models/interpretability/{diffae/figures => _internal}/gen_validation/gen_real_distinct.py (100%) rename src/ops_model/models/interpretability/{diffae/figures => _internal}/gen_validation/ntc_inverse_gap.py (100%) rename src/ops_model/models/interpretability/{diffae/figures => _internal}/gen_validation/patch_cache_real.py (100%) rename src/ops_model/models/interpretability/{diffae/figures => _internal}/gen_validation/publish_multibag_page.py (100%) rename src/ops_model/models/interpretability/{diffae/figures => _internal}/gen_validation/rank_summary.py (100%) rename src/ops_model/models/interpretability/{diffae/figures => _internal}/gen_validation/st_halves_score.py (100%) rename src/ops_model/models/interpretability/{diffae/figures => _internal}/gen_validation/std_anchor_test.py (100%) rename src/ops_model/models/interpretability/{diffae/figures => _internal}/gen_validation/stepabl_compare.py (100%) rename src/ops_model/models/interpretability/{diffae/figures => _internal}/gen_validation/valid200_alphastep.py (100%) rename src/ops_model/models/interpretability/{diffae/figures => _internal}/gen_validation/valid200_cache_build.py (100%) rename src/ops_model/models/interpretability/{diffae/figures => _internal}/gen_validation/valid200_capcheck.py (100%) rename src/ops_model/models/interpretability/{diffae/figures => _internal}/gen_validation/valid200_map_compare.py (100%) rename src/ops_model/models/interpretability/{diffae/figures => _internal}/gen_validation/valid200_metrics.py (100%) diff --git a/src/ops_model/models/interpretability/diffae/figures/gen_validation/bag_sweep_plots.py b/src/ops_model/models/interpretability/_internal/gen_validation/bag_sweep_plots.py similarity index 100% rename from src/ops_model/models/interpretability/diffae/figures/gen_validation/bag_sweep_plots.py rename to src/ops_model/models/interpretability/_internal/gen_validation/bag_sweep_plots.py diff --git a/src/ops_model/models/interpretability/diffae/figures/gen_validation/bag_sweep_score.py b/src/ops_model/models/interpretability/_internal/gen_validation/bag_sweep_score.py similarity index 100% rename from src/ops_model/models/interpretability/diffae/figures/gen_validation/bag_sweep_score.py rename to src/ops_model/models/interpretability/_internal/gen_validation/bag_sweep_score.py diff --git a/src/ops_model/models/interpretability/diffae/figures/gen_validation/centroid_bagsweep.py b/src/ops_model/models/interpretability/_internal/gen_validation/centroid_bagsweep.py similarity index 100% rename from src/ops_model/models/interpretability/diffae/figures/gen_validation/centroid_bagsweep.py rename to src/ops_model/models/interpretability/_internal/gen_validation/centroid_bagsweep.py diff --git a/src/ops_model/models/interpretability/diffae/figures/gen_validation/centroid_halves.py b/src/ops_model/models/interpretability/_internal/gen_validation/centroid_halves.py similarity index 100% rename from src/ops_model/models/interpretability/diffae/figures/gen_validation/centroid_halves.py rename to src/ops_model/models/interpretability/_internal/gen_validation/centroid_halves.py diff --git a/src/ops_model/models/interpretability/diffae/figures/gen_validation/centroid_pooled_bagsweep.py b/src/ops_model/models/interpretability/_internal/gen_validation/centroid_pooled_bagsweep.py similarity index 100% rename from src/ops_model/models/interpretability/diffae/figures/gen_validation/centroid_pooled_bagsweep.py rename to src/ops_model/models/interpretability/_internal/gen_validation/centroid_pooled_bagsweep.py diff --git a/src/ops_model/models/interpretability/diffae/figures/gen_validation/control_halves_zscore.py b/src/ops_model/models/interpretability/_internal/gen_validation/control_halves_zscore.py similarity index 100% rename from src/ops_model/models/interpretability/diffae/figures/gen_validation/control_halves_zscore.py rename to src/ops_model/models/interpretability/_internal/gen_validation/control_halves_zscore.py diff --git a/src/ops_model/models/interpretability/diffae/figures/gen_validation/embcheck.py b/src/ops_model/models/interpretability/_internal/gen_validation/embcheck.py similarity index 100% rename from src/ops_model/models/interpretability/diffae/figures/gen_validation/embcheck.py rename to src/ops_model/models/interpretability/_internal/gen_validation/embcheck.py diff --git a/src/ops_model/models/interpretability/diffae/figures/gen_validation/embedding_diagnostics.py b/src/ops_model/models/interpretability/_internal/gen_validation/embedding_diagnostics.py similarity index 100% rename from src/ops_model/models/interpretability/diffae/figures/gen_validation/embedding_diagnostics.py rename to src/ops_model/models/interpretability/_internal/gen_validation/embedding_diagnostics.py diff --git a/src/ops_model/models/interpretability/diffae/figures/gen_validation/figure4_bagsize_reachfrac.py b/src/ops_model/models/interpretability/_internal/gen_validation/figure4_bagsize_reachfrac.py similarity index 100% rename from src/ops_model/models/interpretability/diffae/figures/gen_validation/figure4_bagsize_reachfrac.py rename to src/ops_model/models/interpretability/_internal/gen_validation/figure4_bagsize_reachfrac.py diff --git a/src/ops_model/models/interpretability/diffae/figures/gen_validation/figure4_v5_accuracy_summary.py b/src/ops_model/models/interpretability/_internal/gen_validation/figure4_v5_accuracy_summary.py similarity index 100% rename from src/ops_model/models/interpretability/diffae/figures/gen_validation/figure4_v5_accuracy_summary.py rename to src/ops_model/models/interpretability/_internal/gen_validation/figure4_v5_accuracy_summary.py diff --git a/src/ops_model/models/interpretability/diffae/figures/gen_validation/gen_alpha_embedding.py b/src/ops_model/models/interpretability/_internal/gen_validation/gen_alpha_embedding.py similarity index 100% rename from src/ops_model/models/interpretability/diffae/figures/gen_validation/gen_alpha_embedding.py rename to src/ops_model/models/interpretability/_internal/gen_validation/gen_alpha_embedding.py diff --git a/src/ops_model/models/interpretability/diffae/figures/gen_validation/gen_embed_refit.py b/src/ops_model/models/interpretability/_internal/gen_validation/gen_embed_refit.py similarity index 100% rename from src/ops_model/models/interpretability/diffae/figures/gen_validation/gen_embed_refit.py rename to src/ops_model/models/interpretability/_internal/gen_validation/gen_embed_refit.py diff --git a/src/ops_model/models/interpretability/diffae/figures/gen_validation/gen_phate_passthrough.py b/src/ops_model/models/interpretability/_internal/gen_validation/gen_phate_passthrough.py similarity index 100% rename from src/ops_model/models/interpretability/diffae/figures/gen_validation/gen_phate_passthrough.py rename to src/ops_model/models/interpretability/_internal/gen_validation/gen_phate_passthrough.py diff --git a/src/ops_model/models/interpretability/diffae/figures/gen_validation/gen_real_centroid.py b/src/ops_model/models/interpretability/_internal/gen_validation/gen_real_centroid.py similarity index 100% rename from src/ops_model/models/interpretability/diffae/figures/gen_validation/gen_real_centroid.py rename to src/ops_model/models/interpretability/_internal/gen_validation/gen_real_centroid.py diff --git a/src/ops_model/models/interpretability/diffae/figures/gen_validation/gen_real_distinct.py b/src/ops_model/models/interpretability/_internal/gen_validation/gen_real_distinct.py similarity index 100% rename from src/ops_model/models/interpretability/diffae/figures/gen_validation/gen_real_distinct.py rename to src/ops_model/models/interpretability/_internal/gen_validation/gen_real_distinct.py diff --git a/src/ops_model/models/interpretability/diffae/figures/gen_validation/ntc_inverse_gap.py b/src/ops_model/models/interpretability/_internal/gen_validation/ntc_inverse_gap.py similarity index 100% rename from src/ops_model/models/interpretability/diffae/figures/gen_validation/ntc_inverse_gap.py rename to src/ops_model/models/interpretability/_internal/gen_validation/ntc_inverse_gap.py diff --git a/src/ops_model/models/interpretability/diffae/figures/gen_validation/patch_cache_real.py b/src/ops_model/models/interpretability/_internal/gen_validation/patch_cache_real.py similarity index 100% rename from src/ops_model/models/interpretability/diffae/figures/gen_validation/patch_cache_real.py rename to src/ops_model/models/interpretability/_internal/gen_validation/patch_cache_real.py diff --git a/src/ops_model/models/interpretability/diffae/figures/gen_validation/publish_multibag_page.py b/src/ops_model/models/interpretability/_internal/gen_validation/publish_multibag_page.py similarity index 100% rename from src/ops_model/models/interpretability/diffae/figures/gen_validation/publish_multibag_page.py rename to src/ops_model/models/interpretability/_internal/gen_validation/publish_multibag_page.py diff --git a/src/ops_model/models/interpretability/diffae/figures/gen_validation/rank_summary.py b/src/ops_model/models/interpretability/_internal/gen_validation/rank_summary.py similarity index 100% rename from src/ops_model/models/interpretability/diffae/figures/gen_validation/rank_summary.py rename to src/ops_model/models/interpretability/_internal/gen_validation/rank_summary.py diff --git a/src/ops_model/models/interpretability/diffae/figures/gen_validation/st_halves_score.py b/src/ops_model/models/interpretability/_internal/gen_validation/st_halves_score.py similarity index 100% rename from src/ops_model/models/interpretability/diffae/figures/gen_validation/st_halves_score.py rename to src/ops_model/models/interpretability/_internal/gen_validation/st_halves_score.py diff --git a/src/ops_model/models/interpretability/diffae/figures/gen_validation/std_anchor_test.py b/src/ops_model/models/interpretability/_internal/gen_validation/std_anchor_test.py similarity index 100% rename from src/ops_model/models/interpretability/diffae/figures/gen_validation/std_anchor_test.py rename to src/ops_model/models/interpretability/_internal/gen_validation/std_anchor_test.py diff --git a/src/ops_model/models/interpretability/diffae/figures/gen_validation/stepabl_compare.py b/src/ops_model/models/interpretability/_internal/gen_validation/stepabl_compare.py similarity index 100% rename from src/ops_model/models/interpretability/diffae/figures/gen_validation/stepabl_compare.py rename to src/ops_model/models/interpretability/_internal/gen_validation/stepabl_compare.py diff --git a/src/ops_model/models/interpretability/diffae/figures/gen_validation/valid200_alphastep.py b/src/ops_model/models/interpretability/_internal/gen_validation/valid200_alphastep.py similarity index 100% rename from src/ops_model/models/interpretability/diffae/figures/gen_validation/valid200_alphastep.py rename to src/ops_model/models/interpretability/_internal/gen_validation/valid200_alphastep.py diff --git a/src/ops_model/models/interpretability/diffae/figures/gen_validation/valid200_cache_build.py b/src/ops_model/models/interpretability/_internal/gen_validation/valid200_cache_build.py similarity index 100% rename from src/ops_model/models/interpretability/diffae/figures/gen_validation/valid200_cache_build.py rename to src/ops_model/models/interpretability/_internal/gen_validation/valid200_cache_build.py diff --git a/src/ops_model/models/interpretability/diffae/figures/gen_validation/valid200_capcheck.py b/src/ops_model/models/interpretability/_internal/gen_validation/valid200_capcheck.py similarity index 100% rename from src/ops_model/models/interpretability/diffae/figures/gen_validation/valid200_capcheck.py rename to src/ops_model/models/interpretability/_internal/gen_validation/valid200_capcheck.py diff --git a/src/ops_model/models/interpretability/diffae/figures/gen_validation/valid200_map_compare.py b/src/ops_model/models/interpretability/_internal/gen_validation/valid200_map_compare.py similarity index 100% rename from src/ops_model/models/interpretability/diffae/figures/gen_validation/valid200_map_compare.py rename to src/ops_model/models/interpretability/_internal/gen_validation/valid200_map_compare.py diff --git a/src/ops_model/models/interpretability/diffae/figures/gen_validation/valid200_metrics.py b/src/ops_model/models/interpretability/_internal/gen_validation/valid200_metrics.py similarity index 100% rename from src/ops_model/models/interpretability/diffae/figures/gen_validation/valid200_metrics.py rename to src/ops_model/models/interpretability/_internal/gen_validation/valid200_metrics.py From 90a80f6ab19435d709b49d89ccefbc814fa4a9a5 Mon Sep 17 00:00:00 2001 From: Gav Sturm Date: Tue, 11 Aug 2026 09:44:55 -0700 Subject: [PATCH 07/13] figures: group the 4 figure4_setacc_panel* into figures/setacc/ subpackage MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit No duplicate code to extract — make_panel (in figure4_setacc_panel) is already the single shared renderer; the 3 variants are thin config wrappers. Moved all 4 into figures/setacc/ with __init__; converted their cross-refs (make_panel, _setacc_common, _setacc_phase) + cis_golgi_alternatives' make_panel import + _setacc_phase's own _setacc_common import to absolute (dual-mode: works under -m and flat). Shared _setacc_common/_setacc_phase stay at figures/ level (used by ~11 scripts). Verified: all 4 panels import as package modules. --- .../models/interpretability/diffae/figures/__init__.py | 1 + .../models/interpretability/diffae/figures/_setacc_phase.py | 2 +- .../interpretability/diffae/figures/cis_golgi_alternatives.py | 2 +- .../models/interpretability/diffae/figures/setacc/__init__.py | 2 ++ .../diffae/figures/{ => setacc}/figure4_setacc_panel.py | 2 +- .../figures/{ => setacc}/figure4_setacc_panel_fluorB.py | 4 ++-- .../figures/{ => setacc}/figure4_setacc_panel_newpheno.py | 4 ++-- .../diffae/figures/{ => setacc}/figure4_setacc_panel_phase.py | 4 ++-- 8 files changed, 12 insertions(+), 9 deletions(-) create mode 100644 src/ops_model/models/interpretability/diffae/figures/__init__.py create mode 100644 src/ops_model/models/interpretability/diffae/figures/setacc/__init__.py rename src/ops_model/models/interpretability/diffae/figures/{ => setacc}/figure4_setacc_panel.py (96%) rename src/ops_model/models/interpretability/diffae/figures/{ => setacc}/figure4_setacc_panel_fluorB.py (87%) rename src/ops_model/models/interpretability/diffae/figures/{ => setacc}/figure4_setacc_panel_newpheno.py (93%) rename src/ops_model/models/interpretability/diffae/figures/{ => setacc}/figure4_setacc_panel_phase.py (70%) diff --git a/src/ops_model/models/interpretability/diffae/figures/__init__.py b/src/ops_model/models/interpretability/diffae/figures/__init__.py new file mode 100644 index 0000000..75118bf --- /dev/null +++ b/src/ops_model/models/interpretability/diffae/figures/__init__.py @@ -0,0 +1 @@ +"""Paper Figure-4 generation scripts for the DiffAE interpretability pipeline.""" diff --git a/src/ops_model/models/interpretability/diffae/figures/_setacc_phase.py b/src/ops_model/models/interpretability/diffae/figures/_setacc_phase.py index 3a6b228..374dbe4 100644 --- a/src/ops_model/models/interpretability/diffae/figures/_setacc_phase.py +++ b/src/ops_model/models/interpretability/diffae/figures/_setacc_phase.py @@ -5,7 +5,7 @@ import numpy as np import pandas as pd -from _setacc_common import crop_pick_from_df, tile_at +from ops_model.models.interpretability.diffae.figures._setacc_common import crop_pick_from_df, tile_at RANKS = "/hpc/projects/icd.fast.ops/models/diffex/viewer_assets_v5/_rankings" PHASE_CH = "Phase2D" diff --git a/src/ops_model/models/interpretability/diffae/figures/cis_golgi_alternatives.py b/src/ops_model/models/interpretability/diffae/figures/cis_golgi_alternatives.py index 4e5f509..9cb170f 100644 --- a/src/ops_model/models/interpretability/diffae/figures/cis_golgi_alternatives.py +++ b/src/ops_model/models/interpretability/diffae/figures/cis_golgi_alternatives.py @@ -1,6 +1,6 @@ """Candidate strip for the panel-D Rab-slot (alternatives to COPI·cis-Golgi) — KO vs NTC at rank 1 (most distinctive) for a few trafficking/organelle complex+marker pairs, to pick the most obvious.""" -from figure4_setacc_panel import make_panel +from ops_model.models.interpretability.diffae.figures.setacc.figure4_setacc_panel import make_panel CANDS = [ dict(slug="cis_Golgi_mStayGold_CENPRaltORF", mc="cis-Golgi_mStayGold-CENPRaltORF", ch="GFP", diff --git a/src/ops_model/models/interpretability/diffae/figures/setacc/__init__.py b/src/ops_model/models/interpretability/diffae/figures/setacc/__init__.py new file mode 100644 index 0000000..dea21ab --- /dev/null +++ b/src/ops_model/models/interpretability/diffae/figures/setacc/__init__.py @@ -0,0 +1,2 @@ +"""Figure-4 set-accuracy panels (top-predictive KO-vs-NTC cell grids). +Built on the shared make_panel renderer + figures/_setacc_common.""" diff --git a/src/ops_model/models/interpretability/diffae/figures/figure4_setacc_panel.py b/src/ops_model/models/interpretability/diffae/figures/setacc/figure4_setacc_panel.py similarity index 96% rename from src/ops_model/models/interpretability/diffae/figures/figure4_setacc_panel.py rename to src/ops_model/models/interpretability/diffae/figures/setacc/figure4_setacc_panel.py index d773a85..5adf16b 100644 --- a/src/ops_model/models/interpretability/diffae/figures/figure4_setacc_panel.py +++ b/src/ops_model/models/interpretability/diffae/figures/setacc/figure4_setacc_panel.py @@ -16,7 +16,7 @@ matplotlib.use("Agg") import matplotlib.pyplot as plt -from _setacc_common import COMPLEX_COLS, GENE_COLS, OUT, column_tiles +from ops_model.models.interpretability.diffae.figures._setacc_common import COMPLEX_COLS, GENE_COLS, OUT, column_tiles plt.rcParams["pdf.fonttype"] = 42 plt.rcParams["svg.fonttype"] = "none" diff --git a/src/ops_model/models/interpretability/diffae/figures/figure4_setacc_panel_fluorB.py b/src/ops_model/models/interpretability/diffae/figures/setacc/figure4_setacc_panel_fluorB.py similarity index 87% rename from src/ops_model/models/interpretability/diffae/figures/figure4_setacc_panel_fluorB.py rename to src/ops_model/models/interpretability/diffae/figures/setacc/figure4_setacc_panel_fluorB.py index 6830206..3c1b84b 100644 --- a/src/ops_model/models/interpretability/diffae/figures/figure4_setacc_panel_fluorB.py +++ b/src/ops_model/models/interpretability/diffae/figures/setacc/figure4_setacc_panel_fluorB.py @@ -2,8 +2,8 @@ CFL1 (actin FastAct), mTOR (LysoTracker). NTC top / KO bottom, per-column KO+NTC intensity window. Run: python figure4_setacc_panel_fluorB.py""" -from figure4_setacc_panel import make_panel -from _setacc_common import column_tiles +from ops_model.models.interpretability.diffae.figures.setacc.figure4_setacc_panel import make_panel +from ops_model.models.interpretability.diffae.figures._setacc_common import column_tiles COLS = [ dict(slug="Mitochondria_TOMM20", mc="Mitochondria_TOMM20", ch="CP1_mitochondria_TOMM20", diff --git a/src/ops_model/models/interpretability/diffae/figures/figure4_setacc_panel_newpheno.py b/src/ops_model/models/interpretability/diffae/figures/setacc/figure4_setacc_panel_newpheno.py similarity index 93% rename from src/ops_model/models/interpretability/diffae/figures/figure4_setacc_panel_newpheno.py rename to src/ops_model/models/interpretability/diffae/figures/setacc/figure4_setacc_panel_newpheno.py index 163c3db..d4ab71d 100644 --- a/src/ops_model/models/interpretability/diffae/figures/figure4_setacc_panel_newpheno.py +++ b/src/ops_model/models/interpretability/diffae/figures/setacc/figure4_setacc_panel_newpheno.py @@ -9,8 +9,8 @@ import numpy as np import pandas as pd -from _setacc_common import crop_pick_from_df, tile_at -from figure4_setacc_panel import make_panel +from ops_model.models.interpretability.diffae.figures._setacc_common import crop_pick_from_df, tile_at +from ops_model.models.interpretability.diffae.figures.setacc.figure4_setacc_panel import make_panel RANK = "/hpc/projects/icd.fast.ops/models/diffex/viewer_assets_v5/_rankings/pma_shap_phase_geneKO.parquet" PHASE_CH = "Phase2D" diff --git a/src/ops_model/models/interpretability/diffae/figures/figure4_setacc_panel_phase.py b/src/ops_model/models/interpretability/diffae/figures/setacc/figure4_setacc_panel_phase.py similarity index 70% rename from src/ops_model/models/interpretability/diffae/figures/figure4_setacc_panel_phase.py rename to src/ops_model/models/interpretability/diffae/figures/setacc/figure4_setacc_panel_phase.py index a8e4f28..70012a7 100644 --- a/src/ops_model/models/interpretability/diffae/figures/figure4_setacc_panel_phase.py +++ b/src/ops_model/models/interpretability/diffae/figures/setacc/figure4_setacc_panel_phase.py @@ -1,8 +1,8 @@ """Panel-E-style figure — top set-accuracy cells in label-free 2D phase, KO vs NTC. Groups: TIMM23 & TIPARP (gene-level), Arp2/3 & Core Mediator (complex). Picks set in _setacc_phase.COLS_PHASE. Vector output (SVG + PNG).""" -from figure4_setacc_panel import make_panel -from _setacc_phase import COLS_PHASE, column_tiles_phase +from ops_model.models.interpretability.diffae.figures.setacc.figure4_setacc_panel import make_panel +from ops_model.models.interpretability.diffae.figures._setacc_phase import COLS_PHASE, column_tiles_phase if __name__ == "__main__": make_panel(COLS_PHASE, "Top-predictive cells (phase)", "panelE_phase_setacc", From 79a3a4439631744a15c00a18cac9c95ef1c3524f Mon Sep 17 00:00:00 2001 From: Gav Sturm Date: Tue, 11 Aug 2026 09:46:12 -0700 Subject: [PATCH 08/13] pyproject: exclude interpretability/_internal from the published wheel hatch wheel target now excludes the not-for-release analysis/paper tooling. Verified via 'uv build --wheel': 0 _internal entries, 73 core diffae files shipped. --- pyproject.toml | 3 +++ 1 file changed, 3 insertions(+) diff --git a/pyproject.toml b/pyproject.toml index ed3724a..b0b0985 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -84,6 +84,9 @@ source = "vcs" [tool.hatch.build.targets.wheel] only-include = ["src"] sources = ["src"] +# not-for-release: internal analysis/paper tooling (atlas, shap, titration, embedding, +# weighted_aggregation, viewer, kyle_pcs, gen_validation) lives under interpretability/_internal +exclude = ["src/ops_model/models/interpretability/_internal/**"] [tool.hatch.metadata] allow-direct-references = true From a3b301ce25ce19af6afb01798f2abca713f31b42 Mon Sep 17 00:00:00 2001 From: Gav Sturm Date: Tue, 11 Aug 2026 10:13:18 -0700 Subject: [PATCH 09/13] rename interpretability/_internal -> interpretability/toolkit Match the ops_process convention: not-for-release tooling lives in toolkit/ on the refactor line (present in source, excluded from the wheel). 20 intra refs repointed, pyproject exclude updated. The public-release branch deletes toolkit/ entirely. --- pyproject.toml | 4 ++-- .../{_internal => toolkit}/RUNBOOK.md | 0 .../atlas/attention_accuracy_umap_animation.py | 0 .../{_internal => toolkit}/atlas/attention_atlas.py | 0 .../atlas/attention_atlas_shap.py | 0 .../atlas/low_attention_phase_atlas.py | 0 .../{_internal => toolkit}/atlas/make_scale_bar.py | 0 .../atlas/marker_selection_distribution.py | 0 .../atlas/plot_eval_accuracy_curves.py | 0 .../embedding/generate_ko_violin_plots.py | 0 .../embedding/regen_umap_gav.py | 0 .../embedding/regen_umap_html.py | 0 .../embedding/run_all_atlases.py | 0 .../embedding/top_attention_embed_and_score.py | 0 .../gen_validation/bag_sweep_plots.py | 0 .../gen_validation/bag_sweep_score.py | 4 ++-- .../gen_validation/centroid_bagsweep.py | 0 .../gen_validation/centroid_halves.py | 0 .../gen_validation/centroid_pooled_bagsweep.py | 0 .../gen_validation/control_halves_zscore.py | 0 .../gen_validation/embcheck.py | 2 +- .../gen_validation/embedding_diagnostics.py | 0 .../gen_validation/figure4_bagsize_reachfrac.py | 0 .../gen_validation/figure4_v5_accuracy_summary.py | 0 .../gen_validation/gen_alpha_embedding.py | 0 .../gen_validation/gen_embed_refit.py | 0 .../gen_validation/gen_phate_passthrough.py | 0 .../gen_validation/gen_real_centroid.py | 0 .../gen_validation/gen_real_distinct.py | 0 .../gen_validation/ntc_inverse_gap.py | 0 .../gen_validation/patch_cache_real.py | 0 .../gen_validation/publish_multibag_page.py | 0 .../gen_validation/rank_summary.py | 0 .../gen_validation/st_halves_score.py | 4 ++-- .../gen_validation/std_anchor_test.py | 0 .../gen_validation/stepabl_compare.py | 0 .../gen_validation/valid200_alphastep.py | 0 .../gen_validation/valid200_cache_build.py | 0 .../gen_validation/valid200_capcheck.py | 0 .../gen_validation/valid200_map_compare.py | 0 .../gen_validation/valid200_metrics.py | 0 .../kyle_pcs/build_static_explorer.py | 0 .../kyle_pcs/compute_pc_strips.py | 0 .../shap/analyze_chad_variants.py | 0 .../shap/generate_shap_captions_combined.py | 0 .../{_internal => toolkit}/shap/ko_shap_features.py | 0 .../shap/merge_shap_shards.py | 0 .../shap/ntc_attention_compare.py | 0 .../{_internal => toolkit}/shap/ntc_pick_cells.py | 0 .../shap/ntc_shap_features.py | 0 .../{_internal => toolkit}/shap/run_all_shap.py | 0 .../shap/run_shap_pipeline.py | 0 .../shap/shap_approach_compare.py | 0 .../titration/decay/map_attention_decay.py | 0 .../titration/decay/phate_peak_groups.py | 0 .../titration/decay/plot_3way_summary_bars.py | 0 .../decay/plot_all_cells_correction_bars.py | 0 .../expansion/count_genes_above_threshold.py | 4 ++-- .../expansion/map_attention_expansion_v4.py | 0 .../expansion/plot_sgrna_coverage_sweep.py | 0 .../titration/expansion/run_percentile_sweep.py | 0 .../{_internal => toolkit}/viewer/__init__.py | 0 .../viewer/_altanchor_build.py | 0 .../{_internal => toolkit}/viewer/_anchortest.py | 0 .../viewer/_build_stepablation.py | 0 .../viewer/_build_v5_inverted.py | 2 +- .../viewer/_build_v5_montages.py | 2 +- .../viewer/_build_valid200.py | 6 +++--- .../viewer/_consolidate_cells.py | 0 .../viewer/_fluor_complex_build.py | 0 .../viewer/_fluor_v5_build.py | 0 .../viewer/_migrate_v4_to_v5.py | 0 .../{_internal => toolkit}/viewer/_phase_vs.py | 0 .../{_internal => toolkit}/viewer/_rebuild_v5.py | 0 .../{_internal => toolkit}/viewer/_rescore_rank.py | 2 +- .../{_internal => toolkit}/viewer/_score_v4.py | 0 .../{_internal => toolkit}/viewer/_v4acc_test.py | 0 .../viewer/_verify_pt_space.py | 0 .../viewer/_verify_score_bridge.py | 0 .../viewer/altanchor_pairs.json | 0 .../{_internal => toolkit}/viewer/anchor_cells.py | 0 .../viewer/build_attention_heads.py | 8 ++++---- .../viewer/build_complex_ebi_map.py | 0 .../viewer/build_fluor_shap_rankings.py | 4 ++-- .../viewer/build_montage_features.py | 2 +- .../viewer/build_pc_features.py | 2 +- .../{_internal => toolkit}/viewer/build_pc_walks.py | 4 ++-- .../{_internal => toolkit}/viewer/build_pcs.py | 4 ++-- .../viewer/build_pcs_marker.py | 2 +- .../viewer/build_phase_shap_rankings.py | 4 ++-- .../viewer/build_phate_figure.py | 2 +- .../viewer/build_setacc_bins.py | 0 .../viewer/build_setacc_bymarker.py | 0 .../viewer/build_top_cells.py | 4 ++-- .../viewer/build_umap_montage.py | 0 .../{_internal => toolkit}/viewer/deploy/README.md | 0 .../{_internal => toolkit}/viewer/marker_leaves.py | 0 .../viewer/mimic_alex_embed.py | 0 .../{_internal => toolkit}/viewer/morphometrics.py | 0 .../{_internal => toolkit}/viewer/nway_clf.py | 0 .../viewer/phenotype_cells.py | 0 .../viewer/render_montage_scales.py | 2 +- .../viewer/score_generated.py | 0 .../{_internal => toolkit}/viewer/set_classifier.py | 0 .../{_internal => toolkit}/viewer/submit.py | 8 ++++---- .../{_internal => toolkit}/viewer/webapp/app.js | 0 .../viewer/webapp/biohub-mark.png | Bin .../viewer/webapp/biohub-wordmark.png | Bin .../viewer/webapp/build_gene_narratives.py | 0 .../{_internal => toolkit}/viewer/webapp/gif.js | 0 .../viewer/webapp/gif.worker.js | 0 .../{_internal => toolkit}/viewer/webapp/index.html | 0 .../{_internal => toolkit}/viewer/webapp/methods.js | 0 .../viewer/webapp/morpho_demo.html | 0 .../viewer/webapp/openseadragon.min.js | 0 .../viewer/webapp/opsin-eyes.svg | 0 .../{_internal => toolkit}/viewer/webapp/style.css | 0 .../weighted_aggregation/_v4_attn_worker.py | 0 .../weighted_aggregation/analyze_v3_acc_bins.py | 0 .../weighted_aggregation/plot_v4_attn_comparison.py | 0 .../run_v3_pipeline_on_v4_attn_weighted.py | 0 .../run_v3_pipeline_on_v4_features.py | 0 122 files changed, 38 insertions(+), 38 deletions(-) rename src/ops_model/models/interpretability/{_internal => toolkit}/RUNBOOK.md (100%) rename src/ops_model/models/interpretability/{_internal => toolkit}/atlas/attention_accuracy_umap_animation.py (100%) rename src/ops_model/models/interpretability/{_internal => toolkit}/atlas/attention_atlas.py (100%) rename src/ops_model/models/interpretability/{_internal => toolkit}/atlas/attention_atlas_shap.py (100%) rename src/ops_model/models/interpretability/{_internal => toolkit}/atlas/low_attention_phase_atlas.py (100%) rename src/ops_model/models/interpretability/{_internal => toolkit}/atlas/make_scale_bar.py (100%) rename src/ops_model/models/interpretability/{_internal => toolkit}/atlas/marker_selection_distribution.py (100%) rename src/ops_model/models/interpretability/{_internal => toolkit}/atlas/plot_eval_accuracy_curves.py (100%) rename src/ops_model/models/interpretability/{_internal => toolkit}/embedding/generate_ko_violin_plots.py (100%) rename src/ops_model/models/interpretability/{_internal => toolkit}/embedding/regen_umap_gav.py (100%) rename src/ops_model/models/interpretability/{_internal => toolkit}/embedding/regen_umap_html.py (100%) rename src/ops_model/models/interpretability/{_internal => toolkit}/embedding/run_all_atlases.py (100%) rename src/ops_model/models/interpretability/{_internal => toolkit}/embedding/top_attention_embed_and_score.py (100%) rename src/ops_model/models/interpretability/{_internal => toolkit}/gen_validation/bag_sweep_plots.py (100%) rename src/ops_model/models/interpretability/{_internal => toolkit}/gen_validation/bag_sweep_score.py (92%) rename src/ops_model/models/interpretability/{_internal => toolkit}/gen_validation/centroid_bagsweep.py (100%) rename src/ops_model/models/interpretability/{_internal => toolkit}/gen_validation/centroid_halves.py (100%) rename src/ops_model/models/interpretability/{_internal => toolkit}/gen_validation/centroid_pooled_bagsweep.py (100%) rename src/ops_model/models/interpretability/{_internal => toolkit}/gen_validation/control_halves_zscore.py (100%) rename src/ops_model/models/interpretability/{_internal => toolkit}/gen_validation/embcheck.py (95%) rename src/ops_model/models/interpretability/{_internal => toolkit}/gen_validation/embedding_diagnostics.py (100%) rename src/ops_model/models/interpretability/{_internal => toolkit}/gen_validation/figure4_bagsize_reachfrac.py (100%) rename src/ops_model/models/interpretability/{_internal => toolkit}/gen_validation/figure4_v5_accuracy_summary.py (100%) rename src/ops_model/models/interpretability/{_internal => toolkit}/gen_validation/gen_alpha_embedding.py (100%) rename src/ops_model/models/interpretability/{_internal => toolkit}/gen_validation/gen_embed_refit.py (100%) rename src/ops_model/models/interpretability/{_internal => toolkit}/gen_validation/gen_phate_passthrough.py (100%) rename src/ops_model/models/interpretability/{_internal => toolkit}/gen_validation/gen_real_centroid.py (100%) rename src/ops_model/models/interpretability/{_internal => toolkit}/gen_validation/gen_real_distinct.py (100%) rename src/ops_model/models/interpretability/{_internal => toolkit}/gen_validation/ntc_inverse_gap.py (100%) rename src/ops_model/models/interpretability/{_internal => toolkit}/gen_validation/patch_cache_real.py (100%) rename src/ops_model/models/interpretability/{_internal => toolkit}/gen_validation/publish_multibag_page.py (100%) rename src/ops_model/models/interpretability/{_internal => toolkit}/gen_validation/rank_summary.py (100%) rename src/ops_model/models/interpretability/{_internal => toolkit}/gen_validation/st_halves_score.py (92%) rename src/ops_model/models/interpretability/{_internal => toolkit}/gen_validation/std_anchor_test.py (100%) rename src/ops_model/models/interpretability/{_internal => toolkit}/gen_validation/stepabl_compare.py (100%) rename src/ops_model/models/interpretability/{_internal => toolkit}/gen_validation/valid200_alphastep.py (100%) rename src/ops_model/models/interpretability/{_internal => toolkit}/gen_validation/valid200_cache_build.py (100%) rename src/ops_model/models/interpretability/{_internal => toolkit}/gen_validation/valid200_capcheck.py (100%) rename src/ops_model/models/interpretability/{_internal => toolkit}/gen_validation/valid200_map_compare.py (100%) rename src/ops_model/models/interpretability/{_internal => toolkit}/gen_validation/valid200_metrics.py (100%) rename src/ops_model/models/interpretability/{_internal => toolkit}/kyle_pcs/build_static_explorer.py (100%) rename src/ops_model/models/interpretability/{_internal => toolkit}/kyle_pcs/compute_pc_strips.py (100%) rename src/ops_model/models/interpretability/{_internal => toolkit}/shap/analyze_chad_variants.py (100%) rename src/ops_model/models/interpretability/{_internal => toolkit}/shap/generate_shap_captions_combined.py (100%) rename src/ops_model/models/interpretability/{_internal => toolkit}/shap/ko_shap_features.py (100%) rename src/ops_model/models/interpretability/{_internal => toolkit}/shap/merge_shap_shards.py (100%) rename src/ops_model/models/interpretability/{_internal => toolkit}/shap/ntc_attention_compare.py (100%) rename src/ops_model/models/interpretability/{_internal => toolkit}/shap/ntc_pick_cells.py (100%) rename src/ops_model/models/interpretability/{_internal => toolkit}/shap/ntc_shap_features.py (100%) rename src/ops_model/models/interpretability/{_internal => toolkit}/shap/run_all_shap.py (100%) rename src/ops_model/models/interpretability/{_internal => toolkit}/shap/run_shap_pipeline.py (100%) rename src/ops_model/models/interpretability/{_internal => toolkit}/shap/shap_approach_compare.py (100%) rename src/ops_model/models/interpretability/{_internal => toolkit}/titration/decay/map_attention_decay.py (100%) rename src/ops_model/models/interpretability/{_internal => toolkit}/titration/decay/phate_peak_groups.py (100%) rename src/ops_model/models/interpretability/{_internal => toolkit}/titration/decay/plot_3way_summary_bars.py (100%) rename src/ops_model/models/interpretability/{_internal => toolkit}/titration/decay/plot_all_cells_correction_bars.py (100%) rename src/ops_model/models/interpretability/{_internal => toolkit}/titration/expansion/count_genes_above_threshold.py (98%) rename src/ops_model/models/interpretability/{_internal => toolkit}/titration/expansion/map_attention_expansion_v4.py (100%) rename src/ops_model/models/interpretability/{_internal => toolkit}/titration/expansion/plot_sgrna_coverage_sweep.py (100%) rename src/ops_model/models/interpretability/{_internal => toolkit}/titration/expansion/run_percentile_sweep.py (100%) rename src/ops_model/models/interpretability/{_internal => toolkit}/viewer/__init__.py (100%) rename src/ops_model/models/interpretability/{_internal => toolkit}/viewer/_altanchor_build.py (100%) rename src/ops_model/models/interpretability/{_internal => toolkit}/viewer/_anchortest.py (100%) rename src/ops_model/models/interpretability/{_internal => toolkit}/viewer/_build_stepablation.py (100%) rename src/ops_model/models/interpretability/{_internal => toolkit}/viewer/_build_v5_inverted.py (99%) rename src/ops_model/models/interpretability/{_internal => toolkit}/viewer/_build_v5_montages.py (98%) rename src/ops_model/models/interpretability/{_internal => toolkit}/viewer/_build_valid200.py (95%) rename src/ops_model/models/interpretability/{_internal => toolkit}/viewer/_consolidate_cells.py (100%) rename src/ops_model/models/interpretability/{_internal => toolkit}/viewer/_fluor_complex_build.py (100%) rename src/ops_model/models/interpretability/{_internal => toolkit}/viewer/_fluor_v5_build.py (100%) rename src/ops_model/models/interpretability/{_internal => toolkit}/viewer/_migrate_v4_to_v5.py (100%) rename src/ops_model/models/interpretability/{_internal => toolkit}/viewer/_phase_vs.py (100%) rename src/ops_model/models/interpretability/{_internal => toolkit}/viewer/_rebuild_v5.py (100%) rename src/ops_model/models/interpretability/{_internal => toolkit}/viewer/_rescore_rank.py (96%) rename src/ops_model/models/interpretability/{_internal => toolkit}/viewer/_score_v4.py (100%) rename src/ops_model/models/interpretability/{_internal => toolkit}/viewer/_v4acc_test.py (100%) rename src/ops_model/models/interpretability/{_internal => toolkit}/viewer/_verify_pt_space.py (100%) rename src/ops_model/models/interpretability/{_internal => toolkit}/viewer/_verify_score_bridge.py (100%) rename src/ops_model/models/interpretability/{_internal => toolkit}/viewer/altanchor_pairs.json (100%) rename src/ops_model/models/interpretability/{_internal => toolkit}/viewer/anchor_cells.py (100%) rename src/ops_model/models/interpretability/{_internal => toolkit}/viewer/build_attention_heads.py (95%) rename src/ops_model/models/interpretability/{_internal => toolkit}/viewer/build_complex_ebi_map.py (100%) rename src/ops_model/models/interpretability/{_internal => toolkit}/viewer/build_fluor_shap_rankings.py (96%) rename src/ops_model/models/interpretability/{_internal => toolkit}/viewer/build_montage_features.py (96%) rename src/ops_model/models/interpretability/{_internal => toolkit}/viewer/build_pc_features.py (99%) rename src/ops_model/models/interpretability/{_internal => toolkit}/viewer/build_pc_walks.py (97%) rename src/ops_model/models/interpretability/{_internal => toolkit}/viewer/build_pcs.py (97%) rename src/ops_model/models/interpretability/{_internal => toolkit}/viewer/build_pcs_marker.py (99%) rename src/ops_model/models/interpretability/{_internal => toolkit}/viewer/build_phase_shap_rankings.py (95%) rename src/ops_model/models/interpretability/{_internal => toolkit}/viewer/build_phate_figure.py (99%) rename src/ops_model/models/interpretability/{_internal => toolkit}/viewer/build_setacc_bins.py (100%) rename src/ops_model/models/interpretability/{_internal => toolkit}/viewer/build_setacc_bymarker.py (100%) rename src/ops_model/models/interpretability/{_internal => toolkit}/viewer/build_top_cells.py (97%) rename src/ops_model/models/interpretability/{_internal => toolkit}/viewer/build_umap_montage.py (100%) rename src/ops_model/models/interpretability/{_internal => toolkit}/viewer/deploy/README.md (100%) rename src/ops_model/models/interpretability/{_internal => toolkit}/viewer/marker_leaves.py (100%) rename src/ops_model/models/interpretability/{_internal => toolkit}/viewer/mimic_alex_embed.py (100%) rename src/ops_model/models/interpretability/{_internal => toolkit}/viewer/morphometrics.py (100%) rename src/ops_model/models/interpretability/{_internal => toolkit}/viewer/nway_clf.py (100%) rename src/ops_model/models/interpretability/{_internal => toolkit}/viewer/phenotype_cells.py (100%) rename src/ops_model/models/interpretability/{_internal => toolkit}/viewer/render_montage_scales.py (99%) rename src/ops_model/models/interpretability/{_internal => toolkit}/viewer/score_generated.py (100%) rename src/ops_model/models/interpretability/{_internal => toolkit}/viewer/set_classifier.py (100%) rename src/ops_model/models/interpretability/{_internal => toolkit}/viewer/submit.py (97%) rename src/ops_model/models/interpretability/{_internal => toolkit}/viewer/webapp/app.js (100%) rename src/ops_model/models/interpretability/{_internal => toolkit}/viewer/webapp/biohub-mark.png (100%) rename src/ops_model/models/interpretability/{_internal => toolkit}/viewer/webapp/biohub-wordmark.png (100%) rename src/ops_model/models/interpretability/{_internal => toolkit}/viewer/webapp/build_gene_narratives.py (100%) rename src/ops_model/models/interpretability/{_internal => toolkit}/viewer/webapp/gif.js (100%) rename src/ops_model/models/interpretability/{_internal => toolkit}/viewer/webapp/gif.worker.js (100%) rename src/ops_model/models/interpretability/{_internal => toolkit}/viewer/webapp/index.html (100%) rename src/ops_model/models/interpretability/{_internal => toolkit}/viewer/webapp/methods.js (100%) rename src/ops_model/models/interpretability/{_internal => toolkit}/viewer/webapp/morpho_demo.html (100%) rename src/ops_model/models/interpretability/{_internal => toolkit}/viewer/webapp/openseadragon.min.js (100%) rename src/ops_model/models/interpretability/{_internal => toolkit}/viewer/webapp/opsin-eyes.svg (100%) rename src/ops_model/models/interpretability/{_internal => toolkit}/viewer/webapp/style.css (100%) rename src/ops_model/models/interpretability/{_internal => toolkit}/weighted_aggregation/_v4_attn_worker.py (100%) rename src/ops_model/models/interpretability/{_internal => toolkit}/weighted_aggregation/analyze_v3_acc_bins.py (100%) rename src/ops_model/models/interpretability/{_internal => toolkit}/weighted_aggregation/plot_v4_attn_comparison.py (100%) rename src/ops_model/models/interpretability/{_internal => toolkit}/weighted_aggregation/run_v3_pipeline_on_v4_attn_weighted.py (100%) rename src/ops_model/models/interpretability/{_internal => toolkit}/weighted_aggregation/run_v3_pipeline_on_v4_features.py (100%) diff --git a/pyproject.toml b/pyproject.toml index b0b0985..9a76857 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -85,8 +85,8 @@ source = "vcs" only-include = ["src"] sources = ["src"] # not-for-release: internal analysis/paper tooling (atlas, shap, titration, embedding, -# weighted_aggregation, viewer, kyle_pcs, gen_validation) lives under interpretability/_internal -exclude = ["src/ops_model/models/interpretability/_internal/**"] +# weighted_aggregation, viewer, kyle_pcs, gen_validation) lives under interpretability/toolkit +exclude = ["src/ops_model/models/interpretability/toolkit/**"] [tool.hatch.metadata] allow-direct-references = true diff --git a/src/ops_model/models/interpretability/_internal/RUNBOOK.md b/src/ops_model/models/interpretability/toolkit/RUNBOOK.md similarity index 100% rename from src/ops_model/models/interpretability/_internal/RUNBOOK.md rename to src/ops_model/models/interpretability/toolkit/RUNBOOK.md diff --git a/src/ops_model/models/interpretability/_internal/atlas/attention_accuracy_umap_animation.py b/src/ops_model/models/interpretability/toolkit/atlas/attention_accuracy_umap_animation.py similarity index 100% rename from src/ops_model/models/interpretability/_internal/atlas/attention_accuracy_umap_animation.py rename to src/ops_model/models/interpretability/toolkit/atlas/attention_accuracy_umap_animation.py diff --git a/src/ops_model/models/interpretability/_internal/atlas/attention_atlas.py b/src/ops_model/models/interpretability/toolkit/atlas/attention_atlas.py similarity index 100% rename from src/ops_model/models/interpretability/_internal/atlas/attention_atlas.py rename to src/ops_model/models/interpretability/toolkit/atlas/attention_atlas.py diff --git a/src/ops_model/models/interpretability/_internal/atlas/attention_atlas_shap.py b/src/ops_model/models/interpretability/toolkit/atlas/attention_atlas_shap.py similarity index 100% rename from src/ops_model/models/interpretability/_internal/atlas/attention_atlas_shap.py rename to src/ops_model/models/interpretability/toolkit/atlas/attention_atlas_shap.py diff --git a/src/ops_model/models/interpretability/_internal/atlas/low_attention_phase_atlas.py b/src/ops_model/models/interpretability/toolkit/atlas/low_attention_phase_atlas.py similarity index 100% rename from src/ops_model/models/interpretability/_internal/atlas/low_attention_phase_atlas.py rename to src/ops_model/models/interpretability/toolkit/atlas/low_attention_phase_atlas.py diff --git a/src/ops_model/models/interpretability/_internal/atlas/make_scale_bar.py b/src/ops_model/models/interpretability/toolkit/atlas/make_scale_bar.py similarity index 100% rename from src/ops_model/models/interpretability/_internal/atlas/make_scale_bar.py rename to src/ops_model/models/interpretability/toolkit/atlas/make_scale_bar.py diff --git a/src/ops_model/models/interpretability/_internal/atlas/marker_selection_distribution.py b/src/ops_model/models/interpretability/toolkit/atlas/marker_selection_distribution.py similarity index 100% rename from src/ops_model/models/interpretability/_internal/atlas/marker_selection_distribution.py rename to src/ops_model/models/interpretability/toolkit/atlas/marker_selection_distribution.py diff --git a/src/ops_model/models/interpretability/_internal/atlas/plot_eval_accuracy_curves.py b/src/ops_model/models/interpretability/toolkit/atlas/plot_eval_accuracy_curves.py similarity index 100% rename from src/ops_model/models/interpretability/_internal/atlas/plot_eval_accuracy_curves.py rename to src/ops_model/models/interpretability/toolkit/atlas/plot_eval_accuracy_curves.py diff --git a/src/ops_model/models/interpretability/_internal/embedding/generate_ko_violin_plots.py b/src/ops_model/models/interpretability/toolkit/embedding/generate_ko_violin_plots.py similarity index 100% rename from src/ops_model/models/interpretability/_internal/embedding/generate_ko_violin_plots.py rename to src/ops_model/models/interpretability/toolkit/embedding/generate_ko_violin_plots.py diff --git a/src/ops_model/models/interpretability/_internal/embedding/regen_umap_gav.py b/src/ops_model/models/interpretability/toolkit/embedding/regen_umap_gav.py similarity index 100% rename from src/ops_model/models/interpretability/_internal/embedding/regen_umap_gav.py rename to src/ops_model/models/interpretability/toolkit/embedding/regen_umap_gav.py diff --git a/src/ops_model/models/interpretability/_internal/embedding/regen_umap_html.py b/src/ops_model/models/interpretability/toolkit/embedding/regen_umap_html.py similarity index 100% rename from src/ops_model/models/interpretability/_internal/embedding/regen_umap_html.py rename to src/ops_model/models/interpretability/toolkit/embedding/regen_umap_html.py diff --git a/src/ops_model/models/interpretability/_internal/embedding/run_all_atlases.py b/src/ops_model/models/interpretability/toolkit/embedding/run_all_atlases.py similarity index 100% rename from src/ops_model/models/interpretability/_internal/embedding/run_all_atlases.py rename to src/ops_model/models/interpretability/toolkit/embedding/run_all_atlases.py diff --git a/src/ops_model/models/interpretability/_internal/embedding/top_attention_embed_and_score.py b/src/ops_model/models/interpretability/toolkit/embedding/top_attention_embed_and_score.py similarity index 100% rename from src/ops_model/models/interpretability/_internal/embedding/top_attention_embed_and_score.py rename to src/ops_model/models/interpretability/toolkit/embedding/top_attention_embed_and_score.py diff --git a/src/ops_model/models/interpretability/_internal/gen_validation/bag_sweep_plots.py b/src/ops_model/models/interpretability/toolkit/gen_validation/bag_sweep_plots.py similarity index 100% rename from src/ops_model/models/interpretability/_internal/gen_validation/bag_sweep_plots.py rename to src/ops_model/models/interpretability/toolkit/gen_validation/bag_sweep_plots.py diff --git a/src/ops_model/models/interpretability/_internal/gen_validation/bag_sweep_score.py b/src/ops_model/models/interpretability/toolkit/gen_validation/bag_sweep_score.py similarity index 92% rename from src/ops_model/models/interpretability/_internal/gen_validation/bag_sweep_score.py rename to src/ops_model/models/interpretability/toolkit/gen_validation/bag_sweep_score.py index 9a231d5..1e4e5aa 100644 --- a/src/ops_model/models/interpretability/_internal/gen_validation/bag_sweep_score.py +++ b/src/ops_model/models/interpretability/toolkit/gen_validation/bag_sweep_score.py @@ -16,8 +16,8 @@ def score_shard(genes): import torch - from ops_model.models.interpretability._internal.viewer.score_generated import score_embs_v5 - from ops_model.models.interpretability._internal.viewer.set_classifier import load_set_classifier, V5_CKPT_ROOT, V5_RUNS + from ops_model.models.interpretability.toolkit.viewer.score_generated import score_embs_v5 + from ops_model.models.interpretability.toolkit.viewer.set_classifier import load_set_classifier, V5_CKPT_ROOT, V5_RUNS os.makedirs(OUT, exist_ok=True) dev = "cuda" if torch.cuda.is_available() else "cpu" from ops_model.models.interpretability.diffae.classifier.config import slugify diff --git a/src/ops_model/models/interpretability/_internal/gen_validation/centroid_bagsweep.py b/src/ops_model/models/interpretability/toolkit/gen_validation/centroid_bagsweep.py similarity index 100% rename from src/ops_model/models/interpretability/_internal/gen_validation/centroid_bagsweep.py rename to src/ops_model/models/interpretability/toolkit/gen_validation/centroid_bagsweep.py diff --git a/src/ops_model/models/interpretability/_internal/gen_validation/centroid_halves.py b/src/ops_model/models/interpretability/toolkit/gen_validation/centroid_halves.py similarity index 100% rename from src/ops_model/models/interpretability/_internal/gen_validation/centroid_halves.py rename to src/ops_model/models/interpretability/toolkit/gen_validation/centroid_halves.py diff --git a/src/ops_model/models/interpretability/_internal/gen_validation/centroid_pooled_bagsweep.py b/src/ops_model/models/interpretability/toolkit/gen_validation/centroid_pooled_bagsweep.py similarity index 100% rename from src/ops_model/models/interpretability/_internal/gen_validation/centroid_pooled_bagsweep.py rename to src/ops_model/models/interpretability/toolkit/gen_validation/centroid_pooled_bagsweep.py diff --git a/src/ops_model/models/interpretability/_internal/gen_validation/control_halves_zscore.py b/src/ops_model/models/interpretability/toolkit/gen_validation/control_halves_zscore.py similarity index 100% rename from src/ops_model/models/interpretability/_internal/gen_validation/control_halves_zscore.py rename to src/ops_model/models/interpretability/toolkit/gen_validation/control_halves_zscore.py diff --git a/src/ops_model/models/interpretability/_internal/gen_validation/embcheck.py b/src/ops_model/models/interpretability/toolkit/gen_validation/embcheck.py similarity index 95% rename from src/ops_model/models/interpretability/_internal/gen_validation/embcheck.py rename to src/ops_model/models/interpretability/toolkit/gen_validation/embcheck.py index 087c6e0..f0f2df1 100644 --- a/src/ops_model/models/interpretability/_internal/gen_validation/embcheck.py +++ b/src/ops_model/models/interpretability/toolkit/gen_validation/embcheck.py @@ -8,7 +8,7 @@ def check(gene="AACS", ai=6): - from ops_model.models.interpretability._internal.viewer.score_generated import _emb_frames + from ops_model.models.interpretability.toolkit.viewer.score_generated import _emb_frames from ops_model.models.interpretability.diffae.classifier.celldino_features import embed_crops from ops_model.models.interpretability.diffae.directions.config import DirConfig trav = f"{B}/viewer_assets_valid200/phase/geneKO/{gene}" diff --git a/src/ops_model/models/interpretability/_internal/gen_validation/embedding_diagnostics.py b/src/ops_model/models/interpretability/toolkit/gen_validation/embedding_diagnostics.py similarity index 100% rename from src/ops_model/models/interpretability/_internal/gen_validation/embedding_diagnostics.py rename to src/ops_model/models/interpretability/toolkit/gen_validation/embedding_diagnostics.py diff --git a/src/ops_model/models/interpretability/_internal/gen_validation/figure4_bagsize_reachfrac.py b/src/ops_model/models/interpretability/toolkit/gen_validation/figure4_bagsize_reachfrac.py similarity index 100% rename from src/ops_model/models/interpretability/_internal/gen_validation/figure4_bagsize_reachfrac.py rename to src/ops_model/models/interpretability/toolkit/gen_validation/figure4_bagsize_reachfrac.py diff --git a/src/ops_model/models/interpretability/_internal/gen_validation/figure4_v5_accuracy_summary.py b/src/ops_model/models/interpretability/toolkit/gen_validation/figure4_v5_accuracy_summary.py similarity index 100% rename from src/ops_model/models/interpretability/_internal/gen_validation/figure4_v5_accuracy_summary.py rename to src/ops_model/models/interpretability/toolkit/gen_validation/figure4_v5_accuracy_summary.py diff --git a/src/ops_model/models/interpretability/_internal/gen_validation/gen_alpha_embedding.py b/src/ops_model/models/interpretability/toolkit/gen_validation/gen_alpha_embedding.py similarity index 100% rename from src/ops_model/models/interpretability/_internal/gen_validation/gen_alpha_embedding.py rename to src/ops_model/models/interpretability/toolkit/gen_validation/gen_alpha_embedding.py diff --git a/src/ops_model/models/interpretability/_internal/gen_validation/gen_embed_refit.py b/src/ops_model/models/interpretability/toolkit/gen_validation/gen_embed_refit.py similarity index 100% rename from src/ops_model/models/interpretability/_internal/gen_validation/gen_embed_refit.py rename to src/ops_model/models/interpretability/toolkit/gen_validation/gen_embed_refit.py diff --git a/src/ops_model/models/interpretability/_internal/gen_validation/gen_phate_passthrough.py b/src/ops_model/models/interpretability/toolkit/gen_validation/gen_phate_passthrough.py similarity index 100% rename from src/ops_model/models/interpretability/_internal/gen_validation/gen_phate_passthrough.py rename to src/ops_model/models/interpretability/toolkit/gen_validation/gen_phate_passthrough.py diff --git a/src/ops_model/models/interpretability/_internal/gen_validation/gen_real_centroid.py b/src/ops_model/models/interpretability/toolkit/gen_validation/gen_real_centroid.py similarity index 100% rename from src/ops_model/models/interpretability/_internal/gen_validation/gen_real_centroid.py rename to src/ops_model/models/interpretability/toolkit/gen_validation/gen_real_centroid.py diff --git a/src/ops_model/models/interpretability/_internal/gen_validation/gen_real_distinct.py b/src/ops_model/models/interpretability/toolkit/gen_validation/gen_real_distinct.py similarity index 100% rename from src/ops_model/models/interpretability/_internal/gen_validation/gen_real_distinct.py rename to src/ops_model/models/interpretability/toolkit/gen_validation/gen_real_distinct.py diff --git a/src/ops_model/models/interpretability/_internal/gen_validation/ntc_inverse_gap.py b/src/ops_model/models/interpretability/toolkit/gen_validation/ntc_inverse_gap.py similarity index 100% rename from src/ops_model/models/interpretability/_internal/gen_validation/ntc_inverse_gap.py rename to src/ops_model/models/interpretability/toolkit/gen_validation/ntc_inverse_gap.py diff --git a/src/ops_model/models/interpretability/_internal/gen_validation/patch_cache_real.py b/src/ops_model/models/interpretability/toolkit/gen_validation/patch_cache_real.py similarity index 100% rename from src/ops_model/models/interpretability/_internal/gen_validation/patch_cache_real.py rename to src/ops_model/models/interpretability/toolkit/gen_validation/patch_cache_real.py diff --git a/src/ops_model/models/interpretability/_internal/gen_validation/publish_multibag_page.py b/src/ops_model/models/interpretability/toolkit/gen_validation/publish_multibag_page.py similarity index 100% rename from src/ops_model/models/interpretability/_internal/gen_validation/publish_multibag_page.py rename to src/ops_model/models/interpretability/toolkit/gen_validation/publish_multibag_page.py diff --git a/src/ops_model/models/interpretability/_internal/gen_validation/rank_summary.py b/src/ops_model/models/interpretability/toolkit/gen_validation/rank_summary.py similarity index 100% rename from src/ops_model/models/interpretability/_internal/gen_validation/rank_summary.py rename to src/ops_model/models/interpretability/toolkit/gen_validation/rank_summary.py diff --git a/src/ops_model/models/interpretability/_internal/gen_validation/st_halves_score.py b/src/ops_model/models/interpretability/toolkit/gen_validation/st_halves_score.py similarity index 92% rename from src/ops_model/models/interpretability/_internal/gen_validation/st_halves_score.py rename to src/ops_model/models/interpretability/toolkit/gen_validation/st_halves_score.py index 3e41527..bb97a10 100644 --- a/src/ops_model/models/interpretability/_internal/gen_validation/st_halves_score.py +++ b/src/ops_model/models/interpretability/toolkit/gen_validation/st_halves_score.py @@ -12,8 +12,8 @@ def score_shard(genes): import torch - from ops_model.models.interpretability._internal.viewer.score_generated import score_embs_v5 - from ops_model.models.interpretability._internal.viewer.set_classifier import load_set_classifier, V5_CKPT_ROOT, V5_RUNS + from ops_model.models.interpretability.toolkit.viewer.score_generated import score_embs_v5 + from ops_model.models.interpretability.toolkit.viewer.set_classifier import load_set_classifier, V5_CKPT_ROOT, V5_RUNS from ops_model.models.interpretability.diffae.classifier.config import slugify os.makedirs(OUT, exist_ok=True) dev = "cuda" if torch.cuda.is_available() else "cpu" diff --git a/src/ops_model/models/interpretability/_internal/gen_validation/std_anchor_test.py b/src/ops_model/models/interpretability/toolkit/gen_validation/std_anchor_test.py similarity index 100% rename from src/ops_model/models/interpretability/_internal/gen_validation/std_anchor_test.py rename to src/ops_model/models/interpretability/toolkit/gen_validation/std_anchor_test.py diff --git a/src/ops_model/models/interpretability/_internal/gen_validation/stepabl_compare.py b/src/ops_model/models/interpretability/toolkit/gen_validation/stepabl_compare.py similarity index 100% rename from src/ops_model/models/interpretability/_internal/gen_validation/stepabl_compare.py rename to src/ops_model/models/interpretability/toolkit/gen_validation/stepabl_compare.py diff --git a/src/ops_model/models/interpretability/_internal/gen_validation/valid200_alphastep.py b/src/ops_model/models/interpretability/toolkit/gen_validation/valid200_alphastep.py similarity index 100% rename from src/ops_model/models/interpretability/_internal/gen_validation/valid200_alphastep.py rename to src/ops_model/models/interpretability/toolkit/gen_validation/valid200_alphastep.py diff --git a/src/ops_model/models/interpretability/_internal/gen_validation/valid200_cache_build.py b/src/ops_model/models/interpretability/toolkit/gen_validation/valid200_cache_build.py similarity index 100% rename from src/ops_model/models/interpretability/_internal/gen_validation/valid200_cache_build.py rename to src/ops_model/models/interpretability/toolkit/gen_validation/valid200_cache_build.py diff --git a/src/ops_model/models/interpretability/_internal/gen_validation/valid200_capcheck.py b/src/ops_model/models/interpretability/toolkit/gen_validation/valid200_capcheck.py similarity index 100% rename from src/ops_model/models/interpretability/_internal/gen_validation/valid200_capcheck.py rename to src/ops_model/models/interpretability/toolkit/gen_validation/valid200_capcheck.py diff --git a/src/ops_model/models/interpretability/_internal/gen_validation/valid200_map_compare.py b/src/ops_model/models/interpretability/toolkit/gen_validation/valid200_map_compare.py similarity index 100% rename from src/ops_model/models/interpretability/_internal/gen_validation/valid200_map_compare.py rename to src/ops_model/models/interpretability/toolkit/gen_validation/valid200_map_compare.py diff --git a/src/ops_model/models/interpretability/_internal/gen_validation/valid200_metrics.py b/src/ops_model/models/interpretability/toolkit/gen_validation/valid200_metrics.py similarity index 100% rename from src/ops_model/models/interpretability/_internal/gen_validation/valid200_metrics.py rename to src/ops_model/models/interpretability/toolkit/gen_validation/valid200_metrics.py diff --git a/src/ops_model/models/interpretability/_internal/kyle_pcs/build_static_explorer.py b/src/ops_model/models/interpretability/toolkit/kyle_pcs/build_static_explorer.py similarity index 100% rename from src/ops_model/models/interpretability/_internal/kyle_pcs/build_static_explorer.py rename to src/ops_model/models/interpretability/toolkit/kyle_pcs/build_static_explorer.py diff --git a/src/ops_model/models/interpretability/_internal/kyle_pcs/compute_pc_strips.py b/src/ops_model/models/interpretability/toolkit/kyle_pcs/compute_pc_strips.py similarity index 100% rename from src/ops_model/models/interpretability/_internal/kyle_pcs/compute_pc_strips.py rename to src/ops_model/models/interpretability/toolkit/kyle_pcs/compute_pc_strips.py diff --git a/src/ops_model/models/interpretability/_internal/shap/analyze_chad_variants.py b/src/ops_model/models/interpretability/toolkit/shap/analyze_chad_variants.py similarity index 100% rename from src/ops_model/models/interpretability/_internal/shap/analyze_chad_variants.py rename to src/ops_model/models/interpretability/toolkit/shap/analyze_chad_variants.py diff --git a/src/ops_model/models/interpretability/_internal/shap/generate_shap_captions_combined.py b/src/ops_model/models/interpretability/toolkit/shap/generate_shap_captions_combined.py similarity index 100% rename from src/ops_model/models/interpretability/_internal/shap/generate_shap_captions_combined.py rename to src/ops_model/models/interpretability/toolkit/shap/generate_shap_captions_combined.py diff --git a/src/ops_model/models/interpretability/_internal/shap/ko_shap_features.py b/src/ops_model/models/interpretability/toolkit/shap/ko_shap_features.py similarity index 100% rename from src/ops_model/models/interpretability/_internal/shap/ko_shap_features.py rename to src/ops_model/models/interpretability/toolkit/shap/ko_shap_features.py diff --git a/src/ops_model/models/interpretability/_internal/shap/merge_shap_shards.py b/src/ops_model/models/interpretability/toolkit/shap/merge_shap_shards.py similarity index 100% rename from src/ops_model/models/interpretability/_internal/shap/merge_shap_shards.py rename to src/ops_model/models/interpretability/toolkit/shap/merge_shap_shards.py diff --git a/src/ops_model/models/interpretability/_internal/shap/ntc_attention_compare.py b/src/ops_model/models/interpretability/toolkit/shap/ntc_attention_compare.py similarity index 100% rename from src/ops_model/models/interpretability/_internal/shap/ntc_attention_compare.py rename to src/ops_model/models/interpretability/toolkit/shap/ntc_attention_compare.py diff --git a/src/ops_model/models/interpretability/_internal/shap/ntc_pick_cells.py b/src/ops_model/models/interpretability/toolkit/shap/ntc_pick_cells.py similarity index 100% rename from src/ops_model/models/interpretability/_internal/shap/ntc_pick_cells.py rename to src/ops_model/models/interpretability/toolkit/shap/ntc_pick_cells.py diff --git a/src/ops_model/models/interpretability/_internal/shap/ntc_shap_features.py b/src/ops_model/models/interpretability/toolkit/shap/ntc_shap_features.py similarity index 100% rename from src/ops_model/models/interpretability/_internal/shap/ntc_shap_features.py rename to src/ops_model/models/interpretability/toolkit/shap/ntc_shap_features.py diff --git a/src/ops_model/models/interpretability/_internal/shap/run_all_shap.py b/src/ops_model/models/interpretability/toolkit/shap/run_all_shap.py similarity index 100% rename from src/ops_model/models/interpretability/_internal/shap/run_all_shap.py rename to src/ops_model/models/interpretability/toolkit/shap/run_all_shap.py diff --git a/src/ops_model/models/interpretability/_internal/shap/run_shap_pipeline.py b/src/ops_model/models/interpretability/toolkit/shap/run_shap_pipeline.py similarity index 100% rename from src/ops_model/models/interpretability/_internal/shap/run_shap_pipeline.py rename to src/ops_model/models/interpretability/toolkit/shap/run_shap_pipeline.py diff --git a/src/ops_model/models/interpretability/_internal/shap/shap_approach_compare.py b/src/ops_model/models/interpretability/toolkit/shap/shap_approach_compare.py similarity index 100% rename from src/ops_model/models/interpretability/_internal/shap/shap_approach_compare.py rename to src/ops_model/models/interpretability/toolkit/shap/shap_approach_compare.py diff --git a/src/ops_model/models/interpretability/_internal/titration/decay/map_attention_decay.py b/src/ops_model/models/interpretability/toolkit/titration/decay/map_attention_decay.py similarity index 100% rename from src/ops_model/models/interpretability/_internal/titration/decay/map_attention_decay.py rename to src/ops_model/models/interpretability/toolkit/titration/decay/map_attention_decay.py diff --git a/src/ops_model/models/interpretability/_internal/titration/decay/phate_peak_groups.py b/src/ops_model/models/interpretability/toolkit/titration/decay/phate_peak_groups.py similarity index 100% rename from src/ops_model/models/interpretability/_internal/titration/decay/phate_peak_groups.py rename to src/ops_model/models/interpretability/toolkit/titration/decay/phate_peak_groups.py diff --git a/src/ops_model/models/interpretability/_internal/titration/decay/plot_3way_summary_bars.py b/src/ops_model/models/interpretability/toolkit/titration/decay/plot_3way_summary_bars.py similarity index 100% rename from src/ops_model/models/interpretability/_internal/titration/decay/plot_3way_summary_bars.py rename to src/ops_model/models/interpretability/toolkit/titration/decay/plot_3way_summary_bars.py diff --git a/src/ops_model/models/interpretability/_internal/titration/decay/plot_all_cells_correction_bars.py b/src/ops_model/models/interpretability/toolkit/titration/decay/plot_all_cells_correction_bars.py similarity index 100% rename from src/ops_model/models/interpretability/_internal/titration/decay/plot_all_cells_correction_bars.py rename to src/ops_model/models/interpretability/toolkit/titration/decay/plot_all_cells_correction_bars.py diff --git a/src/ops_model/models/interpretability/_internal/titration/expansion/count_genes_above_threshold.py b/src/ops_model/models/interpretability/toolkit/titration/expansion/count_genes_above_threshold.py similarity index 98% rename from src/ops_model/models/interpretability/_internal/titration/expansion/count_genes_above_threshold.py rename to src/ops_model/models/interpretability/toolkit/titration/expansion/count_genes_above_threshold.py index e6e1b45..d8a2559 100644 --- a/src/ops_model/models/interpretability/_internal/titration/expansion/count_genes_above_threshold.py +++ b/src/ops_model/models/interpretability/toolkit/titration/expansion/count_genes_above_threshold.py @@ -17,10 +17,10 @@ Usage:: # Submit one SLURM task per K (9 tasks; ~5 min wall once they land) - uv run python -m ops_model.models.interpretability._internal.titration.expansion.count_genes_above_threshold --slurm + uv run python -m ops_model.models.interpretability.toolkit.titration.expansion.count_genes_above_threshold --slurm # Replot from cached per-gene CSVs (no SLURM) - uv run python -m ops_model.models.interpretability._internal.titration.expansion.count_genes_above_threshold --replot + uv run python -m ops_model.models.interpretability.toolkit.titration.expansion.count_genes_above_threshold --replot """ from __future__ import annotations diff --git a/src/ops_model/models/interpretability/_internal/titration/expansion/map_attention_expansion_v4.py b/src/ops_model/models/interpretability/toolkit/titration/expansion/map_attention_expansion_v4.py similarity index 100% rename from src/ops_model/models/interpretability/_internal/titration/expansion/map_attention_expansion_v4.py rename to src/ops_model/models/interpretability/toolkit/titration/expansion/map_attention_expansion_v4.py diff --git a/src/ops_model/models/interpretability/_internal/titration/expansion/plot_sgrna_coverage_sweep.py b/src/ops_model/models/interpretability/toolkit/titration/expansion/plot_sgrna_coverage_sweep.py similarity index 100% rename from src/ops_model/models/interpretability/_internal/titration/expansion/plot_sgrna_coverage_sweep.py rename to src/ops_model/models/interpretability/toolkit/titration/expansion/plot_sgrna_coverage_sweep.py diff --git a/src/ops_model/models/interpretability/_internal/titration/expansion/run_percentile_sweep.py b/src/ops_model/models/interpretability/toolkit/titration/expansion/run_percentile_sweep.py similarity index 100% rename from src/ops_model/models/interpretability/_internal/titration/expansion/run_percentile_sweep.py rename to src/ops_model/models/interpretability/toolkit/titration/expansion/run_percentile_sweep.py diff --git a/src/ops_model/models/interpretability/_internal/viewer/__init__.py b/src/ops_model/models/interpretability/toolkit/viewer/__init__.py similarity index 100% rename from src/ops_model/models/interpretability/_internal/viewer/__init__.py rename to src/ops_model/models/interpretability/toolkit/viewer/__init__.py diff --git a/src/ops_model/models/interpretability/_internal/viewer/_altanchor_build.py b/src/ops_model/models/interpretability/toolkit/viewer/_altanchor_build.py similarity index 100% rename from src/ops_model/models/interpretability/_internal/viewer/_altanchor_build.py rename to src/ops_model/models/interpretability/toolkit/viewer/_altanchor_build.py diff --git a/src/ops_model/models/interpretability/_internal/viewer/_anchortest.py b/src/ops_model/models/interpretability/toolkit/viewer/_anchortest.py similarity index 100% rename from src/ops_model/models/interpretability/_internal/viewer/_anchortest.py rename to src/ops_model/models/interpretability/toolkit/viewer/_anchortest.py diff --git a/src/ops_model/models/interpretability/_internal/viewer/_build_stepablation.py b/src/ops_model/models/interpretability/toolkit/viewer/_build_stepablation.py similarity index 100% rename from src/ops_model/models/interpretability/_internal/viewer/_build_stepablation.py rename to src/ops_model/models/interpretability/toolkit/viewer/_build_stepablation.py diff --git a/src/ops_model/models/interpretability/_internal/viewer/_build_v5_inverted.py b/src/ops_model/models/interpretability/toolkit/viewer/_build_v5_inverted.py similarity index 99% rename from src/ops_model/models/interpretability/_internal/viewer/_build_v5_inverted.py rename to src/ops_model/models/interpretability/toolkit/viewer/_build_v5_inverted.py index 6a7bba9..db9101c 100644 --- a/src/ops_model/models/interpretability/_internal/viewer/_build_v5_inverted.py +++ b/src/ops_model/models/interpretability/toolkit/viewer/_build_v5_inverted.py @@ -9,7 +9,7 @@ complex) whose per-cell rankings exist. That is Alex Lin's top1_acc>0.5@100-cell distinctiveness filter (see _fluor_v5_build.py) — NOT missing data (cells exist for all 1000×55). Lower-acc combos need Alex to gen more. - python -m ops_model.models.interpretability._internal.viewer._build_v5_inverted markers + python -m ops_model.models.interpretability.toolkit.viewer._build_v5_inverted markers """ import json import os diff --git a/src/ops_model/models/interpretability/_internal/viewer/_build_v5_montages.py b/src/ops_model/models/interpretability/toolkit/viewer/_build_v5_montages.py similarity index 98% rename from src/ops_model/models/interpretability/_internal/viewer/_build_v5_montages.py rename to src/ops_model/models/interpretability/toolkit/viewer/_build_v5_montages.py index 290880d..2edd73f 100644 --- a/src/ops_model/models/interpretability/_internal/viewer/_build_v5_montages.py +++ b/src/ops_model/models/interpretability/toolkit/viewer/_build_v5_montages.py @@ -5,7 +5,7 @@ phase embedding. Reads the merged inverted frames from viewer_assets_v5//geneKO and writes tiles to viewer_assets_v5/_montage/. Only builds markers whose geneKO is 100% present in viewer_assets_v5. - python -m ops_model.models.interpretability._internal.viewer._build_v5_montages + python -m ops_model.models.interpretability.toolkit.viewer._build_v5_montages """ import glob import os diff --git a/src/ops_model/models/interpretability/_internal/viewer/_build_valid200.py b/src/ops_model/models/interpretability/toolkit/viewer/_build_valid200.py similarity index 95% rename from src/ops_model/models/interpretability/_internal/viewer/_build_valid200.py rename to src/ops_model/models/interpretability/toolkit/viewer/_build_valid200.py index bdebc51..ecacbbb 100644 --- a/src/ops_model/models/interpretability/_internal/viewer/_build_valid200.py +++ b/src/ops_model/models/interpretability/toolkit/viewer/_build_valid200.py @@ -8,9 +8,9 @@ Reuses the v5 per-class directions (d_vec + gap are the PRODUCTION values we are validating) via a symlinked _directions tree, so no direction re-fit — only the 200-anchor inversion + 200×7 decodes per target. - python -m ops_model.models.interpretability._internal.viewer._build_valid200 anchor # 1 GPU: build the 200-cell NTC anchor - python -m ops_model.models.interpretability._internal.viewer._build_valid200 submit # shard all 1000 geneKO (after anchor) - python -m ops_model.models.interpretability._internal.viewer._build_valid200 all # anchor job -> shards (afterok dep) + python -m ops_model.models.interpretability.toolkit.viewer._build_valid200 anchor # 1 GPU: build the 200-cell NTC anchor + python -m ops_model.models.interpretability.toolkit.viewer._build_valid200 submit # shard all 1000 geneKO (after anchor) + python -m ops_model.models.interpretability.toolkit.viewer._build_valid200 all # anchor job -> shards (afterok dep) """ import os import sys diff --git a/src/ops_model/models/interpretability/_internal/viewer/_consolidate_cells.py b/src/ops_model/models/interpretability/toolkit/viewer/_consolidate_cells.py similarity index 100% rename from src/ops_model/models/interpretability/_internal/viewer/_consolidate_cells.py rename to src/ops_model/models/interpretability/toolkit/viewer/_consolidate_cells.py diff --git a/src/ops_model/models/interpretability/_internal/viewer/_fluor_complex_build.py b/src/ops_model/models/interpretability/toolkit/viewer/_fluor_complex_build.py similarity index 100% rename from src/ops_model/models/interpretability/_internal/viewer/_fluor_complex_build.py rename to src/ops_model/models/interpretability/toolkit/viewer/_fluor_complex_build.py diff --git a/src/ops_model/models/interpretability/_internal/viewer/_fluor_v5_build.py b/src/ops_model/models/interpretability/toolkit/viewer/_fluor_v5_build.py similarity index 100% rename from src/ops_model/models/interpretability/_internal/viewer/_fluor_v5_build.py rename to src/ops_model/models/interpretability/toolkit/viewer/_fluor_v5_build.py diff --git a/src/ops_model/models/interpretability/_internal/viewer/_migrate_v4_to_v5.py b/src/ops_model/models/interpretability/toolkit/viewer/_migrate_v4_to_v5.py similarity index 100% rename from src/ops_model/models/interpretability/_internal/viewer/_migrate_v4_to_v5.py rename to src/ops_model/models/interpretability/toolkit/viewer/_migrate_v4_to_v5.py diff --git a/src/ops_model/models/interpretability/_internal/viewer/_phase_vs.py b/src/ops_model/models/interpretability/toolkit/viewer/_phase_vs.py similarity index 100% rename from src/ops_model/models/interpretability/_internal/viewer/_phase_vs.py rename to src/ops_model/models/interpretability/toolkit/viewer/_phase_vs.py diff --git a/src/ops_model/models/interpretability/_internal/viewer/_rebuild_v5.py b/src/ops_model/models/interpretability/toolkit/viewer/_rebuild_v5.py similarity index 100% rename from src/ops_model/models/interpretability/_internal/viewer/_rebuild_v5.py rename to src/ops_model/models/interpretability/toolkit/viewer/_rebuild_v5.py diff --git a/src/ops_model/models/interpretability/_internal/viewer/_rescore_rank.py b/src/ops_model/models/interpretability/toolkit/viewer/_rescore_rank.py similarity index 96% rename from src/ops_model/models/interpretability/_internal/viewer/_rescore_rank.py rename to src/ops_model/models/interpretability/toolkit/viewer/_rescore_rank.py index 6676e4e..963291e 100644 --- a/src/ops_model/models/interpretability/_internal/viewer/_rescore_rank.py +++ b/src/ops_model/models/interpretability/toolkit/viewer/_rescore_rank.py @@ -11,7 +11,7 @@ def _retarget(assets): """Point the score module at `assets` and return the module (V5_BASE is read at call time).""" - import ops_model.models.interpretability._internal.viewer.score_generated as SG + import ops_model.models.interpretability.toolkit.viewer.score_generated as SG SG.V5_BASE = f"{BASE}/{assets}/phase" return SG diff --git a/src/ops_model/models/interpretability/_internal/viewer/_score_v4.py b/src/ops_model/models/interpretability/toolkit/viewer/_score_v4.py similarity index 100% rename from src/ops_model/models/interpretability/_internal/viewer/_score_v4.py rename to src/ops_model/models/interpretability/toolkit/viewer/_score_v4.py diff --git a/src/ops_model/models/interpretability/_internal/viewer/_v4acc_test.py b/src/ops_model/models/interpretability/toolkit/viewer/_v4acc_test.py similarity index 100% rename from src/ops_model/models/interpretability/_internal/viewer/_v4acc_test.py rename to src/ops_model/models/interpretability/toolkit/viewer/_v4acc_test.py diff --git a/src/ops_model/models/interpretability/_internal/viewer/_verify_pt_space.py b/src/ops_model/models/interpretability/toolkit/viewer/_verify_pt_space.py similarity index 100% rename from src/ops_model/models/interpretability/_internal/viewer/_verify_pt_space.py rename to src/ops_model/models/interpretability/toolkit/viewer/_verify_pt_space.py diff --git a/src/ops_model/models/interpretability/_internal/viewer/_verify_score_bridge.py b/src/ops_model/models/interpretability/toolkit/viewer/_verify_score_bridge.py similarity index 100% rename from src/ops_model/models/interpretability/_internal/viewer/_verify_score_bridge.py rename to src/ops_model/models/interpretability/toolkit/viewer/_verify_score_bridge.py diff --git a/src/ops_model/models/interpretability/_internal/viewer/altanchor_pairs.json b/src/ops_model/models/interpretability/toolkit/viewer/altanchor_pairs.json similarity index 100% rename from src/ops_model/models/interpretability/_internal/viewer/altanchor_pairs.json rename to src/ops_model/models/interpretability/toolkit/viewer/altanchor_pairs.json diff --git a/src/ops_model/models/interpretability/_internal/viewer/anchor_cells.py b/src/ops_model/models/interpretability/toolkit/viewer/anchor_cells.py similarity index 100% rename from src/ops_model/models/interpretability/_internal/viewer/anchor_cells.py rename to src/ops_model/models/interpretability/toolkit/viewer/anchor_cells.py diff --git a/src/ops_model/models/interpretability/_internal/viewer/build_attention_heads.py b/src/ops_model/models/interpretability/toolkit/viewer/build_attention_heads.py similarity index 95% rename from src/ops_model/models/interpretability/_internal/viewer/build_attention_heads.py rename to src/ops_model/models/interpretability/toolkit/viewer/build_attention_heads.py index c6cd9d6..1712e4e 100644 --- a/src/ops_model/models/interpretability/_internal/viewer/build_attention_heads.py +++ b/src/ops_model/models/interpretability/toolkit/viewer/build_attention_heads.py @@ -18,10 +18,10 @@ {AH}/index.json {global_max, assets:{:{:[keys]}}} where modality = "phase" | slugify(marker_channel), grain = geneKO|complex, key = gene | complex-slug. - python -m ops_model.models.interpretability._internal.viewer.build_attention_heads render # SLURM (all trees) - python -m ops_model.models.interpretability._internal.viewer.build_attention_heads render --local # serial, no SLURM - python -m ops_model.models.interpretability._internal.viewer.build_attention_heads render --dry-run - python -m ops_model.models.interpretability._internal.viewer.build_attention_heads index # (re)aggregate index.json + python -m ops_model.models.interpretability.toolkit.viewer.build_attention_heads render # SLURM (all trees) + python -m ops_model.models.interpretability.toolkit.viewer.build_attention_heads render --local # serial, no SLURM + python -m ops_model.models.interpretability.toolkit.viewer.build_attention_heads render --dry-run + python -m ops_model.models.interpretability.toolkit.viewer.build_attention_heads index # (re)aggregate index.json """ from __future__ import annotations diff --git a/src/ops_model/models/interpretability/_internal/viewer/build_complex_ebi_map.py b/src/ops_model/models/interpretability/toolkit/viewer/build_complex_ebi_map.py similarity index 100% rename from src/ops_model/models/interpretability/_internal/viewer/build_complex_ebi_map.py rename to src/ops_model/models/interpretability/toolkit/viewer/build_complex_ebi_map.py diff --git a/src/ops_model/models/interpretability/_internal/viewer/build_fluor_shap_rankings.py b/src/ops_model/models/interpretability/toolkit/viewer/build_fluor_shap_rankings.py similarity index 96% rename from src/ops_model/models/interpretability/_internal/viewer/build_fluor_shap_rankings.py rename to src/ops_model/models/interpretability/toolkit/viewer/build_fluor_shap_rankings.py index 496f7dd..d16bbd3 100644 --- a/src/ops_model/models/interpretability/_internal/viewer/build_fluor_shap_rankings.py +++ b/src/ops_model/models/interpretability/toolkit/viewer/build_fluor_shap_rankings.py @@ -5,8 +5,8 @@ new CSV cols: gene, channel_name, rank, shap, ..., experiment, well, x_pheno, y_pheno, segmentation_id old schema: channel_name, gene, rank, pma_attention, experiment, well, x_pheno, y_pheno, segmentation, rank_type - python -m ops_model.models.interpretability._internal.viewer.build_fluor_shap_rankings # local (needs ~64GB) - python -m ops_model.models.interpretability._internal.viewer.build_fluor_shap_rankings --submit # SLURM cpu, mem 96 + python -m ops_model.models.interpretability.toolkit.viewer.build_fluor_shap_rankings # local (needs ~64GB) + python -m ops_model.models.interpretability.toolkit.viewer.build_fluor_shap_rankings --submit # SLURM cpu, mem 96 """ from __future__ import annotations diff --git a/src/ops_model/models/interpretability/_internal/viewer/build_montage_features.py b/src/ops_model/models/interpretability/toolkit/viewer/build_montage_features.py similarity index 96% rename from src/ops_model/models/interpretability/_internal/viewer/build_montage_features.py rename to src/ops_model/models/interpretability/toolkit/viewer/build_montage_features.py index b1fafaa..2eea3cb 100644 --- a/src/ops_model/models/interpretability/_internal/viewer/build_montage_features.py +++ b/src/ops_model/models/interpretability/toolkit/viewer/build_montage_features.py @@ -7,7 +7,7 @@ viewer_assets/montage_features.json {"features": [base names], "range": {feat: [lo, hi]}, "values": {gene: [0..1 per feature | null]}} - python -m ops_model.models.interpretability._internal.viewer.build_montage_features + python -m ops_model.models.interpretability.toolkit.viewer.build_montage_features """ from __future__ import annotations diff --git a/src/ops_model/models/interpretability/_internal/viewer/build_pc_features.py b/src/ops_model/models/interpretability/toolkit/viewer/build_pc_features.py similarity index 99% rename from src/ops_model/models/interpretability/_internal/viewer/build_pc_features.py rename to src/ops_model/models/interpretability/toolkit/viewer/build_pc_features.py index dee29f8..97ac53f 100644 --- a/src/ops_model/models/interpretability/_internal/viewer/build_pc_features.py +++ b/src/ops_model/models/interpretability/toolkit/viewer/build_pc_features.py @@ -14,7 +14,7 @@ to make "+corr features" line up with the strip's high bins. Compositions are unsigned (|r| / tf-idf) so they need no flip. - python -m ops_model.models.interpretability._internal.viewer.build_pc_features + python -m ops_model.models.interpretability.toolkit.viewer.build_pc_features """ from __future__ import annotations diff --git a/src/ops_model/models/interpretability/_internal/viewer/build_pc_walks.py b/src/ops_model/models/interpretability/toolkit/viewer/build_pc_walks.py similarity index 97% rename from src/ops_model/models/interpretability/_internal/viewer/build_pc_walks.py rename to src/ops_model/models/interpretability/toolkit/viewer/build_pc_walks.py index 322bbc2..143db8a 100644 --- a/src/ops_model/models/interpretability/_internal/viewer/build_pc_walks.py +++ b/src/ops_model/models/interpretability/toolkit/viewer/build_pc_walks.py @@ -8,8 +8,8 @@ the PC score std; v_p is the (z-scored) eigenvector, mapped back to raw CellDINO space by the mean per-exp sd. Output: one composite figure per marker (rows = PCs, cols = α). - python -m ops_model.models.interpretability._internal.viewer.build_pc_walks --markers "Mitochondria_TOMM20" - python -m ops_model.models.interpretability._internal.viewer.build_pc_walks --all # SLURM, every marker + python -m ops_model.models.interpretability.toolkit.viewer.build_pc_walks --markers "Mitochondria_TOMM20" + python -m ops_model.models.interpretability.toolkit.viewer.build_pc_walks --all # SLURM, every marker """ from __future__ import annotations diff --git a/src/ops_model/models/interpretability/_internal/viewer/build_pcs.py b/src/ops_model/models/interpretability/toolkit/viewer/build_pcs.py similarity index 97% rename from src/ops_model/models/interpretability/_internal/viewer/build_pcs.py rename to src/ops_model/models/interpretability/toolkit/viewer/build_pcs.py index ff4aa06..0abdd9f 100644 --- a/src/ops_model/models/interpretability/_internal/viewer/build_pcs.py +++ b/src/ops_model/models/interpretability/toolkit/viewer/build_pcs.py @@ -9,8 +9,8 @@ the raw artifacts dir is gone). If Kyle regenerates artifacts, re-run his build_static_explorer.py then point --html at the fresh output. - python -m ops_model.models.interpretability._internal.viewer.build_pcs - python -m ops_model.models.interpretability._internal.viewer.build_pcs --html /path/to/pc_explorer_static.html + python -m ops_model.models.interpretability.toolkit.viewer.build_pcs + python -m ops_model.models.interpretability.toolkit.viewer.build_pcs --html /path/to/pc_explorer_static.html """ from __future__ import annotations diff --git a/src/ops_model/models/interpretability/_internal/viewer/build_pcs_marker.py b/src/ops_model/models/interpretability/toolkit/viewer/build_pcs_marker.py similarity index 99% rename from src/ops_model/models/interpretability/_internal/viewer/build_pcs_marker.py rename to src/ops_model/models/interpretability/toolkit/viewer/build_pcs_marker.py index 68b15eb..9c563c3 100644 --- a/src/ops_model/models/interpretability/_internal/viewer/build_pcs_marker.py +++ b/src/ops_model/models/interpretability/toolkit/viewer/build_pcs_marker.py @@ -9,7 +9,7 @@ viewer_assets/pcs/markers//index.json (same schema as the phase pcs/index.json) viewer_assets/pcs/markers//crops/pc###_bin##_row#.png - python -m ops_model.models.interpretability._internal.viewer.build_pcs_marker --marker "autophagosome_MAP1LC3B" + python -m ops_model.models.interpretability.toolkit.viewer.build_pcs_marker --marker "autophagosome_MAP1LC3B" """ from __future__ import annotations diff --git a/src/ops_model/models/interpretability/_internal/viewer/build_phase_shap_rankings.py b/src/ops_model/models/interpretability/toolkit/viewer/build_phase_shap_rankings.py similarity index 95% rename from src/ops_model/models/interpretability/_internal/viewer/build_phase_shap_rankings.py rename to src/ops_model/models/interpretability/toolkit/viewer/build_phase_shap_rankings.py index 760ac48..279888f 100644 --- a/src/ops_model/models/interpretability/_internal/viewer/build_phase_shap_rankings.py +++ b/src/ops_model/models/interpretability/toolkit/viewer/build_phase_shap_rankings.py @@ -6,8 +6,8 @@ pma_attention, rank, rank_type) complex → pma_shap_phase_complex.parquet (adds predicted_class=complex, gene=member gene; EBI-pooled) - python -m ops_model.models.interpretability._internal.viewer.build_phase_shap_rankings --geneko --submit - python -m ops_model.models.interpretability._internal.viewer.build_phase_shap_rankings --complex --submit + python -m ops_model.models.interpretability.toolkit.viewer.build_phase_shap_rankings --geneko --submit + python -m ops_model.models.interpretability.toolkit.viewer.build_phase_shap_rankings --complex --submit """ from __future__ import annotations diff --git a/src/ops_model/models/interpretability/_internal/viewer/build_phate_figure.py b/src/ops_model/models/interpretability/toolkit/viewer/build_phate_figure.py similarity index 99% rename from src/ops_model/models/interpretability/_internal/viewer/build_phate_figure.py rename to src/ops_model/models/interpretability/toolkit/viewer/build_phate_figure.py index 0e99bab..bddceec 100644 --- a/src/ops_model/models/interpretability/_internal/viewer/build_phate_figure.py +++ b/src/ops_model/models/interpretability/toolkit/viewer/build_phate_figure.py @@ -5,7 +5,7 @@ Each panel: the same PHATE scatter (grey), that panel's groups colored + leader-labelled with the single-cell generated morph (NTC cell1 → group, alpha=+5). NTC original shown top-left of panel E. - python -m ops_model.models.interpretability._internal.viewer.build_phate_figure + python -m ops_model.models.interpretability.toolkit.viewer.build_phate_figure """ from __future__ import annotations diff --git a/src/ops_model/models/interpretability/_internal/viewer/build_setacc_bins.py b/src/ops_model/models/interpretability/toolkit/viewer/build_setacc_bins.py similarity index 100% rename from src/ops_model/models/interpretability/_internal/viewer/build_setacc_bins.py rename to src/ops_model/models/interpretability/toolkit/viewer/build_setacc_bins.py diff --git a/src/ops_model/models/interpretability/_internal/viewer/build_setacc_bymarker.py b/src/ops_model/models/interpretability/toolkit/viewer/build_setacc_bymarker.py similarity index 100% rename from src/ops_model/models/interpretability/_internal/viewer/build_setacc_bymarker.py rename to src/ops_model/models/interpretability/toolkit/viewer/build_setacc_bymarker.py diff --git a/src/ops_model/models/interpretability/_internal/viewer/build_top_cells.py b/src/ops_model/models/interpretability/toolkit/viewer/build_top_cells.py similarity index 97% rename from src/ops_model/models/interpretability/_internal/viewer/build_top_cells.py rename to src/ops_model/models/interpretability/toolkit/viewer/build_top_cells.py index c99c04e..2d8213c 100644 --- a/src/ops_model/models/interpretability/_internal/viewer/build_top_cells.py +++ b/src/ops_model/models/interpretability/toolkit/viewer/build_top_cells.py @@ -7,8 +7,8 @@ viewer_assets_v5/top_cells/index.json {"top_n", "genes"|"complexes": {CLASS: {"accuracy": [rec...]}}} viewer_assets_v5/top_cells/crops/.png - python -m ops_model.models.interpretability._internal.viewer.build_top_cells geneKO # SLURM crop shards + finalize - python -m ops_model.models.interpretability._internal.viewer.build_top_cells complex --finalize # rebuild index only + python -m ops_model.models.interpretability.toolkit.viewer.build_top_cells geneKO # SLURM crop shards + finalize + python -m ops_model.models.interpretability.toolkit.viewer.build_top_cells complex --finalize # rebuild index only """ from __future__ import annotations diff --git a/src/ops_model/models/interpretability/_internal/viewer/build_umap_montage.py b/src/ops_model/models/interpretability/toolkit/viewer/build_umap_montage.py similarity index 100% rename from src/ops_model/models/interpretability/_internal/viewer/build_umap_montage.py rename to src/ops_model/models/interpretability/toolkit/viewer/build_umap_montage.py diff --git a/src/ops_model/models/interpretability/_internal/viewer/deploy/README.md b/src/ops_model/models/interpretability/toolkit/viewer/deploy/README.md similarity index 100% rename from src/ops_model/models/interpretability/_internal/viewer/deploy/README.md rename to src/ops_model/models/interpretability/toolkit/viewer/deploy/README.md diff --git a/src/ops_model/models/interpretability/_internal/viewer/marker_leaves.py b/src/ops_model/models/interpretability/toolkit/viewer/marker_leaves.py similarity index 100% rename from src/ops_model/models/interpretability/_internal/viewer/marker_leaves.py rename to src/ops_model/models/interpretability/toolkit/viewer/marker_leaves.py diff --git a/src/ops_model/models/interpretability/_internal/viewer/mimic_alex_embed.py b/src/ops_model/models/interpretability/toolkit/viewer/mimic_alex_embed.py similarity index 100% rename from src/ops_model/models/interpretability/_internal/viewer/mimic_alex_embed.py rename to src/ops_model/models/interpretability/toolkit/viewer/mimic_alex_embed.py diff --git a/src/ops_model/models/interpretability/_internal/viewer/morphometrics.py b/src/ops_model/models/interpretability/toolkit/viewer/morphometrics.py similarity index 100% rename from src/ops_model/models/interpretability/_internal/viewer/morphometrics.py rename to src/ops_model/models/interpretability/toolkit/viewer/morphometrics.py diff --git a/src/ops_model/models/interpretability/_internal/viewer/nway_clf.py b/src/ops_model/models/interpretability/toolkit/viewer/nway_clf.py similarity index 100% rename from src/ops_model/models/interpretability/_internal/viewer/nway_clf.py rename to src/ops_model/models/interpretability/toolkit/viewer/nway_clf.py diff --git a/src/ops_model/models/interpretability/_internal/viewer/phenotype_cells.py b/src/ops_model/models/interpretability/toolkit/viewer/phenotype_cells.py similarity index 100% rename from src/ops_model/models/interpretability/_internal/viewer/phenotype_cells.py rename to src/ops_model/models/interpretability/toolkit/viewer/phenotype_cells.py diff --git a/src/ops_model/models/interpretability/_internal/viewer/render_montage_scales.py b/src/ops_model/models/interpretability/toolkit/viewer/render_montage_scales.py similarity index 99% rename from src/ops_model/models/interpretability/_internal/viewer/render_montage_scales.py rename to src/ops_model/models/interpretability/toolkit/viewer/render_montage_scales.py index 27d519e..8feed4f 100644 --- a/src/ops_model/models/interpretability/_internal/viewer/render_montage_scales.py +++ b/src/ops_model/models/interpretability/toolkit/viewer/render_montage_scales.py @@ -2,7 +2,7 @@ embedding legend (leiden_r4, big dots, NTC as a dark labelled circle). The montage image and its baked viewer-style gene names come straight from the built tiles — finer levels give crisper text. - python -m ops_model.models.interpretability._internal.viewer.render_montage_scales --alphas 1-5 --levels 3,4 + python -m ops_model.models.interpretability.toolkit.viewer.render_montage_scales --alphas 1-5 --levels 3,4 Each level of `_montage/phase_geneKO_phate_cell1_a_tiles/L/` is a level-of-detail montage (coarse levels show a decimated non-overlapping subset; finer levels fill in more cells at higher res). diff --git a/src/ops_model/models/interpretability/_internal/viewer/score_generated.py b/src/ops_model/models/interpretability/toolkit/viewer/score_generated.py similarity index 100% rename from src/ops_model/models/interpretability/_internal/viewer/score_generated.py rename to src/ops_model/models/interpretability/toolkit/viewer/score_generated.py diff --git a/src/ops_model/models/interpretability/_internal/viewer/set_classifier.py b/src/ops_model/models/interpretability/toolkit/viewer/set_classifier.py similarity index 100% rename from src/ops_model/models/interpretability/_internal/viewer/set_classifier.py rename to src/ops_model/models/interpretability/toolkit/viewer/set_classifier.py diff --git a/src/ops_model/models/interpretability/_internal/viewer/submit.py b/src/ops_model/models/interpretability/toolkit/viewer/submit.py similarity index 97% rename from src/ops_model/models/interpretability/_internal/viewer/submit.py rename to src/ops_model/models/interpretability/toolkit/viewer/submit.py index 18d3244..a9982d8 100644 --- a/src/ops_model/models/interpretability/_internal/viewer/submit.py +++ b/src/ops_model/models/interpretability/toolkit/viewer/submit.py @@ -1,10 +1,10 @@ """Build the DiffEx viewer cache — reproducible, version-controlled entrypoint (replaces the one-off scratchpad drivers). All target selection comes from `catalog.py`. - python -m ops_model.models.interpretability._internal.viewer.submit seed # per-marker NTC traversals - python -m ops_model.models.interpretability._internal.viewer.submit anchors --k 5 # A→B anchor pairs - python -m ops_model.models.interpretability._internal.viewer.submit manifest # rebuild manifest.json (local) - python -m ops_model.models.interpretability._internal.viewer.submit montage --cell 0 --alpha 2 # harvest cache -> UMAP montage zarr + python -m ops_model.models.interpretability.toolkit.viewer.submit seed # per-marker NTC traversals + python -m ops_model.models.interpretability.toolkit.viewer.submit anchors --k 5 # A→B anchor pairs + python -m ops_model.models.interpretability.toolkit.viewer.submit manifest # rebuild manifest.json (local) + python -m ops_model.models.interpretability.toolkit.viewer.submit montage --cell 0 --alpha 2 # harvest cache -> UMAP montage zarr """ from __future__ import annotations diff --git a/src/ops_model/models/interpretability/_internal/viewer/webapp/app.js b/src/ops_model/models/interpretability/toolkit/viewer/webapp/app.js similarity index 100% rename from src/ops_model/models/interpretability/_internal/viewer/webapp/app.js rename to src/ops_model/models/interpretability/toolkit/viewer/webapp/app.js diff --git a/src/ops_model/models/interpretability/_internal/viewer/webapp/biohub-mark.png b/src/ops_model/models/interpretability/toolkit/viewer/webapp/biohub-mark.png similarity index 100% rename from src/ops_model/models/interpretability/_internal/viewer/webapp/biohub-mark.png rename to src/ops_model/models/interpretability/toolkit/viewer/webapp/biohub-mark.png diff --git a/src/ops_model/models/interpretability/_internal/viewer/webapp/biohub-wordmark.png b/src/ops_model/models/interpretability/toolkit/viewer/webapp/biohub-wordmark.png similarity index 100% rename from src/ops_model/models/interpretability/_internal/viewer/webapp/biohub-wordmark.png rename to src/ops_model/models/interpretability/toolkit/viewer/webapp/biohub-wordmark.png diff --git a/src/ops_model/models/interpretability/_internal/viewer/webapp/build_gene_narratives.py b/src/ops_model/models/interpretability/toolkit/viewer/webapp/build_gene_narratives.py similarity index 100% rename from src/ops_model/models/interpretability/_internal/viewer/webapp/build_gene_narratives.py rename to src/ops_model/models/interpretability/toolkit/viewer/webapp/build_gene_narratives.py diff --git a/src/ops_model/models/interpretability/_internal/viewer/webapp/gif.js b/src/ops_model/models/interpretability/toolkit/viewer/webapp/gif.js similarity index 100% rename from src/ops_model/models/interpretability/_internal/viewer/webapp/gif.js rename to src/ops_model/models/interpretability/toolkit/viewer/webapp/gif.js diff --git a/src/ops_model/models/interpretability/_internal/viewer/webapp/gif.worker.js b/src/ops_model/models/interpretability/toolkit/viewer/webapp/gif.worker.js similarity index 100% rename from src/ops_model/models/interpretability/_internal/viewer/webapp/gif.worker.js rename to src/ops_model/models/interpretability/toolkit/viewer/webapp/gif.worker.js diff --git a/src/ops_model/models/interpretability/_internal/viewer/webapp/index.html b/src/ops_model/models/interpretability/toolkit/viewer/webapp/index.html similarity index 100% rename from src/ops_model/models/interpretability/_internal/viewer/webapp/index.html rename to src/ops_model/models/interpretability/toolkit/viewer/webapp/index.html diff --git a/src/ops_model/models/interpretability/_internal/viewer/webapp/methods.js b/src/ops_model/models/interpretability/toolkit/viewer/webapp/methods.js similarity index 100% rename from src/ops_model/models/interpretability/_internal/viewer/webapp/methods.js rename to src/ops_model/models/interpretability/toolkit/viewer/webapp/methods.js diff --git a/src/ops_model/models/interpretability/_internal/viewer/webapp/morpho_demo.html b/src/ops_model/models/interpretability/toolkit/viewer/webapp/morpho_demo.html similarity index 100% rename from src/ops_model/models/interpretability/_internal/viewer/webapp/morpho_demo.html rename to src/ops_model/models/interpretability/toolkit/viewer/webapp/morpho_demo.html diff --git a/src/ops_model/models/interpretability/_internal/viewer/webapp/openseadragon.min.js b/src/ops_model/models/interpretability/toolkit/viewer/webapp/openseadragon.min.js similarity index 100% rename from src/ops_model/models/interpretability/_internal/viewer/webapp/openseadragon.min.js rename to src/ops_model/models/interpretability/toolkit/viewer/webapp/openseadragon.min.js diff --git a/src/ops_model/models/interpretability/_internal/viewer/webapp/opsin-eyes.svg b/src/ops_model/models/interpretability/toolkit/viewer/webapp/opsin-eyes.svg similarity index 100% rename from src/ops_model/models/interpretability/_internal/viewer/webapp/opsin-eyes.svg rename to src/ops_model/models/interpretability/toolkit/viewer/webapp/opsin-eyes.svg diff --git a/src/ops_model/models/interpretability/_internal/viewer/webapp/style.css b/src/ops_model/models/interpretability/toolkit/viewer/webapp/style.css similarity index 100% rename from src/ops_model/models/interpretability/_internal/viewer/webapp/style.css rename to src/ops_model/models/interpretability/toolkit/viewer/webapp/style.css diff --git a/src/ops_model/models/interpretability/_internal/weighted_aggregation/_v4_attn_worker.py b/src/ops_model/models/interpretability/toolkit/weighted_aggregation/_v4_attn_worker.py similarity index 100% rename from src/ops_model/models/interpretability/_internal/weighted_aggregation/_v4_attn_worker.py rename to src/ops_model/models/interpretability/toolkit/weighted_aggregation/_v4_attn_worker.py diff --git a/src/ops_model/models/interpretability/_internal/weighted_aggregation/analyze_v3_acc_bins.py b/src/ops_model/models/interpretability/toolkit/weighted_aggregation/analyze_v3_acc_bins.py similarity index 100% rename from src/ops_model/models/interpretability/_internal/weighted_aggregation/analyze_v3_acc_bins.py rename to src/ops_model/models/interpretability/toolkit/weighted_aggregation/analyze_v3_acc_bins.py diff --git a/src/ops_model/models/interpretability/_internal/weighted_aggregation/plot_v4_attn_comparison.py b/src/ops_model/models/interpretability/toolkit/weighted_aggregation/plot_v4_attn_comparison.py similarity index 100% rename from src/ops_model/models/interpretability/_internal/weighted_aggregation/plot_v4_attn_comparison.py rename to src/ops_model/models/interpretability/toolkit/weighted_aggregation/plot_v4_attn_comparison.py diff --git a/src/ops_model/models/interpretability/_internal/weighted_aggregation/run_v3_pipeline_on_v4_attn_weighted.py b/src/ops_model/models/interpretability/toolkit/weighted_aggregation/run_v3_pipeline_on_v4_attn_weighted.py similarity index 100% rename from src/ops_model/models/interpretability/_internal/weighted_aggregation/run_v3_pipeline_on_v4_attn_weighted.py rename to src/ops_model/models/interpretability/toolkit/weighted_aggregation/run_v3_pipeline_on_v4_attn_weighted.py diff --git a/src/ops_model/models/interpretability/_internal/weighted_aggregation/run_v3_pipeline_on_v4_features.py b/src/ops_model/models/interpretability/toolkit/weighted_aggregation/run_v3_pipeline_on_v4_features.py similarity index 100% rename from src/ops_model/models/interpretability/_internal/weighted_aggregation/run_v3_pipeline_on_v4_features.py rename to src/ops_model/models/interpretability/toolkit/weighted_aggregation/run_v3_pipeline_on_v4_features.py From 91570bc4e413472f8eebf40adf1176250e2727fd Mon Sep 17 00:00:00 2001 From: Gav Sturm Date: Wed, 12 Aug 2026 09:37:06 -0700 Subject: [PATCH 10/13] public-release: genericize minibinder (config-driven guide_col) + drop PLAN.md Minibinder-free public release: - data_loader: guide-col back-compat bridge -> self.guide_col (was minibinder_perturbation) - cp_extraction / anndata_validator: generic docstrings - diffae core: removed the 'minibinder' GRAIN + its catalog/precompute handling (_minibinder_meta.json / grain=='minibinder' branches) -> geneKO+complex only - 3 guide_col tests: fixtures use a generic custom_perturbation column - dropped diffae/PLAN.md (internal living design doc; also carried latent-lens prose) Entire public-release repo is now minibinder-free (git grep -il minibinder = 0). --- src/ops_model/data/data_loader.py | 20 +- .../models/cellprofiler/cp_extraction.py | 4 +- .../models/interpretability/diffae/PLAN.md | 712 ------------------ .../diffae/classifier/config.py | 1 - .../diffae/traversal/catalog.py | 9 +- .../diffae/traversal/precompute.py | 11 +- .../anndata_processing/anndata_validator.py | 8 +- .../features/test_anndata_utils_guide_col.py | 34 +- .../test_anndata_validator_guide_col.py | 30 +- tests/test_dataloader.py | 28 +- 10 files changed, 64 insertions(+), 793 deletions(-) delete mode 100644 src/ops_model/models/interpretability/diffae/PLAN.md diff --git a/src/ops_model/data/data_loader.py b/src/ops_model/data/data_loader.py index bcbd270..2320516 100644 --- a/src/ops_model/data/data_loader.py +++ b/src/ops_model/data/data_loader.py @@ -39,8 +39,8 @@ warnings.filterwarnings("ignore", category=zarr.errors.ZarrUserWarning) # Default name of the per-construct identifier column in adata.obs. -# Override per experiment via OpsDataManager(guide_col=...) (e.g. -# "minibinder_perturbation" for minibinder experiments). +# Override per experiment via OpsDataManager(guide_col=...) (e.g. a custom +# perturbation column for non-CRISPR libraries). DEFAULT_GUIDE_COL = "sgRNA" @@ -592,18 +592,18 @@ def get_labels(self): print(f"Reading link CSV from {csv_path}") labels_tmp = pd.read_csv(csv_path) - # Minibinder back-compat: link CSVs from minibinder experiments - # don't have a "gene_name" column. Copy minibinder_perturbation - # into gene_name so downstream gene_name-aware code (e.g. the - # balanced-sampling and gene-label LUT helpers) keeps working. - # Follow-up: minibinder gene-level should ultimately use - # gene_target, not the construct id — tracked separately. + # Custom-perturbation back-compat: some link CSVs use a + # non-standard guide column (self.guide_col) and have no + # "gene_name" column. Copy the guide column into gene_name so + # downstream gene_name-aware code (balanced sampling, gene-label + # LUT helpers) keeps working. if ( "gene_name" not in labels_tmp.columns and "Gene name" not in labels_tmp.columns - and "minibinder_perturbation" in labels_tmp.columns + and self.guide_col != "gene_name" + and self.guide_col in labels_tmp.columns ): - labels_tmp["gene_name"] = labels_tmp["minibinder_perturbation"] + labels_tmp["gene_name"] = labels_tmp[self.guide_col] if self.guide_col not in labels_tmp.columns: raise ValueError( diff --git a/src/ops_model/models/cellprofiler/cp_extraction.py b/src/ops_model/models/cellprofiler/cp_extraction.py index ff24c3c..5faad4a 100644 --- a/src/ops_model/models/cellprofiler/cp_extraction.py +++ b/src/ops_model/models/cellprofiler/cp_extraction.py @@ -243,8 +243,8 @@ def create_subset( bounds: [start, end] index range out_channels: List of channel names (default: ["Phase2D", "mCherry"]) guide_col: Name of the per-construct identifier column in the link CSV - (default: "sgRNA"; e.g. "minibinder_perturbation" for minibinder - experiments) + (default: "sgRNA"; e.g. a custom perturbation column for + non-CRISPR libraries) Returns: Tuple of (Subset dataset, label lookup table) diff --git a/src/ops_model/models/interpretability/diffae/PLAN.md b/src/ops_model/models/interpretability/diffae/PLAN.md deleted file mode 100644 index e5cafd5..0000000 --- a/src/ops_model/models/interpretability/diffae/PLAN.md +++ /dev/null @@ -1,712 +0,0 @@ -# DiffEx interpretability — plan - -Living design doc. Goal: interpret geneKO / protein-complex phenotypes **into image space** -using a DiffEx-style diffusion counterfactual (arXiv:2502.09663), since OP/CP classical -features are judged too weak to describe the phenotypes. - -## What is DiffEx? -DiffEx (*Explaining a Classifier with Diffusion Models to Identify Microscopic Cellular -Variations*, arXiv:2502.09663) explains any image classifier by generating visually interpretable -**counterfactuals** — showing, in pixel space, what about an image drives the classifier. -- **Architecture:** a **diffusion autoencoder (DiffAE)** — semantic encoder → low-dim latent - `z_sem`; conditional diffusion decoder reconstructs the image from `z_sem`. The classifier score - is concatenated onto `z_sem`; a bank of MLP **direction models** is trained in that latent with a - **contrastive loss** → distinct, disentangled directions. Explain class k = shift `z_sem` along a - direction and decode. -- **Interpretability features we exploit:** counterfactual morphs; a *global, reusable* attribute - vocabulary (directions shared across all classes — an image-grounded OP/CP replacement); - per-class attribute ranking; classifier-agnostic & forward-only (no retraining); continuous edit - strength α (dose-like morphs); quantitative faithfulness via re-encoding. - -Status (historical): *designing the per-cell classifier* — that phase is long done; see -ACTIVE EFFORTS below for the current state. - ---- - -## ACTIVE EFFORTS (dashboard — updated 2026-07-08) - -### LATEST (2026-07-08) — phenotype-cell handoff, v2 mAP, EBI matrix, viewer embedding tab -- **Phenotype-cell CSV for Ritvik** (`viewer/phenotype_cells.py`) → `viewer_assets/phenotype_cells_for_attention.csv`. - 20 cells × (geneKO + EBI complex) × marker, for SetTransformer attention pixel-patches on the REAL - phenotype cells. **160,420 cells / 53 markers** (phase + 52 fluor). Cols incl `map_score`, `geneKO`, - `ebi_complex`, `rank_source`, `segmentation_id` (=pma `segmentation`), `x/y_pheno`, `rank`, `pma_attention`. - - **Per-marker top-20, NOT the model's global top-20** — the pma `rank` is GLOBAL per geneKO (across all - 56 channels), so `_csv_top` re-ranks WITHIN each (channel, perturbation) and takes the 20 highest-attention - cells present in that channel (two-pass chunked `head` keeps memory bounded). - - **Fluor filtered by mAP ≥ 0.2** (phase = ALL perturbations): geneKO by distinctiveness, complex by EBI mAP. - - **`rank_source` col:** `"model"` (all current cells). Reserved `"fallback"` for markers not in the model. -- **v2 distinctiveness switch:** `catalog.dist_matrix` now reads `paper_v2/with_cp/with_4i/all_livecell` - (single 56-reporter matrix: 43 live + 7 CP + 6 4i), replacing the paper_v1 3-way split. Added 4i - `FIXED_REP` mappings (p53/pRb/pS6/p21/b-catenin/c-Myc). **52/56 pma channels now map**; the 4 - excluded (NFkB, RSP6, Rb, gH2AX) are EXPECTED — genuinely absent from the v2 matrix. -- **EBI complex mAP matrix** (`viewer/build_complex_ebi_map.py`) → `complex_reporter_ebi_map.csv` (98×56). - Runs copairs `phenotypic_consistency_ebi` per-marker on the v2 `with_cp/with_4i` per_signal gene - embeddings, over ALL perturbations (activity_map=None), NOT the wrong `complex_reporter_chad_consistency`. - `catalog.complex_dist()` reads it. Also wired into the aggregation pipeline - (`post_process/combination/pca_optimization/aggregation.py` → `complex_reporter_ebi_consistency.csv`). -- **3 new live-cell markers (cisGolgi, VIM, LMNB1):** HAVE v2 distinctiveness + EBI mAP, but are NOT in - Alex's pma cell CSVs yet (his attention output predates them) → no attention-ranked cells with crop - metadata. **FALLBACK (per user, TODO):** select cells around the CENTROID of the existing CellDINO - embeddings. Source found: `{exp}/3-assembly/cell_dino_features_v2/anndata_objects/features_processed_.h5ad` - (e.g. `mStayGold-CENPRaltORF`, `VIM`, `LMNB1`) — 1024-d embedding + crop metadata (`label_int`=segmentation, - `x/y_position`, `well`, `experiment`, `perturbation`). Per (marker, pert): centroid → 20 closest → - `rank_source="fallback"`, `pma_attention`/`rank` null. (Or wait for Alex's reprocessed pma CSV.) -- **SetTransformer accuracy scoring (`viewer/set_classifier.py` + `viewer/mimic_alex_embed.py`) — PARKED, - waiting on SetTransformer v2 (no-mask classifier).** Full journey + why: - - Reconstructed Alex's cellstate-set-classifier (ISAB/PMA/cosine head); 5 ckpts in `v4/wandb/cellstate_set_classifier/`. - Real-bag CEILING validated: feeding Alex's own `.pt` embeddings → P(target) 0.90–0.999 (HSPA5/KIF11/POLR1B/TIMM23 all hit). - - Raw `embed_crops` → classifier FAILS (OOD, cos 0.47, constant argmax). Built `mimic_alex_embed` to reproduce Alex's - exact pipeline: **128 Phase2D crop → seg-mask (`cell_seg`) → percentile-norm → CellDINO → z-std(control)**. Findings: - the **segmentation mask is the load-bearing step** (cos 0.47→0.91); percentile-norm is canceled by CellDINO's z-score. - At realistic bag sizes (100 cells, per-experiment z-std) the mimic matches the ceiling: **95–100% hit-rate** on real cells. - - Generated-cell per-α curve works end-to-end for **POLR1B** (P 0.01→0.98, flips to target at α≥1.5) — proof the pipeline - is correct — but **the segmentation of GENERATED cells is the blocker**: cellpose on fake crops unreliably captures the - cell (latches onto the bright nucleolus, not the whole body), so only nucleolar-phenotype genes (POLR1B) score; HSPA5/ - KIF11/TIMM23 stay at P≈0 despite healthy embeddings. Diameter/centroid tuning didn't fix it robustly. - - **DECISION:** masking generated cells is too fragile to rely on. **Wait for SetTransformer v2 — a classifier trained - WITHOUT masks** → then our unmasked `embed_crops` is in-distribution, no cellpose needed, and the per-α bag score works - for all genes. The mimic (mask path) + POLR1B validation are kept as a reference/cross-check. -- **Model-metrics curves** (`diffae/plot_metrics.py`): loss + cond_ratio over epochs, one line per DiffAE - → `model_metrics_curves.png/.svg`. -- **Viewer embedding tab** (`build_umap_montage.py` + `webapp/`): OSD montage, UMAP↔PHATE, points/images - toggle, 44 anndata color-by fields, opacity/zoom sliders, click→perturbation sidebar; gene descriptions - from the gene-embedding h5ad (`gene_desc.json`, fixes VAMP2-style blanks). - -### STATUS SUMMARY (historical detail condensed 2026-07-08) -- **Generators:** phase `phase_v1` = PRODUCTION (0.468); 500k warm retrain PARKED (peaked 0.542 then - declined). Fluor **50/50 markers trained** (ep≥98). v2/v3 aug did NOT beat v1. Directions default = - deterministic **mean_diff** α (see build log for the full DiffAE saga). -- **Viewer** (`viewer/` — `submit.py`, `catalog.py`, `precompute.py`, `build_umap_montage.py`, `webapp/`): - static precompute → dependency-free web app; per-marker driver shares the NTC gather + dedups real - cells; embedding tab (see LATEST). Live demo `login-01:8765`. -- **Score:** authoritative = Alex's **SetTransformer** bag `P(target)` (§7 ckpts downloaded) — supersedes - the per-cell N-way MLP (`nway_clf.py`) and the old binary LR badge. Generated-cell bag scoring PARKED: the - mask-mimic works on real cells (95–100%) + POLR1B generated, but segmenting fake cells is too fragile → - waiting for SetTransformer v2 (no-mask classifier). See LATEST. -- **Infra PR (#51): OPENED + ACCEPTED/MERGED** — `diffex-viewer-dev.tf` in `sfbiohub-infra` (S3 bucket - `diffex-viewer-dev` + nonprod read-only IRSA role `biohub-nonprod-diffex-viewer` for SA `diffex-viewer` in - ns `argus-diffex-viewer-rdev` + read-write uploader role; mirrors `proteohub-argus-s3-reader-dev.tf`, 1 TB - ceiling). `terraform apply` provisions the bucket/roles → then `aws s3 sync viewer_assets/ s3://diffex-viewer-dev/` - → Argus boot-download. Next: create the app repo + `argus register` (see App-staging build-log entry). - -### OPEN BUILDOUT -1. **Full NTC drain — LAUNCHED 2026-07-11** (master job `34826112`): all ~1000 geneKO genes/marker for the - 46 valid-`rep` fluor markers (was top-8 seed; 42 were partial ~100–194, 4 hub markers already ~complete). - Command = `submit seed --map-thr 0 --timeout 720` (**not** `--all-genes`; that flag never existed — the - `--map-thr 0` = every gene with distinctiveness ≥ 0 = all ~1000). Resume is automatic (skips built targets). - The 4 rep=None markers (NFkB/RSP6/Rb/gH2AX) are intentionally excluded. ~500 GB. See build-log 2026-07-11. -2. **Full A→B anchors** — `submit anchors --k 10` across all markers + complexes. -3. **Fluor complex traversals — DONE (resume 2026-07-11, master `34826156`)**: 98 EBI complexes × 50 markers - were already ~complete; only 5 markers partial (peroxisome_Peroxi, pS6, pRb, NPM3, SRRM2) → `submit - fluor-complex` resume finishes them. (phase complex = 190 = 98 NTC-anchored + 92 complex→complex anchor pairs.) -4. **Wire SetTransformer bag score** into `precompute` + a per-α curve panel — BLOCKED on **SetTransformer v2 - (no-mask classifier)**; the mask path is too fragile on generated cells (see LATEST). Once v2 lands, unmasked - `embed_crops` bags score directly (no cellpose). -5. **S3 hosting** — infra PR #51 MERGED. Argus app scaffold built (`/hpc/mydata/gav.sturm/diffex-viewer`, forked - from `czbiohub-sf/mops-viewer`). Remaining: `terraform apply` → create `czbiohub-sf/diffex-viewer` repo → - `argus register`/bootstrap (needs argus CLI) → upload `viewer_assets/` (or hand to Kyle) → PR+`stack` label. See build-log. -6. **Centroid fallback** for cisGolgi/VIM/LMNB1 (see LATEST). -7. **Multi-α montage** — `--alphas` flag looping per-α decodes + an α switch in the explorer. -8. **Image-UMAP montage** (`czi-ai/latent-lens`) — idea track; needs full-gene coverage. -9. **Attention-head tab + viewer reorientation** (DONE) — new tab overlaying CellDINO attention-head - pixel weights (inferno) on the real phenotype cells; Browse (marker+perturbation) drives ALL views. - See `### 2026-07-08 — Attention-head tab` build-log entry. -10. **Per-marker embedding montages** (FUTURE) — the Embedding tab now switches its IMAGES per marker but - reuses the SHARED phase gene-UMAP LAYOUT for every marker (`submit montage` loops modalities, all with - the phase `UMAP_H5AD`; only tiles swap). Fluor markers currently place only their ~8 seed geneKO frames - (sparse). FUTURE: give each marker its OWN gene embedding (per-reporter CellDINO gene UMAP/PHATE), so - genes sit at that marker's coordinates. Needs a per-marker gene-embedding h5ad (`obsm X_umap/X_phate` per - gene) — does NOT exist yet (`pca_optimized_v0.3/.../paper_v2/` only has aggregate combos: phase_only, - all_livecell, with_cp, no_phase, only_cp — no per-single-marker layout). Build = aggregate per-marker - CellDINO gene embeddings → PCA → UMAP/PHATE per reporter, then `build_montage_web(modality=, - h5ad=)`. Also depends on fuller fluor geneKO traversals (item 1) to be non-sparse. - ---- - -### Scope (locked) -- **4 classifiers** = 2 modalities × 2 grains: - 1. phase-only · geneKO (1001 KO + NTC-way) - 2. phase-only · complex (98-way, EBI) - 3. all-fluorescence · geneKO - 4. all-fluorescence · complex -- **2 diffusion models** (phase-only, all-fluorescence) — unconditional, shared across grains. -- **Negative contrast = `distinct`** (vs all other geneKOs / complexes), to isolate exactly - what is unique to each perturbation. Falls out of the multi-class softmax for free (§1.3). - ---- - -## 0. Why this shape (decisions already made) - -- **Drop OP/CP.** DiffEx discovers its explanatory vocabulary directly in image space, so we - don't need to translate CellDINO → OP/CP → language. The diffusion model *is* the vocabulary. -- **One unconditional diffusion model, not one per gene.** It only learns to generate realistic - single-cell crops. DiffEx's *classifier guidance* does all per-gene steering at inference. -- **Diffusion autoencoder (DiffAE), faithful to the DiffEx paper.** Semantic encoder → latent - `z_sem` + conditional diffusion decoder, trained jointly. NOT latent diffusion (no separate VAE) - and NOT a bare guided DDPM. The contrastive direction discovery (§3) requires this semantic - latent — it's the core of the method, not an add-on. -- **A small per-cell classifier, NOT the SetTransformer.** The SetTransformer never sees pixels - (input = bag of CellDINO embeddings), so it gives DiffEx no `image→logit` gradient. Rather than - bolt CellDINO on the front and wrestle the per-bag→per-image mismatch, we train a clean - per-image classifier. The SetTransformer still does what it's proven at — **selecting the cells** - the classifier trains on. (Full-stack `image→CellDINO→SetTransformer` faithfulness check is a - later reviewer-defense nicety, not on the critical path.) - -## Final product (per gene-KO / per complex atlas page) - -1. **Evidence cells** — top-X attention cells passing the mAP/accuracy threshold. X is variable - per gene; the value itself is the **penetrance readout** (HSPA5 ≈ 10, RPL10 ≈ 800). Printed. -2. **Counterfactual morph** — NTC → KO (and optionally KO → NTC, often cleaner). Averaged across - the top-X cells, not one cherry-picked cell. -3. **Per-channel difference heatmap** — always show **phase + the top highest-mAP channel(s) for - that class** (NOT every channel); localizes the phenotype to the relevant marker. -4. **Faithfulness number** — % of counterfactuals whose re-encoded logit actually flipped to KO. - ---- - -## 1. THE CLASSIFIER (current focus) - -**3 candidate classifiers (DiffEx target) — per Alex, 2026-06-16.** Cell selection is settled -(top-X attention cells, §1.7); the options differ in WHAT classifier scores them — the source of -DiffEx's differentiable `image → class-k logit`. All train/score on top-attention cells (which -classify accurately). - -- **A. SetTransformer native** — score a generated cell by inserting it into a real NTC reference - bag and reading the bag logit. Most faithful (actual deployed model, no new training). - **Risk (Alex):** the SetTransformer may not treat one synthetic cell as real, and/or +1 cell - won't move the bag logit → weak faithfulness signal. Cross-check, don't depend on it. -- **B. ResNet CNN on single cells**, trained only on top-attention cells. DiffEx-standard. - **Only option needing NO CellDINO encoder** for generated images (scores pixels directly). - Simplest, cleanest gradients; not SOTA representation. Detail in §1.1–1.6 below. -- **C. CellDINO features + MLP**, trained only on top-attention cells. More SOTA / accurate, light - to train, reuses the already-provided CellDINO features. Scoring *generated* images runs the - **local** CellDINO encoder (`ops_model/models/cell_dino.py` → `CellDinoModel`: channel-adaptive - DINO ViT-L/16, resize 224 + per-image z-score, `in_channels=1`); heavier per step than B, NOT a blocker. - -**Recommendation:** prototype **B and C on HSPA5 (phase)** — **fully unblocked locally** (we have the -phase attention rankings, the crops, and the CellDINO encoder). Pick by per-cell classification -accuracy. B = lightest (small CNN on pixels, no CellDINO in the loop); C = stronger classifier but -runs ViT-L per generated image. A = faithfulness cross-check only. - -DiffEx needs a differentiable `image → class-k logit`. Option B (CNN) design detail: - -### 1.1 Input / modality — **DECISION NEEDED** -- SetTransformer regime is **phase-only** (paper-v1). To stay faithful + keep the diffusion model - to 1 channel, **recommend starting phase-only**. Extend to phase+fluor later (fluor carries the - ER/mito visual signal the reader wants, but multiplies diffusion difficulty). -- Crop: reuse `ops_utils.data.bbox_utils.BaseDataset` — 128×128, multi-channel, cell mask - available. Same loader the atlas uses. **No new data infra.** -- Open: feed the cell mask as an extra channel (focus model on the cell, suppress neighbours)? - Likely yes — cheap confound reduction. - -### 1.2 Architecture — DECISION: backbone (start CNN, escalate if needed) -DiffEx is **classifier-agnostic** (it explains a plain supervised classifier; doesn't need -internal layers). So this slot is a free choice. What DiffEx needs is NOT top accuracy but: -clean image→logit gradient, confound-robustness (keys on biology, not plate/intensity), and -faithfulness to the real phenotype. - -- **v1 (recommended): small from-scratch CNN** (ResNet18-ish, N-channel stem). Clean gradient, - no giant frozen encoder in the path, easy to keep confound-robust, fail-fast for the PoC. - Honest caveat: a from-scratch ResNet is **standard, not SOTA** for cell-image representation. -- **Escalation path if CNN counterfactuals are too weak/insensitive:** fine-tuned **CellDINO** - (your near-SOTA self-supervised ViT) or a channel-aware ViT (ChannelViT / DINO4Cells family). - SOTA representation → more sensitive classifier → more sensitive counterfactuals, at the cost - of a heavier gradient path and higher confound risk (frozen SSL features can encode batch). -- Tradeoff is real because phenotypes are subtle. Decide empirically on HSPA5: if the CNN can't - separate HSPA5 cleanly cross-experiment, escalate the backbone before blaming the generator. - -### 1.3 Task framing — multi-class, distinct contrast for free -- **One N-way softmax classifier** per (modality × grain): 1001-KO+NTC-way for geneKO, - 98-way for complex. Shared backbone, per-class linear heads. -- DiffEx guides toward `logit_k`. Because softmax is normalized against all other classes, - **guiding toward class k IS the `distinct` contrast** (what's unique to k vs everything else) — - no separate binary models needed. -- PoC = train the geneKO N-way model, then run DiffEx toward the HSPA5 logit. -- Caveat (logged, not a blocker): distinct **suppresses shared phenotypes** by construction - (two ER-stress genes won't show their shared ER signature, only their difference). Keep - **vs-NTC as a complementary second pass** for genes where the absolute phenotype is wanted. - -### 1.4 Training set (a query, not new infra) -- Positives for class k: top-X attention cells for k, ranked by `attn_geneko` / `pma_attention`, - passing the mAP/accuracy threshold. X variable per gene (→ penetrance). -- Negatives = the **top-attention cells of the OTHER classes** (like-with-like), NOT random/weak - cells of them — else the model relearns "has *any* phenotype" instead of "has *this* one". - With a softmax over top-attention cells per class this is automatic. -- Sources: - - attention sidecar: `/hpc/projects/icd.fast.ops/models/alex_lin_attention/v4/expansion_v1/per_experiment_v4_attn.parquet` - (cols: experiment, well, segmentation_id, attn_ebi, attn_geneko, attn_chad, ...) - - PMA parquets: `.../v4/pma_*` (cols: gene, experiment, well, segmentation, pma_attention) - - cell-set builder already exists: `organelle_profiler.feature_extraction.consolidate_top_attention_cells` -- Join `(experiment, well, segmentation_id)` → bbox → `BaseDataset` crops. - -### 1.5 Confound guardrail (critical — or DiffEx faithfully explains the batch effect) -- Balance / stratify NTC negatives across the same experiments+wells as the positives. -- **Validate cross-experiment** (train on subset of experiments, test on held-out) — biological - signal generalizes, plate/intensity artifacts don't. -- Sanity bar: per-cell classifier AUROC should track the SetTransformer's per-gene mAP ordering - (sharp genes like HSPA5 easy, diffuse genes like RPL10 hard). - -### 1.6 Success criteria -- HSPA5-vs-NTC held-out AUROC clearly > 0.5 and > a same-data **all-cells** classifier - (proves attention selection removes confounds). -- Cross-experiment generalization holds. - ---- - -## 2. Diffusion autoencoder (DiffAE) — the real cost/risk -Faithful to DiffEx. Two jointly-trained parts on single-cell crops (one DiffAE per modality: -phase-only, all-fluor): -- **Semantic encoder** → `z_sem` (a low-dim semantic latent capturing cell appearance). -- **Conditional diffusion decoder** (`UNet2D`, N input channels) that reconstructs the crop - conditioned on `z_sem` (+ stochastic DDIM latent for detail). -- Data is NOT the constraint (millions of cells); compute/engineering is. -- De-risk on HSPA5 PoC before scaling. First sanity check: round-trip reconstruction quality - (encode→decode) on held-out cells — if it can't reconstruct, directions are meaningless. - -## 3. Contrastive direction discovery (the interpretability core) -Faithful to DiffEx — this is where the explanation comes from, NOT per-image classifier guidance: -- Concatenate the §1 classifier score onto `z_sem` → semantic code. -- Train a bank of MLP direction models `D_1…D_N` that each shift the code by `α·Δz_k`, with a - **contrastive loss**: edits from the same direction stay similar, edits across directions stay - dissimilar → distinct, disentangled, reusable attributes. -- **Distinct contrast (per §1.3):** select the direction(s) that move the classifier toward - class k's logit → "what is distinct about geneKO/complex k". The discovered `D_1…D_N` form a - **shared attribute vocabulary across all 1001 genes / 98 complexes** — the real payoff vs a - one-off morph, and an image-grounded replacement for OP/CP. -- Faithfulness check, baked in day one: re-encode the edited image through §1 classifier, confirm - the logit actually moved toward k; report % success on each atlas page. - ---- - -## Milestones / de-risk order -1. **Classifier on HSPA5** (this task) — fail-fast signal that the cell sets are learnable. -2. **DiffAE** on the same crop set; gate on round-trip reconstruction quality (§2). -3. **Contrastive direction discovery** (§3); pick HSPA5's distinct direction; eyeball morph + - per-channel heatmap + flip-rate. -4. If HSPA5 works → scale to atlas (reuse shared directions). If not → it won't work anywhere; stop. - -## Open questions -- Phase-only vs phase+fluor for v1? (recommend phase-only) -- Mask as extra input channel? (recommend yes) -- Threshold definition for "X cells that pass": fixed mAP cutoff vs per-gene accuracy knee? -- Negative contrast for v1: NTC only, or also distinct/global? - ---- - -## §1.7 Cell selection (settled, all options) + Option-A bag scoring -**Settled — cell selection (all 3 classifier options):** cells fed to DiffEx / used to train the -classifier = the **top-X attention cells** (PMA attention rank), exactly as used for the atlas. - -**Option A only — scoring a *generated* image with the bag model:** insert the generated cell into -a fixed real NTC reference bag, read logit_k. Alex's concern: the SetTransformer may not treat the -synthetic cell as real and +1 cell may not move the bag logit. → why A is a cross-check, and why -B/C (self-contained single-cell classifiers) are the primary path. - -## Assets inventory (checked 2026-06-16) -**Have locally:** -- **CellDINO encoder IS local**: `ops_model/models/cell_dino.py` → `CellDinoModel` (channel-adaptive - DINO ViT-L/16, ckpt `channel_adaptive_dino_vitl16_pretrain_cells-…pth`, resize 224 + per-image - z-score, `in_channels=1`). Plus the precomputed CellDINO feature dumps (below). → image→embedding - for *generated* cells is available locally; no encoder request needed. -- **MixedChannelClassifier** code (`train_set_classifier.py`) + inference (`export_pma_attention.py`). -- Per-gene **embedding dumps WITH cell_metadata**: `v4/{train,val}_ops_zstdcontrol_cdino_v2/` - (+ `metadata.pt` w/ gene_to_idx, channel_to_idx). Metadata → zarr crop mapping works - (`_load_cell_crop`: experiment/well/x_pheno/y_pheno/segmentation_id/zarr_channel_index). -- Attention outputs (`per_experiment_v4_attn.parquet`, `pma_phase_cells_*`). -- katamari clone on branch `main` (esmc_paper commit) — likely NOT Alex's attn-classifier branch. - -## Requests for Alex L. (to scale beyond the phase·geneKO PoC) -**Note:** the B/C HSPA5 phase prototype needs NOTHING from Alex — phase attention rankings, crops, -and the CellDINO encoder are all local. -1. **Trained MixedChannelClassifier checkpoints** for the 4 models (phase·geneKO, phase·complex/EBI, - fluor·geneKO, fluor·complex) — the actual `.pt` files or the **wandb artifact IDs** - (`alex-lin/cellstate-set-classifier/model-…`). Needed for option A, and to generate attention - rankings for the models we don't already have rankings for. -2. **Fluorescent + complex (EBI/CHAD) embedding dumps + label maps** — only phase·geneKO dumps - (`*_cdino_v2`) confirmed local. Need the fluor dumps and the complex `label_map_path`. -3. **If pursuing option A:** how to score a single generated cell inside a real reference bag (his - concern: one synthetic cell may not move the bag logit). -4. **The right git branch** to check out for the latest attn-classifier code (clone is on `main`). - ---- - -## Build log - -### 2026-06-16 — classifier B/C package built -Package: `ops_model/models/interpretability/diffae/classifier/` (config, data, models, -celldino_features, train, run, submit, README). Locked params: binary HSPA5-vs-rest, -negatives = other genes' top-5 (distinct), 1000/class, 160×160 phase crops (no mask), -3-way train/val/test split grouped by experiment (val=selection, test=clean reported AUROC; -train+val AUROC logged per epoch to watch over/under-fit). Outputs under -`/hpc/projects/icd.fast.ops/models/diffex//`. -- **Decision:** option C reuses the **local** CellDINO encoder (`cell_dino.py`) on the SAME - crops B uses (cached) — no dump-join; B and C see identical cells. -- **Verified:** full B pipeline end-to-end on CPU (tiny config) — filtered parquet read (no OOM), - store resolution, crops materialized non-degenerate, train→AUROC→artifacts. SLURM submitter - dry-run OK (2 GPU jobs). (Crop cache key includes mask state so masked/unmasked don't collide.) -- **Next (GPU):** `python -m ops_model.models.interpretability.diffae.classifier.submit --gene HSPA5` - → compare B vs C held-out AUROC, pick the DiffEx target. C needs GPU (CellDINO ViT-L). - -### 2026-06-16 — HSPA5 PoC results (job 34280479, experiment-grouped split, 1000/class) -| model | test AUROC | val | train@best | -|---|---|---|---| -| B (ResNet on crops) | 0.80 | 0.85 | 1.00 (overfits) | -| **C (CellDINO+MLP)** | **0.96** | 0.95 | 0.999 | -- PoC validated: HSPA5 top-attention cells are cleanly + cross-experiment classifiable → real DiffEx target. -- **C is the chosen DiffEx target** (far more sensitive/generalizing; B memorizes). C scores generated - counterfactuals via the local `cell_dino.py` encoder. -- Outputs: `/hpc/projects/icd.fast.ops/models/diffex/HSPA5/`. -- **Next:** (a) sanity-check a diffuse gene (e.g. RPL10, top-800 penetrance) to confirm the approach - holds across penetrance; (b) begin the DiffAE stage (§2) with C as the classifier. - -### 2026-06-16 — 98 EBI complexes + NTC sweep (model C, jobs 34285052 + 34285381) -All 99 bins ran. Per-class test AUROC (distinct vs pooled top cells of other perturbations), -experiment-grouped split. Range 0.748–0.958, median 0.860. -- Most distinct: Chaperonin-containing T-complex 0.96, DNA pol α:primase 0.95, eIF4F 0.95, - replication fork protection 0.94, COPI 0.93. -- Least distinct: RNA Pol II 0.75, U1 snRNP 0.75, SRP 0.76, NSL HAT 0.78. -- **NTC = 0.86 (mid-pack) is EXPECTED, not a confound** — NTC lacks the CRISPR cut, so it's a real - distinct (no-DSB) state. (Earlier "red flag" retracted.) -- Artifacts: `…/diffex/complex/auroc_ranking_C.csv`, `auroc_hist_C.png`. -- Takeaway: every complex's top-attention cells carry a separable phenotype → strong DiffEx targets - across the board; clear biologically-sensible ranking. - -### 2026-06-16 — DiffAE stage (§2) built + PoC launched (job 34292127) -Package: `diffex/diffae/` (config, data, model, train, recon, run, submit). Faithful DiffAE: -ResNet18 semantic encoder → `z_sem` (512), conditional `diffusers.UNet2DModel` decoder with -`z_sem` injected via `class_embed_type="identity"` (→ time embedding), trained jointly with DDPM -denoising loss. Gate = DDIM-invert→reverse reconstruction (PSNR + montage; uses DDIMInverseScheduler). -- Decisions (locked): broad training set (all classes incl NTC, all ranks, ~50k crops), 160×160 - phase, per-image z-score/3 normalization, PoC-first. -- Reuses the classifier crop pipeline (`materialize_crops`). Verified end-to-end on CPU (synthetic): - conditioning, DDPM step, DDIM recon, checkpoints all work. -- PoC launched: 50k crops, 80 epochs, batch 32, 1 GPU, 720min. Outputs → `…/diffex/diffae/phase_v1/`. -- **Gate to watch:** reconstruction PSNR (recon montages every 10 epochs). If it reconstructs cells - faithfully → proceed to §3 (contrastive direction discovery). If not → fix before directions. -- **Next stage (§3):** contrastive direction discovery on `z_sem` + the option-C classifier score. - -### 2026-06-17 — DiffAE v1 result + ARCHITECTURE SWITCH to Alex's design -- v1 run (job 34298429, jointly-trained encoder): trained healthily, **reconstruction PSNR ~33–34 dB, - visually faithful** (gate passed), converged ~ep9, but hit the 12h wall at ep37/80 (too slow/many - epochs). First quota failure (34292127) fixed by freeing disk + wrapping recon writes in try/except. -- **Alex (EvolutionaryScale) already has a working DiffEx** (Notion: Imaging AI). Key design: condition - the image-DDPM on the **FROZEN backbone embedding** (not a learned encoder); discover K direction - MLPs **unsupervised** (InfoNCE + VICReg); rank directions **post-hoc** with a logistic-regression - classifier (control vs KD); traverse α∈[−3,+3], DDIM-sample per step; verify by monotonic score. -- **SWITCHED to Alex's design** (job 34312003): DiffAE now conditions on the **frozen CellDINO - embedding** (dropped the learned encoder; `cond_proj` injects it into the UNet time-embedding). - Generator, option-C classifier, and SetTransformer now all live in the SAME CellDINO space → - Stage-3 directions are discovered there and ranked by option-C directly. epochs 20, reuses crop - cache + new CellDINO-embedding cache. -- **§3 plan (Alex's recipe):** K direction MLPs (InfoNCE+VICReg, unsupervised) on CellDINO embeddings - → rank by option-C classifier score shift → α-traversal → DDIM-sample images → verify monotonicity. - -### 2026-06-18 — Stage 3 built + first HSPA5 traversal (job 34385092) -Package `diffex/directions/` (config, model=DirectionBank, losses=InfoNCE+decorrelation, -train_directions, rank=LR score-shift, data=gather target+NTC, traverse=DDIM invert→reverse + -re-encode verify, run, submit). Verified 2a+2b on synthetic; full pipeline ran on GPU in 5 min. -- HSPA5: LR acc 0.999, selected direction #6 (shift 1.31), **6/8 traversals monotonic**, mean - score Δ +0.71 (correct sign). **Full DiffEx machine works end-to-end.** -- **BUT effect is weak**: re-encoded scores stay on NTC side (−6..−14), visual change subtle. - Cause: unit direction × α≤3 ≪ the real control→KD gap (clusters far apart, LR acc 0.999); plus - x_T anchoring from DDIM inversion. -- **Fix (next): scale α to ‖mean(KD)−mean(NTC)‖** (likely ~10–30, not 3); optionally reduce x_T - anchoring; train DiffAE longer. Outputs: `…/diffex/directions/geneKO/HSPA5/`. - -### 2026-06-18 — Stage 3 retry: gap-scaled α + Δ-pixel heatmap overlay (job 34386543) -- Added: heatmap overlay (Δpixels vs α=0, red/blue) on the traversal montage; α now in units of - the control→KD gap (Δ=9.64). 6/8→1/8 monotonic, score Δ 0.71→0.28 — **gap-scaled α OVERSHOT**. -- **Diagnosis from montage:** at large α the edits land on crop BORDERS/background, not the cell - → embedding goes off-manifold, DiffAE renders boundary artifacts. Two root issues: - (a) crops are **unmasked** → direction may exploit context (confluency/neighbors), not the cell; - (b) DiffAE **edit-fidelity** limited (reconstructs well but doesn't render edits onto the cell). -- **Next options:** α-magnitude sweep for the on-manifold regime (~0.25–0.5×gap); try masked crops; - train DiffAE longer / stronger conditioning; reduce x_T anchoring. Pipeline is correct; counterfactual - QUALITY needs iteration (the hard part of DiffEx). - -### 2026-06-18 — BUG FOUND & FIXED: generate-from-noise (job 34388244) -- **Bug:** traversal DDIM-INVERTED the real cell → x_T, which encodes the image and overrides the - embedding → editing the embedding barely moved the picture (and recon was a too-good 34 dB). Alex's - spec: α=0 is the DDPM *reconstruction*, i.e. **generate from noise conditioned on the embedding** (no - inversion). Switched to fixed random noise per cell (constant across α), conditioned on z0+α·d. -- **Result:** mean re-encoded score Δ **0.28 → 5.3**, 6/8 monotonic, correct sign — embedding now drives - generation, direction validated. BUT the visual morph is still **subtle** (CellDINO registers texture - the eye misses; HSPA5 phase phenotype may be genuinely subtle). Δ-pixel heatmap localizes the - changing cell regions = a useful interpretability output on its own. -- **Next levers:** push α to 2–3×gap; strengthen DiffAE (classifier-free guidance / longer); test a - gross-morphology target (complex) to tell if subtlety is biology vs method. - -### 2026-06-18 — obvious targets (TOMM20, TIMM23, Arp2/3): diagnosis = METHOD-limited -- TOMM20 (job 34392814): lr 0.96, gap 6.2, score Δ −3.1, 7/8 monotonic. TIMM23 (34392821): lr 0.99, - Δ −3.1, 7/8. (sign arbitrary per Alex.) Arp2/3 complex: filename bug (target "2/3" has a slash → - unslugified PNG name) — FIXED (slugify filenames in traverse._plot); re-run 34392960. -- **Key finding:** TOMM20 (obvious mito phenotype) morphs just as SUBTLY as HSPA5, edge-concentrated - Δ. → the visual subtlety is **METHOD-limited, not biology**. -- **Root cause:** DiffAE under-utilizes the embedding — the noise latent dominates DDIM generation - (inverted OR random), so embedding edits weakly change pixels even though CellDINO/classifier - register them (score moves, monotonic). -- **Fix: classifier-free / edit guidance** at sampling: ε̃ = ε(c0) + w·(ε(c_α) − ε(c0)), w≈3–5 (no - retrain). If insufficient → retrain DiffAE with conditioning dropout for proper CFG. - -### 2026-06-18 — edit-guidance w-sweep (TOMM20, job 34393392): INSUFFICIENT → must retrain DiffAE -- w=1/3/5 score Δ = 0.55/1.83/2.23 (guidance amplifies) BUT monotonic 0.38/0.25/0.12 (degrades), - and the **cell still does not transform by eye even at w=5** (edge-concentrated Δ only). -- **Conclusion:** sampling-time guidance cannot fix an under-conditioned model. The DiffAE generates - from the NOISE latent and only weakly uses the embedding (why recon hit a too-good 34 dB). -- **SOLID FIX (next): retrain DiffAE with conditioning dropout** (~10–20%, learned null embedding) → - forces embedding use + enables true CFG ε̃=ε(∅)+w(ε(c)−ε(∅)). If still weak → cross-attention - conditioning (spatial) instead of global FiLM. -- Also fix: (a) direction discovery is run-to-run unstable (fix seed + more epochs); (b) replace the - recon gate with an UNCONDITIONAL-generation check (null-embedding samples should look generic; - conditional should match target) — recon PSNR was misleading. - -### 2026-06-18 — conditioning diagnostic = DEAD (0.008), then proper DiffAE rebuild (job 34394595) -- **Diagnostic** (`diffae/diagnose_conditioning.py`, job 34393969): same fixed noise under - null/ctrl-centroid/KD-centroid. MSE(ctrl-vs-kd)=0.0010, MSE(noise-vs-noise)=0.133 → - **emb/noise ratio = 0.008**. The embedding has <1% control; the DiffAE generates from noise and - ignores the embedding. (null-vs-ctrl 0.042 ≫ ctrl-vs-kd 0.001 → reacts to embedding *presence*, - not *content*.) Confirms: cells don't change because conditioning is ~dead. -- **Rebuild** (`diffae/train.py` rewritten): conditioning dropout (0.15, learned null_emb) + EMA - (0.9995) + resume-across-jobs + deeper cond MLP. **Gate = conditioning ratio** (not recon PSNR), - logged every 5 epochs; EMA-best saved by ratio. Target: ratio ≫ 0.008 (→ ~0.3+). -- Retrain 34394595 (phase_v1, reuses caches, 120 epochs, batch 48, resume). Watching cond_ratio - trajectory. **After it works:** switch directions/traverse to true CFG ε(∅)+w(ε(c)−ε(∅)). -- If time-embedding conditioning still can't climb → escalate to cross-attention (UNet2DConditionModel). - -### 2026-06-29 — Plan C implemented (deterministic direction) + v2_aug retrain -- **Reproducibility root cause:** unsupervised InfoNCE direction bank is GPU/seed-nondeterministic - and `best_k = argmax|shift|` flips run-to-run → same cell highlighted different regions. -- **Fix (plan C, implemented):** `directions/config.py` `direction_method` — default **`mean_diff`** - (deterministic control→KD centroid vector; also `lr_weight`) as PRIMARY; the paper's unsupervised - bank kept as `direction_method="unsupervised"` secondary track. `traverse(fixed_dir=…)` uses the - global deterministic direction; `deterministic=True` sets seeds + cuDNN-deterministic. `+α = toward - KO` by construction (no more sign flip). `rank.supervised_direction()` computes it; LR kept for the - re-encode score check only. -- Reproducibility proof in progress: TIMM23 run twice (jobs 34654037/34654040) → pixel-diff strips. -- **Next model (v2_aug)**: orientation-aug DiffAE retraining (job 34651296), cond_ratio climbing - 0.04→0.14 @ep19/120 (aug ramps slower); resume to convergence, then switch directions default to it. - -### Active experiments (2026-07-03) -- **v1 remains the best model.** v2 (dihedral) and v3 (continuous rot+scale) augmentation did NOT - beat v1 — not more orientation-stable, weaker/less-convincing phenotypes; cond_ratio ceiling falls - with aug (v1 0.47 → v2 0.25 → v3 0.20; curves in `coding_exps/diffex/diffae_training_curves.png`). - Flow-matching transport (CellFlow-style, `directions/flow.py`) also explored → smoother but less - clean phenotypes, noisy negative extreme → NOT adopted. Reverted default to v1 + mean-diff α. -- **Generator data-scaling test (RUNNING):** does 50k→**500k** crops help? Two no-aug chains, 24 ep: - - scratch `phase_v1_500k` — jobs `34667092→34667174→34667175` - - warm-start from v1 `phase_v1_500k_warm` — jobs `34667176→34667177→34667178` - - Compare cond_ratio/loss vs v1 (0.47) + visual morphs. mem_gb=400 (500k float32 crops ≈ 51GB each). -- **Direction depth test (pending):** gather 1k→**~12k**/class (the distinctiveness peak) for a tighter - mean-diff centroid — cheap, per-target (~30-40 min CellDINO/target), no retraining. - -### Future direction — per-fluorescent-marker models (2026-07-03) -Reproduce the best phase pipeline **per fluorescent marker** (~60 live-cell markers) — a per-marker -counterfactual view of each gene-KO / complex phenotype in the channel where attention is most -informative. **~60 models** (one DiffAE + direction set per marker). -- **Attention source EXISTS:** `…/alex_lin_attention/v4/pma_fluorescent_cells_all.csv` - (+ `pma_fluorescent_cells_ebi_all.csv` for complexes) — the fluor analog of `pma_phase_cells_v2_all.parquet`. - CellDINO fluor train/val sets also present (`train/val_ops_zstdcontrol_cdino_fluorescent`). -- **Scope:** `good_experiment_list_v2.yml` (87 exps; fluor channels GFP×74, mCherry×23, Cy5×2). Each - experiment's channel→biological-marker label is in `ops_process/ops_analysis/configs/ops_channel_maps.yaml`. - Marker/experiment enumeration tooling: `ops_utils/data/feature_discovery.py`, - `ops_utils/analysis/embedding_discovery.py`. -- **Per-marker pieces:** gather top-attention fluor cells (control + KD) → CellDINO embed the MARKER - channel → mean-diff direction → **per-marker DiffAE** (phase generator can't decode fluor) → traverse. - Direction/traverse code unchanged; needs a per-marker DiffAE + a marker→(experiments, channel) map. -- **Complication (deferred):** the v2 list is **live-cell fluor only** — 4i / Cell-Painting (fixed-cell) - channels are excluded. If added later they need their **own per-round link CSVs** (`link_csv_dir` in - `ops_model/data/data_loader.py`), not the default live 3-assembly link. -- **NEEDS DESIGN CONFIRM before building** (60 DiffAE trainings is a large program). - -### Future direction — attention-informed cell selection (2026-06-29) -Currently we take a flat top-1000 attention-ranked cells per class and pick traversal/feature -cells by index. To exploit attention ranking more (only touches `directions/data.py gather` + -`rank.supervised_direction`, not the DiffAE): -- **Pick highest-attention cells for the traversal/featured strips** (most representative morph), - not an arbitrary cell index. -- **Per-target penetrance depth** — use the attention-accuracy knee (HSPA5≈top10, RPL10≈top800) - instead of a flat top-1000, so sharp phenotypes aren't diluted by the diffuse tail. -- **Attention-weight** the mean_diff / classifier so the most-phenotypic cells dominate the axis. - -### Future direction — orientation-invariance via augmentation (2026-06-29) -Observation: along a traversal the cell often spuriously rotates/transposes (orientation is -encoded in the CellDINO embedding, so the discovered direction carries an orientation component the -DiffAE renders). Fix: during DiffAE training, augment the TARGET image with the dihedral group -(4 rotations × flip = 8 views, incl. transpose) while conditioning on the embedding of the CANONICAL -(un-augmented) cell. Teaches the model orientation ≠ embedding-determined → orientation absorbed by -the (fixed) noise latent, phenotype carried by the embedding → traversals stop rotating. Do NOT -recompute CellDINO on the augmented crop (defeats the decoupling). Also serves as general aug to -sharpen conditioning. - -### 2026-07-08 — Attention-head tab + viewer reorientation (BUILT) -**Status: built + deployed to `viewer_assets/`.** `viewer/build_attention_heads.py` rendered **984/1000** -phase geneKO genes (16 npz still corrupt/mid-write by Kevin — `BTF3L4, BUB1B, DAD1, DHRS9, FECH, -FOXD4L1, GTPBP4, INO80D, MTOR, NCBP2, NRAS, POLR2F, RPS19BP1, TWF1, TYK2, YIPF5` — builder is idempotent, -skips-loud, re-run picks them up). `global_max=2.44`. Webapp reoriented (`webapp/{index.html,app.js, -style.css}` v36, copied to `viewer_assets/`): persistent `#browse` block (marker+grain+perturbation+ -cells/page) drives 3 view tabs — Traversal / Embedding / **Attention heads**. Attn view = inferno LUT + -live per-map/per-gene/fixed normalization + opacity, head dropdown from `heads.json`; greys out for -non-phase / complex / missing-gene. Embedding now rings + pans to the selection. **Availability decoupled -from manifest** — app fetches `attention_heads/phase/index.json` (no `precompute.build_manifest` change). -- **ALL 4 modality×grain combos now IN (2026-07-08 pm).** Kevin dumped pixel_attribution (maps+crops+ - patch_masks) for fluor geneKO (`fluorescence_pixel_attribution//`), phase complex + fluor complex - (`complex_pixel_attribution/{phase,fluorescence/}/`) — same npz schema as phase. Builder rewritten - **SLURM-parallel** (`build_attention_heads.py render` → `submit_parallel_jobs`, 40 shards, ~1 min) into a - uniform `attention_heads////` layout + single `attention_heads/index.json` - ({global_max, assets:{modality:{grain:[keys]}}}). **23 modalities, 1455 keys** (phase geneKO 984 + phase - complex 93 + 16 fluor markers). Webapp resolves assets by (marker→`jsSlug(marker_channel)`|"phase", grain, - target→gene|slug); non-phase-geneKO lack ranking metrics (auroc/spec) but render fine from the npz `heads`. - (`fluorescence_attention/.npz` = the older ranking-features-only dump, superseded.) -- **16 corrupt phase genes: left as-is.** Full-size but bad-zip at SOURCE — needs Kevin to regenerate; 984/1000. -- **Attn viewer UX (2026-07-08 pm):** top-crossbar selection (marker+perturbation comboboxes w/ search, mAP|A–Z - sort) drives Traversal/Embedding/Attention-heads tabs; attn view = per-perturbation color-coded blocks - (rows=heads, cols=cells), pin/reset controls, per-cell/gene/fixed norm + dual clim + opacity sliders, - Ritvik-faithful overlay (σ=2 smooth + cell mask + inferno@α0.6); embedding click → selects in search box. - -### 2026-07-08 — Attention-head tab + viewer reorientation (design) -New viewer view: overlay CellDINO **attention-head pixel weights** (inferno) on the **real phenotype -cells** (`viewer/phenotype_cells.py` output), so you can see WHERE in each cell each top attention head -looks — the classifier's spatial evidence, alongside the generative counterfactual morph. - -**Data (Kevin L., already under `viewer_assets/attention_heads/`, verified 2026-07-08):** -- `phase/pixel_attribution_cache/.npz` (1000 geneKO genes): `maps (20,6,128,128) f16` ∈ [0,~0.34] - (20 cells × top-6 heads × pixel attribution), `crops (20,128,128) f32` (z-scored cell crops), - `heads (6,2) int32` = the ranked (layer,head) pairs, `patch_masks (20,196) bool` (14×14 ViT patches). -- `phase/head_rankings_per_gene.json`: per-gene ranked heads + metrics (`layer,head,feature,spec_p10, - spec_min,auroc_vs_ntc`); order matches the npz `heads` array. **The 20 cells ARE the phase·geneKO - top-20 phenotype cells** from `phenotype_cells.py` (same selection). -- `celldino_attention_head_analysis/fluorescence_attention/.npz` — per-marker fluor analog - (structure TBD) → follow-on. **v1 scope = phase · geneKO only** (no complexes: head_rankings is gene-keyed). - -**Precompute (`viewer/build_attention_heads.py`, to write):** ship **raw** data so normalization + inferno -+ opacity are LIVE display options (user-selectable, per decision). Per gene → per cell: write -`cell/crop.webp` (grayscale, per-crop robust min-max) + per head `cell/head.webp` (grayscale -attribution scaled by a FIXED global max so absolute intensity is preserved). Per-gene `heads.json` = -ranked-head metrics (`layer,head,feature,spec_p10,spec_min,auroc_vs_ntc`) + `n_cells` + `gene_max` + -`global_max`. The webapp applies a 256-entry **inferno LUT** in a canvas and composites over the crop — -so a Display dropdown offers **per-map / per-gene / fixed** normalization live (per-map = rescale by the -loaded tile's own max; per-gene = by `gene_max`; fixed = by `global_max`), plus an opacity slider, without -re-fetching. Count ≈ 1000×20×(6+1) ≈ 140k grayscale WebP (traversal already ~370k). Keeps the app -dependency-free; no 3×-image blowup from baking each norm. - -**Viewer reorientation (`webapp/index.html` + `app.js`):** today the left panel is 3 *control* tabs -(Browse / Anchor / Embedding), each with its own selectors, and the Embedding montage ignores the -Browse (marker,perturbation) selection. **Reorient:** make Browse (marker + grain + perturbation) a -PERSISTENT selector block = single source of truth (`state.marker`, `state.target`); below it a **View -switcher** — Traversal | Embedding | Attention heads — each rendering the main stage for the CURRENT -selection. Fold today's Anchor/display controls under Traversal; α/cell/embedding-mode under Embedding; -head selector + overlay-opacity under Attention heads. Embedding also gains browse→highlight/pan and -keeps montage-click→browse select (closes the selection loop). - -**Attention-head view UI:** grid of the 20 phenotype crops with the selected head's inferno overlay; -head dropdown lists the 6 ranked heads with `(L,H) · AUROC·NTC / spec`; overlay-opacity slider; raw-crop -toggle. Reuses Browse's cells-per-page paging. - -**Manifest:** add per-(marker,target) `attn` availability + `n_heads` so the View switcher greys out -Attention heads where absent (v1: present only for phase geneKO genes with an npz). - -**Decisions (locked 2026-07-08):** (a) normalization = **live user option** in Display settings -(per-map / per-gene / fixed) via the grayscale+LUT approach above; (b) full reorientation approved -(persistent Browse + view switcher); (c) **phase-only v1**, fluor (Kevin's per-marker npz) is a follow-on. - -### 2026-07-08 — phenotype-cell handoff CSV + v2 mAP + EBI matrix (see dashboard LATEST) -- **`viewer/phenotype_cells.py`** — the Ritvik handoff: top-20 REAL phenotype cells per (marker × perturbation) - → `viewer_assets/phenotype_cells_for_attention.csv` (160,420 cells / 53 markers). Key correctness fix: - the pma `rank` is GLOBAL per geneKO (across all channels), so `_csv_top` re-ranks WITHIN each - (channel, perturbation) and takes each marker's own top-20 by attention (two-pass chunked `head`). - Fluor filtered by mAP ≥ 0.2; phase = ALL. Added `map_score`, `geneKO`, `ebi_complex`, `rank_source` cols. -- **v2 mAP:** `catalog.dist_matrix` → `paper_v2/with_cp/with_4i/all_livecell` (56 reporters, live+CP+4i); - added 4i `FIXED_REP`; 52/56 pma channels map (NFkB/RSP6/Rb/gH2AX expected-excluded). -- **`viewer/build_complex_ebi_map.py`** — complex×reporter EBI mAP (98×56) via copairs - `phenotypic_consistency_ebi`, all-perturbation, on v2 per_signal; also emitted by the aggregation pipeline. -- **Pending:** centroid fallback for cisGolgi/VIM/LMNB1 (mAP present, no pma cells) from - `cell_dino_features_v2/features_processed_.h5ad` (embedding + crop metadata); SetTransformer - bag scoring parked on Alex's v2 CellDINO extraction. - -### 2026-07-09 — Argus app staging (S3 hosting) -- **Infra PR #51 OPENED + ACCEPTED/MERGED** (`sfbiohub-infra`, branch `diffex-viewer-dev`, - `terraform/accounts/biohub-nonprod/diffex-viewer-dev.tf`): S3 bucket `diffex-viewer-dev`, read-only IRSA - role `biohub-nonprod-diffex-viewer` (trusts SA `diffex-viewer` in ns `argus-diffex-viewer-rdev`), read-write - uploader role `diffex-viewer-dev-readwrite`. 1 TB ceiling. → `terraform apply` provisions bucket+roles. -- **App scaffold built** at `/hpc/mydata/gav.sturm/diffex-viewer` — forked from `czbiohub-sf/mops-viewer` - (same Argus+S3 pattern; names already align with the `.tf`). Serving layer swapped Gradio → **static nginx**: - - `Dockerfile` (nginx, bakes `webapp/` shell) + `nginx.conf` (port 8080, `/healthz`) + `docker-entrypoint-diffex.sh` - (drops shell into the S3-populated webroot, starts nginx). - - `.infra/common.yaml`: aws-cli `fetch-assets` init container (`aws s3 cp s3://diffex-viewer-dev/ → /usr/share/nginx/html/`), - `web` emptyDir, serviceAccount `diffex-viewer`, OIDC proxy; `.infra/rdev/values.yaml` carries the read-only role ARN. - - `.argus-ci.yaml` app=`diffex-viewer`; `scripts/uploader/*` sync `viewer_assets/` (manifest.json at bucket ROOT) via the readwrite role; `DEPLOY.md` runbook. - - Static because the webapp reads assets by relative path from `manifest.json`; the S3 sync drops assets as siblings of the baked shell. -- **Remaining to deploy:** `terraform apply` → create `czbiohub-sf/diffex-viewer` repo + push → `argus register app` - (team-sci-biohub) + bootstrap/reconcile `.github` (needs argus CLI, interactive) → upload `viewer_assets/` (our - role, or **Kyle from HPC** — he offered) → PR + `stack` label → Argus builds + deploys rdev behind Okta. Confirm - the registered namespace/SA matches `argus-diffex-viewer-rdev`/`diffex-viewer` before relying on the IRSA trust. - -### 2026-07-11 — full 1k-geneKO + complex buildout LAUNCHED -- **geneKO (master `34826112`, 49 jobs):** `submit seed --map-thr 0 --timeout 720` — all ~1000 geneKO - genes/marker for the 46 valid-`rep` fluor markers (49,000 targets). Resume skips the ~100–194 already - seeded per marker; the 4 hub markers (ChromaLIVE_561 957, LysoTracker 932, NPM3 930, NucleoLIVE 780) - were already near-complete. Runs with **no concurrency cap**. -- **complex (master `34826156`, 50 jobs):** `submit fluor-complex` resume — nearly all 98 EBI complexes - were already built for ~45 markers; only 5 partial (peroxisome_Peroxi 35, pS6 35, pRb 59, NPM3 86, SRRM2 88). -- **`submit.py` fixes made for this run:** - - **BUG:** `cmd_seed` used `if args.map_thr:` — `0` is falsy, so `--map-thr 0` silently fell back to top-8 - (a wrong launch, 34826071, was cancelled). Changed all gates to `args.map_thr is not None` → `--map-thr 0` - now correctly means all ~1000 genes. - - **`--parallel` now defaults to `None` = no concurrency cap** (only sets `slurm_array_parallelism` when given), - so full buildouts don't need `--parallel 100`. Added `--timeout` (default 180) to override the per-marker wall. -- **After the builds land:** `submit sync` to refresh manifest + attention + montages from the new cache. -- The 4 rep=None markers (NFkB/RSP6/Rb/gH2AX) remain intentionally excluded (no v2 distinctiveness reporter column). - -### 2026-07-12 — DiffAE cond_ratio PEAKS EARLY then declines; check before extending/rebuilding -Resumed the 6 under-trained (ep55) fluor DiffAEs to ep120 (dihedral). Key lesson: **most peaked their -conditioning ratio around ep39–55 and then DECLINED** — extending to 120 did NOT improve them: -| marker | best cond_ratio | peaked @ | verdict | -|---|---|---|---| -| CLTA | 0.390 | ep54 | plateaued/declined → no gain | -| ATP1B3 | 0.171 | ep54 | plateaued/declined → no gain | -| PSMB7 | 0.830 | ep54 | declined to ~0.49 → **stopped** | -| TFRC | 0.320 | ep39 | declined to ~0.21 → **stopped** | -| VAMP3 | 0.231 | ep54 | declined to ~0.17 → **stopped** | -| **SLC3A2** | **0.277** | **ep109** | still climbing (>pre55 0.259) → **kept running** | - -**RULES (to not repeat the wasted compute):** -1. `diffae_best.pt` is saved on BEST cond_ratio, so it already captures the peak regardless of final epoch — - a longer run does NOT give a better checkpoint unless best_ratio actually advanced. -2. **Before extending training** past its current point, check the cond_ratio trajectory - (`torch.load(train_state)['history']`). If it peaked early and is declining, stop — the best is already banked. -3. **Before clearing + rebuilding a marker's traversals** (1k geneKO / 98 complex / anchors), confirm the model - IMPROVED: compare `diffae_best.pt` mtime + best_ratio vs the value the existing traversals were built on. - Only rebuild if the best genuinely advanced PAST what the current traversals used. -- **Rebuild status:** CLTA/ATP1B3/PSMB7/TFRC/VAMP3 = NO rebuild (best@≤ep54 already used by the Jul-11 traversals). - **SLC3A2 = the one rebuild candidate** — its best advanced to 0.277@ep109 (Jul 12) vs the Jul-11 traversal - checkpoint (~0.259); once it finishes, clear + rebuild ONLY SLC3A2's traversals with the improved model. -### 2026-07-12 — accuracy-selected cell variant (Kevin's accuracy_ranking CSVs) -Cell selection can now use **classifier-accuracy rank** instead of **attention rank**. Source = -`…/alex_lin_attention/v4/accuracy_ranking/`: `pergene_phase_cell_rankings.csv` (geneKO·phase), -`ebi_pergene_phase_cell_rankings.csv` (complex·phase), `ebi_class_channel_cell_rankings.csv` (complex·fluor, -55 marker channels). NTC has NO accuracy data → NTC anchor always stays attention-sourced (shared/cached). -- **phase** = side-by-side A/B: modality `phase` (attention, "phase_attention") vs `phase_topacc` - ("phase_accuracy"). Full 1k geneKO + 98 complex + 182 anchors built for phase_topacc. Hooks: - `_gather_class(parquet=…)`, `precompute_marker(accuracy_parquet=, variant=)`, and the anchor path - `_gather_df`/`_setup`/`precompute_target(accuracy_parquet=, variant=)`. Accuracy dirs have LARGER - control→KD gaps (cleaner separation) than attention. -- **fluor** = REPLACED IN PLACE (no `_acc` duplicate), via `precompute_marker(accuracy_fluor_csv=, force=True)`. - **⚠️ COVERAGE SPLIT:** the fluor accuracy CSV only covers the complexes each marker actually distinguishes - (a SUBSET of the 98 — **median 13/marker, range 1–32, 53 distinct total**), so each fluor marker's `complex/` set is now MIXED: - accuracy-selected for its covered complexes, still attention-selected for the rest. This is intentional/interim - — we will get full 98-complex + 1k-geneKO accuracy coverage for every marker later and **rebuild the entire - cache anyway**, at which point the split disappears. geneKO fluor has NO accuracy CSV yet (phase-only + complex-fluor). -- **fluor complex→complex ANCHORS: NET-NEW with accuracy** (didn't exist in attention). Per marker: top-5 - accuracy-covered complexes (by `class_channel_acc`) → 20 A→B pairs, accuracy cells for both, via per-channel - parquets in `accuracy_ranking/fluor_complex_by_channel/`. 54 markers × ~20 ≈ 982 pairs → `/complex/`. -- **PHASE SWAP DONE (2026-07-12): accuracy is now the canonical `phase`.** `viewer_assets/phase` (attention) + - its `_directions/phase` archived to `viewer_assets_backup/{phase_attention,_directions_phase_attention}` (same FS, - reversible); `phase_topacc` → `phase`. Manifest label reverted to plain "Phase". NOTE: the separate build-cache - `…/diffex/directions/{phase,phase_topacc}/` (anchor gather cache, OUTSIDE viewer_assets) was NOT renamed — only the - served `viewer_assets` traversals + `_directions` were swapped; a future full rebuild regenerates it anyway. - -- **min-ep gate REMOVED as default** (`submit seed --min-ep` default 98→**0**; `catalog.complete_markers` default - 98→**0**). Epoch count is NOT a quality signal — a marker peaking at ep54 is as usable as one at ep120, and - `diffae_best.pt` banks the peak. Inclusion now gates on **checkpoint presence** (diffae_best.pt + train_state), - not epoch. Pass `--min-ep N` only to re-impose a floor. (Right now all 56 markers pass either way since the - Jul-11 resume pushed the 6 past ep98, but the default now won't silently drop a future under-trained marker.) diff --git a/src/ops_model/models/interpretability/diffae/classifier/config.py b/src/ops_model/models/interpretability/diffae/classifier/config.py index 2d53d7d..9cb91f2 100644 --- a/src/ops_model/models/interpretability/diffae/classifier/config.py +++ b/src/ops_model/models/interpretability/diffae/classifier/config.py @@ -28,7 +28,6 @@ GRAINS = { "geneKO": {"parquet": PMA_PHASE_GENEKO, "class_col": "gene"}, "complex": {"parquet": PMA_PHASE_EBI, "class_col": "predicted_class"}, # v5 complex parquet is complex-labeled (member cells pooled), like v4 - "minibinder": {"parquet": PMA_PHASE_GENEKO, "class_col": "gene"}, # NTC anchor from phase geneKO; targets supplied via accuracy_parquet } # Default output root; per-run results land under ///. diff --git a/src/ops_model/models/interpretability/diffae/traversal/catalog.py b/src/ops_model/models/interpretability/diffae/traversal/catalog.py index f503476..820c906 100644 --- a/src/ops_model/models/interpretability/diffae/traversal/catalog.py +++ b/src/ops_model/models/interpretability/diffae/traversal/catalog.py @@ -173,18 +173,11 @@ def dist_map_for_assets(viewer_assets): cx = complex_dist() # complex × reporter EBI mAP (complexes aren't in the gene matrix) except Exception: cx = None - try: - mb = json.load(open(f"{viewer_assets}/_minibinder_meta.json")) # minibinders → per-binder cell_score - except Exception: - mb = {} out = {} for mj in glob.glob(f"{viewer_assets}/*/*/*/meta.json"): m = json.load(open(mj)) rep = (rep_of(dist, m["marker_channel"]) if m.get("marker_channel") else "Phase") - if m["grain"] == "minibinder": # minibinders → cell_score (no mAP) - if m["slug"] in mb: - out[(m["modality"], m["grain"], m["slug"])] = float(mb[m["slug"]]["cell_score"]) - elif m["grain"] == "complex": # complexes → EBI complex mAP + if m["grain"] == "complex": # complexes → EBI complex mAP if cx is not None and m["target"] in cx.index and rep in cx.columns: v = cx.at[m["target"], rep] if pd.notna(v): diff --git a/src/ops_model/models/interpretability/diffae/traversal/precompute.py b/src/ops_model/models/interpretability/diffae/traversal/precompute.py index 4b37968..9282b55 100644 --- a/src/ops_model/models/interpretability/diffae/traversal/precompute.py +++ b/src/ops_model/models/interpretability/diffae/traversal/precompute.py @@ -462,16 +462,10 @@ def build_manifest(out_root, dist_map=None, desc_map=None): dist_map: optional {(modality, grain, slug): mAP} to attach for sorting targets. desc_map: optional {target_name: description} (gene function / complex members).""" root = Path(out_root) / _ASSETS - try: - mb = json.loads((root / "_minibinder_meta.json").read_text()) # per-binder cell_score/binder_prob/gene_target - except Exception: - mb = {} markers = {} for mj in sorted(root.glob("*/*/*/meta.json")): m = json.loads(mj.read_text()) mod = m["modality"] - if mod == "phase_minibinder": # orphan from a cancelled mis-structured run (couldn't rm on shared FS); minibinders live under phase/minibinder - continue label = m["marker_channel"] or "Phase" # phase = canonical (accuracy-selected as of 2026-07-12 swap) mk = markers.setdefault(mod, {"modality": mod, "marker_channel": m["marker_channel"], "label": label, "channel": m["channel"], "targets": []}) @@ -485,10 +479,7 @@ def build_manifest(out_root, dist_map=None, desc_map=None): "n_cells": m["n_cells"], "asset_dir": adir, "alphas": m["alphas"], "dist_map": (dist_map or {}).get(key), "explained_variance": m.get("explained_variance"), # PC grain: % variance - "desc": (desc_map or {}).get(m["target"]), - **({"binder_prob": mb[m["slug"]]["binder_prob"], "gene_target": mb[m["slug"]]["gene_target"], - "phenotype": mb[m["slug"]]["phenotype"], "cell_score": mb[m["slug"]]["cell_score"]} - if m["grain"] == "minibinder" and m["slug"] in mb else {})}) + "desc": (desc_map or {}).get(m["target"])}) for mk in markers.values(): mk["targets"].sort(key=lambda t: (-(t["dist_map"] or -1), t["target"])) manifest = {"alphas": list(VIEWER_ALPHAS), "w": 2.0, diff --git a/src/ops_model/post_process/anndata_processing/anndata_validator.py b/src/ops_model/post_process/anndata_processing/anndata_validator.py index 8f49314..8e3f628 100644 --- a/src/ops_model/post_process/anndata_processing/anndata_validator.py +++ b/src/ops_model/post_process/anndata_processing/anndata_validator.py @@ -442,7 +442,7 @@ class AnndataSpec: Dictionary mapping schema level names to their specifications guide_col : str Name of the .obs column holding per-construct identifiers - (e.g. "sgRNA" for CRISPR, "minibinder_perturbation" for minibinder). + (e.g. "sgRNA" for CRISPR, or a custom perturbation column). """ def __init__(self, guide_col: str = DEFAULT_GUIDE_COL): @@ -618,7 +618,7 @@ def _define_cell_schema(self) -> Dict[str, Any]: name=self.guide_col, dtype=str, required=True, - description="Per-construct identifier (e.g. sgRNA, minibinder peptide)", + description="Per-construct identifier (e.g. sgRNA or a custom construct id)", suggestion=f"Add {self.guide_col} column with construct identifiers", ), FieldSpec( @@ -673,7 +673,7 @@ def _define_guide_schema(self) -> Dict[str, Any]: dtype=str, required=True, unique=False, - description="Per-construct identifier (e.g. sgRNA, minibinder peptide)", + description="Per-construct identifier (e.g. sgRNA or a custom construct id)", suggestion=f"Add {self.guide_col} column with construct identifiers", ), ], @@ -844,7 +844,7 @@ def __init__(self, strict: bool = True, guide_col: Optional[str] = None): failure. If False, returns ValidationReport without raising. guide_col : Optional[str], default=None Name of the .obs column holding per-construct identifiers - (e.g. "sgRNA" for CRISPR, "minibinder_perturbation" for minibinder). + (e.g. "sgRNA" for CRISPR, or a custom perturbation column). If None, validate() will read it from adata.uns["guide_col"] when available, falling back to the default ("sgRNA"). """ diff --git a/tests/features/test_anndata_utils_guide_col.py b/tests/features/test_anndata_utils_guide_col.py index 56f73a4..bb3825c 100644 --- a/tests/features/test_anndata_utils_guide_col.py +++ b/tests/features/test_anndata_utils_guide_col.py @@ -1,7 +1,7 @@ """Tests for configurable guide_col in anndata_utils aggregation paths. Covers PR3 of the refactor — replacing hardcoded "sgRNA" with the -adata.uns["guide_col"] lookup so minibinder-style experiments aggregate +adata.uns["guide_col"] lookup so custom-guide-style experiments aggregate correctly without aliasing peptide identifiers onto an "sgRNA" column. """ @@ -66,18 +66,18 @@ def test_guide_col_helper_default(): def test_guide_col_helper_reads_uns(): adata = ad.AnnData(X=np.zeros((1, 1), dtype=np.float32)) - adata.uns["guide_col"] = "minibinder_perturbation" - assert _guide_col(adata) == "minibinder_perturbation" + adata.uns["guide_col"] = "custom_perturbation" + assert _guide_col(adata) == "custom_perturbation" -def test_aggregate_cell_to_guide_minibinder(): - """Cell→guide aggregation uses the minibinder column, not sgRNA.""" - adata = _make_cell_adata("minibinder_perturbation", set_uns_guide_col=True) +def test_aggregate_cell_to_guide_custom_guide(): + """Cell→guide aggregation uses the custom_guide column, not sgRNA.""" + adata = _make_cell_adata("custom_perturbation", set_uns_guide_col=True) adata_guide = aggregate_to_level(adata, level="guide") - assert "minibinder_perturbation" in adata_guide.obs.columns + assert "custom_perturbation" in adata_guide.obs.columns assert "sgRNA" not in adata_guide.obs.columns - assert adata_guide.uns["guide_col"] == "minibinder_perturbation" - assert sorted(adata_guide.obs["minibinder_perturbation"].unique()) == [ + assert adata_guide.uns["guide_col"] == "custom_perturbation" + assert sorted(adata_guide.obs["custom_perturbation"].unique()) == [ "CDC5L_g0", "CDC5L_g1", "POLR2C_g0", @@ -86,13 +86,13 @@ def test_aggregate_cell_to_guide_minibinder(): ] -def test_aggregate_cell_to_gene_minibinder_no_nan_bug(): +def test_aggregate_cell_to_gene_custom_guide_no_nan_bug(): """The original NaN bug: gene-level aggregation must not choke on the neg_ctrl 'no-peptide' group when guide_col is properly configured.""" - adata = _make_cell_adata("minibinder_perturbation", set_uns_guide_col=True) + adata = _make_cell_adata("custom_perturbation", set_uns_guide_col=True) adata_gene = aggregate_to_level(adata, level="gene") assert "guides" in adata_gene.obs.columns - assert adata_gene.uns["guide_col"] == "minibinder_perturbation" + assert adata_gene.uns["guide_col"] == "custom_perturbation" # The negative-control row contains the literal "no-peptide" string, # not a NaN, so pipe-joining works. neg_row = adata_gene.obs[adata_gene.obs["perturbation"] == "NEG_CTRL"].iloc[0] @@ -109,18 +109,18 @@ def test_aggregate_cell_to_gene_crispr_default(): assert (adata_gene.obs["perturbation"] == "NEG_CTRL").any() -def test_aggregate_guide_to_gene_minibinder(): +def test_aggregate_guide_to_gene_custom_guide(): """Guide → gene aggregation also threads guide_col through.""" - adata_cell = _make_cell_adata("minibinder_perturbation", set_uns_guide_col=True) + adata_cell = _make_cell_adata("custom_perturbation", set_uns_guide_col=True) adata_guide = aggregate_to_level(adata_cell, level="guide") adata_gene = aggregate_to_level(adata_guide, level="gene") - assert adata_gene.uns["guide_col"] == "minibinder_perturbation" + assert adata_gene.uns["guide_col"] == "custom_perturbation" assert "guides" in adata_gene.obs.columns def test_hconcat_by_perturbation_uses_guide_col(): """Horizontal-concat join key follows adata.uns['guide_col'].""" - a = _make_cell_adata("minibinder_perturbation", set_uns_guide_col=True) + a = _make_cell_adata("custom_perturbation", set_uns_guide_col=True) a_guide = aggregate_to_level(a, level="guide") # Build a second block with a different feature set but same guides b_guide = a_guide.copy() @@ -129,4 +129,4 @@ def test_hconcat_by_perturbation_uses_guide_col(): merged = hconcat_by_perturbation([a_guide, b_guide], level="guide") # 5 unique guides aligned, 10 stacked features assert merged.shape == (5, 10) - assert "minibinder_perturbation" in merged.obs.columns + assert "custom_perturbation" in merged.obs.columns diff --git a/tests/post_process/test_anndata_validator_guide_col.py b/tests/post_process/test_anndata_validator_guide_col.py index f1a21c8..ffad480 100644 --- a/tests/post_process/test_anndata_validator_guide_col.py +++ b/tests/post_process/test_anndata_validator_guide_col.py @@ -1,7 +1,7 @@ """Tests for configurable guide_col in AnndataValidator. Covers the PR2 refactor that lets the per-construct identifier column be -named anything (e.g. "minibinder_perturbation") instead of hardcoded "sgRNA". +named anything (e.g. "custom_perturbation") instead of hardcoded "sgRNA". """ import warnings @@ -61,15 +61,15 @@ def test_anndataspec_defaults_to_sgrna(): def test_anndataspec_uses_custom_guide_col(): - spec = AnndataSpec(guide_col="minibinder_perturbation") + spec = AnndataSpec(guide_col="custom_perturbation") cell_fields = spec.get_schema("cell")["required_fields"] names = [f.name for f in cell_fields] - assert "minibinder_perturbation" in names + assert "custom_perturbation" in names assert "sgRNA" not in names guide_fields = spec.get_schema("guide")["required_fields"] g_names = [f.name for f in guide_fields] - assert "minibinder_perturbation" in g_names + assert "custom_perturbation" in g_names assert "sgRNA" not in g_names @@ -80,10 +80,10 @@ def test_legacy_crispr_anndata_validates_with_default(): assert report.is_valid, [(e.field, e.message) for e in report.errors] -def test_minibinder_anndata_validates_when_uns_guide_col_set(): - """AnnData with obs['minibinder_perturbation'] + uns['guide_col'] passes.""" +def test_custom_guide_anndata_validates_when_uns_guide_col_set(): + """AnnData with obs['custom_perturbation'] + uns['guide_col'] passes.""" adata = _make_cell_adata( - "minibinder_perturbation", + "custom_perturbation", ["2_1921", "2_1010", "no-peptide"], set_uns_guide_col=True, ) @@ -91,10 +91,10 @@ def test_minibinder_anndata_validates_when_uns_guide_col_set(): assert report.is_valid, [(e.field, e.message) for e in report.errors] -def test_minibinder_data_without_uns_key_fails_on_default_column(): +def test_custom_guide_data_without_uns_key_fails_on_default_column(): """If uns['guide_col'] is absent, validator falls back to 'sgRNA' and fails.""" adata = _make_cell_adata( - "minibinder_perturbation", + "custom_perturbation", ["2_1921", "2_1010", "no-peptide"], set_uns_guide_col=False, ) @@ -106,11 +106,11 @@ def test_minibinder_data_without_uns_key_fails_on_default_column(): def test_validator_can_be_initialized_with_explicit_guide_col(): """Explicit guide_col on the validator overrides the default fallback.""" adata = _make_cell_adata( - "minibinder_perturbation", + "custom_perturbation", ["2_1921", "2_1010", "no-peptide"], set_uns_guide_col=False, ) - v = AnndataValidator(strict=False, guide_col="minibinder_perturbation") + v = AnndataValidator(strict=False, guide_col="custom_perturbation") report = v.validate(adata, level="cell") assert report.is_valid, [(e.field, e.message) for e in report.errors] @@ -120,7 +120,7 @@ def test_guide_level_uses_dynamic_column(): n = 6 obs = pd.DataFrame( { - "minibinder_perturbation": [f"p{i}" for i in range(n)], + "custom_perturbation": [f"p{i}" for i in range(n)], "perturbation": ["G"] * n, "reporter": ["GFP"] * n, "n_cells": [10] * n, @@ -129,7 +129,7 @@ def test_guide_level_uses_dynamic_column(): adata = ad.AnnData(X=np.random.rand(n, 4).astype(np.float32), obs=obs) adata.uns["cell_type"] = "A549" adata.uns["embedding_type"] = "dinov3" - adata.uns["guide_col"] = "minibinder_perturbation" + adata.uns["guide_col"] = "custom_perturbation" adata.uns["aggregation_method"] = "mean" report = AnndataValidator(strict=False).validate(adata, level="guide") @@ -141,14 +141,14 @@ def test_infer_schema_level_uses_guide_col(): n = 10 obs = pd.DataFrame( { - "minibinder_perturbation": [f"p{i}" for i in range(n)], + "custom_perturbation": [f"p{i}" for i in range(n)], "perturbation": ["G"] * n, "reporter": ["GFP"] * n, "n_cells": [10] * n, } ) adata = ad.AnnData(X=np.random.rand(n, 4).astype(np.float32), obs=obs) - adata.uns["guide_col"] = "minibinder_perturbation" + adata.uns["guide_col"] = "custom_perturbation" level = AnndataValidator(strict=False).infer_schema_level(adata) assert level == "guide" diff --git a/tests/test_dataloader.py b/tests/test_dataloader.py index 0936110..b133516 100644 --- a/tests/test_dataloader.py +++ b/tests/test_dataloader.py @@ -107,9 +107,9 @@ def test_ops_data_manager_stores_guide_col(): """OpsDataManager records the configured guide_col on the instance.""" dm = data_loader.OpsDataManager( experiments={"ops0031_20250424": ["A/1/0"]}, - guide_col="minibinder_perturbation", + guide_col="custom_perturbation", ) - assert dm.guide_col == "minibinder_perturbation" + assert dm.guide_col == "custom_perturbation" def test_ops_data_manager_default_guide_col(): @@ -122,24 +122,24 @@ def test_base_dataset_stores_guide_col(): """BaseDataset captures guide_col on the instance for feature extractors to read.""" df = pd.DataFrame( { - "minibinder_perturbation": ["mb_001"], + "custom_perturbation": ["cp_001"], "gene_name": ["EGFR"], "bbox": ["[0,0,10,10]"], } ) ds = data_loader.BaseDataset( - stores={}, labels_df=df, guide_col="minibinder_perturbation" + stores={}, labels_df=df, guide_col="custom_perturbation" ) - assert ds.guide_col == "minibinder_perturbation" + assert ds.guide_col == "custom_perturbation" -def test_get_labels_minibinder_fallback(tmp_path): - """When the link CSV is minibinder-style (no gene_name column), get_labels - copies minibinder_perturbation into gene_name so the downstream gene_name +def test_get_labels_custom_guide_fallback(tmp_path): + """When the link CSV is custom-guide-style (no gene_name column), get_labels + copies custom_perturbation into gene_name so the downstream gene_name flow keeps working.""" df = pd.DataFrame( { - "minibinder_perturbation": ["mb_001", "mb_002"], + "custom_perturbation": ["cp_001", "cp_002"], "AA_sequence": ["MASTK...", "ABCDE..."], "gene_target": ["EGFR", "BRCA1"], "segmentation_id": [1, 2], @@ -151,15 +151,15 @@ def test_get_labels_minibinder_fallback(tmp_path): dm = data_loader.OpsDataManager( experiments={"ops_test": ["A/1/0"]}, link_csv_dir=str(tmp_path), - guide_col="minibinder_perturbation", + guide_col="custom_perturbation", ) labels = dm.get_labels() # The guide column is preserved with its original name (not aliased). - assert "minibinder_perturbation" in labels.columns + assert "custom_perturbation" in labels.columns assert "sgRNA" not in labels.columns - # gene_name has been copied from minibinder_perturbation. + # gene_name has been copied from custom_perturbation. assert "gene_name" in labels.columns - assert list(labels["gene_name"]) == ["mb_001", "mb_002"] + assert list(labels["gene_name"]) == ["cp_001", "cp_002"] def test_get_labels_fails_loudly_when_guide_col_missing(tmp_path): @@ -167,7 +167,7 @@ def test_get_labels_fails_loudly_when_guide_col_missing(tmp_path): immediately rather than letting NaNs propagate downstream.""" df = pd.DataFrame( { - "minibinder_perturbation": ["mb_001"], + "custom_perturbation": ["cp_001"], "segmentation_id": [1], "bbox": ["[0,0,10,10]"], } From c395a424ca5cbdc547545ce2d2782ca5ab511302 Mon Sep 17 00:00:00 2001 From: Gav Sturm Date: Wed, 12 Aug 2026 09:41:41 -0700 Subject: [PATCH 11/13] public-release: drop personal /hpc/mydata/gav.sturm paths + strip latent-lens from METHODS - virtstain_multi: removed the personal slurm_logs glob fallback (NaN->None instead) - figures/METHODS_*: montage described generically (dropped latent-lens package name+URL) --- .../interpretability/diffae/figures/METHODS_final.txt | 2 +- .../diffae/figures/METHODS_traversal_montage.md | 4 ++-- .../diffae/figures/METHODS_traversal_montage.txt | 2 +- .../interpretability/diffae/generator/virtstain_multi.py | 7 ++----- 4 files changed, 6 insertions(+), 9 deletions(-) diff --git a/src/ops_model/models/interpretability/diffae/figures/METHODS_final.txt b/src/ops_model/models/interpretability/diffae/figures/METHODS_final.txt index 54691e8..4d7a17e 100644 --- a/src/ops_model/models/interpretability/diffae/figures/METHODS_final.txt +++ b/src/ops_model/models/interpretability/diffae/figures/METHODS_final.txt @@ -56,7 +56,7 @@ Distributions are shown as violins of the per-cell values with a bar at the medi EXCLUDE -The montage was assembled with the open-source latent-lens package (https://github.com/czi-ai/latent-lens; multiscale montages of image crops laid out by an embedding), which performs the grid-based, density-prioritized decimation and the multiscale tiling. +The montage was assembled by grid-based, density-prioritized decimation and multiscale tiling of the image crops laid out by an embedding. Quantifying the accuracy of generated phenotypes To test whether the synthesized cells reproduce the intended perturbation phenotype, each traversal was scored by three complementary metrics, all computed as a function of α and interpreted relative to the value attainable on real cells of the same class. diff --git a/src/ops_model/models/interpretability/diffae/figures/METHODS_traversal_montage.md b/src/ops_model/models/interpretability/diffae/figures/METHODS_traversal_montage.md index 5a8c7c6..881874f 100644 --- a/src/ops_model/models/interpretability/diffae/figures/METHODS_traversal_montage.md +++ b/src/ops_model/models/interpretability/diffae/figures/METHODS_traversal_montage.md @@ -106,7 +106,7 @@ cluster to convey the local phenotypic neighborhood, with individual complex mem ribo40S, ribo60S). The result is a single view in which each region of the phenotypic embedding is illustrated by a representative generated cell. -The montage was assembled with the open-source latent-lens package -(https://github.com/czi-ai/latent-lens; multiscale montages of image crops laid out by an embedding), +The montage was assembled by grid-based, density-prioritized decimation and +multiscale tiling of the image crops laid out by an embedding, which performs the grid-based, density-prioritized decimation and the multiscale tiling. diff --git a/src/ops_model/models/interpretability/diffae/figures/METHODS_traversal_montage.txt b/src/ops_model/models/interpretability/diffae/figures/METHODS_traversal_montage.txt index 8a5a616..70b1674 100644 --- a/src/ops_model/models/interpretability/diffae/figures/METHODS_traversal_montage.txt +++ b/src/ops_model/models/interpretability/diffae/figures/METHODS_traversal_montage.txt @@ -35,4 +35,4 @@ Gene embedding and montage A gene-level phenotypic embedding was constructed by aggregating the single-cell Cell-DINO features of each perturbation, and a two-dimensional layout was computed with PHATE. To build the montage, each gene's generated cell (at a chosen α) was placed at that gene's coordinate in this embedding. Because many genes occupy dense regions, the embedding plane was tiled with a regular grid and a single representative gene was retained per grid cell — chosen by local density — so that adjacent cells do not overlap; coarser grids show fewer, larger cells and finer grids fill in more. Each tile was outlined by the color of its gene's Leiden cluster to convey the local phenotypic neighborhood, with individual complex members highlighted (e.g. ribo40S, ribo60S). The result is a single view in which each region of the phenotypic embedding is illustrated by a representative generated cell. -The montage was assembled with the open-source latent-lens package (https://github.com/czi-ai/latent-lens; multiscale montages of image crops laid out by an embedding), which performs the grid-based, density-prioritized decimation and the multiscale tiling. +The montage was assembled by grid-based, density-prioritized decimation and multiscale tiling of the image crops laid out by an embedding. diff --git a/src/ops_model/models/interpretability/diffae/generator/virtstain_multi.py b/src/ops_model/models/interpretability/diffae/generator/virtstain_multi.py index 5be487e..f21dee5 100644 --- a/src/ops_model/models/interpretability/diffae/generator/virtstain_multi.py +++ b/src/ops_model/models/interpretability/diffae/generator/virtstain_multi.py @@ -250,11 +250,8 @@ def eval_only(cap=2500, batch=48, device="cuda", subdir="eval"): st = torch.load(Path(OUT) / "train_state.pt", map_location=dev) ema.load_state_dict(st["ema"]); print(f"[eval-only] loaded EMA @ epoch {st['epoch']}") loss = st.get("loss") - if loss is None or (isinstance(loss, float) and loss != loss): # missing/NaN → read latest from the train log - import glob, os, re - for f in sorted(glob.glob("/hpc/mydata/gav.sturm/ops_mono/slurm_logs/diffae/*/*.out"), key=os.path.getmtime)[::-1][:8]: - m = re.findall(r"epoch \d+: loss=([\d.]+)", open(f).read()) - if m: loss = float(m[-1]); break + if loss is not None and isinstance(loss, float) and loss != loss: # NaN → unknown + loss = None (Path(OUT) / "markers.json").write_text(json.dumps([mc for _, mc, _ in kept], indent=2)) eval_multi(ema, X, E, P, M, kept, cfg, OUT, dev, epoch=st["epoch"], loss=loss, n_train=int(len(X)), eval_name=subdir) return {"epoch": st["epoch"], "n_markers": len(kept)} From 82147b0e87cf311b6c5245872d461299982c7100 Mon Sep 17 00:00:00 2001 From: Gav Sturm Date: Wed, 12 Aug 2026 10:00:11 -0700 Subject: [PATCH 12/13] paths: normalize legacy roots to canonical icd.fast.ops (icd.ops removed) --- src/ops_model/data/labels.py | 4 ++-- src/ops_model/data/paths.py | 2 +- src/ops_model/features/anndata_utils.py | 4 ++-- src/ops_model/features/batch_process_embeddings.py | 2 +- .../diffae/traversal/build_pc_crops_masked.py | 2 +- .../models/interpretability/toolkit/atlas/attention_atlas.py | 2 +- .../interpretability/toolkit/viewer/build_pcs_marker.py | 2 +- .../models/interpretability/toolkit/viewer/morphometrics.py | 2 +- tests/features/test_comprehensive_combination.py | 2 +- 9 files changed, 11 insertions(+), 11 deletions(-) diff --git a/src/ops_model/data/labels.py b/src/ops_model/data/labels.py index 1c384c6..2747699 100644 --- a/src/ops_model/data/labels.py +++ b/src/ops_model/data/labels.py @@ -11,7 +11,7 @@ import numpy as np import pandas as pd -_DEFAULT_BASE_PATH = "/hpc/projects/intracellular_dashboard/fast_ops" +_DEFAULT_BASE_PATH = "/hpc/projects/icd.fast.ops" # Backward-compatible filename templates for legacy csv_source values SOURCE_FILENAME_TEMPLATES = { @@ -71,7 +71,7 @@ def load_immunostaining_labels( filename_template: Filename pattern with {well} placeholder, e.g. "cell_painting_linked_{well}.csv" or "four_i_linked_{well}.csv" base_path: Base directory containing per-experiment subdirectories. - Defaults to /hpc/projects/intracellular_dashboard/fast_ops. + Defaults to /hpc/projects/icd.fast.ops. Returns: labels_df ready to pass to OpsDataManager.construct_dataloaders() diff --git a/src/ops_model/data/paths.py b/src/ops_model/data/paths.py index d795adc..b334efd 100644 --- a/src/ops_model/data/paths.py +++ b/src/ops_model/data/paths.py @@ -85,7 +85,7 @@ def __init__(self, experiment: str, well: str = None): } self.other = { - "gene_library": "/hpc/projects/intracellular_dashboard/ops/configs/annotated_guide_library_123-UpdateJuly28_2025.csv", + "gene_library": "/hpc/projects/icd.fast.ops/configs/annotated_guide_library_123-UpdateJuly28_2025.csv", } def reformat_well_name(self, well: str) -> str: diff --git a/src/ops_model/features/anndata_utils.py b/src/ops_model/features/anndata_utils.py index aadba39..edd7c83 100644 --- a/src/ops_model/features/anndata_utils.py +++ b/src/ops_model/features/anndata_utils.py @@ -30,7 +30,7 @@ DEFAULT_SEARCH_DIRS = [ Path("/hpc/projects/icd.fast.ops"), - Path("/hpc/projects/icd.ops"), + Path("/hpc/projects/icd.fast.ops"), ] DEFAULT_GUIDE_COL = "sgRNA" @@ -1441,7 +1441,7 @@ def load_multiple_experiments( List of paths to .h5ad files Example: - >>> base_dir = "/hpc/projects/intracellular_dashboard/ops" + >>> base_dir = "/hpc/projects/icd.fast.ops" >>> experiments = ["ops0089_20251119", "ops0084_20250101"] >>> paths = load_multiple_experiments(base_dir, experiments) >>> adata_combined = concatenate_anndata_objects(paths) diff --git a/src/ops_model/features/batch_process_embeddings.py b/src/ops_model/features/batch_process_embeddings.py index 683001d..8bc28c6 100644 --- a/src/ops_model/features/batch_process_embeddings.py +++ b/src/ops_model/features/batch_process_embeddings.py @@ -32,7 +32,7 @@ # Base directory for OPS experiments -BASE_DIR = Path("/hpc/projects/intracellular_dashboard/ops") +BASE_DIR = Path("/hpc/projects/icd.fast.ops") def check_csv_exists( diff --git a/src/ops_model/models/interpretability/diffae/traversal/build_pc_crops_masked.py b/src/ops_model/models/interpretability/diffae/traversal/build_pc_crops_masked.py index 9172482..36b928e 100644 --- a/src/ops_model/models/interpretability/diffae/traversal/build_pc_crops_masked.py +++ b/src/ops_model/models/interpretability/diffae/traversal/build_pc_crops_masked.py @@ -23,7 +23,7 @@ from . import catalog as C -BASE = "/hpc/projects/intracellular_dashboard/fast_ops" +BASE = "/hpc/projects/icd.fast.ops" PCS_OUT = f"{C.OUT}/viewer_assets/pcs" CROP_SIZE = 160 # native px re-crop (was 96); crisper at the 150px display + shows surround PHASE_CHANNEL = 0 diff --git a/src/ops_model/models/interpretability/toolkit/atlas/attention_atlas.py b/src/ops_model/models/interpretability/toolkit/atlas/attention_atlas.py index a681d44..99d6b78 100644 --- a/src/ops_model/models/interpretability/toolkit/atlas/attention_atlas.py +++ b/src/ops_model/models/interpretability/toolkit/atlas/attention_atlas.py @@ -2161,7 +2161,7 @@ def _parse_ops_channel_maps_yaml(yaml_path=None): yaml_path = Path(OpsDataset(resolve_experiment_name("ops0107_20251208")).channel_maps) except Exception: yaml_path = Path( - "/hpc/projects/intracellular_dashboard/fast_ops/configs/ops_channel_maps.yaml" + "/hpc/projects/icd.fast.ops/configs/ops_channel_maps.yaml" ) try: with open(yaml_path) as f: diff --git a/src/ops_model/models/interpretability/toolkit/viewer/build_pcs_marker.py b/src/ops_model/models/interpretability/toolkit/viewer/build_pcs_marker.py index 9c563c3..8f71f5f 100644 --- a/src/ops_model/models/interpretability/toolkit/viewer/build_pcs_marker.py +++ b/src/ops_model/models/interpretability/toolkit/viewer/build_pcs_marker.py @@ -24,7 +24,7 @@ from . import marker_leaves as ML from ops_model.models.interpretability.diffae.traversal.build_pc_crops_masked import CROP_SIZE, _crop, _is_blank, _render, _zarr_patch -FOPS = "/hpc/projects/intracellular_dashboard/fast_ops" +FOPS = "/hpc/projects/icd.fast.ops" PCS_OUT = f"{C.OUT}/viewer_assets/pcs/markers" N_PCS, N_BINS, N_ROWS = 40, 15, 3 # PCs shown; strip bins; cells per bin FIT_N, SEL_N = 120_000, 300_000 # cells subsampled per experiment for PCA fit / representative selection diff --git a/src/ops_model/models/interpretability/toolkit/viewer/morphometrics.py b/src/ops_model/models/interpretability/toolkit/viewer/morphometrics.py index 44cbf4c..0436c4b 100644 --- a/src/ops_model/models/interpretability/toolkit/viewer/morphometrics.py +++ b/src/ops_model/models/interpretability/toolkit/viewer/morphometrics.py @@ -11,7 +11,7 @@ import numpy as np import yaml -ORG_SEG_YAML = "/hpc/projects/intracellular_dashboard/fast_ops/configs/org_seg_params.yaml" +ORG_SEG_YAML = "/hpc/projects/icd.fast.ops/configs/org_seg_params.yaml" CACHE = "/hpc/projects/icd.fast.ops/models/diffex/viewer_assets" PIXEL_UM = 0.325 # phenotype native pixel size diff --git a/tests/features/test_comprehensive_combination.py b/tests/features/test_comprehensive_combination.py index 57ee9a7..39afd84 100644 --- a/tests/features/test_comprehensive_combination.py +++ b/tests/features/test_comprehensive_combination.py @@ -385,7 +385,7 @@ def test_missing_metadata_raises(self): @pytest.mark.integration @pytest.mark.skipif( - not Path("/hpc/projects/intracellular_dashboard/ops").exists(), + not Path("/hpc/projects/icd.fast.ops").exists(), reason="Requires access to HPC data", ) class TestIntegration: From 99a20884dfd8ba1c01e8178178daea12fdc4873f Mon Sep 17 00:00:00 2001 From: Gav Sturm Date: Wed, 12 Aug 2026 10:07:46 -0700 Subject: [PATCH 13/13] paths: centralize on OPS_BASE_PATH via src/ops_model/paths.py Single env-overridable BASE_PATH (default /hpc/projects/icd.fast.ops); 25 files' path literals -> f"{BASE_PATH}/..." + import. One OPS_BASE_PATH swap relocates the whole data/model/analysis tree. Defaults byte-identical (verified); env-swap verified. (4 residuals are docstring examples, left as illustrative defaults.) --- src/ops_model/data/labels.py | 3 ++- src/ops_model/data/paths.py | 4 +-- src/ops_model/features/anndata_utils.py | 7 +++--- .../features/batch_process_embeddings.py | 3 ++- .../diffae/classifier/config.py | 7 +++--- .../diffae/directions/config.py | 3 ++- .../diffae/directions/proto_ddim_anchors.py | 7 +++--- .../diffae/figures/_setacc_common.py | 5 ++-- .../diffae/figures/_setacc_phase.py | 3 ++- .../diffae/figures/auto_pick_and_plot.py | 3 ++- .../figures/figure4_morpho_traversal.py | 5 ++-- .../diffae/figures/figure4_morpho_violin.py | 3 ++- .../figures/figure_ebi_morpho_violin.py | 5 ++-- .../figures/figure_multirank_ebi_grid.py | 7 +++--- .../diffae/figures/fluor_panel_montages.py | 3 ++- .../diffae/figures/fluor_shap_montages.py | 7 +++--- .../diffae/figures/nc_ratio.py | 5 ++-- .../diffae/figures/ntc_anchor_compare.py | 5 ++-- .../diffae/figures/phase_multibag_montages.py | 5 ++-- .../setacc/figure4_setacc_panel_newpheno.py | 3 ++- .../figures/virtual_staining_schematic.py | 3 ++- .../diffae/generator/plot_metrics.py | 5 ++-- .../diffae/generator/virtstain_multi.py | 5 ++-- .../diffae/traversal/build_pc_crops_masked.py | 3 ++- .../diffae/traversal/catalog.py | 13 +++++----- .../diffae/traversal/morpho_pipeline.py | 25 ++++++++++--------- src/ops_model/paths.py | 9 +++++++ 27 files changed, 95 insertions(+), 61 deletions(-) create mode 100644 src/ops_model/paths.py diff --git a/src/ops_model/data/labels.py b/src/ops_model/data/labels.py index 2747699..7feb64b 100644 --- a/src/ops_model/data/labels.py +++ b/src/ops_model/data/labels.py @@ -10,8 +10,9 @@ import numpy as np import pandas as pd +from ops_model.paths import BASE_PATH -_DEFAULT_BASE_PATH = "/hpc/projects/icd.fast.ops" +_DEFAULT_BASE_PATH = f"{BASE_PATH}" # Backward-compatible filename templates for legacy csv_source values SOURCE_FILENAME_TEMPLATES = { diff --git a/src/ops_model/data/paths.py b/src/ops_model/data/paths.py index b334efd..a347c1b 100644 --- a/src/ops_model/data/paths.py +++ b/src/ops_model/data/paths.py @@ -14,7 +14,7 @@ def _resolve_base() -> Path: return Path( os.environ.get( "OPS_OUTPUT_BASE_DIR", - "/hpc/projects/icd.fast.ops", + f"{BASE_PATH}", ) ) @@ -85,7 +85,7 @@ def __init__(self, experiment: str, well: str = None): } self.other = { - "gene_library": "/hpc/projects/icd.fast.ops/configs/annotated_guide_library_123-UpdateJuly28_2025.csv", + "gene_library": f"{BASE_PATH}/configs/annotated_guide_library_123-UpdateJuly28_2025.csv", } def reformat_well_name(self, well: str) -> str: diff --git a/src/ops_model/features/anndata_utils.py b/src/ops_model/features/anndata_utils.py index edd7c83..a18e283 100644 --- a/src/ops_model/features/anndata_utils.py +++ b/src/ops_model/features/anndata_utils.py @@ -26,11 +26,12 @@ import matplotlib.pyplot as plt from ops_utils.data.feature_metadata import FeatureMetadata +from ops_model.paths import BASE_PATH DEFAULT_SEARCH_DIRS = [ - Path("/hpc/projects/icd.fast.ops"), - Path("/hpc/projects/icd.fast.ops"), + Path(f"{BASE_PATH}"), + Path(f"{BASE_PATH}"), ] DEFAULT_GUIDE_COL = "sgRNA" @@ -1441,7 +1442,7 @@ def load_multiple_experiments( List of paths to .h5ad files Example: - >>> base_dir = "/hpc/projects/icd.fast.ops" + >>> base_dir = f"{BASE_PATH}" >>> experiments = ["ops0089_20251119", "ops0084_20250101"] >>> paths = load_multiple_experiments(base_dir, experiments) >>> adata_combined = concatenate_anndata_objects(paths) diff --git a/src/ops_model/features/batch_process_embeddings.py b/src/ops_model/features/batch_process_embeddings.py index 8bc28c6..9ba6865 100644 --- a/src/ops_model/features/batch_process_embeddings.py +++ b/src/ops_model/features/batch_process_embeddings.py @@ -29,10 +29,11 @@ from ops_model.features.processing_common import process_features_csv +from ops_model.paths import BASE_PATH # Base directory for OPS experiments -BASE_DIR = Path("/hpc/projects/icd.fast.ops") +BASE_DIR = Path(f"{BASE_PATH}") def check_csv_exists( diff --git a/src/ops_model/models/interpretability/diffae/classifier/config.py b/src/ops_model/models/interpretability/diffae/classifier/config.py index 9cb91f2..c384c77 100644 --- a/src/ops_model/models/interpretability/diffae/classifier/config.py +++ b/src/ops_model/models/interpretability/diffae/classifier/config.py @@ -7,8 +7,9 @@ import os from dataclasses import dataclass +from ops_model.paths import BASE_PATH -_V4 = "/hpc/projects/icd.fast.ops/models/alex_lin_attention/v4" +_V4 = f"{BASE_PATH}/models/alex_lin_attention/v4" # Phase per-cell ranking exports. v4 = pma_attention-ranked (masked SetTransformer). v5 (paper-v2) = # set-accuracy-`score`-ranked, no-mask/160px, and now includes an NTC group (so NTC anchors come from the @@ -16,7 +17,7 @@ # v4 schema (segmentation_id→segmentation, score→pma_attention, rank_type="top"). Toggle via OPS_DIFFEX_V5=1. # Built + served side-by-side under viewer_assets_v5 so the live v4 viewer is untouched until the final swap. _USE_V5 = os.environ.get("OPS_DIFFEX_V5", "0") == "1" -_V5_RANK = "/hpc/projects/icd.fast.ops/models/diffex/viewer_assets_v5/_rankings" +_V5_RANK = f"{BASE_PATH}/models/diffex/viewer_assets_v5/_rankings" # geneKO: class is `gene`. complex (EBI): class is `predicted_class` (complex name). if _USE_V5: PMA_PHASE_GENEKO = f"{_V5_RANK}/pma_v5_phase_geneKO.parquet" @@ -31,7 +32,7 @@ } # Default output root; per-run results land under ///. -DEFAULT_OUT_ROOT = "/hpc/projects/icd.fast.ops/models/diffex" +DEFAULT_OUT_ROOT = f"{BASE_PATH}/models/diffex" def slugify(name: str) -> str: diff --git a/src/ops_model/models/interpretability/diffae/directions/config.py b/src/ops_model/models/interpretability/diffae/directions/config.py index bbcea42..12a44dc 100644 --- a/src/ops_model/models/interpretability/diffae/directions/config.py +++ b/src/ops_model/models/interpretability/diffae/directions/config.py @@ -10,6 +10,7 @@ from dataclasses import dataclass, field from ..classifier.config import DEFAULT_OUT_ROOT, GRAINS, PMA_PHASE_GENEKO # noqa: F401 +from ops_model.paths import BASE_PATH @dataclass @@ -25,7 +26,7 @@ class DirConfig: # (e.g. "nucleus_NucleoLIVE Live Cell dye"); gather() then pulls that marker's top cells # from fluor_csv and reads the raw `channel` above. None = phase mode (uses the grain parquet). marker_channel: str | None = None - fluor_csv: str = "/hpc/projects/icd.fast.ops/models/alex_lin_attention/v4/pma_fluorescent_cells_all.csv" + fluor_csv: str = f"{BASE_PATH}/models/alex_lin_attention/v4/pma_fluorescent_cells_all.csv" mask_cell: bool = False seed: int = 0 diff --git a/src/ops_model/models/interpretability/diffae/directions/proto_ddim_anchors.py b/src/ops_model/models/interpretability/diffae/directions/proto_ddim_anchors.py index 5b99ee1..34a76d6 100644 --- a/src/ops_model/models/interpretability/diffae/directions/proto_ddim_anchors.py +++ b/src/ops_model/models/interpretability/diffae/directions/proto_ddim_anchors.py @@ -25,10 +25,11 @@ from .config import DirConfig from .rank import supervised_direction from .traverse import _ddim_guided, _sample_guided, load_diffae +from ops_model.paths import BASE_PATH -ANALYSIS = "/hpc/projects/icd.fast.ops/analysis" -DD = "/hpc/projects/icd.fast.ops/models/diffex/diffae" -DIR_CACHE = "/hpc/projects/icd.fast.ops/models/diffex/viewer_assets/_directions" +ANALYSIS = f"{BASE_PATH}/analysis" +DD = f"{BASE_PATH}/models/diffex/diffae" +DIR_CACHE = f"{BASE_PATH}/models/diffex/viewer_assets/_directions" # (label, marker_channel|None, raw_channel, diffae_ckpt, gene). marker_channel=None → phase. FLUOR_SPECS = [ diff --git a/src/ops_model/models/interpretability/diffae/figures/_setacc_common.py b/src/ops_model/models/interpretability/diffae/figures/_setacc_common.py index b7f3560..12d7dbd 100644 --- a/src/ops_model/models/interpretability/diffae/figures/_setacc_common.py +++ b/src/ops_model/models/interpretability/diffae/figures/_setacc_common.py @@ -15,9 +15,10 @@ from ops_model.models.interpretability.diffae.directions.config import DirConfig from ops_model.models.interpretability.diffae.traversal._fluor_topcells import _overlay_rgba from ops_model.models.interpretability.diffae.traversal.build_pc_crops_masked import BASE, CROP_SIZE, _crop, _zarr_patch +from ops_model.paths import BASE_PATH -OUT = "/hpc/projects/icd.fast.ops/analysis/figure4_setacc_panel" -RANK_BASE = "/hpc/projects/icd.fast.ops/models/diffex/viewer_assets_v5/_rankings/fluor_shap" +OUT = f"{BASE_PATH}/analysis/figure4_setacc_panel" +RANK_BASE = f"{BASE_PATH}/models/diffex/viewer_assets_v5/_rankings/fluor_shap" TIM23 = "TIM23 mitochondrial inner membrane pre-sequence translocase complex, TIM17A variant" COPI = "COPI vesicle coat complex, COPG1-COPZ1 variant" diff --git a/src/ops_model/models/interpretability/diffae/figures/_setacc_phase.py b/src/ops_model/models/interpretability/diffae/figures/_setacc_phase.py index 374dbe4..4aa543f 100644 --- a/src/ops_model/models/interpretability/diffae/figures/_setacc_phase.py +++ b/src/ops_model/models/interpretability/diffae/figures/_setacc_phase.py @@ -6,8 +6,9 @@ import pandas as pd from ops_model.models.interpretability.diffae.figures._setacc_common import crop_pick_from_df, tile_at +from ops_model.paths import BASE_PATH -RANKS = "/hpc/projects/icd.fast.ops/models/diffex/viewer_assets_v5/_rankings" +RANKS = f"{BASE_PATH}/models/diffex/viewer_assets_v5/_rankings" PHASE_CH = "Phase2D" # panel-E groups (published: TIMM23/Arp2-3 top, TIPARP/Core Mediator bottom); ko_rank/ntc_rank picked diff --git a/src/ops_model/models/interpretability/diffae/figures/auto_pick_and_plot.py b/src/ops_model/models/interpretability/diffae/figures/auto_pick_and_plot.py index 4f2b704..2053713 100644 --- a/src/ops_model/models/interpretability/diffae/figures/auto_pick_and_plot.py +++ b/src/ops_model/models/interpretability/diffae/figures/auto_pick_and_plot.py @@ -12,8 +12,9 @@ os.environ.setdefault("OPS_DIFFEX_ASSETS", "viewer_assets_v5") from ops_model.models.interpretability.diffae.traversal.morpho_pipeline import MORPHO_TARGETS from ops_model.models.interpretability.diffae.classifier.config import slugify +from ops_model.paths import BASE_PATH -VA = "/hpc/projects/icd.fast.ops/models/diffex/viewer_assets_v5/_morphometrics" +VA = f"{BASE_PATH}/models/diffex/viewer_assets_v5/_morphometrics" BAD = re.compile("moment|hu_|inertia|eigval|intensity|haralick|zernike|glcm|orientation|centroid|_timing") BATCH = ["KIF11_PHASE", "ATP6V1B2_PHASE", "HGS_PHASE", "RRM1_PHASE", "RRN3_PHASE", "SEC61A1_PHASE", diff --git a/src/ops_model/models/interpretability/diffae/figures/figure4_morpho_traversal.py b/src/ops_model/models/interpretability/diffae/figures/figure4_morpho_traversal.py index 5c689a4..419d80a 100644 --- a/src/ops_model/models/interpretability/diffae/figures/figure4_morpho_traversal.py +++ b/src/ops_model/models/interpretability/diffae/figures/figure4_morpho_traversal.py @@ -22,6 +22,7 @@ import numpy as np from matplotlib.colors import Normalize from PIL import Image +from ops_model.paths import BASE_PATH KEYLABEL = {"area": "object area (px²)", "area_filled": "filled area (px²)", "mean_int": "object intensity", "ecc": "eccentricity", "skel": "skeleton length", "circularity": "circularity", @@ -33,8 +34,8 @@ plt.rcParams["font.family"] = "sans-serif" plt.rcParams["font.sans-serif"] = ["Arial", "Helvetica", "DejaVu Sans"] -VA = f"/hpc/projects/icd.fast.ops/models/diffex/{os.environ.get('OPS_DIFFEX_ASSETS', 'viewer_assets')}" -OUT = "/hpc/projects/icd.fast.ops/analysis/figure4_traversals" +VA = f"{BASE_PATH}/models/diffex/{os.environ.get('OPS_DIFFEX_ASSETS', 'viewer_assets')}" +OUT = f"{BASE_PATH}/analysis/figure4_traversals" def _objkey(feature, avail): diff --git a/src/ops_model/models/interpretability/diffae/figures/figure4_morpho_violin.py b/src/ops_model/models/interpretability/diffae/figures/figure4_morpho_violin.py index a569244..7286f39 100644 --- a/src/ops_model/models/interpretability/diffae/figures/figure4_morpho_violin.py +++ b/src/ops_model/models/interpretability/diffae/figures/figure4_morpho_violin.py @@ -19,13 +19,14 @@ from figure4_morpho_traversal import FIGURES, VA, image_panels, render_images from ops_model.models.interpretability.diffae.traversal.morpho_pipeline import MORPHO_TARGETS, real_percell from ops_model.models.interpretability.diffae.classifier.config import slugify +from ops_model.paths import BASE_PATH plt.rcParams["pdf.fonttype"] = 42 plt.rcParams["svg.fonttype"] = "none" plt.rcParams["font.family"] = "sans-serif" plt.rcParams["font.sans-serif"] = ["Arial", "Helvetica", "DejaVu Sans"] -OUT = "/hpc/projects/icd.fast.ops/analysis/figure4_traversals_violin" +OUT = f"{BASE_PATH}/analysis/figure4_traversals_violin" COLORS = {"real": "#999999", "KO": "#2e8b57", "α=0": "#c6dbef", "α=1": "#6baed6", "α=3": "#08519c"} ALPHAS_SHOW = [0, 1, 3] # image panel columns (α=3 = exaggeration, not α=5) CELLS = list(range(30)) # render 30 example cells to pick from diff --git a/src/ops_model/models/interpretability/diffae/figures/figure_ebi_morpho_violin.py b/src/ops_model/models/interpretability/diffae/figures/figure_ebi_morpho_violin.py index 3bbed48..9f2cf36 100644 --- a/src/ops_model/models/interpretability/diffae/figures/figure_ebi_morpho_violin.py +++ b/src/ops_model/models/interpretability/diffae/figures/figure_ebi_morpho_violin.py @@ -35,10 +35,11 @@ from figure_multirank_ebi_grid import (BGX, CACHE, COMBINED_ORDER, FOOT, GAP, OUT, SUP, T, TITLE, block_h, block_w, build_blocks, draw_block, ebi_rows, top_rows, windows) +from ops_model.paths import BASE_PATH # paper-v2 stores first; the v2 dir is fluor-only, so the phase store still comes from the v1 dir (loud). -OPCP_DIRS = ["/hpc/projects/icd.fast.ops/analysis/op_cp_features_paper_v2", - "/hpc/projects/icd.fast.ops/analysis/op_cp_features"] +OPCP_DIRS = [f"{BASE_PATH}/analysis/op_cp_features_paper_v2", + f"{BASE_PATH}/analysis/op_cp_features"] VOUT = f"{OUT}/morpho" plt.rcParams["pdf.fonttype"] = 42 plt.rcParams["svg.fonttype"] = "none" diff --git a/src/ops_model/models/interpretability/diffae/figures/figure_multirank_ebi_grid.py b/src/ops_model/models/interpretability/diffae/figures/figure_multirank_ebi_grid.py index dd75ea5..80b5e68 100644 --- a/src/ops_model/models/interpretability/diffae/figures/figure_multirank_ebi_grid.py +++ b/src/ops_model/models/interpretability/diffae/figures/figure_multirank_ebi_grid.py @@ -24,10 +24,11 @@ import yaml from _setacc_common import CROP_SIZE, _materialize, composite, seg_crop +from ops_model.paths import BASE_PATH -MR = "/hpc/projects/icd.fast.ops/models/alex_lin_attention/v5/multi_rank" -EBI_YAML = "/hpc/projects/icd.fast.ops/configs/gene_clusters/EBI_complexes_v1_old_gene_names.yaml" -OUT = "/hpc/projects/icd.fast.ops/analysis/figure4_multirank_ebi" +MR = f"{BASE_PATH}/models/alex_lin_attention/v5/multi_rank" +EBI_YAML = f"{BASE_PATH}/configs/gene_clusters/EBI_complexes_v1_old_gene_names.yaml" +OUT = f"{BASE_PATH}/analysis/figure4_multirank_ebi" CACHE = f"{OUT}/_cache" plt.rcParams["pdf.fonttype"] = 42 plt.rcParams["svg.fonttype"] = "none" diff --git a/src/ops_model/models/interpretability/diffae/figures/fluor_panel_montages.py b/src/ops_model/models/interpretability/diffae/figures/fluor_panel_montages.py index a24be1c..78d90f2 100644 --- a/src/ops_model/models/interpretability/diffae/figures/fluor_panel_montages.py +++ b/src/ops_model/models/interpretability/diffae/figures/fluor_panel_montages.py @@ -11,8 +11,9 @@ from _setacc_common import GENE_COLS, COMPLEX_COLS, _materialize, slugify from fluor_shap_montages import render_montage # reuses OUT=figure4_shap_montages + seg overlay +from ops_model.paths import BASE_PATH -R = "/hpc/projects/icd.fast.ops/models/diffex/viewer_assets_v5/_rankings/fluor_shap" +R = f"{BASE_PATH}/models/diffex/viewer_assets_v5/_rankings/fluor_shap" COLS = GENE_COLS + COMPLEX_COLS N = 100 diff --git a/src/ops_model/models/interpretability/diffae/figures/fluor_shap_montages.py b/src/ops_model/models/interpretability/diffae/figures/fluor_shap_montages.py index 4aad69a..09eb340 100644 --- a/src/ops_model/models/interpretability/diffae/figures/fluor_shap_montages.py +++ b/src/ops_model/models/interpretability/diffae/figures/fluor_shap_montages.py @@ -17,10 +17,11 @@ from ops_model.models.interpretability.diffae.classifier.config import slugify from _setacc_common import _materialize, seg_crop, composite, CROP_SIZE +from ops_model.paths import BASE_PATH -MR = "/hpc/projects/icd.fast.ops/models/alex_lin_attention/v5/multi_rank/shap_screen/shap_screen_fluor_all.compact.parquet" -OUT = "/hpc/projects/icd.fast.ops/analysis/figure4_shap_montages" -PQ = "/hpc/projects/icd.fast.ops/models/diffex/viewer_assets_v5/_rankings/fluor_multirank/geneKO" +MR = f"{BASE_PATH}/models/alex_lin_attention/v5/multi_rank/shap_screen/shap_screen_fluor_all.compact.parquet" +OUT = f"{BASE_PATH}/analysis/figure4_shap_montages" +PQ = f"{BASE_PATH}/models/diffex/viewer_assets_v5/_rankings/fluor_multirank/geneKO" plt.rcParams["pdf.fonttype"] = 42 # fig-4 fluor groups: multi_rank channel_name -> (gene, zarr channel) diff --git a/src/ops_model/models/interpretability/diffae/figures/nc_ratio.py b/src/ops_model/models/interpretability/diffae/figures/nc_ratio.py index 649b585..4611185 100644 --- a/src/ops_model/models/interpretability/diffae/figures/nc_ratio.py +++ b/src/ops_model/models/interpretability/diffae/figures/nc_ratio.py @@ -26,12 +26,13 @@ from skimage.morphology import binary_closing, disk from _setacc_common import _materialize +from ops_model.paths import BASE_PATH plt.rcParams["pdf.fonttype"] = 42 plt.rcParams["svg.fonttype"] = "none" -VA = "/hpc/projects/icd.fast.ops/models/diffex/viewer_assets_v5" -OUT = "/hpc/projects/icd.fast.ops/analysis/figure4_nc_ratio" +VA = f"{BASE_PATH}/models/diffex/viewer_assets_v5" +OUT = f"{BASE_PATH}/analysis/figure4_nc_ratio" MARKER_DIR = "proteasome_PSMB7" MC = "proteasome_PSMB7" # marker_channel (DirConfig); channel to READ is passed separately TARGET = "PSMB6" diff --git a/src/ops_model/models/interpretability/diffae/figures/ntc_anchor_compare.py b/src/ops_model/models/interpretability/diffae/figures/ntc_anchor_compare.py index b3a4727..2018e1b 100644 --- a/src/ops_model/models/interpretability/diffae/figures/ntc_anchor_compare.py +++ b/src/ops_model/models/interpretability/diffae/figures/ntc_anchor_compare.py @@ -7,8 +7,9 @@ import matplotlib.pyplot as plt import zarr from ..viewer.build_pc_crops_masked import BASE, CROP_SIZE, PHASE_CHANNEL, _crop, _render_gray, _zarr_patch +from ops_model.paths import BASE_PATH -R = "/hpc/projects/icd.fast.ops/models/diffex/viewer_assets_v5/_rankings" +R = f"{BASE_PATH}/models/diffex/viewer_assets_v5/_rankings" N = 50 @@ -58,6 +59,6 @@ def panel(ax_grid, imgs, title): axs = [fig.add_subplot(inner[j]) for j in range(N)] panel(axs, imgs, title) fig.text(0.5, 0.905 - row * 0.485, title, ha="center", fontsize=15, fontweight="bold") -out = "/hpc/projects/icd.fast.ops/analysis/ntc_anchor_old_vs_multirank.png" +out = f"{BASE_PATH}/analysis/ntc_anchor_old_vs_multirank.png" fig.savefig(out, dpi=110, bbox_inches="tight"); plt.close(fig) print("wrote", out) diff --git a/src/ops_model/models/interpretability/diffae/figures/phase_multibag_montages.py b/src/ops_model/models/interpretability/diffae/figures/phase_multibag_montages.py index e7fa696..79f066b 100644 --- a/src/ops_model/models/interpretability/diffae/figures/phase_multibag_montages.py +++ b/src/ops_model/models/interpretability/diffae/figures/phase_multibag_montages.py @@ -10,9 +10,10 @@ from _setacc_common import _materialize import debug_setacc_top100 as D +from ops_model.paths import BASE_PATH -OUT_DIR = "/hpc/projects/icd.fast.ops/analysis/figure4_shap_montages" # shared SHAP-montage review dir (phase + fluor) -RANK = "/hpc/projects/icd.fast.ops/models/diffex/viewer_assets_v5/_rankings/pma_shap_phase_geneKO.parquet" +OUT_DIR = f"{BASE_PATH}/analysis/figure4_shap_montages" # shared SHAP-montage review dir (phase + fluor) +RANK = f"{BASE_PATH}/models/diffex/viewer_assets_v5/_rankings/pma_shap_phase_geneKO.parquet" PHASE_CH = "Phase2D" N = 100 diff --git a/src/ops_model/models/interpretability/diffae/figures/setacc/figure4_setacc_panel_newpheno.py b/src/ops_model/models/interpretability/diffae/figures/setacc/figure4_setacc_panel_newpheno.py index d4ab71d..f7ba323 100644 --- a/src/ops_model/models/interpretability/diffae/figures/setacc/figure4_setacc_panel_newpheno.py +++ b/src/ops_model/models/interpretability/diffae/figures/setacc/figure4_setacc_panel_newpheno.py @@ -11,8 +11,9 @@ from ops_model.models.interpretability.diffae.figures._setacc_common import crop_pick_from_df, tile_at from ops_model.models.interpretability.diffae.figures.setacc.figure4_setacc_panel import make_panel +from ops_model.paths import BASE_PATH -RANK = "/hpc/projects/icd.fast.ops/models/diffex/viewer_assets_v5/_rankings/pma_shap_phase_geneKO.parquet" +RANK = f"{BASE_PATH}/models/diffex/viewer_assets_v5/_rankings/pma_shap_phase_geneKO.parquet" PHASE_CH = "Phase2D" COLS = [ # KO rank = montage pick; NTC rank distinct per column (1-5) diff --git a/src/ops_model/models/interpretability/diffae/figures/virtual_staining_schematic.py b/src/ops_model/models/interpretability/diffae/figures/virtual_staining_schematic.py index ea3c1fe..8a90a16 100644 --- a/src/ops_model/models/interpretability/diffae/figures/virtual_staining_schematic.py +++ b/src/ops_model/models/interpretability/diffae/figures/virtual_staining_schematic.py @@ -19,6 +19,7 @@ import matplotlib.pyplot as plt import numpy as np from matplotlib.patches import FancyArrowPatch, FancyBboxPatch, Rectangle, Circle, Ellipse +from ops_model.paths import BASE_PATH plt.rcParams["pdf.fonttype"] = 42 plt.rcParams["svg.fonttype"] = "none" @@ -103,7 +104,7 @@ def _patch_cell(ax, cx, cy, s, edge=INK): ax.plot([cx + g, cx + g], [cy - s, cy + s], color=PURPLE, lw=0.5, alpha=0.55, zorder=7) -def build(outstem="/hpc/projects/icd.fast.ops/analysis/figure4_schematic/virtual_staining_schematic"): +def build(outstem=f"{BASE_PATH}/analysis/figure4_schematic/virtual_staining_schematic"): fig, ax = plt.subplots(figsize=(14.6, 4.9)) ax.set_xlim(0, 14.6); ax.set_ylim(0, 4.9); ax.axis("off") ax.text(0.12, 4.72, "C", fontsize=28, fontweight="bold", va="top") diff --git a/src/ops_model/models/interpretability/diffae/generator/plot_metrics.py b/src/ops_model/models/interpretability/diffae/generator/plot_metrics.py index bf7e10f..84c87ce 100644 --- a/src/ops_model/models/interpretability/diffae/generator/plot_metrics.py +++ b/src/ops_model/models/interpretability/diffae/generator/plot_metrics.py @@ -8,11 +8,12 @@ import matplotlib.pyplot as plt import numpy as np import torch +from ops_model.paths import BASE_PATH plt.rcParams["pdf.fonttype"] = 42 -DD = "/hpc/projects/icd.fast.ops/models/diffex/diffae" -OUT = "/hpc/projects/icd.fast.ops/models/diffex/model_metrics" +DD = f"{BASE_PATH}/models/diffex/diffae" +OUT = f"{BASE_PATH}/models/diffex/model_metrics" def collect(): diff --git a/src/ops_model/models/interpretability/diffae/generator/virtstain_multi.py b/src/ops_model/models/interpretability/diffae/generator/virtstain_multi.py index f21dee5..a681915 100644 --- a/src/ops_model/models/interpretability/diffae/generator/virtstain_multi.py +++ b/src/ops_model/models/interpretability/diffae/generator/virtstain_multi.py @@ -24,8 +24,9 @@ from .data import load_diffae_crops from .model import DiffAE from .train import _pearson +from ops_model.paths import BASE_PATH -OUT = "/hpc/projects/icd.fast.ops/analysis/virtual_staining/multi_marker" +OUT = f"{BASE_PATH}/analysis/virtual_staining/multi_marker" def markers_list(): @@ -46,7 +47,7 @@ def markers_list(): # Every scored cell per marker (concat of Alex's 6 per-cell shards, incl NTC) — NOT the acc>0.5 qualifying # subset. Virtual staining needs a representative pool of paired (phase, marker) crops, so no distinctiveness # filtering: use all cells (74k–1.3M/marker) and sample ≤cap per marker. Built by scratchpad/build_allcells.py. -ALL_CELLS = ("/hpc/projects/icd.fast.ops/models/alex_lin_attention/v5/fluorescence/" +ALL_CELLS = (f"{BASE_PATH}/models/alex_lin_attention/v5/fluorescence/" "misc/all_cells_bychannel.parquet") diff --git a/src/ops_model/models/interpretability/diffae/traversal/build_pc_crops_masked.py b/src/ops_model/models/interpretability/diffae/traversal/build_pc_crops_masked.py index 36b928e..e291066 100644 --- a/src/ops_model/models/interpretability/diffae/traversal/build_pc_crops_masked.py +++ b/src/ops_model/models/interpretability/diffae/traversal/build_pc_crops_masked.py @@ -22,8 +22,9 @@ import numpy as np from . import catalog as C +from ops_model.paths import BASE_PATH -BASE = "/hpc/projects/icd.fast.ops" +BASE = f"{BASE_PATH}" PCS_OUT = f"{C.OUT}/viewer_assets/pcs" CROP_SIZE = 160 # native px re-crop (was 96); crisper at the 150px display + shows surround PHASE_CHANNEL = 0 diff --git a/src/ops_model/models/interpretability/diffae/traversal/catalog.py b/src/ops_model/models/interpretability/diffae/traversal/catalog.py index 820c906..eb803b1 100644 --- a/src/ops_model/models/interpretability/diffae/traversal/catalog.py +++ b/src/ops_model/models/interpretability/diffae/traversal/catalog.py @@ -11,17 +11,18 @@ import pandas as pd from ..classifier.config import slugify +from ops_model.paths import BASE_PATH -OUT = "/hpc/projects/icd.fast.ops/models/diffex" +OUT = f"{BASE_PATH}/models/diffex" DD = f"{OUT}/diffae" -_DIST_BASE = ("/hpc/projects/icd.fast.ops/organelle_attribution/pca_optimized_v0.3/cell_dino/" +_DIST_BASE = (f"{BASE_PATH}/organelle_attribution/pca_optimized_v0.3/cell_dino/" "zscore_per_exp/paper_v2/with_cp/with_4i") _DIST_RELS = ["all_livecell"] # v2 with_cp/with_4i: single 56-reporter matrix (live + CP + 4i) LAUNCH_JSON = f"{OUT}/directions/_ranking/fluor_marker_launch.json" -GENE_PANEL = "/hpc/projects/icd.fast.ops/configs/annotated_gene_panel_July2025.csv" -EBI_YAML = "/hpc/projects/icd.fast.ops/configs/gene_clusters/EBI_complexes_v1_updated_gene_names.yaml" -EBI_FLUOR_CSV = "/hpc/projects/icd.fast.ops/models/alex_lin_attention/v4/pma_fluorescent_cells_ebi_all.csv" -GENE_EMB_H5AD = ("/hpc/projects/icd.fast.ops/organelle_attribution/pca_optimized_v0.3/cell_dino/" +GENE_PANEL = f"{BASE_PATH}/configs/annotated_gene_panel_July2025.csv" +EBI_YAML = f"{BASE_PATH}/configs/gene_clusters/EBI_complexes_v1_updated_gene_names.yaml" +EBI_FLUOR_CSV = f"{BASE_PATH}/models/alex_lin_attention/v4/pma_fluorescent_cells_ebi_all.csv" +GENE_EMB_H5AD = (f"{BASE_PATH}/organelle_attribution/pca_optimized_v0.3/cell_dino/" "zscore_per_exp/paper_v2/phase_only/fixed_80%/cosine/gene_embedding_pca_optimized.h5ad") # fixed-cell reporter (distinctiveness matrix col) for CP/4i marker_channels diff --git a/src/ops_model/models/interpretability/diffae/traversal/morpho_pipeline.py b/src/ops_model/models/interpretability/diffae/traversal/morpho_pipeline.py index e5addf4..ef4d555 100644 --- a/src/ops_model/models/interpretability/diffae/traversal/morpho_pipeline.py +++ b/src/ops_model/models/interpretability/diffae/traversal/morpho_pipeline.py @@ -14,9 +14,10 @@ import os import numpy as np +from ops_model.paths import BASE_PATH -CACHE = f"/hpc/projects/icd.fast.ops/models/diffex/{os.environ.get('OPS_DIFFEX_ASSETS', 'viewer_assets')}" -SYNTH_BASE = "/hpc/projects/icd.fast.ops/models/diffex/morpho_synth" +CACHE = f"{BASE_PATH}/models/diffex/{os.environ.get('OPS_DIFFEX_ASSETS', 'viewer_assets')}" +SYNTH_BASE = f"{BASE_PATH}/models/diffex/morpho_synth" PAD = 24 CROP = 256 GEN_CROP = 160 # DiffEx cfg.crop_size: the generated crops are 160 px (native) upsized to 256 — real ref crops must match this window @@ -106,7 +107,7 @@ def _vs_h2b_nucleus_npz(marker_dir, target, grain, out_npz, n_cells, force=False from diffusers import DDIMScheduler from cellpose import models as cpm from skimage.transform import resize - VOUT = "/hpc/projects/icd.fast.ops/analysis/virtual_staining/multi_marker" + VOUT = f"{BASE_PATH}/analysis/virtual_staining/multi_marker" dev = torch.device("cuda") markers = json.load(open(f"{VOUT}/markers.json")); h2b = markers.index("chromatin_H2BC21") cfg = DiffAEConfig(spatial_cond=True, n_markers=len(markers), device="cuda"); Hg = cfg.crop_size @@ -345,7 +346,7 @@ def run_seg(real_exp, marker_channel, n_alpha, structure_type=None, base_dir=SYN return out -OPCP = "/hpc/projects/icd.fast.ops/analysis/op_cp_features/op_cp_features_{store}.h5ad" +OPCP = f"{BASE_PATH}/analysis/op_cp_features/op_cp_features_{store}.h5ad" REF_SUFFIX = {"count": "count", "total_area": "area_sum", "mean_area": "area_mean", "mean_int": "intensity_mean_mean", "mean_ecc": "eccentricity_mean"} @@ -520,7 +521,7 @@ def reference_from_direction(marker_dir, target, real_exp, marker_channel, struc from PIL import Image from skimage.measure import label as relabel, regionprops from skimage.transform import resize as skresize - npz = glob.glob(f"/hpc/projects/icd.fast.ops/models/diffex/directions/{marker_dir}/{grain}/{target}/cache/crops_{target}_*.npz") + npz = glob.glob(f"{BASE_PATH}/models/diffex/directions/{marker_dir}/{grain}/{target}/cache/crops_{target}_*.npz") if not npz: print(f"[ref-dir] no crops npz for {marker_dir}/{target}"); return d = np.load(npz[0], allow_pickle=True) @@ -597,7 +598,7 @@ def reference_from_store(marker_dir, target, grain, image_channel, org_label, n_ from ..classifier.config import GRAINS from skimage.morphology import skeletonize GK = GRAINS["geneKO"]["parquet"] # attention parquet (has NTC + coords) - ACC = "/hpc/projects/icd.fast.ops/models/diffex/accuracy_ranking/phase_geneKO_topacc_ALL_top1000.parquet" + ACC = f"{BASE_PATH}/models/diffex/accuracy_ranking/phase_geneKO_topacc_ALL_top1000.parquet" COLS = ["gene", "experiment", "well", "segmentation", "x_pheno", "y_pheno", "rank", "rank_type"] def _cells(pq, genes, n): # top-n ranked 'top' cells for gene(s) @@ -609,7 +610,7 @@ def _cells(pq, genes, n): # top- if grain == "complex": import yaml from ..classifier.config import slugify - y = yaml.safe_load(open("/hpc/projects/icd.fast.ops/configs/gene_clusters/EBI_complexes_v1_updated_gene_names.yaml")) + y = yaml.safe_load(open(f"{BASE_PATH}/configs/gene_clusters/EBI_complexes_v1_updated_gene_names.yaml")) members = next((e["genes"] for e in y.values() if slugify(e["name"]) == target), []) ko = _cells(ACC, members, n_cells) else: @@ -626,7 +627,7 @@ def _cells(pq, genes, n): # top- for c, (_, row) in enumerate(df.iterrows()): exp = str(row["experiment"]); w = str(row["well"]).strip() m = re.match(r"^([A-Za-z]+)(\d+)$", w); pos = w if w.count("/") == 2 else (f"{m.group(1)}/{m.group(2)}/0" if m else w) - zp = f"/hpc/projects/icd.fast.ops/{exp}/3-assembly/phenotyping_v3.zarr" + zp = f"{BASE_PATH}/{exp}/3-assembly/phenotyping_v3.zarr" if not os.path.exists(zp): continue try: @@ -664,7 +665,7 @@ def _cells(pq, genes, n): # top- def sweep_grid(marker_dir, target, real_exp, marker_channel, pix=(0.065, 0.13, 0.185, 0.37), - thr=(0.1, 0.5, 1.0), ai=8, out="/hpc/projects/icd.fast.ops/models/diffex/morpho_grid_sweep.png"): + thr=(0.1, 0.5, 1.0), ai=8, out=f"{BASE_PATH}/models/diffex/morpho_grid_sweep.png"): """Sweep the two frangi knobs — pixel_size_um (rows) × threshold_mult (cols) — on one generated crop, render seg boundaries + object counts, so we can pick the cleanest (signal, no noise). No postprocess.""" import matplotlib @@ -900,7 +901,7 @@ def _topn(idx, pq, val, n): ko_val = target if grain == "complex": import yaml - y = yaml.safe_load(open("/hpc/projects/icd.fast.ops/configs/gene_clusters/EBI_complexes_v1_updated_gene_names.yaml")) + y = yaml.safe_load(open(f"{BASE_PATH}/configs/gene_clusters/EBI_complexes_v1_updated_gene_names.yaml")) ent = next((e for e in y.values() if slugify(e["name"]) == target), None) members = ent["genes"] if ent else [] ko_val = ent["name"] if ent else target # complex rank parquet keys by full complex name @@ -909,7 +910,7 @@ def _topn(idx, pq, val, n): i_ko = np.where(gn == target)[0] if store_marker == "phase": # phase: accuracy (KO) + top-attention (NTC) - ko_pq = "/hpc/projects/icd.fast.ops/models/diffex/accuracy_ranking/phase_geneKO_topacc_ALL_top1000.parquet" + ko_pq = f"{BASE_PATH}/models/diffex/accuracy_ranking/phase_geneKO_topacc_ALL_top1000.parquet" ntc_pq = GRAINS["geneKO"]["parquet"]; ko_val = target else: # fluor: the marker's set-accuracy rank parquet ko_pq = ntc_pq = f"{CACHE}/_rankings/fluor/{'complex' if grain == 'complex' else 'geneKO'}/{marker_dir}.parquet" @@ -989,7 +990,7 @@ def msem(v): rng = np.random.default_rng(0) if grain == "complex": import yaml - y = yaml.safe_load(open("/hpc/projects/icd.fast.ops/configs/gene_clusters/EBI_complexes_v1_updated_gene_names.yaml")) + y = yaml.safe_load(open(f"{BASE_PATH}/configs/gene_clusters/EBI_complexes_v1_updated_gene_names.yaml")) members = next((e["genes"] for e in y.values() if slugify(e["name"]) == target), []) i_ko_all = np.where(np.isin(gn, members))[0] else: diff --git a/src/ops_model/paths.py b/src/ops_model/paths.py new file mode 100644 index 0000000..398f26c --- /dev/null +++ b/src/ops_model/paths.py @@ -0,0 +1,9 @@ +"""Central base paths for ops_model data/model/analysis roots. + +Every hardcoded storage location derives from ``BASE_PATH``; override the whole +tree with the ``OPS_BASE_PATH`` env var. The default is the current shared store. +""" +import os +from pathlib import Path + +BASE_PATH = os.environ.get("OPS_BASE_PATH", "/hpc/projects/icd.fast.ops")