Problem 730: Shifted Pythagorean Triples

View on Project Euler

Project Euler Problem 730 Solution

EulerSolve provides an optimized solution for Project Euler Problem 730, Shifted Pythagorean Triples, with C++, Python, Java, and a step-by-step mathematical explanation.

Problem Summary For each shift \(k\) with \(0\le k\le m\), let \(P_k(n)\) be the number of unordered primitive positive integer triples \((p,q,r)\) such that $$p^2+q^2+k=r^2,\qquad p+q+r\le n,\qquad \gcd(p,q,r)=1.$$ The required quantity is $$S(m,n)=\sum_{k=0}^{m} P_k(n).$$ A brute-force scan over all triples up to the perimeter bound would be far too slow, so the implementation organizes the search through a finite set of roots and a tree of shift-preserving linear transformations. Mathematical Approach Fix one value of \(k\), and define the quadratic form $$Q(p,q,r)=r^2-p^2-q^2.$$ Then the triples relevant to this shift are exactly the primitive positive integer points on the level set \(Q(p,q,r)=k\). Step 1: Use transformations that preserve the shift The search is built from three linear maps: $$A(p,q,r)=\bigl(p-2q+2r,\ 2p-q+2r,\ 2p-2q+3r\bigr),$$ $$B(p,q,r)=\bigl(p+2q+2r,\ 2p+q+2r,\ 2p+2q+3r\bigr),$$ $$C(p,q,r)=\bigl(-p+2q+2r,\ -2p+q+2r,\ -2p+2q+3r\bigr).$$ A direct expansion shows that each map preserves the quadratic form: $$Q(A(p,q,r))=Q(B(p,q,r))=Q(C(p,q,r))=Q(p,q,r).$$ So if one triple satisfies \(p^2+q^2+k=r^2\), then all of its descendants under these maps satisfy the same equation with the same shift \(k\). The maps are also invertible over the integers, so primitiveness is preserved: any common divisor of a child would also divide its parent....

Detailed mathematical approach

Problem Summary

For each shift \(k\) with \(0\le k\le m\), let \(P_k(n)\) be the number of unordered primitive positive integer triples \((p,q,r)\) such that

$$p^2+q^2+k=r^2,\qquad p+q+r\le n,\qquad \gcd(p,q,r)=1.$$

The required quantity is

$$S(m,n)=\sum_{k=0}^{m} P_k(n).$$

A brute-force scan over all triples up to the perimeter bound would be far too slow, so the implementation organizes the search through a finite set of roots and a tree of shift-preserving linear transformations.

Mathematical Approach

Fix one value of \(k\), and define the quadratic form

$$Q(p,q,r)=r^2-p^2-q^2.$$

Then the triples relevant to this shift are exactly the primitive positive integer points on the level set \(Q(p,q,r)=k\).

Step 1: Use transformations that preserve the shift

The search is built from three linear maps:

$$A(p,q,r)=\bigl(p-2q+2r,\ 2p-q+2r,\ 2p-2q+3r\bigr),$$

$$B(p,q,r)=\bigl(p+2q+2r,\ 2p+q+2r,\ 2p+2q+3r\bigr),$$

$$C(p,q,r)=\bigl(-p+2q+2r,\ -2p+q+2r,\ -2p+2q+3r\bigr).$$

A direct expansion shows that each map preserves the quadratic form:

$$Q(A(p,q,r))=Q(B(p,q,r))=Q(C(p,q,r))=Q(p,q,r).$$

So if one triple satisfies \(p^2+q^2+k=r^2\), then all of its descendants under these maps satisfy the same equation with the same shift \(k\). The maps are also invertible over the integers, so primitiveness is preserved: any common divisor of a child would also divide its parent.

Step 2: Extract the roots by checking the inverse maps

The inverse transformations are

$$A^{-1}(p,q,r)=\bigl(p+2q-2r,\ -2p-q+2r,\ -2p-2q+3r\bigr),$$

$$B^{-1}(p,q,r)=\bigl(p+2q-2r,\ 2p+q-2r,\ -2p-2q+3r\bigr),$$

$$C^{-1}(p,q,r)=\bigl(-p-2q+2r,\ 2p+q-2r,\ -2p-2q+3r\bigr).$$

