Problem 535: Fractal Sequence

View on Project Euler

Project Euler Problem 535 Solution

EulerSolve provides an optimized solution for Project Euler Problem 535, Fractal Sequence, with C++, Python, Java, and a step-by-step mathematical explanation.

Problem Summary The task is to compute the sum of the first \(10^{18}\) terms of the fractal sequence, modulo \(10^9\). The key observation used by the implementations is that a prefix of length \(n\) does not need to be generated term by term: it can be decomposed into an inherited prefix plus a fresh consecutive tail. The difficulty is therefore not summing values directly, but locating the exact split point where the inherited self-similar part ends and the new tail begins. Mathematical Approach Let \(T(n)\) denote the sum of the first \(n\) terms of the sequence. We also introduce two auxiliary quantities: \(B(n)\), the length of the largest fully resolved inherited prefix contained in the first \(n\) terms; \(R(n)\), the total number of copied terms forced by the fresh values \(1,2,\dots,n\). Set the base values $$B(0)=0,\qquad R(0)=0,\qquad T(0)=0.$$ Once \(B(n)\) is known, write $$c=n-B(n).$$ The first \(n\) terms then split into the first \(B(n)\) terms of the same sequence, followed by the clean block \(1,2,\dots,c\). Step 1: Locate the Split Point After the fresh values \(1,2,\dots,k\) have all appeared, they have also generated \(R(k)\) copied terms....

Detailed mathematical approach

Problem Summary

The task is to compute the sum of the first \(10^{18}\) terms of the fractal sequence, modulo \(10^9\). The key observation used by the implementations is that a prefix of length \(n\) does not need to be generated term by term: it can be decomposed into an inherited prefix plus a fresh consecutive tail. The difficulty is therefore not summing values directly, but locating the exact split point where the inherited self-similar part ends and the new tail begins.

Mathematical Approach

Let \(T(n)\) denote the sum of the first \(n\) terms of the sequence. We also introduce two auxiliary quantities:

\(B(n)\), the length of the largest fully resolved inherited prefix contained in the first \(n\) terms;

\(R(n)\), the total number of copied terms forced by the fresh values \(1,2,\dots,n\).

Set the base values

$$B(0)=0,\qquad R(0)=0,\qquad T(0)=0.$$

Once \(B(n)\) is known, write

$$c=n-B(n).$$

The first \(n\) terms then split into the first \(B(n)\) terms of the same sequence, followed by the clean block \(1,2,\dots,c\).

Step 1: Locate the Split Point

After the fresh values \(1,2,\dots,k\) have all appeared, they have also generated \(R(k)\) copied terms. So the total occupied prefix length at that stage is

$$k+R(k).$$

Therefore the correct inherited prefix length is

$$B(n)=\max\{k\ge 0:\ k+R(k)\le n\}.$$

This formula means that \(B(n)\) is the largest fully completed stage that still fits inside the first \(n\) positions. Because \(R(k)\) is nondecreasing, the predicate \(k+R(k)\le n\) is monotone in \(k\), which is why the implementations can find \(B(n)\) by search rather than by scanning.

Step 2: Count How Many Copied Terms a New Tail Creates

Suppose the fresh tail has length \(c=n-B(n)\). The inherited part has already contributed \(R(B(n))\) copied terms. The new tail is exactly the consecutive block \(1,2,\dots,c\), and each fresh value \(i\) contributes a copied prefix of length

$$\lfloor\sqrt{i}\rfloor.$$

Hence the copied-term count satisfies the recurrence

$$R(n)=R(B(n))+\sum_{i=1}^{c}\lfloor\sqrt{i}\rfloor.$$

This is the structural heart of the solution: the self-similarity is pushed into the smaller argument \(B(n)\), while the new work depends only on the simple consecutive tail.

Step 3: Replace the Floor-Square-Root Sum by a Closed Form

Computing \(\sum_{i=1}^{c}\lfloor\sqrt{i}\rfloor\) term by term would be too slow. Let

$$s=\lfloor\sqrt{c}\rfloor.$$

For a fixed threshold \(t\), the inequality \(\lfloor\sqrt{i}\rfloor\ge t\) is equivalent to \(i\ge t^2\). So for each \(t\in\{1,\dots,s\}\), exactly \(c-t^2+1\) indices contribute at least \(t\). Counting by thresholds gives

$$\sum_{i=1}^{c}\lfloor\sqrt{i}\rfloor=\sum_{t=1}^{s}(c-t^2+1).$$

Using

$$\sum_{t=1}^{s} t^2=\frac{s(s+1)(2s+1)}{6},$$

we obtain the closed form

