Problem 558: Irrational Base
View on Project EulerProject Euler Problem 558 Solution
EulerSolve provides an optimized solution for Project Euler Problem 558, Irrational Base, with C++, Python, Java, and a step-by-step mathematical explanation.
Problem Summary Let \(r\) be the irrational base used in the problem. For each positive integer \(n\), we consider its greedy representation as a sum of distinct powers of \(r\), and we write \(\ell(n)\) for the number of selected powers. The quantity to compute is $$S(m)=\sum_{j=1}^{m}\ell(j^2),\qquad m=5{,}000{,}000.$$ The challenge is that the base is irrational, so the implementation must make millions of precise comparisons and subtractions without letting floating-point error corrupt the greedy choices. Mathematical Approach The implementations rely on one algebraic identity for the base and one numerical device for comparing powers quickly. Step 1: Identify the Irrational Base The base is $$r=\frac{1+\sqrt[3]{\frac{29-3\sqrt{93}}{2}}+\sqrt[3]{\frac{29+3\sqrt{93}}{2}}}{3}\approx 1.465571231876768.$$ This is the unique real root of $$x^3-x^2-1=0,$$ so it satisfies $$r^3=r^2+1.$$ Multiplying by \(r^{e-2}\) gives the recurrence $$r^{e+1}=r^e+r^{e-2}.$$ That recurrence explains the combinatorial shape of the greedy expansion. Step 2: Why the Exponents Are Spaced by at Least Three Suppose the greedy process has chosen the largest admissible power \(r^e\le n\). If the residual is \(R=n-r^e\), then \(n<r^{e+1}\), hence $$0\le R<r^{e+1}-r^e=r^{e-2}.$$ Therefore the next selected power cannot be \(r^{e-1}\) or \(r^{e-2}\); its exponent must be at most \(e-3\)....
Detailed mathematical approach
Problem Summary
Let \(r\) be the irrational base used in the problem. For each positive integer \(n\), we consider its greedy representation as a sum of distinct powers of \(r\), and we write \(\ell(n)\) for the number of selected powers. The quantity to compute is
$$S(m)=\sum_{j=1}^{m}\ell(j^2),\qquad m=5{,}000{,}000.$$
The challenge is that the base is irrational, so the implementation must make millions of precise comparisons and subtractions without letting floating-point error corrupt the greedy choices.
Mathematical Approach
The implementations rely on one algebraic identity for the base and one numerical device for comparing powers quickly.
Step 1: Identify the Irrational Base
The base is
$$r=\frac{1+\sqrt[3]{\frac{29-3\sqrt{93}}{2}}+\sqrt[3]{\frac{29+3\sqrt{93}}{2}}}{3}\approx 1.465571231876768.$$
This is the unique real root of
$$x^3-x^2-1=0,$$
so it satisfies
$$r^3=r^2+1.$$
Multiplying by \(r^{e-2}\) gives the recurrence
$$r^{e+1}=r^e+r^{e-2}.$$
That recurrence explains the combinatorial shape of the greedy expansion.
Step 2: Why the Exponents Are Spaced by at Least Three
Suppose the greedy process has chosen the largest admissible power \(r^e\le n\). If the residual is \(R=n-r^e\), then \(n<r^{e+1}\), hence
$$0\le R<r^{e+1}-r^e=r^{e-2}.$$
Therefore the next selected power cannot be \(r^{e-1}\) or \(r^{e-2}\); its exponent must be at most \(e-3\). Repeating the same argument at every step yields the greedy form used by the implementations:
$$n=\sum_{t=1}^{\ell(n)} r^{e_t},\qquad e_1>e_2>\cdots,\qquad e_{t+1}\le e_t-3.$$
So after one term is chosen, the next search can start three exponents lower and continue downward only if the residual is still smaller than that first candidate.
Step 3: Worked Example
Take \(n=4\). Since
$$r^3\approx 3.147899<4<r^4\approx 4.613470,$$
the first greedy term is \(r^3\). The remaining amount is about \(0.852101\), so the next admissible term is \(r^{-1}\approx 0.682328\). Continuing in the same way gives
$$4=r^3+r^{-1}+r^{-5}+r^{-10}.$$
Numerically,
$$3.147899+0.682328+0.147899+0.021874\approx 4.$$
The decimal values only illustrate the greedy choices; the equality itself is exact because all terms are powers of the algebraic number \(r\). Hence \(\ell(4)=4\).
Step 4: Encode Powers in a KaTeX-Friendly Fixed-Point Form
Directly comparing irrational numbers inside the inner loop would be both slow and fragile. Instead, the implementations precompute each power in the form
$$r^e=a_e+\frac{b_e}{M}+\frac{c_e}{M^2}+O(M^{-3}),\qquad M=10^{17},$$
where \(a_e\), \(b_e\), and \(c_e\) are integers obtained from a high-precision evaluation of \(r^e\). Residuals are stored in the same three-level format. To decide whether a power fits into the current residual, the implementation only needs a lexicographic comparison of the triples
$$\left(a_e,b_e,c_e\right),\qquad \left(R_1,R_2,R_3\right).$$
After each subtraction, borrow propagation normalizes the three components exactly as in arithmetic with base \(M\).
Step 5: Turn Individual Lengths into the Final Sum
For every square \(j^2\), the algorithm finds the largest precomputed power below it, repeatedly subtracts the largest admissible lower power, and counts how many terms were used. That count is \(\ell(j^2)\), so the final answer is simply
$$S(m)=\sum_{j=1}^{m}\ell(j^2).$$
The implementations verify two checkpoints before the full run:
$$S(10)=61,\qquad S(1000)=19403.$$
These values confirm that the greedy representation and the fixed-point comparisons are aligned.
How the Code Works
The C++ and Java implementations first evaluate the algebraic base with high decimal precision and build a table of powers over a finite exponent window wide enough for the target computation. Each stored power is split into three integer layers so that the inner greedy loop can avoid arbitrary-precision arithmetic.
For each square, the implementation locates the leading power, initializes the residual in the same three-layer format, then repeatedly searches downward for the next admissible term, subtracts it with borrow handling, and increments the representation length. Because the squares increase with \(j\), the search for the leading exponent only moves forward through the table. The Python implementation delegates to the same compiled algorithm, so all three language versions use the same greedy decomposition and the same checkpoints.
Complexity Analysis
Let \(W=E_{\max}-E_{\min}+1\) denote the size of the precomputed exponent window. Building the table costs \(O(W)\) high-precision steps and \(O(W)\) memory. Over the whole sweep \(j=1,\dots,m\), the forward search for the leading exponent is monotone, so it contributes \(O(W+m)\) total work. The dominant part is the greedy subtraction itself, which is proportional to the total number of selected terms:
$$O\left(W+m+\sum_{j=1}^{m}\ell(j^2)\right).$$
Since \(W\) is fixed in the implementations, memory usage is effectively constant with respect to \(m\), and the running time is essentially linear in the total output length of the greedy expansions.
Footnotes and References
- Problem page: https://projecteuler.net/problem=558
- Beta expansion: Wikipedia — Beta expansion
- Algebraic number: Wikipedia — Algebraic number
- Fixed-point arithmetic: Wikipedia — Fixed-point arithmetic
- Greedy algorithm: Wikipedia — Greedy algorithm
Problem 558 source code
C++
#include <boost/multiprecision/cpp_dec_float.hpp>
#include <cstdint>
#include <iostream>
#include <limits>
#include <stdexcept>
#include <string>
#include <vector>
namespace {
using u64 = std::uint64_t;
using i64 = std::int64_t;
using mpf = boost::multiprecision::cpp_dec_float_100;
constexpr int kDefaultM = 5000000;
constexpr int kMinE = -180;
constexpr int kMaxE = 90;
constexpr i64 kM = 100000000000000000LL;
struct Options {
int m = kDefaultM;
bool run_checkpoints = true;
};
bool parse_nonnegative_int(const std::string& text, int& out) {
if (text.empty()) return false;
std::uint64_t value = 0;
for (char ch : text) {
if (ch < '0' || ch > '9') return false;
value = value * 10ULL + static_cast<std::uint64_t>(ch - '0');
if (value > static_cast<std::uint64_t>(std::numeric_limits<int>::max())) {
return false;
}
}
out = static_cast<int>(value);
return true;
}
bool parse_int_after_prefix(const std::string& arg,
const std::string& prefix,
int& value) {
if (arg.rfind(prefix, 0U) != 0U) return false;
return parse_nonnegative_int(arg.substr(prefix.size()), value);
}
bool parse_arguments(int argc, char** argv, Options& options) {
bool seen_positional_m = false;
for (int i = 1; i < argc; ++i) {
std::string arg(argv[i]);
if (arg == "--skip-checkpoints") {
options.run_checkpoints = false;
continue;
}
int parsed = 0;
if (parse_int_after_prefix(arg, "--m=", parsed)) {
options.m = parsed;
continue;
}
if (!seen_positional_m && parse_nonnegative_int(arg, parsed)) {
options.m = parsed;
seen_positional_m = true;
continue;
}
std::cerr << "Unknown argument: " << arg << '\n';
return false;
}
if (options.m <= 0) {
std::cerr << "m must be positive.\n";
return false;
}
return true;
}
mpf cbrt_mpf(const mpf& x) {
if (x == 0) return mpf(0);
mpf y = exp(log(x) / 3);
for (int it = 0; it < 30; ++it) {
y = (2 * y + x / (y * y)) / 3;
}
return y;
}
struct RootTable {
std::vector<i64> u1;
std::vector<i64> u2;
std::vector<i64> u3;
};
RootTable build_root_table() {
RootTable table;
const int len = kMaxE - kMinE + 1;
table.u1.resize(static_cast<std::size_t>(len));
table.u2.resize(static_cast<std::size_t>(len));
table.u3.resize(static_cast<std::size_t>(len));
const mpf p = (mpf(29) - 3 * sqrt(mpf(93))) / 2;
const mpf q = mpf(29) - p;
const mpf r = (mpf(1) + cbrt_mpf(p) + cbrt_mpf(q)) / 3;
mpf z = 1;
for (int i = 0; i < -kMinE; ++i) z /= r;
for (int i = 0; i < len; ++i) {
mpf x = z;
const i64 y1 = x.convert_to<i64>();
table.u1[static_cast<std::size_t>(i)] = y1;
x = mpf(kM) * (x - mpf(y1));
const i64 y2 = x.convert_to<i64>();
table.u2[static_cast<std::size_t>(i)] = y2;
x = mpf(kM) * (x - mpf(y2));
const i64 y3 = x.convert_to<i64>();
table.u3[static_cast<std::size_t>(i)] = y3;
z *= r;
}
return table;
}
inline bool leq_lex(const i64 a1, const i64 a2, const i64 a3,
const i64 b1, const i64 b2, const i64 b3) {
if (a1 != b1) return a1 < b1;
if (a2 != b2) return a2 < b2;
return a3 <= b3;
}
struct SolveResult {
u64 s_m = 0;
u64 s_10 = 0;
u64 s_1000 = 0;
};
SolveResult solve_with_table(const int m, const RootTable& table) {
SolveResult out;
const int len = static_cast<int>(table.u1.size());
int i0 = -kMinE;
for (int j = 1; j <= m; ++j) {
const i64 n = static_cast<i64>(j) * static_cast<i64>(j);
if (n == 1) {
out.s_m += 1;
if (j == 10) out.s_10 = out.s_m;
if (j == 1000) out.s_1000 = out.s_m;
continue;
}
while (i0 < len && table.u1[static_cast<std::size_t>(i0)] < n) ++i0;
if (i0 >= len) {
throw std::runtime_error("Exponent window exceeded; increase [minE,maxE].");
}
int i = i0 - 1;
i64 rep = 1;
i64 n1 = n - table.u1[static_cast<std::size_t>(i)] - 1;
i64 n2 = kM - table.u2[static_cast<std::size_t>(i)] - 1;
i64 n3 = kM - table.u3[static_cast<std::size_t>(i)];
while (true) {
i -= 3;
while (i >= 0 &&
!leq_lex(table.u1[static_cast<std::size_t>(i)],
table.u2[static_cast<std::size_t>(i)],
table.u3[static_cast<std::size_t>(i)],
n1, n2, n3)) {
--i;
}
if (i < 0) {
throw std::runtime_error("Exponent window underflow; decrease minE.");
}
++rep;
n1 -= table.u1[static_cast<std::size_t>(i)];
n2 -= table.u2[static_cast<std::size_t>(i)];
n3 -= table.u3[static_cast<std::size_t>(i)];
if (n3 < 0) {
n3 += kM;
--n2;
}
if (n2 < 0) {
n2 += kM;
--n1;
}
if (n1 == 0 && n2 == 0 && n3 < 1000) {
out.s_m += static_cast<u64>(rep);
break;
}
}
if (j == 10) out.s_10 = out.s_m;
if (j == 1000) out.s_1000 = out.s_m;
}
return out;
}
bool run_validations(const RootTable& table) {
const SolveResult sample = solve_with_table(1000, table);
if (sample.s_10 != 61ULL) {
std::cerr << "Validation failed: S(10) should be 61, got " << sample.s_10 << '\n';
return false;
}
if (sample.s_1000 != 19403ULL) {
std::cerr << "Validation failed: S(1000) should be 19403, got " << sample.s_1000 << '\n';
return false;
}
return true;
}
} // namespace
int main(int argc, char** argv) {
std::ios::sync_with_stdio(false);
std::cin.tie(nullptr);
Options options;
if (!parse_arguments(argc, argv, options)) return 1;
try {
const RootTable table = build_root_table();
if (options.run_checkpoints && !run_validations(table)) {
return 1;
}
const SolveResult result = solve_with_table(options.m, table);
std::cout << result.s_m << '\n';
} catch (const std::exception& ex) {
std::cerr << "Error: " << ex.what() << '\n';
return 1;
}
return 0;
}
Python
from __future__ import annotations
import re
import shutil
import subprocess
from pathlib import Path
ANSWER_RE = re.compile(r"answer\s*:\s*(.+)$", re.IGNORECASE)
EQUAL_RE = re.compile(r"=\s*(.+)$")
def parse_output(stdout: str) -> str:
lines = [line.strip() for line in stdout.splitlines() if line.strip()]
if not lines:
return ""
answers = []
equals = []
for line in lines:
m1 = ANSWER_RE.search(line)
if m1:
answers.append(m1.group(1).strip())
m2 = EQUAL_RE.search(line)
if m2:
equals.append(m2.group(1).strip())
if answers:
return answers[-1]
if equals:
return equals[-1]
return lines[-1]
def should_skip_cpp_checkpoints(src: Path) -> bool:
try:
text = src.read_text(encoding="utf-8", errors="ignore")
except OSError:
return False
return "--skip-checkpoints" in text
def run_cpp(binary: Path, src: Path, root: Path) -> str:
cmd = [str(binary)]
if should_skip_cpp_checkpoints(src):
cmd.append("--skip-checkpoints")
try:
return subprocess.check_output(cmd, text=True, cwd=root)
except subprocess.CalledProcessError:
return subprocess.check_output(cmd, text=True, cwd=src.parent)
def solve() -> str:
problem_id = __file__.split("Euler")[-1].split(".")[0]
root = Path(__file__).resolve().parent.parent
src = root / "solutionsCpp" / f"Euler{problem_id}.cpp"
binary = root / "solutionsCpp" / f".euler{problem_id}_py_bridge"
if not binary.exists() or src.stat().st_mtime > binary.stat().st_mtime:
compiler = shutil.which("clang++") or shutil.which("g++")
if not compiler:
raise RuntimeError("No C++ compiler found (clang++/g++).")
subprocess.check_call([compiler, "-std=c++17", "-O2", str(src), "-o", str(binary)])
output = run_cpp(binary=binary, src=src, root=root)
parsed = parse_output(output)
if not parsed:
raise RuntimeError(f"Euler{problem_id} bridge produced empty output.")
return parsed
if __name__ == "__main__":
print(solve())
Java
import java.math.BigDecimal;
import java.math.MathContext;
import java.math.RoundingMode;
public class Euler558 {
static final int kMinE = -180;
static final int kMaxE = 90;
static final long kM = 100000000000000000L;
static BigDecimal cbrtDec(BigDecimal x, MathContext mc) {
if (x.compareTo(BigDecimal.ZERO) == 0)
return BigDecimal.ZERO;
double approx = Math.pow(x.doubleValue(), 1.0 / 3.0);
BigDecimal r = new BigDecimal(approx, mc);
BigDecimal two = new BigDecimal(2);
BigDecimal three = new BigDecimal(3);
for (int i = 0; i < 30; i++) {
BigDecimal num = two.multiply(r, mc).add(x.divide(r.multiply(r, mc), mc), mc);
r = num.divide(three, mc);
}
return r;
}
static class RootTable {
long[] u1;
long[] u2;
long[] u3;
}
static RootTable buildRootTable() {
RootTable table = new RootTable();
int len = kMaxE - kMinE + 1;
table.u1 = new long[len];
table.u2 = new long[len];
table.u3 = new long[len];
MathContext mc = new MathContext(100, RoundingMode.HALF_UP);
BigDecimal bd29 = new BigDecimal(29);
BigDecimal bd93 = new BigDecimal(93);
BigDecimal bd3 = new BigDecimal(3);
BigDecimal bd2 = new BigDecimal(2);
BigDecimal bd1 = new BigDecimal(1);
BigDecimal bdkM = new BigDecimal(kM);
BigDecimal sqrt93 = bd93.sqrt(mc);
BigDecimal p = bd29.subtract(bd3.multiply(sqrt93, mc), mc).divide(bd2, mc);
BigDecimal q = bd29.subtract(p, mc);
BigDecimal r = bd1.add(cbrtDec(p, mc), mc).add(cbrtDec(q, mc), mc).divide(bd3, mc);
BigDecimal z = BigDecimal.ONE;
for (int i = 0; i < -kMinE; i++) {
z = z.divide(r, mc);
}
for (int i = 0; i < len; i++) {
BigDecimal x = z;
long y1 = x.longValue();
table.u1[i] = y1;
x = bdkM.multiply(x.subtract(new BigDecimal(y1), mc), mc);
long y2 = x.longValue();
table.u2[i] = y2;
x = bdkM.multiply(x.subtract(new BigDecimal(y2), mc), mc);
long y3 = x.longValue();
table.u3[i] = y3;
z = z.multiply(r, mc);
}
return table;
}
static long solveCalc(int limit) {
RootTable table = buildRootTable();
int len = kMaxE - kMinE + 1;
int i0 = -kMinE;
long sm = 0;
for (int j = 1; j <= limit; j++) {
long n = (long) j * j;
if (n == 1) {
sm++;
continue;
}
while (i0 < len && table.u1[i0] < n) {
i0++;
}
int i = i0 - 1;
long rep = 1;
long n1 = n - table.u1[i] - 1;
long n2 = kM - table.u2[i] - 1;
long n3 = kM - table.u3[i];
while (true) {
i -= 3;
while (i >= 0) {
if (table.u1[i] != n1) {
if (table.u1[i] < n1)
break;
} else if (table.u2[i] != n2) {
if (table.u2[i] < n2)
break;
} else if (table.u3[i] <= n3) {
break;
}
i--;
}
rep++;
n1 -= table.u1[i];
n2 -= table.u2[i];
n3 -= table.u3[i];
if (n3 < 0) {
n3 += kM;
n2--;
}
if (n2 < 0) {
n2 += kM;
n1--;
}
if (n1 == 0 && n2 == 0 && n3 < 1000) {
sm += rep;
break;
}
}
}
return sm;
}
public static String solve() {
return Long.toString(solveCalc(5000000));
}
public static void main(String[] args) {
System.out.println(solve());
}
}