A primitive triple is called a root if none of these inverse images has all coordinates positive. Equivalently, a root is a solution with no parent inside the same search forest.

The implementations first enumerate primitive candidates with

$$1\le p\le q\le \max(4m,10),$$

test whether \(p^2+q^2+k\) is a square, and then discard every candidate that has a positive inverse parent. The underlying solver states that for \(k>0\), every root satisfies \(p,q\le 4k\), so scanning up to \(\max(4m,10)\) covers every shift \(0\le k\le m\).

Step 3: Traverse each rooted component with a perimeter cutoff

Once the roots for a fixed \(k\) are known, each root generates one component under repeated application of \(A\), \(B\), and \(C\). The implementations explore that component with an explicit depth-first stack.

The perimeter function

$$\pi(p,q,r)=p+q+r$$

provides the pruning rule: if \(\pi(p,q,r)>n\), that node contributes nothing and its branch is discarded. Therefore only triples that can affect \(P_k(n)\) are ever visited.

Because non-roots were removed in the previous step, every admissible triple belongs to exactly one rooted component, so starting DFS from all roots visits each relevant ordered triple exactly once.

Step 4: Convert ordered triples to unordered triples

The equation and the primitiveness condition are unchanged by the swap

$$\sigma(p,q,r)=(q,p,r).$$

The root search only seeds triples with \(p\le q\), so the component structure already avoids double-starting symmetric families. Inside a component, the implementation records two values:

$$T=\text{number of visited triples},\qquad D=\text{number of visited triples with }p=q.$$

If the component contains diagonal states \(p=q\), then \(\sigma\) keeps that component, fixes exactly the diagonal states, and pairs all off-diagonal states. The unordered contribution is therefore

$$U=\frac{T+D}{2}.$$

If no diagonal state appears, the selected rooted component already contributes the correct unordered count, so the implementation simply uses \(U=T\).

Step 5: Worked example for \(k=7\)

The triple \((1,1,3)\) is primitive and satisfies

$$1^2+1^2+7=9=3^2,$$

so it belongs to the shift \(k=7\). Applying the three transformations gives

$$A(1,1,3)=(5,7,9),\qquad B(1,1,3)=(9,9,13),\qquad C(1,1,3)=(7,5,9).$$

All three children still satisfy

$$r^2-p^2-q^2=7.$$

With perimeter bound \(n=31\), exactly four states from this component are counted:

$$ (1,1,3),\ (5,7,9),\ (7,5,9),\ (9,9,13). $$

Here \(T=4\), and the diagonal states are \((1,1,3)\) and \((9,9,13)\), so \(D=2\). Hence the unordered contribution is

$$U=\frac{4+2}{2}=3,$$

corresponding to the three unordered triples \((1,1,3)\), \((5,7,9)\), and \((9,9,13)\).

Step 6: Sum over all shifts

After computing the unordered count \(P_k(n)\) for each fixed \(k\), the final result is simply

$$S(m,n)=\sum_{k=0}^{m} P_k(n).$$

So the full problem is reduced to a finite root search for every shift and a pruned tree traversal from each root.

How the Code Works

The C++, Python, and Java implementations all follow the same counting strategy. For every \(k\), they enumerate candidate primitive triples in the bounded root window, keep only those whose inverse images are not positive, and store those survivors as roots. Then each root is explored with an explicit stack, repeatedly applying the three forward maps and pruning as soon as the perimeter exceeds \(n\).

For every rooted component, the implementation accumulates the total number of visited triples and, separately, the number with \(p=q\). That pair is converted into an unordered contribution using the formula from the previous section. Summing these component contributions yields \(P_k(n)\), and summing \(P_k(n)\) over \(k=0,\dots,m\) gives the final answer. The C++ and Java implementations additionally distribute independent root tasks across worker threads, while the Python implementation delegates to the same underlying search strategy rather than re-deriving different mathematics.

Complexity Analysis

Let \(R_k\) be the number of roots for shift \(k\), and let \(T_k(n)\) be the total number of triples visited across all rooted components after perimeter pruning. The preliminary root scan uses the bounded box

$$1\le p\le q\le \max(4m,10),$$

so its cost depends only on \(m\), not on the large perimeter bound \(n\). The dominant work is the traversal itself, which gives total running time

