77 throw std::invalid_argument(
"DecodingCsp::fit: epochs must be non-empty");
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");
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");
91 auto it = classSet.begin();
95 std::vector<MatrixXd> epochs1, epochs2;
96 for (
int i = 0; i < y.size(); ++i) {
98 epochs1.push_back(epochs[
static_cast<size_t>(i)]);
100 epochs2.push_back(epochs[
static_cast<size_t>(i)]);
104 const auto n_ch = epochs1[0].rows();
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);
112 cov1 /=
static_cast<double>(epochs1.size());
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);
119 cov2 /=
static_cast<double>(epochs2.size());
122 MatrixXd cov_comp = cov1 + cov2;
125 SelfAdjointEigenSolver<MatrixXd> eig_comp(cov_comp);
126 VectorXd d = eig_comp.eigenvalues();
127 MatrixXd U = eig_comp.eigenvectors();
129 const double d_min = d.maxCoeff() * 1e-10;
130 for (Index i = 0; i < d.size(); ++i) {
135 MatrixXd W = d.array().sqrt().inverse().matrix().asDiagonal() * U.transpose();
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();
144 MatrixXd all_filters = B.transpose() * W;
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);
151 m_filters = MatrixXd(n_total, n_ch);
152 VectorXd eigenvalues(n_total);
155 for (
int i = 0; i < n_per_class; ++i) {
156 m_filters.row(i) = all_filters.row(i);
157 eigenvalues(i) = lambdas(i);
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);
168 auto svd = m_filters.bdcSvd<ComputeThinU | ComputeThinV>();
169 m_patterns =
svd.solve(MatrixXd::Identity(n_total, n_total));
172 MatrixXd powerFeatures = computePowerFeatures(epochs);
173 m_mean = powerFeatures.colwise().mean();
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()));
190 throw std::runtime_error(
"DecodingCsp::transform: not fitted");
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());
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)];
207 MatrixXd
X = computePowerFeatures(epochs);
210 X =
X.array().max(1e-30).log().matrix();
213 for (
int c = 0; c <
X.cols(); ++c) {
217 X.col(c) = (
X.col(c).array() - m_mean(c)) / s;