diff --git a/CHANGELOG.md b/CHANGELOG.md index 3563251..3a5d0e4 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -1,6 +1,32 @@ # Changelog -## Unreleased +## v0.8.0 — unreleased (the vertex analyses move to a plugin) + +### Behaviour changes (read these before upgrading) + +- **The vertex analyses are no longer part of source-analytics.** vertex_cluster, + vertex_connectivity, vertex_cross_freq, vertex_directed, vertex_evoked, vertex_graph, + vertex_nbs, vertex_network, vertex_signature, vertex_spatial, vertex_specparam, and + fcd_comparison (which reads vertex_connectivity's output) moved to the private + `source-analytics-vertex` package. So did `spectral.vertex`, `spectral.vertex_aperiodic`, + the five `R/vertex_*.R` scripts, the vertex figure-registry schemas and the glass-brain + summary figures. A config or `--analysis` that names one of them now fails with a + message naming the plugin. To keep running them, install the plugin, or pin v0.7.1. + Vertex maps depend on where one set of sources sits. The ROI analyses on Monte Carlo + source operators replace them. +- `source_analytics.analyses` no longer exports the vertex classes or their old aliases + (`WholebrainAnalysis`, `MVPAAnalysis`, `VertexMVPAAnalysis`, `SpecparamVertexAnalysis`, + `SpatialLMMAnalysis`). The plugin exports them. + +### Added + +- **Analysis plugins** (`source_analytics.plugins`). A package adds analyses through the + `source_analytics.plugins` entry-point group. It provides `ANALYSES`, `METADATA` and + `ALIASES`, and optionally `register_figures(registry)`. `source-analytics run`/`list`/ + `figure` and `analysis_meta()` pick plugin analyses up. A plugin that fails to import + is logged and skipped. One that reuses an existing analysis name raises. + +### Fixed - **nibabel is a core dependency.** `viz/__init__` imports `viz.brain_roi`, which imports nibabel at module level. Without the `atlas` extra, `import @@ -8,6 +34,15 @@ package imports without any extra. The ROI modules' atlas readers need it anyway. The `atlas` extra is kept, so existing install commands still work. +### Unchanged + +- Kept in core because core modules use them: `spectral.vertex_connectivity` + (electrode_connectivity's kernels), `viz.glass_brain`, `analyses._network_base`, the + cluster-permutation statistics, and the `BaseAnalysis` helpers the vertex modules call + (`_vertex_epoch_config`, `_label_vertex_regions`, cluster-state persistence). +- electrode_signature still compares against a `vertex_signature` table in its own + paradigm, if the plugin wrote one. + ## v0.7.1 — 2026-09-11 (R step failures fail the run; PAC mosaics) ### Behaviour changes (read these before re-running a study) diff --git a/R/vertex_cluster_analysis.R b/R/vertex_cluster_analysis.R deleted file mode 100644 index 0ff5a2f..0000000 --- a/R/vertex_cluster_analysis.R +++ /dev/null @@ -1,252 +0,0 @@ -#!/usr/bin/env Rscript -# vertex_cluster_analysis.R — Summary report generation for vertex-level cluster analysis -# -# Called by Python: Rscript R/vertex_cluster_analysis.R --data-dir ... --config ... --output-dir ... -# -# Reads pre-computed CSV tables from Python (voxelwise_stats.csv, cluster_results.csv, -# vertex_cluster_values.csv, vertex_cluster_features.csv) and generates formatted summary -# tables + ANALYSIS_SUMMARY.md. -# -# Statistics and figures are handled in Python (cluster permutation + glass brain), -# so this script focuses on report generation. - -library(argparse) -library(yaml) -library(readr) - -# Resolve script directory for sourcing helpers -script_dir <- if (exists("script.dir")) { - script.dir -} else { - tryCatch({ - args <- commandArgs(trailingOnly = FALSE) - file_arg <- grep("^--file=", args, value = TRUE) - if (length(file_arg) > 0) { - dirname(normalizePath(sub("^--file=", "", file_arg))) - } else { - "R" - } - }, error = function(e) "R") -} - -# Source report utilities if available -report_file <- file.path(script_dir, "report.R") -if (file.exists(report_file)) { - source(report_file) -} - -# --- Argument parsing --- -parser <- ArgumentParser(description = "Vertex cluster analysis report (R)") -parser$add_argument("--data-dir", required = TRUE, - help = "Directory containing vertex cluster CSVs") -parser$add_argument("--config", required = TRUE, - help = "Path to study YAML config") -parser$add_argument("--output-dir", required = TRUE, - help = "Root output directory") -parser$add_argument("--fig-dir", default = NULL, - help = "Directory for figures (default: output-dir/figures)") -parser$add_argument("--tbl-dir", default = NULL, - help = "Directory for tables (default: output-dir/tables)") -parser$add_argument("--no-figures", action = "store_true", default = FALSE, - help = "Skip all figure generation (stats/tables only)") -args <- parser$parse_args() - -data_dir <- args$data_dir -config_path <- args$config -output_dir <- args$output_dir -no_figures <- args$no_figures - -if (no_figures) { - ggsave <- function(...) invisible(NULL) -} - -tbl_dir <- if (!is.null(args$tbl_dir)) args$tbl_dir else file.path(output_dir, "tables") -fig_dir <- if (!is.null(args$fig_dir)) args$fig_dir else file.path(output_dir, "figures") -dir.create(tbl_dir, showWarnings = FALSE, recursive = TRUE) - -# --- Load config --- -config <- read_yaml(config_path) -group_labels <- unlist(config$groups) -group_order <- config$group_order -wb_config <- config$vertex - -message("Study: ", config$name) -message("Vertex cluster report generation") - -# --- Load data --- -cluster_file <- file.path(tbl_dir, "cluster_results.csv") -voxelwise_file <- file.path(tbl_dir, "voxelwise_stats.csv") -values_file <- file.path(data_dir, "vertex_cluster_values.csv") -features_file <- file.path(data_dir, "vertex_cluster_features.csv") - -has_clusters <- file.exists(cluster_file) -has_voxelwise <- file.exists(voxelwise_file) -has_values <- file.exists(values_file) -has_features <- file.exists(features_file) - -if (has_clusters) { - cluster_df <- read_csv(cluster_file, show_col_types = FALSE) - message(" cluster_results.csv: ", nrow(cluster_df), " rows") -} - -if (has_voxelwise) { - voxelwise_df <- read_csv(voxelwise_file, show_col_types = FALSE) - message(" voxelwise_stats.csv: ", nrow(voxelwise_df), " rows") -} - -if (has_values) { - values_df <- read_csv(values_file, show_col_types = FALSE) - message(" vertex_cluster_values.csv: ", nrow(values_df), " rows") - n_subjects <- length(unique(values_df$subject)) - n_vertices <- length(unique(values_df$vertex_idx)) - n_bands <- length(unique(values_df$band)) - message(" Subjects: ", n_subjects, ", Vertices: ", n_vertices, ", Bands: ", n_bands) -} - -if (has_features) { - features_df <- read_csv(features_file, show_col_types = FALSE) - message(" vertex_cluster_features.csv: ", nrow(features_df), " rows") -} - -# --- Compute effect size summaries --- -effect_summary <- NULL -if (has_voxelwise) { - effect_summary <- voxelwise_df %>% - dplyr::group_by(contrast, band, metric) %>% - dplyr::summarise( - mean_t = mean(t, na.rm = TRUE), - max_abs_t = max(abs(t), na.rm = TRUE), - mean_hedges_g = mean(hedges_g, na.rm = TRUE), - max_abs_hedges_g = max(abs(hedges_g), na.rm = TRUE), - n_nominal_sig = sum(p < 0.05, na.rm = TRUE), - .groups = "drop" - ) - - write_csv(effect_summary, file.path(tbl_dir, "effect_size_summary.csv")) - message("Exported effect_size_summary.csv") -} - -# --- Generate ANALYSIS_SUMMARY.md --- -n_perms <- ifelse(is.null(wb_config$n_permutations), 1000, wb_config$n_permutations) -cluster_thresh <- ifelse(is.null(wb_config$cluster_threshold), 2.0, wb_config$cluster_threshold) -adj_dist <- ifelse(is.null(wb_config$adjacency_distance_mm), 5.0, wb_config$adjacency_distance_mm) - -md_lines <- c( - "# Vertex Cluster Analysis Summary", - "", - paste0("**Study**: ", config$name), - paste0("**Analysis**: Vertex-level spectral analysis with cluster permutation testing"), - paste0("**Date**: ", Sys.Date()), - "" -) - -if (has_values) { - md_lines <- c(md_lines, - "## Data Overview", - "", - paste0("- **Subjects**: ", n_subjects), - paste0("- **Vertices**: ", n_vertices), - paste0("- **Frequency bands**: ", n_bands, " (", paste(unique(values_df$band), collapse = ", "), ")"), - paste0("- **Groups**: ", paste(sapply(group_order, function(g) { - n <- sum(values_df$group == g) / (n_vertices * n_bands) - paste0(group_labels[g], " (n=", n, ")") - }), collapse = ", ")), - "" - ) -} - -md_lines <- c(md_lines, - "## Methods", - "", - "Power spectral density (PSD) was computed for each source vertex using Welch's", - "method (2-second Hann windows, 50% overlap). The following metrics were extracted:", - "", - "- **Relative band power**: Band power / total power per vertex", - "- **Absolute band power (dB)**: 10 * log10(band power)", - "- **fALFF**: High-gamma (65-100 Hz) / total (1-100 Hz) power ratio", - "- **Spectral slope**: 1/f exponent via log-log regression (2-50 Hz)", - "- **Peak alpha frequency**: Argmax in 6-13 Hz range", - "", - "Group differences were tested using independent-samples Welch's t-tests at each", - paste0("vertex, with cluster-based permutation correction (", n_perms, " permutations,"), - paste0("t-threshold = ", cluster_thresh, ", adjacency distance = ", adj_dist, " mm;"), - "Maris & Oostenveld, 2007).", - "" -) - -if (has_clusters) { - sig_clusters <- cluster_df[cluster_df$p_corrected < 0.05, ] - - md_lines <- c(md_lines, - "## Results", - "" - ) - - if (nrow(sig_clusters) > 0) { - md_lines <- c(md_lines, - paste0("### Significant Clusters (", nrow(sig_clusters), " at p < 0.05)"), - "", - "| Contrast | Band/Metric | Type | Vertices | Peak t | Cluster stat | p_corrected |", - "|----------|-------------|------|----------|--------|--------------|-------------|" - ) - for (i in seq_len(nrow(sig_clusters))) { - row <- sig_clusters[i, ] - md_lines <- c(md_lines, sprintf( - "| %s | %s | %s | %d | %.2f | %.2f | %.4f |", - row$contrast, row$band, row$metric, - row$n_vertices, row$peak_t, row$cluster_stat, row$p_corrected - )) - } - md_lines <- c(md_lines, "") - } else { - md_lines <- c(md_lines, - "No significant clusters at p < 0.05.", - "" - ) - } - - md_lines <- c(md_lines, - paste0("Total clusters tested: ", nrow(cluster_df)), - "" - ) -} - -if (!is.null(effect_summary)) { - md_lines <- c(md_lines, - "### Effect Size Summary (across all vertices)", - "", - "| Contrast | Band | Metric | Mean |t| | Max |t| | Mean |g| | Max |g| | N sig (uncorr) |", - "|----------|------|--------|---------|---------|---------|---------|----------------|" - ) - for (i in seq_len(nrow(effect_summary))) { - row <- effect_summary[i, ] - md_lines <- c(md_lines, sprintf( - "| %s | %s | %s | %.2f | %.2f | %.2f | %.2f | %d |", - row$contrast, row$band, row$metric, - abs(row$mean_t), row$max_abs_t, - abs(row$mean_hedges_g), row$max_abs_hedges_g, - row$n_nominal_sig - )) - } - md_lines <- c(md_lines, "") -} - -md_lines <- c(md_lines, - "## Output Files", - "", - "- `data/vertex_cluster_values.csv` -- per-subject per-vertex band power (long format)", - "- `data/vertex_cluster_features.csv` -- per-subject per-vertex fALFF, spectral slope, peak alpha", - "- `data/source_coords.csv` -- vertex coordinates in mm", - "- `tables/voxelwise_stats.csv` -- per-vertex t-statistics, p-values, Hedges' g", - "- `tables/cluster_results.csv` -- cluster summaries with permutation-corrected p-values", - "- `tables/effect_size_summary.csv` -- aggregated effect sizes across vertices", - "- `figures/vertex_cluster_*.png` -- glass brain visualizations per band/metric", - "- `figures/vertex_cluster_summary.png` -- summary figure across all metrics", - "" -) - -summary_path <- file.path(output_dir, "ANALYSIS_SUMMARY.md") -writeLines(md_lines, summary_path) -message("Wrote ", summary_path) - -message("Done.") diff --git a/R/vertex_connectivity_analysis.R b/R/vertex_connectivity_analysis.R deleted file mode 100644 index 0b65b7f..0000000 --- a/R/vertex_connectivity_analysis.R +++ /dev/null @@ -1,139 +0,0 @@ -#!/usr/bin/env Rscript -# Vertex Connectivity Analysis Report Generator -# Reads pre-computed FCD data and statistics, generates ANALYSIS_SUMMARY.md - -suppressPackageStartupMessages({ - library(optparse) - library(yaml) -}) - -option_list <- list( - make_option("--data-dir", type = "character", help = "Path to data/ directory"), - make_option("--config", type = "character", help = "Path to study_config.yaml"), - make_option("--output-dir", type = "character", help = "Path to output directory"), - make_option("--no-figures", action = "store_true", default = FALSE, - help = "Skip all figure generation") -) -opts <- parse_args(OptionParser(option_list = option_list)) - -no_figures <- isTRUE(opts[["no-figures"]]) - -if (no_figures) { - ggsave <- function(...) invisible(NULL) -} - -data_dir <- opts[["data-dir"]] -config_path <- opts[["config"]] -output_dir <- opts[["output-dir"]] - -config <- read_yaml(config_path) - -# --- Load data ---------------------------------------------------------------- -fcd_path <- file.path(data_dir, "vertex_fcd.csv") -stats_path <- file.path(output_dir, "tables", "vertex_connectivity_stats.csv") - -if (!file.exists(fcd_path)) { - cat("No vertex_fcd.csv found.\n") - quit(status = 0) -} - -fcd <- read.csv(fcd_path, stringsAsFactors = FALSE) - -# --- Config ------------------------------------------------------------------- -vc_cfg <- config$vertex_connectivity %||% list() -metric <- vc_cfg$metric %||% "imag_coherence" -fcd_threshold <- vc_cfg$fcd_threshold %||% 0.05 -n_perm <- vc_cfg$n_permutations %||% 1000 - -wb_cfg <- config$vertex %||% list() -adj_dist <- wb_cfg$adjacency_distance_mm %||% 5.0 - -# --- Summaries ---------------------------------------------------------------- -n_subjects <- length(unique(fcd$subject)) -n_vertices <- length(unique(fcd$vertex_idx)) -bands <- unique(fcd$band) -groups <- unique(fcd$group) - -# Per-band FCD summary -band_summary <- do.call(rbind, lapply(bands, function(b) { - sub <- fcd[fcd$band == b, ] - do.call(rbind, lapply(groups, function(g) { - g_sub <- sub[sub$group == g, ] - data.frame( - Band = b, - Group = g, - Mean_FCD = round(mean(g_sub$fcd), 4), - SD_FCD = round(sd(g_sub$fcd), 4), - stringsAsFactors = FALSE - ) - })) -})) - -# --- Write ANALYSIS_SUMMARY.md ----------------------------------------------- -lines <- c( - "# Vertex Connectivity Analysis Summary", - "", - sprintf("**Study**: %s", config$name), - sprintf("**Analysis**: All-to-all vertex connectivity (%s) + FCD", metric), - sprintf("**FCD threshold**: %.3f", fcd_threshold), - sprintf("**Permutations**: %d", n_perm), - sprintf("**Subjects**: %d (%s)", n_subjects, paste(groups, collapse = ", ")), - sprintf("**Vertices**: %d", n_vertices), - "" -) - -# Epoch info -epoch_cfg <- wb_cfg$epoch_sampling -if (!is.null(epoch_cfg) && isTRUE(epoch_cfg$enabled)) { - lines <- c(lines, - sprintf("**Epoch sampling**: %d epochs of %.1fs", - epoch_cfg$n_epochs, epoch_cfg$epoch_duration_sec), - "" - ) -} - -lines <- c(lines, - "## Methods", - "", - sprintf("Imaginary coherence was computed between all %d vertex pairs ", n_vertices * (n_vertices - 1) / 2), - sprintf("using cross-spectral density (Welch's method). FCD was derived "), - sprintf("by counting the fraction of connections > %.3f per vertex. ", fcd_threshold), - "Group differences in FCD were tested with cluster-based permutation.", - "", - "## FCD Summary by Group", - "", - "| Band | Group | Mean FCD | SD FCD |", - "|------|-------|----------|--------|" -) - -for (i in seq_len(nrow(band_summary))) { - r <- band_summary[i, ] - lines <- c(lines, sprintf("| %s | %s | %.4f | %.4f |", - r$Band, r$Group, r$Mean_FCD, r$SD_FCD)) -} - -# Statistics if available -if (file.exists(stats_path)) { - stats <- read.csv(stats_path, stringsAsFactors = FALSE) - lines <- c(lines, "", "## Cluster Statistics", "") - - for (b in bands) { - sub <- stats[stats$band == b, ] - n_clust <- length(unique(sub$cluster_id[sub$cluster_id > 0])) - lines <- c(lines, sprintf("- **%s**: %d clusters identified", b, n_clust)) - } - lines <- c(lines, "") -} - -lines <- c(lines, - "## Output Files", - "", - "- `data/vertex_fcd.csv` — FCD per subject per vertex per band", - "- `data/vertex_connectivity_matrices.pkl` — full connectivity matrices", - "- `tables/vertex_connectivity_stats.csv` — cluster statistics", - "- `figures/fcd_*.png` — FCD glass brain maps", - "" -) - -writeLines(lines, file.path(output_dir, "ANALYSIS_SUMMARY.md")) -cat("Wrote ANALYSIS_SUMMARY.md\n") diff --git a/R/vertex_signature_analysis.R b/R/vertex_signature_analysis.R deleted file mode 100644 index 32729c2..0000000 --- a/R/vertex_signature_analysis.R +++ /dev/null @@ -1,130 +0,0 @@ -#!/usr/bin/env Rscript -# vertex_signature_analysis.R — Report Generator -# Reads MVPA classification results and generates ANALYSIS_SUMMARY.md - -suppressPackageStartupMessages({ - library(optparse) - library(yaml) -}) - -option_list <- list( - make_option("--data-dir", type = "character", help = "Path to data/ directory"), - make_option("--config", type = "character", help = "Path to study_config.yaml"), - make_option("--output-dir", type = "character", help = "Path to output directory"), - make_option("--no-figures", action = "store_true", default = FALSE, - help = "Skip all figure generation") -) -opts <- parse_args(OptionParser(option_list = option_list)) - -no_figures <- isTRUE(opts[["no-figures"]]) - -if (no_figures) { - ggsave <- function(...) invisible(NULL) -} - -data_dir <- opts[["data-dir"]] -config_path <- opts[["config"]] -output_dir <- opts[["output-dir"]] - -config <- read_yaml(config_path) - -# --- Load data ---------------------------------------------------------------- -results_path <- file.path(output_dir, "tables", "vertex_signature_results.csv") -if (!file.exists(results_path)) { - cat("No vertex_signature_results.csv found.\n") - quit(status = 0) -} - -results <- read.csv(results_path, stringsAsFactors = FALSE) - -sig_cfg <- config$vertex_signature %||% list() -classifiers <- sig_cfg$classifiers %||% list(sig_cfg$classifier %||% "svm_linear") -classifiers <- unlist(classifiers) -cv_method <- sig_cfg$cv_method %||% "loocv" -n_perm <- sig_cfg$n_permutations %||% 1000 - -# --- Features info ----------------------------------------------------------- -features_path <- file.path(data_dir, "vertex_signature_features.csv") -n_subjects <- 0 -n_features <- 0 -if (file.exists(features_path)) { - feats <- read.csv(features_path, stringsAsFactors = FALSE) - n_subjects <- length(unique(feats$subject)) - n_features <- length(unique(feats$vertex_idx)) -} - -# --- Write ANALYSIS_SUMMARY.md ----------------------------------------------- -lines <- c( - "# Neural Signature Analysis Summary", - "", - sprintf("**Study**: %s", config$name), - "**Analysis**: Whole-brain vertex-level neural signature (classification)", - sprintf("**Classifiers**: %s", paste(classifiers, collapse = ", ")), - sprintf("**CV method**: %s", cv_method), - sprintf("**Permutations**: %d", n_perm), - sprintf("**Subjects**: %d", n_subjects), - sprintf("**Features (vertices)**: %d", n_features), - "", - "## Methods", - "", - "Each classifier, with LOOCV, was trained to distinguish groups from whole-brain", - "spatial patterns of relative band power. Statistical significance was assessed", - "via permutation testing (shuffled group labels). Linear models report per-vertex", - "feature importance; non-linear models report accuracy only.", - "" -) - -# Epoch info -wb_cfg <- config$vertex %||% list() -epoch_cfg <- wb_cfg$epoch_sampling -if (!is.null(epoch_cfg) && isTRUE(epoch_cfg$enabled)) { - lines <- c(lines, - sprintf("**Epoch sampling**: %d epochs of %.1fs", - epoch_cfg$n_epochs, epoch_cfg$epoch_duration_sec), - "" - ) -} - -has_model <- "model" %in% names(results) -lines <- c(lines, - "## Classification Results", - "", - "| Model | Band | Accuracy | p-value | Sensitivity | Specificity | AUC | 95% CI |", - "|-------|------|----------|---------|-------------|-------------|-----|--------|" -) - -for (i in seq_len(nrow(results))) { - r <- results[i, ] - model <- if (has_model) r$model else "—" - lines <- c(lines, sprintf( - "| %s | %s | %.1f%% | %.4f | %.1f%% | %.1f%% | %.3f | [%.1f%%, %.1f%%] |", - model, r$band, r$accuracy * 100, r$p_value, - r$sensitivity * 100, r$specificity * 100, r$auc, - r$ci_lower * 100, r$ci_upper * 100 - )) -} - -# Highlight significant model x band cells -sig <- results[results$p_value < 0.05, ] -if (nrow(sig) > 0) { - lab <- if (has_model) paste(sig$model, sig$band) else sig$band - lines <- c(lines, "", - sprintf("**Significant (p < 0.05)**: %s", paste(lab, collapse = ", "))) -} else { - lines <- c(lines, "", "No model reached significance at p < 0.05.") -} - -lines <- c(lines, - "", - "## Output Files", - "", - "- `data/vertex_signature_features.csv` — feature matrix", - "- `tables/vertex_signature_results.csv` — classification accuracy per band", - "- `figures/vertex_signature_importance_*.png` — feature importance glass brains", - "- `figures/vertex_signature_null_*.png` — permutation null distributions", - "- `figures/vertex_signature_confusion_*.png` — confusion matrices", - "" -) - -writeLines(lines, file.path(output_dir, "ANALYSIS_SUMMARY.md")) -cat("Wrote ANALYSIS_SUMMARY.md\n") diff --git a/R/vertex_spatial_analysis.R b/R/vertex_spatial_analysis.R deleted file mode 100644 index 2c004ad..0000000 --- a/R/vertex_spatial_analysis.R +++ /dev/null @@ -1,497 +0,0 @@ -#!/usr/bin/env Rscript -# vertex_spatial_analysis.R — primary computation module -# Fits nlme::lme with random subject intercept + exponential spatial correlation -# per contrast x band x metric. Compares to non-spatial LME via AIC/BIC, -# generates variograms, exports residuals, writes ANALYSIS_SUMMARY.md -# -# Multi-group design: iterates over contrasts from config (subsetting to 2 groups -# per contrast), consistent with stats_utils.R pattern. - -suppressPackageStartupMessages({ - library(optparse) - library(yaml) - library(lme4) - library(lmerTest) - library(dplyr) - library(emmeans) -}) - -option_list <- list( - make_option("--data-dir", type = "character", help = "Path to data/ directory"), - make_option("--config", type = "character", help = "Path to study_config.yaml"), - make_option("--output-dir", type = "character", help = "Path to output directory"), - make_option("--fig-dir", type = "character", default = NULL, - help = "Directory for figures (default: output-dir/figures)"), - make_option("--tbl-dir", type = "character", default = NULL, - help = "Directory for tables (default: output-dir/tables)"), - make_option("--no-figures", action = "store_true", default = FALSE, - help = "Skip all figure generation") -) -opts <- parse_args(OptionParser(option_list = option_list)) - -no_figures <- isTRUE(opts[["no-figures"]]) - -data_dir <- opts[["data-dir"]] -config_path <- opts[["config"]] -output_dir <- opts[["output-dir"]] - -config <- read_yaml(config_path) - -# --- Load data ---------------------------------------------------------------- -data_path <- file.path(data_dir, "vertex_spatial_data.csv") -if (!file.exists(data_path)) { - cat("No vertex_spatial_data.csv found.\n") - quit(status = 0) -} - -dat <- read.csv(data_path, stringsAsFactors = FALSE) - -slmm_cfg <- config$vertex_spatial %||% config$spatial_lmm %||% list() -stat_method <- slmm_cfg$stat_method %||% "gls" -corr_struct <- slmm_cfg$correlation_structure %||% "exponential" -range_mm <- slmm_cfg$spatial_range_mm %||% 3.0 - -cat(sprintf("stat_method: %s\n", stat_method)) - -bands <- unique(dat$band) -metrics <- c("relative", "absolute") -contrasts <- config$contrasts -n_subjects_total <- length(unique(dat$subject)) -n_vertices <- length(unique(dat$vertex_idx)) - -cat(sprintf("Vertex spatial: %d subjects total, %d bands, %d metrics, %d vertices per subject\n", - n_subjects_total, length(bands), length(metrics), n_vertices)) -cat(sprintf("Contrasts: %d\n", length(contrasts))) - -# --- Fit models per contrast x band x metric --------------------------------- -fig_dir <- if (!is.null(opts[["fig-dir"]])) opts[["fig-dir"]] else file.path(output_dir, "figures") -tbl_dir <- if (!is.null(opts[["tbl-dir"]])) opts[["tbl-dir"]] else file.path(output_dir, "tables") -dir.create(fig_dir, showWarnings = FALSE, recursive = TRUE) -dir.create(tbl_dir, showWarnings = FALSE, recursive = TRUE) - -# ============================================================================= -# RETIRED (design-spec migration, 2026-06). vertex_spatial fit a per-contrast -# GLS spatial-covariance model (corExp + nugget) as a robustness check on the -# vertex group difference, iterating `config$contrasts`. The contrasts:/ -# hypothesis_testing: blocks were replaced by the declarative design:/hypotheses: -# spec, so `config$contrasts` is now NULL and this module has no contrasts to fit. -# It is retired rather than migrated: spatially-resolved vertex inference is -# delivered by vertex_cluster (cluster-based permutation glass-brain maps) and -# vertex_nbs (network-based statistic); the spatial-covariance robustness table -# was never a manuscript result. We emit empty result/residual frames + a note so -# downstream consumers find a well-formed (empty) output, then exit cleanly. -.retire_note <- paste0( - "vertex_spatial is RETIRED (design-spec migration). The per-contrast GLS ", - "spatial-covariance robustness model iterated config$contrasts, which the ", - "declarative design:/hypotheses: spec no longer populates. Spatially-resolved ", - "vertex inference is provided by vertex_cluster (cluster-permutation glass-brain ", - "maps) and vertex_nbs (network-based statistic).") -write.csv(data.frame(), file.path(tbl_dir, "vertex_spatial_results.csv"), row.names = FALSE) -write.csv(data.frame(), file.path(tbl_dir, "vertex_spatial_residuals.csv"), row.names = FALSE) -writeLines(c("# Vertex Spatial Analysis — RETIRED", "", .retire_note), - file.path(output_dir, "ANALYSIS_SUMMARY.md")) -cat("\n", .retire_note, "\n", sep = "") -quit(status = 0) - -# ---- Dead code below (pre-retirement GLS machinery; left for reference) ------ -model_results <- list() -all_residuals <- data.frame() -result_idx <- 0 - -for (contrast in contrasts) { - cname <- contrast$name - ga <- contrast$group_a - gb <- contrast$group_b - - cat(sprintf("\n========== Contrast: %s (%s vs %s) ==========\n", cname, ga, gb)) - - # Subset to the two groups for this contrast - cdat <- dat[dat$group %in% c(ga, gb), ] - cdat$group <- factor(cdat$group, levels = c(ga, gb)) - - n_a <- length(unique(cdat$subject[cdat$group == ga])) - n_b <- length(unique(cdat$subject[cdat$group == gb])) - cat(sprintf(" n(%s)=%d, n(%s)=%d\n", ga, n_a, gb, n_b)) - - for (band in bands) { - for (metric in metrics) { - result_idx <- result_idx + 1 - cat(sprintf("\n--- %s: %s [%s] ---\n", cname, band, metric)) - - band_dat <- cdat[cdat$band == band, ] - - if (nrow(band_dat) < 10) { - cat(sprintf(" Skipping: too few rows (%d)\n", nrow(band_dat))) - next - } - - if (!(metric %in% names(band_dat))) { - cat(sprintf(" Skipping: column '%s' not found\n", metric)) - next - } - - band_dat$response <- band_dat[[metric]] - - coef_val <- NA; se_val <- NA; t_val <- NA; p_val <- NA - f_val <- NA; df1 <- NA; df2 <- NA; interaction_p <- NA - estimated_range <- NA - aic_spatial <- NA; bic_spatial <- NA - aic_nonspatial <- NA; bic_nonspatial <- NA - convergence <- "failed" - fit_model <- NULL - - # ── GLS: spatial correlation, no random effect (replicates antwerp) ────── - if (stat_method == "gls") { - - # Baseline non-spatial GLS for AIC comparison - tryCatch({ - fit_base <- gls(response ~ group, data = band_dat, - control = glsControl(opt = "optim")) - aic_nonspatial <- AIC(fit_base) - bic_nonspatial <- BIC(fit_base) - cat(sprintf(" Non-spatial GLS: AIC=%.1f\n", aic_nonspatial)) - }, error = function(e) cat(sprintf(" Non-spatial GLS failed: %s\n", e$message))) - - tryCatch({ - fit_model <<- gls( - response ~ group, - correlation = corExp(value = range_mm, form = ~ x + y + z | subject, nugget = TRUE), - data = band_dat, - control = glsControl(opt = "optim", maxIter = 200, msMaxIter = 200, tolerance = 1e-4) - ) - aic_spatial <<- AIC(fit_model) - bic_spatial <<- BIC(fit_model) - convergence <<- "converged" - cat(sprintf(" GLS (corExp+nugget): AIC=%.1f\n", aic_spatial)) - }, error = function(e) { - cat(sprintf(" GLS (corExp+nugget) failed: %s\n", e$message)) - tryCatch({ - fit_model <<- gls( - response ~ group, - correlation = corExp(value = range_mm, form = ~ x + y + z | subject), - data = band_dat, - control = glsControl(opt = "optim", maxIter = 200) - ) - aic_spatial <<- AIC(fit_model) - bic_spatial <<- BIC(fit_model) - convergence <<- "converged (no nugget)" - cat(sprintf(" GLS fallback (no nugget): AIC=%.1f\n", aic_spatial)) - }, error = function(e2) cat(sprintf(" GLS fallback failed: %s\n", e2$message))) - }) - - if (!is.null(fit_model)) { - tryCatch({ - s <- summary(fit_model) - tbl <- s$tTable - if (nrow(tbl) >= 2) { - coef_val <- tbl[2, "Value"] - se_val <- tbl[2, "Std.Error"] - t_val <- tbl[2, "t-value"] - p_val <- tbl[2, "p-value"] - } - cs <- coef(fit_model$modelStruct$corStruct, unconstrained = FALSE) - estimated_range <- if ("range" %in% names(cs)) cs["range"] else cs[1] - }, error = function(e) cat(sprintf(" GLS summary failed: %s\n", e$message))) - - if (!no_figures && metric == "relative" && convergence != "failed") { - tryCatch({ - safe_band <- gsub(" ", "_", tolower(band)) - safe_cname <- gsub(" ", "_", cname) - png(file.path(fig_dir, sprintf("variogram_%s_%s.png", safe_cname, safe_band)), - width = 800, height = 500) - plot(Variogram(fit_model, form = ~ x + y + z | subject, maxDist = 8), - main = sprintf("Variogram — %s (relative) [%s]", band, cname)) - dev.off() - }, error = function(e) { - cat(sprintf(" Variogram failed: %s\n", e$message)) - tryCatch(dev.off(), error = function(e2) {}) - }) - } - } - - # ── Spatial LME: random subject intercept + spatial correlation ─────────── - } else if (stat_method == "spatial_lme") { - - tryCatch({ - fit_base <- lme(response ~ group, random = ~ 1 | subject, data = band_dat, - control = lmeControl(opt = "optim")) - aic_nonspatial <<- AIC(fit_base) - bic_nonspatial <<- BIC(fit_base) - cat(sprintf(" Non-spatial LME: AIC=%.1f\n", aic_nonspatial)) - }, error = function(e) cat(sprintf(" Non-spatial LME failed: %s\n", e$message))) - - tryCatch({ - fit_model <<- lme( - response ~ group, - random = ~ 1 | subject, - correlation = corExp(value = range_mm, form = ~ x + y + z, nugget = TRUE), - data = band_dat, - control = lmeControl(opt = "optim", maxIter = 200, msMaxIter = 200, tolerance = 1e-4) - ) - aic_spatial <<- AIC(fit_model) - bic_spatial <<- BIC(fit_model) - convergence <<- "converged" - cat(sprintf(" Spatial LME (corExp): AIC=%.1f\n", aic_spatial)) - }, error = function(e) { - cat(sprintf(" Spatial LME failed: %s\n", e$message)) - tryCatch({ - fit_model <<- lme( - response ~ group, - random = ~ 1 | subject, - correlation = corExp(value = range_mm, form = ~ x + y + z), - data = band_dat, - control = lmeControl(opt = "optim", maxIter = 100, msMaxIter = 100) - ) - aic_spatial <<- AIC(fit_model) - bic_spatial <<- BIC(fit_model) - convergence <<- "converged (no nugget)" - cat(sprintf(" Spatial LME fallback: AIC=%.1f\n", aic_spatial)) - }, error = function(e2) cat(sprintf(" Spatial LME fallback failed: %s\n", e2$message))) - }) - - if (!is.null(fit_model)) { - tryCatch({ - s <- summary(fit_model) - tbl <- s$tTable - if (nrow(tbl) >= 2) { - coef_val <- tbl[2, "Value"] - se_val <- tbl[2, "Std.Error"] - t_val <- tbl[2, "t-value"] - p_val <- tbl[2, "p-value"] - } - cs <- coef(fit_model$modelStruct$corStruct, unconstrained = FALSE) - estimated_range <- if ("range" %in% names(cs)) cs["range"] else cs[1] - }, error = function(e) cat(sprintf(" Spatial LME summary failed: %s\n", e$message))) - } - - # ── LMM: omnibus group × vertex LMM, consistent with ROI framework ──────── - } else if (stat_method == "lmm") { - - # Treat vertex_idx as a factor so we get group*vertex interaction - band_dat$vertex_f <- factor(band_dat$vertex_idx) - convergence <- "failed" - - tryCatch({ - fit_model <<- lmer( - response ~ group * vertex_f + (1 | subject), - data = band_dat, - REML = FALSE - ) - convergence <<- "converged" - aic_spatial <<- AIC(fit_model) - bic_spatial <<- BIC(fit_model) - - # Baseline: no vertex interaction - fit_base <- lmer(response ~ group + vertex_f + (1 | subject), - data = band_dat, REML = FALSE) - aic_nonspatial <<- AIC(fit_base) - bic_nonspatial <<- BIC(fit_base) - - cat(sprintf(" LMM converged: AIC=%.1f (interaction) vs %.1f (additive)\n", - aic_spatial, aic_nonspatial)) - }, error = function(e) cat(sprintf(" LMM failed: %s\n", e$message))) - - if (!is.null(fit_model) && convergence == "converged") { - tryCatch({ - # Type III ANOVA for group main effect and group:vertex interaction - an <- anova(fit_model, type = "III") - - # Group main effect - if ("group" %in% rownames(an)) { - f_val <<- an["group", "F value"] - df1 <<- an["group", "NumDF"] - df2 <<- an["group", "DenDF"] - p_val <<- an["group", "Pr(>F)"] - } - # Group × vertex interaction - int_row <- grep("group:vertex_f|vertex_f:group", rownames(an), value = TRUE)[1] - if (!is.na(int_row)) { - interaction_p <<- an[int_row, "Pr(>F)"] - } - - # Marginal group contrast (emmeans), averaged over vertices - em <- emmeans(fit_model, ~ group) - ct <- contrast(em, method = "pairwise") - ct_df <- as.data.frame(ct) - row1 <- ct_df[1, ] - coef_val <<- row1$estimate - se_val <<- row1$SE - t_val <<- row1$t.ratio - # Use interaction p if group main p not more informative - if (!is.na(interaction_p)) p_val <<- interaction_p - - cat(sprintf(" LMM group F=%.3f p=%.4f; interaction p=%.4f\n", - ifelse(is.na(f_val), NA, f_val), p_val, - ifelse(is.na(interaction_p), NA, interaction_p))) - }, error = function(e) cat(sprintf(" LMM summary failed: %s\n", e$message))) - } - } - - # Extract residuals for spatial methods - if (!is.null(fit_model) && stat_method %in% c("gls", "spatial_lme")) { - tryCatch({ - resids <- data.frame( - contrast = cname, subject = band_dat$subject, - vertex_idx = band_dat$vertex_idx, band = band, metric = metric, - residual = residuals(fit_model), stringsAsFactors = FALSE - ) - all_residuals <- rbind(all_residuals, resids) - }, error = function(e) {}) - } - - cat(sprintf(" Result: coef=%.4f, t=%.3f, p=%.4f\n", - ifelse(is.na(coef_val), NA, coef_val), - ifelse(is.na(t_val), NA, t_val), - ifelse(is.na(p_val), NA, p_val))) - - model_results[[result_idx]] <- data.frame( - contrast = cname, group_a = ga, group_b = gb, - n_a = n_a, n_b = n_b, - band = band, metric = metric, - stat_method = stat_method, convergence = convergence, - aic_spatial = aic_spatial, bic_spatial = bic_spatial, - aic_nonspatial = aic_nonspatial, bic_nonspatial = bic_nonspatial, - aic_improvement = ifelse(!is.na(aic_nonspatial) & !is.na(aic_spatial), - aic_nonspatial - aic_spatial, NA), - coefficient = coef_val, std_error = se_val, - t_value = t_val, p_value = p_val, - interaction_p = interaction_p, - estimated_range_mm = estimated_range, - stringsAsFactors = FALSE - ) - } - } -} - -# --- Compile results and apply FDR correction -------------------------------- -results_df <- do.call(rbind, model_results) - -# FDR correction across bands within each contrast x metric -if (nrow(results_df) > 0) { - results_df <- do.call(rbind, lapply(split(results_df, - interaction(results_df$contrast, results_df$metric, drop = TRUE)), function(sub) { - sub$q_value <- p.adjust(sub$p_value, method = "BH") - sub$significant <- !is.na(sub$q_value) & sub$q_value < 0.05 - sub - })) - rownames(results_df) <- NULL -} - -write.csv(results_df, file.path(tbl_dir, "vertex_spatial_results.csv"), row.names = FALSE) -cat(sprintf("\nExported vertex_spatial_results.csv (%d rows)\n", nrow(results_df))) - -if (nrow(all_residuals) > 0) { - write.csv(all_residuals, file.path(tbl_dir, "vertex_spatial_residuals.csv"), row.names = FALSE) - cat(sprintf("Exported vertex_spatial_residuals.csv (%d rows)\n", nrow(all_residuals))) -} - -# --- Write ANALYSIS_SUMMARY.md ----------------------------------------------- -groups_all <- unique(dat$group) -lines <- c( - "# Vertex Spatial Analysis Summary", - "", - sprintf("**Study**: %s", config$name), - sprintf("**Analysis**: Vertex Spatial (%s)", stat_method), - sprintf("**stat_method**: %s", stat_method), - if (stat_method %in% c("gls", "spatial_lme")) sprintf("**Correlation structure**: %s", corr_struct) else NULL, - if (stat_method %in% c("gls", "spatial_lme")) sprintf("**Initial spatial range**: %.1f mm", range_mm) else NULL, - sprintf("**Subjects**: %d total (%s)", n_subjects_total, paste(groups_all, collapse = ", ")), - sprintf("**Vertices**: %d (dorsal, z >= 0)", n_vertices), - sprintf("**Contrasts**: %d", length(contrasts)), - "" -) - -for (contrast in contrasts) { - n_a <- length(unique(dat$subject[dat$group == contrast$group_a])) - n_b <- length(unique(dat$subject[dat$group == contrast$group_b])) - lines <- c(lines, sprintf("- **%s**: %s (n=%d) vs %s (n=%d)", - contrast$name, contrast$group_a, n_a, contrast$group_b, n_b)) -} - -lines <- c(lines, - "", - "## Methods", - "", - if (stat_method == "gls") paste( - "Spatial GLS (nlme::gls) was used to model vertex-level band power as a function", - "of group, with an exponential spatial correlation structure", - "(`corExp(form = ~x+y+z|subject, nugget=TRUE)`). The `|subject` grouping allows", - "a separate correlation matrix per subject, effectively controlling for", - "between-subject variance without a separate random effect. This approach", - "replicates the antwerp manuscript analysis. Models were compared to a", - "non-spatial GLS (identity correlation) via AIC/BIC." - ) else if (stat_method == "spatial_lme") paste( - "Spatial linear mixed effects models (nlme::lme) were used with a random", - "subject intercept (`random = ~1|subject`) and an exponential spatial", - "correlation structure (`corExp(form = ~x+y+z, nugget=TRUE)`). The random", - "subject intercept partitions between-subject variance, while the spatial", - "correlation structure accounts for autocorrelation between nearby vertices." - ) else paste( - "Omnibus LMM (lmerTest::lmer) treating vertex location as a fixed categorical", - "factor: `power ~ group * vertex + (1|subject)`. Tests both the group main", - "effect (averaged over vertices) and the group x vertex interaction", - "(spatially heterogeneous group differences). FDR (BH) correction applied", - "across bands within each contrast x metric. Consistent with the ROI-level", - "analysis framework." - ), - "" -) - -lines <- c(lines, - "## Model Results", - "", - "| Contrast | Band | Metric | Convergence | AIC Improvement | Coef | SE | t | p | q |", - "|----------|------|--------|-------------|-----------------|------|----|---|---|---|" -) - -for (i in seq_len(nrow(results_df))) { - r <- results_df[i, ] - lines <- c(lines, sprintf( - "| %s | %s | %s | %s | %.1f | %.4f | %.4f | %.3f | %.4f | %.4f |", - r$contrast, r$band, r$metric, r$convergence, - ifelse(is.na(r$aic_improvement), NA, r$aic_improvement), - ifelse(is.na(r$coefficient), NA, r$coefficient), - ifelse(is.na(r$std_error), NA, r$std_error), - ifelse(is.na(r$t_value), NA, r$t_value), - ifelse(is.na(r$p_value), NA, r$p_value), - ifelse(is.na(r$q_value), NA, r$q_value) - )) -} - -# Significant results -sig_results <- results_df[!is.na(results_df$q_value) & results_df$q_value < 0.05, ] -if (nrow(sig_results) > 0) { - lines <- c(lines, "", - sprintf("**Significant results (q < 0.05): %d**", nrow(sig_results)), "") - for (i in seq_len(nrow(sig_results))) { - r <- sig_results[i, ] - lines <- c(lines, sprintf("- **%s** %s [%s]: t=%.3f, p=%.4f, q=%.4f", - r$band, r$metric, r$contrast, r$t_value, r$p_value, r$q_value)) - } -} else { - lines <- c(lines, "", "No results reached significance at q < 0.05 (FDR-corrected).") -} - -# Spatial range estimates (relative metric only, for clarity) -range_results <- results_df[!is.na(results_df$estimated_range_mm) & results_df$metric == "relative", ] -if (nrow(range_results) > 0) { - lines <- c(lines, "", "## Estimated Spatial Ranges (relative metric)", "") - for (i in seq_len(nrow(range_results))) { - r <- range_results[i, ] - lines <- c(lines, sprintf("- **%s** [%s]: %.2f mm", r$band, r$contrast, r$estimated_range_mm)) - } -} - -lines <- c(lines, - "", - "## Output Files", - "", - "- `data/vertex_spatial_data.csv` — per-subject per-vertex band power with coordinates", - "- `tables/vertex_spatial_results.csv` — GLS model results per contrast x band x metric", - "- `tables/vertex_spatial_residuals.csv` — spatial model residuals", - "- `figures/variogram_*.png` — empirical vs fitted variograms", - "" -) - -writeLines(lines, file.path(output_dir, "ANALYSIS_SUMMARY.md")) -cat("Wrote ANALYSIS_SUMMARY.md\n") diff --git a/R/vertex_specparam_analysis.R b/R/vertex_specparam_analysis.R deleted file mode 100644 index 6d42b7a..0000000 --- a/R/vertex_specparam_analysis.R +++ /dev/null @@ -1,275 +0,0 @@ -#!/usr/bin/env Rscript -# vertex_specparam_analysis.R — Report Generator -# Reads vertex-level spectral parameterization results, generates ANALYSIS_SUMMARY.md - -suppressPackageStartupMessages({ - library(optparse) - library(yaml) -}) - -option_list <- list( - make_option("--data-dir", type = "character", help = "Path to data/ directory"), - make_option("--config", type = "character", help = "Path to study_config.yaml"), - make_option("--output-dir", type = "character", help = "Path to output directory"), - make_option("--no-figures", action = "store_true", default = FALSE, - help = "Skip all figure generation") -) -opts <- parse_args(OptionParser(option_list = option_list)) - -no_figures <- isTRUE(opts[["no-figures"]]) - -if (no_figures) { - ggsave <- function(...) invisible(NULL) -} - -data_dir <- opts[["data-dir"]] -config_path <- opts[["config"]] -output_dir <- opts[["output-dir"]] - -config <- read_yaml(config_path) - -# --- Load data ---------------------------------------------------------------- -param_path <- file.path(data_dir, "vertex_specparam.csv") -if (!file.exists(param_path)) { - cat("No vertex_specparam.csv found.\n") - quit(status = 0) -} - -params <- read.csv(param_path, stringsAsFactors = FALSE) - -sp_cfg <- config$vertex_specparam %||% list() -# Keep in step with spectral/aperiodic.py::DEFAULT_FREQ_RANGE (12-45 Hz); see -# docs/APERIODIC_FIT_WINDOW.md for why. This is only the label fallback — the -# actual fit happens in Python, which stamps fit_fmin/fit_fmax into the params. -freq_range <- sp_cfg$freq_range %||% c(12, 45) -if (all(c("fit_fmin", "fit_fmax") %in% names(params))) { - freq_range <- c(params$fit_fmin[1], params$fit_fmax[1]) # authoritative -} -max_peaks <- sp_cfg$max_n_peaks %||% 6 - -# Peak-detection window. When it differs from the aperiodic window the run used -# a two-fit design: narrow window for the exponent, wider one for the peaks, so -# the narrow window's borders can be checked against where the peaks actually -# are instead of asserted. Python stamps peak_fmin/peak_fmax into the params. -peak_range <- freq_range -if (all(c("peak_fmin", "peak_fmax") %in% names(params))) { - peak_range <- c(params$peak_fmin[1], params$peak_fmax[1]) -} -two_fit <- !isTRUE(all.equal(as.numeric(peak_range), as.numeric(freq_range))) - -# --- Summaries ---------------------------------------------------------------- -n_subjects <- length(unique(params$subject)) -n_vertices <- length(unique(params$vertex_idx)) -groups <- unique(params$group) - -# Detect per-band peak columns dynamically -peak_cols <- grep("^has_.*_peak$", names(params), value = TRUE) -band_labels <- sub("^has_(.*)_peak$", "\\1", peak_cols) - -# Per-group summary of aperiodic parameters + per-band peak rates -group_summary <- do.call(rbind, lapply(groups, function(g) { - sub <- params[params$group == g, ] - base <- data.frame( - Group = g, - Mean_Exponent = round(mean(sub$exponent, na.rm = TRUE), 3), - SD_Exponent = round(sd(sub$exponent, na.rm = TRUE), 3), - Mean_Offset = round(mean(sub$offset, na.rm = TRUE), 3), - SD_Offset = round(sd(sub$offset, na.rm = TRUE), 3), - Mean_R2 = round(mean(sub$r_squared, na.rm = TRUE), 3), - stringsAsFactors = FALSE - ) - for (col in peak_cols) { - label <- paste0(sub("^has_(.*)_peak$", "\\1", col), "_Peak_Rate") - base[[label]] <- round(mean(sub[[col]], na.rm = TRUE), 3) - } - base -})) - -# Method distribution -method_table <- table(params$method) - -# --- Write ANALYSIS_SUMMARY.md ----------------------------------------------- -lines <- c( - "# Spectral Parameterization (Vertex-Level) Summary", - "", - sprintf("**Study**: %s", config$name), - "**Analysis**: Vertex-level spectral parameterization (aperiodic + peaks)", - sprintf("**Aperiodic fit range**: %g-%g Hz", freq_range[1], freq_range[2]), - sprintf("**Peak detection range**: %g-%g Hz%s", - peak_range[1], peak_range[2], - if (two_fit) " (separate wider fit)" else " (same fit)"), - sprintf("**Max peaks**: %d", max_peaks), - sprintf("**Subjects**: %d (%s)", n_subjects, paste(groups, collapse = ", ")), - sprintf("**Vertices**: %d", n_vertices), - "" -) - -# Epoch info -wb_cfg <- config$vertex %||% list() -epoch_cfg <- wb_cfg$epoch_sampling -if (!is.null(epoch_cfg) && isTRUE(epoch_cfg$enabled)) { - lines <- c(lines, - sprintf("**Epoch sampling**: %d epochs of %.1fs", - epoch_cfg$n_epochs, epoch_cfg$epoch_duration_sec), - "" - ) -} - -lines <- c(lines, - "## Methods", - "", - "Spectral parameterization (specparam/FOOOF) was applied to the PSD at each vertex.", - "The aperiodic component (1/f slope and offset) and oscillatory peaks were extracted.", - "Group differences in aperiodic parameters were tested with cluster-based permutation.", - "Per-band peak presence rates were compared with per-vertex chi-squared tests.", - "" -) - -if (two_fit) { - lines <- c(lines, - sprintf(paste0( - "Aperiodic parameters come from a %g-%g Hz fit; peaks come from a ", - "separate %g-%g Hz fit. The narrow window is what makes the exponent ", - "unbiased, but it can only find peaks inside itself, so it cannot be ", - "used to check its own borders. Peak columns are emitted only for bands ", - "the peak window can reach: a band outside it would otherwise report a ", - "0%% detection rate that is structural rather than measured."), - freq_range[1], freq_range[2], peak_range[1], peak_range[2]), - "") -} - -lines <- c(lines, - "## Fitting Methods Used", - "", - sprintf("- specparam: %d fits", method_table["specparam"] %||% 0), - sprintf("- linreg: %d fits", method_table["linreg"] %||% 0), - sprintf("- failed: %d fits", method_table["failed"] %||% 0), - "", - "## Group Summary", - "" -) - -# Build dynamic table header with per-band peak rate columns -rate_cols <- grep("_Peak_Rate$", names(group_summary), value = TRUE) -rate_headers <- sub("_Peak_Rate$", "", rate_cols) -header_line <- paste0("| Group | Mean Exp | SD Exp | Mean Offset | SD Offset | Mean R\u00b2 |", - paste0(" ", rate_headers, " Rate |", collapse = "")) -sep_line <- paste0("|-------|----------|--------|-------------|-----------|---------|", - paste0(rep("------|", length(rate_cols)), collapse = "")) -lines <- c(lines, header_line, sep_line) - -for (i in seq_len(nrow(group_summary))) { - r <- group_summary[i, ] - base_fmt <- sprintf("| %s | %.3f | %.3f | %.3f | %.3f | %.3f |", - r$Group, r$Mean_Exponent, r$SD_Exponent, - r$Mean_Offset, r$SD_Offset, r$Mean_R2) - rate_vals <- paste0(sprintf(" %.3f |", unlist(r[rate_cols])), collapse = "") - lines <- c(lines, paste0(base_fmt, rate_vals)) -} - -# Fit-window diagnostic — does the aperiodic window satisfy Gerster's border -# rule on THIS data? Reported before the results it underwrites. -diag_path <- file.path(output_dir, "tables", "fit_window_diagnostic.csv") -if (file.exists(diag_path)) { - diag <- read.csv(diag_path, stringsAsFactors = FALSE) - all_row <- diag[diag$band == "ALL", ] - lines <- c(lines, "", "## Fit-Window Diagnostic", "", - "Gerster et al. (2022): oscillations crossing the fit borders must be", - "avoided, since a peak on a border inflates exponent error. A peak counts", - "as crossing when its support (centre frequency +/- half the specparam", - "bandwidth) straddles a border.", - "") - if (nrow(all_row) == 1) { - verdict <- if (all_row$frac_crossing < 0.05) { - "SATISFIED - the borders sit in spectral gaps on this data." - } else if (all_row$frac_crossing < 0.15) { - "MARGINAL - a minority of peaks touch a border; exponents carry some added error." - } else { - "VIOLATED - peaks sit on the borders; the window needs revisiting for this dataset." - } - lines <- c(lines, - sprintf("**%d peaks** detected over %g-%g Hz. **%d (%.1f%%)** cross an aperiodic border (%g / %g Hz).", - all_row$n_peaks, all_row$peak_fmin, all_row$peak_fmax, - all_row$n_cross_fmin + all_row$n_cross_fmax, - 100 * all_row$frac_crossing, - all_row$aperiodic_fmin, all_row$aperiodic_fmax), - "", - sprintf("**Verdict: %s**", verdict), - "") - } - band_rows <- diag[diag$band != "ALL", ] - if (nrow(band_rows) > 0) { - lines <- c(lines, - "| Band | Range (Hz) | Reachable | Censored | Peaks | Median CF | Crossing |", - "|------|-----------|-----------|----------|-------|-----------|----------|") - for (i in seq_len(nrow(band_rows))) { - r <- band_rows[i, ] - lines <- c(lines, sprintf("| %s | %g-%g | %s | %s | %d | %s | %.1f%% |", - r$band, r$band_lo, r$band_hi, - if (isTRUE(as.logical(r$reachable))) "yes" else "**no**", - if (isTRUE(as.logical(r$censored))) "**yes**" else "no", - r$n_peaks, - if (is.na(r$cf_median)) "--" else sprintf("%.1f", r$cf_median), - 100 * r$frac_crossing)) - } - lines <- c(lines, "", - "Unreachable bands emit no peak columns at all - their absence is a", - "property of the window, not a measurement. Censored bands extend past", - "the peak window, so their detection rates are a lower bound.", - "") - } -} - -# Cluster stats -stats_path <- file.path(output_dir, "tables", "vertex_specparam_stats.csv") -if (file.exists(stats_path)) { - stats <- read.csv(stats_path, stringsAsFactors = FALSE) - lines <- c(lines, "", "## Cluster Permutation Results", "") - - for (param in unique(stats$parameter)) { - sub <- stats[stats$parameter == param, ] - cluster_ids <- unique(sub$cluster_id[sub$cluster_id > 0]) - lines <- c(lines, sprintf("- **%s**: %d clusters identified", param, length(cluster_ids))) - } - lines <- c(lines, "") -} - -# Per-band chi-squared results -chi2_path <- file.path(output_dir, "tables", "band_peak_chi2.csv") -if (!file.exists(chi2_path)) { - chi2_path <- file.path(output_dir, "tables", "gamma_peak_chi2.csv") -} -if (file.exists(chi2_path)) { - chi2 <- read.csv(chi2_path, stringsAsFactors = FALSE) - lines <- c(lines, "## Peak Presence by Band (Chi-squared)", "") - if ("band" %in% names(chi2)) { - for (band in unique(chi2$band)) { - sub <- chi2[chi2$band == band, ] - n_sig <- sum(sub$p < 0.05) - lines <- c(lines, sprintf("- **%s**: %d/%d vertices significant (uncorrected p<0.05)", - band, n_sig, nrow(sub))) - } - } else { - n_sig <- sum(chi2$p < 0.05) - lines <- c(lines, sprintf("- %d/%d vertices significant (uncorrected p<0.05)", - n_sig, nrow(chi2))) - } - lines <- c(lines, "") -} - -lines <- c(lines, - "## Output Files", - "", - "- `data/vertex_specparam.csv` — per-subject per-vertex specparam fit parameters", - "- `data/peak_inventory.csv` — every peak found by the peak-window fit (long format)", - "- `tables/vertex_specparam_stats.csv` — cluster permutation statistics", - "- `tables/fit_window_diagnostic.csv` — aperiodic border check + per-band reachability", - "- `tables/band_peak_chi2.csv` — per-band peak presence chi-squared tests", - "- `figures/fit_window_diagnostic.png` — peak centre frequencies vs the fit borders", - "- `figures/specparam_*.png` — aperiodic parameter glass brain maps", - "- `figures/{band}_peak_presence.png` — per-band peak prevalence maps", - "" -) - -writeLines(lines, file.path(output_dir, "ANALYSIS_SUMMARY.md")) -cat("Wrote ANALYSIS_SUMMARY.md\n") diff --git a/pyproject.toml b/pyproject.toml index d01fdcd..067434d 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -4,7 +4,7 @@ build-backend = "setuptools.build_meta" [project] name = "source-analytics" -version = "0.7.1" +version = "0.8.0" description = "Statistical analysis toolkit for source-localized EEG data (Python + R)" readme = "README.md" requires-python = ">=3.10" @@ -30,9 +30,9 @@ dependencies = [ # Optional extras. The package imports without any of them; the modules that # need one raise a clear ImportError naming the extra to install: -# mne -> the evoked / TFR modules (roi_evoked, vertex_evoked, electrode_evoked) -# mvpa -> vertex_signature / electrode_signature (scikit-learn) -# network -> roi_graph / vertex_graph / *_network / *_nbs (networkx) +# mne -> the evoked / TFR modules (roi_evoked, electrode_evoked) +# mvpa -> roi_signature / electrode_signature (scikit-learn) +# network -> roi_graph / roi_nbs / roi_network (networkx) # atlas -> (kept so old install commands work; nibabel is now core) [project.optional-dependencies] mne = ["mne>=1.5"] diff --git a/src/source_analytics/analyses/__init__.py b/src/source_analytics/analyses/__init__.py index a617e03..3225c78 100644 --- a/src/source_analytics/analyses/__init__.py +++ b/src/source_analytics/analyses/__init__.py @@ -5,32 +5,18 @@ from .roi_aperiodic_analysis import ROIAperiodicAnalysis from .roi_connectivity_analysis import ConnectivityAnalysis from .roi_cross_freq_analysis import ROICrossFreqAnalysis -from .vertex_cross_freq_analysis import VertexCrossFreqAnalysis -from .vertex_cluster_analysis import VertexClusterAnalysis from .electrode_analysis import ElectrodeAnalysis from .electrode_comparison_analysis import ElectrodeComparisonAnalysis from .electrode_connectivity_analysis import ElectrodeConnectivityAnalysis -from .vertex_connectivity_analysis import VertexConnectivityAnalysis -from .vertex_specparam_analysis import VertexSpecparamAnalysis -from .vertex_signature_analysis import VertexSignatureAnalysis from .roi_network_analysis import ROINetworkAnalysis -from .vertex_network_analysis import VertexNetworkAnalysis -from .vertex_spatial_analysis import VertexSpatialAnalysis from .roi_evoked_analysis import ROIEvokedAnalysis -from .vertex_evoked_analysis import VertexEvokedAnalysis from .roi_directed_analysis import ROIDirectedAnalysis -from .vertex_directed_analysis import VertexDirectedAnalysis # Backward-compatible aliases PSDAnalysis = ROIPsdAnalysis AperiodicAnalysis = ROIAperiodicAnalysis ROIPacAnalysis = ROICrossFreqAnalysis # renamed -> roi_cross_freq PACAnalysis = ROICrossFreqAnalysis -WholebrainAnalysis = VertexClusterAnalysis -MVPAAnalysis = VertexSignatureAnalysis -VertexMVPAAnalysis = VertexSignatureAnalysis # renamed -> vertex_signature -SpecparamVertexAnalysis = VertexSpecparamAnalysis -SpatialLMMAnalysis = VertexSpatialAnalysis EvokedAnalysis = ROIEvokedAnalysis ROITransferEntropyAnalysis = ROIDirectedAnalysis # renamed -> roi_directed TransferEntropyAnalysis = ROIDirectedAnalysis @@ -42,32 +28,18 @@ "ROIAperiodicAnalysis", "ConnectivityAnalysis", "ROICrossFreqAnalysis", - "VertexCrossFreqAnalysis", "ROIPacAnalysis", - "VertexClusterAnalysis", "ElectrodeAnalysis", "ElectrodeComparisonAnalysis", "ElectrodeConnectivityAnalysis", - "VertexConnectivityAnalysis", - "VertexSpecparamAnalysis", - "VertexSignatureAnalysis", - "VertexMVPAAnalysis", "ROINetworkAnalysis", - "VertexNetworkAnalysis", - "VertexSpatialAnalysis", "ROIEvokedAnalysis", - "VertexEvokedAnalysis", "ROIDirectedAnalysis", - "VertexDirectedAnalysis", "ROITransferEntropyAnalysis", # Backward-compatible aliases "PSDAnalysis", "AperiodicAnalysis", "PACAnalysis", - "WholebrainAnalysis", - "MVPAAnalysis", - "SpecparamVertexAnalysis", - "SpatialLMMAnalysis", "EvokedAnalysis", "TransferEntropyAnalysis", ] diff --git a/src/source_analytics/analyses/electrode_connectivity_analysis.py b/src/source_analytics/analyses/electrode_connectivity_analysis.py index 5540774..88ab34e 100644 --- a/src/source_analytics/analyses/electrode_connectivity_analysis.py +++ b/src/source_analytics/analyses/electrode_connectivity_analysis.py @@ -2,7 +2,8 @@ Computes all-to-all functional connectivity between the raw scalp electrodes (the 30-channel MEA) and derives per-channel Functional Connectivity Density -(FCD), mirroring :class:`VertexConnectivityAnalysis` at the sensor level. +(FCD), mirroring vertex_connectivity (now in the source-analytics-vertex plugin) +at the sensor level. This is the **comparator** for the MS2 connectivity-methods thesis: source- localized (vertex) connectivity recovers spatial structure that sensor-level diff --git a/src/source_analytics/analyses/fcd_comparison_analysis.py b/src/source_analytics/analyses/fcd_comparison_analysis.py deleted file mode 100644 index 1ebd5b2..0000000 --- a/src/source_analytics/analyses/fcd_comparison_analysis.py +++ /dev/null @@ -1,293 +0,0 @@ -"""Source-vs-sensor functional-connectivity-density (FCD) comparison. - -Compares vertex-level (source) FCD against electrode-level (sensor) FCD — the -head-to-head behind the MS2 thesis that source connectivity recovers spatial -structure that sensor connectivity blurs. Reads the per-element FCD tables that -``vertex_connectivity`` and ``electrode_connectivity`` already persist, and for -each subject × band × metric derives two per-subject summaries: - -* ``fcd_mean`` — mean FCD over the map's units (global connectivity density). - FCD is degree/(n-1)-normalized, so ``fcd_mean`` is directly comparable across - the two resolutions (30 channels vs. hundreds of vertices). - [Tomasi & Volkow 2010, PNAS 107(21):9885-9890 — FCD mapping.] -* ``fcd_cv`` — coefficient of variation (SD/mean) of FCD across units: the - *spatial heterogeneity* of connectivity density. Higher CV = more structure - (hubs vs. a flat field); the "source resolves, sensor blurs" quantity. - [Standard CV; see e.g. Everitt & Skrondal, Cambridge Dict. of Statistics.] - -For each (band, metric) it reports the cross-subject concordance (Pearson r, -source vs. sensor) of each summary, and the per-contrast group effect (Hedges g -+ 95% CI) at BOTH levels with a sign-concordance flag. NB the absolute CV is not -comparable across resolutions (more units sample the FCD field more finely), but -each group contrast is WITHIN a level, so the source-vs-sensor comparison of the -group EFFECT is resolution-fair — the same logic as ``electrode_comparison`` for -spectral power. -""" - -from __future__ import annotations - -import logging -from pathlib import Path - -import numpy as np -import pandas as pd - -from ..config import StudyConfig -from ..io.discovery import SubjectInfo -from ..viz.constants import metric_display, order_bands -from .base import BaseAnalysis -from .electrode_comparison_analysis import _hedges_g_ci - -logger = logging.getLogger(__name__) - - -def _fcd_summaries(fcd_df: pd.DataFrame, prefix: str) -> pd.DataFrame: - """Per-subject FCD summaries from a long per-element FCD table. - - ``fcd_df`` has columns subject, group, band, metric, fcd (one row per spatial - unit — channel or vertex). Returns one row per (subject, group, band, metric) - with ``_mean`` (mean FCD) and ``_cv`` (SD/mean over units). - """ - def _agg(g: pd.DataFrame) -> pd.Series: - vals = g["fcd"].to_numpy(dtype=float) - vals = vals[~np.isnan(vals)] - mean = float(np.mean(vals)) if vals.size else np.nan - if vals.size < 2 or not np.isfinite(mean) or mean == 0: - cv = np.nan - else: - cv = float(np.std(vals, ddof=1) / mean) - return pd.Series({f"{prefix}_mean": mean, f"{prefix}_cv": cv}) - - return ( - fcd_df.groupby(["subject", "group", "band", "metric"], sort=False) - .apply(_agg) - .reset_index() - ) - - -class FCDComparisonAnalysis(BaseAnalysis): - """Compare source (vertex) vs sensor (electrode) FCD, per band × metric.""" - - name = "fcd_comparison" - SELECTABLE = {"metric": "connectivity metric", "band": "frequency band"} - - # (prefix, native summary column) — the two per-subject FCD summaries compared - _SUMMARIES = ("mean", "cv") - - def __init__(self, config: StudyConfig, output_dir: Path): - super().__init__(config, output_dir) - self._sensor_df: pd.DataFrame | None = None - self._source_df: pd.DataFrame | None = None - self._comparison_df: pd.DataFrame | None = None - cfg = config.raw.get(self.name, {}) - self._metric_filter = cfg.get("metrics") - - def _find_upstream_csv(self, module: str, filename: str, override_key: str) -> Path: - """Locate a primary module's data CSV, in this paradigm or a sibling one. - - The two primaries usually live in *different* paradigms (the canonical - layout runs ``electrode_connectivity`` under ``resting`` and - ``vertex_connectivity`` under ``vertex``), so a same-paradigm lookup is - not enough. Search order: - - 1. an explicit ``: `` in this module's config block - (a directory containing ``data/`` or the CSV itself); - 2. this paradigm's working dir (``analytics///data``); - 3. every sibling paradigm dir under the same analytics root. - """ - cfg = self.config.raw.get(self.name, {}) or {} - override = cfg.get(override_key) - if override: - p = Path(override) - if not p.is_absolute(): - p = (self.config.output_dir / p).resolve() - candidate = p if p.suffix == ".csv" else p / "data" / filename - if candidate.exists(): - return candidate - raise FileNotFoundError( - f"{self.name}: {override_key}={override!r} does not contain {filename}" - ) - - base = self.config.output_dir - same = base / module / "data" / filename - if same.exists(): - return same - - searched = [same] - root = base.parent if self.config.paradigm_name else base - for sibling in sorted(p for p in root.iterdir() if p.is_dir()) if root.is_dir() else []: - candidate = sibling / module / "data" / filename - searched.append(candidate) - if candidate.exists(): - logger.info("%s: using %s from paradigm dir %s", self.name, module, sibling.name) - return candidate - raise FileNotFoundError( - f"{filename} not found — run '{module}' first (any paradigm). Searched: " - + ", ".join(str(s) for s in searched) - + f". Or set {self.name}.{override_key} to its output directory." - ) - - def setup(self) -> None: - sensor_csv = self._find_upstream_csv( - "electrode_connectivity", "electrode_fcd.csv", "sensor_dir", - ) - source_csv = self._find_upstream_csv( - "vertex_connectivity", "vertex_fcd.csv", "source_dir", - ) - self._sensor_df = pd.read_csv(sensor_csv) - self._source_df = pd.read_csv(source_csv) - logger.info( - "Loaded FCD: sensor=%d rows, source=%d rows", - len(self._sensor_df), len(self._source_df), - ) - - def process_subject(self, subject: SubjectInfo) -> None: # noqa: D401 - """No-op — reads persisted FCD tables.""" - - def aggregate(self) -> None: - if self._sensor_df is None or self._source_df is None: - return - sensor = _fcd_summaries(self._sensor_df, "sensor") - source = _fcd_summaries(self._source_df, "source") - comp = source.merge(sensor, on=["subject", "group", "band", "metric"], how="inner") - - # metric intersection, honoring --metric / config metrics filter - metrics = sorted(comp["metric"].unique()) - if self._metric_filter: - metrics = [m for m in metrics if m in set(self._metric_filter)] - metrics = self._select("metric", metrics) - comp = comp[comp["metric"].isin(metrics)] - - bands = self._select("band", order_bands(comp["band"].unique(), self.config)) - comp = comp[comp["band"].isin(bands)] - - self._comparison_df = comp.reset_index(drop=True) - out = self.output_dir / "data" / "fcd_subject_summary.csv" - self._comparison_df.to_csv(out, index=False) - logger.info("Exported fcd_subject_summary.csv (%d rows)", len(self._comparison_df)) - - def statistics(self) -> None: - from scipy import stats as sp_stats - - comp = self._comparison_df - if comp is None or comp.empty: - logger.warning("No FCD comparison data — skipping statistics") - return - - rows = [] - for (band, metric), bdata in comp.groupby(["band", "metric"], sort=False): - row: dict = {"band": band, "metric": metric, "n_subjects": len(bdata)} - # cross-subject concordance (source vs sensor) of each summary - for s in self._SUMMARIES: - valid = bdata[[f"sensor_{s}", f"source_{s}"]].dropna() - if len(valid) > 2: - r, p = sp_stats.pearsonr(valid[f"sensor_{s}"], valid[f"source_{s}"]) - else: - r, p = np.nan, np.nan - row[f"corr_{s}_r"] = r - row[f"corr_{s}_p"] = p - - base_row = dict(row) - for contrast in self._pairwise_contrasts(): - ga = bdata[bdata["group"] == contrast.group_a] - gb = bdata[bdata["group"] == contrast.group_b] - r = dict(base_row) - r["contrast"] = contrast.name - for s in self._SUMMARIES: - sen_g, sen_lo, sen_hi = _hedges_g_ci( - ga[f"sensor_{s}"].values, gb[f"sensor_{s}"].values) - src_g, src_lo, src_hi = _hedges_g_ci( - ga[f"source_{s}"].values, gb[f"source_{s}"].values) - r[f"sensor_{s}_g"] = sen_g - r[f"sensor_{s}_ci_lo"] = sen_lo - r[f"sensor_{s}_ci_hi"] = sen_hi - r[f"source_{s}_g"] = src_g - r[f"source_{s}_ci_lo"] = src_lo - r[f"source_{s}_ci_hi"] = src_hi - r[f"{s}_concordant"] = bool( - np.isfinite(sen_g) and np.isfinite(src_g) - and np.sign(sen_g) == np.sign(src_g) - ) - rows.append(r) - - stats_df = pd.DataFrame(rows) - if not stats_df.empty: - path = self.tbl_dir / "fcd_comparison_stats.csv" - stats_df.to_csv(path, index=False) - logger.info("Exported fcd_comparison_stats.csv (%d rows)", len(stats_df)) - self._stats_df = stats_df - - def figures(self) -> None: - # Regenerable from persisted data: reload the per-subject summary if the - # in-memory frame is absent (e.g. `--steps figures` standalone). - comp = self._comparison_df - if comp is None or comp.empty: - summ = self.output_dir / "data" / "fcd_subject_summary.csv" - if summ.exists(): - comp = pd.read_csv(summ) - if comp is None or comp.empty: - return - import matplotlib - matplotlib.use("Agg") - import matplotlib.pyplot as plt - - for metric, mdata in comp.groupby("metric", sort=False): - # (1) concordance scatter: source vs sensor mean FCD, colored by band - fig, ax = plt.subplots(figsize=(5, 5)) - bands = order_bands(mdata["band"].unique(), self.config) - cmap = plt.get_cmap("viridis", max(len(bands), 1)) - for i, band in enumerate(bands): - d = mdata[mdata["band"] == band][["sensor_mean", "source_mean"]].dropna() - if not d.empty: - ax.scatter(d["sensor_mean"], d["source_mean"], s=18, - color=cmap(i), alpha=0.7, label=band) - lims = [ - float(np.nanmin([mdata["sensor_mean"].min(), mdata["source_mean"].min()])), - float(np.nanmax([mdata["sensor_mean"].max(), mdata["source_mean"].max()])), - ] - if np.all(np.isfinite(lims)): - ax.plot(lims, lims, ls="--", c="grey", lw=0.8) - ax.set_xlabel("Sensor mean FCD") - ax.set_ylabel("Source mean FCD") - ax.set_title(f"Source vs sensor mean FCD — {metric_display(metric)}") - ax.legend(fontsize=7, title="Band") - fig.tight_layout() - fig.savefig(self.fig_dir / f"fcd_concordance_mean_{metric}.png", dpi=200) - plt.close(fig) - - # (2) spatial heterogeneity: mean CV by group, source vs sensor, per band - fig, ax = plt.subplots(figsize=(6, 4)) - grp = (mdata.groupby("band")[["sensor_cv", "source_cv"]] - .mean().reindex(bands)) - x = np.arange(len(bands)) - ax.bar(x - 0.2, grp["sensor_cv"], width=0.4, label="Sensor", color="#888") - ax.bar(x + 0.2, grp["source_cv"], width=0.4, label="Source", color="#B2182B") - ax.set_xticks(x) - ax.set_xticklabels(bands, rotation=45, ha="right", fontsize=8) - ax.set_ylabel("FCD spatial CV (SD/mean)") - ax.set_title(f"FCD spatial heterogeneity — {metric_display(metric)}") - ax.legend(fontsize=8) - fig.tight_layout() - fig.savefig(self.fig_dir / f"fcd_heterogeneity_{metric}.png", dpi=200) - plt.close(fig) - - def summary(self) -> None: - stats_df = getattr(self, "_stats_df", None) - lines = ["# FCD Source-vs-Sensor Comparison\n"] - if stats_df is None or stats_df.empty: - lines.append("*No comparison statistics computed.*\n") - else: - n_mean_conc = int(stats_df.get("mean_concordant", pd.Series(dtype=bool)).sum()) - n_cv_conc = int(stats_df.get("cv_concordant", pd.Series(dtype=bool)).sum()) - lines.append( - f"{len(stats_df)} band × metric × contrast cells. " - f"Group-effect sign-concordance source-vs-sensor: " - f"mean FCD {n_mean_conc}/{len(stats_df)}, " - f"spatial CV {n_cv_conc}/{len(stats_df)}.\n" - ) - lines.append( - "Tables: `fcd_comparison_stats.csv` (concordance r + per-contrast " - "Hedges g at both levels), `fcd_subject_summary.csv` (per-subject " - "mean + CV).\n" - ) - (self.output_dir / "ANALYSIS_SUMMARY.md").write_text("".join(lines), encoding="utf-8") diff --git a/src/source_analytics/analyses/vertex_cluster_analysis.py b/src/source_analytics/analyses/vertex_cluster_analysis.py deleted file mode 100644 index 9e1ae13..0000000 --- a/src/source_analytics/analyses/vertex_cluster_analysis.py +++ /dev/null @@ -1,826 +0,0 @@ -"""Vertex cluster spectral analysis with cluster-based permutation testing. - -Computes PSD once per subject for all 154 shell vertices, then extracts: -- Relative and absolute band power per vertex per band -- Spectral slope (1/f exponent) -- Peak alpha frequency - -Statistics use voxel-wise t-tests with spatial cluster-based permutation -correction (Maris & Oostenveld, 2007). Visualization via glass brain plots. -""" - -from __future__ import annotations - -import logging -import pickle -import subprocess -from pathlib import Path - -import numpy as np -import pandas as pd -import yaml - -from ..config import StudyConfig -from ..io.discovery import SubjectInfo -from ..io.loader import SubjectLoader -from ..spectral.vertex import ( - compute_psd_vertices, - extract_band_power_vertices, - compute_spectral_slope, - compute_peak_frequency, -) -from ..spectral.epoch_sampler import sample_epochs -from ..stats.cluster_permutation import ( - cluster_permutation_test, - has_significant_cluster as _has_significant_cluster, - hedges_g, - voxelwise_ttest, -) -from ..stats.tfce import tfce_permutation_test -from ..viz.glass_brain import ( - plot_band_comparison, - plot_glass_brain, - plot_vertex_cluster_summary, -) -from .base import BaseAnalysis - -logger = logging.getLogger(__name__) - - -def _find_r_script_dir() -> Path: - """Locate the R/ directory relative to this package.""" - pkg_root = Path(__file__).resolve().parent.parent.parent.parent - r_dir = pkg_root / "R" - if r_dir.is_dir(): - return r_dir - for candidate in [Path.cwd() / "R", Path(__file__).parent.parent.parent / "R"]: - if candidate.is_dir(): - return candidate - raise FileNotFoundError( - "Cannot find R/ scripts directory. Expected at: " + str(pkg_root / "R") - ) - - -class VertexClusterAnalysis(BaseAnalysis): - """Vertex cluster spectral analysis with cluster permutation testing. - - Processes shell_ellipsoid source data (154 vertices), computes spectral - metrics, runs voxel-wise statistics with cluster correction, and generates - glass brain visualizations. - """ - - name = "vertex_cluster" - SELECTABLE = {"band": "frequency band", "hypothesis": "declared hypothesis"} - - def __init__(self, config: StudyConfig, output_dir: Path): - super().__init__(config, output_dir) - self._band_power_rows: list[dict] = [] - self._feature_rows: list[dict] = [] - self._source_coords: np.ndarray | None = None - self._vertex_indices: np.ndarray | None = None # original indices of kept vertices - self._sfreq: float | None = None - # Per-subject arrays for statistics: {subject_uid: {band: {metric: array}}} - self._subject_data: dict[str, dict] = {} - self._subject_groups: dict[str, str] = {} - - # Vertex cluster-specific config - wb_cfg = config.vertex - self._cluster_threshold = float(wb_cfg.get("cluster_threshold", 2.0)) - self._n_permutations = int(wb_cfg.get("n_permutations", 1000)) - self._adjacency_distance = float(wb_cfg.get("adjacency_distance_mm", 5.0)) - self._noise_exclude = wb_cfg.get("noise_exclude_hz") - if self._noise_exclude is not None: - self._noise_exclude = tuple(self._noise_exclude) - - # Correction method: "cluster" (default) or "tfce" - self._correction_method = wb_cfg.get("correction_method", "cluster") - if self._correction_method == "tfce": - tfce_cfg = wb_cfg.get("tfce", {}) - self._tfce_E = float(tfce_cfg.get("E", 0.5)) - self._tfce_H = float(tfce_cfg.get("H", 2.0)) - self._tfce_dh = float(tfce_cfg.get("dh", 0.1)) - elif self._correction_method != "cluster": - logger.warning( - "Unknown correction_method '%s', using 'cluster'", - self._correction_method, - ) - self._correction_method = "cluster" - - # Global epoch_sampling → vertex: block → per-analysis block (see base). - # None = full continuous PSD (the historical vertex_cluster behaviour). - self._epoch_config = self._vertex_epoch_config() - - def setup(self) -> None: - self._band_power_rows.clear() - self._feature_rows.clear() - self._subject_data.clear() - self._subject_groups.clear() - self._source_coords = None - self._vertex_indices = None - - def _compute_subject(self, subject: SubjectInfo): - """Pure per-subject band-power / slope / peak-alpha compute (parallel-safe).""" - loader = SubjectLoader(subject.data_dir) - uid = f"{subject.group}_{subject.subject_id}" - - # Load source time courses: (n_vertices, n_times) - stc_data = loader.load_source_timecourses() - sfreq = loader.load_sfreq() - coords = loader.load_source_coords() - - mask = self.config.get_vertex_mask(coords) - vertex_indices = np.where(mask)[0] - source_coords = coords[mask] - stc_data = stc_data[vertex_indices] - n_vertices = stc_data.shape[0] - - # Compute PSD for all vertices: (n_vertices, n_freqs) - fmax = max(hi for _, hi in self.config.bands.values()) + 10 - if self._epoch_config is not None: - epochs = sample_epochs( - stc_data, sfreq, - epoch_duration_sec=self._epoch_config.get("epoch_duration_sec", 2.0), - n_epochs=self._epoch_config.get("n_epochs", 80), - seed=self._epoch_config.get("seed", 42), - n_bootstrap=self._epoch_config.get("n_bootstrap", 1), - ) - all_psd = [] - for ep in epochs: - freqs, p = compute_psd_vertices(ep, sfreq, fmax=fmax) - all_psd.append(p) - psd = np.mean(all_psd, axis=0) - else: - freqs, psd = compute_psd_vertices(stc_data, sfreq, fmax=fmax) - - # Extract band power metrics - band_power = extract_band_power_vertices( - freqs, psd, self._selected_bands(), - noise_exclude=self._noise_exclude, - ) - - # Compute additional features - slope = compute_spectral_slope(freqs, psd) - peak_alpha = compute_peak_frequency(freqs, psd, search_range=(6, 13)) - - # Accumulate rows for CSV export (use original vertex indices) - band_power_rows: list[dict] = [] - for band_name, bp in band_power.items(): - for vi in range(n_vertices): - band_power_rows.append({ - "subject": uid, - "group": subject.group, - "vertex_idx": int(vertex_indices[vi]), - "band": band_name, - "absolute": float(bp["absolute"][vi]), - "relative": float(bp["relative"][vi]), - }) - - feature_rows: list[dict] = [] - for vi in range(n_vertices): - feature_rows.append({ - "subject": uid, - "group": subject.group, - "vertex_idx": int(vertex_indices[vi]), - "spectral_slope": float(slope[vi]), - "peak_alpha_freq": float(peak_alpha[vi]), - }) - - return { - "uid": uid, "group": subject.group, "sfreq": float(sfreq), - "vertex_indices": vertex_indices, "source_coords": source_coords, - "subject_data": {"band_power": band_power, "slope": slope, - "peak_alpha": peak_alpha}, - "band_power_rows": band_power_rows, "feature_rows": feature_rows, - } - - def _merge_subject(self, payload) -> None: - uid = payload["uid"] - if self._sfreq is None: - self._sfreq = payload["sfreq"] - elif payload["sfreq"] != self._sfreq: - logger.warning("Subject %s has sfreq=%.0f, expected %.0f", - uid, payload["sfreq"], self._sfreq) - if self._vertex_indices is None: - self._vertex_indices = payload["vertex_indices"] - self._source_coords = payload["source_coords"] - if self.config.has_vertex_filter: - logger.info("Vertex filter: %d vertices retained", - len(self._vertex_indices)) - self._subject_groups[uid] = payload["group"] - self._subject_data[uid] = payload["subject_data"] - self._band_power_rows.extend(payload["band_power_rows"]) - self._feature_rows.extend(payload["feature_rows"]) - - def aggregate(self) -> None: - """Export CSVs.""" - data_dir = self.output_dir / "data" - - # Band power CSV - band_df = pd.DataFrame(self._band_power_rows) - if band_df.empty: - logger.warning("No vertex cluster band power data collected") - return - band_df.to_csv(data_dir / "vertex_cluster_values.csv", index=False) - logger.info("Exported vertex_cluster_values.csv (%d rows)", len(band_df)) - - # Features CSV - feat_df = pd.DataFrame(self._feature_rows) - if not feat_df.empty: - feat_df.to_csv(data_dir / "vertex_cluster_features.csv", index=False) - logger.info("Exported vertex_cluster_features.csv (%d rows)", len(feat_df)) - - # Source coordinates CSV - if self._source_coords is not None: - coords_df = pd.DataFrame( - self._source_coords, - columns=["x", "y", "z"], - ) - coords_df.index.name = "vertex_idx" - coords_df.to_csv(data_dir / "source_coords.csv") - logger.info("Exported source_coords.csv (%d rows)", len(coords_df)) - - def _run_test(self, data_a, data_b, coords): - """Run the configured correction method and return standardized results.""" - if self._correction_method == "tfce": - result = tfce_permutation_test( - data_a, data_b, coords, - n_perms=self._n_permutations, - E=self._tfce_E, H=self._tfce_H, dh=self._tfce_dh, - distance_mm=self._adjacency_distance, seed=42, - ) - return { - "t_map": result.t_map, - "hedges_g": result.hedges_g_map, - "p_corrected": result.p_corrected, - "tfce_scores": result.tfce_scores, - } - else: - result = cluster_permutation_test( - data_a, data_b, coords, - n_perms=self._n_permutations, - threshold=self._cluster_threshold, - distance_mm=self._adjacency_distance, seed=42, - ) - return { - "t_map": result.t_map, - "hedges_g": hedges_g(data_a, data_b), - "p_map": result.p_map, - "cluster_labels": result.cluster_labels, - "cluster_pvalues": result.cluster_pvalues, - "cluster_stats": result.cluster_stats, - } - - def statistics(self) -> None: - """Run voxel-wise t-tests with cluster permutation or TFCE correction.""" - if self._source_coords is None: - logger.error("No source coordinates — cannot run statistics") - return - - coords = self._source_coords - tbl_dir = self.tbl_dir - data_dir = self.output_dir / "data" - is_tfce = self._correction_method == "tfce" - - logger.info("Correction method: %s", self._correction_method) - - all_voxelwise = [] - all_cluster = [] - - # Per-vertex ROI labels (for the anatomical coverage of each significant - # cluster). Computed once — coords are fixed across all contrasts/bands. - vertex_rois = self._label_vertex_regions(self._source_coords) - - # Per-contrast figure state (keyed by contrast name) so EVERY declared - # contrast gets its own glass brains — not just whichever ran last. - self._band_cluster_results = {} - self._feature_cluster_results = {} - self._contrast_labels = {} - - for contrast in self._pairwise_contrasts(): - group_a_uids = [ - uid for uid, g in self._subject_groups.items() - if g == contrast.group_a - ] - group_b_uids = [ - uid for uid, g in self._subject_groups.items() - if g == contrast.group_b - ] - - if not group_a_uids or not group_b_uids: - logger.warning( - "Contrast %s: missing subjects (a=%d, b=%d)", - contrast.name, len(group_a_uids), len(group_b_uids), - ) - continue - - label_a = self.config.get_group_label(contrast.group_a) - label_b = self.config.get_group_label(contrast.group_b) - logger.info( - "Contrast '%s': %s (n=%d) vs %s (n=%d)", - contrast.name, label_a, len(group_a_uids), - label_b, len(group_b_uids), - ) - - # --- Band power metrics --- - band_cluster_results = {} - for band_name in self._selected_bands(): - for metric in ["relative", "absolute"]: - data_a = np.array([ - self._subject_data[uid]["band_power"][band_name][metric] - for uid in group_a_uids - ]) - data_b = np.array([ - self._subject_data[uid]["band_power"][band_name][metric] - for uid in group_b_uids - ]) - - res = self._run_test(data_a, data_b, coords) - - # Store for plotting (use relative power for band plots) - if metric == "relative": - plot_res = { - "t_map": res["t_map"], - "mean_a": data_a.mean(axis=0), - "mean_b": data_b.mean(axis=0), - } - if is_tfce: - plot_res["p_corrected"] = res["p_corrected"] - plot_res["tfce_scores"] = res["tfce_scores"] - plot_res["hedges_g_map"] = res["hedges_g"] - else: - plot_res["cluster_labels"] = res["cluster_labels"] - plot_res["cluster_pvalues"] = res["cluster_pvalues"] - band_cluster_results[band_name] = plot_res - - # Voxelwise stats rows - for vi in range(len(res["t_map"])): - row = { - "contrast": contrast.name, - "band": band_name, - "metric": metric, - "vertex_idx": vi, - "t": float(res["t_map"][vi]), - "hedges_g": float(res["hedges_g"][vi]), - } - if is_tfce: - row["tfce_score"] = float(res["tfce_scores"][vi]) - row["p_corrected"] = float(res["p_corrected"][vi]) - else: - row["p"] = float(res["p_map"][vi]) - row["cluster_id"] = int(res["cluster_labels"][vi]) - all_voxelwise.append(row) - - # Cluster summary rows (cluster method only) - if not is_tfce: - for ci, (cs, cp) in enumerate( - zip(res["cluster_stats"], res["cluster_pvalues"]), - start=1, - ): - n_verts = int(np.sum(res["cluster_labels"] == ci)) - mask = res["cluster_labels"] == ci - peak_t = float( - res["t_map"][mask][np.argmax(np.abs(res["t_map"][mask]))] - ) - all_cluster.append({ - "contrast": contrast.name, - "band": band_name, - "metric": metric, - "cluster_id": ci, - "n_vertices": n_verts, - "cluster_stat": float(cs), - "peak_t": peak_t, - "p_corrected": float(cp), - "region": self._cluster_region(vertex_rois, mask), - }) - - # --- Feature metrics (slope, peak_alpha) --- - feature_cluster_results = {} - for feat_name, feat_key in [ - ("spectral_slope", "slope"), - ("peak_alpha", "peak_alpha"), - ]: - data_a = np.array([ - self._subject_data[uid][feat_key] for uid in group_a_uids - ]) - data_b = np.array([ - self._subject_data[uid][feat_key] for uid in group_b_uids - ]) - - res = self._run_test(data_a, data_b, coords) - - feat_res = { - "t_map": res["t_map"], - "mean_a": data_a.mean(axis=0), - "mean_b": data_b.mean(axis=0), - } - if is_tfce: - feat_res["p_corrected"] = res["p_corrected"] - feat_res["tfce_scores"] = res["tfce_scores"] - feat_res["hedges_g_map"] = res["hedges_g"] - else: - feat_res["cluster_labels"] = res["cluster_labels"] - feat_res["cluster_pvalues"] = res["cluster_pvalues"] - feature_cluster_results[feat_name] = feat_res - - for vi in range(len(res["t_map"])): - row = { - "contrast": contrast.name, - "band": feat_name, - "metric": feat_name, - "vertex_idx": vi, - "t": float(res["t_map"][vi]), - "hedges_g": float(res["hedges_g"][vi]), - } - if is_tfce: - row["tfce_score"] = float(res["tfce_scores"][vi]) - row["p_corrected"] = float(res["p_corrected"][vi]) - else: - row["p"] = float(res["p_map"][vi]) - row["cluster_id"] = int(res["cluster_labels"][vi]) - all_voxelwise.append(row) - - if not is_tfce: - for ci, (cs, cp) in enumerate( - zip(res["cluster_stats"], res["cluster_pvalues"]), - start=1, - ): - n_verts = int(np.sum(res["cluster_labels"] == ci)) - mask = res["cluster_labels"] == ci - peak_t = float( - res["t_map"][mask][np.argmax(np.abs(res["t_map"][mask]))] - ) - all_cluster.append({ - "contrast": contrast.name, - "band": feat_name, - "metric": feat_name, - "cluster_id": ci, - "n_vertices": n_verts, - "cluster_stat": float(cs), - "peak_t": peak_t, - "p_corrected": float(cp), - "region": self._cluster_region(vertex_rois, mask), - }) - - # Store results for figures phase, keyed by contrast. - self._band_cluster_results[contrast.name] = band_cluster_results - self._feature_cluster_results[contrast.name] = feature_cluster_results - self._contrast_labels[contrast.name] = (label_a, label_b) - - # Export CSVs - if all_voxelwise: - vox_df = pd.DataFrame(all_voxelwise) - vox_df.to_csv(tbl_dir / "voxelwise_stats.csv", index=False) - logger.info("Exported voxelwise_stats.csv (%d rows)", len(vox_df)) - - if is_tfce: - # Log TFCE summary - if all_voxelwise: - vox_df = pd.DataFrame(all_voxelwise) - sig = vox_df[vox_df["p_corrected"] < 0.05] - if len(sig) > 0: - logger.info("TFCE significant vertices (p<0.05): %d", len(sig)) - for band in sig["band"].unique(): - n = len(sig[sig["band"] == band]) - logger.info(" %s: %d vertices", band, n) - else: - logger.info("TFCE: no significant vertices at p<0.05") - elif all_cluster: - clust_df = pd.DataFrame(all_cluster) - clust_df.to_csv(tbl_dir / "cluster_results.csv", index=False) - logger.info("Exported cluster_results.csv (%d rows)", len(clust_df)) - - sig_clusters = clust_df[clust_df["p_corrected"] < 0.05] - if len(sig_clusters) > 0: - logger.info( - "Significant clusters (p<0.05): %d", len(sig_clusters), - ) - for _, row in sig_clusters.iterrows(): - logger.info( - " %s / %s (%s): %d vertices, peak t=%.2f, p=%.4f", - row["band"], row["metric"], row["contrast"], - row["n_vertices"], row["peak_t"], row["p_corrected"], - ) - else: - logger.info("No significant clusters at p<0.05") - - # --- Declarative hypotheses (hypothesis layer; additive, map+cluster contract) --- - # Run every declared hypothesis over per-subject vertex maps via the permutation - # adapter (pairwise contrast == legacy cluster test bit-exact). Manual control: - # --hypothesis NAME. Additive — the legacy per-contrast tables above are untouched. - from ..hypothesis import write_module_hypotheses_perm - - maps_by_cell = { - (band_name, metric): { - uid: self._subject_data[uid]["band_power"][band_name][metric] - for uid in self._subject_groups - } - for band_name in self._selected_bands() - for metric in ["relative", "absolute"] - } - wanted_hyp = self._selection.get("hypothesis") - write_module_hypotheses_perm( - maps_by_cell, self._subject_groups, coords, self.config, tbl_dir, - prefix="vertex_cluster", - n_perms=self._n_permutations, threshold=self._cluster_threshold, - distance_mm=self._adjacency_distance, - hypothesis=",".join(sorted(wanted_hyp)) if wanted_hyp else None, - atlas_dir=self._atlas_dir, - ) - - # Save full results dict for reuse - results_pkl = { - "band_cluster_results": getattr(self, "_band_cluster_results", {}), - "feature_cluster_results": getattr(self, "_feature_cluster_results", {}), - "source_coords": self._source_coords, - "subject_groups": self._subject_groups, - "correction_method": self._correction_method, - "config": { - "correction_method": self._correction_method, - "n_permutations": self._n_permutations, - "adjacency_distance_mm": self._adjacency_distance, - }, - } - if is_tfce: - results_pkl["config"]["tfce_E"] = self._tfce_E - results_pkl["config"]["tfce_H"] = self._tfce_H - results_pkl["config"]["tfce_dh"] = self._tfce_dh - else: - results_pkl["config"]["cluster_threshold"] = self._cluster_threshold - with open(data_dir / "vertex_cluster_results.pkl", "wb") as f: - pickle.dump(results_pkl, f) - logger.info("Saved vertex_cluster_results.pkl") - - def _load_state_from_disk(self) -> bool: - """Load saved state from pickle for --steps figures/summary support.""" - pkl_path = self.output_dir / "data" / "vertex_cluster_results.pkl" - if not pkl_path.exists(): - logger.warning("No saved state at %s; skipping figures", pkl_path) - return False - with open(pkl_path, "rb") as f: - saved = pickle.load(f) - self._band_cluster_results = saved.get("band_cluster_results", {}) - self._feature_cluster_results = saved.get("feature_cluster_results", {}) - self._source_coords = saved.get("source_coords") - # Restore contrast labels from config (pickle doesn't store them directly), - # keyed by contrast name to match the per-contrast results dicts. - self._contrast_labels = { - c.name: ( - self.config.get_group_label(c.group_a), - self.config.get_group_label(c.group_b), - ) - for c in self._pairwise_contrasts() - } - logger.info("Loaded vertex cluster state from %s", pkl_path) - return True - - def figures(self) -> None: - """Generate glass brain figures.""" - # Load from disk if in-memory state is missing (--steps support) - if self._source_coords is None or not getattr(self, "_band_cluster_results", {}): - if not self._load_state_from_disk(): - return - - if self._source_coords is None: - logger.warning("No source coordinates — skipping figures") - return - - coords = self._source_coords - fig_dir = self.fig_dir - is_tfce = self._correction_method == "tfce" - - # Results are keyed by contrast name -> {band/feature: result}. Render a - # glass-brain set per contrast, with the contrast in the filename. - band_by_contrast = getattr(self, "_band_cluster_results", {}) - feat_by_contrast = getattr(self, "_feature_cluster_results", {}) - labels_by_contrast = getattr(self, "_contrast_labels", {}) - - for contrast_name, band_results in band_by_contrast.items(): - safe_contrast = contrast_name.lower().replace(" ", "_") - group_labels = labels_by_contrast.get(contrast_name, ("Group A", "Group B")) - feature_results = feat_by_contrast.get(contrast_name, {}) - logger.info("Rendering glass brains for contrast '%s'", contrast_name) - - # Band power figures — only where the band has a significant cluster. - sig_band_results = {} - for band_name, res in band_results.items(): - if not _has_significant_cluster(res): - continue - sig_band_results[band_name] = res - safe_name = band_name.lower().replace(" ", "_") - plot_band_comparison( - coords=coords, - mean_a=res["mean_a"], - mean_b=res["mean_b"], - t_map=res["t_map"], - cluster_labels=res.get("cluster_labels"), - cluster_pvalues=res.get("cluster_pvalues"), - band_name=band_name, - group_labels=group_labels, - output_path=fig_dir / f"vertex_cluster_{safe_contrast}_{safe_name}.png", - p_corrected=res.get("p_corrected"), - ) - # TFCE score maps - if is_tfce and "tfce_scores" in res: - plot_glass_brain( - coords=coords, - values=res["tfce_scores"], - title=f"TFCE Scores — {band_name}", - output_path=fig_dir / f"tfce_scores_{safe_contrast}_{safe_name}.png", - cmap="RdBu_r", - ) - - # Feature figures — only where significant. - sig_feature_results = {} - for feat_name, res in feature_results.items(): - if not _has_significant_cluster(res): - continue - sig_feature_results[feat_name] = res - plot_band_comparison( - coords=coords, - mean_a=res["mean_a"], - mean_b=res["mean_b"], - t_map=res["t_map"], - cluster_labels=res.get("cluster_labels"), - cluster_pvalues=res.get("cluster_pvalues"), - band_name=feat_name, - group_labels=group_labels, - output_path=fig_dir / f"vertex_cluster_{safe_contrast}_{feat_name}.png", - p_corrected=res.get("p_corrected"), - ) - - # Summary figure — only when the contrast has ≥1 significant map. - all_results = {**sig_band_results, **sig_feature_results} - if all_results: - plot_vertex_cluster_summary( - band_results=all_results, - coords=coords, - output_path=fig_dir / f"vertex_cluster_{safe_contrast}_summary.png", - group_labels=group_labels, - ) - - def summary(self) -> None: - """Call R script for formatted tables and ANALYSIS_SUMMARY.md.""" - data_dir = self.output_dir / "data" - - # Write study config for R - config_path = data_dir / "study_config.yaml" - config_data = self._r_config_data() - if self._sfreq is not None: - config_data["sfreq"] = self._sfreq - with open(config_path, "w") as f: - yaml.dump(config_data, f, default_flow_style=False) - - # Try R script - try: - r_dir = _find_r_script_dir() - except FileNotFoundError as e: - logger.warning(str(e)) - self._write_python_summary() - return - - r_script = r_dir / "vertex_cluster_analysis.R" - if not r_script.exists(): - logger.warning("R script not found: %s — writing Python summary", r_script) - self._write_python_summary() - return - - cmd = [ - "Rscript", str(r_script), - "--data-dir", str(data_dir), - "--config", str(config_path), - "--output-dir", str(self.output_dir), - "--fig-dir", str(self.fig_dir), - "--tbl-dir", str(self.tbl_dir), - ] - cmd.extend(self._r_no_figures_flags()) - - logger.info("Calling R: %s", " ".join(cmd)) - try: - result = subprocess.run( - cmd, capture_output=True, text=True, timeout=self._r_timeout, - ) - if result.stdout: - for line in result.stdout.strip().split("\n"): - logger.info("[R] %s", line) - if result.stderr: - for line in result.stderr.strip().split("\n"): - if line.strip(): - logger.info("[R] %s", line) - if result.returncode != 0: - self._r_step_failed("R script failed with exit code %d", result.returncode) - self._write_python_summary() - except FileNotFoundError: - logger.warning("Rscript not found — writing Python summary") - self._write_python_summary() - except subprocess.TimeoutExpired: - self._r_step_failed("R script timed out after %s s", self._r_timeout) - self._write_python_summary() - - def _write_python_summary(self) -> None: - """Fallback summary when R is not available.""" - tbl_dir = self.tbl_dir - is_tfce = self._correction_method == "tfce" - - if is_tfce: - correction_desc = ( - "TFCE (Smith & Nichols, 2009) was applied to vertex-level band power maps. " - "TFCE integrates cluster extent and height across all possible thresholds " - f"(E={self._tfce_E}, H={self._tfce_H}, dh={self._tfce_dh}), " - "eliminating the need for an arbitrary cluster-forming threshold. " - "Statistical significance was assessed via permutation testing." - ) - method_label = "Vertex-level spectral analysis with TFCE correction" - else: - correction_desc = ( - "Group differences were tested using independent-samples t-tests at each vertex, " - "with cluster-based permutation correction for multiple comparisons " - "(Maris & Oostenveld, 2007)." - ) - method_label = "Vertex-level spectral analysis with cluster permutation testing" - - lines = [ - "# Vertex Cluster Analysis Summary", - "", - f"**Study**: {self.config.name}", - f"**Analysis**: {method_label}", - f"**Correction method**: {self._correction_method}", - f"**Permutations**: {self._n_permutations}", - f"**Adjacency distance**: {self._adjacency_distance} mm", - ] - if not is_tfce: - lines.append(f"**Cluster threshold**: t = {self._cluster_threshold}") - lines.extend(["", "## Methods", ""]) - lines.append( - "Power spectral density was computed for each source vertex using Welch's method " - "(2-second windows, 50% overlap). Band power (absolute and relative), " - "spectral slope, and peak alpha frequency were extracted per vertex. " - + correction_desc - ) - lines.append("") - - # Results section - voxelwise_csv = tbl_dir / "voxelwise_stats.csv" - cluster_csv = tbl_dir / "cluster_results.csv" - - if is_tfce and voxelwise_csv.exists(): - vox_df = pd.read_csv(voxelwise_csv) - sig = vox_df[vox_df["p_corrected"] < 0.05] - - lines.append("## Results") - lines.append("") - if len(sig) > 0: - lines.append(f"**{len(sig)} significant vertices** (TFCE p < 0.05):") - lines.append("") - for band in vox_df["band"].unique(): - for metric in vox_df["metric"].unique(): - subset = sig[(sig["band"] == band) & (sig["metric"] == metric)] - total = len(vox_df[ - (vox_df["band"] == band) & (vox_df["metric"] == metric) - ]) - lines.append(f"- **{band}** ({metric}): {len(subset)}/{total} vertices") - lines.append("") - else: - lines.append("No significant vertices at TFCE p < 0.05.") - lines.append("") - - elif not is_tfce and cluster_csv.exists(): - clust_df = pd.read_csv(cluster_csv) - sig = clust_df[clust_df["p_corrected"] < 0.05] - - lines.append("## Results") - lines.append("") - - if len(sig) > 0: - lines.append(f"**{len(sig)} significant clusters** (p < 0.05):") - lines.append("") - lines.append("| Band/Metric | Metric | Vertices | Peak t | p_corrected |") - lines.append("|-------------|--------|----------|--------|-------------|") - for _, row in sig.iterrows(): - lines.append( - f"| {row['band']} | {row['metric']} | {row['n_vertices']} | " - f"{row['peak_t']:.2f} | {row['p_corrected']:.4f} |" - ) - lines.append("") - else: - lines.append("No significant clusters at p < 0.05.") - lines.append("") - - lines.append(f"Total clusters tested: {len(clust_df)}") - lines.append("") - - lines.append("## Output Files") - lines.append("") - lines.append("- `data/vertex_cluster_values.csv` — per-subject per-vertex band power") - lines.append("- `data/vertex_cluster_features.csv` — per-subject per-vertex slope, peak alpha") - lines.append("- `data/source_coords.csv` — vertex coordinates (mm)") - lines.append("- `tables/voxelwise_stats.csv` — per-vertex statistics") - if not is_tfce: - lines.append("- `tables/cluster_results.csv` — cluster summaries with corrected p-values") - lines.append("- `figures/vertex_cluster_*.png` — glass brain visualizations") - if is_tfce: - lines.append("- `figures/tfce_scores_*.png` — TFCE score glass brains") - lines.append("") - - summary_path = self.output_dir / "ANALYSIS_SUMMARY.md" - summary_path.write_text("\n".join(lines)) - logger.info("Wrote %s", summary_path) diff --git a/src/source_analytics/analyses/vertex_connectivity_analysis.py b/src/source_analytics/analyses/vertex_connectivity_analysis.py deleted file mode 100644 index adcdb95..0000000 --- a/src/source_analytics/analyses/vertex_connectivity_analysis.py +++ /dev/null @@ -1,507 +0,0 @@ -"""Vertex-level connectivity analysis: multiple metrics + FCD maps. - -Computes all-to-all connectivity between source vertices using one or more -metrics, derives Functional Connectivity Density (FCD) maps, and tests for -group differences using cluster-based permutation testing. -""" - -from __future__ import annotations - -import logging -import pickle -import subprocess -from pathlib import Path - -import numpy as np -import pandas as pd -import yaml - -from ..config import StudyConfig -from ..io.discovery import SubjectInfo -from ..io.loader import SubjectLoader -from ..spectral.vertex_connectivity import ( - compute_vertex_connectivity_matrix, - compute_vertex_connectivity_matrix_multi, - compute_vertex_connectivity_matrix_epochs, - compute_vertex_connectivity_matrix_epochs_multi, - compute_fcd, - FCD_CENTER, -) -from ..spectral.epoch_sampler import sample_epochs -from ..stats.cluster_permutation import ( - cluster_permutation_test, - has_significant_cluster as _has_significant_cluster, - hedges_g, -) -from ..viz.glass_brain import plot_glass_brain, plot_band_comparison -from .base import BaseAnalysis, find_r_script_dir - -logger = logging.getLogger(__name__) - - -class VertexConnectivityAnalysis(BaseAnalysis): - """All-to-all vertex connectivity with FCD mapping (multi-metric).""" - - name = "vertex_connectivity" - SELECTABLE = {"metric": "connectivity metric", "band": "frequency band", - "hypothesis": "declared hypothesis"} - - def __init__(self, config: StudyConfig, output_dir: Path): - super().__init__(config, output_dir) - self._fcd_rows: list[dict] = [] - self._source_coords: np.ndarray | None = None - self._vertex_indices: np.ndarray | None = None - self._sfreq: float | None = None - self._subject_data: dict[str, dict] = {} - self._subject_groups: dict[str, str] = {} - # conn_matrices: uid -> band -> metric -> matrix - self._conn_matrices: dict[str, dict[str, dict[str, np.ndarray]]] = {} - - # Config - vc_cfg = config.raw.get("vertex_connectivity", {}) - # Support both single metric (legacy) and multi-metric list - metrics_cfg = vc_cfg.get("metrics") - if metrics_cfg is not None: - self._metrics = list(metrics_cfg) - else: - self._metrics = [vc_cfg.get("metric", "imag_coherence")] - - self._fcd_threshold = float(vc_cfg.get("fcd_threshold", 0.05)) - self._n_permutations = int(vc_cfg.get("n_permutations", 1000)) - - wb_cfg = config.vertex - self._adjacency_distance = float(wb_cfg.get("adjacency_distance_mm", 5.0)) - self._cluster_threshold = float(wb_cfg.get("cluster_threshold", 2.0)) - - # Global epoch_sampling → vertex: block → per-analysis block (see base). - self._epoch_config = self._vertex_epoch_config() - - self._cluster_results: dict = {} - - def setup(self) -> None: - # Restrict to --metric / --select metric=... if requested (shared STFT - # pass is preserved — fewer metrics are emitted from the same pass). - self._metrics = self._select("metric", self._metrics) - self._fcd_rows.clear() - self._subject_data.clear() - self._subject_groups.clear() - self._conn_matrices.clear() - self._source_coords = None - self._vertex_indices = None - self._cluster_results.clear() - - def _compute_subject(self, subject: SubjectInfo): - """Pure per-subject connectivity + FCD compute (parallel-safe).""" - loader = SubjectLoader(subject.data_dir) - uid = f"{subject.group}_{subject.subject_id}" - - # Load signed data for phase-preserving connectivity - stc_data = loader.load_source_timecourses() - sfreq = loader.load_sfreq() - coords = loader.load_source_coords() - - mask = self.config.get_vertex_mask(coords) - vertex_indices = np.where(mask)[0] - source_coords = coords[mask] - stc_data = stc_data[vertex_indices] - - subject_fcd: dict[str, dict[str, np.ndarray]] = {} - subject_conn: dict[str, dict[str, np.ndarray]] = {} - fcd_rows: list[dict] = [] - - use_multi = len(self._metrics) > 1 - - for band_name, (fmin, fmax) in self._selected_bands().items(): - logger.info( - " Computing %s connectivity (%s)...", - band_name, ", ".join(self._metrics), - ) - - if use_multi: - # Compute all metrics in a single pass - if self._epoch_config is not None: - epochs = sample_epochs( - stc_data, sfreq, - epoch_duration_sec=self._epoch_config.get( - "epoch_duration_sec", 2.0, - ), - n_epochs=self._epoch_config.get("n_epochs", 80), - seed=self._epoch_config.get("seed", 42), - n_bootstrap=self._epoch_config.get("n_bootstrap", 1), - ) - conn_results = compute_vertex_connectivity_matrix_epochs_multi( - epochs, sfreq, (fmin, fmax), metrics=self._metrics, - ) - else: - conn_results = compute_vertex_connectivity_matrix_multi( - stc_data, sfreq, (fmin, fmax), metrics=self._metrics, - ) - else: - # Single metric — use original function - metric = self._metrics[0] - if self._epoch_config is not None: - epochs = sample_epochs( - stc_data, sfreq, - epoch_duration_sec=self._epoch_config.get( - "epoch_duration_sec", 2.0, - ), - n_epochs=self._epoch_config.get("n_epochs", 80), - seed=self._epoch_config.get("seed", 42), - n_bootstrap=self._epoch_config.get("n_bootstrap", 1), - ) - conn_mat = compute_vertex_connectivity_matrix_epochs( - epochs, sfreq, (fmin, fmax), metric=metric, - ) - else: - conn_mat = compute_vertex_connectivity_matrix( - stc_data, sfreq, (fmin, fmax), metric=metric, - ) - conn_results = {metric: conn_mat} - - band_fcd = {} - band_conn = {} - for metric, conn_mat in conn_results.items(): - fcd = compute_fcd( - conn_mat, threshold=self._fcd_threshold, - center=FCD_CENTER.get(metric), - ) - band_fcd[metric] = fcd - band_conn[metric] = conn_mat - - n_vertices = len(fcd) - for vi in range(n_vertices): - fcd_rows.append({ - "subject": uid, - "group": subject.group, - "vertex_idx": int(vertex_indices[vi]), - "band": band_name, - "metric": metric, - "fcd": float(fcd[vi]), - }) - - subject_fcd[band_name] = band_fcd - subject_conn[band_name] = band_conn - - return { - "uid": uid, "group": subject.group, "sfreq": float(sfreq), - "vertex_indices": vertex_indices, "source_coords": source_coords, - "subject_fcd": subject_fcd, "subject_conn": subject_conn, - "fcd_rows": fcd_rows, - } - - def _merge_subject(self, payload) -> None: - uid = payload["uid"] - if self._sfreq is None: - self._sfreq = payload["sfreq"] - if self._vertex_indices is None: - self._vertex_indices = payload["vertex_indices"] - self._source_coords = payload["source_coords"] - if self.config.has_vertex_filter: - logger.info("Vertex filter: %d vertices retained", - len(self._vertex_indices)) - self._subject_groups[uid] = payload["group"] - self._subject_data[uid] = {"fcd": payload["subject_fcd"]} - self._conn_matrices[uid] = payload["subject_conn"] - self._fcd_rows.extend(payload["fcd_rows"]) - - def aggregate(self) -> None: - data_dir = self.output_dir / "data" - - fcd_df = pd.DataFrame(self._fcd_rows) - if fcd_df.empty: - logger.warning("No vertex connectivity data collected") - return - fcd_df.to_csv(data_dir / "vertex_fcd.csv", index=False) - logger.info("Exported vertex_fcd.csv (%d rows)", len(fcd_df)) - - if self._source_coords is not None: - coords_df = pd.DataFrame( - self._source_coords, columns=["x", "y", "z"], - ) - coords_df.index.name = "vertex_idx" - coords_df.to_csv(data_dir / "source_coords.csv") - - # Save connectivity matrices for downstream use (vertex_network) - if self._conn_matrices: - pkl_path = data_dir / "vertex_connectivity_matrices.pkl" - with open(pkl_path, "wb") as f: - pickle.dump(self._conn_matrices, f) - logger.info("Saved connectivity matrices to %s", pkl_path) - - def _reload_maps_from_disk(self) -> bool: - """Reconstruct per-subject FCD maps + coords + groups from persisted CSVs, - so statistics (and thence figures) is regenerable via --steps without - re-running the expensive process/aggregate. Mirrors vertex_cluster's - reload discipline — figures must be a pure function of persisted data.""" - data_dir = self.output_dir / "data" - fcd_csv = data_dir / "vertex_fcd.csv" - coords_csv = data_dir / "source_coords.csv" - if not fcd_csv.exists(): - logger.warning("No persisted FCD at %s; cannot reload", fcd_csv) - return False - if coords_csv.exists(): - cdf = pd.read_csv(coords_csv) - self._source_coords = cdf[["x", "y", "z"]].to_numpy(dtype=float) - df = pd.read_csv(fcd_csv) - self._subject_data = {} - self._subject_groups = {} - for (uid, group), g in df.groupby(["subject", "group"], sort=False): - self._subject_groups[uid] = group - fcd: dict = {} - for (band, metric), gg in g.groupby(["band", "metric"], sort=False): - arr = gg.sort_values("vertex_idx")["fcd"].to_numpy(dtype=float) - fcd.setdefault(band, {})[metric] = arr - self._subject_data[uid] = {"fcd": fcd} - logger.info("Reloaded %d subjects' FCD maps from %s", - len(self._subject_data), fcd_csv) - return True - - def statistics(self) -> None: - if not self._subject_data: - self._reload_maps_from_disk() - if self._source_coords is None: - logger.error("No source coordinates — cannot run statistics") - return - - coords = self._source_coords - tbl_dir = self.tbl_dir - all_stats = [] - - # Per-vertex t-maps drive the glass-brain figures; iterate the declared - # pairwise contrast hypotheses directly (no config.contrasts bridge). - for contrast in self._pairwise_contrasts(): - group_a_uids = [ - uid for uid, g in self._subject_groups.items() - if g == contrast.group_a - ] - group_b_uids = [ - uid for uid, g in self._subject_groups.items() - if g == contrast.group_b - ] - - if not group_a_uids or not group_b_uids: - continue - - label_a = self.config.get_group_label(contrast.group_a) - label_b = self.config.get_group_label(contrast.group_b) - - for band_name in self._selected_bands(): - for metric in self._metrics: - data_a = np.array([ - self._subject_data[uid]["fcd"][band_name][metric] - for uid in group_a_uids - if band_name in self._subject_data.get(uid, {}).get("fcd", {}) - and metric in self._subject_data[uid]["fcd"].get(band_name, {}) - ]) - data_b = np.array([ - self._subject_data[uid]["fcd"][band_name][metric] - for uid in group_b_uids - if band_name in self._subject_data.get(uid, {}).get("fcd", {}) - and metric in self._subject_data[uid]["fcd"].get(band_name, {}) - ]) - - if data_a.size == 0 or data_b.size == 0: - continue - - result = cluster_permutation_test( - data_a, data_b, coords, - n_perms=self._n_permutations, - threshold=self._cluster_threshold, - distance_mm=self._adjacency_distance, - seed=42, - ) - - g_map = hedges_g(data_a, data_b) - - key = f"{contrast.name}_{band_name}_{metric}" - self._cluster_results[key] = { - "result": result, - "mean_a": data_a.mean(axis=0), - "mean_b": data_b.mean(axis=0), - "group_labels": (label_a, label_b), - "contrast": contrast.name, - "band": band_name, - "metric": metric, - } - - for vi in range(len(result.t_map)): - all_stats.append({ - "contrast": contrast.name, - "band": band_name, - "metric": metric, - "vertex_idx": vi, - "fcd_a": float(data_a.mean(axis=0)[vi]), - "fcd_b": float(data_b.mean(axis=0)[vi]), - "t": float(result.t_map[vi]), - "p": float(result.p_map[vi]), - "hedges_g": float(g_map[vi]), - "cluster_id": int(result.cluster_labels[vi]), - }) - - if all_stats: - stats_df = pd.DataFrame(all_stats) - stats_df.to_csv( - tbl_dir / "vertex_connectivity_stats.csv", index=False, - ) - logger.info( - "Exported vertex_connectivity_stats.csv (%d rows)", - len(stats_df), - ) - - # --- Declarative hypotheses (hypothesis layer; additive, map+cluster) --- - from ..hypothesis import write_module_hypotheses_perm - - if self._source_coords is not None and self._subject_groups: - maps_by_cell = { - (band_name, metric): { - uid: self._subject_data[uid]["fcd"][band_name][metric] - for uid in self._subject_groups - } - for band_name in self._selected_bands() - for metric in self._metrics - } - wanted_hyp = self._selection.get("hypothesis") - write_module_hypotheses_perm( - maps_by_cell, self._subject_groups, self._source_coords, self.config, - tbl_dir, prefix="vertex_connectivity", - n_perms=self._n_permutations, threshold=self._cluster_threshold, - distance_mm=self._adjacency_distance, - hypothesis=",".join(sorted(wanted_hyp)) if wanted_hyp else None, - atlas_dir=self._atlas_dir, - ) - - # Persist the cluster-test results so figures() can regenerate from disk. - self._save_cluster_state() - - def figures(self) -> None: - if not self._cluster_results: - self._load_cluster_state() - if self._source_coords is None: - return - - coords = self._source_coords - fig_dir = self.fig_dir - - # Emit a glass brain ONLY where the contrast has a significant cluster — - # one figure per (contrast, band, metric), named with the contrast so - # nothing is overwritten. Non-significant cells produce no figure. - n = 0 - for key, info in self._cluster_results.items(): - result = info["result"] - if not _has_significant_cluster(result): - continue - contrast = info.get("contrast", key) - band = info["band"] - metric = info.get("metric", "imag_coherence") - safe_name = f"{contrast}_{band}_{metric}".lower().replace(" ", "_") - group_labels = info["group_labels"] - - plot_band_comparison( - coords=coords, - mean_a=info["mean_a"], - mean_b=info["mean_b"], - t_map=result.t_map, - cluster_labels=result.cluster_labels, - cluster_pvalues=result.cluster_pvalues, - band_name=f"FCD ({metric}) — {band} — {contrast}", - group_labels=group_labels, - output_path=fig_dir / f"fcd_{safe_name}.png", - ) - n += 1 - logger.info("vertex_connectivity: %d significant-cluster figures", n) - - def summary(self) -> None: - data_dir = self.output_dir / "data" - - config_path = data_dir / "study_config.yaml" - config_data = self._r_config_data() - if self._sfreq is not None: - config_data["sfreq"] = self._sfreq - with open(config_path, "w") as f: - yaml.dump(config_data, f, default_flow_style=False) - - try: - r_dir = find_r_script_dir() - r_script = r_dir / "vertex_connectivity_analysis.R" - if r_script.exists(): - cmd = [ - "Rscript", str(r_script), - "--data-dir", str(data_dir), - "--config", str(config_path), - "--output-dir", str(self.output_dir), - "--fig-dir", str(self.fig_dir), - "--tbl-dir", str(self.tbl_dir), - ] - cmd.extend(self._r_no_figures_flags()) - result = subprocess.run( - cmd, capture_output=True, text=True, timeout=self._r_timeout, - ) - if result.returncode == 0: - return - except (FileNotFoundError, subprocess.TimeoutExpired): - pass - - self._write_python_summary() - - def _write_python_summary(self) -> None: - tbl_dir = self.tbl_dir - - lines = [ - "# Vertex Connectivity Analysis Summary", - "", - f"**Study**: {self.config.name}", - f"**Analysis**: All-to-all vertex connectivity + FCD", - f"**Metrics**: {', '.join(self._metrics)}", - f"**FCD threshold**: {self._fcd_threshold}", - f"**Permutations**: {self._n_permutations}", - "", - "## Methods", - "", - f"Connectivity was computed between all pairs of source vertices " - f"using {', '.join(self._metrics)}. Functional Connectivity " - "Density (FCD) was derived by counting the fraction of connections " - f"exceeding {self._fcd_threshold} per vertex. Group differences in FCD " - "were tested using cluster-based permutation testing.", - "", - ] - - if self._epoch_config is not None: - lines.append( - f"**Epoch sampling**: {self._epoch_config.get('n_epochs', 80)} epochs " - f"of {self._epoch_config.get('epoch_duration_sec', 2.0)}s" - ) - lines.append("") - - stats_csv = tbl_dir / "vertex_connectivity_stats.csv" - if stats_csv.exists(): - stats_df = pd.read_csv(stats_csv) - lines.append("## Results") - lines.append("") - group_cols = ["band"] - if "metric" in stats_df.columns: - group_cols.append("metric") - for keys, sub in stats_df.groupby(group_cols): - if isinstance(keys, str): - label = keys - else: - label = " / ".join(str(k) for k in keys) - n_sig = len(sub[sub["p"] < 0.05]) - lines.append( - f"- **{label}**: {n_sig}/{len(sub)} vertices nominally significant" - ) - lines.append("") - - lines.extend([ - "## Output Files", - "", - "- `data/vertex_fcd.csv` — per-subject per-vertex FCD values", - "- `data/vertex_connectivity_matrices.pkl` — full connectivity matrices", - "- `data/source_coords.csv` — vertex coordinates (mm)", - "- `tables/vertex_connectivity_stats.csv` — FCD statistics", - "- `figures/fcd_*.png` — FCD glass brain maps", - "", - ]) - - summary_path = self.output_dir / "ANALYSIS_SUMMARY.md" - summary_path.write_text("\n".join(lines)) - logger.info("Wrote %s", summary_path) diff --git a/src/source_analytics/analyses/vertex_cross_freq_analysis.py b/src/source_analytics/analyses/vertex_cross_freq_analysis.py deleted file mode 100644 index 941c1d9..0000000 --- a/src/source_analytics/analyses/vertex_cross_freq_analysis.py +++ /dev/null @@ -1,351 +0,0 @@ -"""Vertex-level cross-frequency coupling: local PAC maps + AAC/PPC (full resolution). - -Vertex mirror of ``roi_cross_freq``. ``--metric`` selects the measure; all run at -**full vertex resolution** (no parcellation — see MS2_CONNECTIVITY_PLAN.md D1): - -- **pac** — local within-vertex PAC z-map (one Modulation Index per vertex per - cross-frequency band pair → whole-brain map). The Tier-A measure for the - source spatial advantage in cross-frequency coupling. -- **aac** — cross-frequency power–power coupling, vertex×vertex per band pair. -- **ppc** — n:m phase–phase coupling (PLF + surrogate z), vertex×vertex per pair. - -For AAC/PPC the per-vertex statistical map is the node coupling strength (mean -off-diagonal); the full matrices are stored for downstream edge-level analysis. -Group differences in the per-vertex maps are tested with cluster-based -permutation testing. References: see CONNECTIVITY_METHODS.md. -""" - -from __future__ import annotations - -import logging -import pickle - -import numpy as np -import pandas as pd -import yaml - -from ..config import StudyConfig -from ..io.discovery import SubjectInfo -from ..io.loader import SubjectLoader -from ..spectral.pac import get_valid_pac_pairs, compute_local_pac_vertices -from ..spectral.cross_freq import compute_aac, compute_ppc -from .roi_cross_freq_analysis import _nm_ratio -from ..stats.cluster_permutation import ( - cluster_permutation_test, - has_significant_cluster as _has_significant_cluster, - hedges_g, -) -from ..viz.glass_brain import plot_band_comparison -from pathlib import Path -from .base import BaseAnalysis - -logger = logging.getLogger(__name__) - - -def _node_strength(mat: np.ndarray) -> np.ndarray: - """Per-node coupling strength = mean off-diagonal value over the row.""" - n = mat.shape[0] - if n <= 1: - return np.zeros(n) - off = mat.copy() - np.fill_diagonal(off, 0.0) - return off.sum(axis=1) / (n - 1) - - -class VertexCrossFreqAnalysis(BaseAnalysis): - """Vertex-level cross-frequency coupling (local PAC, AAC, n:m PPC).""" - - name = "vertex_cross_freq" - SELECTABLE = {"metric": "coupling measure", "band": "frequency band", - "hypothesis": "declared hypothesis"} - - _CROSS_FREQ_METRICS = ["pac", "aac", "ppc"] - - def __init__(self, config: StudyConfig, output_dir: Path): - super().__init__(config, output_dir) - self._map_rows: list[dict] = [] - self._sfreq: float | None = None - self._source_coords: np.ndarray | None = None - self._vertex_indices: np.ndarray | None = None - self._subject_groups: dict[str, str] = {} - # uid -> "{metric}|{freq_pair}" -> per-vertex map - self._subject_maps: dict[str, dict[str, np.ndarray]] = {} - # uid -> "{metric}|{freq_pair}" -> full matrix (aac/ppc only) - self._matrices: dict[str, dict[str, np.ndarray]] = {} - self._cluster_results: dict = {} - - cfg = config.raw.get(self.name, {}) - self._metrics: list[str] = list(self._CROSS_FREQ_METRICS) - self._ppc_surrogates = int(cfg.get("ppc_surrogates", 100)) - self._pac_surrogates = int(cfg.get("pac_surrogates", 100)) - self._n_permutations = int(cfg.get("n_permutations", 1000)) - wb_cfg = config.vertex - self._adjacency_distance = float(wb_cfg.get("adjacency_distance_mm", 5.0)) - self._cluster_threshold = float(wb_cfg.get("cluster_threshold", 2.0)) - - def setup(self) -> None: - self._metrics = self._select("metric", self._CROSS_FREQ_METRICS) - self._map_rows.clear() - self._subject_groups.clear() - self._subject_maps.clear() - self._matrices.clear() - self._cluster_results.clear() - self._source_coords = None - self._vertex_indices = None - - # ------------------------------------------------------------ processing - def _compute_subject(self, subject: SubjectInfo): - """Pure per-subject cross-frequency compute (parallel-safe): loads the - subject, computes its maps/matrices, returns a payload. No self mutation.""" - loader = SubjectLoader(subject.data_dir) - uid = f"{subject.group}_{subject.subject_id}" - - stc_data = loader.load_source_timecourses() - sfreq = loader.load_sfreq() - coords = loader.load_source_coords() - - mask = self.config.get_vertex_mask(coords) - vertex_indices = np.where(mask)[0] - source_coords = coords[mask] - stc_data = stc_data[vertex_indices] - - pairs = get_valid_pac_pairs(self._selected_bands()) - if not pairs: - logger.warning("No valid cross-frequency band pairs for %s", subject.subject_id) - return None - - bands = self.config.bands - subject_maps: dict[str, np.ndarray] = {} - matrices: dict[str, np.ndarray] = {} - map_rows: list[dict] = [] - - for phase_band, amp_band in pairs: - band_x, band_y = bands[phase_band], bands[amp_band] - freq_pair = f"{phase_band}-{amp_band}" - - for metric in self._metrics: - if metric == "pac": - z, _mi = compute_local_pac_vertices( - stc_data, sfreq, band_x, band_y, - n_surrogates=self._pac_surrogates, - ) - vmap = z - elif metric == "aac": - mat = compute_aac(stc_data, sfreq, band_x, band_y) - matrices[f"aac|{freq_pair}"] = mat - vmap = _node_strength(mat) - else: # ppc - if self._ppc_surrogates > 0: - plf, _z = compute_ppc( - stc_data, sfreq, band_x, band_y, - n=_nm_ratio(band_x, band_y)[0], m=1, - n_surrogates=self._ppc_surrogates, seed=42, - ) - else: - plf = compute_ppc(stc_data, sfreq, band_x, band_y, - n=_nm_ratio(band_x, band_y)[0], m=1) - matrices[f"ppc|{freq_pair}"] = plf - vmap = _node_strength(plf) - - subject_maps[f"{metric}|{freq_pair}"] = vmap - for vi in range(len(vmap)): - map_rows.append({ - "subject": uid, "group": subject.group, - "vertex_idx": int(vertex_indices[vi]), - "metric": metric, "freq_pair": freq_pair, - "value": float(vmap[vi]), - }) - - return { - "uid": uid, "group": subject.group, "sfreq": float(sfreq), - "vertex_indices": vertex_indices, "source_coords": source_coords, - "subject_maps": subject_maps, "matrices": matrices, "map_rows": map_rows, - } - - def _merge_subject(self, payload) -> None: - uid = payload["uid"] - if self._sfreq is None: - self._sfreq = payload["sfreq"] - if self._vertex_indices is None: - self._vertex_indices = payload["vertex_indices"] - self._source_coords = payload["source_coords"] - if self.config.has_vertex_filter: - logger.info("Vertex filter: %d vertices retained", - len(self._vertex_indices)) - self._subject_groups[uid] = payload["group"] - self._subject_maps[uid] = payload["subject_maps"] - self._matrices[uid] = payload["matrices"] - self._map_rows.extend(payload["map_rows"]) - - # ----------------------------------------------------------- aggregate - def aggregate(self) -> None: - data_dir = self.output_dir / "data" - if not self._map_rows: - logger.warning("No vertex cross-frequency data collected") - return - pd.DataFrame(self._map_rows).to_csv( - data_dir / "vertex_cross_freq_maps.csv", index=False) - logger.info("Exported vertex_cross_freq_maps.csv (%d rows)", len(self._map_rows)) - - if self._source_coords is not None: - cdf = pd.DataFrame(self._source_coords, columns=["x", "y", "z"]) - cdf.index.name = "vertex_idx" - cdf.to_csv(data_dir / "source_coords.csv") - if any(self._matrices.values()): - with open(data_dir / "vertex_cross_freq_matrices.pkl", "wb") as f: - pickle.dump(self._matrices, f) - logger.info("Saved AAC/PPC matrices pkl") - - # ----------------------------------------------------------- statistics - def _reload_maps_from_disk(self) -> bool: - """Reconstruct per-subject CFC maps + coords + groups from persisted CSVs - so statistics/figures are regenerable via --steps (no reprocessing).""" - data_dir = self.output_dir / "data" - maps_csv = data_dir / "vertex_cross_freq_maps.csv" - coords_csv = data_dir / "source_coords.csv" - if not maps_csv.exists(): - logger.warning("No persisted maps at %s; cannot reload", maps_csv) - return False - if coords_csv.exists(): - self._source_coords = pd.read_csv(coords_csv)[["x", "y", "z"]].to_numpy(dtype=float) - df = pd.read_csv(maps_csv) - self._subject_maps = {} - self._subject_groups = {} - for (uid, group), g in df.groupby(["subject", "group"], sort=False): - self._subject_groups[uid] = group - m: dict = {} - for (metric, freq_pair), gg in g.groupby(["metric", "freq_pair"], sort=False): - gg = gg.sort_values("vertex_idx") - m[f"{metric}|{freq_pair}"] = gg["value"].to_numpy(dtype=float) - self._subject_maps[uid] = m - logger.info("Reloaded %d subjects' CFC maps from %s", - len(self._subject_maps), maps_csv) - return True - - def statistics(self) -> None: - if not self._subject_maps: - self._reload_maps_from_disk() - if self._source_coords is None: - logger.error("No source coordinates — cannot run statistics") - return - coords = self._source_coords - all_stats = [] - - # keys are "{metric}|{freq_pair}" — test each across each contrast - keys = sorted({k for m in self._subject_maps.values() for k in m}) - for contrast in self._pairwise_contrasts(): - uids_a = [u for u, g in self._subject_groups.items() if g == contrast.group_a] - uids_b = [u for u, g in self._subject_groups.items() if g == contrast.group_b] - if not uids_a or not uids_b: - continue - label_a = self.config.get_group_label(contrast.group_a) - label_b = self.config.get_group_label(contrast.group_b) - - for key in keys: - metric, freq_pair = key.split("|", 1) - data_a = np.array([self._subject_maps[u][key] for u in uids_a - if key in self._subject_maps.get(u, {})]) - data_b = np.array([self._subject_maps[u][key] for u in uids_b - if key in self._subject_maps.get(u, {})]) - if data_a.size == 0 or data_b.size == 0: - continue - - result = cluster_permutation_test( - data_a, data_b, coords, - n_perms=self._n_permutations, - threshold=self._cluster_threshold, - distance_mm=self._adjacency_distance, seed=42, - ) - g_map = hedges_g(data_a, data_b) - self._cluster_results[f"{contrast.name}_{key}"] = { - "result": result, "mean_a": data_a.mean(axis=0), - "mean_b": data_b.mean(axis=0), "group_labels": (label_a, label_b), - "metric": metric, "freq_pair": freq_pair, - } - for vi in range(len(result.t_map)): - all_stats.append({ - "contrast": contrast.name, "metric": metric, - "freq_pair": freq_pair, "vertex_idx": vi, - "value_a": float(data_a.mean(axis=0)[vi]), - "value_b": float(data_b.mean(axis=0)[vi]), - "t": float(result.t_map[vi]), "p": float(result.p_map[vi]), - "hedges_g": float(g_map[vi]), - "cluster_id": int(result.cluster_labels[vi]), - }) - - if all_stats: - pd.DataFrame(all_stats).to_csv( - self.tbl_dir / "vertex_cross_freq_stats.csv", index=False) - logger.info("Exported vertex_cross_freq_stats.csv (%d rows)", len(all_stats)) - - # --- Declarative hypotheses (hypothesis layer; additive, map+cluster) --- - from ..hypothesis import write_module_hypotheses_perm - - if self._subject_groups: - # cells keyed by (freq_pair, metric); the per-vertex coupling map is - # the unit-of-test, same contract as vertex_connectivity FCD. - maps_by_cell: dict[tuple[str, str], dict] = {} - for key in keys: - metric, freq_pair = key.split("|", 1) - cell = {uid: m[key] for uid, m in self._subject_maps.items() if key in m} - if cell: - maps_by_cell[(freq_pair, metric)] = cell - wanted_hyp = self._selection.get("hypothesis") - write_module_hypotheses_perm( - maps_by_cell, self._subject_groups, coords, self.config, - self.tbl_dir, prefix="vertex_cross_freq", - n_perms=self._n_permutations, threshold=self._cluster_threshold, - distance_mm=self._adjacency_distance, - hypothesis=",".join(sorted(wanted_hyp)) if wanted_hyp else None, - atlas_dir=self._atlas_dir, - ) - - self._save_cluster_state() - - def figures(self) -> None: - if not self._cluster_results: - self._load_cluster_state() - if self._source_coords is None: - return - coords = self._source_coords - n = 0 - for key, info in self._cluster_results.items(): - result = info["result"] - if not _has_significant_cluster(result): - continue - contrast = info.get("contrast", key) - metric, freq_pair = info["metric"], info["freq_pair"] - safe = f"{contrast}_{metric}_{freq_pair}".lower().replace(" ", "_") - plot_band_comparison( - coords=coords, mean_a=info["mean_a"], mean_b=info["mean_b"], - t_map=result.t_map, cluster_labels=result.cluster_labels, - cluster_pvalues=result.cluster_pvalues, - band_name=f"{metric.upper()} — {freq_pair} — {contrast}", - group_labels=info["group_labels"], - output_path=self.fig_dir / f"cfc_{safe}.png", - ) - n += 1 - logger.info("vertex_cross_freq: %d significant-cluster figures", n) - - def summary(self) -> None: - data_dir = self.output_dir / "data" - cfg = self._r_config_data() - if self._sfreq is not None: - cfg["sfreq"] = self._sfreq - with open(data_dir / "study_config.yaml", "w") as f: - yaml.dump(cfg, f, default_flow_style=False) - - lines = [ - "# Vertex Cross-Frequency Coupling Summary", "", - f"**Study**: {self.config.name}", - f"**Metrics**: {', '.join(self._metrics)}", - f"**Resolution**: full vertex (no parcellation)", "", - "## Methods", "", - "Cross-frequency coupling at full vertex resolution. PAC = local " - "within-vertex Modulation Index z-map (Tort 2010); AAC = cross-" - "frequency power-power coupling (Bruns 2000 / Masimore 2004); PPC = " - "n:m phase-phase coupling with surrogate z (Palva 2005). Per-vertex " - "maps tested with cluster-based permutation. See CONNECTIVITY_METHODS.md.", - "", - ] - (self.tbl_dir / "vertex_cross_freq_summary.md").write_text("\n".join(lines)) diff --git a/src/source_analytics/analyses/vertex_directed_analysis.py b/src/source_analytics/analyses/vertex_directed_analysis.py deleted file mode 100644 index 0fb4d63..0000000 --- a/src/source_analytics/analyses/vertex_directed_analysis.py +++ /dev/null @@ -1,344 +0,0 @@ -"""Vertex-level directed connectivity: DTF outflow / inflow / net flow. - -The vertex companion to :class:`ROIDirectedAnalysis`. Fits one ridge-regularized -MVAR per subject over the dorsal source vertices and reads out the directed -transfer function (DTF), then reduces the directed matrix to three per-vertex -summary maps that are tested with the same spatial cluster-permutation machinery -as the other vertex modules: - - - **outflow** — mean DTF *from* a vertex to the rest (causal-source / driver strength) - - **inflow** — mean DTF *into* a vertex from the rest (receiver strength) - - **netflow** — outflow − inflow (positive = net driver, negative = net receiver) - -Ridge regularization is mandatory here: dorsal source vertices are strongly -collinear (mean inter-vertex |corr| ≈ 0.64 from source leakage), which makes the -plain least-squares MVAR explosive — see :mod:`..spectral.directed`. -""" - -from __future__ import annotations - -import logging -import pickle -from pathlib import Path - -import numpy as np -import pandas as pd - -from ..config import StudyConfig -from ..io.discovery import SubjectInfo -from ..io.loader import SubjectLoader -from ..spectral.directed import ( - fit_mvar, - dtf_spectrum, - mvar_spectral_radius, - DEFAULT_ORDER, - DEFAULT_RIDGE, -) -from ..stats.cluster_permutation import ( - cluster_permutation_test, - has_significant_cluster as _has_significant_cluster, - hedges_g, -) -from ..viz.glass_brain import plot_band_comparison -from .base import BaseAnalysis - -logger = logging.getLogger(__name__) - -_MEASURES = ["outflow", "inflow", "netflow"] - - -class VertexDirectedAnalysis(BaseAnalysis): - """All-to-all vertex DTF reduced to per-vertex outflow/inflow/netflow maps.""" - - name = "vertex_directed" - SELECTABLE = {"measure": "directed summary", "band": "frequency band", - "hypothesis": "declared hypothesis"} - - def __init__(self, config: StudyConfig, output_dir: Path): - super().__init__(config, output_dir) - self._sfreq: float | None = None - self._source_coords: np.ndarray | None = None - self._vertex_indices: np.ndarray | None = None - # uid -> band -> measure -> per-vertex array - self._subject_maps: dict[str, dict[str, dict[str, np.ndarray]]] = {} - self._subject_groups: dict[str, str] = {} - # uid -> band -> directed matrix - self._dtf_matrices: dict[str, dict[str, np.ndarray]] = {} - self._rows: list[dict] = [] - self._cluster_results: dict = {} - - cfg = config.raw.get(self.name, {}) - self._mvar_order = int(cfg.get("mvar_order", DEFAULT_ORDER)) - self._mvar_ridge = float(cfg.get("mvar_ridge", DEFAULT_RIDGE)) - self._n_freqs = int(cfg.get("n_freqs", 128)) - self._n_permutations = int(cfg.get("n_permutations", 1000)) - self._measures = list(_MEASURES) - - wb_cfg = config.vertex - self._cluster_threshold = float(wb_cfg.get("cluster_threshold", 2.0)) - self._adjacency_distance = float(wb_cfg.get("adjacency_distance_mm", 5.0)) - - def setup(self) -> None: - self._measures = self._select("measure", _MEASURES) - self._subject_maps.clear() - self._subject_groups.clear() - self._dtf_matrices.clear() - self._rows.clear() - self._cluster_results.clear() - self._source_coords = None - self._vertex_indices = None - - @staticmethod - def _reduce(mat: np.ndarray) -> dict[str, np.ndarray]: - """Per-vertex directed summaries from a DTF matrix (mat[i,j] = j -> i).""" - n = mat.shape[0] - m = mat.copy() - np.fill_diagonal(m, 0.0) - inflow = m.sum(axis=1) / (n - 1) # row i: total inflow to vertex i - outflow = m.sum(axis=0) / (n - 1) # col i: total outflow from vertex i - return {"outflow": outflow, "inflow": inflow, "netflow": outflow - inflow} - - def _compute_subject(self, subject: SubjectInfo): - """Pure per-subject ridge-MVAR DTF compute (parallel-safe).""" - loader = SubjectLoader(subject.data_dir) - uid = f"{subject.group}_{subject.subject_id}" - - stc_data = loader.load_source_timecourses() # signed - sfreq = loader.load_sfreq() - coords = loader.load_source_coords() - - mask = self.config.get_vertex_mask(coords) - vertex_indices = np.where(mask)[0] - source_coords = coords[mask] - stc_data = stc_data[vertex_indices] - - # One ridge-MVAR fit per subject; DTF read out across all bands. - A, _ = fit_mvar(stc_data, order=self._mvar_order, ridge=self._mvar_ridge) - radius = mvar_spectral_radius(A) - if radius >= 1.0: - logger.warning( - "%s: MVAR unstable (spectral radius %.2f) — DTF unreliable; " - "raise vertex_directed.mvar_ridge or lower mvar_order.", uid, radius, - ) - freqs = np.linspace(0, sfreq / 2, self._n_freqs) - dtf = dtf_spectrum(A, sfreq, freqs) # (n, n, n_freqs), [i,j] = j -> i - - subj_maps: dict[str, dict[str, np.ndarray]] = {} - subj_mats: dict[str, np.ndarray] = {} - rows: list[dict] = [] - - for band_name, (fmin, fmax) in self._selected_bands().items(): - idx = (freqs >= fmin) & (freqs <= fmax) - if not idx.any(): - continue - mat = dtf[:, :, idx].mean(axis=2) - subj_mats[band_name] = mat - maps = self._reduce(mat) - subj_maps[band_name] = maps - for vi in range(len(vertex_indices)): - row = { - "subject": uid, "group": subject.group, - "vertex_idx": int(vertex_indices[vi]), "band": band_name, - } - for meas in self._measures: - row[meas] = float(maps[meas][vi]) - rows.append(row) - - return { - "uid": uid, "group": subject.group, "sfreq": float(sfreq), - "vertex_indices": vertex_indices, "source_coords": source_coords, - "subject_maps": subj_maps, "dtf_matrices": subj_mats, "rows": rows, - } - - def _merge_subject(self, payload) -> None: - uid = payload["uid"] - if self._sfreq is None: - self._sfreq = payload["sfreq"] - if self._vertex_indices is None: - self._vertex_indices = payload["vertex_indices"] - self._source_coords = payload["source_coords"] - if self.config.has_vertex_filter: - logger.info("Vertex filter: %d vertices retained", - len(self._vertex_indices)) - self._subject_groups[uid] = payload["group"] - self._subject_maps[uid] = payload["subject_maps"] - self._dtf_matrices[uid] = payload["dtf_matrices"] - self._rows.extend(payload["rows"]) - - def aggregate(self) -> None: - data_dir = self.output_dir / "data" - df = pd.DataFrame(self._rows) - if df.empty: - logger.warning("No vertex directed data collected") - return - df.to_csv(data_dir / "vertex_directed.csv", index=False) - logger.info("Exported vertex_directed.csv (%d rows)", len(df)) - - if self._source_coords is not None: - cdf = pd.DataFrame(self._source_coords, columns=["x", "y", "z"]) - cdf.index.name = "vertex_idx" - cdf.to_csv(data_dir / "source_coords.csv") - - if self._dtf_matrices: - with open(data_dir / "vertex_dtf_matrices.pkl", "wb") as f: - pickle.dump(self._dtf_matrices, f) - logger.info("Saved vertex_dtf_matrices.pkl") - - def _reload_maps_from_disk(self) -> bool: - """Reconstruct per-subject directed maps + coords + groups from the - persisted CSVs, so statistics/figures are regenerable via --steps - without re-running process/aggregate.""" - data_dir = self.output_dir / "data" - maps_csv = data_dir / "vertex_directed.csv" - coords_csv = data_dir / "source_coords.csv" - if not maps_csv.exists(): - logger.warning("No persisted maps at %s; cannot reload", maps_csv) - return False - if coords_csv.exists(): - self._source_coords = pd.read_csv(coords_csv)[["x", "y", "z"]].to_numpy(dtype=float) - df = pd.read_csv(maps_csv) - self._subject_maps = {} - self._subject_groups = {} - for (uid, group), g in df.groupby(["subject", "group"], sort=False): - self._subject_groups[uid] = group - bands: dict = {} - for band, gb in g.groupby("band", sort=False): - gb = gb.sort_values("vertex_idx") - bands[band] = {m: gb[m].to_numpy(dtype=float) for m in self._measures} - self._subject_maps[uid] = bands - logger.info("Reloaded %d subjects' directed maps from %s", - len(self._subject_maps), maps_csv) - return True - - def statistics(self) -> None: - if not self._subject_maps: - self._reload_maps_from_disk() - if self._source_coords is None: - logger.error("No source coordinates — cannot run statistics") - return - coords = self._source_coords - all_stats = [] - - for contrast in self._pairwise_contrasts(): - uids_a = [u for u, g in self._subject_groups.items() if g == contrast.group_a] - uids_b = [u for u, g in self._subject_groups.items() if g == contrast.group_b] - if not uids_a or not uids_b: - continue - label_a = self.config.get_group_label(contrast.group_a) - label_b = self.config.get_group_label(contrast.group_b) - - for band_name in self._selected_bands(): - for meas in self._measures: - data_a = np.array([ - self._subject_maps[u][band_name][meas] for u in uids_a - if band_name in self._subject_maps.get(u, {}) - ]) - data_b = np.array([ - self._subject_maps[u][band_name][meas] for u in uids_b - if band_name in self._subject_maps.get(u, {}) - ]) - if data_a.size == 0 or data_b.size == 0: - continue - - result = cluster_permutation_test( - data_a, data_b, coords, - n_perms=self._n_permutations, - threshold=self._cluster_threshold, - distance_mm=self._adjacency_distance, - seed=42, - ) - g_map = hedges_g(data_a, data_b) - self._cluster_results[f"{contrast.name}_{band_name}_{meas}"] = { - "result": result, - "mean_a": data_a.mean(axis=0), - "mean_b": data_b.mean(axis=0), - "group_labels": (label_a, label_b), - "band": band_name, "measure": meas, - } - for vi in range(len(result.t_map)): - all_stats.append({ - "contrast": contrast.name, "band": band_name, - "measure": meas, "vertex_idx": vi, - "value_a": float(data_a.mean(axis=0)[vi]), - "value_b": float(data_b.mean(axis=0)[vi]), - "t": float(result.t_map[vi]), - "p": float(result.p_map[vi]), - "hedges_g": float(g_map[vi]), - "cluster_id": int(result.cluster_labels[vi]), - }) - - if all_stats: - pd.DataFrame(all_stats).to_csv( - self.tbl_dir / "vertex_directed_stats.csv", index=False, - ) - logger.info("Exported vertex_directed_stats.csv (%d rows)", len(all_stats)) - - # --- Declarative hypotheses (hypothesis layer; additive, map+cluster) --- - from ..hypothesis import write_module_hypotheses_perm - - if self._source_coords is not None and self._subject_groups: - maps_by_cell = { - (band_name, measure): { - uid: self._subject_maps[uid][band_name][measure] - for uid in self._subject_groups - } - for band_name in self._selected_bands() - for measure in self._measures - } - wanted_hyp = self._selection.get("hypothesis") - write_module_hypotheses_perm( - maps_by_cell, self._subject_groups, self._source_coords, self.config, - self.tbl_dir, prefix="vertex_directed", - n_perms=self._n_permutations, threshold=self._cluster_threshold, - distance_mm=self._adjacency_distance, - hypothesis=",".join(sorted(wanted_hyp)) if wanted_hyp else None, - atlas_dir=self._atlas_dir, - ) - - self._save_cluster_state() - - def figures(self) -> None: - if not self._cluster_results: - self._load_cluster_state() - if self._source_coords is None: - return - coords = self._source_coords - n = 0 - for key, info in self._cluster_results.items(): - result = info["result"] - if not _has_significant_cluster(result): - continue - safe = key.lower().replace(" ", "_") - plot_band_comparison( - coords=coords, - mean_a=info["mean_a"], mean_b=info["mean_b"], - t_map=result.t_map, - cluster_labels=result.cluster_labels, - cluster_pvalues=result.cluster_pvalues, - band_name=f"DTF {info['measure']} — {info['band']} — {info.get('contrast', '')}", - group_labels=info["group_labels"], - output_path=self.fig_dir / f"dtf_{safe}.png", - ) - n += 1 - logger.info("vertex_directed: %d significant-cluster figures", n) - - def summary(self) -> None: - lines = [ - "# Vertex Directed (DTF) Analysis Summary", - "", - f"**Study**: {self.config.name}", - "**Analysis**: Vertex DTF — outflow / inflow / netflow", - f"**MVAR**: order {self._mvar_order}, ridge {self._mvar_ridge}", - f"**Permutations**: {self._n_permutations}", - "", - "## Output Files", - "", - "- `data/vertex_directed.csv` — per-subject per-vertex outflow/inflow/netflow", - "- `data/vertex_dtf_matrices.pkl` — full directed matrices", - "- `data/source_coords.csv` — vertex coordinates (mm)", - "- `tables/vertex_directed_stats.csv` — cluster-permutation statistics", - "- `figures/dtf_*.png` — glass-brain directed maps", - "", - ] - (self.output_dir / "ANALYSIS_SUMMARY.md").write_text("\n".join(lines)) - logger.info("Wrote %s", self.output_dir / "ANALYSIS_SUMMARY.md") diff --git a/src/source_analytics/analyses/vertex_evoked_analysis.py b/src/source_analytics/analyses/vertex_evoked_analysis.py deleted file mode 100644 index 24a315d..0000000 --- a/src/source_analytics/analyses/vertex_evoked_analysis.py +++ /dev/null @@ -1,366 +0,0 @@ -"""Vertex Evoked response analysis: per-vertex ITC, ERSP, STP for trial paradigms. - -The vertex-level companion to :class:`ROIEvokedAnalysis`. Computes Morlet -time-frequency measures (inter-trial coherence, event-related spectral -perturbation, single-trial power) for every source vertex of a trial-based -(evoked) paradigm, then tests group differences with the same spatial -cluster-permutation machinery used by the other vertex modules. - -Requires an ``evoked`` section in the study YAML (epoch_samples, sfreq, baseline, -tf_params, measures) — identical schema to ``roi_evoked`` — and a vertex paradigm -whose source reconstructions hold concatenated trial epochs. -""" - -from __future__ import annotations - -import logging -from pathlib import Path - -import numpy as np -import pandas as pd - -from ..config import StudyConfig -from ..io.discovery import SubjectInfo -from ..io.loader import SubjectLoader -from ..spectral.tfr import ( - morlet_tfr_avg_power_itc, - compute_ersp, - debias_itc, - extract_measure_in_band, - extract_measure_in_tiles, - resolve_n_cycles, -) -from ..stats.cluster_permutation import cluster_permutation_test, hedges_g -from ..viz.glass_brain import plot_band_comparison -from .base import BaseAnalysis - -logger = logging.getLogger(__name__) - - -class VertexEvokedAnalysis(BaseAnalysis): - """Per-vertex ITC / ERSP / STP for evoked (trial-based) paradigms. - - Python computes the Morlet TFR and extracts a scalar measure per vertex; - group differences are tested per measure with cluster-based permutation over - the source grid (no R — the maps are vertex-level, like vertex_connectivity). - """ - - name = "vertex_evoked" - SELECTABLE = {"hypothesis": "declared hypothesis"} - - def __init__(self, config: StudyConfig, output_dir: Path): - super().__init__(config, output_dir) - self._measure_rows: list[dict] = [] - self._sfreq: float | None = None - self._source_coords: np.ndarray | None = None - self._vertex_indices: np.ndarray | None = None - # uid -> measure_name -> per-vertex value array - self._subject_measures: dict[str, dict[str, np.ndarray]] = {} - self._subject_groups: dict[str, str] = {} - self._cluster_results: dict = {} - - wb_cfg = config.vertex - self._n_permutations = int(config.raw.get("vertex_evoked", {}).get( - "n_permutations", 1000)) - self._cluster_threshold = float(wb_cfg.get("cluster_threshold", 2.0)) - self._adjacency_distance = float(wb_cfg.get("adjacency_distance_mm", 5.0)) - - def _get_evoked_config(self) -> dict: - """Get and validate the evoked config section (same schema as roi_evoked).""" - evoked = self.config.evoked - if not evoked: - evoked = self.config.raw.get(self.name, {}) - if not evoked: - evoked = self.config.raw.get("evoked", {}) - if not evoked: - raise ValueError( - "No 'evoked' section in study config. Evoked analysis requires " - "epoch_samples, sfreq, baseline, tf_params, and measures." - ) - required = ["epoch_samples", "sfreq", "baseline", "tf_params", "measures"] - missing = [k for k in required if k not in evoked] - if missing: - raise ValueError(f"Evoked config missing required keys: {missing}") - return evoked - - def setup(self) -> None: - self._measure_rows.clear() - self._subject_measures.clear() - self._subject_groups.clear() - self._cluster_results.clear() - self._source_coords = None - self._vertex_indices = None - self._get_evoked_config() - - def process_subject(self, subject: SubjectInfo) -> None: - evoked_cfg = self._get_evoked_config() - epoch_samples = int(evoked_cfg["epoch_samples"]) - sfreq = float(evoked_cfg["sfreq"]) - baseline = tuple(evoked_cfg["baseline"]) - tf_params = evoked_cfg["tf_params"] - measures = evoked_cfg["measures"] - self._sfreq = sfreq - - fmin, fmax = tf_params["freq_range"] - freqs = np.arange(fmin, fmax + 1, 1.0) - # n_cycles: scalar, "adaptive", or [lo, hi] as a linear ramp (see - # spectral.tfr.resolve_n_cycles). Shared with roi_evoked so the three - # modules cannot drift apart on the config contract. - n_cycles = resolve_n_cycles(freqs, tf_params.get("n_cycles", 7)) - xmin = baseline[0] - - loader = SubjectLoader(subject.data_dir) - # (n_vertices, n_epochs, epoch_samples) — signed for phase-based ITC - epochs = loader.load_source_epochs(epoch_samples, magnitude=False) - coords = loader.load_source_coords() - - # Vertex filter (compute mask once from the first subject) - if self._vertex_indices is None: - mask = self.config.get_vertex_mask(coords) - self._vertex_indices = np.where(mask)[0] - self._source_coords = coords[mask] - if self.config.has_vertex_filter: - logger.info( - "Vertex filter: %d/%d vertices retained", - len(self._vertex_indices), len(coords), - ) - epochs = epochs[self._vertex_indices] - - uid = f"{subject.group}_{subject.subject_id}" - self._subject_groups[uid] = subject.group - n_vertices = epochs.shape[0] - n_epochs = epochs.shape[1] - - # Per-vertex value array for each configured measure - subj = {m["name"]: np.full(n_vertices, np.nan) for m in measures} - - for vi in range(n_vertices): - avg_power, itc_map = morlet_tfr_avg_power_itc( - epochs[vi], sfreq, freqs, n_cycles, - ) - stp_map = avg_power - ersp_map = compute_ersp(avg_power, sfreq, baseline, xmin=xmin) - measure_maps = {"itc": itc_map, "ersp": ersp_map, "stp": stp_map} - - # ITC is biased upward at low trial counts — pure noise gives about - # 1/sqrt(n) — so subjects with different trial counts are not on the - # same scale. Offered alongside raw ITC rather than replacing it. - if n_epochs >= 2: - measure_maps["itc_debiased"] = debias_itc(itc_map, n_epochs) - - for mdef in measures: - mtype, mname = mdef["type"], mdef["name"] - if mtype not in measure_maps: - logger.warning("Unknown measure type '%s' — skipping", mtype) - continue - # One rectangle, or a union of tiles for a response that - # sweeps in frequency over time (the chirp). - tiles = mdef.get("tiles") - if tiles: - value = extract_measure_in_tiles( - measure_maps[mtype], freqs, sfreq, tiles, xmin=xmin - ) - band = (min(t["band"][0] for t in tiles), - max(t["band"][1] for t in tiles)) - time_window = (min(t["time_window"][0] for t in tiles), - max(t["time_window"][1] for t in tiles)) - else: - band = tuple(mdef["band"]) - time_window = tuple(mdef["time_window"]) - value = extract_measure_in_band( - measure_maps[mtype], freqs, sfreq, band, time_window, - xmin=xmin, - ) - subj[mname][vi] = value - self._measure_rows.append({ - "subject": uid, - "group": subject.group, - "vertex_idx": int(self._vertex_indices[vi]), - "measure_name": mname, - "measure_type": mtype, - "band_lo": band[0], "band_hi": band[1], - "time_lo": time_window[0], "time_hi": time_window[1], - "value": float(value), - "n_epochs": n_epochs, - }) - del avg_power, itc_map, ersp_map, stp_map - - self._subject_measures[uid] = subj - - def aggregate(self) -> None: - data_dir = self.output_dir / "data" - df = pd.DataFrame(self._measure_rows) - if df.empty: - logger.warning("No vertex evoked measure data collected") - return - df.to_csv(data_dir / "vertex_evoked_measures.csv", index=False) - logger.info("Exported vertex_evoked_measures.csv (%d rows)", len(df)) - - if self._source_coords is not None: - coords_df = pd.DataFrame(self._source_coords, columns=["x", "y", "z"]) - coords_df.index.name = "vertex_idx" - coords_df.to_csv(data_dir / "source_coords.csv") - - def _reload_maps_from_disk(self) -> bool: - """Reconstruct per-subject evoked measures + coords + groups from the - persisted CSVs, so statistics/figures are regenerable via --steps.""" - data_dir = self.output_dir / "data" - csv = data_dir / "vertex_evoked_measures.csv" - coords_csv = data_dir / "source_coords.csv" - if not csv.exists(): - logger.warning("No persisted measures at %s; cannot reload", csv) - return False - if coords_csv.exists(): - self._source_coords = pd.read_csv(coords_csv)[["x", "y", "z"]].to_numpy(dtype=float) - df = pd.read_csv(csv) - self._measure_rows = df.to_dict("records") - self._subject_measures = {} - self._subject_groups = {} - for (uid, group), g in df.groupby(["subject", "group"], sort=False): - self._subject_groups[uid] = group - m: dict = {} - for mname, gg in g.groupby("measure_name", sort=False): - m[mname] = gg.sort_values("vertex_idx")["value"].to_numpy(dtype=float) - self._subject_measures[uid] = m - logger.info("Reloaded %d subjects' evoked measures from %s", - len(self._subject_measures), csv) - return True - - def statistics(self) -> None: - if not self._subject_measures: - self._reload_maps_from_disk() - if self._source_coords is None: - logger.error("No source coordinates — cannot run statistics") - return - coords = self._source_coords - measure_names = sorted({m["measure_name"] for m in self._measure_rows}) - all_stats = [] - - for contrast in self._pairwise_contrasts(): - uids_a = [u for u, g in self._subject_groups.items() if g == contrast.group_a] - uids_b = [u for u, g in self._subject_groups.items() if g == contrast.group_b] - if not uids_a or not uids_b: - continue - label_a = self.config.get_group_label(contrast.group_a) - label_b = self.config.get_group_label(contrast.group_b) - - for mname in measure_names: - data_a = np.array([ - self._subject_measures[u][mname] for u in uids_a - if mname in self._subject_measures.get(u, {}) - ]) - data_b = np.array([ - self._subject_measures[u][mname] for u in uids_b - if mname in self._subject_measures.get(u, {}) - ]) - if data_a.size == 0 or data_b.size == 0: - continue - - result = cluster_permutation_test( - data_a, data_b, coords, - n_perms=self._n_permutations, - threshold=self._cluster_threshold, - distance_mm=self._adjacency_distance, - seed=42, - ) - g_map = hedges_g(data_a, data_b) - self._cluster_results[f"{contrast.name}_{mname}"] = { - "result": result, - "mean_a": data_a.mean(axis=0), - "mean_b": data_b.mean(axis=0), - "group_labels": (label_a, label_b), - "measure": mname, - } - for vi in range(len(result.t_map)): - all_stats.append({ - "contrast": contrast.name, - "measure": mname, - "vertex_idx": vi, - "value_a": float(data_a.mean(axis=0)[vi]), - "value_b": float(data_b.mean(axis=0)[vi]), - "t": float(result.t_map[vi]), - "p": float(result.p_map[vi]), - "hedges_g": float(g_map[vi]), - "cluster_id": int(result.cluster_labels[vi]), - }) - - if all_stats: - pd.DataFrame(all_stats).to_csv( - self.tbl_dir / "vertex_evoked_stats.csv", index=False, - ) - logger.info("Exported vertex_evoked_stats.csv (%d rows)", len(all_stats)) - - # --- Declared hypotheses (hypothesis layer; additive, map+cluster contract) --- - # Same wiring as vertex_cluster: every declared hypothesis runs over the - # per-subject vertex maps through the permutation adapter and lands in - # tables/vertex_evoked_hypotheses.csv. The measure is the cell's "band" - # coordinate (each measure fixes its own band + time window, so there is - # no separate band axis — see analyses/_evoked_hypotheses.py) and the - # dv is the measure type (itc / ersp / stp / ...). --hypothesis narrows. - from ..hypothesis import write_module_hypotheses_perm - - measure_types = { - m["measure_name"]: m["measure_type"] for m in self._measure_rows - } - maps_by_cell = { - (mname, measure_types.get(mname, "value")): { - uid: self._subject_measures[uid][mname] - for uid in self._subject_groups - if mname in self._subject_measures.get(uid, {}) - } - for mname in measure_names - } - wanted_hyp = self._selection.get("hypothesis") - write_module_hypotheses_perm( - maps_by_cell, self._subject_groups, coords, self.config, self.tbl_dir, - prefix="vertex_evoked", - n_perms=self._n_permutations, threshold=self._cluster_threshold, - distance_mm=self._adjacency_distance, - hypothesis=",".join(sorted(wanted_hyp)) if wanted_hyp else None, - atlas_dir=self._atlas_dir, - ) - - self._save_cluster_state() - - def figures(self) -> None: - if not self._cluster_results: - self._load_cluster_state() - if self._source_coords is None: - return - coords = self._source_coords - for key, info in self._cluster_results.items(): - result = info["result"] - measure = info["measure"] - safe = key.lower().replace(" ", "_") - plot_band_comparison( - coords=coords, - mean_a=info["mean_a"], - mean_b=info["mean_b"], - t_map=result.t_map, - cluster_labels=result.cluster_labels, - cluster_pvalues=result.cluster_pvalues, - band_name=f"{measure.upper()}", - group_labels=info["group_labels"], - output_path=self.fig_dir / f"evoked_{safe}.png", - ) - - def summary(self) -> None: - lines = [ - "# Vertex Evoked Analysis Summary", - "", - f"**Study**: {self.config.name}", - "**Analysis**: Per-vertex ITC / ERSP / STP (trial paradigm)", - f"**Permutations**: {self._n_permutations}", - "", - "## Output Files", - "", - "- `data/vertex_evoked_measures.csv` — per-subject per-vertex measures", - "- `data/source_coords.csv` — vertex coordinates (mm)", - "- `tables/vertex_evoked_stats.csv` — cluster-permutation statistics (pairwise contrasts)", - "- `tables/vertex_evoked_hypotheses.csv` — declared hypotheses (permutation adapter; " - "band = measure name, dv = measure type)", - "- `figures/evoked_*.png` — glass-brain measure maps", - "", - ] - (self.output_dir / "ANALYSIS_SUMMARY.md").write_text("\n".join(lines)) - logger.info("Wrote %s", self.output_dir / "ANALYSIS_SUMMARY.md") diff --git a/src/source_analytics/analyses/vertex_network_analysis.py b/src/source_analytics/analyses/vertex_network_analysis.py deleted file mode 100644 index c4977d4..0000000 --- a/src/source_analytics/analyses/vertex_network_analysis.py +++ /dev/null @@ -1,589 +0,0 @@ -"""Vertex network layer: multi-density AUC graph metrics + NBS. - -Loads shell-based source timecourses (filtered to dorsal vertices), builds -all-to-all vertex connectivity per band × connectivity metric (loading the -pre-computed matrices from vertex_connectivity when available), and runs two -independent analyses that share those matrices: - -* :class:`VertexGraphAnalysis` (``vertex_graph``) — multi-density AUC graph - metrics with group permutation tests. -* :class:`VertexNBSAnalysis` (``vertex_nbs``) — the Network-Based Statistic. - -:class:`VertexNetworkAnalysis` (``vertex_network``) is the back-compat combined -alias that runs both, with the original output filenames. -""" - -from __future__ import annotations - -import logging -import pickle -from functools import lru_cache -from pathlib import Path - -import numpy as np -import pandas as pd -import yaml - -from ..config import StudyConfig -from ..io.discovery import SubjectInfo -from ..io.loader import SubjectLoader -from ..spectral.vertex_connectivity import ( - compute_vertex_connectivity_matrix, - compute_vertex_connectivity_matrix_epochs, -) -from ..spectral.epoch_sampler import sample_epochs -from ..stats.graph_metrics import ( - GLOBAL_METRIC_NAMES, - compute_auc, - auc_permutation_test, -) -from ..viz.glass_brain import ( - plot_glass_brain_edges, - plot_nbs_roi_bar, - plot_nbs_roi_circos, -) -from ._network_base import NetworkAnalysisBase - -logger = logging.getLogger(__name__) - - -@lru_cache(maxsize=2) -def _load_conn_pkl(path: str, _mtime_ns: int): - """Unpickle a precomputed-connectivity dict ONCE per process (cached). - - ``_mtime_ns`` is part of the cache key so a rebuilt pickle invalidates the - cache. Returns the dict, or None if it can't be read. - """ - try: - with open(path, "rb") as f: - return pickle.load(f) - except Exception: # noqa: BLE001 - return None - - -def _generate_vertex_labels( - coords: np.ndarray, - atlas_labels: list[str] | None = None, -) -> list[str]: - """Descriptive vertex labels: atlas ROI where available, else spatial.""" - n = len(coords) - labels = [] - centroid = coords.mean(axis=0) - for i in range(n): - if atlas_labels and atlas_labels[i] not in ("Exterior", "Unknown_0"): - labels.append(atlas_labels[i]) - continue - x, y, z = coords[i] - parts = [] - parts.append("Anterior" if y > centroid[1] + 1.0 else - "Posterior" if y < centroid[1] - 1.0 else "Central_AP") - parts.append("Left" if x < centroid[0] - 0.5 else - "Right" if x > centroid[0] + 0.5 else "Midline") - parts.append("Dorsal" if z > centroid[2] + 0.5 else - "Ventral" if z < centroid[2] - 0.5 else "Central_DV") - labels.append("_".join(parts)) - return labels - - -class _VertexNetworkBase(NetworkAnalysisBase): - """Shared vertex machinery: matrix build/load, AUC graph metrics, NBS.""" - - _default_nbs_threshold = 3.0 - _nbs_results_filename = "vertex_nbs_results.csv" - _fallback_config_key = "vertex_network" - - def __init__(self, config: StudyConfig, output_dir: Path): - super().__init__(config, output_dir) - self._auc_rows: list[dict] = [] - self._density_rows: list[dict] = [] - self._source_coords: np.ndarray | None = None - self._vertex_indices: np.ndarray | None = None - self._vertex_labels: list[str] = [] - self._sfreq: float | None = None - # uid -> band -> conn_metric -> {graph_metric: auc} - self._subject_aucs: dict[str, dict[str, dict[str, dict[str, float]]]] = {} - - self._init_network_config() - cfg = self._net_cfg - self._auc_permutations = int(cfg.get("auc_permutations", 5000)) - self._density_min = float(cfg.get("density_min", 0.05)) - self._density_max = float(cfg.get("density_max", 0.40)) - self._density_step = float(cfg.get("density_step", 0.01)) - # Global epoch_sampling → vertex: block → per-analysis block (see base). - self._epoch_config = self._vertex_epoch_config() - - def setup(self) -> None: - self._auc_rows.clear() - self._density_rows.clear() - self._subject_aucs.clear() - self._subject_groups.clear() - self._conn_matrices.clear() - self._source_coords = None - self._vertex_indices = None - self._vertex_labels.clear() - self._nbs_results.clear() - - # ----------------------------------------------------- process steps --- # - def _compute_matrices(self, subject: SubjectInfo) -> dict: - """Pure: build/load per band × metric vertex connectivity. No self mutation.""" - loader = SubjectLoader(subject.data_dir) - uid = f"{subject.group}_{subject.subject_id}" - stc_data = loader.load_source_timecourses() - sfreq = loader.load_sfreq() - coords = loader.load_source_coords() - - mask = self.config.get_vertex_mask(coords) - vertex_indices = np.where(mask)[0] - source_coords = coords[mask] - stc_data = stc_data[vertex_indices] - - subject_conn: dict[str, dict[str, np.ndarray]] = {} - for band_name, (fmin, fmax) in self._selected_bands().items(): - band_conn: dict[str, np.ndarray] = {} - for metric in self._connectivity_metrics: - conn_mat = self._load_precomputed_conn(uid, band_name, metric) - if conn_mat is None: - logger.info(" %s / %s: computing connectivity...", band_name, metric) - if self._epoch_config is not None: - epochs = sample_epochs( - stc_data, sfreq, - epoch_duration_sec=self._epoch_config.get("epoch_duration_sec", 2.0), - n_epochs=self._epoch_config.get("n_epochs", 80), - seed=self._epoch_config.get("seed", 42), - n_bootstrap=self._epoch_config.get("n_bootstrap", 1), - ) - conn_mat = compute_vertex_connectivity_matrix_epochs( - epochs, sfreq, (fmin, fmax), metric=metric) - else: - conn_mat = compute_vertex_connectivity_matrix( - stc_data, sfreq, (fmin, fmax), metric=metric) - band_conn[metric] = conn_mat - subject_conn[band_name] = band_conn - - return { - "uid": uid, "group": subject.group, "sfreq": float(sfreq), - "vertex_indices": vertex_indices, "source_coords": source_coords, - "conn": subject_conn, - } - - def _merge_matrices(self, payload: dict) -> None: - """Store a matrices payload; set lazy vertex/atlas state on first subject.""" - uid = payload["uid"] - if self._sfreq is None: - self._sfreq = payload["sfreq"] - if self._vertex_indices is None: - self._vertex_indices = payload["vertex_indices"] - self._source_coords = payload["source_coords"] - if self.config.has_vertex_filter: - logger.info("Vertex filter: %d vertices retained", - len(self._vertex_indices)) - try: - from ..atlas import find_atlas_dir, load_vertex_roi_labels - atlas_labels = load_vertex_roi_labels( - self._source_coords, - self._atlas_dir if self._atlas_dir is not None else find_atlas_dir()) - except Exception: - atlas_labels = None - self._vertex_labels = _generate_vertex_labels(self._source_coords, atlas_labels) - self._subject_groups[uid] = payload["group"] - self._conn_matrices[uid] = payload["conn"] - - def _process_matrices(self, subject: SubjectInfo) -> None: - """Serial build/load per band × metric vertex connectivity (graph + NBS).""" - self._merge_matrices(self._compute_matrices(subject)) - - def _compute_graph_metrics(self, uid: str, group: str, - conn: dict[str, dict[str, np.ndarray]]) -> dict: - """Pure: multi-density AUC graph metrics from a subject's matrices.""" - subject_auc_by_band: dict[str, dict[str, dict[str, float]]] = {} - auc_rows: list[dict] = [] - density_rows: list[dict] = [] - for band_name, band_conn in conn.items(): - band_auc: dict[str, dict[str, float]] = {} - for metric, conn_mat in band_conn.items(): - auc_result = compute_auc( - conn_mat, density_min=self._density_min, - density_max=self._density_max, density_step=self._density_step) - band_auc[metric] = auc_result.auc - row = {"subject": uid, "group": group, - "band": band_name, "conn_metric": metric} - row.update(auc_result.auc) - auc_rows.append(row) - for gm in auc_result.metrics_by_density: - drow = {"subject": uid, "group": group, "band": band_name, - "conn_metric": metric, "density": gm.density} - for mn in GLOBAL_METRIC_NAMES: - drow[mn] = getattr(gm, mn) - density_rows.append(drow) - subject_auc_by_band[band_name] = band_auc - return {"subject_aucs": subject_auc_by_band, - "auc_rows": auc_rows, "density_rows": density_rows} - - def _process_graph(self, subject: SubjectInfo) -> None: - """Serial AUC graph metrics from the already-built matrices.""" - uid = f"{subject.group}_{subject.subject_id}" - gp = self._compute_graph_metrics(uid, subject.group, self._conn_matrices.get(uid, {})) - self._subject_aucs[uid] = gp["subject_aucs"] - self._auc_rows.extend(gp["auc_rows"]) - self._density_rows.extend(gp["density_rows"]) - - def _load_precomputed_conn(self, uid: str, band_name: str, metric: str) -> np.ndarray | None: - vc_pkl = (self.config.output_dir / "vertex_connectivity" / "data" - / "vertex_connectivity_matrices.pkl") - if not vc_pkl.exists(): - return None - try: - # Cached per process: the whole dict was previously unpickled on every - # (uid, band, metric) call — ~2100 full reloads for one graph run. - all_conn = _load_conn_pkl(str(vc_pkl), vc_pkl.stat().st_mtime_ns) - if all_conn is None: - return None - val = all_conn.get(uid, {}).get(band_name) - if val is None: - return None - conn_mat = val.get(metric) if isinstance(val, dict) else ( - val if metric == "imag_coherence" else None) - if conn_mat is not None: - logger.info(" Loaded pre-computed connectivity for %s / %s", band_name, metric) - return conn_mat - except Exception: - return None - - # --------------------------------------------------- graph aggregate --- # - def _graph_aggregate(self) -> None: - data_dir = self.output_dir / "data" - if self._auc_rows: - pd.DataFrame(self._auc_rows).to_csv(data_dir / f"{self.name}_auc.csv", index=False) - if self._density_rows: - pd.DataFrame(self._density_rows).to_csv( - data_dir / f"{self.name}_density_curves.csv", index=False) - self._write_source_coords(data_dir) - - def _write_source_coords(self, data_dir: Path) -> None: - if self._source_coords is None: - return - coords_df = pd.DataFrame(self._source_coords, columns=["x", "y", "z"]) - if self._vertex_labels: - coords_df["label"] = self._vertex_labels - coords_df.index.name = "vertex_idx" - coords_df.to_csv(data_dir / "source_coords.csv") - - # --------------------------------------------------- graph statistics -- # - def _graph_statistics(self) -> None: - all_auc_stats = [] - for contrast in self._pairwise_contrasts(): - group_a = [u for u, g in self._subject_groups.items() if g == contrast.group_a] - group_b = [u for u, g in self._subject_groups.items() if g == contrast.group_b] - if not group_a or not group_b: - continue - label_a = self.config.get_group_label(contrast.group_a) - label_b = self.config.get_group_label(contrast.group_b) - for band_name in self._selected_bands(): - for metric in self._connectivity_metrics: - auc_a = [self._subject_aucs[u][band_name][metric] for u in group_a - if metric in self._subject_aucs.get(u, {}).get(band_name, {})] - auc_b = [self._subject_aucs[u][band_name][metric] for u in group_b - if metric in self._subject_aucs.get(u, {}).get(band_name, {})] - if not auc_a or not auc_b: - continue - perm_results = auc_permutation_test( - auc_a, auc_b, n_permutations=self._auc_permutations, seed=42) - for metric_name, res in perm_results.items(): - all_auc_stats.append({ - "contrast": contrast.name, "group_a": label_a, "group_b": label_b, - "band": band_name, "conn_metric": metric, "metric": metric_name, - "mean_a": res["mean_a"], "mean_b": res["mean_b"], - "observed_diff": res["observed_diff"], "p_value": res["p_value"], - "hedges_g": res["hedges_g"], "significant": res["p_value"] < 0.05, - }) - if all_auc_stats: - pd.DataFrame(all_auc_stats).to_csv( - self.tbl_dir / f"{self.name}_stats.csv", index=False) - logger.info("Exported %s_stats.csv (%d rows)", self.name, len(all_auc_stats)) - - self._write_graph_hypotheses() - - def _write_graph_hypotheses(self) -> None: - """Additive declarative-hypothesis CSV (the scalar tabular contract). - - Between-subjects contrast on the global AUC for every declared hypothesis, - faceted by connectivity × graph metric (no spatial unit), with declarative - FDR. Written alongside the legacy permutation ``_stats.csv``. - """ - spec = self.config.design_spec - if spec is None or not spec.hypotheses: - return - from ..hypothesis import write_module_hypotheses_tabular - - long_rows: list[dict] = [] - for uid, group in self._subject_groups.items(): - for band_name in self._selected_bands(): - for metric in self._connectivity_metrics: - aucs = self._subject_aucs.get(uid, {}).get(band_name, {}).get(metric) - if not aucs: - continue - for graph_metric, auc in aucs.items(): - long_rows.append({ - "subject": uid, spec.factor: group, "band": band_name, - "conn_metric": metric, "graph_metric": graph_metric, - "value": float(auc), - }) - if not long_rows: - return - wanted = self._selection.get("hypothesis") - write_module_hypotheses_tabular( - pd.DataFrame(long_rows), self.config, self.tbl_dir, prefix=self.name, - value_col="value", spatial_col=None, - facet_cols=("conn_metric", "graph_metric"), band_col="band", - hypothesis=",".join(sorted(wanted)) if wanted else None, - ) - - # -------------------------------------------------------- nbs figures -- # - def _save_nbs_state(self) -> None: - """Persist the significant subnetworks' t-matrices so _nbs_figures can be - regenerated with `--steps figures` (source_coords already saved). Only - significant keys are stored (float32), keeping the pickle small.""" - sig = {k: v.t_matrix.astype("float32") - for k, v in self._nbs_results.items() - if getattr(v, "n_significant_components", 0) > 0} - if not sig: - return - path = self.output_dir / "data" / f"{self.name}_nbs_state.pkl" - with open(path, "wb") as f: - pickle.dump(sig, f) - logger.info("Saved NBS figure state (%d significant key(s))", len(sig)) - - def _load_nbs_state(self) -> bool: - """Reload persisted NBS figure state (t-matrices + source coords).""" - from types import SimpleNamespace - path = self.output_dir / "data" / f"{self.name}_nbs_state.pkl" - if not path.exists(): - return False - try: - with open(path, "rb") as f: - sig = pickle.load(f) - except Exception: # noqa: BLE001 - return False - # every stored key is significant by construction - self._nbs_results = { - k: SimpleNamespace(t_matrix=t, n_significant_components=1) - for k, t in sig.items() - } - if self._source_coords is None: - cc = self.output_dir / "data" / "source_coords.csv" - if cc.exists(): - df = pd.read_csv(cc) - self._source_coords = df[["x", "y", "z"]].to_numpy() - if "label" in df.columns and not self._vertex_labels: - self._vertex_labels = df["label"].astype(str).tolist() - return True - - def _nbs_figures(self) -> None: - if not self._nbs_results: - self._load_nbs_state() - if self._source_coords is None: - return - n_nodes = next(iter(self._nbs_results.values())).t_matrix.shape[0] - # Guard: the coords MUST align with the connectivity matrix (same vertex - # set/order). If a filter mismatch leaves them different lengths, node - # indices would map to the wrong coordinates — skip coord-based views - # rather than mislabel them. - if len(self._source_coords) != n_nodes: - logger.warning( - " NBS coords (%d) != matrix nodes (%d) — skipping glass-brain/ROI " - "figures (check vertex_filter matches vertex_connectivity)", - len(self._source_coords), n_nodes) - return - # Allen ROI per vertex via the module's OWN atlas (self._atlas_dir, the - # configured allen32) — the same labeler + granularity the NBS table uses, - # so the figure legend matches the table/digest region coverage exactly. - roi_labels = self._label_vertex_regions(self._source_coords) - if roi_labels and not any(r for r in roi_labels): - roi_labels = self._vertex_labels or None # atlas unavailable → fallback - for key, nbs in self._nbs_results.items(): - if nbs.n_significant_components == 0: - continue - sig_mask = np.abs(nbs.t_matrix) > self._nbs_threshold - edge_pairs = np.argwhere(np.triu(sig_mask, k=1)) - if len(edge_pairs) == 0: - continue - # Signed t so direction (group_a>group_b vs <) is visible, not |t|. - edge_t = np.array([nbs.t_matrix[i, j] for i, j in edge_pairs]) - plot_glass_brain_edges( - coords=self._source_coords, edges=edge_pairs, - output_path=self.fig_dir / f"vertex_nbs_edges_{key}.png", - edge_values=edge_t, signed=True, roi_labels=roi_labels, - title=f"NBS Edges — {key}") - # ROI-grouped views of the same subnetwork (topology + involvement). - if roi_labels is not None: - plot_nbs_roi_circos( - edge_pairs, roi_labels, - self.fig_dir / f"vertex_nbs_roicircos_{key}.png", - edge_values=edge_t, signed=True, title=f"NBS ROI — {key}") - plot_nbs_roi_bar( - edge_pairs, roi_labels, - self.fig_dir / f"vertex_nbs_roibar_{key}.png", - edge_values=edge_t, title=f"NBS ROI involvement — {key}") - - # --------------------------------------------------------- summaries --- # - def _write_summary(self, graph: bool = True, nbs: bool = True) -> None: - lines = [f"# {self.name} Analysis Summary", "", - f"**Study**: {self.config.name}", - f"**Connectivity metrics**: {', '.join(self._connectivity_metrics)}"] - if graph: - lines += [f"**AUC permutations**: {self._auc_permutations}", - f"**Density range**: {self._density_min:.0%}–{self._density_max:.0%}"] - auc_csv = self.tbl_dir / f"{self.name}_stats.csv" - if auc_csv.exists(): - df = pd.read_csv(auc_csv) - sig = df[df["significant"] == True] - lines += ["", "## AUC Permutation Results", "", - (f"**{len(sig)} significant AUC differences** (p<0.05)" - if len(sig) else "No significant AUC differences at p < 0.05.")] - if nbs: - lines += ["", f"**NBS threshold**: t = {self._nbs_threshold}", - f"**NBS permutations**: {self._nbs_permutations}"] - nbs_csv = self.tbl_dir / self._nbs_results_filename - if nbs_csv.exists(): - ndf = pd.read_csv(nbs_csv) - sig = ndf[ndf["p_corrected"] < 0.05] - lines += ["", "## NBS Results", ""] - if len(sig): - lines.append(f"**{len(sig)} significant subnetworks** (p<0.05)") - for _, row in sig.iterrows(): - lines.append(f"- {row['key']}: {row['n_edges']} edges, p={row['p_corrected']:.4f}") - else: - lines.append("No significant NBS subnetworks at p < 0.05.") - lines.append("") - (self.output_dir / "ANALYSIS_SUMMARY.md").write_text("\n".join(lines)) - - -class VertexGraphAnalysis(_VertexNetworkBase): - """Vertex multi-density AUC graph metrics, no NBS.""" - - name = "vertex_graph" - - def _compute_subject(self, subject: SubjectInfo) -> dict: - """Pure per-subject compute (matrices + AUC graph metrics) — parallel-safe.""" - mp = self._compute_matrices(subject) - gp = self._compute_graph_metrics(mp["uid"], mp["group"], mp["conn"]) - return {**mp, **gp} - - def _merge_subject(self, payload: dict) -> None: - self._merge_matrices(payload) - uid = payload["uid"] - self._subject_aucs[uid] = payload["subject_aucs"] - self._auc_rows.extend(payload["auc_rows"]) - self._density_rows.extend(payload["density_rows"]) - - def aggregate(self) -> None: - self._graph_aggregate() - - def statistics(self) -> None: - self._graph_statistics() - - def figures(self) -> None: - """AUC group-effect heatmaps, regenerated from the persisted stats table - (band x connectivity-metric per graph metric, per contrast).""" - import matplotlib - matplotlib.use("Agg") - import matplotlib.pyplot as plt - - stats_csv = self.tbl_dir / "vertex_graph_stats.csv" - if not stats_csv.exists(): - logger.warning("No vertex_graph_stats.csv — skipping figures") - return - df = pd.read_csv(stats_csv) - if df.empty: - return - bands = list(dict.fromkeys(df["band"])) - conns = sorted(df["conn_metric"].unique()) - for (contrast, gmetric), sub in df.groupby(["contrast", "metric"], sort=False): - mat = np.full((len(conns), len(bands)), np.nan) - sig = np.zeros_like(mat, dtype=bool) - for _, r in sub.iterrows(): - if r["band"] not in bands: - continue - i, j = conns.index(r["conn_metric"]), bands.index(r["band"]) - mat[i, j] = r["hedges_g"] - sig[i, j] = bool(r["significant"]) - vmax = float(np.nanmax(np.abs(mat))) if np.isfinite(mat).any() else 1.0 - vmax = vmax or 1.0 - fig, ax = plt.subplots(figsize=(1.1 * len(bands) + 2, 0.5 * len(conns) + 2)) - im = ax.imshow(mat, cmap="RdBu_r", vmin=-vmax, vmax=vmax, aspect="auto") - ax.set_xticks(range(len(bands))) - ax.set_xticklabels(bands, rotation=45, ha="right", fontsize=8) - ax.set_yticks(range(len(conns))) - ax.set_yticklabels(conns, fontsize=8) - for i in range(len(conns)): - for j in range(len(bands)): - if not np.isnan(mat[i, j]): - ax.text(j, i, f"{mat[i, j]:.2f}" + ("*" if sig[i, j] else ""), - ha="center", va="center", fontsize=7) - ax.set_title(f"{gmetric} AUC effect (Hedges g) — {contrast}", fontsize=10) - fig.colorbar(im, ax=ax, label="Hedges g") - fig.tight_layout() - fname = f"graph_auc_{gmetric}_{contrast}".lower().replace(" ", "_") + ".png" - fig.savefig(self.fig_dir / fname, dpi=200) - plt.close(fig) - - def summary(self) -> None: - self._write_summary(graph=True, nbs=False) - - -class VertexNBSAnalysis(_VertexNetworkBase): - """Vertex Network-Based Statistic, no graph metrics.""" - - name = "vertex_nbs" - - def process_subject(self, subject: SubjectInfo) -> None: - self._process_matrices(subject) - - def aggregate(self) -> None: - self._write_source_coords(self.output_dir / "data") - - def statistics(self) -> None: - self._run_nbs() - self._run_nbs_hypotheses() - self._save_nbs_state() - - def figures(self) -> None: - self._nbs_figures() - - def summary(self) -> None: - self._write_summary(graph=False, nbs=True) - - -class VertexNetworkAnalysis(_VertexNetworkBase): - """Combined alias: AUC graph metrics + NBS (back-compat filenames).""" - - name = "vertex_network" - _fallback_config_key = None - - def process_subject(self, subject: SubjectInfo) -> None: - self._process_matrices(subject) - self._process_graph(subject) - - def aggregate(self) -> None: - self._graph_aggregate() - - def statistics(self) -> None: - self._graph_statistics() - self._run_nbs() - self._run_nbs_hypotheses() - self._save_nbs_state() - - def figures(self) -> None: - self._nbs_figures() - - def summary(self) -> None: - data_dir = self.output_dir / "data" - config_path = data_dir / "study_config.yaml" - config_data = self._r_config_data() - if self._sfreq is not None: - config_data["sfreq"] = self._sfreq - with open(config_path, "w") as f: - yaml.dump(config_data, f, default_flow_style=False) - # The combined alias has no R report of its own (graph + NBS statistics - # run in Python via the hypothesis layer); write the Python summary. - self._write_summary(graph=True, nbs=True) diff --git a/src/source_analytics/analyses/vertex_signature_analysis.py b/src/source_analytics/analyses/vertex_signature_analysis.py deleted file mode 100644 index 95ea920..0000000 --- a/src/source_analytics/analyses/vertex_signature_analysis.py +++ /dev/null @@ -1,465 +0,0 @@ -"""Vertex neural-signature (whole-brain classification) analysis. - -Classifies groups from whole-brain spatial patterns of band power with LOOCV + -permutation testing, across one or more classifiers (the interpretable linear -trio svm_linear/logistic/lda give feature-importance maps; svm_rbf is -accuracy-only). Provides an omnibus test per band × classifier: can the spatial -pattern of activity distinguish KO from WT? -""" - -from __future__ import annotations - -import logging -import pickle -import subprocess -from pathlib import Path - -import numpy as np -import pandas as pd -import yaml -import matplotlib - -matplotlib.use("Agg") -import matplotlib.pyplot as plt - -from ..config import StudyConfig -from ..io.discovery import SubjectInfo -from ..io.loader import SubjectLoader -from ..spectral.vertex import compute_psd_vertices, extract_band_power_vertices -from ..spectral.epoch_sampler import sample_epochs -from ..stats.signature import ( - SignatureResult, - classifier_label, - normalize_classifier, - run_signature, -) -from ..viz.glass_brain import plot_glass_brain -from .base import BaseAnalysis - -logger = logging.getLogger(__name__) - - -def _find_r_script_dir() -> Path: - pkg_root = Path(__file__).resolve().parent.parent.parent.parent - r_dir = pkg_root / "R" - if r_dir.is_dir(): - return r_dir - for candidate in [Path.cwd() / "R", Path(__file__).parent.parent.parent / "R"]: - if candidate.is_dir(): - return candidate - raise FileNotFoundError("Cannot find R/ scripts directory") - - -class VertexSignatureAnalysis(BaseAnalysis): - """Whole-brain vertex-level neural-signature (classification) analysis.""" - - name = "vertex_signature" - SELECTABLE = {"band": "frequency band"} - - def __init__(self, config: StudyConfig, output_dir: Path): - super().__init__(config, output_dir) - self._feature_rows: list[dict] = [] - self._source_coords: np.ndarray | None = None - self._sfreq: float | None = None - self._subject_data: dict[str, dict] = {} - self._subject_groups: dict[str, str] = {} - self._subject_order: list[str] = [] - - # Config. `classifiers:` (list) drives the multi-model run; `classifier:` - # (scalar) is the single-model fallback. Names are normalised/deduped in - # config order. (vertex_mvpa/mvpa keys kept as back-compat aliases.) - sig_cfg = config.raw.get( - "vertex_signature", - config.raw.get("vertex_mvpa", config.raw.get("mvpa", {}))) - raw_clfs = sig_cfg.get("classifiers") or [sig_cfg.get("classifier", "svm_linear")] - seen: set[str] = set() - self._classifiers: list[str] = [] - for c in raw_clfs: - key = normalize_classifier(c) - if key not in seen: - seen.add(key) - self._classifiers.append(key) - self._cv_method = sig_cfg.get("cv_method", "loocv") - self._n_permutations = int(sig_cfg.get("n_permutations", 1000)) - - wb_cfg = config.vertex - self._noise_exclude = wb_cfg.get("noise_exclude_hz") - if self._noise_exclude is not None: - self._noise_exclude = tuple(self._noise_exclude) - - # Global epoch_sampling → vertex: block → per-analysis block (see base). - self._epoch_config = self._vertex_epoch_config() - self._signature_results: dict[str, object] = {} - - def setup(self) -> None: - self._feature_rows.clear() - self._subject_data.clear() - self._subject_groups.clear() - self._subject_order.clear() - self._source_coords = None - self._signature_results.clear() - - def process_subject(self, subject: SubjectInfo) -> None: - loader = SubjectLoader(subject.data_dir) - uid = f"{subject.group}_{subject.subject_id}" - - stc_data = loader.load_source_timecourses() - sfreq = loader.load_sfreq() - coords = loader.load_source_coords() - - if self._sfreq is None: - self._sfreq = sfreq - if self._source_coords is None: - self._source_coords = coords - - # Compute PSD - fmax = max(hi for _, hi in self.config.bands.values()) + 10 - if self._epoch_config is not None: - epochs = sample_epochs( - stc_data, sfreq, - epoch_duration_sec=self._epoch_config.get("epoch_duration_sec", 2.0), - n_epochs=self._epoch_config.get("n_epochs", 80), - seed=self._epoch_config.get("seed", 42), - n_bootstrap=self._epoch_config.get("n_bootstrap", 1), - ) - all_psd = [] - for ep in epochs: - f, p = compute_psd_vertices(ep, sfreq, fmax=fmax) - all_psd.append(p) - freqs = f - psd = np.mean(all_psd, axis=0) - else: - freqs, psd = compute_psd_vertices(stc_data, sfreq, fmax=fmax) - - band_power = extract_band_power_vertices( - freqs, psd, self._selected_bands(), noise_exclude=self._noise_exclude, - ) - - self._subject_groups[uid] = subject.group - self._subject_order.append(uid) - self._subject_data[uid] = {"band_power": band_power} - - n_vertices = stc_data.shape[0] - for band_name, bp in band_power.items(): - for vi in range(n_vertices): - self._feature_rows.append({ - "subject": uid, - "group": subject.group, - "vertex_idx": vi, - "band": band_name, - "relative": float(bp["relative"][vi]), - }) - - def aggregate(self) -> None: - data_dir = self.output_dir / "data" - - feat_df = pd.DataFrame(self._feature_rows) - if feat_df.empty: - logger.warning("No signature feature data collected") - return - feat_df.to_csv(data_dir / "vertex_signature_features.csv", index=False) - logger.info("Exported vertex_signature_features.csv (%d rows)", len(feat_df)) - - if self._source_coords is not None: - coords_df = pd.DataFrame(self._source_coords, columns=["x", "y", "z"]) - coords_df.index.name = "vertex_idx" - coords_df.to_csv(data_dir / "source_coords.csv") - - def statistics(self) -> None: - if not self._subject_data: - logger.error("No subject data for signature analysis") - return - - tbl_dir = self.tbl_dir - all_results = [] - - for contrast in self._pairwise_contrasts(): - group_a_uids = [ - uid for uid in self._subject_order - if self._subject_groups[uid] == contrast.group_a - ] - group_b_uids = [ - uid for uid in self._subject_order - if self._subject_groups[uid] == contrast.group_b - ] - - if not group_a_uids or not group_b_uids: - continue - - ordered_uids = group_a_uids + group_b_uids - labels = np.array( - [0] * len(group_a_uids) + [1] * len(group_b_uids) - ) - - for band_name in self._selected_bands(): - # Build feature matrix: (n_subjects, n_vertices) - features = np.array([ - self._subject_data[uid]["band_power"][band_name]["relative"] - for uid in ordered_uids - ]) - - for clf in self._classifiers: - result = run_signature( - features, labels, - classifier=clf, - cv_method=self._cv_method, - n_permutations=self._n_permutations, - seed=42, - ) - - key = f"{contrast.name}_{band_name}_{clf}" - self._signature_results[key] = result - - all_results.append({ - "contrast": contrast.name, - "band": band_name, - "classifier": clf, - "model": classifier_label(clf), - "accuracy": result.accuracy, - "p_value": result.p_value, - "balanced_accuracy": result.balanced_accuracy, - "balanced_p_value": result.balanced_p_value, - "sensitivity": result.sensitivity, - "specificity": result.specificity, - "auc": result.auc, - "ci_lower": result.accuracy_ci[0], - "ci_upper": result.accuracy_ci[1], - "balanced_ci_lower": result.balanced_accuracy_ci[0], - "balanced_ci_upper": result.balanced_accuracy_ci[1], - "n_permutations": result.n_permutations, - }) - - if all_results: - results_df = pd.DataFrame(all_results) - results_df.to_csv(tbl_dir / "vertex_signature_results.csv", index=False) - logger.info("Exported vertex_signature_results.csv") - - # Save full results for --steps figures support - if self._signature_results: - data_dir = self.output_dir / "data" - pkl_data = {} - for key, result in self._signature_results.items(): - pkl_data[key] = { - "feature_weights": result.feature_weights, - "null_distribution": result.null_distribution, - "predictions": result.predictions, - "true_labels": result.true_labels, - "accuracy": result.accuracy, - "p_value": result.p_value, - "sensitivity": result.sensitivity, - "specificity": result.specificity, - "auc": result.auc, - "accuracy_ci": result.accuracy_ci, - "balanced_accuracy": result.balanced_accuracy, - "balanced_p_value": result.balanced_p_value, - "balanced_accuracy_ci": result.balanced_accuracy_ci, - "n_permutations": result.n_permutations, - "classifier": result.classifier, - "has_weights": result.has_weights, - } - with open(data_dir / "vertex_signature_results.pkl", "wb") as f: - pickle.dump(pkl_data, f) - logger.info("Saved vertex_signature_results.pkl") - - def _load_state_from_disk(self) -> bool: - """Load saved signature state from pickle for --steps figures support.""" - data_dir = self.output_dir / "data" - pkl_path = data_dir / "vertex_signature_results.pkl" - if not pkl_path.exists(): - logger.warning("No saved signature state at %s; skipping figures", pkl_path) - return False - - with open(pkl_path, "rb") as f: - saved = pickle.load(f) - - for key, d in saved.items(): - self._signature_results[key] = SignatureResult( - accuracy=d["accuracy"], - p_value=d["p_value"], - sensitivity=d["sensitivity"], - specificity=d["specificity"], - auc=d["auc"], - accuracy_ci=tuple(d["accuracy_ci"]), - feature_weights=d["feature_weights"], - null_distribution=d["null_distribution"], - predictions=d["predictions"], - true_labels=d["true_labels"], - n_permutations=d["n_permutations"], - classifier=d.get("classifier", "svm_linear"), - has_weights=d.get("has_weights", True), - # .get() so a pkl written before the balanced-metric change still loads. - balanced_accuracy=d.get("balanced_accuracy", float("nan")), - balanced_p_value=d.get("balanced_p_value", float("nan")), - balanced_accuracy_ci=tuple( - d.get("balanced_accuracy_ci", (float("nan"), float("nan")))), - ) - - # Load source coords - coords_csv = data_dir / "source_coords.csv" - if coords_csv.exists(): - coords_df = pd.read_csv(coords_csv) - self._source_coords = coords_df[["x", "y", "z"]].values - - logger.info("Loaded signature state from %s", pkl_path) - return True - - def figures(self) -> None: - # Load from disk if in-memory state is missing (--steps support) - if not self._signature_results or self._source_coords is None: - if not self._load_state_from_disk(): - return - - if self._source_coords is None: - return - - coords = self._source_coords - fig_dir = self.fig_dir - - for key, result in self._signature_results.items(): - safe_name = key.lower().replace(" ", "_") - - # Feature importance glass brain — only for linear models that expose - # coef_ (non-linear e.g. svm_rbf has no per-vertex weight map). - if getattr(result, "has_weights", True) and not np.all(np.isnan(result.feature_weights)): - plot_glass_brain( - coords=coords, - values=result.feature_weights, - title=f"Feature Importance — {key}", - output_path=fig_dir / f"vertex_signature_importance_{safe_name}.png", - cmap="YlOrRd", - ) - - # Null distribution histogram - fig, ax = plt.subplots(figsize=(8, 5)) - ax.hist(result.null_distribution, bins=30, color="#3498DB", - alpha=0.7, edgecolor="white", label="Null distribution") - ax.axvline(result.accuracy, color="#E74C3C", linewidth=2, - linestyle="--", label=f"Observed: {result.accuracy:.1%}") - ax.set_xlabel("Accuracy") - ax.set_ylabel("Count") - ax.set_title(f"Signature Permutation Test — {key}") - ax.legend() - fig.tight_layout() - fig.savefig(fig_dir / f"vertex_signature_null_{safe_name}.png", dpi=150) - plt.close(fig) - - # Confusion matrix - fig, ax = plt.subplots(figsize=(5, 4)) - preds = result.predictions - true = result.true_labels - cm = np.array([ - [(true == 0) & (preds == 0), (true == 0) & (preds == 1)], - [(true == 1) & (preds == 0), (true == 1) & (preds == 1)], - ]) - cm_counts = np.array([[s.sum() for s in row] for row in cm]) - ax.imshow(cm_counts, cmap="Blues") - for i in range(2): - for j in range(2): - ax.text(j, i, str(cm_counts[i, j]), - ha="center", va="center", fontsize=16) - ax.set_xticks([0, 1]) - ax.set_yticks([0, 1]) - ax.set_xticklabels(["Pred 0", "Pred 1"]) - ax.set_yticklabels(["True 0", "True 1"]) - ax.set_title(f"Confusion Matrix — {key}") - fig.tight_layout() - fig.savefig(fig_dir / f"vertex_signature_confusion_{safe_name}.png", dpi=150) - plt.close(fig) - - def summary(self) -> None: - data_dir = self.output_dir / "data" - - config_path = data_dir / "study_config.yaml" - config_data = self._r_config_data() - if self._sfreq is not None: - config_data["sfreq"] = self._sfreq - with open(config_path, "w") as f: - yaml.dump(config_data, f, default_flow_style=False) - - try: - r_dir = _find_r_script_dir() - r_script = r_dir / "vertex_signature_analysis.R" - if r_script.exists(): - cmd = [ - "Rscript", str(r_script), - "--data-dir", str(data_dir), - "--config", str(config_path), - "--output-dir", str(self.output_dir), - "--fig-dir", str(self.fig_dir), - "--tbl-dir", str(self.tbl_dir), - ] - cmd.extend(self._r_no_figures_flags()) - result = subprocess.run(cmd, capture_output=True, text=True, timeout=self._r_timeout) - if result.returncode == 0: - return - except (FileNotFoundError, subprocess.TimeoutExpired): - pass - - self._write_python_summary() - - def _write_python_summary(self) -> None: - tbl_dir = self.tbl_dir - - models = ", ".join(classifier_label(c) for c in self._classifiers) - lines = [ - "# Neural Signature Analysis Summary", - "", - f"**Study**: {self.config.name}", - "**Analysis**: Whole-brain vertex-level neural signature (classification)", - f"**Classifiers**: {models}", - f"**CV method**: {self._cv_method}", - f"**Permutations**: {self._n_permutations}", - "", - "## Methods", - "", - "Each classifier, with Leave-One-Out Cross-Validation (LOOCV), was trained to " - "distinguish groups from the spatial pattern of vertex-level relative band " - "power. Significance was assessed by permutation testing: group labels were " - "shuffled and LOOCV accuracy recomputed to build a null distribution. For " - "linear models (SVM/logistic/LDA), feature importance is the mean |coefficient| " - "across folds; non-linear models (RBF SVM) report accuracy only.", - "", - ] - - if self._epoch_config is not None: - lines.append( - f"**Epoch sampling**: {self._epoch_config.get('n_epochs', 80)} epochs " - f"of {self._epoch_config.get('epoch_duration_sec', 2.0)}s" - ) - lines.append("") - - results_csv = tbl_dir / "vertex_signature_results.csv" - if results_csv.exists(): - results_df = pd.read_csv(results_csv) - has_model = "model" in results_df.columns - lines.append("## Results") - lines.append("") - lines.append( - "| Model | Band | Accuracy | p-value | Sensitivity | Specificity | AUC | 95% CI |" - ) - lines.append( - "|-------|------|----------|---------|-------------|-------------|-----|--------|" - ) - for _, row in results_df.iterrows(): - model = row["model"] if has_model else "—" - lines.append( - f"| {model} | {row['band']} | {row['accuracy']:.1%} | {row['p_value']:.4f} | " - f"{row['sensitivity']:.1%} | {row['specificity']:.1%} | " - f"{row['auc']:.3f} | [{row['ci_lower']:.1%}, {row['ci_upper']:.1%}] |" - ) - lines.append("") - - lines.extend([ - "## Output Files", - "", - "- `data/vertex_signature_features.csv` — feature matrix (per-subject per-vertex band power)", - "- `tables/vertex_signature_results.csv` — classification results per band", - "- `figures/vertex_signature_importance_*.png` — feature importance glass brains", - "- `figures/vertex_signature_null_*.png` — permutation null distribution histograms", - "- `figures/vertex_signature_confusion_*.png` — confusion matrices", - "", - ]) - - summary_path = self.output_dir / "ANALYSIS_SUMMARY.md" - summary_path.write_text("\n".join(lines)) - logger.info("Wrote %s", summary_path) diff --git a/src/source_analytics/analyses/vertex_spatial_analysis.py b/src/source_analytics/analyses/vertex_spatial_analysis.py deleted file mode 100644 index ecfde3e..0000000 --- a/src/source_analytics/analyses/vertex_spatial_analysis.py +++ /dev/null @@ -1,93 +0,0 @@ -"""Vertex-level spatial analysis — RETIRED. - -This module used to fit a per-contrast spatial-covariance GLS (nlme::gls, -corExp + nugget) on per-vertex band power as a robustness check on the vertex -group difference. It iterated the legacy ``config$contrasts`` block, which the -declarative ``design:``/``hypotheses:`` spec no longer populates, and the -spatial-covariance table was never a manuscript result. It was retired rather -than migrated (2026-06): spatially-resolved vertex inference is delivered by -``vertex_cluster`` (cluster-based permutation glass-brain maps) and -``vertex_nbs`` (network-based statistic). - -The module is kept in the registry so old configs and ``--analysis -vertex_spatial`` (and its ``spatial_lmm`` alias) still resolve, but it does no -work: it neither loads source estimates nor calls R. It writes well-formed -empty result tables and a retirement note so downstream consumers (the -gallery) find the expected files, then exits cleanly. ``R/vertex_spatial_analysis.R`` -is retained for reference only and is not invoked. -""" - -from __future__ import annotations - -import logging -from pathlib import Path - -import pandas as pd - -from ..config import StudyConfig -from ..io.discovery import SubjectInfo -from .base import BaseAnalysis - -logger = logging.getLogger(__name__) - -RETIRE_NOTE = ( - "vertex_spatial is RETIRED (design-spec migration, 2026-06). The per-contrast " - "GLS spatial-covariance robustness model iterated config$contrasts, which the " - "declarative design:/hypotheses: spec no longer populates. Spatially-resolved " - "vertex inference is provided by vertex_cluster (cluster-permutation glass-brain " - "maps) and vertex_nbs (network-based statistic). No data was processed." -) - - -class VertexSpatialAnalysis(BaseAnalysis): - """RETIRED — writes empty result tables + a note; processes no subjects.""" - - name = "vertex_spatial" - SELECTABLE = {"band": "frequency band"} - - def __init__(self, config: StudyConfig, output_dir: Path): - super().__init__(config, output_dir) - self._warned = False - - def _warn_once(self) -> None: - if not self._warned: - logger.warning(RETIRE_NOTE) - self._warned = True - - def setup(self) -> None: - self._warn_once() - - def process_subject(self, subject: SubjectInfo) -> None: # noqa: D401 - """No-op: the retired module loads nothing.""" - - def aggregate(self) -> None: # noqa: D401 - """No-op: nothing was computed.""" - - def statistics(self) -> None: - """Write the well-formed empty tables downstream consumers expect.""" - tbl_dir = self.tbl_dir - tbl_dir.mkdir(parents=True, exist_ok=True) - for name in ("vertex_spatial_results.csv", "vertex_spatial_residuals.csv"): - pd.DataFrame().to_csv(tbl_dir / name, index=False) - logger.info("vertex_spatial (retired): wrote empty result tables to %s", tbl_dir) - - def figures(self) -> None: # noqa: D401 - """No figures: there are no results.""" - - def summary(self) -> None: - lines = [ - "# Vertex Spatial Analysis — RETIRED", - "", - f"**Study**: {self.config.name}", - "", - RETIRE_NOTE, - "", - "## Output Files", - "", - "- `tables/vertex_spatial_results.csv` — empty (retired)", - "- `tables/vertex_spatial_residuals.csv` — empty (retired)", - "", - ] - path = self.output_dir / "ANALYSIS_SUMMARY.md" - path.write_text("\n".join(lines)) - logger.info("Wrote %s", path) diff --git a/src/source_analytics/analyses/vertex_specparam_analysis.py b/src/source_analytics/analyses/vertex_specparam_analysis.py deleted file mode 100644 index a63bcb5..0000000 --- a/src/source_analytics/analyses/vertex_specparam_analysis.py +++ /dev/null @@ -1,875 +0,0 @@ -"""Vertex-level spectral parameterization analysis. - -Fits aperiodic (1/f) models at each vertex to decompose the power spectrum -into aperiodic (1/f) and oscillatory components. Detects peaks in every -configured frequency band and tests group differences in aperiodic -exponent, offset, and per-band peak presence using cluster permutation -and chi-squared tests. -""" - -from __future__ import annotations - -import logging -import pickle -import subprocess -from pathlib import Path - -import numpy as np -import pandas as pd -import yaml -from scipy import stats as sp_stats - -from ..config import StudyConfig -from ..io.discovery import SubjectInfo -from ..io.loader import SubjectLoader -from ..spectral.aperiodic import band_peak_reachability, resolve_freq_range -from ..spectral.vertex import compute_psd_vertices -from ..spectral.vertex_aperiodic import fit_aperiodic_vertices -from ..spectral.epoch_sampler import sample_epochs -from ..stats.cluster_permutation import ( - cluster_permutation_test, - has_significant_cluster as _has_significant_cluster, - hedges_g, -) -from ..viz.glass_brain import plot_glass_brain, plot_band_comparison -from .base import BaseAnalysis - -logger = logging.getLogger(__name__) - - -def _find_r_script_dir() -> Path: - pkg_root = Path(__file__).resolve().parent.parent.parent.parent - r_dir = pkg_root / "R" - if r_dir.is_dir(): - return r_dir - for candidate in [Path.cwd() / "R", Path(__file__).parent.parent.parent / "R"]: - if candidate.is_dir(): - return candidate - raise FileNotFoundError("Cannot find R/ scripts directory") - - -class VertexSpecparamAnalysis(BaseAnalysis): - """Vertex-level spectral parameterization analysis.""" - - name = "vertex_specparam" - SELECTABLE = {"hypothesis": "declared hypothesis"} - - def __init__(self, config: StudyConfig, output_dir: Path): - super().__init__(config, output_dir) - self._param_rows: list[dict] = [] - self._peak_rows: list[dict] = [] - self._source_coords: np.ndarray | None = None - self._sfreq: float | None = None - self._subject_data: dict[str, dict] = {} - self._subject_groups: dict[str, str] = {} - - # Config - sp_cfg = config.raw.get("vertex_specparam", {}) - # Shared package default (2-50 Hz). The previous default here was - # 1-100 Hz, which spans the 57-63 Hz notch and the >80 Hz roll-off and - # collapsed the vertex fits to r^2~0.16 / exponent~0.04 (flat). - self._freq_range = resolve_freq_range(sp_cfg) - self._peak_width_limits = tuple(sp_cfg.get("peak_width_limits", [1.0, 12.0])) - self._max_n_peaks = int(sp_cfg.get("max_n_peaks", 6)) - # Separate, wider window for PEAK detection. The narrow aperiodic window - # makes every band outside it structurally undetectable, which would turn - # the fit-window choice into an unfalsifiable assertion: the borders are - # justified by where the peaks are, so the peaks must be measured over a - # window that does not presuppose the answer. Unset => single fit. - self._peak_freq_range = ( - resolve_freq_range(sp_cfg, key="peak_freq_range") - if sp_cfg.get("peak_freq_range") else self._freq_range - ) - - # Frequency bands for peak detection - self._bands = dict(config.bands) - # Only bands the PEAK window can reach get peak columns — an unreachable - # band's absence is structural, and emitting False for it fabricates a - # measured null (see spectral.aperiodic.band_peak_reachability). - self._band_reach = band_peak_reachability(self._bands, self._peak_freq_range) - self._peak_bands = { - n: b for n, b in self._bands.items() if self._band_reach[n]["reachable"] - } - self._band_keys = { - name: name.lower().replace(" ", "_") for name in self._peak_bands - } - - wb_cfg = config.vertex - self._n_permutations = int(wb_cfg.get("n_permutations", 1000)) - self._adjacency_distance = float(wb_cfg.get("adjacency_distance_mm", 5.0)) - self._cluster_threshold = float(wb_cfg.get("cluster_threshold", 2.0)) - self._noise_exclude = wb_cfg.get("noise_exclude_hz") - if self._noise_exclude is not None: - self._noise_exclude = tuple(self._noise_exclude) - - # Global epoch_sampling → vertex: block → per-analysis block (see base). - self._epoch_config = self._vertex_epoch_config() - self._cluster_results: dict = {} - - def setup(self) -> None: - self._param_rows.clear() - self._peak_rows.clear() - self._subject_data.clear() - self._subject_groups.clear() - self._source_coords = None - self._cluster_results.clear() - - def _compute_subject(self, subject: SubjectInfo): - """Pure per-subject specparam/FOOOF compute (parallel-safe).""" - loader = SubjectLoader(subject.data_dir) - uid = f"{subject.group}_{subject.subject_id}" - - stc_data = loader.load_source_timecourses() - sfreq = loader.load_sfreq() - coords = loader.load_source_coords() - - # Compute PSD — must span the WIDER of the two windows - fmax = max(self._freq_range[1], self._peak_freq_range[1]) + 10 - if self._epoch_config is not None: - epochs = sample_epochs( - stc_data, sfreq, - epoch_duration_sec=self._epoch_config.get("epoch_duration_sec", 2.0), - n_epochs=self._epoch_config.get("n_epochs", 80), - seed=self._epoch_config.get("seed", 42), - n_bootstrap=self._epoch_config.get("n_bootstrap", 1), - ) - all_psd = [] - for ep in epochs: - f, p = compute_psd_vertices(ep, sfreq, fmax=fmax) - all_psd.append(p) - freqs = f - psd = np.mean(all_psd, axis=0) - else: - freqs, psd = compute_psd_vertices(stc_data, sfreq, fmax=fmax) - - # Fit specparam at each vertex - params = fit_aperiodic_vertices( - freqs, psd, - freq_range=self._freq_range, - max_n_peaks=self._max_n_peaks, - peak_width_limits=self._peak_width_limits, - bands=self._bands, - peak_freq_range=self._peak_freq_range, - ) - - param_rows: list[dict] = [] - n_vertices = psd.shape[0] - for vi in range(n_vertices): - row = { - "subject": uid, - "group": subject.group, - "vertex_idx": vi, - "exponent": float(params["exponent"][vi]), - "offset": float(params["offset"][vi]), - # Offset at the fit-window centre — the one to report alongside - # exponent (the 1 Hz-referenced offset is mechanically coupled to - # the slope). Matches roi_aperiodic/electrode_aperiodic. - "offset_centered": float(params["offset_centered"][vi]), - "r_squared": float(params["r_squared"][vi]), - "n_peaks": int(params["n_peaks"][vi]), - "n_peaks_wide": int(params["n_peaks_wide"][vi]), - "method": params["method"][vi], - "fit_fmin": float(self._freq_range[0]), - "fit_fmax": float(self._freq_range[1]), - "peak_fmin": float(self._peak_freq_range[0]), - "peak_fmax": float(self._peak_freq_range[1]), - } - for key in self._band_keys.values(): - row[f"has_{key}_peak"] = bool(params[f"has_{key}_peak"][vi]) - row[f"{key}_peak_freq"] = float(params[f"{key}_peak_freq"][vi]) - row[f"{key}_peak_power"] = float(params[f"{key}_peak_power"][vi]) - param_rows.append(row) - - # Long-format inventory of every peak found by the peak-window fit. - # This is the raw material for the fit-window diagnostic: where the - # oscillations actually are, independent of any band definition. - peak_rows = [ - {"subject": uid, "group": subject.group, **pk} - for vertex_peaks in params["peaks_all"] for pk in vertex_peaks - ] - - return { - "uid": uid, "group": subject.group, "sfreq": float(sfreq), - "source_coords": coords, "params": params, "param_rows": param_rows, - "peak_rows": peak_rows, - } - - def _merge_subject(self, payload) -> None: - uid = payload["uid"] - if self._sfreq is None: - self._sfreq = payload["sfreq"] - if self._source_coords is None: - self._source_coords = payload["source_coords"] - self._subject_groups[uid] = payload["group"] - self._subject_data[uid] = payload["params"] - self._param_rows.extend(payload["param_rows"]) - self._peak_rows.extend(payload.get("peak_rows", [])) - - def aggregate(self) -> None: - data_dir = self.output_dir / "data" - - param_df = pd.DataFrame(self._param_rows) - if param_df.empty: - logger.warning("No specparam data collected") - return - param_df.to_csv(data_dir / "vertex_specparam.csv", index=False) - logger.info("Exported vertex_specparam.csv (%d rows)", len(param_df)) - - if self._peak_rows: - peak_df = pd.DataFrame(self._peak_rows) - peak_df.to_csv(data_dir / "peak_inventory.csv", index=False) - logger.info("Exported peak_inventory.csv (%d peaks)", len(peak_df)) - self._write_fit_window_diagnostic(peak_df) - - if self._source_coords is not None: - coords_df = pd.DataFrame(self._source_coords, columns=["x", "y", "z"]) - coords_df.index.name = "vertex_idx" - coords_df.to_csv(data_dir / "source_coords.csv") - - def _write_fit_window_diagnostic(self, peak_df: pd.DataFrame) -> None: - """Check the aperiodic fit window against where the oscillations are. - - Gerster et al. 2022: "Oscillations crossing the fitting range borders - must be avoided for all investigated power spectra" — a peak sitting on - a border produces large exponent error. That rule is testable, and this - table tests it on the study's own spectra instead of asserting it. - - A peak is treated as crossing a border when its support (centre - frequency +/- half the specparam bandwidth, i.e. +/-1 SD of the fitted - Gaussian) straddles that border. This catches both a peak centred inside - the window whose tail leaks out, and one centred outside whose tail - leaks in. - """ - fmin, fmax = self._freq_range - n_fits = len(self._param_rows) or 1 - - cf = peak_df["center_frequency"].to_numpy(dtype=float) - bw = peak_df["bandwidth"].to_numpy(dtype=float) - half = np.where(np.isfinite(bw), bw / 2.0, 0.0) - lo_edge, hi_edge = cf - half, cf + half - cross_lo = (lo_edge < fmin) & (hi_edge > fmin) - cross_hi = (lo_edge < fmax) & (hi_edge > fmax) - - def _row(name, blo, bhi, mask, reach): - sub_cf = cf[mask] - n = int(mask.sum()) - return { - "band": name, - "band_lo": blo, - "band_hi": bhi, - "aperiodic_fmin": float(fmin), - "aperiodic_fmax": float(fmax), - "peak_fmin": float(self._peak_freq_range[0]), - "peak_fmax": float(self._peak_freq_range[1]), - "reachable": reach.get("reachable", True), - "censored": reach.get("censored", False), - "frac_visible": reach.get("frac_visible", 1.0), - "n_peaks": n, - "peaks_per_fit": float(n / n_fits), - "cf_median": float(np.median(sub_cf)) if n else float("nan"), - "cf_p5": float(np.percentile(sub_cf, 5)) if n else float("nan"), - "cf_p95": float(np.percentile(sub_cf, 95)) if n else float("nan"), - "n_cross_fmin": int((cross_lo & mask).sum()), - "n_cross_fmax": int((cross_hi & mask).sum()), - "frac_crossing": float( - ((cross_lo | cross_hi) & mask).sum() / n) if n else 0.0, - } - - rows = [_row( - "ALL", float(self._peak_freq_range[0]), float(self._peak_freq_range[1]), - np.ones(len(cf), dtype=bool), {}, - )] - for name, (blo, bhi) in self._bands.items(): - reach = self._band_reach[name] - in_band = (cf >= blo) & (cf <= bhi) - rows.append(_row(name, float(blo), float(bhi), in_band, reach)) - - # Peaks the aperiodic window deliberately excludes — context for the - # crossing counts: a window that excludes a lot of oscillatory activity - # is only defensible if that activity is clear of the borders. - outside = (cf < fmin) | (cf > fmax) - diag = pd.DataFrame(rows) - diag.to_csv(self.tbl_dir / "fit_window_diagnostic.csv", index=False) - - n_total = len(cf) - n_cross = int((cross_lo | cross_hi).sum()) - logger.info( - "Fit-window diagnostic: %d peaks over %.4g-%.4g Hz; %d (%.1f%%) " - "cross the aperiodic borders %.4g/%.4g Hz; %d (%.1f%%) lie outside " - "the aperiodic window entirely.", - n_total, self._peak_freq_range[0], self._peak_freq_range[1], - n_cross, 100.0 * n_cross / max(n_total, 1), fmin, fmax, - int(outside.sum()), 100.0 * outside.sum() / max(n_total, 1), - ) - - def statistics(self) -> None: - if self._source_coords is None: - logger.error("No source coordinates") - return - - coords = self._source_coords - tbl_dir = self.tbl_dir - all_stats = [] - - for contrast in self._pairwise_contrasts(): - group_a_uids = [ - uid for uid, g in self._subject_groups.items() if g == contrast.group_a - ] - group_b_uids = [ - uid for uid, g in self._subject_groups.items() if g == contrast.group_b - ] - - if not group_a_uids or not group_b_uids: - continue - - label_a = self.config.get_group_label(contrast.group_a) - label_b = self.config.get_group_label(contrast.group_b) - - # Cluster permutation on exponent and offset maps - for param_name in ["exponent", "offset"]: - data_a = np.array([ - self._subject_data[uid][param_name] for uid in group_a_uids - ]) - data_b = np.array([ - self._subject_data[uid][param_name] for uid in group_b_uids - ]) - - result = cluster_permutation_test( - data_a, data_b, coords, - n_perms=self._n_permutations, - threshold=self._cluster_threshold, - distance_mm=self._adjacency_distance, - seed=42, - ) - - g_map = hedges_g(data_a, data_b) - - self._cluster_results[f"{contrast.name}_{param_name}"] = { - "result": result, - "mean_a": data_a.mean(axis=0), - "mean_b": data_b.mean(axis=0), - "group_labels": (label_a, label_b), - "param": param_name, - } - - for vi in range(len(result.t_map)): - cid = int(result.cluster_labels[vi]) - cp = float(result.cluster_pvalues[cid - 1]) if cid > 0 else float("nan") - all_stats.append({ - "contrast": contrast.name, - "parameter": param_name, - "vertex_idx": vi, - "t": float(result.t_map[vi]), - "p": float(result.p_map[vi]), - "hedges_g": float(g_map[vi]), - "cluster_id": cid, - # Corrected per-cluster p and significance — cluster_id alone - # is NOT significance (clusters are candidates pre-permutation). - "cluster_p": cp, - "significant": bool(cid > 0 and cp < 0.05), - }) - - # Per-band chi-squared tests on peak presence + optional - # cluster permutation on peak power - n_a, n_b = len(group_a_uids), len(group_b_uids) - all_chi2_stats: list[dict] = [] - - # Detectable bands only. A band the peak window cannot reach has no - # column to test, and testing it would emit p=1.0 at every vertex — - # a null the data never had the power to reject. - for band_name in self._peak_bands: - key = self._band_keys[band_name] - col = f"has_{key}_peak" - - peak_a = np.array([ - self._subject_data[uid][col] for uid in group_a_uids - ]) # (n_a, n_vertices) bool - peak_b = np.array([ - self._subject_data[uid][col] for uid in group_b_uids - ]) - - rate_a = peak_a.mean(axis=0) - rate_b = peak_b.mean(axis=0) - - for vi in range(len(rate_a)): - a_yes = int(peak_a[:, vi].sum()) - a_no = n_a - a_yes - b_yes = int(peak_b[:, vi].sum()) - b_no = n_b - b_yes - - table = np.array([[a_yes, a_no], [b_yes, b_no]]) - if table.sum() > 0 and np.all(table.sum(axis=0) > 0): - chi2, p_val, _, _ = sp_stats.chi2_contingency( - table, correction=True, - ) - else: - chi2, p_val = 0.0, 1.0 - - all_chi2_stats.append({ - "contrast": contrast.name, - "band": band_name, - "band_key": key, - "vertex_idx": vi, - "rate_a": float(rate_a[vi]), - "rate_b": float(rate_b[vi]), - "chi2": float(chi2), - "p": float(p_val), - }) - - # Cluster permutation on peak power for bands with - # enough detected peaks (>=10% of vertices overall) - overall_rate = np.concatenate([peak_a, peak_b]).mean(axis=0) - if overall_rate.mean() >= 0.10: - power_a = np.array([ - self._subject_data[uid][f"{key}_peak_power"] - for uid in group_a_uids - ]) - power_b = np.array([ - self._subject_data[uid][f"{key}_peak_power"] - for uid in group_b_uids - ]) - power_a = np.nan_to_num(power_a, nan=0.0) - power_b = np.nan_to_num(power_b, nan=0.0) - - result = cluster_permutation_test( - power_a, power_b, coords, - n_perms=self._n_permutations, - threshold=self._cluster_threshold, - distance_mm=self._adjacency_distance, - seed=42, - ) - g_map = hedges_g(power_a, power_b) - - self._cluster_results[ - f"{contrast.name}_{key}_peak_power" - ] = { - "result": result, - "mean_a": power_a.mean(axis=0), - "mean_b": power_b.mean(axis=0), - "group_labels": (label_a, label_b), - "param": f"{key}_peak_power", - } - - for vi in range(len(result.t_map)): - cid = int(result.cluster_labels[vi]) - cp = float(result.cluster_pvalues[cid - 1]) if cid > 0 else float("nan") - all_stats.append({ - "contrast": contrast.name, - "parameter": f"{key}_peak_power", - "vertex_idx": vi, - "t": float(result.t_map[vi]), - "p": float(result.p_map[vi]), - "hedges_g": float(g_map[vi]), - "cluster_id": cid, - "cluster_p": cp, - "significant": bool(cid > 0 and cp < 0.05), - }) - - if all_chi2_stats: - chi2_df = pd.DataFrame(all_chi2_stats) - chi2_df.to_csv(tbl_dir / "band_peak_chi2.csv", index=False) - # Backward compat: gamma-only subset - gamma_keys = [ - k for k in self._band_keys.values() if "gamma" in k - ] - gamma_sub = chi2_df[chi2_df["band_key"].isin(gamma_keys)] - if not gamma_sub.empty: - gamma_sub.to_csv( - tbl_dir / "gamma_peak_chi2.csv", index=False, - ) - - if all_stats: - stats_df = pd.DataFrame(all_stats) - stats_df.to_csv(tbl_dir / "vertex_specparam_stats.csv", index=False) - logger.info("Exported vertex_specparam_stats.csv (%d rows)", len(stats_df)) - - # Save cluster results for --steps figures support - if self._cluster_results: - data_dir = self.output_dir / "data" - pkl_data = {} - for key, info in self._cluster_results.items(): - result = info["result"] - pkl_data[key] = { - "t_map": result.t_map, - "p_map": result.p_map, - "cluster_labels": result.cluster_labels, - "cluster_pvalues": result.cluster_pvalues, - "cluster_stats": result.cluster_stats, - "n_clusters": result.n_clusters, - "n_permutations": result.n_permutations, - "mean_a": info["mean_a"], - "mean_b": info["mean_b"], - "group_labels": info["group_labels"], - "param": info["param"], - } - with open(data_dir / "specparam_cluster_results.pkl", "wb") as f: - pickle.dump(pkl_data, f) - logger.info("Saved specparam_cluster_results.pkl") - - # --- Declarative hypotheses (hypothesis layer; additive, map+cluster) --- - # exponent/offset are broadband per-vertex maps (no band dimension). - from ..hypothesis import write_module_hypotheses_perm - - if self._source_coords is not None and self._subject_groups: - maps_by_cell = { - ("broadband", param): { - uid: self._subject_data[uid][param] - for uid in self._subject_groups - } - for param in ["exponent", "offset"] - } - wanted_hyp = self._selection.get("hypothesis") - write_module_hypotheses_perm( - maps_by_cell, self._subject_groups, self._source_coords, self.config, - self.tbl_dir, prefix="vertex_specparam", - n_perms=self._n_permutations, threshold=self._cluster_threshold, - distance_mm=self._adjacency_distance, - hypothesis=",".join(sorted(wanted_hyp)) if wanted_hyp else None, - atlas_dir=self._atlas_dir, - ) - - def _load_state_from_disk(self) -> bool: - """Load saved state from pickle for --steps figures support.""" - from ..stats.cluster_permutation import ClusterResult - - data_dir = self.output_dir / "data" - pkl_path = data_dir / "specparam_cluster_results.pkl" - if not pkl_path.exists(): - logger.warning("No saved state at %s; skipping figures", pkl_path) - return False - - with open(pkl_path, "rb") as f: - saved = pickle.load(f) - - for key, d in saved.items(): - self._cluster_results[key] = { - "result": ClusterResult( - t_map=d["t_map"], - p_map=d["p_map"], - cluster_labels=d["cluster_labels"], - cluster_pvalues=d["cluster_pvalues"], - cluster_stats=d["cluster_stats"], - n_clusters=d.get("n_clusters", 0), - n_permutations=d.get("n_permutations", 0), - ), - "mean_a": d["mean_a"], - "mean_b": d["mean_b"], - "group_labels": d["group_labels"], - "param": d["param"], - } - - # Load source coords - coords_csv = data_dir / "source_coords.csv" - if coords_csv.exists(): - coords_df = pd.read_csv(coords_csv) - self._source_coords = coords_df[["x", "y", "z"]].values - - logger.info("Loaded vertex_specparam state from %s", pkl_path) - return True - - def figures(self) -> None: - # Load from disk if in-memory state is missing (--steps support) - if not self._cluster_results or self._source_coords is None: - if not self._load_state_from_disk(): - return - - if self._source_coords is None: - return - - coords = self._source_coords - fig_dir = self.fig_dir - - for key, info in self._cluster_results.items(): - result = info["result"] - if not _has_significant_cluster(result): - continue - param = info["param"] - contrast = info.get("contrast", key) - group_labels = info["group_labels"] - - safe = f"{contrast}_{param}".lower().replace(" ", "_") - plot_band_comparison( - coords=coords, - mean_a=info["mean_a"], - mean_b=info["mean_b"], - t_map=result.t_map, - cluster_labels=result.cluster_labels, - cluster_pvalues=result.cluster_pvalues, - band_name=f"{param} — {contrast}", - group_labels=group_labels, - output_path=fig_dir / f"specparam_{safe}.png", - ) - - self._plot_fit_window_diagnostic() - - # Per-band peak presence maps - band_keys = self._band_keys if self._band_keys else {} - # Detect band keys from CSV if in-memory state is empty - if not band_keys: - data_dir = self.output_dir / "data" - csv_path = data_dir / "vertex_specparam.csv" - if csv_path.exists(): - df = pd.read_csv(csv_path) - import re - for col in df.columns: - m = re.match(r"has_(.+)_peak$", col) - if m: - band_keys[m.group(1)] = m.group(1) # key == key - - for band_name, key in ( - self._band_keys.items() if self._band_keys else - {k: k for k in band_keys}.items() - ): - col = f"has_{key}_peak" - if self._subject_data: - try: - all_peaks = np.array([ - d[col].astype(float) for d in self._subject_data.values() - ]) - mean_rate = all_peaks.mean(axis=0) - except KeyError: - continue - else: - data_dir = self.output_dir / "data" - csv_path = data_dir / "vertex_specparam.csv" - if not csv_path.exists(): - continue - df = pd.read_csv(csv_path) - if col not in df.columns: - continue - mean_rate = df.groupby("vertex_idx")[col].mean().values - - plot_glass_brain( - coords=coords, - values=mean_rate, - title=f"{band_name} Peak Presence Rate", - output_path=fig_dir / f"{key}_peak_presence.png", - cmap="YlOrRd", - vlim=(0, 1), - ) - - def _plot_fit_window_diagnostic(self) -> None: - """The fit-window justification figure: peaks vs the aperiodic borders. - - Regenerated from ``data/peak_inventory.csv`` so ``--steps figures`` - alone reproduces it (no in-memory state). - """ - import matplotlib - matplotlib.use("Agg") - import matplotlib.pyplot as plt - - csv_path = self.output_dir / "data" / "peak_inventory.csv" - if not csv_path.exists(): - return - peaks = pd.read_csv(csv_path) - if peaks.empty: - return - - fmin, fmax = self._freq_range - pmin, pmax = self._peak_freq_range - cf = peaks["center_frequency"].to_numpy(dtype=float) - bw = peaks["bandwidth"].to_numpy(dtype=float) - half = np.where(np.isfinite(bw), bw / 2.0, 0.0) - crossing = ((cf - half < fmin) & (cf + half > fmin)) | \ - ((cf - half < fmax) & (cf + half > fmax)) - - fig, ax = plt.subplots(figsize=(9, 4.5)) - bins = np.linspace(pmin, pmax, 80) - ax.hist(cf[~crossing], bins=bins, color="#4C78A8", - label=f"clear of borders (n={int((~crossing).sum())})") - ax.hist(cf[crossing], bins=bins, color="#E45756", - label=f"support crosses a border (n={int(crossing.sum())})") - - ax.axvspan(pmin, fmin, color="0.85", zorder=0) - ax.axvspan(fmax, pmax, color="0.85", zorder=0) - for border in (fmin, fmax): - ax.axvline(border, color="k", ls="--", lw=1.5) - ax.set_xlim(pmin, pmax) - ax.set_xlabel("Peak centre frequency (Hz)") - ax.set_ylabel("Peaks detected") - ax.set_title( - f"Fit-window check — peaks detected over {pmin:g}-{pmax:g} Hz " - f"vs aperiodic window {fmin:g}-{fmax:g} Hz (dashed)" - ) - ax.legend(frameon=False, fontsize=9) - fig.tight_layout() - fig.savefig(self.fig_dir / "fit_window_diagnostic.png", dpi=150) - plt.close(fig) - - def summary(self) -> None: - data_dir = self.output_dir / "data" - - config_path = data_dir / "study_config.yaml" - config_data = self._r_config_data() - if self._sfreq is not None: - config_data["sfreq"] = self._sfreq - with open(config_path, "w") as f: - yaml.dump(config_data, f, default_flow_style=False) - - try: - r_dir = _find_r_script_dir() - r_script = r_dir / "vertex_specparam_analysis.R" - if r_script.exists(): - cmd = [ - "Rscript", str(r_script), - "--data-dir", str(data_dir), - "--config", str(config_path), - "--output-dir", str(self.output_dir), - "--fig-dir", str(self.fig_dir), - "--tbl-dir", str(self.tbl_dir), - ] - cmd.extend(self._r_no_figures_flags()) - result = subprocess.run(cmd, capture_output=True, text=True, timeout=self._r_timeout) - if result.returncode == 0: - return - except (FileNotFoundError, subprocess.TimeoutExpired): - pass - - self._write_python_summary() - - def _fit_window_summary_lines(self) -> list[str]: - """Report the border check before the results it underwrites.""" - diag_csv = self.tbl_dir / "fit_window_diagnostic.csv" - if not diag_csv.exists(): - return [] - diag = pd.read_csv(diag_csv) - all_rows = diag[diag["band"] == "ALL"] - if all_rows.empty: - return [] - a = all_rows.iloc[0] - - frac = float(a["frac_crossing"]) - if frac < 0.05: - verdict = "SATISFIED — the borders sit in spectral gaps on this data." - elif frac < 0.15: - verdict = ("MARGINAL — a minority of peaks touch a border, so the " - "exponents carry some added error.") - else: - verdict = ("VIOLATED — peaks sit on the borders; the window needs " - "revisiting for this dataset.") - - lines = [ - "## Fit-Window Diagnostic", - "", - "Gerster et al. (2022): oscillations crossing the fit borders must be " - "avoided, since a peak on a border inflates exponent error. A peak " - "counts as crossing when its support (centre frequency ± half the " - "specparam bandwidth) straddles a border.", - "", - f"**{int(a['n_peaks'])} peaks** detected over " - f"{a['peak_fmin']:g}–{a['peak_fmax']:g} Hz. " - f"**{int(a['n_cross_fmin']) + int(a['n_cross_fmax'])} ({100 * frac:.1f}%)** " - f"cross an aperiodic border ({a['aperiodic_fmin']:g} / " - f"{a['aperiodic_fmax']:g} Hz).", - "", - f"**Verdict: {verdict}**", - "", - "| Band | Range (Hz) | Reachable | Censored | Peaks | Median CF | Crossing |", - "|------|-----------|-----------|----------|-------|-----------|----------|", - ] - for _, r in diag[diag["band"] != "ALL"].iterrows(): - cf_med = "—" if pd.isna(r["cf_median"]) else f"{r['cf_median']:.1f}" - lines.append( - f"| {r['band']} | {r['band_lo']:g}–{r['band_hi']:g} | " - f"{'yes' if r['reachable'] else '**no**'} | " - f"{'**yes**' if r['censored'] else 'no'} | " - f"{int(r['n_peaks'])} | {cf_med} | {100 * r['frac_crossing']:.1f}% |" - ) - lines.extend([ - "", - "Unreachable bands emit no peak columns at all — their absence is a " - "property of the window, not a measurement. Censored bands extend " - "past the peak window, so their detection rates are a lower bound.", - "", - ]) - return lines - - def _write_python_summary(self) -> None: - tbl_dir = self.tbl_dir - - lines = [ - "# Spectral Parameterization (Vertex-Level) Summary", - "", - f"**Study**: {self.config.name}", - "**Analysis**: Vertex-level spectral parameterization (aperiodic + peaks)", - f"**Aperiodic fit range**: {self._freq_range[0]:g}-{self._freq_range[1]:g} Hz", - f"**Peak detection range**: {self._peak_freq_range[0]:g}-" - f"{self._peak_freq_range[1]:g} Hz" - + (" (separate wider fit)" - if self._peak_freq_range != self._freq_range else " (same fit)"), - f"**Max peaks**: {self._max_n_peaks}", - f"**Peak width limits**: {self._peak_width_limits[0]}-{self._peak_width_limits[1]} Hz", - "", - "## Methods", - "", - "Spectral parameterization (specparam/FOOOF) was applied to the PSD at " - "each source vertex to decompose the spectrum into aperiodic (1/f) and " - "oscillatory components. Peaks were detected in all configured frequency " - "bands to determine whether power elevations reflect true oscillatory " - "peaks or broadband spectral shifts. Group differences in aperiodic " - "exponent and offset were tested using cluster-based permutation testing. " - "Per-band peak presence rates were compared using per-vertex chi-squared " - "tests. For bands with sufficient peak prevalence (>=10% of vertices), " - "cluster permutation was also applied to peak power maps.", - "", - ] - - if self._epoch_config is not None: - lines.append( - f"**Epoch sampling**: {self._epoch_config.get('n_epochs', 80)} epochs " - f"of {self._epoch_config.get('epoch_duration_sec', 2.0)}s" - ) - lines.append("") - - lines.extend(self._fit_window_summary_lines()) - - # Specparam stats - stats_csv = tbl_dir / "vertex_specparam_stats.csv" - if stats_csv.exists(): - stats_df = pd.read_csv(stats_csv) - lines.append("## Aperiodic Parameter Results") - lines.append("") - for param in stats_df["parameter"].unique(): - sub = stats_df[stats_df["parameter"] == param] - n_clust = len(set(sub["cluster_id"]) - {0}) - lines.append(f"- **{param}**: {n_clust} clusters found") - lines.append("") - - # Per-band chi-squared results - chi2_csv = tbl_dir / "band_peak_chi2.csv" - if not chi2_csv.exists(): - chi2_csv = tbl_dir / "gamma_peak_chi2.csv" # backward compat - if chi2_csv.exists(): - chi2_df = pd.read_csv(chi2_csv) - lines.append("## Peak Presence by Band") - lines.append("") - if "band" in chi2_df.columns: - for band_name in chi2_df["band"].unique(): - sub = chi2_df[chi2_df["band"] == band_name] - n_sig = len(sub[sub["p"] < 0.05]) - lines.append( - f"- **{band_name}**: {n_sig}/{len(sub)} vertices with " - "significant group differences (uncorrected p<0.05)" - ) - else: - n_sig = len(chi2_df[chi2_df["p"] < 0.05]) - lines.append( - f"- {n_sig}/{len(chi2_df)} vertices with significant " - "group differences (uncorrected p<0.05)" - ) - lines.append("") - - lines.extend([ - "## Output Files", - "", - "- `data/vertex_specparam.csv` — per-subject per-vertex specparam parameters", - "- `tables/vertex_specparam_stats.csv` — cluster permutation results", - "- `tables/band_peak_chi2.csv` — per-band peak presence chi-squared tests", - "- `figures/specparam_*.png` — aperiodic parameter glass brains", - "- `figures/{band}_peak_presence.png` — per-band peak prevalence maps", - "", - ]) - - summary_path = self.output_dir / "ANALYSIS_SUMMARY.md" - summary_path.write_text("\n".join(lines)) - logger.info("Wrote %s", summary_path) diff --git a/src/source_analytics/cli.py b/src/source_analytics/cli.py index c6f1152..dac4533 100644 --- a/src/source_analytics/cli.py +++ b/src/source_analytics/cli.py @@ -11,10 +11,20 @@ from .config import StudyConfig from .core import StudyAnalyzer, ANALYSIS_REGISTRY, ANALYSIS_METADATA, canonical_analysis_name +from .plugins import missing_analysis_hint from .analyses.base import RStepFailed from .analyses.base import VALID_STEPS, BaseAnalysis +def _analysis_name(name: str) -> str: + """argparse type for --analysis: a registered name, or an error saying where it went.""" + if name not in ANALYSIS_REGISTRY: + raise argparse.ArgumentTypeError( + f"unknown analysis '{name}' (see `source-analytics list`).{missing_analysis_hint(name)}" + ) + return name + + def setup_logging(verbose: bool = False): level = logging.DEBUG if verbose else logging.INFO logging.basicConfig( @@ -660,7 +670,9 @@ def cmd_init(args): analyses = [a.strip() for a in args.analyses.split(",") if a.strip()] unknown = [a for a in analyses if a not in ANALYSIS_REGISTRY] if unknown: - print(f"ERROR: unknown analyses: {', '.join(unknown)} (see `source-analytics list`)", file=err) + hints = "".join(missing_analysis_hint(a) for a in unknown) + print(f"ERROR: unknown analyses: {', '.join(unknown)} (see `source-analytics list`).{hints}", + file=err) sys.exit(1) paradigm_block: dict = { @@ -738,7 +750,8 @@ def main(): p_run = subparsers.add_parser("run", help="Run an analysis") p_run.add_argument("--study", required=True, type=Path, help="Path to study YAML config") p_run.add_argument("--paradigm", help="Paradigm name (multi-paradigm configs)") - p_run.add_argument("--analysis", choices=list(ANALYSIS_REGISTRY.keys()), help="Analysis to run") + p_run.add_argument("--analysis", type=_analysis_name, + help="Analysis to run (see `source-analytics list`)") p_run.add_argument( "--profile", metavar="NAME", help="Run under the top-level ':' profile block, which narrows bands, " diff --git a/src/source_analytics/core.py b/src/source_analytics/core.py index dfd9101..1ce0ab7 100644 --- a/src/source_analytics/core.py +++ b/src/source_analytics/core.py @@ -13,15 +13,9 @@ from .analyses.roi_aperiodic_analysis import ROIAperiodicAnalysis from .analyses.roi_connectivity_analysis import ConnectivityAnalysis from .analyses.roi_cross_freq_analysis import ROICrossFreqAnalysis -from .analyses.vertex_cross_freq_analysis import VertexCrossFreqAnalysis -from .analyses.vertex_cluster_analysis import VertexClusterAnalysis from .analyses.electrode_analysis import ElectrodeAnalysis from .analyses.electrode_comparison_analysis import ElectrodeComparisonAnalysis from .analyses.electrode_connectivity_analysis import ElectrodeConnectivityAnalysis -from .analyses.fcd_comparison_analysis import FCDComparisonAnalysis -from .analyses.vertex_connectivity_analysis import VertexConnectivityAnalysis -from .analyses.vertex_specparam_analysis import VertexSpecparamAnalysis -from .analyses.vertex_signature_analysis import VertexSignatureAnalysis from .analyses.electrode_signature_analysis import ElectrodeSignatureAnalysis from .analyses.roi_signature_analysis import ROISignatureAnalysis from .analyses.roi_network_analysis import ( @@ -29,50 +23,32 @@ ROIGraphAnalysis, ROINBSAnalysis, ) -from .analyses.vertex_network_analysis import ( - VertexNetworkAnalysis, - VertexGraphAnalysis, - VertexNBSAnalysis, -) -from .analyses.vertex_spatial_analysis import VertexSpatialAnalysis from .analyses.roi_directed_analysis import ROIDirectedAnalysis -from .analyses.vertex_directed_analysis import VertexDirectedAnalysis from .analyses.roi_evoked_analysis import ROIEvokedAnalysis -from .analyses.vertex_evoked_analysis import VertexEvokedAnalysis from .analyses.electrode_evoked_analysis import ElectrodeEvokedAnalysis from .analyses.electrode_aperiodic_analysis import ElectrodeAperiodicAnalysis +from .plugins import install_plugins, missing_analysis_hint logger = logging.getLogger(__name__) -# Registry of available analyses +# Registry of the built-in analyses. Installed plugins add theirs (plugins.py); +# the vertex analyses moved to the source-analytics-vertex plugin in v0.8.0. ANALYSIS_REGISTRY: dict[str, type[BaseAnalysis]] = { "roi_psd": ROIPsdAnalysis, "roi_aperiodic": ROIAperiodicAnalysis, "roi_connectivity": ConnectivityAnalysis, "roi_cross_freq": ROICrossFreqAnalysis, - "vertex_cluster": VertexClusterAnalysis, "electrode_psd": ElectrodeAnalysis, "electrode_aperiodic": ElectrodeAperiodicAnalysis, "electrode_comparison": ElectrodeComparisonAnalysis, "electrode_connectivity": ElectrodeConnectivityAnalysis, - "fcd_comparison": FCDComparisonAnalysis, - "vertex_connectivity": VertexConnectivityAnalysis, - "vertex_cross_freq": VertexCrossFreqAnalysis, - "vertex_specparam": VertexSpecparamAnalysis, - "vertex_signature": VertexSignatureAnalysis, "electrode_signature": ElectrodeSignatureAnalysis, "roi_signature": ROISignatureAnalysis, "roi_graph": ROIGraphAnalysis, "roi_nbs": ROINBSAnalysis, - "vertex_graph": VertexGraphAnalysis, - "vertex_nbs": VertexNBSAnalysis, "roi_network": ROINetworkAnalysis, # combined alias (graph + NBS) - "vertex_network": VertexNetworkAnalysis, # combined alias (graph + NBS) - "vertex_spatial": VertexSpatialAnalysis, "roi_directed": ROIDirectedAnalysis, - "vertex_directed": VertexDirectedAnalysis, "roi_evoked": ROIEvokedAnalysis, - "vertex_evoked": VertexEvokedAnalysis, "electrode_evoked": ElectrodeEvokedAnalysis, } @@ -82,21 +58,12 @@ "aperiodic": "roi_aperiodic", "pac": "roi_cross_freq", "roi_pac": "roi_cross_freq", - "wholebrain": "vertex_cluster", - "spatial_lmm": "vertex_spatial", - "specparam_vertex": "vertex_specparam", - "mvpa": "vertex_signature", - "vertex_mvpa": "vertex_signature", "transfer_entropy": "roi_directed", "roi_transfer_entropy": "roi_directed", "evoked": "roi_evoked", "electrode": "electrode_psd", } -# Register aliases so old YAML configs still work -for _old, _new in _DEPRECATED_NAMES.items(): - ANALYSIS_REGISTRY[_old] = ANALYSIS_REGISTRY[_new] - def canonical_analysis_name(name: str) -> str: """Canonical registry name for ``name`` (alias-resolved), without warning.""" @@ -140,28 +107,11 @@ def resolve_analysis_name(name: str) -> str: "about": "Graph-theoretic summaries of the ROI connectivity network, per band and connectivity metric. Each subject's ROI-by-ROI matrix is thresholded to a fixed connection density (proportional, default 15% of edges kept) and turned into a graph, from which it computes nodal metrics per ROI (degree, clustering coefficient, betweenness centrality) and whole-network metrics (global efficiency, modularity, small-worldness). Groups are compared per ROI/metric with a Welch t-test, effect sizes are Hedges' g, and p-values are FDR-corrected (Benjamini-Hochberg, q < 0.05). Read it as: which ROIs act as more/less connected hubs (degree, betweenness) or more clustered (clustering), and whether whole-network integration/segregation shifts between groups (an up arrow means the first-listed group is higher). It asks how the connectivity is organized as a network, beyond individual edge strengths."}, "roi_nbs": {"category": "resting", "level": "roi", "domain": "Connectivity", "supplements": "roi_connectivity", "description": "ROI-level Network-Based Statistic (sub-network test)", "about": "The Network-Based Statistic (Zalesky et al. 2010) applied to the ROI connectivity network, per band and metric -- a connected-subnetwork test that has more power than edge-by-edge correction when a group effect is distributed across many connected edges. Every ROI-pair edge gets a group Welch t-statistic; edges exceeding a primary threshold (default t = 2.5) are retained and grouped into connected components; each component's size (edge count) is compared to a permutation null built by relabeling groups and tracking the largest component per permutation, giving component-level family-wise (FWE) control. Read it as: significant sub-networks -- clusters of ROI connections that jointly differ between groups (a component with p < 0.05), rather than any single edge. It complements roi_graph (network topology) and roi_connectivity (individual edges)."}, - "vertex_graph": {"category": "resting", "level": "vertex", "domain": "Connectivity", "supplements": "vertex_connectivity", "description": "Vertex-level multi-density AUC graph metrics", - "about": "Whole-brain graph-theoretic organization of the vertex connectivity network, per band and connectivity metric. To avoid picking one arbitrary connection density, each subject's vertex network is thresholded across a range of densities (default 0.05-0.40 in 0.01 steps) and each global metric is integrated over that range as an area-under-the-curve (AUC) value: global efficiency, characteristic path length, mean clustering, transitivity, modularity, assortativity, mean local efficiency, and small-worldness. Groups are compared on each AUC metric with a permutation test (5000 permutations of group labels), with Hedges' g effect sizes; permutation p-values are reported per metric (raw, marked significant at p < 0.05). Read it as: whole-brain shifts in network integration (efficiency, path length), segregation (clustering, modularity), or small-world balance between groups (an up arrow means the first-listed group is higher). It is the unparcellated, density-integrated counterpart of roi_graph."}, - "vertex_nbs": {"category": "resting", "level": "vertex", "domain": "Connectivity", "supplements": "vertex_connectivity", "description": "Vertex-level Network-Based Statistic (sub-network test)", - "about": "The Network-Based Statistic (Zalesky et al. 2010) applied to the whole-brain vertex connectivity network, per band and metric -- the unparcellated counterpart of roi_nbs. Every vertex-pair edge gets a group Welch t-statistic; edges above a primary threshold (default t = 3.0) are grouped into connected components, and each component's size is compared to a permutation null (5000 relabelings, largest-component-per-permutation) for component-level family-wise (FWE) control. Read it as: significant sub-networks -- spatially distributed sets of vertex connections that jointly differ between groups (a component with p < 0.05), rather than any single vertex pair. It complements vertex_graph (network topology) and vertex_connectivity (per-vertex FCD)."}, "roi_network": {"category": "resting", "level": "roi", "domain": "Connectivity", "supplements": "roi_connectivity", "description": "ROI-level graph theory + NBS (combined alias of roi_graph + roi_nbs)"}, - "vertex_network": {"category": "resting", "level": "vertex", "domain": "Connectivity", "supplements": "vertex_connectivity", "description": "Vertex-level graph theory + NBS (combined alias of vertex_graph + vertex_nbs)"}, "roi_directed": {"category": "resting", "level": "roi", "domain": "Directed", "description": "Directed connectivity (transfer entropy + DTF)", "about": "Directional (who-drives-whom) connectivity between ROIs, in two flavors. Transfer entropy (TE, Schreiber 2000) is a model-free information-theoretic measure -- a binned lag-1 estimator of how much one ROI's past reduces uncertainty about another's future -- computed for every directed ROI pair (callosal tracts excluded); its net asymmetry (te - te-transpose) gives the dominant direction. The optional Directed Transfer Function (DTF, Kaminski & Blinowska 1991) derives directed influence from a multivariate autoregressive (MVAR) model fit with ridge regularization (order 8) -- a deviation from ordinary-least-squares DTF used to stabilize the fit against collinear channels. Groups are compared on TE three ways: a global per-pair Welch t-test (Hedges' g), a within-group one-sample test of net TE against zero (is the driving direction consistent within a group), and a region-pair linear mixed model (te ~ group * region_pair + (1|subject)). Read it as: which ROI pairs show a group difference in directed influence, and which region drives which (an up arrow means the first-listed group is higher). DTF, when selected, currently emits directed edges without the R group stats."}, - "vertex_directed": {"category": "resting", "level": "vertex", "domain": "Directed", "description": "Vertex DTF outflow/inflow/netflow (ridge-MVAR, cluster-corrected)", - "about": "Whole-brain directed connectivity on the dorsal source surface via the Directed Transfer Function (DTF, Kaminski & Blinowska 1991) from a multivariate autoregressive model. Because source vertices are strongly collinear (mean inter-vertex |r| ~ 0.64), the MVAR is fit with ridge regularization (order 8) rather than ordinary least squares, and the fit's stability (spectral radius) is checked. The full all-to-all directed DTF matrix is reduced to three per-vertex maps: outflow (mean directed influence a vertex sends to all others), inflow (mean it receives), and netflow (outflow minus inflow -- net source vs sink). Group differences in each map are tested with a cluster-based permutation test (per-vertex t-statistics clustered by spatial adjacency, cluster-extent FWE from a permutation null; Maris & Oostenveld 2007), with per-vertex Hedges' g. Read it as: spatially-contiguous clusters where the groups differ in how strongly a region drives (outflow), is driven by (inflow), or net-drives (netflow) the rest of the brain -- a cluster with p_corrected < 0.05 marks a region of difference, the sign of its t-values gives direction."}, "electrode_signature": {"category": "resting", "level": "electrode", "domain": "Source vs Sensor", "display_name": "Neural signature", "description": "Sensor-level neural signature (classification/decoding on electrode band power) — the sensor counterpart of roi_signature / vertex_signature"}, "roi_signature": {"category": "resting", "level": "roi", "domain": "Source vs Sensor", "display_name": "Neural signature", "description": "ROI-level neural signature (classification/decoding on per-parcel band power) — the source-side counterpart of electrode_signature; runs on any ROI output"}, - "vertex_signature": {"category": "resting", "level": "vertex", "domain": "Multivariate", "display_name": "Neural signature", "description": "Multivariate/ML neural signature (classification, decoding; PCA-reduced with back-projection)"}, - "vertex_cluster": {"category": "resting", "level": "vertex", "domain": "Spectral", "description": "Vertex-level cluster permutation", - "about": "Whole-brain resting spectral maps on the dorsal source surface: per vertex it computes band power (absolute as mean density in dB/Hz, the same definition as roi_psd, and relative), the 1/f spectral slope, and the peak alpha frequency, then tests where the groups differ. Inference is a cluster-based permutation test -- per-vertex t-statistics are threshold-clustered over neighbouring vertices and each cluster's extent is compared to a permutation null, giving family-wise (FWE) control (Maris & Oostenveld 2007); a threshold-free TFCE variant (Smith & Nichols 2009) is available. Effect sizes are per-vertex Hedges' g. Read it as: spatially-contiguous clusters where the groups differ in a spectral measure -- a cluster with p_corrected < 0.05 marks a region of difference, and the sign of its t-values gives the direction. This is the whole-brain, unparcellated counterpart to the ROI spectral analyses."}, - "vertex_spatial": {"category": "resting", "level": "vertex", "domain": "Spectral", "description": "RETIRED — was: spatial GLS robustness check; exits with empty tables (use vertex_cluster / vertex_nbs)"}, - "vertex_specparam": {"category": "resting", "level": "vertex", "domain": "Spectral", "description": "Vertex-level spectral parameterization", - "about": "The aperiodic (1/f) spectrum fit per vertex across the dorsal source surface with specparam/FOOOF -- exponent, offset, and per-band oscillatory peaks (presence, frequency, power). Group differences in the exponent and offset maps, and in per-band peak power, are tested with a cluster-based permutation test (threshold-clustered vertex t-statistics with cluster-extent FWE correction by permutation); band peak presence is compared with a per-vertex chi-square test. Effect sizes are per-vertex Hedges' g. Read it as: spatially-contiguous clusters of vertices where the groups differ in spectral slope, broadband power, or an oscillatory peak -- a cluster with p_corrected < 0.05 marks a region of difference, the sign of its t-values gives direction. This is the whole-brain, unparcellated version of roi_aperiodic (plus peaks)."}, - "vertex_connectivity": {"category": "resting", "level": "vertex", "domain": "Connectivity", "description": "Vertex pairwise connectivity", - "about": "Whole-brain resting connectivity between source vertices -- the unparcellated counterpart of roi_connectivity. Per band it computes all-to-all vertex coupling with the same kernels (coherence, imaginary coherence, PLI/wPLI/dwPLI, dPLI, AEC, partial correlation) and condenses each vertex's connectivity to a functional connectivity density (FCD) map: the fraction of other vertices it couples to above a threshold (degree/(n-1), Tomasi & Volkow 2010). Group differences in the FCD maps are tested with a cluster-based permutation test -- per-vertex t-statistics are clustered over neighbouring vertices (adjacency by distance) and cluster extents compared to a permutation null for family-wise (FWE) control (Maris & Oostenveld 2007). Effect sizes are per-vertex Hedges' g. Read it as: spatially-contiguous clusters where the groups differ in how densely a region is functionally connected -- a cluster with p_corrected < 0.05 marks a region of difference, the sign of its t-values gives direction. This whole-brain FCD map is the source-side input to the source-vs-sensor comparison (fcd_comparison)."}, - "vertex_cross_freq": {"category": "resting", "level": "vertex", "domain": "Cross-frequency", "description": "Vertex cross-frequency coupling (local PAC, AAC, n:m PPC)", - "about": "Whole-brain cross-frequency coupling on the dorsal source surface -- the unparcellated counterpart of roi_cross_freq, using the same kernels for each valid slow-phase x fast-amplitude band pair. Phase-amplitude coupling (PAC, Tort et al. 2010 modulation index, surrogate z-scored) is computed locally -- the slow phase and fast amplitude come from the same vertex -- yielding a whole-brain coupling map. Amplitude-amplitude coupling (AAC, power-envelope correlation) and n:m phase-phase coupling (PPC, Palva et al. 2005 phase-locking factor) are computed all-to-all across vertices and summarized to a per-vertex node strength (mean off-diagonal coupling). Group differences in these maps are tested with a cluster-based permutation test (per-vertex t-statistics clustered by spatial adjacency, cluster-extent FWE from a permutation null; Maris & Oostenveld 2007), with per-vertex Hedges' g. Read it as: spatially-contiguous clusters where the groups differ in cross-frequency coupling -- a cluster with p_corrected < 0.05 marks a region of difference, the sign of its t-values gives direction. PAC here is the primary source-spatial-advantage measure (local, no leakage between nodes)."}, "electrode_psd": {"category": "resting", "level": "electrode", "domain": "Sensor-level", "description": "Sensor-level PSD analysis", "about": "Resting band power at each scalp electrode -- the sensor-space counterpart of roi_psd. Per channel and band, power is reported as absolute power density in dB/Hz (10*log10 of the band's integrated power divided by its bandwidth, as in roi_psd) and relative power (the band's fraction of total 1-100 Hz power). Groups are compared per channel and band with a linear mixed model (dv ~ group * channel, subject as a random effect), with an optional region-nested model over the configured electrode regions (channels as replicates); per-contrast effects come from the hypothesis layer as Hedges' g with band-wise Benjamini-Hochberg FDR. Read it as: which electrodes/regions differ in band power, in which bands and direction (up = the first-listed group is higher) -- the sensor-level check against the source (ROI) result."}, "electrode_aperiodic": {"category": "resting", "level": "electrode", "domain": "Sensor-level", "description": "Sensor-level aperiodic (1/f) analysis", @@ -170,13 +120,17 @@ def resolve_analysis_name(name: str) -> str: "about": "A source-versus-sensor check on resting band power: for each subject and band, electrode power (averaged over channels) is compared to source power (averaged over ROIs). It reports (1) the cross-subject concordance between sensor and source power (Pearson r per band) and (2) whether the group effect agrees at both levels -- per contrast, Hedges' g with 95% CIs at the electrode level and the source (ROI/region) level, plus an 'exceeds_electrode' flag where a region's effect is larger than the global sensor effect. There is no cluster/FWE correction here; significance is read from whether the 95% CI excludes zero. Read it as: does the source reconstruction recover the same spectral group effect the scalp shows, and does it localize it more sharply than the sensor average?"}, "electrode_connectivity": {"category": "resting", "level": "electrode", "domain": "Sensor-level", "description": "Sensor pairwise connectivity + FCD (source-vs-sensor comparator)", "about": "Resting connectivity between the 30 scalp electrodes -- the sensor-space comparator for vertex_connectivity. Per band it computes all-to-all channel coupling with the leakage/volume-conduction-robust subset of the same kernels (AEC, imaginary coherence, PLI, wPLI, dwPLI, dPLI) and the per-channel functional connectivity density (FCD; degree/(n-1) above threshold, Tomasi & Volkow 2010). Groups are compared per channel with a Welch t-test and Benjamini-Hochberg FDR across the 30 channels (effect sizes Hedges' g), plus a hypothesis-layer cluster-permutation test over the sensor montage (adjacency from channel coordinates). Read it as: which electrodes differ in connectivity/FCD, in which band and direction (up = the first-listed group is higher) -- the scalp-level check on whether the source FCD effect is also visible without source reconstruction. Because sensor space is blurred by volume conduction, the volume-conduction-sensitive metrics are omitted here."}, - "fcd_comparison": {"category": "resting", "level": "electrode", "domain": "Source vs Sensor", "display_name": "Connectivity", "supplements": "electrode_connectivity", "requires": ["electrode_connectivity", "vertex_connectivity"], "description": "Source vs sensor FCD comparison (mean + spatial CV; needs electrode_connectivity AND vertex_connectivity, cross-paradigm)", - "about": "A source-versus-sensor check on functional connectivity density (FCD), pairing the whole-brain vertex FCD maps (vertex_connectivity) against the scalp channel FCD (electrode_connectivity). FCD is each node's fraction of supra-threshold connections (degree/(n-1), threshold 0.05; for dPLI the deviation from its 0.5 no-lag center; Tomasi & Volkow 2010). Per subject and band it summarizes each map two ways -- mean FCD (overall coupling density) and spatial coefficient of variation (CV = SD/mean, how heterogeneous the map is) -- then reports (1) cross-subject concordance between source and sensor summaries (Pearson r per band) and (2) whether the group effect agrees at both levels: per contrast, Hedges' g with 95% CIs at each level plus a sign-concordance flag. There is no cluster/FWE correction here; significance is read from whether a 95% CI excludes zero. Read it as: does source-space recover the same connectivity-density group effect the scalp shows, and is the spatial pattern preserved?"}, "roi_evoked": {"category": "evoked", "level": "roi", "domain": "Evoked", "description": "ITC, ERSP, STP for trial-based paradigms"}, - "vertex_evoked": {"category": "evoked", "level": "vertex", "domain": "Evoked", "description": "Vertex-level ITC, ERSP, STP (cluster-corrected) for trial-based paradigms"}, "electrode_evoked": {"category": "evoked", "level": "electrode", "domain": "Evoked", "description": "Electrode-level ITC, ERSP, STP for trial-based paradigms"}, } +# Plugins add their analyses, metadata, aliases and figure types. +install_plugins(ANALYSIS_REGISTRY, ANALYSIS_METADATA, _DEPRECATED_NAMES) + +# Register aliases so old YAML configs still work +for _old, _new in _DEPRECATED_NAMES.items(): + ANALYSIS_REGISTRY[_old] = ANALYSIS_REGISTRY[_new] + # Add metadata entries for deprecated aliases (point to same metadata) for _old, _new in _DEPRECATED_NAMES.items(): if _new in ANALYSIS_METADATA: @@ -261,7 +215,10 @@ def run_analysis( available = ", ".join( k for k in ANALYSIS_REGISTRY.keys() if k not in _DEPRECATED_NAMES ) - raise ValueError(f"Unknown analysis '{analysis_name}'. Available: {available}") + raise ValueError( + f"Unknown analysis '{analysis_name}'. Available: {available}." + + missing_analysis_hint(analysis_name) + ) # Resolve deprecated name (with warning) and use canonical output dir canonical_name = resolve_analysis_name(analysis_name) diff --git a/src/source_analytics/plugins.py b/src/source_analytics/plugins.py new file mode 100644 index 0000000..c71b84e --- /dev/null +++ b/src/source_analytics/plugins.py @@ -0,0 +1,104 @@ +"""Analysis plugins: installed packages that add analyses to source-analytics. + +A plugin declares an entry point in the ``source_analytics.plugins`` group, e.g. in +its pyproject.toml:: + + [project.entry-points."source_analytics.plugins"] + vertex = "source_analytics_vertex" + +The entry point names an object (usually the plugin's top-level module) that may +define any of: + +``ANALYSES`` + ``dict[str, type[BaseAnalysis]]``: the analyses it adds, by registry name. +``METADATA`` + ``dict[str, dict]``: their ``ANALYSIS_METADATA`` entries (domain, level, ...). +``ALIASES`` + ``dict[str, str]``: deprecated name -> canonical name, so old configs still run. +``register_figures(registry)`` + Called with the ``viz.figure_registry`` module, to add ``TABLE_SCHEMAS`` + entries and ``register()`` figure types for the plugin's analyses. + +``core`` installs every plugin when it is imported. A plugin that fails to import +is logged and skipped, so a broken install cannot stop the core analyses. A plugin +that reuses an existing analysis name raises: that is a packaging bug, and letting +it shadow a built-in would silently change what a config runs. +""" + +from __future__ import annotations + +import logging +from functools import cache +from importlib.metadata import entry_points + +logger = logging.getLogger(__name__) + +ENTRY_POINT_GROUP = "source_analytics.plugins" + +# Analyses that left the core package, and the plugin that now provides them. +# Only used to turn "unknown analysis" into an instruction. +MOVED_TO_PLUGIN: dict[str, str] = { + name: "source-analytics-vertex" + for name in ( + "vertex_cluster", "vertex_connectivity", "vertex_cross_freq", "vertex_directed", + "vertex_evoked", "vertex_graph", "vertex_nbs", "vertex_network", + "vertex_signature", "vertex_spatial", "vertex_specparam", "fcd_comparison", + # and their deprecated aliases + "wholebrain", "spatial_lmm", "specparam_vertex", "mvpa", "vertex_mvpa", + ) +} + + +@cache +def load_plugins() -> tuple[tuple[str, object], ...]: + """Import every installed plugin once; return ``(entry point name, object)`` pairs.""" + loaded = [] + for ep in sorted(entry_points(group=ENTRY_POINT_GROUP), key=lambda e: e.name): + try: + loaded.append((ep.name, ep.load())) + except Exception: + logger.exception( + "source-analytics plugin %r (%s) failed to import; skipping it", + ep.name, ep.value, + ) + return tuple(loaded) + + +def install_plugins( + registry: dict, metadata: dict, aliases: dict, plugins=None, +) -> None: + """Add each plugin's analyses, metadata, aliases and figure types, in place. + + ``plugins`` defaults to :func:`load_plugins`; tests pass their own. + """ + plugins = load_plugins() if plugins is None else plugins + for name, plugin in plugins: + added = dict(getattr(plugin, "ANALYSES", None) or {}) + new_aliases = dict(getattr(plugin, "ALIASES", None) or {}) + clash = sorted((set(added) | set(new_aliases)) & (set(registry) | set(aliases))) + if clash: + raise ValueError( + f"source-analytics plugin {name!r} reuses existing analysis names: " + f"{', '.join(clash)}" + ) + registry.update(added) + metadata.update(getattr(plugin, "METADATA", None) or {}) + aliases.update(new_aliases) + + hooks = [h for _, p in plugins if (h := getattr(p, "register_figures", None))] + if hooks: + from .viz import figure_registry + + for hook in hooks: + hook(figure_registry) + + +def missing_analysis_hint(name: str) -> str: + """A sentence naming the plugin that provides ``name``, or ``""``.""" + plugin = MOVED_TO_PLUGIN.get(name) + if plugin is None: + return "" + return ( + f" '{name}' moved out of source-analytics in v0.8.0 into the {plugin} " + f"plugin; install {plugin} to run it." + ) diff --git a/src/source_analytics/spectral/__init__.py b/src/source_analytics/spectral/__init__.py index 16bc28c..f8a5ad6 100644 --- a/src/source_analytics/spectral/__init__.py +++ b/src/source_analytics/spectral/__init__.py @@ -5,15 +5,7 @@ from .aperiodic import fit_aperiodic from .connectivity import compute_connectivity_matrix from .pac import compute_pac, compute_pac_zscore, compute_pac_multiroi, get_valid_pac_pairs -from .vertex import ( - compute_psd_vertices, - extract_band_power_vertices, - compute_falff, - compute_spectral_slope, - compute_peak_frequency, -) from .epoch_sampler import sample_epochs, get_epoch_config -from .vertex_aperiodic import fit_aperiodic_vertices from .transfer_entropy import compute_transfer_entropy from .vertex_connectivity import ( compute_vertex_connectivity_matrix, @@ -39,14 +31,8 @@ "compute_pac_zscore", "compute_pac_multiroi", "get_valid_pac_pairs", - "compute_psd_vertices", - "extract_band_power_vertices", - "compute_falff", - "compute_spectral_slope", - "compute_peak_frequency", "sample_epochs", "get_epoch_config", - "fit_aperiodic_vertices", "compute_transfer_entropy", "compute_vertex_connectivity_matrix", "compute_vertex_connectivity_matrix_epochs", diff --git a/src/source_analytics/spectral/vertex.py b/src/source_analytics/spectral/vertex.py deleted file mode 100644 index 56fd5da..0000000 --- a/src/source_analytics/spectral/vertex.py +++ /dev/null @@ -1,256 +0,0 @@ -"""Vectorized PSD and feature extraction for vertex-level source data. - -Operates on (n_vertices, n_times) arrays — all vertices processed at once -via scipy.signal.welch with axis=-1 broadcasting. -""" - -from __future__ import annotations - -import numpy as np -from scipy.signal import welch -from scipy.integrate import trapezoid - - -def compute_psd_vertices( - stc_data: np.ndarray, - sfreq: float, - *, - nperseg: int | None = None, - noverlap: int | None = None, - fmin: float = 0.5, - fmax: float | None = None, - window: str = "hann", -) -> tuple[np.ndarray, np.ndarray]: - """Compute PSD for all vertices simultaneously using Welch's method. - - Parameters - ---------- - stc_data : ndarray, shape (n_vertices, n_times) - Source time courses. - sfreq : float - Sampling frequency in Hz. - nperseg : int, optional - Segment length. Default: 2 * sfreq (2-second windows). - noverlap : int, optional - Overlap between segments. Default: nperseg // 2. - fmin, fmax : float - Frequency range to return. - window : str - Window function. - - Returns - ------- - freqs : ndarray, shape (n_freqs,) - Frequency vector in Hz. - psd : ndarray, shape (n_vertices, n_freqs) - Power spectral density per vertex. - """ - n_vertices, n_times = stc_data.shape - - if nperseg is None: - nperseg = int(2 * sfreq) - nperseg = min(nperseg, n_times) - - if noverlap is None: - noverlap = nperseg // 2 - - if fmax is None: - fmax = sfreq / 2.0 - - freqs, psd = welch( - stc_data, - fs=sfreq, - window=window, - nperseg=nperseg, - noverlap=noverlap, - axis=-1, - ) - - # Crop to requested frequency range - mask = (freqs >= fmin) & (freqs <= fmax) - return freqs[mask], psd[:, mask] - - -def extract_band_power_vertices( - freqs: np.ndarray, - psd: np.ndarray, - bands: dict[str, tuple[float, float]], - noise_exclude: tuple[float, float] | None = None, -) -> dict[str, dict[str, np.ndarray]]: - """Extract band power metrics for all vertices at once. - - Parameters - ---------- - freqs : ndarray, shape (n_freqs,) - Frequency vector. - psd : ndarray, shape (n_vertices, n_freqs) - PSD per vertex. - bands : dict - Band name -> (fmin, fmax). - noise_exclude : tuple, optional - Frequency range to exclude from total power computation - (e.g., (55, 65) for line noise). - - Returns - ------- - dict[str, dict[str, ndarray]] - band_name -> {"absolute": (n_vertices,), "relative": (n_vertices,)} - where absolute is the mean band power density, 10*log10(integrated - power / bandwidth), in dB/Hz — the same definition as the ROI path. - """ - # Total power (optionally excluding noise band) - if noise_exclude is not None: - lo, hi = noise_exclude - total_mask = ~((freqs >= lo) & (freqs <= hi)) - total_power = trapezoid(psd[:, total_mask], freqs[total_mask], axis=-1) - else: - total_power = trapezoid(psd, freqs, axis=-1) - - # Substitute a non-zero placeholder only where total is non-positive, to - # avoid division by zero. Using np.maximum(..., eps) here would silently - # floor every legitimately small total power (source-localized PSDs can - # integrate to values well below machine epsilon ~2e-16), corrupting the - # relative metric. See ROI band_power.extract_band_power for the scalar - # equivalent of this guard. - total_power = np.where(total_power > 0, total_power, np.finfo(float).tiny) - - result = {} - for band_name, (fmin, fmax) in bands.items(): - mask = (freqs >= fmin) & (freqs <= fmax) - if not np.any(mask): - n = psd.shape[0] - result[band_name] = { - "absolute": np.full(n, -np.inf), - "relative": np.zeros(n), - } - continue - - band_psd = psd[:, mask] - band_freqs = freqs[mask] - abs_power = trapezoid(band_psd, band_freqs, axis=-1) - - rel_power = abs_power / total_power - - # "absolute" is the mean power DENSITY over the band (integrated power - # / bandwidth) in dB/Hz — the same definition as the ROI/electrode path - # in band_power.py, so the column is comparable across levels. (A - # per-band constant shift relative to the old integrated-power form; - # within-band group statistics are unaffected.) - bandwidth = fmax - fmin - density = abs_power / bandwidth if bandwidth > 0 else abs_power - db_power = np.where(density > 0, 10.0 * np.log10(density), -np.inf) - - result[band_name] = { - "absolute": db_power, - "relative": rel_power, - } - - return result - - -def compute_falff( - freqs: np.ndarray, - psd: np.ndarray, - gamma_range: tuple[float, float] = (65, 100), - total_range: tuple[float, float] = (1, 100), -) -> np.ndarray: - """Compute fractional amplitude of low-frequency fluctuations (fALFF). - - Here defined as high-gamma power / total broadband power, identifying - vertices with disproportionate high-frequency activity. - - Parameters - ---------- - freqs : ndarray, shape (n_freqs,) - psd : ndarray, shape (n_vertices, n_freqs) - gamma_range : tuple - Frequency range for numerator (high gamma). - total_range : tuple - Frequency range for denominator (total). - - Returns - ------- - ndarray, shape (n_vertices,) - fALFF ratio per vertex. - """ - gamma_mask = (freqs >= gamma_range[0]) & (freqs <= gamma_range[1]) - total_mask = (freqs >= total_range[0]) & (freqs <= total_range[1]) - - gamma_power = trapezoid(psd[:, gamma_mask], freqs[gamma_mask], axis=-1) - total_power = trapezoid(psd[:, total_mask], freqs[total_mask], axis=-1) - # See extract_band_power_vertices: substitute only for non-positive totals - # to avoid division by zero without flooring legitimately small denominators. - total_power = np.where(total_power > 0, total_power, np.finfo(float).tiny) - - return gamma_power / total_power - - -def compute_spectral_slope( - freqs: np.ndarray, - psd: np.ndarray, - fit_range: tuple[float, float] = (2, 50), -) -> np.ndarray: - """Compute 1/f spectral slope via log-log linear regression. - - Parameters - ---------- - freqs : ndarray, shape (n_freqs,) - psd : ndarray, shape (n_vertices, n_freqs) - fit_range : tuple - Frequency range for fitting. - - Returns - ------- - ndarray, shape (n_vertices,) - Spectral slope (exponent) per vertex. Negative values indicate - typical 1/f decay; steeper negative = more aperiodic dominance. - """ - mask = (freqs >= fit_range[0]) & (freqs <= fit_range[1]) - log_f = np.log10(freqs[mask]) - # Floor only non-positive PSDs (which shouldn't occur for Welch output but - # are guarded against -inf). Source-localized PSD values commonly fall - # below np.finfo(float).eps (~2e-16); using eps as the floor would - # collapse the log spectrum onto a constant and zero out every slope. - safe_psd = np.where(psd[:, mask] > 0, psd[:, mask], np.finfo(float).tiny) - log_psd = np.log10(safe_psd) - - # Vectorized least-squares: slope = cov(x, y) / var(x) for each vertex - n = log_f.shape[0] - mean_f = log_f.mean() - mean_psd = log_psd.mean(axis=-1, keepdims=True) - - f_centered = log_f - mean_f # (n_freqs,) - psd_centered = log_psd - mean_psd # (n_vertices, n_freqs) - - var_f = np.sum(f_centered ** 2) - cov = np.sum(psd_centered * f_centered, axis=-1) # (n_vertices,) - - slopes = cov / var_f - return slopes - - -def compute_peak_frequency( - freqs: np.ndarray, - psd: np.ndarray, - search_range: tuple[float, float] = (6, 13), -) -> np.ndarray: - """Find peak frequency in a specified range per vertex. - - Parameters - ---------- - freqs : ndarray, shape (n_freqs,) - psd : ndarray, shape (n_vertices, n_freqs) - search_range : tuple - Frequency range to search for peak. - - Returns - ------- - ndarray, shape (n_vertices,) - Peak frequency in Hz per vertex. - """ - mask = (freqs >= search_range[0]) & (freqs <= search_range[1]) - search_freqs = freqs[mask] - search_psd = psd[:, mask] - - peak_idx = np.argmax(search_psd, axis=-1) - return search_freqs[peak_idx] diff --git a/src/source_analytics/spectral/vertex_aperiodic.py b/src/source_analytics/spectral/vertex_aperiodic.py deleted file mode 100644 index 46bc40e..0000000 --- a/src/source_analytics/spectral/vertex_aperiodic.py +++ /dev/null @@ -1,204 +0,0 @@ -"""Vertex-level spectral parameterization (aperiodic + oscillatory peaks). - -Wraps the existing fit_aperiodic() in a vectorized loop over all source -vertices, extracting aperiodic parameters (exponent, offset) and detecting -oscillatory peaks across all configured frequency bands at each spatial -location. -""" - -from __future__ import annotations - -import logging - -import numpy as np - -from .aperiodic import DEFAULT_FREQ_RANGE, band_peak_reachability, fit_aperiodic - -logger = logging.getLogger(__name__) - - -def _safe_band_key(band_name: str) -> str: - """Sanitize band name for use as a dict/column key.""" - return band_name.lower().replace(" ", "_") - - -def fit_aperiodic_vertices( - freqs: np.ndarray, - psd: np.ndarray, - freq_range: tuple[float, float] = DEFAULT_FREQ_RANGE, - max_n_peaks: int = 6, - peak_width_limits: tuple[float, float] = (1.0, 12.0), - bands: dict[str, tuple[float, float]] | None = None, - peak_freq_range: tuple[float, float] | None = None, -) -> dict[str, np.ndarray]: - """Fit aperiodic (1/f) model at each vertex. - - Aperiodic estimation and peak detection want *different* windows and this - function can use two. The aperiodic window (``freq_range``) is deliberately - narrow — borders clear of oscillatory peaks, roll-off and line noise — which - is what makes the exponent unbiased, but it also makes every band outside it - structurally undetectable. ``peak_freq_range`` runs a second, wider fit whose - ONLY job is to locate oscillations, so that the choice of aperiodic window - can be checked against where the peaks actually are (Gerster et al. 2022: - fit borders must not cross oscillatory peaks) rather than merely asserted. - - Peaks from the wide fit are worse-constrained than the narrow fit's aperiodic - parameters — a wide window is a worse 1/f model, and specparam finds peaks by - subtracting that model. They are a diagnostic and an interpretive guard on - band power, not a precision measurement. - - Parameters - ---------- - freqs : ndarray, shape (n_freqs,) - Frequency vector. - psd : ndarray, shape (n_vertices, n_freqs) - PSD per vertex. - freq_range : tuple - Frequency range for the APERIODIC fit (exponent/offset/r^2). - max_n_peaks : int - Maximum number of peaks to detect per vertex. - peak_width_limits : tuple - Min and max peak width in Hz. - bands : dict mapping band name to (fmin, fmax), optional - Frequency bands for peak detection. When *None*, defaults to - ``{"Gamma": (30, 100)}`` for backward compatibility. - peak_freq_range : tuple, optional - Separate, usually wider window for PEAK detection. When *None* or equal - to ``freq_range`` a single fit is performed (no extra cost) and peaks - come from it, preserving the previous behaviour. - - Returns - ------- - dict[str, ndarray] - Always contains: exponent, offset, offset_centered, r_squared, n_peaks, - method, peaks_all (per-vertex list of every detected peak), and the - reachability/window metadata under ``peak_window`` / ``band_reach``. - Per-band keys ``has_{key}_peak``, ``{key}_peak_freq``, - ``{key}_peak_power`` are emitted ONLY for bands the peak window can - reach — an unreachable band gets no column rather than a fabricated - ``False``. All arrays have shape (n_vertices,) except method/peaks_all. - """ - if bands is None: - bands = {"Gamma": (30, 100)} - - peak_range = tuple(peak_freq_range) if peak_freq_range else tuple(freq_range) - two_fit = peak_range != tuple(freq_range) - - reach = band_peak_reachability(bands, peak_range) - detectable = {n: b for n, b in bands.items() if reach[n]["reachable"]} - dropped = [n for n in bands if not reach[n]["reachable"]] - if dropped: - logger.info( - "Peak window %.4g-%.4g Hz cannot reach %s — no has_*_peak columns " - "emitted for these (absence would be structural, not measured).", - peak_range[0], peak_range[1], ", ".join(dropped), - ) - censored = [n for n in detectable if reach[n]["censored"]] - if censored: - logger.warning( - "Peak window %.4g-%.4g Hz only partially covers %s — detection " - "rates are a LOWER bound and peak frequencies are truncated.", - peak_range[0], peak_range[1], ", ".join(censored), - ) - - n_vertices = psd.shape[0] - - exponents = np.zeros(n_vertices) - offsets = np.zeros(n_vertices) - offsets_centered = np.zeros(n_vertices) - r_squareds = np.zeros(n_vertices) - n_peaks_arr = np.zeros(n_vertices, dtype=int) - methods: list[str] = [] - - n_peaks_wide = np.zeros(n_vertices, dtype=int) - - # Per-band peak arrays — DETECTABLE bands only (see band_peak_reachability) - band_keys = {name: _safe_band_key(name) for name in detectable} - band_has_peak = {key: np.zeros(n_vertices, dtype=bool) for key in band_keys.values()} - band_peak_freq = {key: np.full(n_vertices, np.nan) for key in band_keys.values()} - band_peak_power = {key: np.full(n_vertices, np.nan) for key in band_keys.values()} - peaks_all: list[list[dict]] = [[] for _ in range(n_vertices)] - - for vi in range(n_vertices): - try: - result = fit_aperiodic( - freqs, psd[vi], freq_range=freq_range, max_n_peaks=max_n_peaks, - ) - - exponents[vi] = result["exponent"] - offsets[vi] = result["offset"] - offsets_centered[vi] = result["offset_centered"] - r_squareds[vi] = result["r_squared"] - # n_peaks stays the APERIODIC fit's peak count: it is QC on that - # model (many peaks in a narrow window = peaks papering over a bad - # 1/f fit), not the oscillation inventory. That is n_peaks_wide. - n_peaks_arr[vi] = result.get("n_peaks", 0) - methods.append(result.get("method", "unknown")) - - if two_fit: - peak_fit = fit_aperiodic( - freqs, psd[vi], freq_range=peak_range, max_n_peaks=max_n_peaks, - ) - else: - peak_fit = result - peaks = peak_fit.get("peaks", []) - n_peaks_wide[vi] = peak_fit.get("n_peaks", 0) - - # Match detected peaks to frequency bands - for peak in peaks: - cf = peak.get("center_frequency", 0) - pw = peak.get("power", 0) - peaks_all[vi].append({ - "vertex_idx": vi, - "center_frequency": float(cf), - "power": float(pw), - "bandwidth": float(peak.get("bandwidth", np.nan)), - }) - for band_name, (flo, fhi) in detectable.items(): - key = band_keys[band_name] - if flo <= cf <= fhi: - if np.isnan(band_peak_power[key][vi]) or pw > band_peak_power[key][vi]: - band_has_peak[key][vi] = True - band_peak_freq[key][vi] = cf - band_peak_power[key][vi] = pw - - except Exception as e: - logger.debug("Vertex %d fit failed: %s", vi, e) - methods.append("failed") - - n_specparam = sum(1 for m in methods if m == "specparam") - n_linreg = sum(1 for m in methods if m == "linreg") - logger.info("Specparam fit: %d specparam, %d linreg", n_specparam, n_linreg) - if two_fit: - logger.info( - "Two-fit: aperiodic %.4g-%.4g Hz, peaks %.4g-%.4g Hz", - freq_range[0], freq_range[1], peak_range[0], peak_range[1], - ) - for band_name in detectable: - key = band_keys[band_name] - n_det = int(band_has_peak[key].sum()) - logger.info( - " %s peaks detected: %d/%d vertices%s", - band_name, n_det, n_vertices, - " (CENSORED window)" if reach[band_name]["censored"] else "", - ) - - result_dict: dict[str, np.ndarray | list[str]] = { - "exponent": exponents, - "offset": offsets, - "offset_centered": offsets_centered, - "r_squared": r_squareds, - "n_peaks": n_peaks_arr, - "n_peaks_wide": n_peaks_wide, - "method": methods, - "peaks_all": peaks_all, - "peak_window": peak_range, - "aperiodic_window": tuple(freq_range), - "band_reach": reach, - } - for key in band_keys.values(): - result_dict[f"has_{key}_peak"] = band_has_peak[key] - result_dict[f"{key}_peak_freq"] = band_peak_freq[key] - result_dict[f"{key}_peak_power"] = band_peak_power[key] - - return result_dict diff --git a/src/source_analytics/viz/figure_registry.py b/src/source_analytics/viz/figure_registry.py index d0309e5..cfa9586 100644 --- a/src/source_analytics/viz/figure_registry.py +++ b/src/source_analytics/viz/figure_registry.py @@ -90,57 +90,12 @@ def __init__( q_col="q_value", estimate_label="Difference", ), - "vertex_cluster": TableSchema( - posthoc_file="cluster_results.csv", - estimate_col="cluster_stat", - label_col="metric", - band_col="band", - effect_col="peak_t", - p_col="p_corrected", - q_col="p_corrected", - sig_col="p_corrected", - contrast_col="contrast", - estimate_label="Cluster Stat", - ), - "vertex_spatial": TableSchema( - posthoc_file="vertex_spatial_results.csv", - estimate_col="coefficient", - label_col="metric", - band_col="band", - effect_col="t_value", - p_col="p_value", - q_col="q_value", - estimate_label="Coefficient", - ), - "vertex_specparam": TableSchema( - posthoc_file="vertex_specparam_stats.csv", - estimate_col=None, - label_col="parameter", - band_col=None, - effect_col="hedges_g", - p_col="p", - q_col="p", - estimate_label="Hedges g", - ), - "vertex_signature": TableSchema( - posthoc_file="vertex_signature_results.csv", - estimate_col="accuracy", - label_col="band", - band_col="band", - effect_col="accuracy", - p_col="p_value", - q_col="p_value", - estimate_label="Accuracy", - ), } # Backward-compatible aliases for old analysis names for _old, _new in [ ("psd", "roi_psd"), ("aperiodic", "roi_aperiodic"), ("evoked", "roi_evoked"), ("pac", "roi_cross_freq"), ("roi_pac", "roi_cross_freq"), - ("wholebrain", "vertex_cluster"), ("spatial_lmm", "vertex_spatial"), - ("specparam_vertex", "vertex_specparam"), - ("mvpa", "vertex_signature"), ("vertex_mvpa", "vertex_signature"), ]: if _new in TABLE_SCHEMAS: TABLE_SCHEMAS[_old] = TABLE_SCHEMAS[_new] @@ -235,11 +190,9 @@ def _register_all() -> None: # Analyses that support the standard heatmap + volcano heatmap_analyses = [ "roi_psd", "roi_aperiodic", "roi_evoked", "roi_connectivity", "roi_cross_freq", - "vertex_spatial", ] volcano_analyses = [ "roi_psd", "roi_aperiodic", "roi_evoked", "roi_connectivity", "roi_cross_freq", - "vertex_spatial", ] for a in heatmap_analyses: @@ -250,9 +203,5 @@ def _register_all() -> None: # Connectivity gets circos register("roi_connectivity", "circos", sf.plot_summary_circos) - # Vertex-level analyses get glass_brain - for a in ("vertex_cluster", "vertex_spatial", "vertex_specparam"): - register(a, "glass_brain", sf.plot_summary_glass_brain) - _register_all() diff --git a/src/source_analytics/viz/summary_figures.py b/src/source_analytics/viz/summary_figures.py index 976dcd7..1de5fa9 100644 --- a/src/source_analytics/viz/summary_figures.py +++ b/src/source_analytics/viz/summary_figures.py @@ -26,7 +26,6 @@ COLOR_TREND, FIGSIZE_WIDE, FIGSIZE_CIRCOS, - FIGSIZE_GLASS_BRAIN, FONT_TITLE, FONT_SUBTITLE, FONT_ANNOTATION, @@ -424,226 +423,3 @@ def plot_summary_circos( logger.info("Saved circos: %s", out) return outputs - - -# ── Summary glass brain (vertex_cluster / vertex_spatial) ──────────── - -def plot_summary_glass_brain( - tbl_dir: Path, - fig_dir: Path, - **kwargs, -) -> list[Path]: - """Glass-brain showing significant clusters/vertices. - - For vertex_cluster: reads cluster_results.csv + voxelwise_stats.csv - For vertex_spatial: reads vertex_spatial_results.csv - For vertex_specparam: reads vertex_specparam_stats.csv - """ - from .glass_brain import plot_glass_brain - - analysis = kwargs.pop("analysis", None) or _infer_analysis(tbl_dir) - data_dir = kwargs.get("data_dir") - - apply_theme() - - # Load vertex coordinates - coords = _find_coords(tbl_dir, data_dir, analysis) - if coords is None: - logger.warning("Could not find source coordinates for glass brain") - return [] - - outputs = [] - - if analysis in ("vertex_cluster", "wholebrain"): - outputs.extend(_glass_brain_vertex_cluster(tbl_dir, fig_dir, coords, **kwargs)) - elif analysis in ("vertex_spatial", "spatial_lmm"): - outputs.extend(_glass_brain_vertex_spatial(tbl_dir, fig_dir, coords, **kwargs)) - elif analysis in ("vertex_specparam", "specparam_vertex"): - outputs.extend(_glass_brain_vertex_specparam(tbl_dir, fig_dir, coords, **kwargs)) - - return outputs - - -def _find_coords(tbl_dir: Path, data_dir: Path | None, analysis: str) -> np.ndarray | None: - """Search for source_coords.csv in data_dir or nearby directories.""" - search_paths = [] - if data_dir: - search_paths.append(Path(data_dir) / "source_coords.csv") - - # Mirror tbl_dir's position under results/ into the parallel analytics/ tree: - # results/[/]tables// - # -> analytics/[/]//data/ - # Found by naming the `results` ancestor rather than counting levels, so an - # optional profile segment doesn't shift the walk. (The previous fixed 4-level - # walk reached `results` with zero slack and failed *silently* — returning None - # here just drops the glass brain with a warning.) - results_root = next( - (anc for anc in tbl_dir.parents if anc.name == "results"), None, - ) - if results_root is not None: - analytics_root = results_root.parent / "analytics" - # ("tables", , ) or (, "tables", , ) - rel = [p for p in tbl_dir.relative_to(results_root).parts if p != "tables"] - if rel: - search_paths.append( - analytics_root.joinpath(*rel) / "data" / "source_coords.csv" - ) - # vertex_cluster is the canonical producer of source_coords.csv, so fall - # back to it within the same profile+paradigm. - search_paths.append( - analytics_root.joinpath(*rel[:-1]) - / "vertex_cluster" / "data" / "source_coords.csv" - ) - - for p in search_paths: - if p.exists(): - df = pd.read_csv(p) - return df[["x", "y", "z"]].values - - return None - - -def _glass_brain_vertex_cluster( - tbl_dir: Path, fig_dir: Path, coords: np.ndarray, **kwargs, -) -> list[Path]: - """Glass brain from vertex cluster + voxelwise stats.""" - from .glass_brain import plot_glass_brain - - cluster_file = tbl_dir / "cluster_results.csv" - voxel_file = tbl_dir / "voxelwise_stats.csv" - - if not voxel_file.exists(): - logger.warning("voxelwise_stats.csv not found") - return [] - - vox = pd.read_csv(voxel_file) - clusters = pd.read_csv(cluster_file) if cluster_file.exists() else pd.DataFrame() - - contrast_filter = kwargs.get("contrast") - band_filter = kwargs.get("band") - - # Find significant clusters - if not clusters.empty: - sig_clusters = clusters[clusters["p_corrected"] < 0.05] - if contrast_filter: - sig_clusters = sig_clusters[sig_clusters["contrast"] == contrast_filter] - if band_filter: - sig_clusters = sig_clusters[sig_clusters["band"] == band_filter] - else: - sig_clusters = pd.DataFrame() - - if sig_clusters.empty: - logger.info("No significant vertex clusters; plotting top uncorrected t-map") - - # Group by contrast x band x metric and plot the t-map - if contrast_filter: - vox = vox[vox["contrast"] == contrast_filter] - if band_filter: - vox = vox[vox["band"] == band_filter] - - outputs = [] - for (cname, band, metric), grp in vox.groupby(["contrast", "band", "metric"]): - grp = grp.sort_values("vertex_idx") - t_vals = grp["t"].values - if len(t_vals) != len(coords): - # Vertex count mismatch — skip - continue - - title = f"{cname} | {band} | {metric}" - fname = f"glass_brain_{cname}_{band}_{metric}.png" - out = fig_dir / fname - plot_glass_brain(coords, t_vals, title=title, output_path=out, cmap="RdBu_r") - outputs.append(out) - logger.info("Saved glass brain: %s", out) - - return outputs - - -def _glass_brain_vertex_spatial( - tbl_dir: Path, fig_dir: Path, coords: np.ndarray, **kwargs, -) -> list[Path]: - """Glass brain from vertex spatial results (significant bands/metrics only).""" - from .glass_brain import plot_glass_brain - - lmm_file = tbl_dir / "vertex_spatial_results.csv" - if not lmm_file.exists(): - return [] - - df = pd.read_csv(lmm_file) - # Normalise significance - if "significant" in df.columns: - if df["significant"].dtype == object: - df["_sig"] = df["significant"].str.upper().eq("TRUE") - else: - df["_sig"] = df["significant"].astype(bool) - else: - df["_sig"] = False - - contrast_filter = kwargs.get("contrast") - band_filter = kwargs.get("band") - if contrast_filter: - df = df[df["contrast"] == contrast_filter] - if band_filter: - df = df[df["band"] == band_filter] - - sig_df = df[df["_sig"]] - if sig_df.empty: - logger.info("No significant vertex spatial results for glass brain") - return [] - - # For each significant result, try to find per-vertex residuals or use coefficient - # Since spatial_lmm is a single coefficient per band/metric, we show an info plot - outputs = [] - residual_file = tbl_dir / "vertex_spatial_residuals.csv" - if residual_file.exists(): - resid = pd.read_csv(residual_file) - for _, row in sig_df.iterrows(): - cname, band, metric = row["contrast"], row["band"], row["metric"] - sub = resid[(resid.get("contrast", "") == cname) & - (resid.get("band", "") == band) & - (resid.get("metric", "") == metric)] - if sub.empty or "vertex_idx" not in sub.columns: - continue - sub = sub.sort_values("vertex_idx") - vals = sub.iloc[:, -1].values # last column is residual - if len(vals) != len(coords): - continue - title = f"Vertex Spatial | {cname} | {band} | {metric}" - fname = f"glass_brain_slmm_{cname}_{band}_{metric}.png" - out = fig_dir / fname - plot_glass_brain(coords, vals, title=title, output_path=out, cmap="RdBu_r") - outputs.append(out) - logger.info("Saved vertex spatial glass brain: %s", out) - - return outputs - - -def _glass_brain_vertex_specparam( - tbl_dir: Path, fig_dir: Path, coords: np.ndarray, **kwargs, -) -> list[Path]: - """Glass brain from specparam vertex stats (t-values per vertex).""" - from .glass_brain import plot_glass_brain - - stats_file = tbl_dir / "vertex_specparam_stats.csv" - if not stats_file.exists(): - return [] - - df = pd.read_csv(stats_file) - contrast_filter = kwargs.get("contrast") - if contrast_filter: - df = df[df["contrast"] == contrast_filter] - - outputs = [] - for (cname, param), grp in df.groupby(["contrast", "parameter"]): - grp = grp.sort_values("vertex_idx") - t_vals = grp["t"].values - if len(t_vals) != len(coords): - continue - title = f"Specparam | {cname} | {param}" - fname = f"glass_brain_specparam_{cname}_{param}.png" - out = fig_dir / fname - plot_glass_brain(coords, t_vals, title=title, output_path=out, cmap="PiYG") - outputs.append(out) - logger.info("Saved specparam glass brain: %s", out) - - return outputs diff --git a/tests/test_audit_fixes.py b/tests/test_audit_fixes.py index 4abe400..08ed5fe 100644 --- a/tests/test_audit_fixes.py +++ b/tests/test_audit_fixes.py @@ -13,8 +13,6 @@ from source_analytics.config import StudyConfig from source_analytics.core import canonical_analysis_name, ANALYSIS_METADATA -from source_analytics.spectral.band_power import extract_band_power -from source_analytics.spectral.vertex import extract_band_power_vertices from source_analytics.spectral.epoch_sampler import sample_epochs @@ -38,18 +36,6 @@ def test_tfr_raises_clear_error_without_mne(monkeypatch): tfr.tfr_array_morlet(np.zeros((1, 1, 10)), sfreq=100.0, freqs=[5.0], n_cycles=2) -# ---- #9: ROI and vertex `absolute` share one definition (dB/Hz density) ----- -def test_vertex_absolute_matches_roi_density(): - freqs = np.linspace(1, 100, 397) - psd_1d = 1e-12 / freqs # 1/f - bands = {"Alpha": (8, 13), "Beta": (13, 30)} - roi = extract_band_power(freqs, psd_1d, bands) - vtx = extract_band_power_vertices(freqs, psd_1d[np.newaxis, :], bands, noise_exclude=None) - for b in bands: - assert vtx[b]["absolute"][0] == pytest.approx(roi[b]["absolute"], rel=1e-9) - assert vtx[b]["relative"][0] == pytest.approx(roi[b]["relative"], rel=1e-9) - - # ---- #12: n_bootstrap: 0 means "full timeseries" on the vertex sampler too -- def test_sample_epochs_n_bootstrap_zero_returns_full_data(): data = np.random.default_rng(0).standard_normal((3, 5000)) @@ -60,78 +46,17 @@ def test_sample_epochs_n_bootstrap_zero_returns_full_data(): assert sampled.shape == (4, 3, 500) -# ---- #19: vertex modules see global + per-analysis epoch_sampling ---------- -def test_vertex_epoch_config_merges_global_vertex_and_analysis(): - from source_analytics.analyses.vertex_cluster_analysis import VertexClusterAnalysis - - class _Cfg: - raw = { - "epoch_sampling": {"enabled": True, "n_epochs": 40, "n_bootstrap": 5}, - "vertex_cluster": {"epoch_sampling": {"n_bootstrap": 0}}, - } - vertex = {"epoch_sampling": {"n_epochs": 60}} - - a = VertexClusterAnalysis.__new__(VertexClusterAnalysis) - a.config = _Cfg() - merged = a._vertex_epoch_config() - assert merged == {"enabled": True, "n_epochs": 60, "n_bootstrap": 0} - - class _Off: - raw = {"epoch_sampling": {"n_epochs": 40}} - vertex = {} - - a.config = _Off() - assert a._vertex_epoch_config() is None - - # ---- #17: deprecated names resolve to the canonical output dir ------------- def test_canonical_analysis_name(): assert canonical_analysis_name("psd") == "roi_psd" - assert canonical_analysis_name("vertex_mvpa") == "vertex_signature" assert canonical_analysis_name("roi_psd") == "roi_psd" # ---- #4/#5: dependency metadata is honest -------------------------------- def test_comparison_modules_declare_requires(): - assert ANALYSIS_METADATA["fcd_comparison"]["requires"] == [ - "electrode_connectivity", "vertex_connectivity"] - assert ANALYSIS_METADATA["fcd_comparison"]["supplements"] == "electrode_connectivity" assert ANALYSIS_METADATA["electrode_comparison"]["requires"] == ["electrode_psd", "roi_psd"] -# ---- #4: fcd_comparison finds its primaries across paradigm dirs ---------- -def test_fcd_comparison_cross_paradigm_lookup(tmp_path): - from source_analytics.analyses.fcd_comparison_analysis import FCDComparisonAnalysis - - analytics = tmp_path / "analytics" - (analytics / "resting" / "electrode_connectivity" / "data").mkdir(parents=True) - (analytics / "vertex" / "vertex_connectivity" / "data").mkdir(parents=True) - rows = "subject,group,band,metric,fcd\ns1,A,Alpha,pli,0.5\ns2,B,Alpha,pli,0.4\n" - (analytics / "resting" / "electrode_connectivity" / "data" / "electrode_fcd.csv").write_text(rows) - (analytics / "vertex" / "vertex_connectivity" / "data" / "vertex_fcd.csv").write_text(rows) - - class _Cfg: - raw = {} - output_dir = analytics / "resting" - paradigm_name = "resting" - results_dir = tmp_path / "results" - vertex = {} - rois = None - roi_categories = {} - atlas_dir = None - - a = FCDComparisonAnalysis.__new__(FCDComparisonAnalysis) - a.config = _Cfg() - a._sensor_df = a._source_df = None - a._selection = {} - src = a._find_upstream_csv("vertex_connectivity", "vertex_fcd.csv", "source_dir") - assert src == analytics / "vertex" / "vertex_connectivity" / "data" / "vertex_fcd.csv" - sen = a._find_upstream_csv("electrode_connectivity", "electrode_fcd.csv", "sensor_dir") - assert sen.parent.parent.parent.name == "resting" - with pytest.raises(FileNotFoundError, match="run 'vertex_connectivity' first"): - a._find_upstream_csv("vertex_connectivity", "nope.csv", "source_dir") - - # ---- #1: init writes a parseable design/hypotheses/paradigms config -------- def _fake_reconstruction(root: Path, flat: bool): deriv = root / "derivatives" @@ -232,38 +157,6 @@ class _Cfg: _prepare_output(_Cfg, "roi_psd", strict=True, force=False, steps=None) -# ---- #8: vertex_spatial is retired end to end ------------------------------ -def test_vertex_spatial_processes_nothing(tmp_path): - from source_analytics.analyses.vertex_spatial_analysis import VertexSpatialAnalysis - - class _Cfg: - raw = {} - vertex = {} - name = "t" - results_dir = tmp_path / "results" - paradigm_name = None - roi_categories = {} - atlas_dir = None - - a = VertexSpatialAnalysis.__new__(VertexSpatialAnalysis) - a.config = _Cfg() - a.output_dir = tmp_path / "vertex_spatial" - a.output_dir.mkdir() - a._warned = False - a.setup() - a.process_subject(object()) # must not touch a loader - a.statistics() - a.summary() - assert (tmp_path / "results" / "tables" / "vertex_spatial" / "vertex_spatial_results.csv").exists() - assert "RETIRED" in (a.output_dir / "ANALYSIS_SUMMARY.md").read_text() - - -# ---- #13: vertex_evoked is on the hypothesis contract ---------------------- -def test_vertex_evoked_selectable_hypothesis(): - from source_analytics.analyses.vertex_evoked_analysis import VertexEvokedAnalysis - assert "hypothesis" in VertexEvokedAnalysis.SELECTABLE - - # ---- #7: R scripts are discoverable from an installed prefix -------------- def test_find_r_script_dir_env_override(tmp_path, monkeypatch): from source_analytics.analyses.base import find_r_script_dir diff --git a/tests/test_fcd_comparison.py b/tests/test_fcd_comparison.py deleted file mode 100644 index 1c3cf4e..0000000 --- a/tests/test_fcd_comparison.py +++ /dev/null @@ -1,77 +0,0 @@ -"""Tests for the source-vs-sensor FCD comparison module.""" - -import numpy as np -import pandas as pd - -from source_analytics.config import StudyConfig -from source_analytics.core import StudyAnalyzer -from source_analytics.analyses.fcd_comparison_analysis import _fcd_summaries - - -def test_fcd_summaries_mean_and_cv(): - df = pd.DataFrame({ - "subject": ["A"] * 3 + ["B"] * 3, - "group": ["KO_VEH"] * 3 + ["WT_VEH"] * 3, - "band": ["Alpha"] * 6, "metric": ["aec"] * 6, - "fcd": [0.2, 0.4, 0.6, 0.5, 0.5, 0.5], - }) - s = _fcd_summaries(df, "sensor") - a = s[s.subject == "A"].iloc[0] - b = s[s.subject == "B"].iloc[0] - assert abs(a.sensor_mean - 0.4) < 1e-9 - assert abs(a.sensor_cv - 0.5) < 1e-9 # std=0.2, mean=0.4 - assert abs(b.sensor_cv) < 1e-12 # flat map -> zero heterogeneity - - -def test_fcd_summaries_edge_cases(): - # single unit -> CV NaN; all-zero -> mean 0, CV NaN - df = pd.DataFrame({ - "subject": ["A", "B", "B"], "group": ["KO_VEH", "WT_VEH", "WT_VEH"], - "band": ["Alpha"] * 3, "metric": ["aec"] * 3, "fcd": [0.3, 0.0, 0.0], - }) - s = _fcd_summaries(df, "source") - assert np.isnan(s[s.subject == "A"].iloc[0].source_cv) # <2 units - assert np.isnan(s[s.subject == "B"].iloc[0].source_cv) # mean 0 - - -def _write_fcd(path, subjects, n_units, band, metric, base, rng): - rows = [] - for subj, group, mean in subjects: - vals = np.clip(rng.normal(mean, 0.05, n_units), 0, 1) - for u in range(n_units): - rows.append({"subject": subj, "group": group, "band": band, - "metric": metric, "fcd": float(vals[u])}) - path.parent.mkdir(parents=True, exist_ok=True) - pd.DataFrame(rows).to_csv(path, index=False) - - -def test_fcd_comparison_end_to_end(sample_config_yaml): - config = StudyConfig.from_yaml(sample_config_yaml) - base = config.output_dir - rng = np.random.default_rng(0) - - # KO higher FCD than WT at BOTH levels (concordant group effect) - subs = [(f"KO_VEH_{i}", "KO_VEH", 0.6) for i in range(5)] + \ - [(f"WT_VEH_{i}", "WT_VEH", 0.4) for i in range(5)] - _write_fcd(base / "electrode_connectivity" / "data" / "electrode_fcd.csv", - subs, 30, "Alpha", "aec", 0.5, rng) - _write_fcd(base / "vertex_connectivity" / "data" / "vertex_fcd.csv", - subs, 200, "Alpha", "aec", 0.5, rng) - - StudyAnalyzer(config).run_analysis("fcd_comparison") - - summ = pd.read_csv(base / "fcd_comparison" / "data" / "fcd_subject_summary.csv") - assert set(["sensor_mean", "sensor_cv", "source_mean", "source_cv"]).issubset(summ.columns) - assert len(summ) == 10 # one row per subject (Alpha x aec) - - tbl = config.results_dir / "tables" / (config.paradigm_name or "") / \ - "fcd_comparison" / "fcd_comparison_stats.csv" - s = pd.read_csv(tbl) - for col in ("band", "metric", "contrast", "corr_mean_r", - "sensor_mean_g", "source_mean_g", "sensor_cv_g", "source_cv_g", - "mean_concordant", "cv_concordant"): - assert col in s.columns, f"missing {col}" - row = s[s.contrast == "disease_effect"].iloc[0] - # KO>WT at both levels -> mean effect same sign -> concordant - assert row["mean_concordant"] - assert np.sign(row["sensor_mean_g"]) == np.sign(row["source_mean_g"]) diff --git a/tests/test_figure_regen.py b/tests/test_figure_regen.py index c45f290..e0d5ebd 100644 --- a/tests/test_figure_regen.py +++ b/tests/test_figure_regen.py @@ -1,59 +1,13 @@ """Guard the standard: figures() must be regenerable from PERSISTED data via `--steps figures` alone — never dependent on in-memory state carried within one process. These are structural guards (cheap, catch regressions) complementing the -end-to-end reload check in test_fcd_comparison.py. +end-to-end reload checks; the vertex modules' guards moved to the +source-analytics-vertex plugin. """ import inspect from source_analytics.analyses.base import BaseAnalysis -from source_analytics.analyses import ( - vertex_connectivity_analysis as vc, - vertex_directed_analysis as vd, - vertex_cross_freq_analysis as vcf, - vertex_evoked_analysis as ve, - vertex_network_analysis as vn, - fcd_comparison_analysis as fc, -) - -# Map/cluster modules: figures() renders per-vertex glass brains from cluster -# results that statistics() must persist and figures() must reload. -MAP_MODULES = [ - (vc.VertexConnectivityAnalysis, "vertex_connectivity"), - (vd.VertexDirectedAnalysis, "vertex_directed"), - (vcf.VertexCrossFreqAnalysis, "vertex_cross_freq"), - (ve.VertexEvokedAnalysis, "vertex_evoked"), -] - - def test_base_has_cluster_state_helpers(): assert callable(getattr(BaseAnalysis, "_save_cluster_state", None)) assert callable(getattr(BaseAnalysis, "_load_cluster_state", None)) - - -def test_map_modules_regenerable_from_disk(): - for cls, name in MAP_MODULES: - fig_src = inspect.getsource(cls.figures) - stat_src = inspect.getsource(cls.statistics) - # figures() reloads persisted cluster state when in-memory is empty - assert "_load_cluster_state" in fig_src, \ - f"{name}.figures() must reload persisted state (not use in-memory)" - # statistics() persists that state AND can reload its inputs from disk - assert "_save_cluster_state" in stat_src, \ - f"{name}.statistics() must persist cluster state for figures()" - assert "_reload_maps_from_disk" in stat_src, \ - f"{name}.statistics() must reload per-subject maps from disk" - - -def test_fcd_comparison_figures_reload_from_csv(): - fig_src = inspect.getsource(fc.FCDComparisonAnalysis.figures) - assert "fcd_subject_summary.csv" in fig_src, \ - "fcd_comparison.figures() must reload its per-subject summary CSV" - - -def test_vertex_graph_figures_read_persisted_table(): - # vertex_graph has no in-memory state to rely on; its figures() must render - # from the persisted stats table (not be a `pass` stub). - fig_src = inspect.getsource(vn.VertexGraphAnalysis.figures) - assert "vertex_graph_stats.csv" in fig_src, \ - "vertex_graph.figures() must render from its persisted stats table" diff --git a/tests/test_network_split.py b/tests/test_network_split.py index c16cc7a..8e3129b 100644 --- a/tests/test_network_split.py +++ b/tests/test_network_split.py @@ -11,14 +11,10 @@ ROIGraphAnalysis, ROINBSAnalysis, ) -from source_analytics.analyses.vertex_network_analysis import ( - VertexGraphAnalysis, -) def test_split_analyses_registered(): - for name in ("roi_graph", "roi_nbs", "vertex_graph", "vertex_nbs", - "roi_network", "vertex_network"): + for name in ("roi_graph", "roi_nbs", "roi_network"): assert name in ANALYSIS_REGISTRY @@ -26,9 +22,7 @@ def test_split_metadata_domain_and_supplements(): meta = analysis_meta() assert meta["roi_graph"]["supplements"] == "roi_connectivity" assert meta["roi_nbs"]["supplements"] == "roi_connectivity" - assert meta["vertex_graph"]["supplements"] == "vertex_connectivity" - assert meta["vertex_nbs"]["supplements"] == "vertex_connectivity" - for n in ("roi_graph", "roi_nbs", "vertex_graph", "vertex_nbs"): + for n in ("roi_graph", "roi_nbs"): assert meta[n]["domain"] == "Connectivity" @@ -47,12 +41,6 @@ def _config(tmp_path): roi_network: connectivity_metrics: [imag_coherence, dwpli, pli] nbs_threshold: 2.5 - vertex: - data_dir: ./d - analyses: - vertex_network: - connectivity_metrics: [imag_coherence, aec] - nbs_threshold: 3.0 """ p = tmp_path / "s.yaml" p.write_text(text) @@ -69,7 +57,3 @@ def test_split_inherits_config_via_fallback(tmp_path): n = ROINBSAnalysis(cfg.for_paradigm_analysis("resting", "roi_nbs"), tmp_path / "on") assert n._connectivity_metrics == ["imag_coherence", "dwpli", "pli"] assert n._nbs_results_filename == "roi_nbs_results.csv" - - vg = VertexGraphAnalysis(cfg.for_paradigm_analysis("vertex", "vertex_graph"), tmp_path / "ovg") - assert vg._connectivity_metrics == ["imag_coherence", "aec"] - assert vg._nbs_threshold == 3.0 # vertex default, from the vertex_network block diff --git a/tests/test_plugins.py b/tests/test_plugins.py new file mode 100644 index 0000000..2e36b08 --- /dev/null +++ b/tests/test_plugins.py @@ -0,0 +1,80 @@ +"""The plugin hook: installed packages add analyses through an entry point.""" + +from __future__ import annotations + +import logging +import types + +import pytest + +import source_analytics.analyses +from source_analytics import plugins +from source_analytics.core import ANALYSIS_REGISTRY, StudyAnalyzer +from source_analytics.viz import figure_registry + + +class _Toy: + """Stand-in analysis class; install_plugins only files it.""" + + +def _plugin(**attrs): + return types.SimpleNamespace(**attrs) + + +def test_install_adds_analyses_metadata_and_aliases(): + registry, metadata, aliases = {"roi_psd": object}, {}, {"psd": "roi_psd"} + toy = _plugin(ANALYSES={"toy": _Toy}, METADATA={"toy": {"domain": "Toy"}}, + ALIASES={"old_toy": "toy"}) + plugins.install_plugins(registry, metadata, aliases, plugins=[("toy", toy)]) + assert registry["toy"] is _Toy + assert metadata["toy"] == {"domain": "Toy"} + assert aliases == {"psd": "roi_psd", "old_toy": "toy"} + + +@pytest.mark.parametrize("attrs", [ + {"ANALYSES": {"roi_psd": _Toy}}, # an analysis name core already has + {"ALIASES": {"psd": "toy"}}, # a deprecated alias core already has +]) +def test_a_plugin_cannot_reuse_a_name(attrs): + registry, aliases = {"roi_psd": object}, {"psd": "roi_psd"} + with pytest.raises(ValueError, match="reuses existing analysis names"): + plugins.install_plugins(registry, {}, aliases, plugins=[("bad", _plugin(**attrs))]) + + +def test_register_figures_receives_the_figure_registry(): + seen = [] + plugins.install_plugins({}, {}, {}, plugins=[("toy", _plugin(register_figures=seen.append))]) + assert seen == [figure_registry] + + +def test_a_plugin_that_fails_to_import_is_logged_and_skipped(monkeypatch, caplog): + class _BrokenEntryPoint: + name, value = "broken", "not_a_real_module" + + def load(self): + raise ImportError("no module named not_a_real_module") + + monkeypatch.setattr(plugins, "entry_points", lambda group: [_BrokenEntryPoint()]) + plugins.load_plugins.cache_clear() + try: + with caplog.at_level(logging.ERROR, logger="source_analytics.plugins"): + assert plugins.load_plugins() == () + assert "'broken'" in caplog.text and "failed to import" in caplog.text + finally: + monkeypatch.undo() + plugins.load_plugins.cache_clear() + + +def test_the_vertex_analyses_left_core(): + assert not hasattr(source_analytics.analyses, "VertexClusterAnalysis") + assert "source-analytics-vertex" in plugins.missing_analysis_hint("vertex_cluster") + assert "source-analytics-vertex" in plugins.missing_analysis_hint("wholebrain") + assert plugins.missing_analysis_hint("roi_psd") == "" + + +def test_unknown_analysis_error_names_the_plugin(): + if "vertex_cluster" in ANALYSIS_REGISTRY: + pytest.skip("source-analytics-vertex is installed in this environment") + analyzer = StudyAnalyzer(config=None, subjects=[object()]) + with pytest.raises(ValueError, match="install source-analytics-vertex"): + analyzer.run_analysis("vertex_cluster") diff --git a/tests/test_select.py b/tests/test_select.py index 2f887cc..d24c769 100644 --- a/tests/test_select.py +++ b/tests/test_select.py @@ -142,5 +142,5 @@ def test_parse_dim_not_selectable_for_target_analysis_exits(): def test_parse_valid_dim_for_target_analysis(): - sel = _parse_selection(_args(metric="pli", analysis="vertex_connectivity")) + sel = _parse_selection(_args(metric="pli", analysis="roi_connectivity")) assert sel == {"metric": frozenset({"pli"})} diff --git a/tests/test_spectral.py b/tests/test_spectral.py index 93d95fd..e5b1226 100644 --- a/tests/test_spectral.py +++ b/tests/test_spectral.py @@ -5,11 +5,6 @@ from source_analytics.spectral.psd import compute_psd, compute_psd_multiroi from source_analytics.spectral.band_power import extract_band_power -from source_analytics.spectral.vertex import ( - compute_falff, - compute_spectral_slope, - extract_band_power_vertices, -) from source_analytics.stats.cluster_permutation import hedges_g @@ -54,87 +49,6 @@ def test_extract_band_power(): assert result["Gamma"]["relative"] > result["Alpha"]["relative"] -def test_extract_band_power_vertices_small_scale(): - """Regression test: relative power must not collapse to ~0 for PSDs whose - integrated total falls below np.finfo(float).eps (~2.2e-16). - - Source-localized PSDs commonly integrate to ~1e-18; an over-aggressive - np.maximum(total, eps) clamp on the denominator silently rescaled every - band's relative metric by ~100x and inverted some inter-group directions. - See FORGE manuscript 2 RERUN_PROPOSAL.md, May 2026. - """ - freqs = np.linspace(0.5, 110, 220) - # Flat spectrum at amplitude well below eps — mimics source-localized scale. - psd = np.full((4, len(freqs)), 1e-20) # 4 "vertices" - - bands = {"Delta": (1, 4), "Alpha": (8, 13), "Gamma": (30, 55)} - result = extract_band_power_vertices(freqs, psd, bands, noise_exclude=None) - - # For a flat PSD of amplitude A, band power = A * band_width, total = A * total_width. - # Relative = band_width / total_width — independent of A. - total_width = freqs[-1] - freqs[0] - for band_name, (fmin, fmax) in bands.items(): - rel = result[band_name]["relative"] - expected = (fmax - fmin) / total_width - assert np.all( - np.abs(rel - expected) < 0.01 - ), f"{band_name} relative {rel.mean():.4f} != expected {expected:.4f} — eps-clamp bug regression" - - -def test_extract_band_power_vertices_sums_to_one(): - """A spectrum entirely covered by named bands should produce relative - powers summing to 1.0 (allowing small numerical slack).""" - freqs = np.linspace(1, 100, 200) - psd = np.ones((3, len(freqs))) * 1e-18 # below-eps amplitude - - # Bands exactly covering 1-100 Hz without gaps - bands = {"A": (1, 25), "B": (25, 50), "C": (50, 75), "D": (75, 100)} - result = extract_band_power_vertices(freqs, psd, bands, noise_exclude=None) - total_rel = sum(result[b]["relative"] for b in bands) - # Tolerance accounts for trapezoid integration at sub-bin boundaries; the - # regression test catches the eps-clamp collapse (~0.005), which is two - # orders of magnitude away from this assertion. - assert np.all(total_rel > 0.95), ( - f"Relative powers across gap-free bands should sum to ~1.0, got {total_rel}" - ) - - -def test_compute_falff_small_scale(): - """fALFF must not collapse to ~0 when total integrated power is below eps.""" - freqs = np.linspace(1, 100, 200) - psd = np.ones((3, len(freqs))) * 1e-18 - - # Flat spectrum: gamma (65-100) / total (1-100) = 35 / 99 ≈ 0.354 - falff = compute_falff(freqs, psd, gamma_range=(65, 100), total_range=(1, 100)) - expected = 35 / 99 - assert np.all(np.abs(falff - expected) < 0.01), ( - f"fALFF should be ~{expected:.3f} for flat spectrum, got {falff}" - ) - - -def test_compute_spectral_slope_small_scale(): - """Spectral slope must recover the true 1/f^alpha exponent even when PSD - values are below np.finfo(float).eps. - - Pre-fix, np.maximum(psd, eps) floored every value in the log10 spectrum - to log10(eps) ≈ -15.66, collapsing slope to ~0. Source-localized PSDs - typically integrate to ~1e-18 with per-bin values 1e-19 to 1e-22 — well - below eps. - """ - freqs = np.logspace(0, 2, 200) # 1 to 100 Hz, log-spaced - # Construct PSD = scale * f^(-1.5), small absolute scale - scale = 1e-22 - true_alpha = 1.5 - psd_1d = scale * freqs ** (-true_alpha) - psd = np.tile(psd_1d, (5, 1)) # 5 vertices, identical - - slope = compute_spectral_slope(freqs, psd, fit_range=(2, 50)) - # Slope should be -true_alpha. Without fix it would be ~0. - assert np.all(np.abs(slope - (-true_alpha)) < 0.1), ( - f"Slope should recover {-true_alpha:.2f}; got {slope.mean():.3f}" - ) - - def test_hedges_g_small_scale(): """Hedges' g must not collapse when pooled SD is in physical units smaller than eps. Without the fix, np.maximum(pooled_std, eps) would deflate g to @@ -271,14 +185,6 @@ def test_fit_aperiodic_carries_window_provenance(): # --- Two-fit peak detection / fit-window justification ----------------------- -def _synthetic_psd(freqs, peaks=((6.0, 0.9, 1.5), (22.0, 0.5, 3.0)), exponent=1.0): - """1/f spectrum with Gaussian peaks; one below 12 Hz, one inside 12-45.""" - psd = 10 ** (1.2 - exponent * np.log10(freqs)) - for cf, pw, sd in peaks: - psd += 10 ** pw * np.exp(-((freqs - cf) ** 2) / (2 * sd ** 2)) * 1e-1 - return psd - - BANDS = { "Delta": (1, 4), "Theta": (4, 10), "Alpha": (10, 13), "Beta": (13, 30), "Low Gamma": (30, 55), "High Gamma": (65, 80), "Epsilon": (80, 150), @@ -299,71 +205,3 @@ def test_band_reachability_marks_unreachable_and_censored(): # Partial overlap -> reachable but censored (rates are a lower bound) assert reach["Low Gamma"]["reachable"] and reach["Low Gamma"]["censored"] assert 0.0 < reach["Low Gamma"]["frac_visible"] < 1.0 - - -def test_unreachable_bands_emit_no_peak_columns(): - """A band the window cannot see must be ABSENT, never a False. - - Emitting has_delta_peak=False for a window that starts at 12 Hz fabricates - a measured null: downstream chi-squared tests then report p=1.0 at every - vertex for a comparison the data never had power to make. - """ - from source_analytics.spectral.vertex_aperiodic import fit_aperiodic_vertices - - freqs = np.arange(1, 101, 0.5) - psd = np.array([_synthetic_psd(freqs) for _ in range(3)]) - - out = fit_aperiodic_vertices(freqs, psd, freq_range=(12, 45), bands=BANDS) - - for band in ("delta", "theta", "high_gamma", "epsilon"): - assert f"has_{band}_peak" not in out - assert "has_beta_peak" in out - - -def test_two_fit_recovers_peaks_the_aperiodic_window_cannot_see(): - """The wide peak fit finds the 6 Hz peak; the narrow fit still sets exponent.""" - from source_analytics.spectral.vertex_aperiodic import fit_aperiodic_vertices - - freqs = np.arange(1, 101, 0.5) - psd = np.array([_synthetic_psd(freqs) for _ in range(3)]) - - narrow = fit_aperiodic_vertices(freqs, psd, freq_range=(12, 45), bands=BANDS) - two = fit_aperiodic_vertices( - freqs, psd, freq_range=(12, 45), bands=BANDS, peak_freq_range=(2, 50), - ) - - # Theta is now measurable, and the 6 Hz peak is actually found - assert "has_theta_peak" not in narrow - assert "has_theta_peak" in two - assert two["has_theta_peak"].all() - assert np.allclose(two["theta_peak_freq"], 6.0, atol=1.0) - - # Aperiodic estimates still come from the NARROW window, unchanged - assert np.allclose(two["exponent"], narrow["exponent"]) - assert two["aperiodic_window"] == (12, 45) - assert two["peak_window"] == (2, 50) - # n_peaks stays the narrow fit's QC count; the inventory is n_peaks_wide - assert (two["n_peaks_wide"] > two["n_peaks"]).all() - - -def test_peak_inventory_supports_border_crossing_check(): - """peaks_all carries the bandwidth needed to test Gerster's border rule.""" - from source_analytics.spectral.vertex_aperiodic import fit_aperiodic_vertices - - freqs = np.arange(1, 101, 0.5) - psd = np.array([_synthetic_psd(freqs) for _ in range(2)]) - - out = fit_aperiodic_vertices( - freqs, psd, freq_range=(12, 45), bands=BANDS, peak_freq_range=(2, 50), - ) - - peaks = out["peaks_all"][0] - assert len(peaks) >= 2 - for pk in peaks: - assert {"center_frequency", "power", "bandwidth"} <= set(pk) - assert np.isfinite(pk["bandwidth"]) - # The 6 Hz peak sits well clear of the 12 Hz border on this synthetic data - cf = np.array([p["center_frequency"] for p in peaks]) - bw = np.array([p["bandwidth"] for p in peaks]) - crossing = ((cf - bw / 2 < 12) & (cf + bw / 2 > 12)) - assert not crossing.any() diff --git a/uv.lock b/uv.lock index 4942072..8726330 100644 --- a/uv.lock +++ b/uv.lock @@ -2038,7 +2038,7 @@ wheels = [ [[package]] name = "source-analytics" -version = "0.7.1" +version = "0.8.0" source = { editable = "." } dependencies = [ { name = "joblib" },