$$O\!\left(\sum_{k=0}^{m} T_k(n)\right).$$

Memory usage is linear in the stored roots plus the explicit DFS stack for the currently explored component. Parallel execution reduces wall-clock time, but it does not change the underlying combinatorial work measured by the same sum of visited states.

Footnotes and References

  1. Problem page: https://projecteuler.net/problem=730
  2. Pythagorean triple: Wikipedia — Pythagorean triple
  3. Tree of primitive Pythagorean triples: Wikipedia — Tree of primitive Pythagorean triples
  4. Depth-first search: Wikipedia — Depth-first search
  5. Quadratic form: Wikipedia — Quadratic form

Problem 730 source code

C++

#include <algorithm>
#include <atomic>
#include <cmath>
#include <cstdint>
#include <iostream>
#include <numeric>
#include <thread>
#include <vector>

using int64 = long long;

struct Triple {
    int64 p;
    int64 q;
    int64 r;
};

static bool is_square(int64 v, int64 &root) {
    if (v < 0) return false;
    int64 r = static_cast<int64>(std::sqrt(static_cast<long double>(v)) + 0.5L);
    if (r * r == v) {
        root = r;
        return true;
    }
    return false;
}

static bool has_parent(const Triple &t) {
    const int64 p = t.p;
    const int64 q = t.q;
    const int64 r = t.r;
    const int64 x3 = -2 * p - 2 * q + 3 * r;
    if (x3 <= 0) return false;
    int64 x1 = p + 2 * q - 2 * r;
    int64 x2 = -2 * p - q + 2 * r;
    if (x1 > 0 && x2 > 0) return true;  // A^{-1}
    x2 = 2 * p + q - 2 * r;
    if (x1 > 0 && x2 > 0) return true;  // B^{-1}
    x1 = -p - 2 * q + 2 * r;
    if (x1 > 0 && x2 > 0) return true;  // C^{-1}
    return false;
}

static std::vector<std::vector<Triple>> build_roots(int m) {
    // Roots are primitive solutions with no parent under A/B/C inverses.
    // For fixed k there are finitely many roots, and p,q <= 4k suffices for k>0.
    int limit = std::max(4 * m, 10);
    std::vector<std::vector<Triple>> roots(m + 1);
    for (int k = 0; k <= m; ++k) {
        for (int p = 1; p <= limit; ++p) {
            for (int q = p; q <= limit; ++q) {
                int64 r;
                int64 v = 1LL * p * p + 1LL * q * q + k;
                if (!is_square(v, r)) continue;
                if (std::gcd<int64>(p, std::gcd<int64>(q, r)) != 1) continue;
                Triple t{p, q, r};
                if (!has_parent(t)) roots[k].push_back(t);
            }
        }
    }
    return roots;
}

struct CountPair {
    int64 total;
    int64 diag;
};

static CountPair count_from_root(const Triple &root, int64 n, std::vector<Triple> &stack) {
    stack.clear();
    stack.push_back(root);
    int64 count = 0;
    int64 diag = 0;
    while (!stack.empty()) {
        Triple cur = stack.back();
        stack.pop_back();
        int64 per = cur.p + cur.q + cur.r;
        if (per > n) continue;
        ++count;
        if (cur.p == cur.q) ++diag;

        Triple a{cur.p - 2 * cur.q + 2 * cur.r,
                 2 * cur.p - cur.q + 2 * cur.r,
                 2 * cur.p - 2 * cur.q + 3 * cur.r};
        if (a.p + a.q + a.r <= n) stack.push_back(a);

        Triple b{cur.p + 2 * cur.q + 2 * cur.r,
                 2 * cur.p + cur.q + 2 * cur.r,
                 2 * cur.p + 2 * cur.q + 3 * cur.r};
        if (b.p + b.q + b.r <= n) stack.push_back(b);

        Triple c{-cur.p + 2 * cur.q + 2 * cur.r,
                 -2 * cur.p + cur.q + 2 * cur.r,
                 -2 * cur.p + 2 * cur.q + 3 * cur.r};
        if (c.p + c.q + c.r <= n) stack.push_back(c);
    }
    return {count, diag};
}

