Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
9 changes: 7 additions & 2 deletions neurokit2/microstates/microstates_classify.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,7 +3,7 @@
from ..misc import replace


def microstates_classify(segmentation, microstates):
def microstates_classify(segmentation, microstates, return_order=False):
"""**Reorder (sort) the microstates (experimental)**

Reorder (sort) the microstates (experimental) based on the pattern of values in the vector of
Expand All @@ -15,6 +15,9 @@ def microstates_classify(segmentation, microstates):
Vector containing the segmentation.
microstates : Union[np.array, dict]
Array of microstates maps . Defaults to ``None``.
return_order : bool
If ``True``, also return the indices used to reorder the microstate maps. Defaults to
``False``.

Returns
-------
Expand Down Expand Up @@ -45,9 +48,11 @@ def microstates_classify(segmentation, microstates):
new_order = _microstates_sort(microstates)
microstates = microstates[new_order]

replacement = dict(enumerate(new_order))
replacement = {old: new for new, old in enumerate(new_order)}
segmentation = replace(segmentation, replacement)

if return_order is True:
return segmentation, microstates, new_order
return segmentation, microstates


Expand Down
21 changes: 8 additions & 13 deletions neurokit2/microstates/microstates_clean.py
Original file line number Diff line number Diff line change
@@ -1,9 +1,8 @@
import numpy as np
import pandas as pd

from ..eeg import eeg_gfp
from ..stats import standardize
from .microstates_peaks import microstates_peaks
from .microstates_peaks import _microstates_sanitize_eeg, microstates_peaks


def microstates_clean(eeg, sampling_rate=None, train="gfp", standardize_eeg=True, normalize=True, gfp_method="l1", **kwargs):
Expand Down Expand Up @@ -62,26 +61,22 @@ def microstates_clean(eeg, sampling_rate=None, train="gfp", standardize_eeg=True
.eeg_gfp, microstates_peaks, .microstates_segment

"""
# If MNE object
if isinstance(eeg, (pd.DataFrame, np.ndarray)) is False:
sampling_rate = eeg.info["sfreq"]
info = eeg.info
eeg = eeg.get_data()
else:
info = None
eeg, sampling_rate, info = _microstates_sanitize_eeg(eeg, sampling_rate=sampling_rate)

# Normalization
if standardize_eeg is True:
eeg = standardize(eeg, **kwargs)
standardize_kwargs = {key: kwargs[key] for key in ["robust", "window"] if key in kwargs}
eeg = standardize(eeg.T, **standardize_kwargs).T

# Get GFP
gfp = eeg_gfp(eeg, sampling_rate=sampling_rate, normalize=normalize, method=gfp_method, **kwargs)
gfp_kwargs = {key: kwargs[key] for key in ["robust", "smooth"] if key in kwargs}
gfp = eeg_gfp(eeg, sampling_rate=sampling_rate, normalize=normalize, method=gfp_method, **gfp_kwargs)

# If train is a custom of vector (assume it's the pre-computed peaks)
if isinstance(train, (list, np.ndarray)):
peaks = train
peaks = np.asarray(train, dtype=int)
# Find peaks in the global field power (GFP) or take a given amount of indices
else:
peaks = microstates_peaks(eeg, gfp=train, sampling_rate=sampling_rate, **kwargs)
peaks = microstates_peaks(eeg, gfp=gfp if train == "gfp" else train, sampling_rate=sampling_rate)

return eeg, peaks, gfp, info
8 changes: 2 additions & 6 deletions neurokit2/microstates/microstates_findnumber.py
Original file line number Diff line number Diff line change
Expand Up @@ -4,6 +4,7 @@

from ..misc import find_knee, progress_bar
from ..stats.cluster_quality import _cluster_quality_dispersion
from .microstates_peaks import _microstates_sanitize_eeg
from .microstates_segment import microstates_segment


