Compare roots of quadratic functions

Viewed 192

I need function to fast compare root of quadratic function and a given value and function to fast compare two roots of two quadratic functions.

I write first function

bool isRootLessThanValue (bool sqrtDeltaSign, int a, int b, int c, int value) {
    bool ret;

    if(sqrtDeltaSign){
        if(a < 0){
            ret = (2*a*value + b < 0) || (a*value*value + b*value + c > 0);
        }else{
            ret = (2*a*value + b > 0) && (a*value*value + b*value + c > 0);
        }
    }else{
        if(a < 0){
            ret = (2*a*value + b < 0) && (a*value*value + b*value + c < 0);
        }else{
            ret = (2*a*value + b > 0) || (a*value*value + b*value + c < 0);
        }
    }
    return ret;
};

When i try to write this for second function it grow to very big and complicated...

bool isRoot1LessThanRoot2 (bool sqrtDeltaSign1, int a1, int b1, int c1, bool sqrtDeltaSign2, int a2, int b2, int c2) {
    //...
}

Have u any suggestions how can i simplify this function?

If you think thats stupid idea for optimizations, please tell me why :)

4 Answers

I give a simpified version of the first part of your code by comparing the greater root of the quadratic function with a given value as follows:

#include <iostream>
#include <cmath> // for main testing
int isRootLessThanValue (int a, int b, int c, int value)
  {
    if (a<0){ b *= -1;  c *= -1;   a *= -1;}
    int xt, delta;
    xt = 2 * a * value + b;
    if (xt < 0) return false; // value is left to reflection point
    delta = b*b - 4*a*c;
 // compare square distance between value and the root
    return  ( (xt * xt) > delta )? true: false;
  }

In the test main() program, the roots are first calculate for clarity purpose:

int main()
{
    int a, b, c, v;

    a = -2;
    b = 4;
    c = 3;
    double r1, r2, r, dt;
    dt  = std::sqrt(b*b-4.0*a*c);
    r1 = (-b + dt) / (2.0*a);
    r2 = (-b - dt) / (2.0*a);
    r = (r1>r2)? r1 : r2;
    while (1)
    {
       std::cout << "Input the try value = ";
       std::cin >> v;
       if (isRootLessThanValue(a,b,c,v)) std::cout << v <<" > " << r << std::endl;
       else std::cout << v <<" < " << r  << std::endl;
    }
    return 0;
 }

A test run

enter image description here

We are definitely talking micro-optimization here, but consider making calculations before performing the comparison:

bool isRootLessThanValue (bool sqrtDeltaSign, int a, int b, int c, int value)
{
    const int a_value = a * value;
    const int two_a_b_value = 2 * a_value + b;
    const int a_squared_b = a_value * value + b * value + c;
    const bool two_ab_less_zero = (two_a_b_value < 0);
    bool ret = false;
    if(sqrtDeltaSign)
    {
        const bool a_squared_b_greater_zero = (a_squared_b > 0);
        if (a < 0)
        {
            ret = two_ab_less_zero || a_squared_b_greater_zero;
        }
        else
        {
            ret = !two_ab_less_zero && a_squared_b_greater_zero;//(edited)
        }
    }
    else
    {
        const bool a_squared_b_less_zero = (a_squared_b < 0);
        if (a < 0)
        {
            ret = two_ab_less_zero && a_squared_b_less_zero;
        }
        else
        {
            ret = !two_ab_less_zero || a_squared_b_less_zero;//(edited)
        }
    }
    return ret;
};  

Another note, is that the boolean expression is calculated and stored in a variable, thus could be counted as a data processing instruction (depending on the compiler and processor).

Compare the assembly language of this function to yours. Also benchmark. As I said, I'm not expecting much time savings here, but I don't know how many times this function is called in your code.

The following assumes that both quadratics have real, mutually distinct roots, and a1 = a2 = 1. This keeps the notations simpler, though similar logic can be used in the general case.

Suppose f(x) = x^2 + b1 x + c1 has the real roots u1 < u2, and g(x) = x^2 + b2 x + c2 has the real roots v1 < v2. Then there are 6 possible sort orders.

  • (1)   u1 < u2 < v1 < v2
  • (2)   u1 < v1 < u2 < v2
  • (3)   u1 < v1 < v2 < u2
  • (4)   v1 < u1 < u2 < v2
  • (5)   v1 < u1 < v2 < u2
  • (6)   v1 < v2 < u1 < u2

Let v be a root of g so that g(v) = v^2 + b2 v + c2 = 0 then v^2 = -b2 v - c2 and therefore f(v) = (b1 - b2) v + c1 - c2 = b12 v + c12 where b12 = b1 - b2 and c12 = c1 - c2.

