题解:P17240 [IOI 2026] 弹球机 / ballmachine

· · 题解

题意简述

机器中藏着一棵有根树,共有 N 个节点。根的编号固定为 N-1M 个叶子的编号固定为 0,1,\ldots,M-1,只有内部节点允许重新编号。

我们可以进行两种操作:

我们的任务是还原一棵与原树同构的树,同时保持根和每个叶子的编号不变。

collect 的调用次数为 K,使用过的最大球值为 B,得分所用的资源指标为 C=K+B。满分要求 C\le 44

思路

整套做法可以分成三个阶段:

  1. 先用一次实验把树拆成若干条链,并求出每条链的长度。
  2. 再设计球值,使一次 collect() 的返回序列能够还原整棵无序树。
  3. 最后分批给较难识别的叶子编号,并把不同实验中还原出的树对齐。

前两步负责恢复“树长什么样”,第三步负责回答“每个叶子具体是哪一个”。

1. 先把树拆成链

0,1,\ldots,M-1 的顺序处理叶子。轮到叶子 i 时,不断执行 insert(i,0),直到第一次失败。记这一轮成功插入了 l_i 个球。

来看这些球会停在哪里。处理 i 之前,已经有球的节点连成一棵包含根的子树。第一次从 i 插入的球,会停在路径 i\to\text{root} 上最高的空节点;再插入一个球,它就停在前一个球的正下方。这个过程不断向叶子延伸,最后一个球恰好停在叶子 i

所以,第 i 轮新填入的节点一定连成一条从上到下的链。把它称为链 i,其长度就是 l_i

为什么所有这些链恰好覆盖整棵树?任选一个节点 u,在 u 的子树中找编号最小的叶子 i。处理编号小于 i 的叶子时,它们都不在 u 的子树中,因此不可能经过 u。等到处理 i 时,节点 u 仍为空,于是会被划入链 i

这说明每个节点都会进入某条链。另一方面,一个节点一旦被占据,以后便不会再次进入其他链,所以这些链两两不交。

还可以得到三个很有用的性质:

所有叶子都插到失败后,整棵树已经被球填满。此时调用一次 collect(),返回序列的长度就是 N。这一次只用来求 N 和所有 l_i,不需要分析返回值的具体顺序。

l_i=1,链 i 只有叶子自己,称为单点链;若 l_i\ge2,称为长链。设长链有 D 条,单点链有 S 条,则 D+S=M

2. 一次实验里怎样给节点做标记

接下来希望达到一个目标:填满整棵树后,只看 collect() 返回的球值序列,就能重新拼出树的结构。

先按照叶子编号从小到大排列所有长链,再依次编号为 r=0,1,\ldots,D-1。令

H=\max\left(1,\left\lceil\frac D{11}\right\rceil\right).

在之后的每次实验中,仍按叶子编号依次插入,并完整填满每一条链。我们把球值分成互不重叠的四段:

节点在链中的位置 使用的球值 作用
长链 r 的第一个节点 \lfloor r/11\rfloor 保存 r 除以 11 的商
单点链 [H,H+11) 中的值 标记单点叶子
长链 r 的最后一个节点 H+11+(r\bmod 11) 保存 r 除以 11 的余数
长链的其他节点 H+22 表示继续沿当前长链向下

四类值分别位于 [0,H)[H,H+11)[H+11,H+22) 和单独的值 H+22,所以看到一个球值时,可以立刻判断它属于哪一类节点。

一条长链的开头保存商,结尾保存余数。把两者放在一起,便能唯一还原长链编号:r=11\lfloor r/11\rfloor+(r\bmod 11)。因此也能知道这条链末端是哪一个叶子。

3. 如何从返回序列还原树

现在解释为什么上面的标记足够恢复树。

假设当前正在访问某条长链上的节点。它的儿子可能有三类:

collect() 会按球值递增访问儿子。因此,它会先完整访问挂在当前节点旁边的链和单点叶子,最后才沿当前长链继续向下。返回序列由此具有类似括号嵌套的结构,可以用一个栈从左到右解析:

栈中每一项记录一条尚未解析完的长链,以及这条链当前最下面的节点。新读到的节点应挂在哪里,也就随之确定了。

