Problem 721: High Powers of Irrational Numbers

View on Project Euler

Project Euler Problem 721 Solution

EulerSolve provides an optimized solution for Project Euler Problem 721, High Powers of Irrational Numbers, with C++, Python, Java, and a step-by-step mathematical explanation.

Problem Summary Define $$T(a,n)=\left\lfloor\left(\lceil\sqrt{a}\rceil+\sqrt{a}\right)^n\right\rfloor,\qquad M=999999937.$$ The task is to evaluate $$S(N)=\sum_{a=1}^{N} T(a,a^2)\pmod{M}$$ for \(N=5{,}000{,}000\). The exponent is already quadratic in \(a\), so direct floating-point evaluation is completely impractical; the solution instead converts the irrational power into an exact integer recurrence that can be computed modulo \(M\). Mathematical Approach The core observation is that the expression becomes easy once we pair it with its algebraic conjugate. Step 1: Introduce the Conjugate Pair For a fixed integer \(a\), let $$m=\lceil\sqrt{a}\rceil,\qquad \alpha=m+\sqrt{a},\qquad \beta=m-\sqrt{a}.$$ Then \(\alpha\) and \(\beta\) are conjugates, and they satisfy $$\alpha+\beta=2m,\qquad \alpha\beta=m^2-a.$$ These two symmetric quantities are integers, which is what makes the later recurrence integral....

Detailed mathematical approach

Problem Summary

Define

$$T(a,n)=\left\lfloor\left(\lceil\sqrt{a}\rceil+\sqrt{a}\right)^n\right\rfloor,\qquad M=999999937.$$

The task is to evaluate

$$S(N)=\sum_{a=1}^{N} T(a,a^2)\pmod{M}$$

for \(N=5{,}000{,}000\). The exponent is already quadratic in \(a\), so direct floating-point evaluation is completely impractical; the solution instead converts the irrational power into an exact integer recurrence that can be computed modulo \(M\).

Mathematical Approach

The core observation is that the expression becomes easy once we pair it with its algebraic conjugate.

Step 1: Introduce the Conjugate Pair

For a fixed integer \(a\), let

$$m=\lceil\sqrt{a}\rceil,\qquad \alpha=m+\sqrt{a},\qquad \beta=m-\sqrt{a}.$$

Then \(\alpha\) and \(\beta\) are conjugates, and they satisfy

$$\alpha+\beta=2m,\qquad \alpha\beta=m^2-a.$$

These two symmetric quantities are integers, which is what makes the later recurrence integral.

Step 2: Replace the Floor by an Exact Integer Formula

If \(a\) is not a perfect square, then \(m-1<\sqrt{a}<m\), so

$$0<\beta=m-\sqrt{a}<1.$$

For every positive \(n\), this gives

$$0<\beta^n<1.$$

Now define

$$U_n=\alpha^n+\beta^n.$$

Since \(U_n-\alpha^n=\beta^n\) lies strictly between \(0\) and \(1\), we obtain

$$\left\lfloor\alpha^n\right\rfloor=U_n-1\qquad\text{for non-square }a.$$

If \(a\) is a perfect square, then \(m=\sqrt{a}\) and \(\beta=0\), so instead

$$\left\lfloor\alpha^n\right\rfloor=U_n=(2m)^n.$$

Therefore the entire problem reduces to computing \(U_n\) efficiently.

Step 3: Derive an Integer Recurrence for \(U_n\)

The numbers \(\alpha\) and \(\beta\) are the two roots of

$$x^2-2mx+(m^2-a)=0.$$

Therefore the sequence \(U_n=\alpha^n+\beta^n\) satisfies the standard second-order linear recurrence

$$U_0=2,\qquad U_1=2m,$$

$$U_n=2m\,U_{n-1}-(m^2-a)\,U_{n-2}\qquad(n\ge 2).$$

Because the coefficients and initial values are integers, every \(U_n\) is an integer. This is the exact quantity hidden behind the irrational-looking power.

Step 4: Compute \(U_n\) by Fast Exponentiation in a Quadratic Ring

Instead of iterating the recurrence all the way up to \(n=a^2\), the implementation exponentiates the base element \(m+\sqrt{a}\) directly. Represent

$$x+y\sqrt{a}$$

by the pair \((x,y)\). Then multiplication becomes

$$\left(x_1,y_1\right)\left(x_2,y_2\right)=\left(x_1x_2+a y_1y_2,\ x_1y_2+x_2y_1\right).$$

If

$$\alpha^n=X_n+Y_n\sqrt{a},$$

then its conjugate is

$$\beta^n=X_n-Y_n\sqrt{a},$$

so

$$U_n=\alpha^n+\beta^n=2X_n.$$

Binary exponentiation computes \(X_n\) in \(O(\log n)\) ring multiplications, which is essential because \(n=a^2\) can be enormous.

