73 "pfqn_kt requires transcendental arithmetic (steepest-descent expansion of log G)");
81 for (
const T& v : N0) Nt0 += v;
82 if (L0.
empty() || N0.empty() || Nt0 == zero) {
87 const std::size_t Rorig = L0.
cols(), Morig = L0.
rows();
88 std::vector<T> Zin = Z0;
89 if (Zin.empty()) Zin.assign(Rorig, zero);
90 if (Zin.size() != Rorig || N0.size() != Rorig)
91 throw InputError(
"pfqn_kt: L, N and Z disagree on the class count");
96 std::size_t nkeep = 0;
97 for (std::size_t r = 0; r < Rorig; ++r)
98 if (N0[r] > zero) ++nkeep;
99 if (nkeep > 0 && nkeep < Rorig) {
101 std::vector<T> Nk, Zk;
105 for (std::size_t r = 0; r < Rorig; ++r) {
106 if (!(N0[r] > zero))
continue;
107 for (std::size_t i = 0; i < Morig; ++i) Lk(i, c) = L0(i, r);
109 Zk.push_back(Zin[r]);
116 T slcdemandfactor = zero;
117 std::vector<bool> isslc(Rorig,
false);
118 std::vector<std::vector<T>> rows;
119 for (std::size_t i = 0; i < Morig; ++i) {
120 std::vector<T> row(Rorig);
121 for (std::size_t r = 0; r < Rorig; ++r) row[r] = L0(i, r);
124 std::vector<std::size_t> slcstation(Rorig, 0);
126 for (std::size_t r = 0; r < Rorig; ++r) {
127 std::size_t nnz = 0, ist = 0;
128 for (std::size_t i = 0; i < Morig; ++i)
129 if (L0(i, r) != zero) {
133 if (nnz != 1 || Zin[r] != zero)
continue;
139 std::vector<bool> done(Rorig,
false);
140 for (std::size_t r = 0; r < Rorig; ++r) {
141 if (!isslc[r] || done[r])
continue;
142 const std::size_t ist = slcstation[r];
144 for (std::size_t s = r; s < Rorig; ++s) {
145 if (!isslc[s] || slcstation[s] != ist)
continue;
148 slcdemandfactor += T(N0[s] * log(L0(ist, s))) - detail::num_factln<T>(N0[s]);
150 slcdemandfactor += detail::num_factln<T>(ntot);
152 for (
long k = 0; k < reps; ++k) {
153 std::vector<T> row(Rorig);
154 for (std::size_t s = 0; s < Rorig; ++s) row[s] = L0(ist, s);
159 std::vector<std::size_t> keep;
160 for (std::size_t r = 0; r < Rorig; ++r)
161 if (!isslc[r]) keep.push_back(r);
163 res.
lG = slcdemandfactor;
164 res.
G = exp(slcdemandfactor);
167 const std::size_t M = rows.size(), R = keep.size();
169 for (std::size_t i = 0; i < M; ++i)
170 for (std::size_t r = 0; r < R; ++r) L(i, r) = rows[i][keep[r]];
171 std::vector<T> N(R), Z(R);
172 for (std::size_t r = 0; r < R; ++r) {
177 for (
const T& v : N) Ntot += v;
179 res.
lG = slcdemandfactor;
180 res.
G = exp(slcdemandfactor);
190 std::vector<T> u = amva.
XN;
193 for (std::size_t k = 0; k < M; ++k) {
195 for (std::size_t r = 0; r < R; ++r) s += L(k, r) * u[r];
196 if (s > Umax) Umax = s;
200 for (std::size_t r = 0; r < R; ++r) u[r] = T(u[r] * f);
203 bool converged =
false;
204 for (
int it = 0; it < 200; ++it) {
206 for (std::size_t k = 0; k < M; ++k) {
208 for (std::size_t r = 0; r < R; ++r) s += L(k, r) * u[r];
209 D[k] = T(one / T(one - s));
213 for (std::size_t r = 0; r < R; ++r) {
215 for (std::size_t k = 0; k < M; ++k) s += L(k, r) * D[k];
216 g[r] = T(u[r] * T(Z[r] + s) - N[r]);
225 for (std::size_t r = 0; r < R; ++r) {
227 for (std::size_t k = 0; k < M; ++k) s += L(k, r) * D[k];
228 J(r, r) = T(Z[r] + s);
229 for (std::size_t sIdx = 0; sIdx < R; ++sIdx) {
231 for (std::size_t k = 0; k < M; ++k) acc += L(k, r) * T(D[k] * D[k]) * L(k, sIdx);
232 J(r, sIdx) += u[r] * acc;
235 std::vector<T> rhs(R);
236 for (std::size_t r = 0; r < R; ++r) rhs[r] = T(-g[r]);
237 const std::vector<T> du =
solve(J, rhs);
239 for (
int b = 0; b < 60; ++b) {
241 for (std::size_t r = 0; r < R && ok; ++r)
242 if (T(u[r] + alpha * du[r]) <= zero) ok =
false;
244 for (std::size_t k = 0; k < M && ok; ++k) {
246 for (std::size_t r = 0; r < R; ++r) s += L(k, r) * T(u[r] + alpha * du[r]);
247 if (s >= one) ok =
false;
254 for (std::size_t r = 0; r < R; ++r) u[r] = T(u[r] + alpha * du[r]);
257 std::vector<T> us = converged ? u : amva.
XN;
261 for (std::size_t k = 0; k < M; ++k) {
263 for (std::size_t r = 0; r < R; ++r) s += L(k, r) * u[r];
264 D[k] = T(one / T(one - s));
267 for (std::size_t r = 0; r < R; ++r) {
269 for (std::size_t k = 0; k < M; ++k) s += L(k, r) * D[k];
270 const T gr = T(u[r] * T(Z[r] + s) - N[r]);
278 std::vector<T> Uk(M), D(M);
279 for (std::size_t k = 0; k < M; ++k) {
281 for (std::size_t r = 0; r < R; ++r) s += L(k, r) * us[r];
284 if (den < fineTol) den = fineTol;
288 for (std::size_t r = 0; r < R; ++r) {
289 if (us[r] == zero)
throw NumericError(
"pfqn_kt: zero saddle-point coordinate");
290 H(r, r) = T(N[r] / T(us[r] * us[r]));
291 for (std::size_t s = 0; s < R; ++s) {
293 for (std::size_t k = 0; k < M; ++k) acc += L(k, r) * T(D[k] * D[k]) * L(k, s);
298 for (std::size_t r = 0; r < R; ++r) F += Z[r] * us[r];
299 for (std::size_t k = 0; k < M; ++k) {
300 T den = T(one - Uk[k]);
301 if (den < fineTol) den = fineTol;
304 for (std::size_t r = 0; r < R; ++r) F -= N[r] * log(us[r]);
308 for (std::size_t r = 0; r < R; ++r) lG -= log(us[r]);
311 lG += slcdemandfactor;