From abe2247eb6ad71058ed0affe6bb70242e6475c5b Mon Sep 17 00:00:00 2001 From: Paul Romano Date: Tue, 21 Dec 2021 14:47:48 -0500 Subject: [PATCH] Minimial conversion of quartic_solver.c to C++ --- CMakeLists.txt | 2 +- include/openmc/external/quartic_solver.h | 2 - .../{quartic_solver.c => quartic_solver.cpp} | 215 +++++++++--------- 3 files changed, 107 insertions(+), 112 deletions(-) rename src/external/{quartic_solver.c => quartic_solver.cpp} (73%) diff --git a/CMakeLists.txt b/CMakeLists.txt index 9cfb55b7a..fb5da6f2a 100644 --- a/CMakeLists.txt +++ b/CMakeLists.txt @@ -366,7 +366,7 @@ list(APPEND libopenmc_SOURCES # Add bundled external dependencies list(APPEND libopenmc_SOURCES - src/external/quartic_solver.c) + src/external/quartic_solver.cpp) # For Visual Studio compilers if(MSVC) diff --git a/include/openmc/external/quartic_solver.h b/include/openmc/external/quartic_solver.h index 98159afb0..9438e494a 100644 --- a/include/openmc/external/quartic_solver.h +++ b/include/openmc/external/quartic_solver.h @@ -3,8 +3,6 @@ #include -extern "C" { void oqs_quartic_solver(double coeff[5], std::complex roots[4]); -} #endif // OPENMC_EXTERNAL_QUARTIC_SOLVER_H diff --git a/src/external/quartic_solver.c b/src/external/quartic_solver.cpp similarity index 73% rename from src/external/quartic_solver.c rename to src/external/quartic_solver.cpp index 71b24b8a4..be881ddf3 100644 --- a/src/external/quartic_solver.c +++ b/src/external/quartic_solver.cpp @@ -1,16 +1,8 @@ -#include -#include -#include -#include -#include -#include -#include -#include -#include -#define Sqr(x) ((x) * (x)) -#ifndef CMPLX -#define CMPLX(x, y) (x) + (y)*I -#endif +#include +#include +#include +#include + const double cubic_rescal_fact = 3.488062113727083E+102; //= pow(DBL_MAX,1.0/3.0)/1.618034; const double quart_rescal_fact = @@ -42,35 +34,36 @@ void oqs_solve_cubic_analytic_depressed_handle_inf( R = 0.5 * c; if (R == 0) { if (b <= 0) { - *sol = sqrt(-b); + *sol = std::sqrt(-b); } else { *sol = 0; } return; } - if (fabs(Q) < fabs(R)) { + if (std::fabs(Q) < std::fabs(R)) { QR = Q / R; QRSQ = QR * QR; KK = 1.0 - Q * QRSQ; } else { RQ = R / Q; - KK = copysign(1.0, Q) * (RQ * RQ / Q - 1.0); + KK = std::copysign(1.0, Q) * (RQ * RQ / Q - 1.0); } if (KK < 0.0) { - sqrtQ = sqrt(Q); - theta = acos((R / fabs(Q)) / sqrtQ); + sqrtQ = std::sqrt(Q); + theta = std::acos((R / std::fabs(Q)) / sqrtQ); if (theta < PI2) - *sol = -2.0 * sqrtQ * cos(theta / 3.0); + *sol = -2.0 * sqrtQ * std::cos(theta / 3.0); else - *sol = -2.0 * sqrtQ * cos((theta + TWOPI) / 3.0); + *sol = -2.0 * sqrtQ * std::cos((theta + TWOPI) / 3.0); } else { - if (fabs(Q) < fabs(R)) - A = -copysign(1.0, R) * cbrt(fabs(R) * (1.0 + sqrt(KK))); + if (std::fabs(Q) < std::fabs(R)) + A = -std::copysign(1.0, R) * cbrt(std::fabs(R) * (1.0 + std::sqrt(KK))); else { - A = - -copysign(1.0, R) * cbrt(fabs(R) + sqrt(fabs(Q)) * fabs(Q) * sqrt(KK)); + A = -std::copysign(1.0, R) * + cbrt(std::fabs(R) + + std::sqrt(std::fabs(Q)) * std::fabs(Q) * std::sqrt(KK)); } if (A == 0.0) B = 0.0; @@ -86,21 +79,22 @@ void oqs_solve_cubic_analytic_depressed(double b, double c, double* sol) double Q, R, theta, Q3, R2, A, B, sqrtQ; Q = -b / 3.0; R = 0.5 * c; - if (fabs(Q) > 1E102 || fabs(R) > 1E154) { + if (std::fabs(Q) > 1e102 || std::fabs(R) > 1e154) { oqs_solve_cubic_analytic_depressed_handle_inf(b, c, sol); return; } - Q3 = Sqr(Q) * Q; - R2 = Sqr(R); + Q3 = Q * Q * Q; + R2 = R * R; if (R2 < Q3) { - theta = acos(R / sqrt(Q3)); - sqrtQ = -2.0 * sqrt(Q); + theta = std::acos(R / std::sqrt(Q3)); + sqrtQ = -2.0 * std::sqrt(Q); if (theta < M_PI / 2) - *sol = sqrtQ * cos(theta / 3.0); + *sol = sqrtQ * std::cos(theta / 3.0); else - *sol = sqrtQ * cos((theta + 2.0 * M_PI) / 3.0); + *sol = sqrtQ * std::cos((theta + 2.0 * M_PI) / 3.0); } else { - A = -copysign(1.0, R) * pow(fabs(R) + sqrt(R2 - Q3), 1.0 / 3.0); + A = -std::copysign(1.0, R) * + std::pow(std::fabs(R) + std::sqrt(R2 - Q3), 1.0 / 3.0); if (A == 0.0) B = 0.0; else @@ -120,7 +114,7 @@ void oqs_calc_phi0( diskr = 9 * a * a - 24 * b; /* eq. (87) */ if (diskr > 0.0) { - diskr = sqrt(diskr); + diskr = std::sqrt(diskr); if (a > 0.0) s = -2 * b / (3 * a + diskr); else @@ -139,9 +133,9 @@ void oqs_calc_phi0( g = hh - 4 * dq - 3 * gg; /* eq. (85) */ h = (8 * dq + hh - 2 * gg) * bq / 3 - cq * cq - dq * aq * aq; /* eq. (86) */ oqs_solve_cubic_analytic_depressed(g, h, &rmax); - if (isnan(rmax) || isinf(rmax)) { + if (std::isnan(rmax) || std::isinf(rmax)) { oqs_solve_cubic_analytic_depressed_handle_inf(g, h, &rmax); - if ((isnan(rmax) || isinf(rmax)) && scaled) { + if ((std::isnan(rmax) || std::isinf(rmax)) && scaled) { // try harder: rescale also the depressed cubic if quartic has been // already rescaled rfact = cubic_rescal_fact; @@ -158,7 +152,7 @@ void oqs_calc_phi0( h = (8.0 * dqss + hhss - 2.0 * ggss) * bqs / 3 - cqs * (cqs / rfact) - (dq / rfact) * aqs * aqs; oqs_solve_cubic_analytic_depressed(g, h, &rmax); - if (isnan(rmax) || isinf(rmax)) { + if (std::isnan(rmax) || std::isinf(rmax)) { oqs_solve_cubic_analytic_depressed_handle_inf(g, h, &rmax); } rmax *= rfact; @@ -171,14 +165,14 @@ void oqs_calc_phi0( xxx = x * xsq; gx = g * x; f = x * (xsq + g) + h; - if (fabs(xxx) > fabs(gx)) - maxtt = fabs(xxx); + if (std::fabs(xxx) > std::fabs(gx)) + maxtt = std::fabs(xxx); else - maxtt = fabs(gx); - if (fabs(h) > maxtt) - maxtt = fabs(h); + maxtt = std::fabs(gx); + if (std::fabs(h) > maxtt) + maxtt = std::fabs(h); - if (fabs(f) > macheps * maxtt) { + if (std::fabs(f) > macheps * maxtt) { for (iter = 0; iter < 8; iter++) { df = 3.0 * xsq + g; if (df == 0) { @@ -193,7 +187,7 @@ void oqs_calc_phi0( break; } - if (fabs(f) >= fabs(fold)) { + if (std::fabs(f) >= std::fabs(fold)) { x = xold; break; } @@ -206,26 +200,27 @@ double oqs_calc_err_ldlt( { /* Eqs. (29) and (30) in the manuscript */ double sum; - sum = (b == 0) ? fabs(d2 + l1 * l1 + 2.0 * l3) - : fabs(((d2 + l1 * l1 + 2.0 * l3) - b) / b); - sum += (c == 0) ? fabs(2.0 * d2 * l2 + 2.0 * l1 * l3) - : fabs(((2.0 * d2 * l2 + 2.0 * l1 * l3) - c) / c); - sum += (d == 0) ? fabs(d2 * l2 * l2 + l3 * l3) - : fabs(((d2 * l2 * l2 + l3 * l3) - d) / d); + sum = (b == 0) ? std::fabs(d2 + l1 * l1 + 2.0 * l3) + : std::fabs(((d2 + l1 * l1 + 2.0 * l3) - b) / b); + sum += (c == 0) ? std::fabs(2.0 * d2 * l2 + 2.0 * l1 * l3) + : std::fabs(((2.0 * d2 * l2 + 2.0 * l1 * l3) - c) / c); + sum += (d == 0) ? std::fabs(d2 * l2 * l2 + l3 * l3) + : std::fabs(((d2 * l2 * l2 + l3 * l3) - d) / d); return sum; } double oqs_calc_err_abcd_cmplx(double a, double b, double c, double d, - complex double aq, complex double bq, complex double cq, complex double dq) + std::complex aq, std::complex bq, std::complex cq, + std::complex dq) { /* Eqs. (68) and (69) in the manuscript for complex alpha1 (aq), beta1 (bq), * alpha2 (cq) and beta2 (dq) */ double sum; - sum = (d == 0) ? cabs(bq * dq) : cabs((bq * dq - d) / d); - sum += - (c == 0) ? cabs(bq * cq + aq * dq) : cabs(((bq * cq + aq * dq) - c) / c); - sum += - (b == 0) ? cabs(bq + aq * cq + dq) : cabs(((bq + aq * cq + dq) - b) / b); - sum += (a == 0) ? cabs(aq + cq) : cabs(((aq + cq) - a) / a); + sum = (d == 0) ? std::abs(bq * dq) : std::abs((bq * dq - d) / d); + sum += (c == 0) ? std::abs(bq * cq + aq * dq) + : std::abs(((bq * cq + aq * dq) - c) / c); + sum += (b == 0) ? std::abs(bq + aq * cq + dq) + : std::abs(((bq + aq * cq + dq) - b) / b); + sum += (a == 0) ? std::abs(aq + cq) : std::abs(((aq + cq) - a) / a); return sum; } double oqs_calc_err_abcd(double a, double b, double c, double d, double aq, @@ -234,12 +229,12 @@ double oqs_calc_err_abcd(double a, double b, double c, double d, double aq, /* Eqs. (68) and (69) in the manuscript for real alpha1 (aq), beta1 (bq), * alpha2 (cq) and beta2 (dq)*/ double sum; - sum = (d == 0) ? fabs(bq * dq) : fabs((bq * dq - d) / d); - sum += - (c == 0) ? fabs(bq * cq + aq * dq) : fabs(((bq * cq + aq * dq) - c) / c); - sum += - (b == 0) ? fabs(bq + aq * cq + dq) : fabs(((bq + aq * cq + dq) - b) / b); - sum += (a == 0) ? fabs(aq + cq) : fabs(((aq + cq) - a) / a); + sum = (d == 0) ? std::fabs(bq * dq) : std::fabs((bq * dq - d) / d); + sum += (c == 0) ? std::fabs(bq * cq + aq * dq) + : std::fabs(((bq * cq + aq * dq) - c) / c); + sum += (b == 0) ? std::fabs(bq + aq * cq + dq) + : std::fabs(((bq + aq * cq + dq) - b) / b); + sum += (a == 0) ? std::fabs(aq + cq) : std::fabs(((aq + cq) - a) / a); return sum; } double oqs_calc_err_abc( @@ -247,11 +242,11 @@ double oqs_calc_err_abc( { /* Eqs. (48)-(51) in the manuscript */ double sum; - sum = - (c == 0) ? fabs(bq * cq + aq * dq) : fabs(((bq * cq + aq * dq) - c) / c); - sum += - (b == 0) ? fabs(bq + aq * cq + dq) : fabs(((bq + aq * cq + dq) - b) / b); - sum += (a == 0) ? fabs(aq + cq) : fabs(((aq + cq) - a) / a); + sum = (c == 0) ? std::fabs(bq * cq + aq * dq) + : std::fabs(((bq * cq + aq * dq) - c) / c); + sum += (b == 0) ? std::fabs(bq + aq * cq + dq) + : std::fabs(((bq + aq * cq + dq) - b) / b); + sum += (a == 0) ? std::fabs(aq + cq) : std::fabs(((aq + cq) - a) / a); return sum; } void oqs_NRabcd(double a, double b, double c, double d, double* AQ, double* BQ, @@ -276,7 +271,7 @@ void oqs_NRabcd(double a, double b, double c, double d, double* AQ, double* BQ, fvec[3] = x[0] + x[2] - a; errf = 0; for (k1 = 0; k1 < 4; k1++) { - errf += (vr[k1] == 0) ? fabs(fvec[k1]) : fabs(fvec[k1] / vr[k1]); + errf += (vr[k1] == 0) ? std::fabs(fvec[k1]) : std::fabs(fvec[k1] / vr[k1]); } for (iter = 0; iter < 8; iter++) { x02 = x[0] - x[2]; @@ -318,7 +313,8 @@ void oqs_NRabcd(double a, double b, double c, double d, double* AQ, double* BQ, errfold = errf; errf = 0; for (k1 = 0; k1 < 4; k1++) { - errf += (vr[k1] == 0) ? fabs(fvec[k1]) : fabs(fvec[k1] / vr[k1]); + errf += + (vr[k1] == 0) ? std::fabs(fvec[k1]) : std::fabs(fvec[k1] / vr[k1]); } if (errf == 0) break; @@ -333,15 +329,15 @@ void oqs_NRabcd(double a, double b, double c, double d, double* AQ, double* BQ, *CQ = x[2]; *DQ = x[3]; } -void oqs_solve_quadratic(double a, double b, complex double roots[2]) +void oqs_solve_quadratic(double a, double b, std::complex roots[2]) { double div, sqrtd, diskr, zmax, zmin; diskr = a * a - 4 * b; if (diskr >= 0.0) { if (a >= 0.0) - div = -a - sqrt(diskr); + div = -a - std::sqrt(diskr); else - div = -a + sqrt(diskr); + div = -a + std::sqrt(diskr); zmax = div / 2; @@ -350,15 +346,15 @@ void oqs_solve_quadratic(double a, double b, complex double roots[2]) else zmin = b / zmax; - roots[0] = CMPLX(zmax, 0.0); - roots[1] = CMPLX(zmin, 0.0); + roots[0] = std::complex(zmax, 0.0); + roots[1] = std::complex(zmin, 0.0); } else { - sqrtd = sqrt(-diskr); - roots[0] = CMPLX(-a / 2, sqrtd / 2); - roots[1] = CMPLX(-a / 2, -sqrtd / 2); + sqrtd = std::sqrt(-diskr); + roots[0] = std::complex(-a / 2, sqrtd / 2); + roots[1] = std::complex(-a / 2, -sqrtd / 2); } } -void oqs_quartic_solver(double coeff[5], complex double roots[4]) +void oqs_quartic_solver(double coeff[5], std::complex roots[4]) { /* USAGE: * @@ -371,8 +367,8 @@ void oqs_quartic_solver(double coeff[5], complex double roots[4]) * the four roots will be stored in the complex array roots[] * * */ - complex double acx1, bcx1, ccx1, dcx1, acx, bcx, ccx, dcx, cdiskr, zx1, zx2, - zxmax, zxmin, qroots[2]; + std::complex acx1, bcx1, ccx1, dcx1, acx, bcx, ccx, dcx, cdiskr, zx1, + zx2, zxmax, zxmin, qroots[2]; double l2m[12], d2m[12], res[12], resmin, bl311, dml3l3, err0 = 0, err1 = 0, aq1, bq1, cq1, dq1; double a, b, c, d, phi0, aq, bq, cq, dq, d2, d3, l1, l2, l3, errmin, errv[3], @@ -381,7 +377,7 @@ void oqs_quartic_solver(double coeff[5], complex double roots[4]) double rfactsq, rfact = 1.0; if (coeff[4] == 0.0) { - printf("That's not a quartic!\n"); + std::cout << "That's not a quartic!\n"; return; } a = coeff[3] / coeff[4]; @@ -391,7 +387,7 @@ void oqs_quartic_solver(double coeff[5], complex double roots[4]) oqs_calc_phi0(a, b, c, d, &phi0, 0); // simple polynomial rescaling - if (isnan(phi0) || isinf(phi0)) { + if (std::isnan(phi0) || std::isinf(phi0)) { rfact = quart_rescal_fact; a /= rfact; rfactsq = rfact * rfact; @@ -445,17 +441,17 @@ void oqs_quartic_solver(double coeff[5], complex double roots[4]) whichcase = 0; if (d2 < 0.0) { /* Case I eqs. (37)-(40) */ - gamma = sqrt(-d2); + gamma = std::sqrt(-d2); aq = l1 + gamma; bq = l3 + gamma * l2; cq = l1 - gamma; dq = l3 - gamma * l2; - if (fabs(dq) < fabs(bq)) + if (std::fabs(dq) < std::fabs(bq)) dq = d / bq; - else if (fabs(dq) > fabs(bq)) + else if (std::fabs(dq) > std::fabs(bq)) bq = d / dq; - if (fabs(aq) < fabs(cq)) { + if (std::fabs(aq) < std::fabs(cq)) { nsol = 0; if (dq != 0) { aqv[nsol] = (c - bq * cq) / dq; /* see eqs. (47) */ @@ -507,18 +503,19 @@ void oqs_quartic_solver(double coeff[5], complex double roots[4]) realcase[0] = 1; } else if (d2 > 0) { /* Case II eqs. (53)-(56) */ - gamma = sqrt(d2); - acx = CMPLX(l1, gamma); - bcx = CMPLX(l3, gamma * l2); - ccx = conj(acx); - dcx = conj(bcx); + gamma = std::sqrt(d2); + acx = std::complex(l1, gamma); + bcx = std::complex(l3, gamma * l2); + ccx = std::conj(acx); + dcx = std::conj(bcx); realcase[0] = 0; } else realcase[0] = -1; // d2=0 /* Case III: d2 is 0 or approximately 0 (in this case check which solution is * better) */ - if (realcase[0] == -1 || (fabs(d2) <= macheps * oqs_max3(fabs(2. * b / 3.), - fabs(phi0), l1 * l1))) { + if (realcase[0] == -1 || + (std::fabs(d2) <= + macheps * oqs_max3(std::fabs(2. * b / 3.), std::fabs(phi0), l1 * l1))) { d3 = d - l3 * l3; if (realcase[0] == 1) err0 = oqs_calc_err_abcd(a, b, c, d, aq, bq, cq, dq); @@ -527,21 +524,21 @@ void oqs_quartic_solver(double coeff[5], complex double roots[4]) if (d3 <= 0) { realcase[1] = 1; aq1 = l1; - bq1 = l3 + sqrt(-d3); + bq1 = l3 + std::sqrt(-d3); cq1 = l1; - dq1 = l3 - sqrt(-d3); - if (fabs(dq1) < fabs(bq1)) + dq1 = l3 - std::sqrt(-d3); + if (std::fabs(dq1) < std::fabs(bq1)) dq1 = d / bq1; - else if (fabs(dq1) > fabs(bq1)) + else if (std::fabs(dq1) > std::fabs(bq1)) bq1 = d / dq1; err1 = oqs_calc_err_abcd(a, b, c, d, aq1, bq1, cq1, dq1); /* eq. (68) */ } else /* complex */ { realcase[1] = 0; acx1 = l1; - bcx1 = l3 + I * sqrt(d3); + bcx1 = l3 + std::complex(0., std::sqrt(d3)); ccx1 = l1; - dcx1 = conj(bcx1); + dcx1 = std::conj(bcx1); err1 = oqs_calc_err_abcd_cmplx(a, b, c, d, acx1, bcx1, ccx1, dcx1); } if (realcase[0] == -1 || err1 < err0) { @@ -575,37 +572,37 @@ void oqs_quartic_solver(double coeff[5], complex double roots[4]) /* complex coefficients of p1 and p2 */ if (whichcase == 0) // d2!=0 { - cdiskr = acx * acx / 4 - bcx; + cdiskr = 0.25 * acx * acx - bcx; /* calculate the roots as roots of p1(x) and p2(x) (see end of sec. 2.1) */ - zx1 = -acx / 2 + csqrt(cdiskr); - zx2 = -acx / 2 - csqrt(cdiskr); - if (cabs(zx1) > cabs(zx2)) + zx1 = -0.5 * acx + std::sqrt(cdiskr); + zx2 = -0.5 * acx - std::sqrt(cdiskr); + if (std::abs(zx1) > std::abs(zx2)) zxmax = zx1; else zxmax = zx2; zxmin = bcx / zxmax; roots[0] = zxmin; - roots[1] = conj(zxmin); + roots[1] = std::conj(zxmin); roots[2] = zxmax; - roots[3] = conj(zxmax); + roots[3] = std::conj(zxmax); } else // d2 ~ 0 { /* never gets here! */ - cdiskr = csqrt(acx * acx - 4.0 * bcx); + cdiskr = std::sqrt(acx * acx - 4.0 * bcx); zx1 = -0.5 * (acx + cdiskr); zx2 = -0.5 * (acx - cdiskr); - if (cabs(zx1) > cabs(zx2)) + if (std::abs(zx1) > std::abs(zx2)) zxmax = zx1; else zxmax = zx2; zxmin = bcx / zxmax; roots[0] = zxmax; roots[1] = zxmin; - cdiskr = csqrt(ccx * ccx - 4.0 * dcx); + cdiskr = std::sqrt(ccx * ccx - 4.0 * dcx); zx1 = -0.5 * (ccx + cdiskr); zx2 = -0.5 * (ccx - cdiskr); - if (cabs(zx1) > cabs(zx2)) + if (std::abs(zx1) > std::abs(zx2)) zxmax = zx1; else zxmax = zx2;