static std::vector<int64> count_all(int m, int64 n,
                                    const std::vector<std::vector<Triple>> &roots,
                                    int threads) {
    struct Task {
        int k;
        Triple root;
    };

    std::vector<Task> tasks;
    for (int k = 0; k <= m; ++k) {
        for (const auto &r : roots[k]) tasks.push_back({k, r});
    }

    if (tasks.empty()) return std::vector<int64>(m + 1, 0);

    unsigned hw = std::thread::hardware_concurrency();
    int T = (threads > 0) ? threads : (hw ? static_cast<int>(hw) : 1);
    T = std::max(1, std::min(T, 16));
    T = std::min<int>(T, tasks.size());

    std::vector<std::vector<int64>> local(T, std::vector<int64>(m + 1, 0));
    std::atomic<size_t> next(0);

    auto worker = [&](int tid) {
        std::vector<Triple> stack;
        stack.reserve(1024);
        while (true) {
            size_t idx = next.fetch_add(1);
            if (idx >= tasks.size()) break;
            const Task &task = tasks[idx];
            int64 per = task.root.p + task.root.q + task.root.r;
            if (per > n) continue;
            CountPair cnt = count_from_root(task.root, n, stack);
            // If this component hits the diagonal p==q, swapping p and q keeps the
            // component and pairs off-diagonal nodes; adjust to unordered counts.
            if (cnt.diag > 0) {
                local[tid][task.k] += (cnt.total + cnt.diag) / 2;
            } else {
                local[tid][task.k] += cnt.total;
            }
        }
    };

    std::vector<std::thread> pool;
    pool.reserve(T);
    for (int t = 0; t < T; ++t) pool.emplace_back(worker, t);
    for (auto &th : pool) th.join();

    std::vector<int64> counts(m + 1, 0);
    for (int t = 0; t < T; ++t) {
        for (int k = 0; k <= m; ++k) counts[k] += local[t][k];
    }
    return counts;
}

static bool run_validation(const std::vector<std::vector<Triple>> &roots, int threads) {
    const int64 n = 10000;
    const int m = 20;
    auto counts = count_all(m, n, roots, threads);

    if (counts[0] != 703) {
        std::cerr << "Validation failed: P_0(1e4) expected 703, got " << counts[0] << "\n";
        return false;
    }
    if (counts[20] != 1979) {
        std::cerr << "Validation failed: P_20(1e4) expected 1979, got " << counts[20] << "\n";
        return false;
    }
    int64 s10 = 0;
    for (int k = 0; k <= 10; ++k) s10 += counts[k];
    if (s10 != 10956) {
        std::cerr << "Validation failed: S(10,1e4) expected 10956, got " << s10 << "\n";
        return false;
    }
    std::cerr << "Validation checkpoints passed.\n";
    return true;
}

