Minimial conversion of quartic_solver.c to C++

This commit is contained in:
Paul Romano 2021-12-21 14:47:48 -05:00
parent 377ee77486
commit abe2247eb6
3 changed files with 107 additions and 112 deletions

View file

@ -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)

View file

@ -3,8 +3,6 @@
#include <complex>
extern "C" {
void oqs_quartic_solver(double coeff[5], std::complex<double> roots[4]);
}
#endif // OPENMC_EXTERNAL_QUARTIC_SOLVER_H

View file

@ -1,16 +1,8 @@
#include <complex.h>
#include <float.h>
#include <math.h>
#include <signal.h>
#include <stdio.h>
#include <stdlib.h>
#include <string.h>
#include <time.h>
#include <unistd.h>
#define Sqr(x) ((x) * (x))
#ifndef CMPLX
#define CMPLX(x, y) (x) + (y)*I
#endif
#include <cmath>
#include <complex>
#include <cstdlib>
#include <iostream>
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<double> aq, std::complex<double> bq, std::complex<double> cq,
std::complex<double> 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<double> 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<double>(zmax, 0.0);
roots[1] = std::complex<double>(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<double>(-a / 2, sqrtd / 2);
roots[1] = std::complex<double>(-a / 2, -sqrtd / 2);
}
}
void oqs_quartic_solver(double coeff[5], complex double roots[4])
void oqs_quartic_solver(double coeff[5], std::complex<double> 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<double> 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<double>(l1, gamma);
bcx = std::complex<double>(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<double>(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;