Step 5: Assemble the Final Summation

For each \(a\), the contribution is

$$T(a,a^2)=\begin{cases} U_{a^2}-1, & \text{if } a \text{ is not a perfect square},\\ U_{a^2}, & \text{if } a \text{ is a perfect square}. \end{cases}$$

Hence

$$S(N)=\sum_{a=1}^{N} T(a,a^2)\pmod{M}.$$

The only case distinction is whether \(a\) is square; the same ring exponentiation handles both branches.

Worked Example: \(a=5\)

Here \(m=\lceil\sqrt{5}\rceil=3\), so

$$\alpha=3+\sqrt{5},\qquad \beta=3-\sqrt{5},\qquad \alpha\beta=4.$$

The recurrence becomes

$$U_0=2,\qquad U_1=6,\qquad U_n=6U_{n-1}-4U_{n-2}.$$

Then

$$U_2=6\cdot 6-4\cdot 2=28,$$

so

$$T(5,2)=\left\lfloor(3+\sqrt{5})^2\right\rfloor=28-1=27.$$

Continuing,

$$U_3=144,\qquad U_4=752,\qquad U_5=3936,$$

and since \(0<\beta<1\),

$$T(5,5)=3936-1=3935.$$

This is exactly the kind of identity the implementations use, but computed modulo \(M\) and at much larger exponents.

How the Code Works

The C++, Python, and Java implementations all follow the same plan. They scan \(a\) over a contiguous range, keep track of the current value of \(\lceil\sqrt{a}\rceil\), and update it only when \(a\) crosses a perfect square. For each \(a\), they set \(n=a^2\), exponentiate \(m+\sqrt{a}\) by binary exponentiation in the pair representation above, double the real component to recover \(U_n\), subtract \(1\) when \(a\) is not square, and add the result into a running modular sum.

To speed up the full computation, the interval \([1,N]\) is partitioned into independent chunks. Each worker computes its own partial sum modulo \(M\), and the partial results are combined at the end. The mathematics is identical in all three languages; only the concurrency mechanism differs.

Complexity Analysis

For one value of \(a\), the cost is \(O(\log(a^2))=O(\log a)\) pair multiplications. Summing over all \(1\le a\le N\) gives

$$\sum_{a=1}^{N} O(\log a)=O(N\log N).$$

The working memory inside one worker is \(O(1)\). With parallel execution, total auxiliary memory is linear in the number of workers, but still constant per worker range.

Footnotes and References

  1. Problem page: https://projecteuler.net/problem=721
  2. Lucas sequence: Wikipedia — Lucas sequence
  3. Binary exponentiation: Wikipedia — Exponentiation by squaring
  4. Integer square root: Wikipedia — Integer square root
  5. Quadratic field: Wikipedia — Quadratic field

Problem 721 source code

C++

#include <cassert>
#include <cstdint>
#include <iostream>
#include <pthread.h>
#include <unistd.h>
#include <vector>

