v2.0.0
Loading...
Searching...
No Matches
decoding_ssd.cpp
Go to the documentation of this file.
1//=============================================================================================================
34
35//=============================================================================================================
36// INCLUDES
37//=============================================================================================================
38
39#include "decoding_ssd.h"
40
41//=============================================================================================================
42// EIGEN INCLUDES
43//=============================================================================================================
44
45#include <Eigen/Eigenvalues>
46
47//=============================================================================================================
48// STL INCLUDES
49//=============================================================================================================
50
51#include <algorithm>
52#include <cmath>
53#include <stdexcept>
54
55//=============================================================================================================
56// USED NAMESPACES
57//=============================================================================================================
58
59using namespace DECODINGLIB;
60using namespace Eigen;
61
62//=============================================================================================================
63
64DecodingSsd::DecodingSsd(int nComponents, double regParam)
65: m_nComponents(nComponents)
66, m_regParam(regParam)
67{
68}
69
70//=============================================================================================================
71
72void DecodingSsd::fit(const Ref<const MatrixXd>& data,
73 double sfreq,
74 double signalLow, double signalHigh,
75 double noiseLow, double noiseHigh)
76{
77 if (noiseLow > signalLow || signalHigh > noiseHigh) {
78 throw std::invalid_argument(
79 "DecodingSsd::fit: signal band must be within noise band");
80 }
81
82 const auto n_ch = data.rows();
83 const auto n_times = data.cols();
84
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");
88 }
89
90 int n_comp = std::min(m_nComponents, static_cast<int>(n_ch));
91
92 // Mean-center
93 MatrixXd centered = data.colwise() - data.rowwise().mean();
94
95 // Bandpass filter for signal and noise bands
96 MatrixXd data_signal = bandpassFilter(centered, sfreq, signalLow, signalHigh);
97 MatrixXd data_noise = bandpassFilter(centered, sfreq, noiseLow, noiseHigh);
98
99 // Covariance matrices
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);
102
103 // Regularize noise covariance
104 cov_noise += m_regParam * cov_noise.trace() / static_cast<double>(n_ch) * MatrixXd::Identity(n_ch, n_ch);
105
106 // Generalized eigenvalue problem: C_signal w = λ C_noise w
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");
111 }
112
113 // Eigenvalues ascending → reverse for descending
114 VectorXd evals = ges.eigenvalues().reverse();
115 MatrixXd evecs = ges.eigenvectors().rowwise().reverse();
116
117 m_eigenvalues = evals.head(n_comp);
118 m_filters = evecs.leftCols(n_comp).transpose(); // (n_comp × n_ch)
119
120 // Patterns: A = C W inv(W^T C W)
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()); // (n_ch × n_comp)
124
125 m_fitted = true;
126}
127
128//=============================================================================================================
129
130MatrixXd DecodingSsd::transform(const Ref<const MatrixXd>& data) const
131{
132 if (!m_fitted) {
133 throw std::runtime_error("DecodingSsd::transform: not fitted");
134 }
135 return m_filters * data;
136}
137
138//=============================================================================================================
139
140MatrixXd DecodingSsd::fitTransform(const Ref<const MatrixXd>& data,
141 double sfreq,
142 double signalLow, double signalHigh,
143 double noiseLow, double noiseHigh)
144{
145 fit(data, sfreq, signalLow, signalHigh, noiseLow, noiseHigh);
146 return transform(data);
147}
148
149//=============================================================================================================
150
151MatrixXd DecodingSsd::apply(const Ref<const MatrixXd>& data) const
152{
153 if (!m_fitted) {
154 throw std::runtime_error("DecodingSsd::apply: not fitted");
155 }
156
157 // Denoise: project to SSD space, then reconstruct using patterns
158 // patterns_ is (n_ch × n_comp), ssdSources is (n_comp × n_times)
159 MatrixXd ssdSources = transform(data);
160 return m_patterns * ssdSources;
161}
162
163//=============================================================================================================
164
165const MatrixXd& DecodingSsd::filters() const
166{
167 if (!m_fitted) {
168 throw std::runtime_error("DecodingSsd::filters: not fitted");
169 }
170 return m_filters;
171}
172
173//=============================================================================================================
174
175const MatrixXd& DecodingSsd::patterns() const
176{
177 if (!m_fitted) {
178 throw std::runtime_error("DecodingSsd::patterns: not fitted");
179 }
180 return m_patterns;
181}
182
183//=============================================================================================================
184
185const VectorXd& DecodingSsd::eigenvalues() const
186{
187 if (!m_fitted) {
188 throw std::runtime_error("DecodingSsd::eigenvalues: not fitted");
189 }
190 return m_eigenvalues;
191}
192
193//=============================================================================================================
194
196{
197 return m_fitted;
198}
199
200//=============================================================================================================
201
202MatrixXd DecodingSsd::bandpassFilter(const MatrixXd& data,
203 double sfreq,
204 double lowFreq,
205 double highFreq)
206{
207 const auto n_ch = data.rows();
208 const auto n_times = data.cols();
209
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)
214 ++filter_order;
215 filter_order = std::min(filter_order, 128);
216
217 const int half = filter_order / 2;
218 const double pi = M_PI;
219
220 // Windowed-sinc coefficients
221 VectorXd h(filter_order + 1);
222 double w_low = 2.0 * pi * lowFreq / sfreq;
223 double w_high = 2.0 * pi * highFreq / sfreq;
224
225 for (int i = 0; i <= filter_order; ++i) {
226 int n = i - half;
227 if (n == 0) {
228 h(i) = (w_high - w_low) / pi;
229 } else {
230 h(i) = (std::sin(w_high * static_cast<double>(n)) - std::sin(w_low * static_cast<double>(n))) / (pi * static_cast<double>(n));
231 }
232 // Hamming window
233 h(i) *= 0.54 - 0.46 * std::cos(2.0 * pi * static_cast<double>(i) / static_cast<double>(filter_order));
234 }
235
236 // Normalize to unit gain at center frequency
237 double center = (lowFreq + highFreq) / 2.0;
238 double w_center = 2.0 * pi * center / sfreq;
239 double gain = 0.0;
240 for (int i = 0; i <= filter_order; ++i) {
241 gain += h(i) * std::cos(w_center * static_cast<double>(i - half));
242 }
243 if (std::abs(gain) > 1e-12)
244 h /= std::abs(gain);
245
246 // Forward + reverse convolution (zero-phase)
247 MatrixXd result(n_ch, n_times);
248
249 for (Index ch = 0; ch < n_ch; ++ch) {
250 VectorXd forward(n_times);
251 for (Index t = 0; t < n_times; ++t) {
252 double sum = 0.0;
253 for (int k = 0; k <= filter_order; ++k) {
254 auto idx = t - k;
255 if (idx >= 0 && idx < n_times)
256 sum += h(k) * data(ch, idx);
257 }
258 forward(t) = sum;
259 }
260
261 VectorXd backward(n_times);
262 for (Index t = n_times - 1; t >= 0; --t) {
263 double sum = 0.0;
264 for (int k = 0; k <= filter_order; ++k) {
265 auto idx = t + k;
266 if (idx >= 0 && idx < n_times)
267 sum += h(k) * forward(idx);
268 }
269 backward(t) = sum;
270 }
271
272 result.row(ch) = backward.transpose();
273 }
274
275 return result;
276}
Spatio-Spectral Decomposition (SSD) for noise-aware narrowband enhancement of continuous M/EEG.
#define M_PI
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)