45#include <Eigen/Eigenvalues>
65: m_nComponents(nComponents)
74 double signalLow,
double signalHigh,
75 double noiseLow,
double noiseHigh)
77 if (noiseLow > signalLow || signalHigh > noiseHigh) {
78 throw std::invalid_argument(
79 "DecodingSsd::fit: signal band must be within noise band");
82 const auto n_ch = data.rows();
83 const auto n_times = data.cols();
85 if (n_ch < 2 || n_times < 2) {
86 throw std::invalid_argument(
87 "DecodingSsd::fit: data must have >= 2 channels and >= 2 time points");
90 int n_comp = std::min(m_nComponents,
static_cast<int>(n_ch));
93 MatrixXd centered = data.colwise() - data.rowwise().mean();
96 MatrixXd data_signal = bandpassFilter(centered, sfreq, signalLow, signalHigh);
97 MatrixXd data_noise = bandpassFilter(centered, sfreq, noiseLow, noiseHigh);
100 MatrixXd cov_signal = (data_signal * data_signal.transpose()) /
static_cast<double>(n_times - 1);
101 MatrixXd cov_noise = (data_noise * data_noise.transpose()) /
static_cast<double>(n_times - 1);
104 cov_noise += m_regParam * cov_noise.trace() /
static_cast<double>(n_ch) * MatrixXd::Identity(n_ch, n_ch);
107 GeneralizedSelfAdjointEigenSolver<MatrixXd> ges(cov_signal, cov_noise);
108 if (ges.info() != Eigen::Success) {
109 throw std::runtime_error(
110 "DecodingSsd::fit: eigenvalue decomposition failed");
114 VectorXd evals = ges.eigenvalues().reverse();
115 MatrixXd evecs = ges.eigenvectors().rowwise().reverse();
117 m_eigenvalues = evals.head(n_comp);
118 m_filters = evecs.leftCols(n_comp).transpose();
121 MatrixXd cov_full = (centered * centered.transpose()) /
static_cast<double>(n_times - 1);
122 MatrixXd WtCW = m_filters * cov_full * m_filters.transpose();
123 m_patterns = (cov_full * m_filters.transpose() * WtCW.inverse());
133 throw std::runtime_error(
"DecodingSsd::transform: not fitted");
135 return m_filters * data;
142 double signalLow,
double signalHigh,
143 double noiseLow,
double noiseHigh)
145 fit(data, sfreq, signalLow, signalHigh, noiseLow, noiseHigh);
154 throw std::runtime_error(
"DecodingSsd::apply: not fitted");
160 return m_patterns * ssdSources;
168 throw std::runtime_error(
"DecodingSsd::filters: not fitted");
178 throw std::runtime_error(
"DecodingSsd::patterns: not fitted");
188 throw std::runtime_error(
"DecodingSsd::eigenvalues: not fitted");
190 return m_eigenvalues;
202MatrixXd DecodingSsd::bandpassFilter(
const MatrixXd& data,
207 const auto n_ch = data.rows();
208 const auto n_times = data.cols();
210 int filter_order = std::min(
211 static_cast<int>(std::round(3.0 * sfreq / lowFreq)),
212 static_cast<int>(n_times) - 1);
213 if (filter_order % 2 != 0)
215 filter_order = std::min(filter_order, 128);
217 const int half = filter_order / 2;
218 const double pi =
M_PI;
221 VectorXd h(filter_order + 1);
222 double w_low = 2.0 * pi * lowFreq / sfreq;
223 double w_high = 2.0 * pi * highFreq / sfreq;
225 for (
int i = 0; i <= filter_order; ++i) {
228 h(i) = (w_high - w_low) / pi;
230 h(i) = (std::sin(w_high *
static_cast<double>(n)) - std::sin(w_low *
static_cast<double>(n))) / (pi *
static_cast<double>(n));
233 h(i) *= 0.54 - 0.46 * std::cos(2.0 * pi *
static_cast<double>(i) /
static_cast<double>(filter_order));
237 double center = (lowFreq + highFreq) / 2.0;
238 double w_center = 2.0 * pi * center / sfreq;
240 for (
int i = 0; i <= filter_order; ++i) {
241 gain += h(i) * std::cos(w_center *
static_cast<double>(i - half));
243 if (std::abs(gain) > 1e-12)
247 MatrixXd result(n_ch, n_times);
249 for (Index ch = 0; ch < n_ch; ++ch) {
250 VectorXd forward(n_times);
251 for (Index t = 0; t < n_times; ++t) {
253 for (
int k = 0; k <= filter_order; ++k) {
255 if (idx >= 0 && idx < n_times)
256 sum += h(k) * data(ch, idx);
261 VectorXd backward(n_times);
262 for (Index t = n_times - 1; t >= 0; --t) {
264 for (
int k = 0; k <= filter_order; ++k) {
266 if (idx >= 0 && idx < n_times)
267 sum += h(k) * forward(idx);
272 result.row(ch) = backward.transpose();
Spatio-Spectral Decomposition (SSD) for noise-aware narrowband enhancement of continuous M/EEG.
Supervised and unsupervised spatial-filter decompositions for M/EEG decoding.
DecodingSsd(int nComponents=6, double regParam=0.05)
Eigen::MatrixXd apply(const Eigen::Ref< const Eigen::MatrixXd > &data) const
const Eigen::MatrixXd & filters() const
const Eigen::VectorXd & eigenvalues() const
Eigen::MatrixXd transform(const Eigen::Ref< const Eigen::MatrixXd > &data) const
const Eigen::MatrixXd & patterns() const
void fit(const Eigen::Ref< const Eigen::MatrixXd > &data, double sfreq, double signalLow, double signalHigh, double noiseLow, double noiseHigh)
Eigen::MatrixXd fitTransform(const Eigen::Ref< const Eigen::MatrixXd > &data, double sfreq, double signalLow, double signalHigh, double noiseLow, double noiseHigh)