48 const MatrixXd& matGain,
49 const MatrixXd& matData,
50 const MatrixXd& matNoiseCov,
53 double gammaThreshold)
55 const int nChannels =
static_cast<int>(matGain.rows());
56 const int nSources =
static_cast<int>(matGain.cols());
57 const int nTimes =
static_cast<int>(matData.cols());
60 MatrixXd matNoiseCovInv = matNoiseCov.ldlt().solve(MatrixXd::Identity(nChannels, nChannels));
63 VectorXd vecGamma = VectorXd::Ones(nSources);
64 VectorXd vecGammaOld = vecGamma;
67 std::vector<int> activeIdx(nSources);
68 std::iota(activeIdx.begin(), activeIdx.end(), 0);
71 MatrixXd matX = MatrixXd::Zero(nSources, nTimes);
73 int actualIterations = 0;
75 for (
int iter = 0; iter < nIterations; ++iter) {
76 actualIterations = iter + 1;
78 const int nActive =
static_cast<int>(activeIdx.size());
83 MatrixXd matG_active(nChannels, nActive);
84 VectorXd vecGamma_active(nActive);
85 for (
int i = 0; i < nActive; ++i) {
86 matG_active.col(i) = matGain.col(activeIdx[i]);
87 vecGamma_active(i) = vecGamma(activeIdx[i]);
92 MatrixXd matCm = matG_active * vecGamma_active.asDiagonal() * matG_active.transpose() + matNoiseCov;
94 const auto ldlt = matCm.ldlt();
95 const MatrixXd matCmInvG = ldlt.solve(matG_active);
96 const MatrixXd matA = matCmInvG.transpose() * matData;
99 MatrixXd matX_active = vecGamma_active.asDiagonal() * matA;
103 for (
int i = 0; i < nActive; ++i) {
104 matX.row(activeIdx[i]) = matX_active.row(i);
108 vecGammaOld = vecGamma;
109 for (
int i = 0; i < nActive; ++i) {
110 const double denom = std::max(matG_active.col(i).dot(matCmInvG.col(i)), std::numeric_limits<double>::epsilon());
111 vecGamma(activeIdx[i]) = vecGamma_active(i) * matA.row(i).squaredNorm() /
static_cast<double>(nTimes) / denom;
115 std::vector<int> newActive;
116 newActive.reserve(nActive);
117 for (
int i = 0; i < nActive; ++i) {
118 int srcIdx = activeIdx[i];
119 if (vecGamma(srcIdx) >= gammaThreshold) {
120 newActive.push_back(srcIdx);
122 vecGamma(srcIdx) = 0.0;
125 activeIdx = newActive;
128 double maxGammaOld = std::max(vecGammaOld.cwiseAbs().maxCoeff(), 1e-10);
129 double maxRelChange = 0.0;
130 for (
int idx : activeIdx) {
131 double relChange = std::abs(vecGamma(idx) - vecGammaOld(idx)) / maxGammaOld;
132 maxRelChange = std::max(maxRelChange, relChange);
134 if (maxRelChange < tolerance)
144 QVector<int> finalActive;
145 for (
int i = 0; i < nSources; ++i) {
146 if (vecGamma(i) >= gammaThreshold) {
147 finalActive.append(i);
153 const int nActiveFinal = finalActive.size();
154 MatrixXd matActiveSol(nActiveFinal, nTimes);
155 VectorXi vecActiveVerts(nActiveFinal);
156 for (
int i = 0; i < nActiveFinal; ++i) {
157 matActiveSol.row(i) = matX.row(finalActive[i]);
158 vecActiveVerts(i) = finalActive[i];
165 MatrixXd matResidual = matData - matGain * matX;