$$\sum_{i=1}^{c}\lfloor\sqrt{i}\rfloor=s(c+1)-\frac{s(s+1)(2s+1)}{6}.$$

This turns the expensive tail update into \(O(1)\) arithmetic once \(s\) is known.

Step 4: Derive the Prefix-Sum Recurrence

The tail \(1,2,\dots,c\) contributes the triangular number

$$1+2+\cdots+c=\frac{c(c+1)}{2}.$$

Therefore

$$T(n)=T(B(n))+\frac{c(c+1)}{2}.$$

So the whole problem reduces to repeatedly shrinking \(n\) to the smaller argument \(B(n)\), while adding one triangular contribution per level. No explicit sequence construction is needed.

Step 5: Worked Example for \(n=20\)

The local checkpoint is \(T(20)=86\). The split point is determined by

$$9+R(9)=20,\qquad 10+R(10)=22>20,$$

so

$$B(20)=9,\qquad c=20-9=11.$$

Hence

$$T(20)=T(9)+\frac{11\cdot 12}{2}=T(9)+66.$$

Apply the same decomposition again:

$$B(9)=4,\qquad T(9)=T(4)+\frac{5\cdot 6}{2}=T(4)+15,$$

$$B(4)=2,\qquad T(4)=T(2)+\frac{2\cdot 3}{2}=T(2)+3,$$

$$B(2)=1,\qquad T(2)=T(1)+\frac{1\cdot 2}{2}=T(1)+1,$$

$$B(1)=0,\qquad T(1)=1.$$

Combining everything,

$$T(20)=1+1+3+15+66=86,$$

which matches the checkpoint used by the implementations. The same recursion also gives the larger checkpoints \(T(10^3)=364089\) and \(T(10^9)=498676527978348241\).

How the Code Works

The C++, Python, and Java implementations memoize three kinds of information for previously solved prefix lengths: the split point \(B(n)\), the exact copied-term count \(R(n)\), and the required prefix sum \(T(n)\) modulo \(10^9\). For a new query \(n\), the implementation first brackets the answer for \(B(n)\) by repeated doubling, then performs binary search on the monotone condition \(k+R(k)\le n\).

Once the split point is known, the implementation computes the floor-square-root sum with the closed formula above, so no \(O(c)\) loop appears. The structural quantities used for comparisons are kept exact, while the running prefix sum is reduced modulo \(10^9\) after each triangular contribution. The non-Python versions also use wider intermediate arithmetic for the square-sum and triangular formulas so that large products remain exact before the final reduction.

Complexity Analysis

Let \(M\) be the number of distinct prefix lengths that are actually memoized while evaluating the target. Each such state is solved once. Determining its split point requires exponential bracketing and then binary search, so the search overhead is \(O(\log n_{\max})\) per new state, where \(n_{\max}\) is the largest queried prefix length. The floor-square-root contribution and the triangular contribution are both \(O(1)\).

Therefore the overall running time is \(O(M\log n_{\max})\), and the memory usage is \(O(M)\). In practice the recursion contracts quickly because every call replaces \(n\) by the smaller value \(B(n)\), which is exactly what makes a target as large as \(10^{18}\) feasible.

Footnotes and References

  1. Problem page: https://projecteuler.net/problem=535
  2. Square pyramidal number / sum of squares: Wikipedia — Square pyramidal number
  3. Triangular number: Wikipedia — Triangular number
  4. Binary search algorithm: Wikipedia — Binary search algorithm
  5. Memoization: Wikipedia — Memoization

Problem 535 source code

C++

#include <algorithm>
#include <cmath>
#include <cstdint>
#include <iomanip>
#include <iostream>
#include <unordered_map>

namespace {

using u64 = std::uint64_t;
using u128 = __uint128_t;

constexpr u64 kMod = 1'000'000'000ULL;

u64 isqrt_u64(u64 x) {
    u64 r = static_cast<u64>(std::sqrt(static_cast<long double>(x)));
    while ((r + 1) != 0 && (r + 1) * (r + 1) <= x) {
        ++r;
    }
    while (r * r > x) {
        --r;
    }
    return r;
}

u128 sum_floor_sqrt(u64 n) {
    // Sum_{k=1..n} floor(sqrt(k)).
    // Count by threshold: floor(sqrt(k)) >= t iff k >= t^2.
    // => sum = sum_{t=1..s} (n - t^2 + 1) = s*(n+1) - sum_{t=1..s} t^2.
    const u64 s = isqrt_u64(n);
    const u128 su = static_cast<u128>(s);
    const u128 nn1 = static_cast<u128>(n) + 1;
    const u128 sumsq = su * (su + 1) * (2 * su + 1) / 6;
    return su * nn1 - sumsq;
}

u64 tri_mod(u64 n) {
    // n*(n+1)/2 mod 1e9, computed exactly in 128-bit.
    const u128 t = static_cast<u128>(n) * (static_cast<u128>(n) + 1) / 2;
    return static_cast<u64>(t % static_cast<u128>(kMod));
}

struct Solver {
    std::unordered_map<u64, u64> memoA;
    std::unordered_map<u64, u128> memoU;
    std::unordered_map<u64, u64> memoTmod;