int main() {
    std::ios::sync_with_stdio(false);
    std::cin.tie(nullptr);

    const int m = 100;
    const int64 n = 100000000LL;

    unsigned hw = std::thread::hardware_concurrency();
    int threads = hw ? static_cast<int>(hw) : 1;
    threads = std::max(1, std::min(threads, 16));

    auto roots = build_roots(m);
    if (!run_validation(roots, threads)) return 1;

    auto counts = count_all(m, n, roots, threads);
    int64 total = 0;
    for (int k = 0; k <= m; ++k) total += counts[k];
    std::cout << total << "\n";
    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.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 Euler730 {
    static class Triple {
        long p, q, r;

        Triple(long p, long q, long r) {
            this.p = p;
            this.q = q;
            this.r = r;
        }
    }

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

    static boolean isSquare(long v, long[] root) {
        if (v < 0)
            return false;
        long r = isqrt(v);
        if (r * r == v) {
            root[0] = r;
            return true;
        }
        return false;
    }

    static boolean hasParent(Triple t) {
        long p = t.p, q = t.q, r = t.r;
        long x3 = -2 * p - 2 * q + 3 * r;
        if (x3 <= 0)
            return false;

        long x1 = p + 2 * q - 2 * r;
        long x2 = -2 * p - q + 2 * r;
        if (x1 > 0 && x2 > 0)
            return true;

        x2 = 2 * p + q - 2 * r;
        if (x1 > 0 && x2 > 0)
            return true;

        x1 = -p - 2 * q + 2 * r;
        if (x1 > 0 && x2 > 0)
            return true;

        return false;
    }

    static long gcd(long a, long b) {
        while (b != 0) {
            long temp = b;
            b = a % b;
            a = temp;
        }
        return a;
    }

    static List<List<Triple>> buildRoots(int m) {
        int limit = Math.max(4 * m, 10);
        List<List<Triple>> roots = new ArrayList<>(m + 1);
        for (int k = 0; k <= m; ++k) {
            roots.add(new ArrayList<>());
        }
        long[] root = new long[1];
        for (int k = 0; k <= m; ++k) {
            for (long p = 1; p <= limit; ++p) {
                for (long q = p; q <= limit; ++q) {
                    long v = p * p + q * q + k;
                    if (!isSquare(v, root))
                        continue;
                    long r = root[0];
                    if (gcd(p, gcd(q, r)) != 1)
                        continue;
                    Triple t = new Triple(p, q, r);
                    if (!hasParent(t)) {
                        roots.get(k).add(t);
                    }
                }
            }
        }
        return roots;
    }

    static class CountPair {
        long total, diag;

        CountPair(long total, long diag) {
            this.total = total;
            this.diag = diag;
        }
    }

    static CountPair countFromRoot(Triple root, long n) {
        List<Triple> stack = new ArrayList<>();
        stack.add(root);
        long count = 0;
        long diag = 0;

        while (!stack.isEmpty()) {
            Triple cur = stack.remove(stack.size() - 1);
            long per = cur.p + cur.q + cur.r;
            if (per > n)
                continue;
            count++;
            if (cur.p == cur.q)
                diag++;

            Triple a = new Triple(cur.p - 2 * cur.q + 2 * cur.r,
                    2 * cur.p - cur.q + 2 * cur.r,
                    2 * cur.p - 2 * cur.q + 3 * cur.r);
            if (a.p + a.q + a.r <= n)
                stack.add(a);

            Triple b = new Triple(cur.p + 2 * cur.q + 2 * cur.r,
                    2 * cur.p + cur.q + 2 * cur.r,
                    2 * cur.p + 2 * cur.q + 3 * cur.r);
            if (b.p + b.q + b.r <= n)
                stack.add(b);

            Triple c = new Triple(-cur.p + 2 * cur.q + 2 * cur.r,
                    -2 * cur.p + cur.q + 2 * cur.r,
                    -2 * cur.p + 2 * cur.q + 3 * cur.r);
            if (c.p + c.q + c.r <= n)
                stack.add(c);
        }
        return new CountPair(count, diag);
    }

    static class Task {
        int k;
        Triple root;

        Task(int k, Triple root) {
            this.k = k;
            this.root = root;
        }
    }

    public static String solve() {
        int m = 100;
        long n = 100000000L;

        List<List<Triple>> roots = buildRoots(m);
        List<Task> tasks = new ArrayList<>();
        for (int k = 0; k <= m; ++k) {
            for (Triple r : roots.get(k)) {
                tasks.add(new Task(k, r));
            }
        }

        if (tasks.isEmpty())
            return "0";

        int threads = Math.min(16, Math.max(1, Runtime.getRuntime().availableProcessors()));
        threads = Math.min(threads, tasks.size());

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

        int chunkSize = Math.max(1, tasks.size() / threads);

        for (int i = 0; i < tasks.size(); i += chunkSize) {
            final List<Task> chunk = tasks.subList(i, Math.min(i + chunkSize, tasks.size()));
            futures.add(executor.submit(() -> {
                long[] localCounts = new long[m + 1];
                for (Task task : chunk) {
                    long per = task.root.p + task.root.q + task.root.r;
                    if (per > n)
                        continue;
                    CountPair cnt = countFromRoot(task.root, n);
                    if (cnt.diag > 0) {
                        localCounts[task.k] += (cnt.total + cnt.diag) / 2;
                    } else {
                        localCounts[task.k] += cnt.total;
                    }
                }
                return localCounts;
            }));
        }

        long total = 0;
        try {
            for (Future<long[]> f : futures) {
                long[] res = f.get();
                for (int k = 0; k <= m; ++k) {
                    total += res[k];
                }
            }
        } catch (Exception e) {
            e.printStackTrace();
        }
        executor.shutdown();

        return Long.toString(total);
    }

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