78 if (epochs.isEmpty()) {
79 qWarning() <<
"Xdawn::fit: empty epoch list.";
84 goodIdx.reserve(epochs.size());
85 for (
int i = 0; i < epochs.size(); ++i) {
86 if (!epochs[i].bReject) {
91 if (goodIdx.isEmpty()) {
92 qWarning() <<
"Xdawn::fit: no non-rejected epochs available.";
96 const int nCh =
static_cast<int>(epochs[goodIdx[0]].epoch.rows());
97 const int nSamp =
static_cast<int>(epochs[goodIdx[0]].epoch.cols());
98 if (nCh == 0 || nSamp == 0) {
99 qWarning() <<
"Xdawn::fit: epoch matrices are empty.";
103 for (
int idx : goodIdx) {
104 if (epochs[idx].epoch.rows() != nCh || epochs[idx].epoch.cols() != nSamp) {
105 qWarning() <<
"Xdawn::fit: epoch dimension mismatch.";
110 nComponents = std::max(1, std::min(nComponents, nCh));
112 QHash<int, MatrixXd> classSums;
113 QHash<int, int> classCounts;
114 MatrixXd targetSum = MatrixXd::Zero(nCh, nSamp);
117 for (
int idx : goodIdx) {
119 if (!classSums.contains(ep.
event)) {
120 classSums.insert(ep.
event, MatrixXd::Zero(nCh, nSamp));
121 classCounts.insert(ep.
event, 0);
125 classCounts[ep.
event] += 1;
127 if (ep.
event == iTargetEvent) {
128 targetSum += ep.
epoch;
134 qWarning() <<
"Xdawn::fit: no target epochs found for event" << iTargetEvent;
140 MatrixXd noiseCov = MatrixXd::Zero(nCh, nCh);
141 MatrixXd dataCov = MatrixXd::Zero(nCh, nCh);
142 long long nNoiseSamples = 0;
143 long long nDataSamples = 0;
145 QHash<int, MatrixXd> classMeans;
146 for (
auto it = classSums.constBegin(); it != classSums.constEnd(); ++it) {
147 classMeans.insert(it.key(), it.value() /
static_cast<double>(classCounts.value(it.key())));
150 for (
int idx : goodIdx) {
152 const MatrixXd residual = ep.
epoch - classMeans.value(ep.
event);
155 noiseCov += residual * residual.transpose();
156 nDataSamples += nSamp;
157 nNoiseSamples += nSamp;
160 if (nNoiseSamples <= 0 || nDataSamples <= 0) {
161 qWarning() <<
"Xdawn::fit: failed to accumulate covariance samples.";
166 result.
matNoiseCov = noiseCov /
static_cast<double>(nNoiseSamples);
167 dataCov = dataCov /
static_cast<double>(nDataSamples);
169 const double traceNoise = result.
matNoiseCov.trace();
170 const double regValue = std::max(dReg, 0.0) * ((traceNoise > 0.0) ? traceNoise /
static_cast<double>(nCh) : 1.0);
172 regNoiseCov.diagonal().array() += regValue;
174 SelfAdjointEigenSolver<MatrixXd> noiseEig(regNoiseCov);
175 if (noiseEig.info() != Success) {
176 qWarning() <<
"Xdawn::fit: noise covariance eigendecomposition failed.";
180 VectorXd noiseVals = noiseEig.eigenvalues().cwiseMax(1e-12);
181 MatrixXd noiseVecs = noiseEig.eigenvectors();
182 MatrixXd invSqrtNoise = noiseVecs * noiseVals.cwiseInverse().cwiseSqrt().asDiagonal() * noiseVecs.transpose();
184 MatrixXd whitenedSignal = invSqrtNoise * result.
matSignalCov * invSqrtNoise;
185 SelfAdjointEigenSolver<MatrixXd> signalEig(whitenedSignal);
186 if (signalEig.info() != Success) {
187 qWarning() <<
"Xdawn::fit: signal covariance eigendecomposition failed.";
192 const MatrixXd signalVecsAsc = signalEig.eigenvectors().rightCols(nComponents);
193 MatrixXd signalVecs(signalVecsAsc.rows(), signalVecsAsc.cols());
194 for (
int i = 0; i < nComponents; ++i) {
195 signalVecs.col(i) = signalVecsAsc.col(nComponents - 1 - i);
198 result.
matFilters = invSqrtNoise * signalVecs;
200 for (
int col = 0; col < result.
matFilters.cols(); ++col) {
201 const double noiseNorm = std::sqrt(result.
matFilters.col(col).transpose() * regNoiseCov * result.
matFilters.col(col));
202 if (noiseNorm > 1e-12) {