读到长链结尾时,再把链首记录的商和链尾记录的余数组合起来,就能确定这条长链对应的叶子编号。这样,一次实验便可恢复完整的无序树,并认出所有长链末端的叶子。

这里的“无序树”非常重要:我们只知道一个节点有哪些儿子,不认为这些儿子之间存在固定的先后次序。

4. 单点叶子为什么要分批识别

长链可以利用“链首加链尾”记录编号,单点链却只有一个节点,只能用这一个节点的球值携带信息。

单点链所用的区间 [H,H+11) 一共有 11 个值:

于是,一次实验最多识别 10 个单点叶子。把 S 个单点叶子每 10 个分成一批,共需 \lceil S/10\rceil 次实验。第一批实验可以同时承担上一节的结构恢复工作,不会额外浪费一次 collect()

如果 S=0,没有需要编号的单点叶子,另外进行一次只编码长链的结构实验即可。

5. 为什么不同实验不能按位置对齐

这一步是实现中最容易出错的地方。

每次实验都会独立产生一段 collect() 序列。即使两次实验填满的是同一棵树,也不能认为两个序列中下标相同的元素来自同一个节点。原因是:同一个父节点下,若几个儿子的球值相同,它们的访问顺序可以任意改变;而我们在不同批次中改变了部分球值,这种顺序就更可能发生变化。

因此,正确做法是:每次都独立解析出一棵无序树,再根据树的结构把它映射到第一次实验得到的树。以下把第一次实验得到的树称为基准树

具体来说,对每棵树自底向上计算子树签名。节点 u 的签名包含:

儿子没有固定顺序,所以先对它们的签名排序,再把整个信息离散化为一个整数。当前实验的树和基准树共用同一张离散化表,这样,同构且已知叶子编号相同的两棵子树一定得到相同签名。

随后从两棵树的根开始向下匹配。对于一对已经对应上的节点,把两边的儿子都按签名排序,再把签名相同的儿子配在一起。因为两边本来就是同一棵无序树,且所有已知的长链叶子编号相同,所以每种儿子签名的出现次数也一定相同。

6. 完全对称的位置不必强行区分

上一节还有一个小问题:如果同一个父节点下面有两棵签名完全相同的子树,该把左边的第一棵配到右边的哪一棵?

答案是任选一棵都可以。既然它们的结构以及其中所有已知叶子编号都相同,那么交换这两棵子树不会改变当前已知的整棵树。题目又允许内部节点重新编号,所以没有必要强行区分这两个位置。

不过,任意配对会带来一个后果:某个本批已编号的单点叶子映射回基准树后,它的父节点可能不是唯一的。我们真正需要记录的,不是某个数组下标,而是“这个父节点属于哪一组完全对称的位置”。

例如,某个节点下面挂着两个结构完全相同、目前都没有编号的分支。交换这两个分支后,所有已知信息都不变。那么实验只能确定目标叶子的父节点属于这两个对称位置之一,却无法也无需判断究竟是哪一个。

形式化地说,如果在保持根和所有已知长链叶子编号不变的前提下,可以通过树的一个自同构把位置 u 变成位置 v,就把 u,v 放入同一个自同构轨道。它就是上面所说的“一组无法区分的位置”。

代码通过反复染色求出这些轨道:

为什么稳定后的颜色正好表示自同构轨道?经过 t 轮后,一个节点的颜色已经记录了它周围距离不超过 t 的结构和已知标号。树没有环;当 t\ge N 时,这些信息已经覆盖整棵树。

若两个节点仍然同色,就存在一个保持所有初始颜色的自同构把其中一个映到另一个。反过来,任何保持初始颜色的自同构,在每一轮染色后都会把同色节点映到同色节点。因此,最终同色当且仅当属于同一轨道。

处理每一批单点叶子时,我们先用子树签名把当前实验映射到基准树,再记录这个单点叶子的父节点轨道

全部批次完成后,把基准树中尚未编号的单点位置也按父节点轨道分类。对于每个单点叶子,从对应类别中取一个未使用的位置并赋上它的编号。同一轨道内的位置能够互相交换,因此类别内部怎样一一配对都不会改变树的同构关系。

