v2.0.0
Loading...
Searching...
No Matches
decoding_csp.cpp
Go to the documentation of this file.
1//=============================================================================================================
30
31//=============================================================================================================
32// INCLUDES
33//=============================================================================================================
34
35#include "decoding_csp.h"
36
37//=============================================================================================================
38// EIGEN INCLUDES
39//=============================================================================================================
40
41#include <Eigen/Eigenvalues>
42#include <Eigen/SVD>
43
44//=============================================================================================================
45// STL INCLUDES
46//=============================================================================================================
47
48#include <algorithm>
49#include <cmath>
50#include <set>
51#include <stdexcept>
52
53//=============================================================================================================
54// USED NAMESPACES
55//=============================================================================================================
56
57using namespace DECODINGLIB;
58using namespace Eigen;
59
60//=============================================================================================================
61
63 TransformMode transformInto,
64 bool useLog)
65: m_nComponents(nComponents)
66, m_transformInto(transformInto)
67, m_useLog(useLog)
68{
69}
70
71//=============================================================================================================
72
73void DecodingCsp::fit(const std::vector<MatrixXd>& epochs,
74 const VectorXi& y)
75{
76 if (epochs.empty()) {
77 throw std::invalid_argument("DecodingCsp::fit: epochs must be non-empty");
78 }
79 if (static_cast<int>(epochs.size()) != y.size()) {
80 throw std::invalid_argument(
81 "DecodingCsp::fit: epochs and y must have the same length");
82 }
83
84 // Identify unique classes
85 std::set<int> classSet(y.data(), y.data() + y.size());
86 if (classSet.size() != 2) {
87 throw std::invalid_argument(
88 "DecodingCsp::fit: y must contain exactly 2 unique class labels");
89 }
90
91 auto it = classSet.begin();
92 int label1 = *it++;
93
94 // Split epochs by class
95 std::vector<MatrixXd> epochs1, epochs2;
96 for (int i = 0; i < y.size(); ++i) {
97 if (y(i) == label1) {
98 epochs1.push_back(epochs[static_cast<size_t>(i)]);
99 } else {
100 epochs2.push_back(epochs[static_cast<size_t>(i)]);
101 }
102 }
103
104 const auto n_ch = epochs1[0].rows();
105
106 // Average covariance for each class
107 MatrixXd cov1 = MatrixXd::Zero(n_ch, n_ch);
108 for (const auto& epoch : epochs1) {
109 MatrixXd centered = epoch.colwise() - epoch.rowwise().mean();
110 cov1 += centered * centered.transpose() / static_cast<double>(epoch.cols() - 1);
111 }
112 cov1 /= static_cast<double>(epochs1.size());
113
114 MatrixXd cov2 = MatrixXd::Zero(n_ch, n_ch);
115 for (const auto& epoch : epochs2) {
116 MatrixXd centered = epoch.colwise() - epoch.rowwise().mean();
117 cov2 += centered * centered.transpose() / static_cast<double>(epoch.cols() - 1);
118 }
119 cov2 /= static_cast<double>(epochs2.size());
120
121 // Composite covariance
122 MatrixXd cov_comp = cov1 + cov2;
123
124 // Whitening: W = D^{-1/2} U^T
125 SelfAdjointEigenSolver<MatrixXd> eig_comp(cov_comp);
126 VectorXd d = eig_comp.eigenvalues();
127 MatrixXd U = eig_comp.eigenvectors();
128
129 const double d_min = d.maxCoeff() * 1e-10;
130 for (Index i = 0; i < d.size(); ++i) {
131 if (d(i) < d_min)
132 d(i) = d_min;
133 }
134
135 MatrixXd W = d.array().sqrt().inverse().matrix().asDiagonal() * U.transpose();
136
137 // Whiten class-1 covariance and eigendecompose
138 MatrixXd S1 = W * cov1 * W.transpose();
139 SelfAdjointEigenSolver<MatrixXd> eig_s1(S1);
140 VectorXd lambdas = eig_s1.eigenvalues();
141 MatrixXd B = eig_s1.eigenvectors();
142
143 // All filters in eigenvalue order (ascending)
144 MatrixXd all_filters = B.transpose() * W;
145
146 // Select first and last components
147 int n_per_class = m_nComponents / 2;
148 int n_total = std::min(m_nComponents, static_cast<int>(n_ch));
149 n_per_class = std::min(n_per_class, static_cast<int>(n_ch) / 2);
150
151 m_filters = MatrixXd(n_total, n_ch);
152 VectorXd eigenvalues(n_total);
153
154 // Bottom n_per_class (maximize class-2 variance)
155 for (int i = 0; i < n_per_class; ++i) {
156 m_filters.row(i) = all_filters.row(i);
157 eigenvalues(i) = lambdas(i);
158 }
159 // Top (n_total - n_per_class) (maximize class-1 variance)
160 for (int i = 0; i < n_total - n_per_class; ++i) {
161 m_filters.row(n_per_class + i) =
162 all_filters.row(static_cast<int>(n_ch) - 1 - i);
163 eigenvalues(n_per_class + i) =
164 lambdas(static_cast<int>(n_ch) - 1 - i);
165 }
166
167 // Patterns = pinv(filters) — shape (n_ch × n_comp)
168 auto svd = m_filters.bdcSvd<ComputeThinU | ComputeThinV>();
169 m_patterns = svd.solve(MatrixXd::Identity(n_total, n_total));
170
171 // Compute mean band power for z-score normalisation
172 MatrixXd powerFeatures = computePowerFeatures(epochs);
173 m_mean = powerFeatures.colwise().mean();
174
175 VectorXd centered = VectorXd(powerFeatures.rows());
176 m_std = VectorXd(powerFeatures.cols());
177 for (int c = 0; c < powerFeatures.cols(); ++c) {
178 centered = powerFeatures.col(c).array() - m_mean(c);
179 m_std(c) = std::sqrt(centered.squaredNorm() / static_cast<double>(centered.size()));
180 }
181
182 m_fitted = true;
183}
184
185//=============================================================================================================
186
187MatrixXd DecodingCsp::transform(const std::vector<MatrixXd>& epochs) const
188{
189 if (!m_fitted) {
190 throw std::runtime_error("DecodingCsp::transform: not fitted");
191 }
192
193 if (m_transformInto == TransformMode::CspSpace) {
194 // Return data in CSP space: (n_epochs * n_components, n_times)
195 const int nEpochs = static_cast<int>(epochs.size());
196 const int nComp = static_cast<int>(m_filters.rows());
197 const int nTimes = static_cast<int>(epochs[0].cols());
198
199 MatrixXd result(nEpochs * nComp, nTimes);
200 for (int e = 0; e < nEpochs; ++e) {
201 result.middleRows(static_cast<Eigen::Index>(e) * nComp, nComp) = m_filters * epochs[static_cast<size_t>(e)];
202 }
203 return result;
204 }
205
206 // AveragePower mode
207 MatrixXd X = computePowerFeatures(epochs);
208
209 if (m_useLog) {
210 X = X.array().max(1e-30).log().matrix();
211 } else {
212 // z-score
213 for (int c = 0; c < X.cols(); ++c) {
214 double s = m_std(c);
215 if (s < 1e-15)
216 s = 1.0;
217 X.col(c) = (X.col(c).array() - m_mean(c)) / s;
218 }
219 }
220
221 return X;
222}
223
224//=============================================================================================================
225
226MatrixXd DecodingCsp::fitTransform(const std::vector<MatrixXd>& epochs,
227 const VectorXi& y)
228{
229 fit(epochs, y);
230 return transform(epochs);
231}
232
233//=============================================================================================================
234
235MatrixXd DecodingCsp::inverseTransform(const MatrixXd& X) const
236{
237 if (!m_fitted) {
238 throw std::runtime_error("DecodingCsp::inverseTransform: not fitted");
239 }
240 if (m_transformInto != TransformMode::AveragePower) {
241 throw std::runtime_error(
242 "DecodingCsp::inverseTransform: only valid for AveragePower mode");
243 }
244
245 const int nComp = static_cast<int>(m_filters.rows());
246 if (X.cols() != nComp) {
247 throw std::invalid_argument(
248 "DecodingCsp::inverseTransform: X must have n_components columns");
249 }
250
251 // patterns_ is (n_channels × n_components), X is (n_epochs × n_components)
252 // Result: (n_epochs × n_channels)
253 return X * m_patterns.transpose();
254}
255
256//=============================================================================================================
257
258const MatrixXd& DecodingCsp::filters() const
259{
260 if (!m_fitted) {
261 throw std::runtime_error("DecodingCsp::filters: not fitted");
262 }
263 return m_filters;
264}
265
266//=============================================================================================================
267
268const MatrixXd& DecodingCsp::patterns() const
269{
270 if (!m_fitted) {
271 throw std::runtime_error("DecodingCsp::patterns: not fitted");
272 }
273 return m_patterns;
274}
275
276//=============================================================================================================
277
278const VectorXd& DecodingCsp::mean() const
279{
280 if (!m_fitted) {
281 throw std::runtime_error("DecodingCsp::mean: not fitted");
282 }
283 return m_mean;
284}
285
286//=============================================================================================================
287
288const VectorXd& DecodingCsp::stddev() const
289{
290 if (!m_fitted) {
291 throw std::runtime_error("DecodingCsp::stddev: not fitted");
292 }
293 return m_std;
294}
295
296//=============================================================================================================
297
299{
300 return m_fitted;
301}
302
303//=============================================================================================================
304
305MatrixXd DecodingCsp::computePowerFeatures(
306 const std::vector<MatrixXd>& epochs) const
307{
308 const int nEpochs = static_cast<int>(epochs.size());
309 const int nComp = static_cast<int>(m_filters.rows());
310
311 MatrixXd features(nEpochs, nComp);
312 for (int e = 0; e < nEpochs; ++e) {
313 MatrixXd filtered = m_filters * epochs[static_cast<size_t>(e)];
314 for (int c = 0; c < nComp; ++c) {
315 features(e, c) = filtered.row(c).squaredNorm() / static_cast<double>(filtered.cols());
316 }
317 }
318
319 return features;
320}
Eigen::JacobiSVD< Eigen::Matrix3f > svd(S, Eigen::ComputeFullU|Eigen::ComputeFullV)
Common Spatial Patterns (CSP) for two-class discriminative spatial filtering of band-passed M/EEG.
constexpr int X
Supervised and unsupervised spatial-filter decompositions for M/EEG decoding.
const Eigen::MatrixXd & filters() const
const Eigen::VectorXd & mean() const
const Eigen::VectorXd & stddev() const
Eigen::MatrixXd inverseTransform(const Eigen::MatrixXd &X) const
Eigen::MatrixXd fitTransform(const std::vector< Eigen::MatrixXd > &epochs, const Eigen::VectorXi &y)
void fit(const std::vector< Eigen::MatrixXd > &epochs, const Eigen::VectorXi &y)
DecodingCsp(int nComponents=4, TransformMode transformInto=TransformMode::AveragePower, bool useLog=true)
const Eigen::MatrixXd & patterns() const
Eigen::MatrixXd transform(const std::vector< Eigen::MatrixXd > &epochs) const