"""
Source reconstruction for pyAVS package.
This module provides functions for MEG source reconstruction including
beamforming, minimum norm estimation, and population code analysis.
"""
import os
import numpy as np
import pandas as pd
import mne
import h5py
import json
import hashlib
from typing import List, Optional, Tuple, Dict, Any, Union
from sklearn.preprocessing import StandardScaler
from ..layout import get_layout
from ..utils.validation import validate_subject_id, validate_session
from ..utils.paths import get_default_subjects_dir
from ..utils.derivatives import get_derivatives_manager
from ..utils.logging import get_logger
from .forward import load_forward_model
from ..io.write import save_source_data, save_population_codes_h5
from ..io.read import load_source_data, find_population_codes_files, list_available_parameter_sets
logger = get_logger('source.reconstruction')
[docs]
def setup_source_reconstruction(subject_id: int, session: int,
method: str = 'beamformer',
data_path: Optional[str] = None,
**method_kwargs) -> Dict[str, Any]:
"""
Set up source reconstruction for a subject/session.
Parameters
----------
subject_id : int
Subject ID
session : int
Session number
method : str, optional
Source reconstruction method ('beamformer', 'mne', 'lcmv') (default: 'beamformer')
data_path : str, optional
Path to data directory
**method_kwargs
Additional method-specific parameters
Returns
-------
dict
Source reconstruction setup parameters
"""
validate_subject_id(subject_id)
validate_session(session)
# Load forward model
fwd = load_forward_model(subject_id, session, data_path)
setup = {
'subject_id': subject_id,
'session': session,
'method': method,
'forward_model': fwd,
'method_kwargs': method_kwargs
}
return setup
[docs]
def compute_minimum_norm_estimate(epochs: mne.Epochs,
forward: mne.Forward,
noise_cov: Optional[mne.Covariance] = None,
lambda2: float = 1.0/9.0,
method: str = 'dSPM',
pick_ori: Optional[str] = None,
verbose: bool = True) -> List[mne.SourceEstimate]:
"""
Compute minimum norm estimate for epoched data.
Parameters
----------
epochs : mne.Epochs
Epoched MEG data
forward : mne.Forward
Forward solution
noise_cov : mne.Covariance, optional
Noise covariance matrix (default: None)
lambda2 : float, optional
Regularization parameter (default: 1.0/9.0)
method : str, optional
Inverse method ('MNE', 'dSPM', 'sLORETA') (default: 'dSPM')
pick_ori : str, optional
Orientation selection (default: None)
verbose : bool, optional
Whether to print progress information (default: True)
Returns
-------
list of mne.SourceEstimate
Source estimates for each epoch
"""
if verbose:
logger.info("Computing minimum norm estimate...")
# Compute noise covariance if not provided
if noise_cov is None:
if verbose:
logger.info("Computing noise covariance matrix...")
noise_cov = mne.compute_covariance(
epochs, tmin=None, tmax=0.0, method='empirical', verbose=verbose
)
# Compute inverse operator
if verbose:
logger.info("Computing inverse operator...")
inverse_operator = mne.minimum_norm.make_inverse_operator(
epochs.info, forward, noise_cov, loose=0.2, depth=0.8, verbose=verbose
)
# Apply inverse operator to each epoch
if verbose:
logger.info("Applying inverse operator...")
stc_epochs = []
for i, epoch in enumerate(epochs):
if verbose and i % 50 == 0:
logger.info(f"Processing epoch {i}/{len(epochs)}")
stc = mne.minimum_norm.apply_inverse(
epoch, inverse_operator, lambda2=lambda2, method=method,
pick_ori=pick_ori, verbose=False
)
stc_epochs.append(stc)
if verbose:
logger.info(f"Computed {len(stc_epochs)} source estimates")
return stc_epochs
[docs]
def compute_source_power(source_data: np.ndarray,
method: str = 'mean',
time_window: Optional[Tuple[float, float]] = None,
baseline: Optional[Tuple[float, float]] = None,
times: Optional[np.ndarray] = None,
verbose: bool = True) -> np.ndarray:
"""
Compute source power from source space data.
Parameters
----------
source_data : np.ndarray
Source space data with shape (n_epochs, n_sources, n_times)
method : str, optional
Power computation method ('mean', 'peak', 'rms') (default: 'mean')
time_window : tuple of float, optional
Time window for power computation (default: None, uses all times)
baseline : tuple of float, optional
Baseline time window for normalization (default: None)
times : np.ndarray, optional
Time points in seconds (default: None)
verbose : bool, optional
Whether to print progress information (default: True)
Returns
-------
np.ndarray
Source power with shape (n_epochs, n_sources)
"""
if verbose:
logger.info("Computing source power...")
n_epochs, n_sources, n_times = source_data.shape
# Select time window
if time_window is not None and times is not None:
time_mask = (times >= time_window[0]) & (times <= time_window[1])
data_windowed = source_data[:, :, time_mask]
else:
data_windowed = source_data
# Compute power
if method == 'mean':
power = np.mean(np.abs(data_windowed), axis=2)
elif method == 'peak':
power = np.max(np.abs(data_windowed), axis=2)
elif method == 'rms':
power = np.sqrt(np.mean(data_windowed**2, axis=2))
else:
raise ValueError(f"Unknown power method: {method}")
# Apply baseline correction if requested
if baseline is not None and times is not None:
baseline_mask = (times >= baseline[0]) & (times <= baseline[1])
baseline_data = source_data[:, :, baseline_mask]
if method == 'mean':
baseline_power = np.mean(np.abs(baseline_data), axis=2)
elif method == 'peak':
baseline_power = np.max(np.abs(baseline_data), axis=2)
elif method == 'rms':
baseline_power = np.sqrt(np.mean(baseline_data**2, axis=2))
# Relative change
power = (power - baseline_power) / baseline_power
if verbose:
logger.info(f"Computed power for {n_epochs} epochs, {n_sources} sources")
return power
[docs]
def compute_population_codes(source_data: np.ndarray,
events_metadata: pd.DataFrame,
conditions: List[str],
time_window: Tuple[float, float],
times: np.ndarray,
baseline: Optional[Tuple[float, float]] = None,
normalize: bool = False,
verbose: bool = True) -> Dict[str, np.ndarray]:
"""
Compute population codes for different experimental conditions.
Parameters
----------
source_data : np.ndarray
Source space data with shape (n_epochs, n_sources, n_times)
events_metadata : pd.DataFrame
Metadata for each epoch
conditions : list of str
Column names in metadata defining conditions
time_window : tuple of float
Time window for population code computation
times : np.ndarray
Time points in seconds
baseline : tuple of float, optional
Baseline time window (default: None)
normalize : bool, optional
Whether to normalize population codes (default: True)
verbose : bool, optional
Whether to print progress information (default: True)
Returns
-------
dict
Dictionary mapping conditions to population codes
"""
if verbose:
logger.info("Computing population codes...")
# Select time window
time_mask = (times >= time_window[0]) & (times <= time_window[1])
windowed_data = source_data[:, :, time_mask]
# Average over time window
epoch_data = np.mean(windowed_data, axis=2) # Shape: (n_epochs, n_sources)
# Apply baseline correction if requested
if baseline is not None:
baseline_mask = (times >= baseline[0]) & (times <= baseline[1])
baseline_data = np.mean(source_data[:, :, baseline_mask], axis=2)
epoch_data = epoch_data - baseline_data
# Compute population codes for each condition
population_codes = {}
for condition in conditions:
if condition not in events_metadata.columns:
if verbose:
logger.warning(f"Condition {condition} not found in metadata")
continue
# Get unique values for this condition
unique_values = events_metadata[condition].dropna().unique()
condition_codes = {}
for value in unique_values:
# Select epochs for this condition value
condition_mask = events_metadata[condition] == value
condition_data = epoch_data[condition_mask]
if len(condition_data) == 0:
continue
# Average across epochs
population_code = np.mean(condition_data, axis=0)
# Normalize if requested
if normalize:
scaler = StandardScaler()
population_code = scaler.fit_transform(population_code.reshape(-1, 1)).flatten()
condition_codes[str(value)] = population_code
population_codes[condition] = condition_codes
if verbose:
logger.info(f" {condition}: {len(condition_codes)} conditions")
return population_codes
[docs]
def apply_source_reconstruction(epochs: mne.Epochs,
forward: mne.Forward,
method: str = 'beamformer',
**method_kwargs) -> np.ndarray:
"""
Apply source reconstruction to epoched data.
Parameters
----------
epochs : mne.Epochs
Epoched MEG data
forward : mne.Forward
Forward solution
method : str, optional
Source reconstruction method (default: 'beamformer')
**method_kwargs
Method-specific parameters
Returns
-------
np.ndarray
Source space data
"""
if method == 'beamformer':
filters = compute_beamformer_filters(epochs, forward, **method_kwargs)
source_data = apply_beamformer(epochs, filters)
elif method == 'mne':
stc_epochs = compute_minimum_norm_estimate(epochs, forward, **method_kwargs)
source_data = np.array([stc.data for stc in stc_epochs])
else:
raise ValueError(f"Unknown source reconstruction method: {method}")
return source_data
def _generate_parameter_signature(**params) -> str:
"""
Generate a unique signature string based on processing parameters.
This creates a hash-based identifier that uniquely identifies a set of
processing parameters, allowing for intelligent organization of results.
Parameters
----------
**params
Processing parameters
Returns
-------
str
Unique parameter signature string
"""
# Clean and standardize parameters
clean_params = {}
for key, value in params.items():
if value is None:
continue
elif isinstance(value, (list, tuple, np.ndarray)):
# Convert sequences to sorted tuples for consistent hashing
if isinstance(value, np.ndarray):
value = value.tolist()
if isinstance(value, list) and len(value) > 0 and isinstance(value[0], str):
value = sorted(value) # Sort string lists
clean_params[key] = tuple(value)
elif isinstance(value, dict):
# Convert dicts to sorted tuple of items
clean_params[key] = tuple(sorted(value.items()))
else:
clean_params[key] = value
# Create deterministic string representation
param_string = json.dumps(clean_params, sort_keys=True, separators=(',', ':'))
# Generate hash
param_hash = hashlib.sha256(param_string.encode()).hexdigest()
# Create readable signature: event_type_hash
event_type = clean_params.get('event_type', 'unknown')
signature = f"{event_type}_{param_hash[:16]}"
return signature
def _save_parameter_metadata(metadata_file: str, metadata: Dict[str, Any]) -> None:
"""
Save parameter metadata to JSON file.
Parameters
----------
metadata_file : str
Path to metadata file
metadata : dict
Metadata dictionary to save
"""
# Load existing metadata if file exists
if os.path.exists(metadata_file):
try:
with open(metadata_file, 'r') as f:
existing_metadata = json.load(f)
except (json.JSONDecodeError, FileNotFoundError):
existing_metadata = {}
else:
existing_metadata = {}
# Update with new metadata
existing_metadata.update(metadata)
existing_metadata['last_updated'] = pd.Timestamp.now().isoformat()
# Save updated metadata
with open(metadata_file, 'w') as f:
json.dump(existing_metadata, f, indent=2, default=str)
[docs]
def compute_empty_room_covariance(data_path: str,
subject_id: int,
sessions: List[int],
verbose: bool = True) -> Tuple[mne.Covariance, str]:
"""
Compute noise covariance from empty room recordings.
Parameters
----------
data_path : str
Path to data directory
subject_id : int
Subject ID
sessions : list of int
Session numbers to process
verbose : bool, optional
Whether to print progress information (default: True)
Returns
-------
mne.Covariance
Computed noise covariance matrix
str
Path to saved covariance file
"""
validate_subject_id(subject_id)
if verbose:
logger.info(f"Computing empty room covariance for subject {subject_id}")
# Find Maxwell-filtered empty room files in the derivatives tree
layout = get_layout(data_path)
empty_room_files = []
empty_room_recording_names = ['d', 'b'] # 'danach'/'bevor' — after/before the session
for session in sessions:
validate_session(session)
for recording in empty_room_recording_names:
empty_room_path = layout.meg_sss_empty_room(subject_id, session, recording)
if empty_room_path.exists():
empty_room_files.append(str(empty_room_path))
if verbose:
logger.info(f"Found empty room file: {empty_room_path.name}")
elif verbose:
logger.warning(f"Empty room file not found: {empty_room_path}")
if not empty_room_files:
raise FileNotFoundError(f"No empty room files found for subject {subject_id}")
if verbose:
logger.info(f"Found {len(empty_room_files)} empty room files")
# Load and concatenate empty room data
raw_list = []
for file_path in empty_room_files:
if verbose:
logger.info(f"Loading: {os.path.basename(file_path)}")
raw = mne.io.read_raw_fif(file_path, preload=True, verbose=False)
raw_list.append(raw)
# Concatenate if multiple files
if len(raw_list) > 1:
raw_empty = mne.concatenate_raws(raw_list)
else:
raw_empty = raw_list[0]
# Compute covariance
if verbose:
logger.info("Computing noise covariance matrix...")
noise_cov = mne.compute_raw_covariance(
raw_empty, method='empirical', verbose=verbose
)
# Save covariance
noise_cov_dir = get_derivatives_manager(data_path).get_noise_covariance_path(
subject_id, create=True)
noise_cov_file = noise_cov_dir / f'sub-{subject_id:02d}_task-avs_desc-emptyroom_cov.fif'
mne.write_cov(str(noise_cov_file), noise_cov, verbose=verbose)
if verbose:
logger.info(f"Saved noise covariance: {noise_cov_file}")
return noise_cov, str(noise_cov_file)