namespace {

using u64 = std::uint64_t;
using u128 = unsigned __int128;

constexpr u64 kMod = 999'999'937ULL;

struct Pair {
    u64 x;
    u64 y;
};

u64 isqrt_u64(const u64 n) {
    u64 r = static_cast<u64>(__builtin_sqrtl(static_cast<long double>(n)));
    while ((r + 1ULL) <= n / (r + 1ULL)) {
        ++r;
    }
    while (r > n / r) {
        --r;
    }
    return r;
}

u64 ceil_sqrt_u64(const u64 n) {
    const u64 r = isqrt_u64(n);
    return (r * r == n) ? r : (r + 1ULL);
}

Pair mul_pair(const Pair a, const Pair b, const u64 a_mod) {
    const u64 t1 = (a.x * b.x) % kMod;
    const u64 t2 = ((a.y * b.y) % kMod * a_mod) % kMod;
    const u64 real = (t1 + t2) % kMod;
    const u64 imag = (a.x * b.y + a.y * b.x) % kMod;
    return {real, imag};
}

u64 lucas_sum_mod(const u64 a, const u64 m, u64 n) {
    const u64 a_mod = a % kMod;
    Pair base{m % kMod, 1ULL};
    Pair result{1ULL, 0ULL};

    while (n > 0ULL) {
        if (n & 1ULL) {
            result = mul_pair(result, base, a_mod);
        }
        n >>= 1ULL;
        if (n > 0ULL) {
            base = mul_pair(base, base, a_mod);
        }
    }

    return (2ULL * result.x) % kMod;
}

u64 f_mod(const u64 a, const u64 n) {
    const u64 m = ceil_sqrt_u64(a);
    const bool is_square = (m * m == a);
    u64 value = lucas_sum_mod(a, m, n);
    if (!is_square) {
        value = (value + kMod - 1ULL) % kMod;
    }
    return value;
}

u128 f_exact_small(const u64 a, const u64 n) {
    const u64 m = ceil_sqrt_u64(a);
    const u64 d = m * m - a;

    if (d == 0ULL) {
        u128 p = 1;
        const u128 base = static_cast<u128>(2ULL * m);
        for (u64 i = 0ULL; i < n; ++i) {
            p *= base;
        }
        return p;
    }

    u128 s0 = 2;
    u128 s1 = static_cast<u128>(2ULL * m);
    if (n == 0ULL) {
        return 1;
    }
    if (n == 1ULL) {
        return s1 - 1;
    }

    for (u64 i = 2ULL; i <= n; ++i) {
        const u128 s = static_cast<u128>(2ULL * m) * s1 - static_cast<u128>(d) * s0;
        s0 = s1;
        s1 = s;
    }
    return s1 - 1;
}

u64 G_mod(const int limit) {
    auto sum_range = [](u64 l, u64 r) -> u64 {
        if (l > r) {
            return 0ULL;
        }
        u64 sum = 0ULL;
        u64 m = isqrt_u64(l);
        if (m * m < l) {
            ++m;
        }
        u64 sq = m * m;

        for (u64 a = l; a <= r; ++a) {
            while (sq < a) {
                ++m;
                sq = m * m;
            }
            const bool is_square = (sq == a);
            u64 value = lucas_sum_mod(a, m, a * a);
            if (!is_square) {
                value = (value + kMod - 1ULL) % kMod;
            }

            sum += value;
            if (sum >= kMod) {
                sum -= kMod;
            }
        }
        return sum;
    };

    long cpu_count = ::sysconf(_SC_NPROCESSORS_ONLN);
    int thread_count = (cpu_count > 1) ? static_cast<int>(cpu_count) : 1;
    if (thread_count > 16) {
        thread_count = 16;
    }
    if (thread_count > limit) {
        thread_count = limit;
    }
    if (limit < 200'000) {
        thread_count = 1;
    }

    struct Task {
        u64 l = 0;
        u64 r = 0;
        u64 partial = 0;
    };
    auto worker = [](void* raw) -> void* {
        auto* t = static_cast<Task*>(raw);
        t->partial = 0ULL;
        if (t->l <= t->r) {
            u64 sum = 0ULL;
            u64 m = isqrt_u64(t->l);
            if (m * m < t->l) {
                ++m;
            }
            u64 sq = m * m;
            for (u64 a = t->l; a <= t->r; ++a) {
                while (sq < a) {
                    ++m;
                    sq = m * m;
                }
                const bool is_square = (sq == a);
                u64 value = lucas_sum_mod(a, m, a * a);
                if (!is_square) {
                    value = (value + kMod - 1ULL) % kMod;
                }
                sum += value;
                if (sum >= kMod) {
                    sum -= kMod;
                }
            }
            t->partial = sum;
        }
        return nullptr;
    };

    if (thread_count <= 1) {
        return sum_range(1ULL, static_cast<u64>(limit));
    }

    std::vector<pthread_t> tids(static_cast<std::size_t>(thread_count));
    std::vector<Task> tasks(static_cast<std::size_t>(thread_count));
    const u64 total = static_cast<u64>(limit);
    const u64 base = total / static_cast<u64>(thread_count);
    const u64 rem = total % static_cast<u64>(thread_count);
    u64 cur = 1ULL;
    for (int t = 0; t < thread_count; ++t) {
        const u64 len = base + (static_cast<u64>(t) < rem ? 1ULL : 0ULL);
        tasks[static_cast<std::size_t>(t)].l = cur;
        tasks[static_cast<std::size_t>(t)].r = (len == 0ULL ? 0ULL : cur + len - 1ULL);
        cur += len;
        const int rc = ::pthread_create(&tids[static_cast<std::size_t>(t)], nullptr, worker, &tasks[static_cast<std::size_t>(t)]);
        assert(rc == 0);
    }

    u64 sum = 0ULL;
    for (int t = 0; t < thread_count; ++t) {
        const int rc = ::pthread_join(tids[static_cast<std::size_t>(t)], nullptr);
        assert(rc == 0);
        sum += tasks[static_cast<std::size_t>(t)].partial;
        if (sum >= kMod) {
            sum -= kMod;
        }
    }
    return sum;
}

}  // namespace

int main() {
    assert(f_exact_small(5, 2) == 27);
    assert(f_exact_small(5, 5) == 3935);
    assert(f_mod(5, 2) == 27 % kMod);
    assert(f_mod(5, 5) == 3935 % kMod);
    assert(G_mod(1000) == 163'861'845ULL);

    std::cout << G_mod(5'000'000) << '\n';
    return 0;
}

Python

import math
import multiprocessing