Expand Down Expand Up @@ -63,12 +64,7 @@ def microstates_findnumber(eeg, n_max=12, method="GEV", clustering_method="kmod"

"""
# Retrieve data
if isinstance(eeg, (pd.DataFrame, np.ndarray)) is False:
data = eeg.get_data()
elif isinstance(eeg, pd.DataFrame):
data = eeg.values
else:
data = eeg.copy()
data, _, _ = _microstates_sanitize_eeg(eeg)

# Loop accross number and get indices of fit
n_channel, _ = data.shape
Expand Down
71 changes: 43 additions & 28 deletions neurokit2/microstates/microstates_peaks.py
Original file line number Diff line number Diff line change
Expand Up @@ -53,46 +53,61 @@ def microstates_peaks(eeg, gfp=None, sampling_rate=None, distance_between=0.01,
.eeg_gfp

"""
if isinstance(eeg, (pd.DataFrame, np.ndarray)) is False:
sampling_rate = eeg.info["sfreq"]
eeg = eeg.get_data()

if sampling_rate is None:
raise ValueError(
"NeuroKit error: microstates_peaks(): The sampling_rate is requested ",
"for this function to run. Please provide it as an argument.",
)

# If we don't want to rely on peaks but take uniformly spaced samples
# (used in microstates_clustering)
if isinstance(gfp, (int, float)):
if gfp <= 1: # If fraction
gfp = int(gfp * len(eeg[0, :]))
return np.linspace(0, len(eeg[0, :]), gfp, endpoint=False, dtype=int)
eeg, sampling_rate, _ = _microstates_sanitize_eeg(eeg, sampling_rate=sampling_rate)

# Deal with string inputs
if isinstance(gfp, str):
if gfp == "all":
gfp = False
elif gfp == "gfp":
gfp = True
if gfp.lower() == "all":
return np.arange(eeg.shape[1])
if gfp.lower() == "gfp":
gfp = None
else:
raise ValueError(
"The `gfp` argument was not understood.",
)
raise ValueError("The `gfp` argument was not understood.")

# If we want ALL the indices
if gfp is False:
return np.arange(len(eeg))
# If we don't want to rely on peaks but take uniformly spaced samples
# (used in microstates_clustering)
if isinstance(gfp, (int, float, np.integer, np.floating)) and not isinstance(gfp, (bool, np.bool_)):
if gfp <= 1: # If fraction
gfp = int(gfp * eeg.shape[1])
if not float(gfp).is_integer() or gfp < 1 or gfp > eeg.shape[1]:
raise ValueError("The number of training samples must be between 1 and the number of timepoints.")
return np.linspace(0, eeg.shape[1], int(gfp), endpoint=False, dtype=int)

if gfp is None or gfp is True:
gfp = eeg_gfp(eeg, sampling_rate=sampling_rate, **kwargs)
else:
gfp = np.asarray(gfp)
if gfp.ndim != 1 or len(gfp) != eeg.shape[1]:
raise ValueError("The precomputed `gfp` must contain one value per timepoint.")

# if gfp is True or gfp is None:
gfp = eeg_gfp(eeg, **kwargs)
if sampling_rate is None:
raise ValueError("NeuroKit error: microstates_peaks(): `sampling_rate` is required when detecting GFP peaks.")

peaks = _microstates_peaks_gfp(gfp=gfp, sampling_rate=sampling_rate, distance_between=distance_between)

return peaks


def _microstates_sanitize_eeg(eeg, sampling_rate=None):
"""Return EEG data as a channels-by-timepoints array."""
info = None
if isinstance(eeg, (pd.DataFrame, np.ndarray)) is False:
sampling_rate = eeg.info["sfreq"]
info = eeg.info
eeg = eeg.get_data()
elif isinstance(eeg, pd.DataFrame):
eeg = eeg.values

eeg = np.asarray(eeg)
if eeg.ndim == 3:
# MNE Epochs are ordered as epochs, channels, timepoints.
eeg = eeg.transpose(1, 0, 2).reshape(eeg.shape[1], -1)
if eeg.ndim != 2:
raise ValueError("EEG data must have shape (channels, timepoints) or (epochs, channels, timepoints).")

return eeg, sampling_rate, info


# =============================================================================
# Methods
# =============================================================================
Expand Down
23 changes: 18 additions & 5 deletions neurokit2/microstates/microstates_segment.py
Original file line number Diff line number Diff line change
Expand Up @@ -173,9 +173,16 @@ def microstates_segment(
on Biomedical Engineering.

"""
method = method.lower()
criterion = criterion.lower()
if criterion not in ["gev", "cv"]:
raise ValueError("`criterion` must be one of 'gev' or 'cv'.")
if n_runs < 1:
raise ValueError("`n_runs` must be at least 1.")

# Sanitize input
data, indices, gfp, info_mne = microstates_clean(
eeg, train=train, sampling_rate=sampling_rate, standardize_eeg=standardize_eeg, gfp_method=gfp_method, **kwargs
eeg, train=train, sampling_rate=sampling_rate, standardize_eeg=standardize_eeg, gfp_method=gfp_method
)

# Run clustering algorithm
Expand All @@ -187,7 +194,7 @@ def microstates_segment(
random_state = rng.choice(n_runs * 1000, n_runs, replace=False)

# Initialize values
gev = 0
gev = -np.inf
cv = np.inf
microstates = None
segmentation = None
Expand Down Expand Up @@ -229,7 +236,7 @@ def microstates_segment(
if current_residual < cv:
microstates, segmentation, polarity = current_microstates, s, p
cv, gev, gev_all = current_residual, g, g_all
info -= current_info
info = current_info

else:
# Run clustering algorithm on subset
Expand All @@ -243,7 +250,8 @@ def microstates_segment(
)

# Reorder
segmentation, microstates = microstates_classify(segmentation, microstates)
segmentation, microstates, new_order = microstates_classify(segmentation, microstates, return_order=True)
gev_all = gev_all[new_order]

# CLustering quality
# quality = cluster_quality(data, segmentation, clusters=microstates, info=info, n_random=10, sd=gfp)
Expand All @@ -268,7 +276,12 @@ def microstates_segment(
# =============================================================================
def _microstates_segment_runsegmentation(data, microstates, gfp, n_microstates):
# Find microstate corresponding to each datapoint
activation = microstates.dot(data)
maps_centered = microstates - np.mean(microstates, axis=1, keepdims=True)
map_norms = np.linalg.norm(maps_centered, axis=1, keepdims=True)
map_norms[map_norms == 0] = 1
maps_normalized = maps_centered / map_norms
data_centered = data - np.mean(data, axis=0, keepdims=True)
activation = maps_normalized.dot(data_centered)
segmentation = np.argmax(np.abs(activation), axis=0)
polarity = np.sign(np.choose(segmentation, activation))

Expand Down
2 changes: 1 addition & 1 deletion neurokit2/microstates/microstates_static.py
Original file line number Diff line number Diff line change
Expand Up @@ -192,7 +192,7 @@ def _microstates_lifetime(microstates, out=None):
tau_dict = {s: [] for s in states}
s = microstates[0] # current symbol
tau = 1.0 # current lifetime
for i in range(n):
for i in range(1, n):
if microstates[i] == s:
tau += 1.0
else:
Expand Down
11 changes: 8 additions & 3 deletions neurokit2/stats/cluster.py
Original file line number Diff line number Diff line change
Expand Up @@ -145,14 +145,14 @@ def cluster(data, method="kmeans", n_clusters=2, random_state=None, optimize=Fal

# ICA
elif method in ["ica", "independent", "independent component analysis"]:
out = _cluster_pca(data, n_clusters=n_clusters, random_state=random_state, **kwargs)
out = _cluster_ica(data, n_clusters=n_clusters, random_state=random_state, **kwargs)

# Mixture
elif method in ["mixture", "mixt"]:
out = _cluster_mixture(data, n_clusters=n_clusters, bayesian=False, random_state=random_state, **kwargs)

# Frederic's AAHC
elif method in ["aahc_frederic", "aahc_eegmicrostates"]:
elif method in ["aahc", "aahc_frederic", "aahc_eegmicrostates"]:
out = _cluster_aahc(data, n_clusters=n_clusters, random_state=random_state, **kwargs)

# Bayesian
Expand Down Expand Up @@ -470,7 +470,12 @@ def _cluster_ica(data, n_clusters=2, random_state=None, **kwargs):
"""Independent Component Analysis (ICA) for clustering."""
# Fit ICA
ica = sklearn.decomposition.FastICA(
n_components=n_clusters, algorithm="parallel", whiten=True, fun="exp", random_state=random_state, **kwargs
n_components=n_clusters,
algorithm="parallel",
whiten="unit-variance",
fun="exp",
random_state=random_state,
**kwargs,
)

ica = ica.fit(data)
Expand Down
Loading
Loading