v2.0.0
Loading...
Searching...
No Matches
inv_cmne.cpp
Go to the documentation of this file.
1//=============================================================================================================
20
21//=============================================================================================================
22// INCLUDES
23//=============================================================================================================
24
25#include "inv_cmne.h"
26
27//=============================================================================================================
28// EIGEN INCLUDES
29//=============================================================================================================
30
31#include <Eigen/Eigenvalues>
32
33//=============================================================================================================
34// QT INCLUDES
35//=============================================================================================================
36
37#include <QDebug>
38#include <QCoreApplication>
39#include <QDir>
40#include <QJsonDocument>
41#include <QJsonObject>
42
43#include <algorithm>
44#include <limits>
45
46//=============================================================================================================
47// MNE-CPP INCLUDES
48//=============================================================================================================
49
50#include <ml/ml_onnx_model.h>
51#include <ml/ml_tensor.h>
52
53#ifndef WASMBUILD
54#include <ml/ml_trainer.h>
55#endif
56
57//=============================================================================================================
58// USED NAMESPACES
59//=============================================================================================================
60
61using namespace INVLIB;
62using namespace Eigen;
63
64//=============================================================================================================
65// DEFINE MEMBER METHODS
66//=============================================================================================================
67
69 const MatrixXd& matEvoked,
70 const MatrixXd& matGain,
71 const MatrixXd& matNoiseCov,
72 const MatrixXd& matSrcCov,
73 const InvCMNESettings& settings)
74{
75 InvCMNEResult result;
76 result.matKernelDspm = computeDspmKernel(matGain, matNoiseCov, matSrcCov, settings.lambda2);
77 const MatrixXd matDspm = result.matKernelDspm * matEvoked;
78 const VectorXi vertices = VectorXi::LinSpaced(matDspm.rows(), 0, static_cast<int>(matDspm.rows()) - 1);
79 result.stcDspm = InvSourceEstimate(matDspm, vertices, 0.0f, 1.0f);
80
81 MatrixXd sensing, prediction, cmne;
82 if (settings.onnxModelPath.isEmpty()) {
83 qInfo() << "[InvCMNE] No model: control estimate with look-back" << settings.lookBack;
84 cmne = controlEstimate(matDspm, settings.lookBack);
85 sensing = zScoreRectify(matDspm);
86 prediction = sensing;
87 } else if (!applyCmne(matDspm, settings.onnxModelPath, sensing, prediction, cmne)) {
88 qWarning() << "[InvCMNE] CMNE could not be applied; returning the dSPM estimate only.";
89 return result;
90 }
91 result.stcSensing = InvSourceEstimate(sensing, vertices, 0.0f, 1.0f);
92 result.stcLstmPredict = InvSourceEstimate(prediction, vertices, 0.0f, 1.0f);
93 result.stcCmne = InvSourceEstimate(cmne, vertices, 0.0f, 1.0f);
94 return result;
95}
96
97//=============================================================================================================
98
99MatrixXd InvCMNE::computeDspmKernel(
100 const MatrixXd& matGain,
101 const MatrixXd& matNoiseCov,
102 const MatrixXd& matSrcCov,
103 double lambda2)
104{
105 int nChannels = matGain.rows();
106 int nSources = matGain.cols();
107
108 // Step 1: Whiten noise covariance via eigendecomposition
109 // C_n = V * D * V^T -> C_n^{-1/2} = V * D^{-1/2} * V^T
110 qInfo() << " [dSPM kernel] Eigendecomposition of noise covariance"
111 << "(" << nChannels << "x" << nChannels << ") …";
112 // Decompose the correlation matrix D^-1/2 C D^-1/2: MEG (T^2) and EEG (V^2) variances differ by 1e13,
113 // so a cutoff relative to the largest eigenvalue of C would drop every MEG component.
114 const VectorXd scale = matNoiseCov.diagonal().cwiseMax(std::numeric_limits<double>::min()).cwiseSqrt();
115 const MatrixXd correlation = scale.cwiseInverse().asDiagonal() * matNoiseCov * scale.cwiseInverse().asDiagonal();
116 SelfAdjointEigenSolver<MatrixXd> eigSolver(correlation);
117 VectorXd eigVals = eigSolver.eigenvalues();
118 MatrixXd eigVecs = eigSolver.eigenvectors();
119
120 // Regularize: clamp small eigenvalues
121 double maxEig = eigVals.maxCoeff();
122 double threshold = maxEig * 1e-10;
123 VectorXd eigValsInvSqrt(nChannels);
124 for (int i = 0; i < nChannels; ++i) {
125 eigValsInvSqrt(i) = (eigVals(i) > threshold) ? 1.0 / std::sqrt(eigVals(i)) : 0.0;
126 }
127
128 // W = (R^+)^1/2 D^-1/2 satisfies W C W^T = I on the retained subspace.
129 MatrixXd matWhitener = eigVecs * eigValsInvSqrt.asDiagonal() * eigVecs.transpose() * scale.cwiseInverse().asDiagonal();
130
131 // Step 2: Whiten gain matrix
132 qInfo() << " [dSPM kernel] Whitening gain matrix …";
133 MatrixXd matGainWhitened = matWhitener * matGain; // n_channels x n_sources
134
135 // Step 3: MNE kernel
136 qInfo() << " [dSPM kernel] Computing MNE kernel (LDLT solve," << nChannels << "x" << nChannels << ") …";
137 // K = C_R * G_tilde^T * (G_tilde * C_R * G_tilde^T + lambda2 * I)^{-1}
138 MatrixXd matGCR = matGainWhitened * matSrcCov; // n_channels x n_sources
139 MatrixXd matA = matGCR * matGainWhitened.transpose(); // n_channels x n_channels
140 matA.diagonal().array() += lambda2;
141
142 // Solve once: A^{-1} via LDLT, then K = (C_R * G_tilde^T) * A^{-1}
143 auto ldlt = matA.ldlt();
144 MatrixXd matK = (matSrcCov * matGainWhitened.transpose()) * ldlt.solve(matWhitener);
145
146 // Step 4: dSPM normalization
147 // noise_norm_i = sqrt((K * C_n * K^T)(i,i))
148 // K_dSPM(i,:) = K(i,:) / noise_norm_i
149 qInfo() << " [dSPM kernel] Normalizing" << nSources << "source rows …";
150 MatrixXd matKCn = matK * matNoiseCov; // n_sources x n_channels
151 for (int i = 0; i < nSources; ++i) {
152 double noiseNorm = std::sqrt(matKCn.row(i).dot(matK.row(i)));
153 if (noiseNorm > 1e-10) {
154 matK.row(i) /= noiseNorm;
155 }
156 }
157
158 return matK; // n_sources x n_channels (dSPM kernel)
159}
160
161//=============================================================================================================
162
163MatrixXd InvCMNE::standardize(const MatrixXd& matStcData)
164{
165 // Constant rows are centred but not scaled, as in cmne.standardize.
166 const VectorXd mean = matStcData.rowwise().mean();
167 MatrixXd result = matStcData.colwise() - mean;
168 const VectorXd std = (result.array().square().rowwise().sum() / static_cast<double>(result.cols())).sqrt();
169 for (int i = 0; i < result.rows(); ++i) {
170 if (std(i) > 0.0)
171 result.row(i) /= std(i);
172 }
173 return result;
174}
175
176//=============================================================================================================
177
178MatrixXd InvCMNE::zScoreRectify(const MatrixXd& matStcData)
179{
180 return standardize(matStcData.cwiseAbs()); // Eq. 9
181}
182
183//=============================================================================================================
184
185MatrixXd InvCMNE::controlEstimate(const MatrixXd& matDspmData, int lookBack)
186{
187 // Paper's control: q_t times the mean of the previous k sensing estimates (no LSTM, not recursive).
188 const MatrixXd q = zScoreRectify(matDspmData);
189 MatrixXd result = q;
190 for (int t = lookBack; t < q.cols(); ++t)
191 result.col(t) = q.col(t).cwiseProduct(q.middleCols(t - lookBack, lookBack).rowwise().mean());
192 return result;
193}
194
195//=============================================================================================================
196
197bool InvCMNE::applyCmne(const MatrixXd& matDspmData,
198 const QString& onnxModelPath,
199 MatrixXd& sensing,
200 MatrixXd& prediction,
201 MatrixXd& cmne)
202{
203 MLLIB::MlOnnxModel model;
204 if (!model.load(onnxModelPath)) {
205 qWarning() << "[InvCMNE] Cannot load CMNE model" << onnxModelPath;
206 return false;
207 }
208 const QJsonObject config = QJsonDocument::fromJson(model.metadata(QStringLiteral("cmne_config")).toUtf8()).object();
209 const int k = config.value(QStringLiteral("look_back")).toInt();
210 const int nSources = config.value(QStringLiteral("n_sources")).toInt();
211 if (k <= 0 || nSources != matDspmData.rows()) {
212 qWarning() << "[InvCMNE] Model" << onnxModelPath << "has no cmne_config for" << matDspmData.rows()
213 << "sources (look_back" << k << ", n_sources" << nSources << ").";
214 return false;
215 }
216 if (matDspmData.cols() <= k) {
217 qWarning() << "[InvCMNE] Need more than look_back =" << k << "samples, got" << matDspmData.cols();
218 return false;
219 }
220 sensing = config.value(QStringLiteral("rectify")).toBool(true) ? zScoreRectify(matDspmData) : standardize(matDspmData);
221
222 // The network sees float32 time-major windows of the contextual estimate b (Eq. 13).
223 const MatrixXf q = sensing.cast<float>();
224 MatrixXf b = q;
225 MatrixXf pred = q;
226 std::vector<float> window(static_cast<size_t>(k) * static_cast<size_t>(nSources));
227 for (int t = k; t < q.cols(); ++t) {
228 Map<MatrixXf>(window.data(), nSources, k) = b.middleCols(t - k, k);
229 const MLLIB::MlTensor out = model.predict(MLLIB::MlTensor::view(window.data(), {1, k, nSources}));
230 const Map<const VectorXf> p(out.data(), nSources);
231 pred.col(t) = p;
232 const VectorXf w = p.cwiseAbs();
233 const float maxW = std::max(w.maxCoeff(), std::numeric_limits<float>::min());
234 b.col(t) = (w / maxW).cwiseProduct(q.col(t)); // Eqs. 10-11
235 }
236 prediction = pred.cast<double>();
237 cmne = b.cast<double>();
238 return true;
239}
240
241//=============================================================================================================
242
243#ifndef WASMBUILD
244
246 const QString& fwdPath,
247 const QString& covPath,
248 const QString& epochsPath,
249 const QString& outOnnxPath,
250 const InvCMNESettings& settings,
251 const QString& gtStcPrefix,
252 int hiddenSize,
253 int numLayers,
254 int trainEpochs,
255 double learningRate,
256 int batchSize,
257 const QString& finetuneOnnxPath,
258 const QString& pythonExe)
259{
260 // Resolve training package directory (contains pyproject.toml + script)
261 // Expected layout: <app_dir>/../scripts/ml/training/cmne/
262 QString appDir = QCoreApplication::applicationDirPath();
263 QString cmneDir = QDir(appDir).absoluteFilePath(
264 QStringLiteral("../scripts/ml/training/cmne"));
265
266 // Fallback: source tree relative to working directory
267 if (!QFile::exists(QDir(cmneDir).absoluteFilePath(QStringLiteral("pyproject.toml")))) {
268 cmneDir = QStringLiteral("scripts/ml/training/cmne");
269 }
270
271 QString scriptPath = QDir(cmneDir).absoluteFilePath(QStringLiteral("train_cmne_lstm.py"));
272
273 if (!QFile::exists(scriptPath)) {
275 result.stdErr = QStringLiteral("Training script not found: ") + scriptPath;
276 qWarning() << "[InvCMNE::trainLstm]" << result.stdErr;
277 return result;
278 }
279
280 qDebug() << "[InvCMNE::trainLstm] Script:" << scriptPath;
281 qDebug() << "[InvCMNE::trainLstm] Package dir:" << cmneDir;
282
283 // Map method integer to string
284 QString methodStr;
285 switch (settings.method) {
286 case 0:
287 methodStr = QStringLiteral("MNE");
288 break;
289 case 1:
290 methodStr = QStringLiteral("dSPM");
291 break;
292 case 2:
293 methodStr = QStringLiteral("sLORETA");
294 break;
295 case 3:
296 methodStr = QStringLiteral("eLORETA");
297 break;
298 default:
299 methodStr = QStringLiteral("dSPM");
300 break;
301 }
302
303 double snr = 1.0 / std::sqrt(settings.lambda2);
304
305 // Build argument list matching train_cmne_lstm.py CLI
306 QStringList args;
307 args << QStringLiteral("--fwd") << fwdPath
308 << QStringLiteral("--cov") << covPath
309 << QStringLiteral("--epochs") << epochsPath
310 << QStringLiteral("--out") << outOnnxPath
311 << QStringLiteral("--look-back") << QString::number(settings.lookBack)
312 << QStringLiteral("--method") << methodStr
313 << QStringLiteral("--snr") << QString::number(snr, 'g', 6)
314 << QStringLiteral("--hidden") << QString::number(hiddenSize)
315 << QStringLiteral("--layers") << QString::number(numLayers)
316 << QStringLiteral("--train-epochs") << QString::number(trainEpochs)
317 << QStringLiteral("--lr") << QString::number(learningRate, 'g', 6)
318 << QStringLiteral("--batch") << QString::number(batchSize);
319
320 if (!gtStcPrefix.isEmpty()) {
321 args << QStringLiteral("--gt-stc") << gtStcPrefix;
322 }
323
324 if (!finetuneOnnxPath.isEmpty()) {
325 args << QStringLiteral("--finetune") << finetuneOnnxPath;
326 }
327
328 // Configure PythonRunner with venv + pyproject.toml
329 // Venv lives inside the cmne package directory as .venv/
331 config.pythonExe = pythonExe;
332 config.venvDir = QDir(cmneDir).absoluteFilePath(QStringLiteral(".venv"));
333 config.packageDir = cmneDir;
334
335 MLLIB::MLTrainer trainer(config);
336
337 return trainer.run(scriptPath, args);
338}
339
340#endif // !WASMBUILD
ONNX Runtime backed MLLIB::MlModel implementation for loading and evaluating .onnx graphs.
N-dimensional, row-major, reference-counted float32 tensor used as the universal MLLIB data carrier.
MLLIB::MLTrainer convenience wrapper that drives Python training scripts via UTILSLIB::PythonRunner.
Contextual Minimum-Norm Estimate (CMNE) inverse solver — deep-learning-corrected dSPM (Dinh et al....
Inverse source estimation (MNE, dSPM, sLORETA, dipole fitting).
Source-space inverse-solution container with dense grid plus optional focal-dipole,...
CMNE result.
Definition inv_cmne.h:72
InvSourceEstimate stcSensing
Definition inv_cmne.h:74
InvSourceEstimate stcDspm
Definition inv_cmne.h:73
InvSourceEstimate stcCmne
Definition inv_cmne.h:75
Eigen::MatrixXd matKernelDspm
Definition inv_cmne.h:77
InvSourceEstimate stcLstmPredict
Definition inv_cmne.h:76
static InvCMNEResult compute(const Eigen::MatrixXd &matEvoked, const Eigen::MatrixXd &matGain, const Eigen::MatrixXd &matNoiseCov, const Eigen::MatrixXd &matSrcCov, const InvCMNESettings &settings)
Definition inv_cmne.cpp:68
static Eigen::MatrixXd zScoreRectify(const Eigen::MatrixXd &matStcData)
Definition inv_cmne.cpp:178
static Eigen::MatrixXd controlEstimate(const Eigen::MatrixXd &matDspmData, int lookBack)
Definition inv_cmne.cpp:185
static bool applyCmne(const Eigen::MatrixXd &matDspmData, const QString &onnxModelPath, Eigen::MatrixXd &sensing, Eigen::MatrixXd &prediction, Eigen::MatrixXd &cmne)
Definition inv_cmne.cpp:197
static UTILSLIB::PythonRunnerResult trainLstm(const QString &fwdPath, const QString &covPath, const QString &epochsPath, const QString &outOnnxPath, const InvCMNESettings &settings, const QString &gtStcPrefix={}, int hiddenSize=256, int numLayers=1, int trainEpochs=50, double learningRate=1e-3, int batchSize=64, const QString &finetuneOnnxPath={}, const QString &pythonExe=QStringLiteral("python3"))
Definition inv_cmne.cpp:245
MlModel backend that runs .onnx graphs through ONNX Runtime with a cached CPU session.
QString metadata(const QString &key) const
bool load(const QString &path) override
MlTensor predict(const MlTensor &input) const override
N-dimensional row-major float32 tensor with shared-buffer storage, Eigen Map accessors and a non-owni...
Definition ml_tensor.h:76
static MlTensor view(float *data, std::vector< int64_t > shape)
Launches Python training scripts via UTILSLIB::PythonRunner with automatic venv handling and prerequi...
Definition ml_trainer.h:71
UTILSLIB::PythonRunnerResult run(const QString &scriptPath, const QStringList &args={})
Script execution result container.
Script execution configuration.