def isqrt(n):
    r = int(math.isqrt(n))
    while (r + 1) * (r + 1) <= n:
        r += 1
    while r > 0 and r * r > n:
        r -= 1
    return r

def worker_func(l, r):
    kMod = 999999937
    partial_sum = 0
    
    m = isqrt(l)
    if m * m < l:
        m += 1
    sq = m * m
    
    for a in range(l, r + 1):
        while sq < a:
            m += 1
            sq = m * m
            
        is_square = (sq == a)
        n = a * a
        a_mod = a % kMod
        
        base_x = m % kMod
        base_y = 1
        res_x = 1
        res_y = 0
        
        while n > 0:
            if n & 1:
                nx = (res_x * base_x + res_y * base_y * a_mod) % kMod
                ny = (res_x * base_y + res_y * base_x) % kMod
                res_x, res_y = nx, ny
            n >>= 1
            if n > 0:
                nx = (base_x * base_x + base_y * base_y * a_mod) % kMod
                ny = (2 * base_x * base_y) % kMod
                base_x, base_y = nx, ny
                
        val = (2 * res_x) % kMod
        if not is_square:
            val = (val + kMod - 1) % kMod
            
        partial_sum = (partial_sum + val) % kMod
        
    return partial_sum

def solve():
    kMod = 999999937
    limit = 5000000
    
    threads = max(1, multiprocessing.cpu_count())
    chunk = limit // threads
    rem = limit % threads
    
    ranges = []
    curr = 1
    for t in range(threads):
        length = chunk + (1 if t < rem else 0)
        if length > 0:
            ranges.append((curr, curr + length - 1))
        curr += length
        
    if threads <= 1:
        total = worker_func(1, limit)
    else:
        with multiprocessing.Pool(threads) as pool:
            results = pool.starmap(worker_func, ranges)
            total = sum(results) % kMod
            
    return str(total)

if __name__ == "__main__":
    print(solve())

Java

import java.util.ArrayList;
import java.util.List;
import java.util.concurrent.ExecutorService;
import java.util.concurrent.Executors;
import java.util.concurrent.Future;
import java.util.concurrent.Callable;

public class Euler721 {
    static final long kMod = 999999937L;

    static long isqrt(long n) {
        long r = (long) Math.sqrt(n);
        while ((r + 1) * (r + 1) <= n)
            r++;
        while (r > 0 && r * r > n)
            r--;
        return r;
    }

    static class Worker implements Callable<Long> {
        long l, r;

        Worker(long l, long r) {
            this.l = l;
            this.r = r;
        }

        @Override
        public Long call() {
            long partialSum = 0;
            long m = isqrt(l);
            if (m * m < l)
                m++;
            long sq = m * m;

            for (long a = l; a <= r; ++a) {
                while (sq < a) {
                    m++;
                    sq = m * m;
                }
                boolean isSquare = (sq == a);
                long n = a * a;
                long aMod = a % kMod;

                long baseX = m % kMod;
                long baseY = 1;
                long resX = 1;
                long resY = 0;

                while (n > 0) {
                    if ((n & 1) != 0) {
                        long nx = (resX * baseX % kMod + resY * baseY % kMod * aMod % kMod) % kMod;
                        long ny = (resX * baseY % kMod + resY * baseX % kMod) % kMod;
                        resX = nx;
                        resY = ny;
                    }
                    n >>= 1;
                    if (n > 0) {
                        long nx = (baseX * baseX % kMod + baseY * baseY % kMod * aMod % kMod) % kMod;
                        long ny = (2 * baseX * baseY) % kMod;
                        baseX = nx;
                        baseY = ny;
                    }
                }

                long val = (2 * resX) % kMod;
                if (!isSquare) {
                    val = (val + kMod - 1) % kMod;
                }

                partialSum = (partialSum + val) % kMod;
            }
            return partialSum;
        }
    }

    public static String solve() {
        long limit = 5000000;
        int threads = Runtime.getRuntime().availableProcessors();
        if (threads < 1)
            threads = 1;

        long chunk = limit / threads;
        long rem = limit % threads;

        ExecutorService executor = Executors.newFixedThreadPool(threads);
        List<Future<Long>> futures = new ArrayList<>();

        long curr = 1;
        for (int t = 0; t < threads; ++t) {
            long length = chunk + (t < rem ? 1 : 0);
            if (length > 0) {
                futures.add(executor.submit(new Worker(curr, curr + length - 1)));
            }
            curr += length;
        }

        long total = 0;
        try {
            for (Future<Long> f : futures) {
                total = (total + f.get()) % kMod;
            }
        } catch (Exception e) {
            e.printStackTrace();
        }
        executor.shutdown();

        return Long.toString(total);
    }

    public static void main(String[] args) {
        System.out.println(solve());
    }
}