    Solver() {
        memoA.reserve(1 << 20);
        memoU.reserve(1 << 20);
        memoTmod.reserve(1 << 20);
        memoA[0] = 0;
        memoU[0] = 0;
        memoTmod[0] = 0;
    }

    u64 A(u64 n) {
        auto it = memoA.find(n);
        if (it != memoA.end()) {
            return it->second;
        }
        if (n <= 1) {
            memoA.emplace(n, 0);
            return 0;
        }

        // A(n) is the largest k such that k + U(k) <= n.
        auto f = [&](u64 k) -> u128 { return static_cast<u128>(k) + U(k); };

        const u64 hi_max = n - 1;  // A(n) < n since S starts with a circled 1
        u64 lo = 0;
        u64 hi = 1;
        while (hi < hi_max && f(hi) <= static_cast<u128>(n)) {
            lo = hi;
            const u64 next = hi << 1;
            hi = (next < hi_max) ? next : hi_max;
        }
        if (f(hi) <= static_cast<u128>(n)) {
            memoA.emplace(n, hi);
            return hi;
        }
        while (lo + 1 < hi) {
            const u64 mid = lo + (hi - lo) / 2;
            if (f(mid) <= static_cast<u128>(n)) {
                lo = mid;
            } else {
                hi = mid;
            }
        }

        memoA.emplace(n, lo);
        return lo;
    }

    u128 U(u64 n) {
        auto it = memoU.find(n);
        if (it != memoU.end()) {
            return it->second;
        }
        const u64 a = A(n);
        const u64 c = n - a;
        const u128 res = U(a) + sum_floor_sqrt(c);
        memoU.emplace(n, res);
        return res;
    }

    u64 T_mod(u64 n) {
        auto it = memoTmod.find(n);
        if (it != memoTmod.end()) {
            return it->second;
        }
        const u64 a = A(n);
        const u64 c = n - a;
        const u64 res = (T_mod(a) + tri_mod(c)) % kMod;
        memoTmod.emplace(n, res);
        return res;
    }