7. 完整算法

按实际执行顺序整理如下:

  1. 依次把每个叶子插到失败,得到所有链长 l_i;随后 collect(),用返回序列的长度求出 N
  2. 将链分成长链和单点链,并给所有长链编号。
  3. 按四段值域填满整棵树,调用 collect(),用栈解析出基准树,同时确定所有长链叶子的编号。
  4. 在基准树上反复染色,求出所有节点的自同构轨道。
  5. 每次选择至多 10 个单点叶子赋予不同球值;解析当前实验,通过子树签名映射到基准树,并记录这些叶子的父节点轨道。
  6. 按父节点轨道把所有单点叶子分配给基准树中的匿名单点位置。
  7. 保留根和叶子的规定编号,为其余内部节点依次编号,最后返回每个非根节点的父亲。

8. 为什么算法正确

链分解阶段中,每个节点会在其子树内编号最小的叶子被处理时首次占据,因此所有链两两不交并覆盖整棵树,测得的 l_i 就是各链的真实长度。

结构实验中,四类节点的球值范围互不相交;对于每条长链,链首和链尾又共同唯一确定它的编号。球值大小保证所有侧链都在当前长链继续向下之前访问,因此栈解析得到的父子关系与原树完全一致,并能正确标出每个长链叶子。

不同实验之间只按照无序树的子树签名匹配,不依赖 collect() 中同值儿子的顺序。若一种签名出现多次,任意配对只是在完全同构的子树之间进行交换。父节点的自同构轨道不会因这种交换而改变,所以每个单点叶子记录到的父节点轨道一定正确。

最后,每个单点叶子都被放入父节点轨道相同且尚未使用的位置。同一轨道中的位置可以通过保持所有已知编号的自同构互换,因此这种分配不会改变树的同构关系。故最终返回的树与原树同构,并保持根和所有叶子的编号不变。

资源与复杂度分析

insert 次数

第一次测链长时,每个节点恰好被成功插入一次,并且每个叶子还会产生一次用于结束当前循环的失败插入,共调用 N+Minsert

之后每次结构实验恰好成功插入 N 个球。令 E=\max(1,\lceil S/10\rceil) 为结构实验次数,则总调用数不超过

N(1+E)+M\le1000\times21+200<500000.

资源指标 C

S>0 时,第一次 collect() 用来结束链长实验,之后每批单点叶子调用一次。因此

K=1+\left\lceil\frac S{10}\right\rceil,\qquad B=H+22.

下面证明 \lceil S/10\rceil+H\le21。若 D=0,则 H=1,结论直接成立。若 D>0,令 a=\lceil S/10\rceil。由 S\ge10a-9D+S\le200,可得

D\le209-10a\le231-11a=11(21-a).

其中 a\le20,所以 H=\lceil D/11\rceil\le21-a。于是

C=K+B=23+\left\lceil\frac S{10}\right\rceil+H\le44.

S=0 时,只需调用一次链长实验和一次结构实验,所以 K=2。此时 D=M,从而

C=2+\left\lceil\frac M{11}\right\rceil+22\le43.

因此,所有情况都满足满分限制 C\le44

时间与空间复杂度

设结构实验次数为 E=\max(1,\lceil S/10\rceil)。解析实验、计算子树签名和匹配无序树共需 O(EN\log N) 时间。

求自同构轨道至多细化 N 轮,每轮需要 O(N\log N) 时间。因此,总时间复杂度为 O(N^2\log N+EN\log N),空间复杂度为 O(N+M)

实现中的所有树遍历都采用迭代写法。即使整棵树退化成一条长度为 N 的链,也不会因为递归层数过深而栈溢出。

参考代码

#include "ballmachine.h"

#include <algorithm>
#include <map>
#include <numeric>
#include <utility>
#include <vector>

using namespace std;

namespace {

constexpr int BASE = 11;
constexpr int BATCH = BASE - 1;

struct Node {
    int code = -1;
    int parent = -1;
    bool singleton = false;
    vector<int> children;
};

struct Frame {
    int high;
    vector<int> nodes;
};

}  // namespace

