可以做到 期望 (O(nk\log(nk))),低于 (O(n^2k))。
更准确地,设 (N=\sum_{i=1}^k n_i),复杂度为:
$$ \boxed{ O\!\left(\sum_{i=1}^k n_i\log n_i+N\log(k+1)\right) } $$
空间复杂度 (O(N))。“期望复杂度”只来自随机 Treap;答案不是随机模拟,而是通过分段线性函数计算。
核心是:把每棵树压成一个分段线性函数,用带仿射懒标记的 Treap 合并,最后用堆扫描所有断点。
1. 先把多棵树的问题拆开
题目每一步可以任选一个球移动,任意球到达 (1) 就结束。
先考虑单棵树,引入一个参数 (R):到达根可以得到奖励 (R),每移动一步花费 (1),并且允许随时放弃。
定义
$$ F_u(R)=\text{从 }u\text{ 出发,可以得到的最大期望净收益}. $$
于是
$$ F_1(R)=R, $$
$$ F_u(R)= \max\left( 0,\frac{\sum_{v\sim u}F_v(R)}{\deg(u)}-1 \right). $$
这个函数是连续、凸、分段线性的,斜率在 ([0,1]) 内。
对于第 (i) 棵树的起点 (s_i),记
$$ p_i(R)=F'_{s_i}(R). $$
那么原题答案为
$$ \boxed{ \operatorname{Ans} = \int_0^\infty \prod_{i=1}^k \bigl(1-p_i(R)\bigr)\,dR } $$
因此,只要求出每棵树起点收益函数的斜率变化,就可以计算答案。
为什么这个积分是正确的?
设当前各个球的位置为 (u_1,\ldots,u_k),定义
$$ A(u_1,\ldots,u_k) = \int_0^\infty \prod_i\bigl(1-F'_{u_i}(R)\bigr)\,dR. $$
令 (\gamma_u) 为 (F_u) 开始变为正数的阈值。在阈值处有
$$ \frac1{\deg(u)}\sum_{v\sim u}F_v(\gamma_u)=1. $$
假设选择位于 (u) 的球移动,并记
$$ Q(R)=\prod_{\text{其他球 }j}\bigl(1-F'_{u_j}(R)\bigr). $$
由上面的单树递推式可得:
$$ \begin{aligned} A-\mathbb E[A_{\text{移动后}}] &= \int_0^{\gamma_u} \left(\frac1{\deg(u)}\sum_{v\sim u}F'_v(R)\right)Q(R)\,dR\\ &\le \frac1{\deg(u)}\sum_{v\sim u}F_v(\gamma_u)\\ &=1. \end{aligned} $$
如果选择当前 (\gamma_u) 最小的球,那么在 (0<R<\gamma_u) 内,其他球的收益函数斜率都为 (0),所以 (Q(R)=1),等号成立。
也就是说,任何操作使 (A) 的期望下降量都不超过 (1),而总存在一个操作使它恰好下降 (1)。结束状态的 (A=0),故最优期望步数就是上述积分。
下面重点解决如何快速求出这些斜率。
2. 单棵树的分段线性 DP
以 (1) 为根。
对非根节点 (u),假设其父亲的收益值被固定为 (x),定义:
$$ f_u(x)=u\text{ 的收益值},\qquad g_u(x)=x-f_u(x). $$
设 (u) 自己的收益值为 (t=f_u(x)),孩子集合为 (\operatorname{son}(u)),并记
$$ S(t)=\sum_{v\in\operatorname{son}(u)}g_v(t),\qquad d=\deg(u). $$
当 (t>0) 时,由 Bellman 方程:
$$ dt=x+\sum_v\bigl(t-g_v(t)\bigr)-d. $$
因为非根节点的孩子数为 (d-1),整理得到
$$ \boxed{x=t+S(t)+d}. $$
同时,
$$ \boxed{g_u(x)=S(t)+d}. $$
所以,只要先把孩子的 (g_v) 逐点相加,再对整条曲线执行变换:
$$ \boxed{ (t,y)\longmapsto(t+y+d,\ y+d) } $$
就得到了 (g_u) 在 (x\ge d) 上的部分。
而在 (0\le x\le d) 上:
$$ f_u(x)=0,\qquad g_u(x)=x. $$
因此再补一个断点 ((d,d)) 即可。
每个非根节点只新增一个断点,所以一棵树总共最多有 (n-1) 个断点。
同时维护起点的收益函数
还需要跟踪:当父亲收益为 (x) 时,起点 (s) 的收益是多少。
对于不包含 (s) 的子树,将这个辅助函数定义为 (0)。合并孩子时,辅助函数也逐点相加;由于只有一个孩子可能包含 (s),不会重复计算。
在节点 (u=s) 处,节点自身收益就是 (t),所以将辅助函数设为 (t)。
实现时不需要保存这个辅助函数的值,只需要保存其导数 (p)。
设当前曲线斜率为 (m=S'(t))。执行上述坐标变换后:
$$ \boxed{ m\leftarrow \frac{m}{1+m},\qquad p\leftarrow \frac{p}{1+m} } $$
在 (u=s) 处,变换前先令 (p=1)。
根节点 (1) 的收益已经固定为 (R),因此根只合并孩子,不执行上述变换。最后得到的 (p),就是需要的 (F'_s(R))。
3. 用 Treap 消掉平方复杂度
Treap 按断点横坐标维护,每个断点存:
$$ (x,y,m,p), $$
其中 (m,p) 都表示断点右侧的斜率。
主要操作只有两种。
整条曲线变换。 对整个 Treap 打懒标记:
$$ x\leftarrow x+y+d,\quad y\leftarrow y+d,\quad m\leftarrow\frac{m}{1+m},\quad p\leftarrow\frac{p}{1+m}. $$
这个变换保持横坐标的顺序,不需要枚举断点。
两条曲线逐点相加。 使用 Treap 的整树合并:选择优先级较大的根,按它的横坐标分裂另一棵树,再递归合并左右两侧。
关键是:如果其中一侧已经没有断点,那么它在当前区间上就是一个线性函数,可以一次懒标记加到另一棵 Treap 上,而不是逐点处理。
这种整树合并沿用 Treap union 的分裂递归结构。大小为 (a\le b) 的两棵树,期望合并复杂度为
$$ O\!\left(a\log\left(1+\frac ba\right)\right). $$
Treap 的这一合并界可参考 CMU 的复杂度说明;这里每个递归节点只额外维护常数个曲线参数。(CMU School of Computer Science)
所有孩子合并的代价累计为 (O(n\log n)),而不是逐点插入带来的额外一层对数。
最后,每棵树的断点已经有序,用一个大小为 (k) 的小根堆归并。两个相邻断点之间所有 (p_i) 都不变,因此直接累加矩形面积即可,耗时 (O(N\log(k+1)))。
4. C++17 代码
#include <bits/stdc++.h>
using namespace std;
using Real = long double;
// 一段函数:y(x) = m*x + b;起点收益函数的导数为 p。
struct Line {
Real m = 0, b = 0, p = 0;
};
// x_new = a*x + b*y + c
// y_new = d*x + e*y + f
// m_new = (d + e*m) / (a + b*m)
// p_new = (g + h*m + i*p) / (a + b*m)
struct Tag {
Real a = 1, b = 0, c = 0;
Real d = 0, e = 1, f = 0;
Real g = 0, h = 0, i = 1;
};
struct Node {
int l = 0, r = 0;
uint64_t priority = 0;
Real x = 0, y = 0, m = 0, p = 0;
Tag tag;
bool dirty = false;
};
struct Event {
Real x, p;
};
class CurveTreap {
vector<Node> tr;
mt19937_64 &rng;
static Tag compose(const Tag &u, const Tag &v) {
// 返回复合变换 u(v(.))
Tag w;
w.a = u.a * v.a + u.b * v.d;
w.b = u.a * v.b + u.b * v.e;
w.c = u.a * v.c + u.b * v.f + u.c;
w.d = u.d * v.a + u.e * v.d;
w.e = u.d * v.b + u.e * v.e;
w.f = u.d * v.c + u.e * v.f + u.f;
w.g = u.g * v.a + u.h * v.d + u.i * v.g;
w.h = u.g * v.b + u.h * v.e + u.i * v.h;
w.i = u.i * v.i;
return w;
}
void apply(int u, const Tag &t) {
if (!u) return;
Node &v = tr[u];
Real x = v.x, y = v.y;
Real m = v.m, p = v.p;
Real den = t.a + t.b * m;
v.x = t.a * x + t.b * y + t.c;
v.y = t.d * x + t.e * y + t.f;
v.m = (t.d + t.e * m) / den;
v.p = (t.g + t.h * m + t.i * p) / den;
v.tag = v.dirty ? compose(t, v.tag) : t;
v.dirty = true;
}
void push(int u) {
if (!u || !tr[u].dirty) return;
Tag t = tr[u].tag;
apply(tr[u].l, t);
apply(tr[u].r, t);
tr[u].tag = Tag{};
tr[u].dirty = false;
}
Line rightLine(int u) const {
const Node &v = tr[u];
return {v.m, v.y - v.m * v.x, v.p};
}
// 整棵 Treap:y += s.m*x + s.b,p += s.p。
void addLine(int u, const Line &s) {
if (!u || (s.m == 0 && s.b == 0 && s.p == 0)) {
return;
}
Node &v = tr[u];
v.y += s.m * v.x + s.b;
v.m += s.m;
v.p += s.p;
Tag &t = v.tag;
t.d += s.m * t.a;
t.e += s.m * t.b;
t.f += s.m * t.c + s.b;
t.g += s.p * t.a;
t.h += s.p * t.b;
v.dirty = true;
}
// 按横坐标 x 分裂,同时返回 x 右侧的线性段。
// 恰好位于 x 的断点会与另一棵 Treap 的根合并。
void split(int u, Real x, int &l, int &r, Line &mid) {
if (!u) {
l = r = 0;
return;
}
push(u);
if (tr[u].x == x) {
l = tr[u].l;
r = tr[u].r;
mid = rightLine(u);
} else if (tr[u].x < x) {
mid = rightLine(u);
int a, b;
split(tr[u].r, x, a, b, mid);
tr[u].r = a;
l = u;
r = b;
} else {
int a, b;
split(tr[u].l, x, a, b, mid);
tr[u].l = b;
l = a;
r = u;
}
}
// 要求 a 中所有横坐标都小于 b 中的横坐标。
int join(int a, int b) {
if (!a || !b) return a ? a : b;
if (tr[a].priority > tr[b].priority) {
push(a);
tr[a].r = join(tr[a].r, b);
return a;
}
push(b);
tr[b].l = join(a, tr[b].l);
return b;
}
void collect(int u, vector<Event> &out) {
if (!u) return;
push(u);
collect(tr[u].l, out);
out.push_back({tr[u].x, tr[u].p});
collect(tr[u].r, out);
}
public:
CurveTreap(int n, mt19937_64 &generator) : rng(generator) {
tr.reserve(n + 1);
tr.emplace_back();
}
// 两个分段线性函数逐点相加。
// la、lb 分别是各自第一个断点之前的线性段。
int meld(int a, int b, Line la, Line lb) {
if (!a) {
addLine(b, la);
return b;
}
if (!b) {
addLine(a, lb);
return a;
}
if (tr[a].priority < tr[b].priority) {
swap(a, b);
swap(la, lb);
}
push(a);
// 必须在修改 a 之前保存原来的右侧线性段。
Line ar = rightLine(a);
Line br = lb;
int bl, rr;
split(b, tr[a].x, bl, rr, br);
tr[a].l = meld(tr[a].l, bl, la, lb);
tr[a].r = meld(tr[a].r, rr, ar, br);
tr[a].y += br.m * tr[a].x + br.b;
tr[a].m += br.m;
tr[a].p += br.p;
return a;
}
// 将孩子函数之和 S(t) 转成当前顶点的函数:
// X = t + S(t) + degree,Y = S(t) + degree。
int extend(int root, int degree, int children, bool isStart) {
if (root) {
Node &v = tr[root];
Tag &t = v.tag;
// 当前节点是起点:辅助收益函数设为 t,导数为 1。
if (isStart) {
v.p = 1;
t.g = t.a;
t.h = t.b;
t.i = 0;
}
Real den = 1 + v.m;
v.x += v.y + degree;
v.y += degree;
v.m /= den;
v.p /= den;
t.a += t.d;
t.b += t.e;
t.c += t.f + degree;
t.f += degree;
v.dirty = true;
}
// 补上新断点 (degree, degree)。
// 此前 S(t) 的初始斜率等于孩子数。
Node v;
v.priority = rng();
v.x = v.y = degree;
v.m = Real(children) / (children + 1);
v.p = isStart ? Real(1) / (children + 1) : 0;
tr.push_back(v);
int u = int(tr.size()) - 1;
return join(u, root);
}
vector<Event> events(int root) {
vector<Event> out;
collect(root, out);
// 数学上横坐标和 p 都不下降;
// 修正浮点舍入造成的微小偏差。
Real lastX = 0, lastP = 0;
for (auto &e : out) {
e.x = max(e.x, lastX);
e.p = max(lastP, min(Real(1), max(Real(0), e.p)));
lastX = e.x;
lastP = e.p;
}
// 奖励充分大时不会放弃,最终斜率必为 1。
if (!out.empty()) out.back().p = 1;
return out;
}
};
struct HeapItem {
Real x;
int tree, pos;
bool operator>(const HeapItem &o) const {
return x > o.x;
}
};
int main() {
ios::sync_with_stdio(false);
cin.tie(nullptr);
int k;
if (!(cin >> k)) return 0;
mt19937_64 rng(
chrono::steady_clock::now().time_since_epoch().count()
);
vector<vector<Event>> all(k);
for (int i = 0; i < k; ++i) {
int n, s;
cin >> n >> s;
vector<vector<int>> adj(n + 1);
for (int j = 1; j < n; ++j) {
int u, v;
cin >> u >> v;
adj[u].push_back(v);
adj[v].push_back(u);
}
CurveTreap curves(n, rng);
auto dfs = [&](auto &&self, int u, int fa) -> int {
int root = 0;
int children = 0;
for (int v : adj[u]) {
if (v == fa) continue;
int child = self(self, v, u);
root = curves.meld(
root, child,
{Real(children), 0, 0},
{1, 0, 0}
);
++children;
}
// 根的收益就是外部参数 R,只合并,不执行变换。
if (u == 1) return root;
return curves.extend(
root, int(adj[u].size()), children, u == s
);
};
all[i] = curves.events(dfs(dfs, 1, 0));
}
priority_queue<
HeapItem,
vector<HeapItem>,
greater<HeapItem>
> heap;
for (int i = 0; i < k; ++i) {
heap.push({all[i][0].x, i, 0});
}
vector<Real> q(k, 1);
Real product = 1;
Real previous = 0;
Real answer = 0;
while (!heap.empty()) {
auto [x, i, pos] = heap.top();
heap.pop();
answer += (x - previous) * product;
previous = x;
Real nextQ = 1 - all[i][pos].p;
// 有一个因子变成 0,之后的积分全部为 0。
if (nextQ <= 0) break;
product = product / q[i] * nextQ;
q[i] = nextQ;
if (product == 0) break;
if (pos + 1 < int(all[i].size())) {
heap.push({all[i][pos + 1].x, i, pos + 1});
}
}
cout << fixed << setprecision(15) << answer << '\n';
return 0;
}
本地校验:题面样例输出 4.666666666666667;与小规模完整状态空间的最优策略计算做了 200 组随机对拍,与朴素分段线性函数实现做了 1,200 组随机对拍,结果一致;另外检查了链、星形树、二叉树等结构。