v2.0.0
Loading...
Searching...
No Matches
bad_channel_detect.cpp
Go to the documentation of this file.
1//=============================================================================================================
12
13//=============================================================================================================
14// INCLUDES
15//=============================================================================================================
16
17#include "bad_channel_detect.h"
18
19//=============================================================================================================
20// EIGEN INCLUDES
21//=============================================================================================================
22
23#include <Eigen/Core>
24
25//=============================================================================================================
26// C++ INCLUDES
27//=============================================================================================================
28
29#include <cmath>
30#include <algorithm>
31
32//=============================================================================================================
33// USED NAMESPACES
34//=============================================================================================================
35
36using namespace UTILSLIB;
37using namespace Eigen;
38
39//=============================================================================================================
40// PRIVATE
41//=============================================================================================================
42
43double BadChannelDetect::pearsonCorr(const RowVectorXd& a, const RowVectorXd& b)
44{
45 if (a.size() != b.size() || a.size() == 0)
46 return 0.0;
47
48 RowVectorXd ac = a.array() - a.mean();
49 RowVectorXd bc = b.array() - b.mean();
50
51 double normA = ac.norm();
52 double normB = bc.norm();
53 if (normA < 1e-30 || normB < 1e-30)
54 return 0.0;
55
56 return ac.dot(bc) / (normA * normB);
57}
58
59//=============================================================================================================
60
61double BadChannelDetect::median(QVector<double> values)
62{
63 if (values.isEmpty())
64 return 0.0;
65 std::sort(values.begin(), values.end());
66 const int n = values.size();
67 if (n % 2 == 1)
68 return values[n / 2];
69 return 0.5 * (values[n / 2 - 1] + values[n / 2]);
70}
71
72//=============================================================================================================
73// PUBLIC
74//=============================================================================================================
75
76QVector<int> BadChannelDetect::detect(const MatrixXd& matData,
77 const Params& params)
78{
79 QVector<int> flat = detectFlat(matData, params.dFlatThreshold);
80 QVector<int> noisy = detectHighVariance(matData, params.dVarZThresh);
81 QVector<int> low = detectLowCorrelation(matData, params.dCorrThresh, params.iNeighbours);
82
83 // Union without duplicates
84 QVector<bool> flagged(static_cast<int>(matData.rows()), false);
85 for (int i : flat)
86 if (i >= 0 && i < matData.rows())
87 flagged[i] = true;
88 for (int i : noisy)
89 if (i >= 0 && i < matData.rows())
90 flagged[i] = true;
91 for (int i : low)
92 if (i >= 0 && i < matData.rows())
93 flagged[i] = true;
94
95 QVector<int> result;
96 for (int i = 0; i < static_cast<int>(matData.rows()); ++i) {
97 if (flagged[i])
98 result.append(i);
99 }
100 return result;
101}
102
103//=============================================================================================================
104
105QVector<int> BadChannelDetect::detectFlat(const MatrixXd& matData,
106 double dThreshold)
107{
108 QVector<int> bad;
109 for (int ch = 0; ch < matData.rows(); ++ch) {
110 const RowVectorXd row = matData.row(ch);
111 double ptp = row.maxCoeff() - row.minCoeff();
112 if (ptp < dThreshold)
113 bad.append(ch);
114 }
115 return bad;
116}
117
118//=============================================================================================================
119
120QVector<int> BadChannelDetect::detectHighVariance(const MatrixXd& matData,
121 double dZThresh)
122{
123 const int nCh = static_cast<int>(matData.rows());
124 if (nCh < 3)
125 return {}; // Cannot compute meaningful statistics with < 3 channels
126
127 // Compute per-channel standard deviation
128 QVector<double> stds(nCh);
129 for (int ch = 0; ch < nCh; ++ch) {
130 const RowVectorXd row = matData.row(ch);
131 double mean = row.mean();
132 double var = (row.array() - mean).square().mean();
133 stds[ch] = std::sqrt(var);
134 }
135
136 // Median and MAD of std across channels
137 double med = median(stds);
138
139 QVector<double> absDevs(nCh);
140 for (int ch = 0; ch < nCh; ++ch)
141 absDevs[ch] = std::abs(stds[ch] - med);
142 double mad = median(absDevs);
143
144 // Consistency factor: MAD -> sigma for Gaussian = 1.4826
145 double sigma = 1.4826 * mad;
146
147 // If the robust scale collapses because most channels are identical and only
148 // a few are outliers, fall back to the classical std-dev across channel stds.
149 if (sigma < 1e-30) {
150 double meanStd = 0.0;
151 for (double value : stds) {
152 meanStd += value;
153 }
154 meanStd /= static_cast<double>(nCh);
155
156 double varStd = 0.0;
157 for (double value : stds) {
158 const double diff = value - meanStd;
159 varStd += diff * diff;
160 }
161 sigma = std::sqrt(varStd / static_cast<double>(nCh));
162 }
163
164 if (sigma < 1e-30)
165 return {}; // All channels identical (e.g. all zeros)
166
167 QVector<int> bad;
168 for (int ch = 0; ch < nCh; ++ch) {
169 double z = (stds[ch] - med) / sigma;
170 if (z > dZThresh)
171 bad.append(ch);
172 }
173 return bad;
174}
175
176//=============================================================================================================
177
178QVector<int> BadChannelDetect::detectLowCorrelation(const MatrixXd& matData,
179 double dCorrThresh,
180 int iNeighbours)
181{
182 const int nCh = static_cast<int>(matData.rows());
183 if (nCh < 2)
184 return {};
185
186 QVector<int> bad;
187
188 for (int ch = 0; ch < nCh; ++ch) {
189 int nValid = 0;
190 double sumC = 0.0;
191
192 int lo = std::max(0, ch - iNeighbours);
193 int hi = std::min(nCh - 1, ch + iNeighbours);
194
195 for (int nb = lo; nb <= hi; ++nb) {
196 if (nb == ch)
197 continue;
198 double c = std::abs(pearsonCorr(matData.row(ch), matData.row(nb)));
199 sumC += c;
200 ++nValid;
201 }
202
203 if (nValid == 0)
204 continue;
205
206 double meanCorr = sumC / static_cast<double>(nValid);
207 if (meanCorr < dCorrThresh)
208 bad.append(ch);
209 }
210
211 return bad;
212}
Declaration of BadChannelDetect — automated detection of bad MEG/EEG channels.
Shared utilities (I/O helpers, spectral analysis, layout management, warp algorithms).
BadChannelDetectParams Params
static QVector< int > detect(const Eigen::MatrixXd &matData, const Params &params=Params())
static QVector< int > detectHighVariance(const Eigen::MatrixXd &matData, double dZThresh=4.0)
static QVector< int > detectLowCorrelation(const Eigen::MatrixXd &matData, double dCorrThresh=0.4, int iNeighbours=5)
static QVector< int > detectFlat(const Eigen::MatrixXd &matData, double dThreshold=1e-13)