vector<int> find_structure(int M) {
    vector<int> length(M);

    for (int leaf = 0; leaf < M; ++leaf) {
        while (insert(leaf, 0)) ++length[leaf];
    }
    const int N = static_cast<int>(collect().size());

    vector<int> long_leaves, short_leaves;
    for (int leaf = 0; leaf < M; ++leaf) {
        (length[leaf] == 1 ? short_leaves : long_leaves).push_back(leaf);
    }

    vector<int> rank(M, -1);
    for (int i = 0; i < static_cast<int>(long_leaves.size()); ++i) {
        rank[long_leaves[i]] = i;
    }

    const int high_count = max(
        1, (static_cast<int>(long_leaves.size()) + BASE - 1) / BASE);
    const int singleton_begin = high_count;
    const int tail_begin = singleton_begin + BASE;
    const int middle_value = tail_begin + BASE;
    const int anonymous_value = singleton_begin;

    auto fill_experiment = [&](int short_begin) {
        for (int leaf = 0; leaf < M; ++leaf) {
            if (length[leaf] == 1) {
                int value = anonymous_value;
                auto it = lower_bound(short_leaves.begin(),
                                      short_leaves.end(), leaf);
                int at = static_cast<int>(it - short_leaves.begin());
                if (at < static_cast<int>(short_leaves.size()) &&
                    short_leaves[at] == leaf && short_begin <= at &&
                    at < short_begin + BATCH) {
                    value = singleton_begin + 1 + (at - short_begin);
                }
                insert(leaf, value);
            } else {
                int r = rank[leaf];
                insert(leaf, r / BASE);
                for (int k = 1; k + 1 < length[leaf]; ++k) {
                    insert(leaf, middle_value);
                }
                insert(leaf, tail_begin + r % BASE);
            }
        }
        return collect();
    };

    auto parse = [&](const vector<int>& sequence,
                     vector<Node>& nodes,
                     vector<int>& chain_leaf,
                     vector<pair<int, int>>& marked) {
        nodes.clear();
        chain_leaf.clear();
        marked.clear();
        vector<Frame> stack;

        auto add_node = [&](int code, bool singleton = false) {
            int u = static_cast<int>(nodes.size());
            nodes.push_back(Node{});
            nodes[u].code = code;
            nodes[u].singleton = singleton;
            if (!stack.empty()) {
                nodes[u].parent = stack.back().nodes.back();
                nodes[nodes[u].parent].children.push_back(u);
            }
            return u;
        };

        for (int value : sequence) {
            if (0 <= value && value < high_count) {
                int u = add_node(value);
                stack.push_back({value, {u}});
            } else if (singleton_begin <= value &&
                       value < singleton_begin + BASE) {
                int u = add_node(value, true);
                if (value != anonymous_value) {
                    marked.push_back(
                        {value - singleton_begin - 1, u});
                }
            } else if (tail_begin <= value &&
                       value < tail_begin + BASE) {
                int low = value - tail_begin;
                int u = add_node(value);
                stack.back().nodes.push_back(u);
                int r = stack.back().high * BASE + low;
                chain_leaf.push_back(long_leaves[r]);
                stack.pop_back();
            } else {
                int u = add_node(value, stack.empty());
                if (!stack.empty()) stack.back().nodes.push_back(u);
            }
        }
    };

    vector<int> answer_parent(N, -1), answer_leaf_at_node(N, -1);
    vector<char> answer_long_leaf(N, false);
    vector<Node> answer_nodes;
    vector<int> ignored_chain_leaf;
    vector<pair<int, int>> dummy_marked;

    vector<int> first_sequence =
        fill_experiment(short_leaves.empty() ? M + 1 : 0);
    parse(first_sequence, answer_nodes,
          ignored_chain_leaf, dummy_marked);

    for (int u = 0; u < N; ++u) {
        int value = answer_nodes[u].code;
        if (tail_begin <= value && value < tail_begin + BASE &&
            answer_nodes[u].parent != -1) {
            int p = answer_nodes[u].parent;
            if (answer_nodes[p].code == middle_value ||
                (0 <= answer_nodes[p].code &&
                 answer_nodes[p].code < high_count)) {
                int x = u;
                while (answer_nodes[x].parent != -1 &&
                       answer_nodes[x].code >= high_count) {
                    x = answer_nodes[x].parent;
                }
                if (0 <= answer_nodes[x].code &&
                    answer_nodes[x].code < high_count) {
                    int r = answer_nodes[x].code * BASE
                            + value - tail_begin;
                    int expected =
                        r < static_cast<int>(long_leaves.size())
                            ? length[long_leaves[r]] : -1;
                    int actual = 1;
                    for (int y = answer_nodes[u].parent;
                         y != answer_nodes[x].parent;
                         y = answer_nodes[y].parent) {
                        ++actual;
                    }
                    if (r < static_cast<int>(long_leaves.size()) &&
                        actual == expected) {
                        answer_leaf_at_node[u] = long_leaves[r];
                        answer_long_leaf[u] = true;
                    }
                }
            }
        }
    }

    int root = -1;
    for (int u = 0; u < N; ++u) {
        if (answer_nodes[u].parent == -1) root = u;
    }

    vector<int> orbit(N);
    for (int u = 0; u < N; ++u) {
        if (answer_long_leaf[u]) {
            orbit[u] = 1000 + answer_leaf_at_node[u];
        } else {
            orbit[u] = answer_nodes[u].singleton ? -2 : -1;
        }
    }

    for (int iteration = 0; iteration < N; ++iteration) {
        map<vector<int>, int> ids;
        vector<int> next_orbit(N);
        for (int u = 0; u < N; ++u) {
            vector<int> key = {
                orbit[u],
                answer_nodes[u].parent == -1
                    ? -3 : orbit[answer_nodes[u].parent]
            };
            vector<int> child;
            for (int v : answer_nodes[u].children) {
                child.push_back(orbit[v]);
            }
            sort(child.begin(), child.end());
            key.insert(key.end(), child.begin(), child.end());
            auto it = ids.find(key);
            if (it == ids.end()) {
                it = ids.emplace(
                    key, static_cast<int>(ids.size())).first;
            }
            next_orbit[u] = it->second;
        }
        if (next_orbit == orbit) break;
        orbit.swap(next_orbit);
    }

    vector<int> singleton_leaf_orbit(M, -1);

    auto make_mapping = [&](const vector<Node>& cur,
                            const vector<int>& cur_leaf_at_node,
                            vector<int>& to_answer) {
        map<vector<int>, int> ids;

        auto obtain_sig = [&](const vector<Node>& tree,
                              const vector<int>& marked_leaf) {
            vector<int> sig(N), order, traversal;
            int tree_root = -1;
            for (int u = 0; u < N; ++u) {
                if (tree[u].parent == -1) tree_root = u;
            }

            traversal.push_back(tree_root);
            while (!traversal.empty()) {
                int u = traversal.back();
                traversal.pop_back();
                order.push_back(u);
                for (int v : tree[u].children) {
                    traversal.push_back(v);
                }
            }
            reverse(order.begin(), order.end());

            for (int u : order) {
                int type = tree[u].singleton ? 1 : 0;
                int mark = marked_leaf[u] == -1
                               ? -1 : marked_leaf[u];
                vector<int> key = {type, mark};
                vector<int> child;
                for (int v : tree[u].children) {
                    child.push_back(sig[v]);
                }
                sort(child.begin(), child.end());
                key.insert(key.end(), child.begin(), child.end());
                auto it = ids.find(key);
                if (it == ids.end()) {
                    it = ids.emplace(
                        key, static_cast<int>(ids.size())).first;
                }
                sig[u] = it->second;
            }
            return sig;
        };

        vector<int> cur_sig = obtain_sig(cur, cur_leaf_at_node);
        vector<int> answer_sig =
            obtain_sig(answer_nodes, answer_leaf_at_node);
        cur_sig = obtain_sig(cur, cur_leaf_at_node);

        int cur_root = -1;
        for (int u = 0; u < N; ++u) {
            if (cur[u].parent == -1) cur_root = u;
        }

        to_answer.assign(N, -1);
        vector<pair<int, int>> stack = {{cur_root, root}};
        while (!stack.empty()) {
            auto [u, a] = stack.back();
            stack.pop_back();
            to_answer[u] = a;

            vector<int> left = cur[u].children;
            vector<int> right = answer_nodes[a].children;
            auto cur_less = [&](int x, int y) {
                return cur_sig[x] < cur_sig[y];
            };
            auto ans_less = [&](int x, int y) {
                return answer_sig[x] < answer_sig[y];
            };
            sort(left.begin(), left.end(), cur_less);
            sort(right.begin(), right.end(), ans_less);

            for (int i = 0;
                 i < static_cast<int>(left.size()); ++i) {
                stack.push_back({left[i], right[i]});
            }
        }
    };

    for (int begin = 0;
         begin < static_cast<int>(short_leaves.size());
         begin += BATCH) {
        vector<int> sequence = begin == 0
                                   ? first_sequence
                                   : fill_experiment(begin);
        vector<Node> cur;
        vector<int> cur_chain_leaf;
        vector<pair<int, int>> marked;
        parse(sequence, cur, cur_chain_leaf, marked);

        vector<int> cur_leaf_at_node(N, -1);
        for (int u = 0; u < N; ++u) {
            int value = cur[u].code;
            if (!(tail_begin <= value &&
                  value < tail_begin + BASE) ||
                cur[u].parent == -1) {
                continue;
            }

            int x = u;
            while (cur[x].parent != -1 &&
                   cur[x].code >= high_count) {
                x = cur[x].parent;
            }
            if (0 <= cur[x].code &&
                cur[x].code < high_count) {
                int r = cur[x].code * BASE
                        + value - tail_begin;
                int expected =
                    r < static_cast<int>(long_leaves.size())
                        ? length[long_leaves[r]] : -1;
                int actual = 1;
                for (int y = cur[u].parent;
                     y != cur[x].parent;
                     y = cur[y].parent) {
                    ++actual;
                }
                if (r < static_cast<int>(long_leaves.size()) &&
                    actual == expected) {
                    cur_leaf_at_node[u] = long_leaves[r];
                }
            }
        }

        vector<int> answer_match_leaf(N, -1);
        for (int u = 0; u < N; ++u) {
            if (answer_long_leaf[u]) {
                answer_match_leaf[u] = answer_leaf_at_node[u];
            }
        }

        vector<int> to_answer;
        if (begin == 0) {
            to_answer.resize(N);
            iota(to_answer.begin(), to_answer.end(), 0);
        } else {
            vector<int> saved_answer_leaf = answer_leaf_at_node;
            answer_leaf_at_node = answer_match_leaf;
            make_mapping(cur, cur_leaf_at_node, to_answer);
            answer_leaf_at_node = move(saved_answer_leaf);
        }

        for (auto [low, u] : marked) {
            int at = begin + low;
            if (at < static_cast<int>(short_leaves.size()) &&
                to_answer[u] != -1) {
                int parent = cur[u].parent;
                singleton_leaf_orbit[short_leaves[at]] =
                    orbit[to_answer[parent]];
            }
        }
    }

    map<int, vector<int>> free_slots;
    for (int u = 0; u < N; ++u) {
        if (answer_nodes[u].singleton) {
            free_slots[orbit[answer_nodes[u].parent]].push_back(u);
        }
    }

    for (int leaf : short_leaves) {
        vector<int>& slots =
            free_slots[singleton_leaf_orbit[leaf]];
        int u = slots.back();
        slots.pop_back();
        answer_leaf_at_node[u] = leaf;
    }

    for (int u = 0; u < N; ++u) {
        answer_parent[u] = answer_nodes[u].parent;
    }

    vector<int> label(N, -1);
    for (int u = 0; u < N; ++u) {
        if (answer_leaf_at_node[u] != -1) {
            label[u] = answer_leaf_at_node[u];
        }
    }
    label[root] = N - 1;

    int next_label = M;
    for (int u = 0; u < N; ++u) {
        if (label[u] == -1) label[u] = next_label++;
    }

    vector<int> result(N - 1);
    for (int u = 0; u < N; ++u) {
        if (u != root) {
            result[label[u]] = label[answer_parent[u]];
        }
    }
    return result;
}