Skip to content

[CodeChef] XOR Tree: a trie of the current root path, one query per leaf

Published 6 October 2026

Problem: CodeChef — XOR Tree (difficulty 2697).

Not accepted yet

This scores 60/100. The samples pass, but the later subtasks (from 13 on) give a runtime error. I need to revisit the BinaryTrie code. See Status at the end.

The problem

There’s a tree with \(N\) vertices, rooted at vertex \(1\). Vertex \(i\) has value \(A_i\).

Alice starts at the root. At a vertex \(u\) with \(c\) children, she moves to one of them, each with probability \(\frac{1}{c}\). She stops at a leaf.

Say she visited \(u_1, u_2, \dots, u_k\). She forgets exactly one of them, and her score is the XOR of the other \(k - 1\) values. She forgets the vertex that gives the largest score. Find her expected score modulo \(10^9 + 7\), written as \(P \cdot Q^{-1}\).

Constraints: \(T \le 10^4\), \(N \le 5 \cdot 10^5\), \(1 \le A_i \le 10^9\), \(\sum N \le 5 \cdot 10^5\). Time limit 3 s, memory 1.5 GB.

Sample 3: the root has value \(1\) and three leaf children with values \(2, 3, 4\). Each path is \(1 \to \text{leaf}\), and Alice forgets the \(1\) every time, so the scores are \(2, 3, 4\) and the answer is \(\frac{2 + 3 + 4}{3} = 3\).

Forgetting one vertex is a max-XOR query

Let \(S\) be the XOR of every value on the path. Forgetting \(u_i\) leaves

\[S \oplus A_{u_i},\]

because XORing \(A_{u_i}\) in a second time cancels it. So the best score on a path is

\[\max_{i} \; (S \oplus A_{u_i}),\]

which is the classic “max XOR of a fixed number with any value in a set” question. A binary trie answers it in \(O(\log A)\): walk from the top bit down, and at each bit go to the child with the opposite bit if one exists.

Keep the trie equal to the current path

Do a DFS from the root and keep the trie holding exactly the values on the root-to-current path:

  • on entering a vertex, insert its value and XOR it into path_xor;
  • at a leaf, query max_xor(path_xor);
  • on leaving, remove the value again.

Removal is a count decrement on each node along the value’s path. The query only follows children whose count is positive, so removed values are ignored. Each vertex is inserted and removed once, and each leaf makes one query, so the total is \(O(N \log A)\).

The expectation

The probability of reaching a leaf is the product of \(\frac{1}{c}\) over the vertices above it. Pass it down the DFS: a child gets probability * inverse(c). The answer is

\[\sum_{\text{leaf } \ell} \Pr[\ell] \cdot \text{best}(\ell).\]

Where I went wrong

I printed a float. My first version passed probabilities as 1.0 / c and printed 2.0000000000. The answer has to be \(P \cdot Q^{-1} \bmod (10^9 + 7)\), so the probability needs to stay a modular number. \(\frac{1}{c}\) becomes pow(c, MOD - 2, MOD), by Fermat’s little theorem, since \(10^9 + 7\) is prime. The best XOR is an integer below \(2^{30}\), so it can be multiplied in as is.

bits only needs to be 29. \(A_i \le 10^9 < 2^{30}\), so bits \(29\) down to \(0\) are enough.

Code

This is my current submission. It passes the samples (2, 2, 3) and scores 60/100.

import sys

MOD = 10**9 + 7
sys.setrecursionlimit(1_000_000)


class BinaryTrie:
    def __init__(self, bits=29):  # Ai <= 10^9 < 2^30
        self.bits = bits
        self.children = [[-1, -1]]
        self.count = [0]

    def _new_node(self):
        self.children.append([-1, -1])
        self.count.append(0)
        return len(self.children) - 1

    def add(self, value, delta):
        node = 0
        self.count[node] += delta

        for bit in range(self.bits, -1, -1):
            b = (value >> bit) & 1
            nxt = self.children[node][b]

            if nxt == -1:
                nxt = self._new_node()
                self.children[node][b] = nxt

            node = nxt
            self.count[node] += delta

    def max_xor(self, value):
        node = 0
        result = 0

        for bit in range(self.bits, -1, -1):
            b = (value >> bit) & 1
            preferred = b ^ 1
            nxt = self.children[node][preferred]

            if nxt != -1 and self.count[nxt] > 0:
                result |= 1 << bit
                node = nxt
            else:
                node = self.children[node][b]

        return result


def expected_max_path_xor(values, graph, root=0):
    trie = BinaryTrie()

    def dfs(node, parent, path_xor, probability):
        trie.add(values[node], 1)
        path_xor ^= values[node]

        children = [v for v in graph[node] if v != parent]

        if not children:
            answer = probability * trie.max_xor(path_xor) % MOD
        else:
            child_probability = probability * pow(len(children), MOD - 2, MOD) % MOD
            answer = sum(
                dfs(child, node, path_xor, child_probability)
                for child in children
            ) % MOD

        trie.add(values[node], -1)
        return answer

    return dfs(root, -1, 0, 1)


def solve():
    data = list(map(int, sys.stdin.buffer.read().split()))
    pos = 0
    test_cases = data[pos]
    pos += 1
    answers = []

    for _ in range(test_cases):
        n = data[pos]
        pos += 1

        values = data[pos:pos + n]
        pos += n

        graph = [[] for _ in range(n)]
        for _ in range(n - 1):
            u = data[pos] - 1
            v = data[pos + 1] - 1
            pos += 2
            graph[u].append(v)
            graph[v].append(u)

        answers.append(str(expected_max_path_xor(values, graph)))

    print("\n".join(answers))


if __name__ == "__main__":
    solve()

Status: not accepted

Still not accepted: 60/100, runtime error on the last subtasks. I need to revisit the BinaryTrie code.

What I know so far:

  • A chain of 200,000 vertices reproduces it locally. It crashes with RecursionError: Stack overflow even with setrecursionlimit(1_000_000). Raising the limit doesn’t make the C stack any bigger, and a chain makes the DFS (plus the sum(...) generator frame on every level) about \(N\) calls deep.
  • The trie can grow to about \(30N = 1.5 \cdot 10^7\) nodes, and each node is a separate [-1, -1] Python list. That’s a lot of memory and time, even if it fits in 1.5 GB.

  • Rewrite the DFS iteratively, with an explicit stack and an enter/exit flag per vertex.

  • Rework BinaryTrie to use flat arrays (left, right, count) instead of a list per node.
  • Resubmit and update this post.