Problem 1006: Fibonacci Subwords
View on Project EulerProject Euler Problem 1006 Solution
EulerSolve provides an optimized solution for Project Euler Problem 1006, Fibonacci Subwords, with C++, Python, Java, and a step-by-step mathematical explanation.
Problem Summary Starting from \(S_0=0\), \(S_1=01\), and \(S_n=S_{n-1}S_{n-2}\), every finite factor of the resulting infinite Fibonacci word is called a Fibonacci subword. For each positive \(k\), exactly \(k+1\) distinct factors have length \(k\). Reading each factor as a base-10 integer, leading zeros included harmlessly, we must compute the sum of their squares $$\Psi(k)=\sum_{u\in\mathcal F_k}\operatorname{val}_{10}(u)^2$$ for \(k=10^{18}\), modulo \(M=101001001\). Neither a word of that length nor its \(k+1\) factors can be constructed explicitly. Mathematical Approach Choose the first Fibonacci word longer than the window Let \(f_n=|S_n|\). Then \(f_0=1\), \(f_1=2\), and $$f_n=f_{n-1}+f_{n-2}.$$ Choose the smallest \(n\) for which \(L=f_n\gt k\), and put \(W=S_n=w_0w_1\cdots w_{L-1}\). A standard factor property of the Fibonacci word says that the \(L\) cyclic windows of length \(k\) in \(W\) contain all \(k+1\) distinct length-\(k\) factors. Since \(L\) windows represent only \(k+1\) values, $$\delta=L-k-1$$ of those occurrences are redundant. With the standard indexing used here, one extra copy of each of the first \(\delta\) consecutive windows must be removed....
Detailed mathematical approach
Problem Summary
Starting from \(S_0=0\), \(S_1=01\), and \(S_n=S_{n-1}S_{n-2}\), every finite factor of the resulting infinite Fibonacci word is called a Fibonacci subword. For each positive \(k\), exactly \(k+1\) distinct factors have length \(k\). Reading each factor as a base-10 integer, leading zeros included harmlessly, we must compute the sum of their squares
$$\Psi(k)=\sum_{u\in\mathcal F_k}\operatorname{val}_{10}(u)^2$$
for \(k=10^{18}\), modulo \(M=101001001\). Neither a word of that length nor its \(k+1\) factors can be constructed explicitly.
Mathematical Approach
Choose the first Fibonacci word longer than the window
Let \(f_n=|S_n|\). Then \(f_0=1\), \(f_1=2\), and
$$f_n=f_{n-1}+f_{n-2}.$$
Choose the smallest \(n\) for which \(L=f_n\gt k\), and put \(W=S_n=w_0w_1\cdots w_{L-1}\). A standard factor property of the Fibonacci word says that the \(L\) cyclic windows of length \(k\) in \(W\) contain all \(k+1\) distinct length-\(k\) factors. Since \(L\) windows represent only \(k+1\) values,
$$\delta=L-k-1$$
of those occurrences are redundant. With the standard indexing used here, one extra copy of each of the first \(\delta\) consecutive windows must be removed. Thus the problem becomes
$$\Psi(k)=\text{sum over all cyclic windows of }W -\text{sum of squares of the first }\delta\text{ windows}.$$
For \(k=10^{18}\), the index \(n\) is only \(O(\log k)\), although \(L\) itself is enormous.
Store words as a concatenation DAG
The implementation never expands \(S_n\). A leaf stores one bit, and an internal node stores two child identifiers and their total length. The recurrence \(S_n=S_{n-1}S_{n-2}\) therefore adds just one concatenation node per Fibonacci level. A prefix of any required length is obtained recursively: it lies wholly in the left child, or it is the entire left child followed by a prefix of the right child.
Concatenations and prefixes are memoized. Consequently, a prefix whose numerical length may be near \(10^{18}\) is still represented by only \(O(\log k)\) nodes.
Four summaries of a binary word
Work modulo \(M\), with \(b=10\). Because \(\gcd(10,M)=\gcd(99,M)=1\), both \(b^{-1}\) and \((b^2-1)^{-1}\) exist modulo \(M\). For a word \(X=x_0x_1\cdots x_{m-1}\), define
$$H_b(X)=\sum_{i=0}^{m-1}x_i b^i,\qquad R_b(X)=\sum_{i=0}^{m-1}x_i b^{m-1-i},$$
$$P_b(X)=\sum_{0\le i\lt j\lt m}x_i x_j b^{j-i},\qquad O(X)=\sum_{i=0}^{m-1}x_i.$$
\(R_b(X)\) is exactly the usual decimal value of \(X\) modulo \(M\). The pair summary \(P_b\) groups every pair of 1-bits by their distance.
If \(X=UV\), \(a=|U|\), and \(c=|V|\), all four values combine without inspecting a digit:
$$H_b(UV)=H_b(U)+b^aH_b(V),$$
$$R_b(UV)=b^cR_b(U)+R_b(V),$$
$$P_b(UV)=P_b(U)+P_b(V)+bR_b(U)H_b(V),$$
$$O(UV)=O(U)+O(V).$$
The cross term in \(P_b\) is correct because a bit \(i\) in \(U\) and a bit \(j\) in \(V\) are separated by \(a+j-i\) positions. The code caches these summaries for both bases \(b\) and \(b^{-1}\).
Range correlation queries on compressed words
Two additional recursive queries filter pairs without opening the word. The difference query computes
$$D_b(X,Y;\ell,h)= \sum_{\substack{i,j\\ \ell\le j-i\le h}}x_i y_j b^{j-i},$$
and the sum-index query computes
$$A_b(X,Y;\ell,h)= \sum_{\substack{i,j\\ \ell\le i+j\le h}}x_i y_j b^{i+j}.$$
When the requested range covers an entire node pair, the answer factors into two cached polynomial summaries. Otherwise the longer node is split, the interval is shifted by the child offset, and the two answers are combined. Memoization prevents the same node-pair/range state from being solved twice.
Cyclic correlations and the useful distances
For \(1\le d\lt L\), define the cyclic correlation
$$C_d=\sum_{i=0}^{L-1}w_iw_{(i+d)\bmod L}.$$
All nonzero cyclic distances can be packed into one polynomial:
$$Q_b=\sum_{d=1}^{L-1}C_db^d =P_b(W)+b^LP_{b^{-1}}(W).$$
The first term counts pairs that do not cross the end of \(W\); the second turns a linear distance \(j-i\) into its wrapped distance \(L-(j-i)\).
A length-\(k\) window can contain two positions only when their forward distance is \(1,\dots,k-1\). Put \(g=L-k\). The unwanted distances \(k,\dots,L-1\) are the reversals \(L-r\) of \(r=1,\dots,g\). Difference queries compute
$$E_b=\sum_{r=1}^{g}C_rb^r,$$
including both the ordinary and wrapped pieces. Hence the correlations actually used by windows are
$$U_b=Q_b-b^LE_{b^{-1}},\qquad U_{b^{-1}}=Q_{b^{-1}}-b^{-L}E_b.$$
Sum the squares of every cyclic window
Let \(V_s\) be the decimal value of the cyclic length-\(k\) window starting at \(s\). On expanding \(V_s^2\), diagonal digit terms and pairs of distinct positions separate cleanly.
Every 1-bit of \(W\) occupies every decimal place once across all cyclic windows. Its diagonal contribution is therefore
$$O(W)\sum_{t=0}^{k-1}b^{2t} =O(W)\frac{b^{2k}-1}{b^2-1}.$$
For a fixed cyclic distance \(d\), \(1\le d\lt k\), a pair appears in \(k-d\) relative placements. The sum of its place-value products is
$$\sum_{t=0}^{k-d-1}b^{2t+d} =\frac{b^{2k-d}-b^d}{b^2-1}.$$
After summing over all correlations, the square sum of all \(L\) cyclic windows is
$$T_{\mathrm{cyc}}=O(W)\frac{b^{2k}-1}{b^2-1} +\frac{2}{b^2-1}\left(b^{2k}U_{b^{-1}}-U_b\right)\pmod M.$$
This identity is the central compression step: exponentially many digit products collapse into two correlation evaluations and modular exponentiation.
Subtract the redundant windows without sliding one by one
If \(\delta=0\), the cyclic total already contains each factor once. Otherwise let \(V_0,\dots,V_{\delta-1}\) be the redundant prefix windows. The first value is \(R_b\) of a compressed prefix of length \(k\).
Put \(m=\delta-1\), and call the outgoing prefix \(x_0,\dots,x_{m-1}\). On this redundant run the entering block is the reversal of that prefix, so the bit entering at transition \(t\) is \(x_{m-1-t}\). The rolling decimal recurrence is therefore
$$V_{t+1}=bV_t-b^kx_t+x_{m-1-t},\qquad 0\le t\lt m.$$
Expanding this recurrence expresses every \(V_t\) in terms of \(V_0\), the outgoing prefix, and its reversal. Squaring and summing needs only counts of 1-bits, weighted bit pairs, the reversal diagonal \(\sum_t x_tx_{m-1-t}\), and two triangular index ranges. Those are precisely the cached summaries and \(A_b\) range queries. The method duplicated_window_sum evaluates the resulting telescoping/geometric expression modulo \(M\), then the final answer is
$$\Psi(k)\equiv T_{\mathrm{cyc}}-\sum_{t=0}^{\delta-1}V_t^2\pmod M.$$
Worked example: \(k=3\)
The first standard word longer than 3 is \(W=S_3=01001\), with \(L=5\) and \(\delta=5-3-1=1\). Its five cyclic windows are
$$010, 100, 001, 010, 101.$$
The first window \(010\) is the one redundant occurrence. Subtracting one copy leaves \(001,010,100,101\), so
$$\Psi(3)=1^2+10^2+100^2+101^2=20302.$$
How the Code Works
FibonacciSubwords builds the logarithmic standard-word DAG and caches powers of \(10\) and \(10^{-1}\). summarize implements the four concatenation formulas; difference_query filters correlations by \(j-i\), while sum_query handles the triangular products needed by the duplicate correction.
solve selects \(L\), forms \(Q_b\), removes separations too large for a \(k\)-window, evaluates \(T_{\mathrm{cyc}}\), and subtracts duplicated_window_sum. All arithmetic is reduced modulo \(101001001\).
The checkpoint suite constructs actual Fibonacci strings only for small inputs. It confirms \(\Psi(3)=20302\), compares the compressed solver with direct set enumeration for every \(1\le k\le50\), and verifies the supplied value \(\Psi(10)\equiv10699667\pmod M\).
Complexity Analysis
There are \(O(\log k)\) Fibonacci levels and prefix nodes. The memoized two-word range queries visit at most a quadratic number of relevant node-pair states, giving \(O(\log^2 k)\) structural work and \(O(\log^2 k)\) cached memory; modular exponentiation contributes \(O(\log k)\) multiplications per previously unseen exponent.
No operation is linear in \(k\), and no length-\(k\) string is stored. This is what makes \(k=10^{18}\) practical.
Footnotes and References
- Problem page: Project Euler 1006 - Fibonacci Subwords
- Fibonacci word: Wikipedia - Fibonacci word
- Sturmian word and factor complexity: Wikipedia - Sturmian word
- Modular multiplicative inverse: Wikipedia - Modular multiplicative inverse
Problem 1006 source code
C++
#include <algorithm>
#include <array>
#include <cassert>
#include <cstdint>
#include <cstdlib>
#include <iostream>
#include <map>
#include <set>
#include <string>
#include <unordered_map>
#include <utility>
#include <vector>
namespace {
using i64 = std::int64_t;
using u64 = std::uint64_t;
constexpr i64 MOD = 101'001'001;
constexpr i64 BASE = 10;
constexpr u64 TARGET = 1'000'000'000'000'000'000ULL;
i64 normalize(i64 value) {
value %= MOD;
if (value < 0) {
value += MOD;
}
return value;
}
i64 mod_pow(i64 base, u64 exponent) {
i64 result = 1;
base = normalize(base);
while (exponent > 0) {
if ((exponent & 1ULL) != 0) {
result = result * base % MOD;
}
base = base * base % MOD;
exponent >>= 1ULL;
}
return result;
}
i64 extended_gcd(const i64 a, const i64 b, i64& x, i64& y) {
if (b == 0) {
x = 1;
y = 0;
return a;
}
i64 next_x = 0;
i64 next_y = 0;
const i64 gcd = extended_gcd(b, a % b, next_x, next_y);
x = next_y;
y = next_x - (a / b) * next_y;
return gcd;
}
i64 mod_inverse(const i64 value) {
i64 x = 0;
i64 y = 0;
const i64 gcd = extended_gcd(normalize(value), MOD, x, y);
if (gcd != 1) {
std::cerr << "A required modular inverse does not exist.\n";
std::exit(EXIT_FAILURE);
}
return normalize(x);
}
struct Node {
u64 length = 0;
int left = -1;
int right = -1;
int bit = -1;
};
struct Summary {
i64 forward = 0;
i64 reverse = 0;
i64 pairs = 0;
i64 ones = 0;
bool ready = false;
};
struct QueryKey {
int x = 0;
int y = 0;
i64 low = 0;
i64 high = 0;
int base_id = 0;
bool operator==(const QueryKey& other) const {
return x == other.x && y == other.y && low == other.low && high == other.high &&
base_id == other.base_id;
}
};
struct QueryKeyHash {
std::size_t operator()(const QueryKey& key) const {
std::size_t hash = static_cast<std::size_t>(key.x) * 1'000'003ULL +
static_cast<std::size_t>(key.y);
hash ^= static_cast<std::size_t>(key.low) + 0x9e3779b97f4a7c15ULL + (hash << 6U) +
(hash >> 2U);
hash ^= static_cast<std::size_t>(key.high) + 0x9e3779b97f4a7c15ULL + (hash << 6U) +
(hash >> 2U);
hash ^= static_cast<std::size_t>(key.base_id) + (hash << 6U) + (hash >> 2U);
return hash;
}
};
class FibonacciSubwords {
public:
FibonacciSubwords() {
bases_[0] = BASE;
bases_[1] = mod_inverse(BASE);
zero_ = add_bit(0);
one_ = add_bit(1);
fibonacci_.push_back(1);
fibonacci_.push_back(2);
standard_.push_back(zero_);
standard_.push_back(concatenate(zero_, one_));
while (fibonacci_.back() <= TARGET) {
fibonacci_.push_back(fibonacci_[fibonacci_.size() - 1] +
fibonacci_[fibonacci_.size() - 2]);
standard_.push_back(concatenate(standard_[standard_.size() - 1],
standard_[standard_.size() - 2]));
}
}
i64 solve(const u64 k) {
const auto iterator = std::upper_bound(fibonacci_.begin(), fibonacci_.end(), k);
assert(iterator != fibonacci_.end());
const int index = static_cast<int>(iterator - fibonacci_.begin());
const u64 cycle_length = fibonacci_[static_cast<std::size_t>(index)];
const u64 duplicate_count = cycle_length - 1 - k;
const u64 gap_length = duplicate_count + 1;
const int word = standard_[static_cast<std::size_t>(index)];
const Summary word_base = summarize(word, 0);
const Summary word_inverse = summarize(word, 1);
const i64 cycle_power = power(0, cycle_length);
const i64 inverse_cycle_power = power(1, cycle_length);
const i64 all_correlations_base =
normalize(word_base.pairs + cycle_power * word_inverse.pairs);
const i64 all_correlations_inverse =
normalize(word_inverse.pairs + inverse_cycle_power * word_base.pairs);
const i64 prefix_correlations_base = normalize(
difference_query(word, word, 1, static_cast<i64>(gap_length), 0) +
cycle_power * difference_query(word,
word,
static_cast<i64>(cycle_length - gap_length),
static_cast<i64>(cycle_length - 1),
1));
const i64 prefix_correlations_inverse = normalize(
difference_query(word, word, 1, static_cast<i64>(gap_length), 1) +
inverse_cycle_power * difference_query(word,
word,
static_cast<i64>(cycle_length - gap_length),
static_cast<i64>(cycle_length - 1),
0));
const i64 used_correlations_base = normalize(
all_correlations_base - cycle_power * prefix_correlations_inverse);
const i64 used_correlations_inverse = normalize(
all_correlations_inverse - inverse_cycle_power * prefix_correlations_base);
const i64 inverse_geometric_denominator = mod_inverse(BASE * BASE - 1);
const i64 geometric_squares =
normalize(power(0, 2 * k) - 1) * inverse_geometric_denominator % MOD;
const i64 full_cycle = normalize(
word_base.ones * geometric_squares +
2 * inverse_geometric_denominator % MOD *
normalize(power(0, 2 * k) * used_correlations_inverse -
used_correlations_base));
return normalize(full_cycle -
duplicated_window_sum(index, k, duplicate_count));
}
private:
std::vector<Node> nodes_;
std::vector<std::array<Summary, 2>> summary_cache_;
std::map<std::pair<int, int>, int> concatenation_cache_;
std::map<std::pair<int, u64>, int> prefix_cache_;
std::vector<u64> fibonacci_;
std::vector<int> standard_;
std::array<i64, 2> bases_{};
std::array<std::unordered_map<u64, i64>, 2> power_cache_;
std::unordered_map<QueryKey, i64, QueryKeyHash> sum_query_cache_;
std::unordered_map<QueryKey, i64, QueryKeyHash> difference_query_cache_;
int zero_ = -1;
int one_ = -1;
int add_bit(const int bit) {
const int id = static_cast<int>(nodes_.size());
nodes_.push_back(Node{1, -1, -1, bit});
summary_cache_.push_back({});
return id;
}
int concatenate(const int left, const int right) {
if (left < 0) {
return right;
}
if (right < 0) {
return left;
}
const std::pair<int, int> key{left, right};
const auto found = concatenation_cache_.find(key);
if (found != concatenation_cache_.end()) {
return found->second;
}
const int id = static_cast<int>(nodes_.size());
nodes_.push_back(Node{nodes_[static_cast<std::size_t>(left)].length +
nodes_[static_cast<std::size_t>(right)].length,
left,
right,
-1});
summary_cache_.push_back({});
concatenation_cache_[key] = id;
return id;
}
int prefix(const int standard_index, const u64 length) {
assert(length > 0 && length <= fibonacci_[static_cast<std::size_t>(standard_index)]);
if (length == fibonacci_[static_cast<std::size_t>(standard_index)]) {
return standard_[static_cast<std::size_t>(standard_index)];
}
const std::pair<int, u64> key{standard_index, length};
const auto found = prefix_cache_.find(key);
if (found != prefix_cache_.end()) {
return found->second;
}
assert(standard_index > 0);
int result = -1;
if (length <= fibonacci_[static_cast<std::size_t>(standard_index - 1)]) {
result = prefix(standard_index - 1, length);
} else {
result = concatenate(
standard_[static_cast<std::size_t>(standard_index - 1)],
prefix(standard_index - 2,
length - fibonacci_[static_cast<std::size_t>(standard_index - 1)]));
}
prefix_cache_[key] = result;
return result;
}
i64 power(const int base_id, const u64 exponent) {
auto& cache = power_cache_[static_cast<std::size_t>(base_id)];
const auto found = cache.find(exponent);
if (found != cache.end()) {
return found->second;
}
const i64 result = mod_pow(bases_[static_cast<std::size_t>(base_id)], exponent);
cache[exponent] = result;
return result;
}
i64 signed_power(const int base_id, const i64 exponent) {
if (exponent >= 0) {
return power(base_id, static_cast<u64>(exponent));
}
return power(1 - base_id, static_cast<u64>(-exponent));
}
Summary summarize(const int node_id, const int base_id) {
Summary& cached =
summary_cache_[static_cast<std::size_t>(node_id)][static_cast<std::size_t>(base_id)];
if (cached.ready) {
return cached;
}
const Node node = nodes_[static_cast<std::size_t>(node_id)];
Summary result;
if (node.bit >= 0) {
result.forward = node.bit;
result.reverse = node.bit;
result.ones = node.bit;
result.ready = true;
cached = result;
return result;
}
const Summary left = summarize(node.left, base_id);
const Summary right = summarize(node.right, base_id);
const u64 left_length = nodes_[static_cast<std::size_t>(node.left)].length;
const u64 right_length = nodes_[static_cast<std::size_t>(node.right)].length;
const i64 base = bases_[static_cast<std::size_t>(base_id)];
result.forward = normalize(left.forward + power(base_id, left_length) * right.forward);
result.reverse = normalize(power(base_id, right_length) * left.reverse + right.reverse);
result.pairs =
normalize(left.pairs + right.pairs + base * left.reverse % MOD * right.forward);
result.ones = normalize(left.ones + right.ones);
result.ready = true;
cached = result;
return result;
}
i64 sum_query(int x, int y, const i64 low, const i64 high, const int base_id) {
if (x > y) {
std::swap(x, y);
}
const QueryKey key{x, y, low, high, base_id};
const auto found = sum_query_cache_.find(key);
if (found != sum_query_cache_.end()) {
return found->second;
}
const Node node_x = nodes_[static_cast<std::size_t>(x)];
const Node node_y = nodes_[static_cast<std::size_t>(y)];
const i64 maximum = static_cast<i64>(node_x.length + node_y.length - 2);
i64 result = 0;
if (high < 0 || low > maximum) {
result = 0;
} else if (low <= 0 && maximum <= high) {
result = summarize(x, base_id).forward * summarize(y, base_id).forward % MOD;
} else if (node_x.bit >= 0 && node_y.bit >= 0) {
result = low <= 0 && 0 <= high ? node_x.bit * node_y.bit : 0;
} else if (node_x.length >= node_y.length && node_x.bit < 0) {
const u64 offset = nodes_[static_cast<std::size_t>(node_x.left)].length;
result = normalize(
sum_query(node_x.left, y, low, high, base_id) +
power(base_id, offset) *
sum_query(node_x.right,
y,
low - static_cast<i64>(offset),
high - static_cast<i64>(offset),
base_id));
} else {
const u64 offset = nodes_[static_cast<std::size_t>(node_y.left)].length;
result = normalize(
sum_query(x, node_y.left, low, high, base_id) +
power(base_id, offset) *
sum_query(x,
node_y.right,
low - static_cast<i64>(offset),
high - static_cast<i64>(offset),
base_id));
}
sum_query_cache_[key] = result;
return result;
}
i64 difference_query(const int x,
const int y,
const i64 low,
const i64 high,
const int base_id) {
const QueryKey key{x, y, low, high, base_id};
const auto found = difference_query_cache_.find(key);
if (found != difference_query_cache_.end()) {
return found->second;
}
const Node node_x = nodes_[static_cast<std::size_t>(x)];
const Node node_y = nodes_[static_cast<std::size_t>(y)];
const i64 minimum = -static_cast<i64>(node_x.length - 1);
const i64 maximum = static_cast<i64>(node_y.length - 1);
i64 result = 0;
if (high < minimum || low > maximum) {
result = 0;
} else if (low <= minimum && maximum <= high) {
result = summarize(x, 1 - base_id).forward * summarize(y, base_id).forward % MOD;
} else if (node_x.bit >= 0 && node_y.bit >= 0) {
result = low <= 0 && 0 <= high ? node_x.bit * node_y.bit : 0;
} else if (node_x.length >= node_y.length && node_x.bit < 0) {
const u64 offset = nodes_[static_cast<std::size_t>(node_x.left)].length;
result = normalize(
difference_query(node_x.left, y, low, high, base_id) +
signed_power(base_id, -static_cast<i64>(offset)) *
difference_query(node_x.right,
y,
low + static_cast<i64>(offset),
high + static_cast<i64>(offset),
base_id));
} else {
const u64 offset = nodes_[static_cast<std::size_t>(node_y.left)].length;
result = normalize(
difference_query(x, node_y.left, low, high, base_id) +
power(base_id, offset) *
difference_query(x,
node_y.right,
low - static_cast<i64>(offset),
high - static_cast<i64>(offset),
base_id));
}
difference_query_cache_[key] = result;
return result;
}
i64 duplicated_window_sum(const int standard_index,
const u64 k,
const u64 count) {
if (count == 0) {
return 0;
}
const int initial_word = prefix(standard_index, k);
const i64 initial_value = summarize(initial_word, 0).reverse;
if (count == 1) {
return initial_value * initial_value % MOD;
}
const u64 transition_count = count - 1;
const int outgoing = prefix(standard_index, transition_count);
const Summary outgoing_summary = summarize(outgoing, 0);
const i64 window_power = power(0, k);
const i64 forward_delta = normalize(outgoing_summary.reverse -
window_power * outgoing_summary.forward);
const i64 reverse_delta = normalize(outgoing_summary.forward -
window_power * outgoing_summary.reverse);
const i64 low_inverse = sum_query(outgoing,
outgoing,
0,
static_cast<i64>(transition_count) - 2,
1);
const i64 high_base = sum_query(outgoing,
outgoing,
static_cast<i64>(transition_count),
static_cast<i64>(2 * transition_count - 2),
0);
const i64 diagonal_base = sum_query(outgoing,
outgoing,
static_cast<i64>(transition_count - 1),
static_cast<i64>(transition_count - 1),
0);
const i64 diagonal =
diagonal_base * power(1, transition_count - 1) % MOD;
const i64 lower_cross = power(0, transition_count - 1) * low_inverse % MOD;
const i64 upper_cross = power(1, transition_count - 1) * high_base % MOD;
const i64 delta_pairs = normalize(
normalize(1 + window_power * window_power) * outgoing_summary.pairs -
window_power * normalize(lower_cross + upper_cross));
const i64 delta_squares = normalize(
normalize(1 + window_power * window_power) * outgoing_summary.ones -
2 * window_power % MOD * diagonal);
const i64 last_value = normalize(
power(0, transition_count) * initial_value + reverse_delta);
const i64 value_delta_sum = normalize(
initial_value * forward_delta + bases_[1] * delta_pairs);
const i64 numerator = normalize(
initial_value * initial_value -
BASE * BASE % MOD * last_value % MOD * last_value +
2 * BASE % MOD * value_delta_sum + delta_squares);
return numerator * mod_inverse(1 - BASE * BASE) % MOD;
}
};
i64 brute_force(const int k) {
std::string older = "0";
std::string newer = "01";
while (static_cast<int>(newer.size()) < 10 * k + 20) {
const std::string next = newer + older;
older = newer;
newer = next;
}
std::set<std::string> factors;
for (int start = 0; start + k <= static_cast<int>(newer.size()); ++start) {
factors.insert(newer.substr(static_cast<std::size_t>(start), static_cast<std::size_t>(k)));
}
assert(factors.size() == static_cast<std::size_t>(k + 1));
i64 result = 0;
for (const std::string& factor : factors) {
i64 value = 0;
for (const char digit : factor) {
value = (BASE * value + digit - '0') % MOD;
}
result = (result + value * value) % MOD;
}
return result;
}
void run_checkpoints() {
FibonacciSubwords solver;
if (brute_force(3) != 20'302) {
std::cerr << "Checkpoint failed: Psi(3).\n";
std::exit(EXIT_FAILURE);
}
for (int k = 1; k <= 50; ++k) {
if (solver.solve(static_cast<u64>(k)) != brute_force(k)) {
std::cerr << "Checkpoint failed: brute force comparison for k=" << k << ".\n";
std::exit(EXIT_FAILURE);
}
}
if (solver.solve(10) != 10'699'667) {
std::cerr << "Checkpoint failed: supplied Psi(10).\n";
std::exit(EXIT_FAILURE);
}
std::cerr << "Validation checkpoints passed.\n";
}
}
int main() {
run_checkpoints();
FibonacciSubwords solver;
std::cout << solver.solve(TARGET) << '\n';
return 0;
}
Python
#!/usr/bin/env python3
"""Project Euler Problem 1006 - Fibonacci Subwords."""
from __future__ import annotations
import bisect
import sys
from dataclasses import dataclass
MOD = 101_001_001
BASE = 10
TARGET = 1_000_000_000_000_000_000
def normalize(value: int) -> int:
return value % MOD
def mod_pow(base: int, exponent: int) -> int:
return pow(normalize(base), exponent, MOD)
def extended_gcd(a: int, b: int) -> tuple[int, int, int]:
if b == 0:
return a, 1, 0
gcd, next_x, next_y = extended_gcd(b, a % b)
return gcd, next_y, next_x - (a // b) * next_y
def mod_inverse(value: int) -> int:
gcd, x, _ = extended_gcd(normalize(value), MOD)
if gcd != 1:
raise ValueError("a required modular inverse does not exist")
return normalize(x)
@dataclass
class Node:
length: int
left: int = -1
right: int = -1
bit: int = -1
@dataclass
class Summary:
forward: int = 0
reverse: int = 0
pairs: int = 0
ones: int = 0
ready: bool = False
class FibonacciSubwords:
def __init__(self) -> None:
self.bases = [BASE, mod_inverse(BASE)]
self.nodes: list[Node] = []
self.summary_cache: list[list[Summary]] = []
self.concatenation_cache: dict[tuple[int, int], int] = {}
self.prefix_cache: dict[tuple[int, int], int] = {}
self.power_cache: list[dict[int, int]] = [{}, {}]
self.sum_query_cache: dict[tuple[int, int, int, int, int], int] = {}
self.difference_query_cache: dict[tuple[int, int, int, int, int], int] = {}
self.zero = self.add_bit(0)
self.one = self.add_bit(1)
self.fibonacci = [1, 2]
self.standard = [self.zero, self.concatenate(self.zero, self.one)]
while self.fibonacci[-1] <= TARGET:
self.fibonacci.append(self.fibonacci[-1] + self.fibonacci[-2])
self.standard.append(
self.concatenate(self.standard[-1], self.standard[-2])
)
def add_bit(self, bit: int) -> int:
node_id = len(self.nodes)
self.nodes.append(Node(length=1, bit=bit))
self.summary_cache.append([Summary(), Summary()])
return node_id
def concatenate(self, left: int, right: int) -> int:
if left < 0:
return right
if right < 0:
return left
key = (left, right)
cached = self.concatenation_cache.get(key)
if cached is not None:
return cached
node_id = len(self.nodes)
self.nodes.append(
Node(
length=self.nodes[left].length + self.nodes[right].length,
left=left,
right=right,
)
)
self.summary_cache.append([Summary(), Summary()])
self.concatenation_cache[key] = node_id
return node_id
def prefix(self, standard_index: int, length: int) -> int:
assert 0 < length <= self.fibonacci[standard_index]
if length == self.fibonacci[standard_index]:
return self.standard[standard_index]
key = (standard_index, length)
cached = self.prefix_cache.get(key)
if cached is not None:
return cached
assert standard_index > 0
if length <= self.fibonacci[standard_index - 1]:
result = self.prefix(standard_index - 1, length)
else:
result = self.concatenate(
self.standard[standard_index - 1],
self.prefix(
standard_index - 2,
length - self.fibonacci[standard_index - 1],
),
)
self.prefix_cache[key] = result
return result
def power(self, base_id: int, exponent: int) -> int:
cache = self.power_cache[base_id]
cached = cache.get(exponent)
if cached is not None:
return cached
result = mod_pow(self.bases[base_id], exponent)
cache[exponent] = result
return result
def signed_power(self, base_id: int, exponent: int) -> int:
if exponent >= 0:
return self.power(base_id, exponent)
return self.power(1 - base_id, -exponent)
def summarize(self, node_id: int, base_id: int) -> Summary:
cached = self.summary_cache[node_id][base_id]
if cached.ready:
return cached
node = self.nodes[node_id]
if node.bit >= 0:
result = Summary(
forward=node.bit,
reverse=node.bit,
ones=node.bit,
ready=True,
)
self.summary_cache[node_id][base_id] = result
return result
left = self.summarize(node.left, base_id)
right = self.summarize(node.right, base_id)
left_length = self.nodes[node.left].length
right_length = self.nodes[node.right].length
base = self.bases[base_id]
result = Summary(
forward=normalize(
left.forward + self.power(base_id, left_length) * right.forward
),
reverse=normalize(
self.power(base_id, right_length) * left.reverse + right.reverse
),
pairs=normalize(
left.pairs
+ right.pairs
+ base * left.reverse % MOD * right.forward
),
ones=normalize(left.ones + right.ones),
ready=True,
)
self.summary_cache[node_id][base_id] = result
return result
def sum_query(
self, x: int, y: int, low: int, high: int, base_id: int
) -> int:
if x > y:
x, y = y, x
key = (x, y, low, high, base_id)
cached = self.sum_query_cache.get(key)
if cached is not None:
return cached
node_x = self.nodes[x]
node_y = self.nodes[y]
maximum = node_x.length + node_y.length - 2
if high < 0 or low > maximum:
result = 0
elif low <= 0 and maximum <= high:
result = (
self.summarize(x, base_id).forward
* self.summarize(y, base_id).forward
% MOD
)
elif node_x.bit >= 0 and node_y.bit >= 0:
result = node_x.bit * node_y.bit if low <= 0 <= high else 0
elif node_x.length >= node_y.length and node_x.bit < 0:
offset = self.nodes[node_x.left].length
result = normalize(
self.sum_query(node_x.left, y, low, high, base_id)
+ self.power(base_id, offset)
* self.sum_query(
node_x.right, y, low - offset, high - offset, base_id
)
)
else:
offset = self.nodes[node_y.left].length
result = normalize(
self.sum_query(x, node_y.left, low, high, base_id)
+ self.power(base_id, offset)
* self.sum_query(
x, node_y.right, low - offset, high - offset, base_id
)
)
self.sum_query_cache[key] = result
return result
def difference_query(
self, x: int, y: int, low: int, high: int, base_id: int
) -> int:
key = (x, y, low, high, base_id)
cached = self.difference_query_cache.get(key)
if cached is not None:
return cached
node_x = self.nodes[x]
node_y = self.nodes[y]
minimum = -(node_x.length - 1)
maximum = node_y.length - 1
if high < minimum or low > maximum:
result = 0
elif low <= minimum and maximum <= high:
result = (
self.summarize(x, 1 - base_id).forward
* self.summarize(y, base_id).forward
% MOD
)
elif node_x.bit >= 0 and node_y.bit >= 0:
result = node_x.bit * node_y.bit if low <= 0 <= high else 0
elif node_x.length >= node_y.length and node_x.bit < 0:
offset = self.nodes[node_x.left].length
result = normalize(
self.difference_query(node_x.left, y, low, high, base_id)
+ self.signed_power(base_id, -offset)
* self.difference_query(
node_x.right, y, low + offset, high + offset, base_id
)
)
else:
offset = self.nodes[node_y.left].length
result = normalize(
self.difference_query(x, node_y.left, low, high, base_id)
+ self.power(base_id, offset)
* self.difference_query(
x, node_y.right, low - offset, high - offset, base_id
)
)
self.difference_query_cache[key] = result
return result
def duplicated_window_sum(
self, standard_index: int, k: int, count: int
) -> int:
if count == 0:
return 0
initial_word = self.prefix(standard_index, k)
initial_value = self.summarize(initial_word, 0).reverse
if count == 1:
return initial_value * initial_value % MOD
transition_count = count - 1
outgoing = self.prefix(standard_index, transition_count)
outgoing_summary = self.summarize(outgoing, 0)
window_power = self.power(0, k)
forward_delta = normalize(
outgoing_summary.reverse - window_power * outgoing_summary.forward
)
reverse_delta = normalize(
outgoing_summary.forward - window_power * outgoing_summary.reverse
)
low_inverse = self.sum_query(
outgoing, outgoing, 0, transition_count - 2, 1
)
high_base = self.sum_query(
outgoing, outgoing, transition_count, 2 * transition_count - 2, 0
)
diagonal_base = self.sum_query(
outgoing, outgoing, transition_count - 1, transition_count - 1, 0
)
diagonal = diagonal_base * self.power(1, transition_count - 1) % MOD
lower_cross = (
self.power(0, transition_count - 1) * low_inverse % MOD
)
upper_cross = (
self.power(1, transition_count - 1) * high_base % MOD
)
delta_pairs = normalize(
normalize(1 + window_power * window_power) * outgoing_summary.pairs
- window_power * normalize(lower_cross + upper_cross)
)
delta_squares = normalize(
normalize(1 + window_power * window_power) * outgoing_summary.ones
- 2 * window_power % MOD * diagonal
)
last_value = normalize(
self.power(0, transition_count) * initial_value + reverse_delta
)
value_delta_sum = normalize(
initial_value * forward_delta + self.bases[1] * delta_pairs
)
numerator = normalize(
initial_value * initial_value
- BASE * BASE % MOD * last_value % MOD * last_value
+ 2 * BASE % MOD * value_delta_sum
+ delta_squares
)
return numerator * mod_inverse(1 - BASE * BASE) % MOD
def solve(self, k: int) -> int:
index = bisect.bisect_right(self.fibonacci, k)
assert index < len(self.fibonacci)
cycle_length = self.fibonacci[index]
duplicate_count = cycle_length - 1 - k
gap_length = duplicate_count + 1
word = self.standard[index]
word_base = self.summarize(word, 0)
word_inverse = self.summarize(word, 1)
cycle_power = self.power(0, cycle_length)
inverse_cycle_power = self.power(1, cycle_length)
all_correlations_base = normalize(
word_base.pairs + cycle_power * word_inverse.pairs
)
all_correlations_inverse = normalize(
word_inverse.pairs + inverse_cycle_power * word_base.pairs
)
prefix_correlations_base = normalize(
self.difference_query(word, word, 1, gap_length, 0)
+ cycle_power
* self.difference_query(
word,
word,
cycle_length - gap_length,
cycle_length - 1,
1,
)
)
prefix_correlations_inverse = normalize(
self.difference_query(word, word, 1, gap_length, 1)
+ inverse_cycle_power
* self.difference_query(
word,
word,
cycle_length - gap_length,
cycle_length - 1,
0,
)
)
used_correlations_base = normalize(
all_correlations_base
- cycle_power * prefix_correlations_inverse
)
used_correlations_inverse = normalize(
all_correlations_inverse
- inverse_cycle_power * prefix_correlations_base
)
inverse_geometric_denominator = mod_inverse(BASE * BASE - 1)
geometric_squares = (
normalize(self.power(0, 2 * k) - 1)
* inverse_geometric_denominator
% MOD
)
full_cycle = normalize(
word_base.ones * geometric_squares
+ 2
* inverse_geometric_denominator
% MOD
* normalize(
self.power(0, 2 * k) * used_correlations_inverse
- used_correlations_base
)
)
return normalize(
full_cycle
- self.duplicated_window_sum(index, k, duplicate_count)
)
def brute_force(k: int) -> int:
older = "0"
newer = "01"
while len(newer) < 10 * k + 20:
older, newer = newer, newer + older
factors = {
newer[start : start + k]
for start in range(len(newer) - k + 1)
}
assert len(factors) == k + 1
result = 0
for factor in factors:
value = 0
for digit in factor:
value = (BASE * value + ord(digit) - ord("0")) % MOD
result = (result + value * value) % MOD
return result
def run_checkpoints() -> None:
solver = FibonacciSubwords()
if brute_force(3) != 20_302:
raise AssertionError("checkpoint failed: Psi(3)")
for k in range(1, 51):
if solver.solve(k) != brute_force(k):
raise AssertionError(
f"checkpoint failed: brute-force comparison for k={k}"
)
if solver.solve(10) != 10_699_667:
raise AssertionError("checkpoint failed: supplied Psi(10)")
print("Validation checkpoints passed.", file=sys.stderr)
def main() -> None:
run_checkpoints()
solver = FibonacciSubwords()
print(solver.solve(TARGET))
if __name__ == "__main__":
main()
Java
import java.util.ArrayList;
import java.util.HashMap;
import java.util.HashSet;
import java.util.List;
import java.util.Map;
import java.util.Set;
public class Euler1006 {
private static final long MOD = 101_001_001L;
private static final long BASE = 10L;
private static final long TARGET = 1_000_000_000_000_000_000L;
private static long normalize(long value) {
value %= MOD;
if (value < 0) {
value += MOD;
}
return value;
}
private static long modPow(long base, long exponent) {
long result = 1;
base = normalize(base);
while (exponent > 0) {
if ((exponent & 1L) != 0) {
result = result * base % MOD;
}
base = base * base % MOD;
exponent >>= 1;
}
return result;
}
private static final class Bezout {
long gcd;
long x;
long y;
Bezout(long gcd, long x, long y) {
this.gcd = gcd;
this.x = x;
this.y = y;
}
}
private static Bezout extendedGcd(long a, long b) {
if (b == 0) {
return new Bezout(a, 1, 0);
}
Bezout next = extendedGcd(b, a % b);
return new Bezout(next.gcd, next.y, next.x - (a / b) * next.y);
}
private static long modInverse(long value) {
Bezout result = extendedGcd(normalize(value), MOD);
if (result.gcd != 1) {
throw new ArithmeticException("a required modular inverse does not exist");
}
return normalize(result.x);
}
private static final class Node {
long length;
int left;
int right;
int bit;
Node(long length, int left, int right, int bit) {
this.length = length;
this.left = left;
this.right = right;
this.bit = bit;
}
}
private static final class Summary {
long forward;
long reverse;
long pairs;
long ones;
Summary(long forward, long reverse, long pairs, long ones) {
this.forward = forward;
this.reverse = reverse;
this.pairs = pairs;
this.ones = ones;
}
}
private record PairKey(int left, int right) {}
private record PrefixKey(int standardIndex, long length) {}
private record QueryKey(int x, int y, long low, long high, int baseId) {}
private static final class FibonacciSubwords {
private final List<Node> nodes = new ArrayList<>();
private final List<Summary[]> summaryCache = new ArrayList<>();
private final Map<PairKey, Integer> concatenationCache = new HashMap<>();
private final Map<PrefixKey, Integer> prefixCache = new HashMap<>();
private final List<Long> fibonacci = new ArrayList<>();
private final List<Integer> standard = new ArrayList<>();
private final long[] bases = new long[2];
private final List<Map<Long, Long>> powerCache =
List.of(new HashMap<>(), new HashMap<>());
private final Map<QueryKey, Long> sumQueryCache = new HashMap<>();
private final Map<QueryKey, Long> differenceQueryCache = new HashMap<>();
private final int zero;
private final int one;
FibonacciSubwords() {
bases[0] = BASE;
bases[1] = modInverse(BASE);
zero = addBit(0);
one = addBit(1);
fibonacci.add(1L);
fibonacci.add(2L);
standard.add(zero);
standard.add(concatenate(zero, one));
while (fibonacci.get(fibonacci.size() - 1) <= TARGET) {
fibonacci.add(
fibonacci.get(fibonacci.size() - 1)
+ fibonacci.get(fibonacci.size() - 2));
standard.add(
concatenate(
standard.get(standard.size() - 1),
standard.get(standard.size() - 2)));
}
}
private int addBit(int bit) {
int id = nodes.size();
nodes.add(new Node(1, -1, -1, bit));
summaryCache.add(new Summary[2]);
return id;
}
private int concatenate(int left, int right) {
if (left < 0) {
return right;
}
if (right < 0) {
return left;
}
PairKey key = new PairKey(left, right);
Integer cached = concatenationCache.get(key);
if (cached != null) {
return cached;
}
int id = nodes.size();
nodes.add(
new Node(
nodes.get(left).length + nodes.get(right).length,
left,
right,
-1));
summaryCache.add(new Summary[2]);
concatenationCache.put(key, id);
return id;
}
private int prefix(int standardIndex, long length) {
require(
length > 0 && length <= fibonacci.get(standardIndex),
"invalid Fibonacci-word prefix");
if (length == fibonacci.get(standardIndex)) {
return standard.get(standardIndex);
}
PrefixKey key = new PrefixKey(standardIndex, length);
Integer cached = prefixCache.get(key);
if (cached != null) {
return cached;
}
require(standardIndex > 0, "prefix recursion reached S_0");
int result;
if (length <= fibonacci.get(standardIndex - 1)) {
result = prefix(standardIndex - 1, length);
} else {
result =
concatenate(
standard.get(standardIndex - 1),
prefix(
standardIndex - 2,
length - fibonacci.get(standardIndex - 1)));
}
prefixCache.put(key, result);
return result;
}
private long power(int baseId, long exponent) {
Map<Long, Long> cache = powerCache.get(baseId);
Long cached = cache.get(exponent);
if (cached != null) {
return cached;
}
long result = modPow(bases[baseId], exponent);
cache.put(exponent, result);
return result;
}
private long signedPower(int baseId, long exponent) {
if (exponent >= 0) {
return power(baseId, exponent);
}
return power(1 - baseId, -exponent);
}
private Summary summarize(int nodeId, int baseId) {
Summary cached = summaryCache.get(nodeId)[baseId];
if (cached != null) {
return cached;
}
Node node = nodes.get(nodeId);
Summary result;
if (node.bit >= 0) {
result = new Summary(node.bit, node.bit, 0, node.bit);
summaryCache.get(nodeId)[baseId] = result;
return result;
}
Summary left = summarize(node.left, baseId);
Summary right = summarize(node.right, baseId);
long leftLength = nodes.get(node.left).length;
long rightLength = nodes.get(node.right).length;
long base = bases[baseId];
result =
new Summary(
normalize(left.forward + power(baseId, leftLength) * right.forward),
normalize(power(baseId, rightLength) * left.reverse + right.reverse),
normalize(
left.pairs
+ right.pairs
+ base * left.reverse % MOD * right.forward),
normalize(left.ones + right.ones));
summaryCache.get(nodeId)[baseId] = result;
return result;
}
private long sumQuery(int x, int y, long low, long high, int baseId) {
if (x > y) {
int temporary = x;
x = y;
y = temporary;
}
QueryKey key = new QueryKey(x, y, low, high, baseId);
Long cached = sumQueryCache.get(key);
if (cached != null) {
return cached;
}
Node nodeX = nodes.get(x);
Node nodeY = nodes.get(y);
long maximum = nodeX.length + nodeY.length - 2;
long result;
if (high < 0 || low > maximum) {
result = 0;
} else if (low <= 0 && maximum <= high) {
result = summarize(x, baseId).forward * summarize(y, baseId).forward % MOD;
} else if (nodeX.bit >= 0 && nodeY.bit >= 0) {
result = low <= 0 && 0 <= high ? (long) nodeX.bit * nodeY.bit : 0;
} else if (nodeX.length >= nodeY.length && nodeX.bit < 0) {
long offset = nodes.get(nodeX.left).length;
result =
normalize(
sumQuery(nodeX.left, y, low, high, baseId)
+ power(baseId, offset)
* sumQuery(
nodeX.right,
y,
low - offset,
high - offset,
baseId));
} else {
long offset = nodes.get(nodeY.left).length;
result =
normalize(
sumQuery(x, nodeY.left, low, high, baseId)
+ power(baseId, offset)
* sumQuery(
x,
nodeY.right,
low - offset,
high - offset,
baseId));
}
sumQueryCache.put(key, result);
return result;
}
private long differenceQuery(int x, int y, long low, long high, int baseId) {
QueryKey key = new QueryKey(x, y, low, high, baseId);
Long cached = differenceQueryCache.get(key);
if (cached != null) {
return cached;
}
Node nodeX = nodes.get(x);
Node nodeY = nodes.get(y);
long minimum = -(nodeX.length - 1);
long maximum = nodeY.length - 1;
long result;
if (high < minimum || low > maximum) {
result = 0;
} else if (low <= minimum && maximum <= high) {
result =
summarize(x, 1 - baseId).forward
* summarize(y, baseId).forward
% MOD;
} else if (nodeX.bit >= 0 && nodeY.bit >= 0) {
result = low <= 0 && 0 <= high ? (long) nodeX.bit * nodeY.bit : 0;
} else if (nodeX.length >= nodeY.length && nodeX.bit < 0) {
long offset = nodes.get(nodeX.left).length;
result =
normalize(
differenceQuery(nodeX.left, y, low, high, baseId)
+ signedPower(baseId, -offset)
* differenceQuery(
nodeX.right,
y,
low + offset,
high + offset,
baseId));
} else {
long offset = nodes.get(nodeY.left).length;
result =
normalize(
differenceQuery(x, nodeY.left, low, high, baseId)
+ power(baseId, offset)
* differenceQuery(
x,
nodeY.right,
low - offset,
high - offset,
baseId));
}
differenceQueryCache.put(key, result);
return result;
}
private long duplicatedWindowSum(int standardIndex, long k, long count) {
if (count == 0) {
return 0;
}
int initialWord = prefix(standardIndex, k);
long initialValue = summarize(initialWord, 0).reverse;
if (count == 1) {
return initialValue * initialValue % MOD;
}
long transitionCount = count - 1;
int outgoing = prefix(standardIndex, transitionCount);
Summary outgoingSummary = summarize(outgoing, 0);
long windowPower = power(0, k);
long forwardDelta =
normalize(outgoingSummary.reverse - windowPower * outgoingSummary.forward);
long reverseDelta =
normalize(outgoingSummary.forward - windowPower * outgoingSummary.reverse);
long lowInverse = sumQuery(outgoing, outgoing, 0, transitionCount - 2, 1);
long highBase =
sumQuery(
outgoing,
outgoing,
transitionCount,
2 * transitionCount - 2,
0);
long diagonalBase =
sumQuery(
outgoing,
outgoing,
transitionCount - 1,
transitionCount - 1,
0);
long diagonal = diagonalBase * power(1, transitionCount - 1) % MOD;
long lowerCross = power(0, transitionCount - 1) * lowInverse % MOD;
long upperCross = power(1, transitionCount - 1) * highBase % MOD;
long deltaPairs =
normalize(
normalize(1 + windowPower * windowPower) * outgoingSummary.pairs
- windowPower * normalize(lowerCross + upperCross));
long deltaSquares =
normalize(
normalize(1 + windowPower * windowPower) * outgoingSummary.ones
- 2 * windowPower % MOD * diagonal);
long lastValue =
normalize(power(0, transitionCount) * initialValue + reverseDelta);
long valueDeltaSum =
normalize(initialValue * forwardDelta + bases[1] * deltaPairs);
long numerator =
normalize(
initialValue * initialValue
- BASE * BASE % MOD * lastValue % MOD * lastValue
+ 2 * BASE % MOD * valueDeltaSum
+ deltaSquares);
return numerator * modInverse(1 - BASE * BASE) % MOD;
}
long solve(long k) {
int index = upperBound(fibonacci, k);
require(index < fibonacci.size(), "target exceeds the prepared Fibonacci words");
long cycleLength = fibonacci.get(index);
long duplicateCount = cycleLength - 1 - k;
long gapLength = duplicateCount + 1;
int word = standard.get(index);
Summary wordBase = summarize(word, 0);
Summary wordInverse = summarize(word, 1);
long cyclePower = power(0, cycleLength);
long inverseCyclePower = power(1, cycleLength);
long allCorrelationsBase =
normalize(wordBase.pairs + cyclePower * wordInverse.pairs);
long allCorrelationsInverse =
normalize(wordInverse.pairs + inverseCyclePower * wordBase.pairs);
long prefixCorrelationsBase =
normalize(
differenceQuery(word, word, 1, gapLength, 0)
+ cyclePower
* differenceQuery(
word,
word,
cycleLength - gapLength,
cycleLength - 1,
1));
long prefixCorrelationsInverse =
normalize(
differenceQuery(word, word, 1, gapLength, 1)
+ inverseCyclePower
* differenceQuery(
word,
word,
cycleLength - gapLength,
cycleLength - 1,
0));
long usedCorrelationsBase =
normalize(allCorrelationsBase - cyclePower * prefixCorrelationsInverse);
long usedCorrelationsInverse =
normalize(
allCorrelationsInverse
- inverseCyclePower * prefixCorrelationsBase);
long inverseGeometricDenominator = modInverse(BASE * BASE - 1);
long geometricSquares =
normalize(power(0, 2 * k) - 1) * inverseGeometricDenominator % MOD;
long fullCycle =
normalize(
wordBase.ones * geometricSquares
+ 2
* inverseGeometricDenominator
% MOD
* normalize(
power(0, 2 * k)
* usedCorrelationsInverse
- usedCorrelationsBase));
return normalize(
fullCycle - duplicatedWindowSum(index, k, duplicateCount));
}
}
private static int upperBound(List<Long> values, long target) {
int low = 0;
int high = values.size();
while (low < high) {
int middle = (low + high) >>> 1;
if (values.get(middle) <= target) {
low = middle + 1;
} else {
high = middle;
}
}
return low;
}
private static long bruteForce(int k) {
String older = "0";
String newer = "01";
while (newer.length() < 10 * k + 20) {
String next = newer + older;
older = newer;
newer = next;
}
Set<String> factors = new HashSet<>();
for (int start = 0; start + k <= newer.length(); ++start) {
factors.add(newer.substring(start, start + k));
}
require(factors.size() == k + 1, "factor complexity for k=" + k);
long result = 0;
for (String factor : factors) {
long value = 0;
for (int i = 0; i < factor.length(); ++i) {
value = (BASE * value + factor.charAt(i) - '0') % MOD;
}
result = (result + value * value) % MOD;
}
return result;
}
private static void require(boolean condition, String message) {
if (!condition) {
throw new AssertionError("checkpoint failed: " + message);
}
}
private static void runCheckpoints() {
FibonacciSubwords solver = new FibonacciSubwords();
require(bruteForce(3) == 20_302, "Psi(3)");
for (int k = 1; k <= 50; ++k) {
require(
solver.solve(k) == bruteForce(k),
"brute-force comparison for k=" + k);
}
require(solver.solve(10) == 10_699_667, "supplied Psi(10)");
System.err.println("Validation checkpoints passed.");
}
public static void main(String[] args) {
runCheckpoints();
FibonacciSubwords solver = new FibonacciSubwords();
System.out.println(solver.solve(TARGET));
}
}