Problem 339: Peredur Fab Efrawg
View on Project EulerProject Euler Problem 339 Solution
EulerSolve provides an optimized solution for Project Euler Problem 339, Peredur Fab Efrawg, with C++, Python, Java, and a step-by-step mathematical explanation.
Problem Summary We begin with two flocks, one white and one black, each containing \(n\) sheep. Every sheep is equally likely to bleat next. If a white sheep bleats, one black sheep crosses over and becomes white, so \((w,b)\mapsto (w+1,b-1)\). If a black sheep bleats, one white sheep crosses over and becomes black, so \((w,b)\mapsto (w-1,b+1)\). After this colour change, Peredur may remove any number of white sheep. The goal is to maximize the expected final number of black sheep, denoted \(E(n)\), starting from the balanced position \((n,n)\). The published solver computes this expectation to fixed decimal precision but, as elsewhere on the site, intentionally does not reveal the final Project Euler value. Mathematical Approach Let \(V(w,b)\) be the optimal expected final number of black sheep when the current state contains \(w\) white and \(b\) black sheep and Peredur is allowed to prune white sheep immediately. Terminal States If no white sheep remain, the process is over and all remaining sheep are black, so $$V(0,b)=b.$$ If no black sheep remain, the terminal payoff is zero: $$V(w,0)=0.$$ These boundary values anchor the entire dynamic program. Bellman Equation Suppose Peredur keeps exactly \(k\) white sheep, with \(1\le k\le w\)....
Detailed mathematical approach
Problem Summary
We begin with two flocks, one white and one black, each containing \(n\) sheep. Every sheep is equally likely to bleat next. If a white sheep bleats, one black sheep crosses over and becomes white, so \((w,b)\mapsto (w+1,b-1)\). If a black sheep bleats, one white sheep crosses over and becomes black, so \((w,b)\mapsto (w-1,b+1)\). After this colour change, Peredur may remove any number of white sheep. The goal is to maximize the expected final number of black sheep, denoted \(E(n)\), starting from the balanced position \((n,n)\).
The published solver computes this expectation to fixed decimal precision but, as elsewhere on the site, intentionally does not reveal the final Project Euler value.
Mathematical Approach
Let \(V(w,b)\) be the optimal expected final number of black sheep when the current state contains \(w\) white and \(b\) black sheep and Peredur is allowed to prune white sheep immediately.
Terminal States
If no white sheep remain, the process is over and all remaining sheep are black, so
$$V(0,b)=b.$$
If no black sheep remain, the terminal payoff is zero:
$$V(w,0)=0.$$
These boundary values anchor the entire dynamic program.
Bellman Equation
Suppose Peredur keeps exactly \(k\) white sheep, with \(1\le k\le w\). The next bleat is white with probability \(\frac{k}{k+b}\), leading to \((k+1,b-1)\), and black with probability \(\frac{b}{k+b}\), leading to \((k-1,b+1)\). Hence the one-step continuation value is
$$C(k,b)=\frac{k}{k+b}V(k+1,b-1)+\frac{b}{k+b}V(k-1,b+1).$$
If he removes all white sheep (\(k=0\)), the game stops immediately and the payoff is simply \(b\). Therefore the Bellman optimization is
$$V(w,b)=\max\!\left(b,\ \max_{1\le k\le w} C(k,b)\right).$$
This is a stochastic control problem: the decision is how many white sheep to keep, and the randomness comes from the next bleat.
Shape of the Optimal Policy
The fast solver relies on the policy
$$k^*(w,b)=\min(w,b-1),$$
which is independently checked by the exact solver on all small states used in the C++ checkpoints. In words, once white sheep are at least as numerous as black sheep, it is never beneficial to keep them above \(b-1\).
The intuition is straightforward: only black sheep contribute to the final payoff, while extra white sheep increase the chance of a white bleat, and a white bleat converts one black sheep into a white sheep. So excess white sheep are pure risk and can be removed without sacrificing upside.
Layering by the Fixed Total \(s=w+b\)
Between bleats no sheep are created; colours only swap. This makes the total
$$s=w+b$$
a natural layer index. Define
$$U_s(w)=V(w,s-w).$$
If \(w\le b-1\), equivalently \(1\le w\le t\) with
$$t=\left\lfloor\frac{s-1}{2}\right\rfloor,$$
then no immediate pruning occurs, and the Bellman equation reduces to the continuation recurrence
$$U_s(w)=\frac{w}{s}U_s(w+1)+\frac{s-w}{s}U_s(w-1).$$
If \(w\ge b\), the optimal policy removes white sheep until the state leaves that region. One immediate pruning step gives
$$U_s(w)=U_{s-1}(w-1),\qquad w\ge t+1.$$
Applying this relation repeatedly is exactly the same as capping white sheep to \(b-1\).
Tridiagonal Linear System
For a fixed layer \(s\), the unknown continuation values are \(U_s(1),U_s(2),\dots,U_s(t)\). Rearranging the recurrence gives
$$-\frac{s-w}{s}U_s(w-1)+U_s(w)-\frac{w}{s}U_s(w+1)=0,\qquad 1\le w\le t.$$
The left boundary is known exactly:
$$U_s(0)=V(0,s)=s.$$
The right boundary is inherited from the pruning region:
$$U_s(t+1)=U_{s-1}(t).$$
So each layer becomes a tridiagonal linear system, which the code solves in \(O(s)\) time by forward elimination and back substitution.
Recovering \(E(n)\) from the Balanced Start
The initial position is \((n,n)\), but Peredur may prune only after the first bleat and crossover. By symmetry, the first bleat is white or black with probability \(\frac12\). Thus
$$E(n)=\frac{V(n+1,n-1)+V(n-1,n+1)}{2}=\frac{U_{2n}(n+1)+U_{2n}(n-1)}{2}.$$
This is why the implementation computes the full layer \(s=2n\) and then averages the two neighboring states around the diagonal start.
Worked Example: \(n=1\)
With one white and one black sheep, the first bleat already determines the outcome. If the white sheep bleats, the state becomes \((2,0)\), so the final number of black sheep is \(0\). If the black sheep bleats, the state becomes \((0,2)\), so the final number of black sheep is \(2\). Therefore
$$E(1)=\frac{0+2}{2}=1.$$
The code also checks the published benchmark \(E(5)=6.871346\) (rounded to six decimals) as a checkpoint for the fast recurrence.
How the Code Works
The arrays prev and cur store the values of \(U_{s-1}\) and \(U_s\). For each layer \(s\), the cap region \(w\ge t+1\) is filled immediately by cur[w] = prev[w-1], implementing the pruning rule.
The continuation region \(1\le w\le t\) is then solved as a tridiagonal system using temporary arrays for the diagonal, superdiagonal, and right-hand side coefficients. After the final layer \(s=2n\) is built, the answer is returned as
$$\frac{\texttt{prev}[n-1]+\texttt{prev}[n+1]}{2}.$$
In the C++ version, an exact small-state solver is additionally used to verify the cap policy and to compare the fast method against exact values on sample inputs.
Complexity Analysis
Layer \(s\) contains \(O(s)\) relevant states and is solved in \(O(s)\) time. Summing over \(s=1,2,\dots,2n\) gives
$$O(1+2+\cdots+2n)=O(n^2).$$
Only the current layer, the previous layer, and a few temporary tridiagonal arrays are stored, so the memory usage is \(O(n)\).
Further Reading
- Problem page: https://projecteuler.net/problem=339
- Bellman equations / dynamic programming: https://en.wikipedia.org/wiki/Dynamic_programming
- Tridiagonal systems (Thomas algorithm): https://en.wikipedia.org/wiki/Tridiagonal_matrix_algorithm
Problem 339 source code
C++
#include <algorithm>
#include <chrono>
#include <cmath>
#include <cstdint>
#include <iomanip>
#include <iostream>
#include <limits>
#include <stdexcept>
#include <string>
#include <thread>
#include <vector>
namespace {
struct Options {
int n = 10'000;
bool run_checkpoints = true;
bool allow_multithreading = true;
unsigned requested_threads = 0U;
};
bool parse_u64_after_prefix(const std::string& arg, const char* prefix, std::uint64_t& value) {
const std::string p(prefix);
if (arg.rfind(p, 0) != 0) {
return false;
}
const std::string tail = arg.substr(p.size());
if (tail.empty()) {
return false;
}
std::uint64_t parsed = 0ULL;
for (const char c : tail) {
if (c < '0' || c > '9') {
return false;
}
const std::uint64_t digit = static_cast<std::uint64_t>(c - '0');
if (parsed > (std::numeric_limits<std::uint64_t>::max() - digit) / 10ULL) {
return false;
}
parsed = parsed * 10ULL + digit;
}
value = parsed;
return true;
}
bool parse_unsigned_after_prefix(const std::string& arg,
const char* prefix,
unsigned& value) {
std::uint64_t parsed = 0ULL;
if (!parse_u64_after_prefix(arg, prefix, parsed)) {
return false;
}
if (parsed > static_cast<std::uint64_t>(std::numeric_limits<unsigned>::max())) {
return false;
}
value = static_cast<unsigned>(parsed);
return true;
}
bool parse_arguments(int argc, char** argv, Options& options) {
for (int i = 1; i < argc; ++i) {
const std::string arg(argv[i]);
if (arg == "--skip-checkpoints") {
options.run_checkpoints = false;
continue;
}
if (arg == "--single-thread") {
options.allow_multithreading = false;
continue;
}
std::uint64_t parsed_u64 = 0ULL;
if (parse_u64_after_prefix(arg, "--n=", parsed_u64)) {
if (parsed_u64 > static_cast<std::uint64_t>(std::numeric_limits<int>::max())) {
std::cerr << "--n is too large for this build.\n";
return false;
}
options.n = static_cast<int>(parsed_u64);
continue;
}
unsigned parsed_unsigned = 0U;
if (parse_unsigned_after_prefix(arg, "--threads=", parsed_unsigned)) {
options.requested_threads = parsed_unsigned;
continue;
}
std::cerr << "Unknown argument: " << arg << '\n';
return false;
}
if (options.n <= 0) {
std::cerr << "--n must be >= 1.\n";
return false;
}
return true;
}
unsigned choose_thread_count(bool allow_multithreading,
unsigned requested_threads,
std::size_t workload) {
if (!allow_multithreading || workload < 2ULL) {
return 1U;
}
unsigned threads = requested_threads;
if (threads == 0U) {
threads = std::thread::hardware_concurrency();
if (threads == 0U) {
threads = 1U;
}
}
return std::max(1U, std::min<unsigned>(threads, static_cast<unsigned>(workload)));
}
// Fast solver based on the optimal cap policy: after each bleat, if (white >= black),
// remove white sheep so that white = black - 1.
long double solve_fast(int n) {
const int max_sum = 2 * n;
std::vector<long double> prev(1ULL, 0.0L);
std::vector<long double> cur;
std::vector<long double> b;
std::vector<long double> c;
std::vector<long double> d;
for (int s = 1; s <= max_sum; ++s) {
cur.assign(static_cast<std::size_t>(s) + 1ULL, 0.0L);
cur[0] = static_cast<long double>(s);
cur[static_cast<std::size_t>(s)] = 0.0L;
if (s >= 2) {
const int t = (s - 1) / 2; // continuation region: 1..t
for (int w = t + 1; w <= s - 1; ++w) {
cur[static_cast<std::size_t>(w)] = prev[static_cast<std::size_t>(w - 1)];
}
if (t >= 1) {
const long double s_ld = static_cast<long double>(s);
const long double right_boundary = prev[static_cast<std::size_t>(t)];
b.assign(static_cast<std::size_t>(t) + 1ULL, 0.0L);
c.assign(static_cast<std::size_t>(t) + 1ULL, 0.0L);
d.assign(static_cast<std::size_t>(t) + 1ULL, 0.0L);
b[1] = 1.0L;
c[1] = (t >= 2) ? (-1.0L / s_ld) : 0.0L;
d[1] = s_ld - 1.0L;
if (t == 1) {
d[1] += right_boundary / s_ld;
}
for (int w = 2; w <= t; ++w) {
const long double a = -static_cast<long double>(s - w) / s_ld;
long double bw = 1.0L;
long double cw = (w < t) ? (-static_cast<long double>(w) / s_ld) : 0.0L;
long double dw = 0.0L;
if (w == t) {
dw += static_cast<long double>(w) * right_boundary / s_ld;
}
const long double mult = a / b[static_cast<std::size_t>(w - 1)];
bw -= mult * c[static_cast<std::size_t>(w - 1)];
dw -= mult * d[static_cast<std::size_t>(w - 1)];
b[static_cast<std::size_t>(w)] = bw;
c[static_cast<std::size_t>(w)] = cw;
d[static_cast<std::size_t>(w)] = dw;
}
cur[static_cast<std::size_t>(t)] =
d[static_cast<std::size_t>(t)] / b[static_cast<std::size_t>(t)];
for (int w = t - 1; w >= 1; --w) {
cur[static_cast<std::size_t>(w)] =
(d[static_cast<std::size_t>(w)] -
c[static_cast<std::size_t>(w)] * cur[static_cast<std::size_t>(w + 1)]) /
b[static_cast<std::size_t>(w)];
}
}
}
prev.swap(cur);
}
const long double left = prev[static_cast<std::size_t>(n - 1)];
const long double right = prev[static_cast<std::size_t>(n + 1)];
return 0.5L * (left + right);
}
std::vector<std::vector<long double>> solve_exact_table(int max_sum,
long double tolerance = 1e-18L) {
if (max_sum < 0) {
throw std::runtime_error("max_sum must be non-negative.");
}
constexpr int kMaxIterations = 1'000'000;
std::vector<std::vector<long double>> value(static_cast<std::size_t>(max_sum) + 1ULL);
value[0] = {0.0L};
for (int s = 1; s <= max_sum; ++s) {
std::vector<long double> obstacle(static_cast<std::size_t>(s) + 1ULL, 0.0L);
for (int w = 1; w < s; ++w) {
obstacle[static_cast<std::size_t>(w)] =
value[static_cast<std::size_t>(s - 1)][static_cast<std::size_t>(w - 1)];
}
std::vector<long double> cur(static_cast<std::size_t>(s) + 1ULL, 0.0L);
cur[0] = static_cast<long double>(s);
cur[static_cast<std::size_t>(s)] = 0.0L;
for (int w = 1; w < s; ++w) {
cur[static_cast<std::size_t>(w)] = obstacle[static_cast<std::size_t>(w)];
}
const long double s_ld = static_cast<long double>(s);
for (int iter = 0; iter < kMaxIterations; ++iter) {
long double max_delta = 0.0L;
long double left_new = cur[0];
for (int w = 1; w < s; ++w) {
const std::size_t wi = static_cast<std::size_t>(w);
const long double old_value = cur[wi];
const long double continuation =
(static_cast<long double>(w) / s_ld) * cur[wi + 1ULL] +
(static_cast<long double>(s - w) / s_ld) * left_new;
const long double new_value = std::max(obstacle[wi], continuation);
cur[wi] = new_value;
const long double delta = std::fabsl(new_value - old_value);
max_delta = std::max(max_delta, delta);
left_new = new_value;
}
if (max_delta < tolerance) {
break;
}
if (iter + 1 == kMaxIterations) {
throw std::runtime_error("Exact checkpoint solver did not converge.");
}
}
value[static_cast<std::size_t>(s)] = std::move(cur);
}
return value;
}
long double expectation_from_table(const std::vector<std::vector<long double>>& table, int n) {
const int s = 2 * n;
const std::vector<long double>& diag = table[static_cast<std::size_t>(s)];
return 0.5L * (diag[static_cast<std::size_t>(n - 1)] + diag[static_cast<std::size_t>(n + 1)]);
}
long double solve_exact_small(int n) {
const auto table = solve_exact_table(2 * n);
return expectation_from_table(table, n);
}
bool verify_cap_policy_on_small_domain(int max_sum, long double tolerance) {
const auto table = solve_exact_table(max_sum);
for (int b = 1; b <= max_sum; ++b) {
for (int w = 0; w + b <= max_sum; ++w) {
const int expected_best = std::min(w, b - 1);
int best_k = 0;
long double best_value = static_cast<long double>(b); // k = 0
for (int k = 1; k <= w; ++k) {
const int s = k + b;
const long double s_ld = static_cast<long double>(s);
const long double candidate =
(static_cast<long double>(k) / s_ld) *
table[static_cast<std::size_t>(s)][static_cast<std::size_t>(k + 1)] +
(static_cast<long double>(b) / s_ld) *
table[static_cast<std::size_t>(s)][static_cast<std::size_t>(k - 1)];
if (candidate > best_value + tolerance) {
best_value = candidate;
best_k = k;
}
}
if (best_k != expected_best) {
std::cerr << "Cap-policy checkpoint failed at state (w=" << w
<< ", b=" << b << "). "
<< "Expected best k=" << expected_best
<< ", got k=" << best_k << '\n';
return false;
}
}
}
return true;
}
bool run_checkpoints(const Options& options) {
constexpr long double kGivenE5 = 6.871346L;
constexpr long double kGivenTolerance = 0.5e-6L;
constexpr long double kFastVsExactTolerance = 1e-11L;
if (!verify_cap_policy_on_small_domain(80, 1e-12L)) {
return false;
}
std::cout << "Checkpoint passed: cap policy verified for totals up to 80.\n";
const std::vector<int> test_ns = {1, 5, 10, 20, 30};
struct CheckResult {
int n = 0;
long double fast = 0.0L;
long double exact = 0.0L;
};
std::vector<CheckResult> results(test_ns.size());
const unsigned threads = choose_thread_count(
options.allow_multithreading, options.requested_threads, test_ns.size());
auto worker = [&](unsigned tid) {
for (std::size_t i = tid; i < test_ns.size(); i += threads) {
const int n = test_ns[i];
results[i].n = n;
results[i].fast = solve_fast(n);
results[i].exact = solve_exact_small(n);
}
};
std::vector<std::thread> pool;
pool.reserve(threads > 0U ? threads - 1U : 0U);
for (unsigned t = 1U; t < threads; ++t) {
pool.emplace_back(worker, t);
}
worker(0U);
for (std::thread& thread : pool) {
thread.join();
}
for (const CheckResult& result : results) {
const long double delta = std::fabsl(result.fast - result.exact);
if (delta > kFastVsExactTolerance) {
std::cerr << std::setprecision(18);
std::cerr << "Fast-vs-exact checkpoint failed for n=" << result.n
<< ": fast=" << result.fast
<< ", exact=" << result.exact
<< ", |delta|=" << delta << '\n';
return false;
}
std::cout << "Checkpoint passed: n=" << result.n
<< " fast/exact delta = " << std::setprecision(3) << std::scientific
<< static_cast<double>(delta) << '\n';
std::cout << std::defaultfloat;
}
const long double e5 = results[1].fast; // n = 5
if (std::fabsl(e5 - kGivenE5) > kGivenTolerance) {
std::cerr << std::setprecision(12);
std::cerr << "Given-value checkpoint failed: E(5) expected 6.871346 (rounded), got "
<< e5 << '\n';
return false;
}
std::cout << "Checkpoint passed: E(5) rounds to 6.871346.\n";
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 auto start_time = std::chrono::steady_clock::now();
if (options.run_checkpoints && !run_checkpoints(options)) {
return 1;
}
const long double answer = solve_fast(options.n);
const auto end_time = std::chrono::steady_clock::now();
const std::chrono::duration<double> elapsed = end_time - start_time;
std::cout << std::fixed << std::setprecision(6) << answer << '\n';
std::cout << "Elapsed: " << std::setprecision(3) << std::defaultfloat
<< elapsed.count() << " seconds\n";
std::cout << std::fixed << std::setprecision(6)
<< "Answer: " << answer << '\n';
} catch (const std::exception& ex) {
std::cerr << "Error: " << ex.what() << '\n';
return 1;
}
return 0;
}
Python
def solve():
n = 10000
def solve_fast(n):
max_sum = 2 * n
prev = [0.0]
for s in range(1, max_sum + 1):
cur = [0.0] * (s + 1)
cur[0] = float(s)
cur[s] = 0.0
if s >= 2:
t = (s - 1) // 2
for w in range(t + 1, s):
cur[w] = prev[w - 1]
if t >= 1:
s_f = float(s)
right_boundary = prev[t]
b = [0.0] * (t + 1)
c = [0.0] * (t + 1)
d = [0.0] * (t + 1)
b[1] = 1.0
c[1] = -1.0 / s_f if t >= 2 else 0.0
d[1] = s_f - 1.0
if t == 1:
d[1] += right_boundary / s_f
for w in range(2, t + 1):
a_coeff = -(s - w) / s_f
bw = 1.0
cw = -w / s_f if w < t else 0.0
dw = 0.0
if w == t:
dw += w * right_boundary / s_f
mult = a_coeff / b[w - 1]
bw -= mult * c[w - 1]
dw -= mult * d[w - 1]
b[w] = bw
c[w] = cw
d[w] = dw
cur[t] = d[t] / b[t]
for w in range(t - 1, 0, -1):
cur[w] = (d[w] - c[w] * cur[w + 1]) / b[w]
prev = cur
return 0.5 * (prev[n - 1] + prev[n + 1])
answer = solve_fast(n)
return f"{answer:.6f}"
if __name__ == '__main__':
print(solve())
Java
import java.util.*;
public class Euler339 {
public static String solve() {
int n = 10000;
int max_sum = 2 * n;
double[] prev = new double[1];
prev[0] = 0.0;
for (int s = 1; s <= max_sum; s++) {
double[] cur = new double[s + 1];
cur[0] = (double) s;
cur[s] = 0.0;
if (s >= 2) {
int t = (s - 1) / 2;
for (int w = t + 1; w < s; w++) {
cur[w] = prev[w - 1];
}
if (t >= 1) {
double sLd = (double) s;
double rightBoundary = prev[t];
double[] b = new double[t + 1];
double[] c = new double[t + 1];
double[] d = new double[t + 1];
b[1] = 1.0;
c[1] = (t >= 2) ? (-1.0 / sLd) : 0.0;
d[1] = sLd - 1.0;
if (t == 1) {
d[1] += rightBoundary / sLd;
}
for (int w = 2; w <= t; w++) {
double a = -(double) (s - w) / sLd;
double bw = 1.0;
double cw = (w < t) ? (-(double) w / sLd) : 0.0;
double dw = 0.0;
if (w == t) {
dw += (double) w * rightBoundary / sLd;
}
double mult = a / b[w - 1];
bw -= mult * c[w - 1];
dw -= mult * d[w - 1];
b[w] = bw;
c[w] = cw;
d[w] = dw;
}
cur[t] = d[t] / b[t];
for (int w = t - 1; w >= 1; w--) {
cur[w] = (d[w] - c[w] * cur[w + 1]) / b[w];
}
}
}
prev = cur;
}
double left = prev[n - 1];
double right = prev[n + 1];
double ans = 0.5 * (left + right);
return String.format(Locale.US, "%.6f", ans);
}
public static void main(String[] args) {
System.out.println(solve());
}
}