    u64 T_exact_u64(u64 n) {
        // Exact T(n) for moderate n where it fits in u64; used only for checkpoints.
        if (n == 0) {
            return 0;
        }
        const u64 a = A(n);
        const u64 c = n - a;
        const u128 res = static_cast<u128>(T_exact_u64(a)) +
                         static_cast<u128>(c) * (static_cast<u128>(c) + 1) / 2;
        return static_cast<u64>(res);
    }
};

bool run_checkpoints() {
    Solver s;
    if (s.T_exact_u64(1) != 1ULL) {
        std::cerr << "Checkpoint failed: T(1)\n";
        return false;
    }
    if (s.T_exact_u64(20) != 86ULL) {
        std::cerr << "Checkpoint failed: T(20) got " << s.T_exact_u64(20) << " A(20)=" << s.A(20)
                  << " T(9)=" << s.T_exact_u64(9) << " A(9)=" << s.A(9) << " T(4)=" << s.T_exact_u64(4)
                  << " A(4)=" << s.A(4) << " T(2)=" << s.T_exact_u64(2) << " A(2)=" << s.A(2) << '\n';
        return false;
    }
    if (s.T_exact_u64(1'000) != 364'089ULL) {
        std::cerr << "Checkpoint failed: T(10^3)\n";
        return false;
    }
    if (s.T_exact_u64(1'000'000'000ULL) != 498'676'527'978'348'241ULL) {
        std::cerr << "Checkpoint failed: T(10^9)\n";
        return false;
    }
    return true;
}

}  // namespace

int main() {
    if (!run_checkpoints()) {
        return 1;
    }
    Solver s;
    constexpr u64 n = 1'000'000'000'000'000'000ULL;
    const u64 ans = s.T_mod(n);
    std::cout << std::setw(9) << std::setfill('0') << ans << '\n';
    return 0;
}

Python

import math
import sys

sys.setrecursionlimit(20000)

kMod = 1000000000

def isqrt(x):
    r = int(math.sqrt(x))
    while (r + 1) * (r + 1) <= x:
        r += 1
    while r * r > x:
        r -= 1
    return r

def sum_floor_sqrt(n):
    s = isqrt(n)
    nn1 = n + 1
    sumsq = s * (s + 1) * (2 * s + 1) // 6
    return s * nn1 - sumsq

def tri_mod(n):
    t = n * (n + 1) // 2
    return t % kMod

memoA = {0: 0}
memoU = {0: 0}
memoTmod = {0: 0}

def U(k):
    if k in memoU:
        return memoU[k]
    a = A(k)
    c = k - a
    res = U(a) + sum_floor_sqrt(c)
    memoU[k] = res
    return res

def A(n):
    if n in memoA:
        return memoA[n]
    if n <= 1:
        memoA[n] = 0
        return 0
        
    def f(k):
        return k + U(k)
        
    hi_max = n - 1
    lo = 0
    hi = 1
    
    while hi < hi_max and f(hi) <= n:
        lo = hi
        hi = min(hi * 2, hi_max)
        
    if f(hi) <= n:
        memoA[n] = hi
        return hi
        
    while lo + 1 < hi:
        mid = lo + (hi - lo) // 2
        if f(mid) <= n:
            lo = mid
        else:
            hi = mid
            
    memoA[n] = lo
    return lo

def T_mod(n):
    if n in memoTmod:
        return memoTmod[n]
    a = A(n)
    c = n - a
    res = (T_mod(a) + tri_mod(c)) % kMod
    memoTmod[n] = res
    return res

def solve():
    n = 1000000000000000000
    ans = T_mod(n)
    return f"{ans:09d}"

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

Java

import java.util.HashMap;
import java.util.Map;
import java.math.BigInteger;

public class Euler535 {
    private static final long kMod = 1000000000L;

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

    private static long sumFloorSqrt(long n) {
        long s = isqrt(n);
        BigInteger bs = BigInteger.valueOf(s);
        BigInteger bs1 = BigInteger.valueOf(s + 1);
        BigInteger b2s1 = BigInteger.valueOf(2 * s + 1);
        BigInteger sumsq = bs.multiply(bs1).multiply(b2s1).divide(BigInteger.valueOf(6));
        BigInteger nn1 = BigInteger.valueOf(n + 1);
        BigInteger res = bs.multiply(nn1).subtract(sumsq);
        return res.longValue();
    }

    private static long triMod(long n) {
        BigInteger bn = BigInteger.valueOf(n);
        BigInteger bn1 = BigInteger.valueOf(n + 1);
        BigInteger t = bn.multiply(bn1).divide(BigInteger.valueOf(2));
        long rem = t.remainder(BigInteger.valueOf(kMod)).longValue();
        if (rem < 0)
            rem += kMod;
        return rem;
    }

    private static Map<Long, Long> memoA = new HashMap<>();
    private static Map<Long, Long> memoU = new HashMap<>();
    private static Map<Long, Long> memoTmod = new HashMap<>();

    static {
        memoA.put(0L, 0L);
        memoU.put(0L, 0L);
        memoTmod.put(0L, 0L);
    }

    private static long U(long k) {
        if (memoU.containsKey(k))
            return memoU.get(k);
        long a = A(k);
        long c = k - a;
        long res = U(a) + sumFloorSqrt(c);
        memoU.put(k, res);
        return res;
    }

    private static long f(long k) {
        return k + U(k);
    }

    private static long A(long n) {
        if (memoA.containsKey(n))
            return memoA.get(n);
        if (n <= 1) {
            memoA.put(n, 0L);
            return 0L;
        }

        long hiMax = n - 1;
        long lo = 0;
        long hi = 1;

        while (hi < hiMax && f(hi) <= n) {
            lo = hi;
            long next = hi << 1;
            hi = (next < hiMax && next > 0) ? next : hiMax;
        }
        if (f(hi) <= n) {
            memoA.put(n, hi);
            return hi;
        }

        while (lo + 1 < hi) {
            long mid = lo + (hi - lo) / 2;
            if (f(mid) <= n) {
                lo = mid;
            } else {
                hi = mid;
            }
        }

        memoA.put(n, lo);
        return lo;
    }

    private static long TMod(long n) {
        if (memoTmod.containsKey(n))
            return memoTmod.get(n);
        long a = A(n);
        long c = n - a;
        long res = (TMod(a) + triMod(c)) % kMod;
        memoTmod.put(n, res);
        return res;
    }

    public static void main(String[] args) {
        long n = 1000000000000000000L;
        long ans = TMod(n);
        System.out.printf("%09d\n", ans);
    }
}