31#include <Eigen/Eigenvalues>
38#include <QCoreApplication>
40#include <QJsonDocument>
69 const MatrixXd& matEvoked,
70 const MatrixXd& matGain,
71 const MatrixXd& matNoiseCov,
72 const MatrixXd& matSrcCov,
78 const VectorXi vertices = VectorXi::LinSpaced(matDspm.rows(), 0,
static_cast<int>(matDspm.rows()) - 1);
81 MatrixXd sensing, prediction, cmne;
83 qInfo() <<
"[InvCMNE] No model: control estimate with look-back" << settings.
lookBack;
88 qWarning() <<
"[InvCMNE] CMNE could not be applied; returning the dSPM estimate only.";
99MatrixXd InvCMNE::computeDspmKernel(
100 const MatrixXd& matGain,
101 const MatrixXd& matNoiseCov,
102 const MatrixXd& matSrcCov,
105 int nChannels = matGain.rows();
106 int nSources = matGain.cols();
110 qInfo() <<
" [dSPM kernel] Eigendecomposition of noise covariance"
111 <<
"(" << nChannels <<
"x" << nChannels <<
") …";
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();
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;
129 MatrixXd matWhitener = eigVecs * eigValsInvSqrt.asDiagonal() * eigVecs.transpose() * scale.cwiseInverse().asDiagonal();
132 qInfo() <<
" [dSPM kernel] Whitening gain matrix …";
133 MatrixXd matGainWhitened = matWhitener * matGain;
136 qInfo() <<
" [dSPM kernel] Computing MNE kernel (LDLT solve," << nChannels <<
"x" << nChannels <<
") …";
138 MatrixXd matGCR = matGainWhitened * matSrcCov;
139 MatrixXd matA = matGCR * matGainWhitened.transpose();
140 matA.diagonal().array() += lambda2;
143 auto ldlt = matA.ldlt();
144 MatrixXd matK = (matSrcCov * matGainWhitened.transpose()) * ldlt.solve(matWhitener);
149 qInfo() <<
" [dSPM kernel] Normalizing" << nSources <<
"source rows …";
150 MatrixXd matKCn = matK * matNoiseCov;
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;
163MatrixXd InvCMNE::standardize(
const MatrixXd& matStcData)
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) {
171 result.row(i) /= std(i);
180 return standardize(matStcData.cwiseAbs());
190 for (
int t = lookBack; t < q.cols(); ++t)
191 result.col(t) = q.col(t).cwiseProduct(q.middleCols(t - lookBack, lookBack).rowwise().mean());
198 const QString& onnxModelPath,
200 MatrixXd& prediction,
204 if (!model.
load(onnxModelPath)) {
205 qWarning() <<
"[InvCMNE] Cannot load CMNE model" << onnxModelPath;
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 <<
").";
216 if (matDspmData.cols() <= k) {
217 qWarning() <<
"[InvCMNE] Need more than look_back =" << k <<
"samples, got" << matDspmData.cols();
220 sensing = config.value(QStringLiteral(
"rectify")).toBool(
true) ?
zScoreRectify(matDspmData) : standardize(matDspmData);
223 const MatrixXf q = sensing.cast<
float>();
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);
230 const Map<const VectorXf> p(out.
data(), nSources);
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));
236 prediction = pred.cast<
double>();
237 cmne = b.cast<
double>();
246 const QString& fwdPath,
247 const QString& covPath,
248 const QString& epochsPath,
249 const QString& outOnnxPath,
251 const QString& gtStcPrefix,
257 const QString& finetuneOnnxPath,
258 const QString& pythonExe)
262 QString appDir = QCoreApplication::applicationDirPath();
263 QString cmneDir = QDir(appDir).absoluteFilePath(
264 QStringLiteral(
"../scripts/ml/training/cmne"));
267 if (!QFile::exists(QDir(cmneDir).absoluteFilePath(QStringLiteral(
"pyproject.toml")))) {
268 cmneDir = QStringLiteral(
"scripts/ml/training/cmne");
271 QString scriptPath = QDir(cmneDir).absoluteFilePath(QStringLiteral(
"train_cmne_lstm.py"));
273 if (!QFile::exists(scriptPath)) {
275 result.
stdErr = QStringLiteral(
"Training script not found: ") + scriptPath;
276 qWarning() <<
"[InvCMNE::trainLstm]" << result.
stdErr;
280 qDebug() <<
"[InvCMNE::trainLstm] Script:" << scriptPath;
281 qDebug() <<
"[InvCMNE::trainLstm] Package dir:" << cmneDir;
285 switch (settings.
method) {
287 methodStr = QStringLiteral(
"MNE");
290 methodStr = QStringLiteral(
"dSPM");
293 methodStr = QStringLiteral(
"sLORETA");
296 methodStr = QStringLiteral(
"eLORETA");
299 methodStr = QStringLiteral(
"dSPM");
303 double snr = 1.0 / std::sqrt(settings.
lambda2);
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);
320 if (!gtStcPrefix.isEmpty()) {
321 args << QStringLiteral(
"--gt-stc") << gtStcPrefix;
324 if (!finetuneOnnxPath.isEmpty()) {
325 args << QStringLiteral(
"--finetune") << finetuneOnnxPath;
332 config.
venvDir = QDir(cmneDir).absoluteFilePath(QStringLiteral(
".venv"));
337 return trainer.
run(scriptPath, args);
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,...
InvSourceEstimate stcSensing
InvSourceEstimate stcDspm
InvSourceEstimate stcCmne
Eigen::MatrixXd matKernelDspm
InvSourceEstimate stcLstmPredict
static InvCMNEResult compute(const Eigen::MatrixXd &matEvoked, const Eigen::MatrixXd &matGain, const Eigen::MatrixXd &matNoiseCov, const Eigen::MatrixXd &matSrcCov, const InvCMNESettings &settings)
static Eigen::MatrixXd zScoreRectify(const Eigen::MatrixXd &matStcData)
static Eigen::MatrixXd controlEstimate(const Eigen::MatrixXd &matDspmData, int lookBack)
static bool applyCmne(const Eigen::MatrixXd &matDspmData, const QString &onnxModelPath, Eigen::MatrixXd &sensing, Eigen::MatrixXd &prediction, Eigen::MatrixXd &cmne)
static UTILSLIB::PythonRunnerResult trainLstm(const QString &fwdPath, const QString &covPath, const QString &epochsPath, const QString &outOnnxPath, const InvCMNESettings &settings, const QString >StcPrefix={}, int hiddenSize=256, int numLayers=1, int trainEpochs=50, double learningRate=1e-3, int batchSize=64, const QString &finetuneOnnxPath={}, const QString &pythonExe=QStringLiteral("python3"))
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...
static MlTensor view(float *data, std::vector< int64_t > shape)
Launches Python training scripts via UTILSLIB::PythonRunner with automatic venv handling and prerequi...
UTILSLIB::PythonRunnerResult run(const QString &scriptPath, const QStringList &args={})
Script execution result container.
Script execution configuration.