It follows that Sf = f(v1) + f(v2) = b12(v1 + v2) + 2 c12 and Pf = f(v1) f(v2) = b12^2 v1 v2 + b12 c12 (v1 + v2) + c12^2. Using Vieta's relations v1 v2 = c2 and v1 + v2 = -b2 so in the end Sf = f(v1) + f(v2) = -b12 b2 + 2 c2 and Pf = f(v1) f(v2) = b12^2 c2 - b12 c12 b2 + c12^2. Similar expressions can be calculated for Sg = g(u1) + g(u2) and Pg = g(u1) g(u2).

(Should be noted that Sf, Pf, Sg, Pg above are arithmetic expressions in the coefficients, not involving sqrt square roots. There is, however, the potential for integer overflow. If that is an actual concern, then the calculations would have to be done in floating point instead of integers.)

  • If Pf = f(v1) f(v2) < 0 then exactly one root of f is between the roots v1, v2 of g.

    • If the axis of f is to the left of the g one, meaning -b1 < -b2, then that's the smaller root u1 of f which is between v1, v2 i.e. case (5).
    • Otherwise if -b1 > -b2 then that's the larger root i.e. case (2).
  • If Pf = f(v1) f(v2) > 0 then either both or none of the roots of f are between the roots of g. In this case f(v1) and f(v2) must have the same sign, and they will either be both negative if Sf = f(v1) + f(v2) < 0 or both positive if Sf > 0.

    • If f(v1) < 0 and f(v2) < 0 then both roots v1, v2 of g are between the roots of f i.e. case (3).
    • By symmetry, if Pg > 0 and Sg < 0 then g(u1) < 0 and g(u2) < 0, so both roots u1, u2 of f are between the roots of g i.e. case (4).
    • Otherwise the last combination left is f(v1), f(v2) > 0 and g(u1), g(u2) > 0 where the intervals (u1, u2) and (v1, v2) do not overlap. If -b1 < -b2 the axis of f is to the left of the g one i.e. case (1) else it's case (6).

Once the sort order between all roots is determined, comparing any particular pair of roots follows.

Im reorganising my code and have found some facilities :)

When calculate a, b and c i can keep structure to get only a > 0 :) and i know that i want small or big root :)

so function to compare root to value is regresed to the form below

bool isRootMinLessThanValue (int a, int b, int c, int value) {
    const int a_value = a * value;
    const int u = 2*a_value + b;
    const int v = a_value*value + b*value + c;
    
    return u > 0 || v < 0 ;
};
bool isRootMaxLessThanValue (int a, int b, int c, int value) {
    const int a_value = a*value;
    const int u = 2*a_value + b;
    const int v = a_value*value + b*value + c;
    
    return u > 0 && v > 0;
}

when im testing benchmark its faster than calculate roots traditionaly (by assumptions I cannot say how much)

Below code for fast (and slow traditionaly) compare root to value without assumptions

bool isRootLessThanValue (bool sqrtDeltaSign, int a, int b, int c, int value) {
    const int a_value = a*value;
    const int u = 2*a_value + b;
    const int v = a_value*value + b*value + c;
    const bool s = sqrtDeltaSign;
    
    return      (  a < 0  &&  s && u < 0 ) || 
                (  a < 0  &&  s && v > 0 ) || 
                (  a < 0  && !s && u < 0 && v < 0) ||
                (!(a < 0) && !s && u > 0 ) ||
                (!(a < 0) && !s && v < 0 ) ||
                (!(a < 0) &&  s && u > 0 && v > 0);
};
bool isRootLessThanValueTraditional (bool sqrtDeltaSign, int a, int b, int c, int value) {
    double delta = b*b - 4.0*a*c;
    double calculatedRoot = sqrtDeltaSign ? (-b + sqrt(delta))/(2.0*a) : (-b - sqrt(delta))/(2.0*a);
    return calculatedRoot < value;
};

benchmark results below:

  • isRootLessThanValue (optimized): 10000000000 compares in 152.922s
  • isRootLessThanValueTraditional : 10000000000 compares in 196.168s

Any suggestions how can i simplify even more isRootLessThanValue function? :)

I will try to prepare function to compare two roots of different equations

edited::2020-11-30

bool isRootLessThanValue (bool sqrtDeltaSign, int a, int b, int c, int value) {
const int a_value = a*value;
const int u = 2*a_value + b;
const int v = a_value*value + b*value + c;

return  sqrtDeltaSign ? 
    (( a < 0 && (u < 0 || v > 0) ) || (u > 0 && v > 0)) :
    (( a > 0 && (u > 0 || v < 0) ) || (u < 0 && v < 0));

};

Related